From 111508738309dd8e117d37238ac0a41bfa88e147 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Thu, 12 Mar 2026 23:59:18 +0100 Subject: [PATCH 01/17] feat: implement Mem0 memory backend adapter (#206) Add concrete MemoryBackend implementation using Mem0 (embedded Qdrant + SQLite) as the storage layer, unblocking all downstream memory features. - Mem0MemoryBackend implements MemoryBackend, MemoryCapabilities, and SharedKnowledgeStore protocols with asyncio.to_thread for all sync Mem0 calls - Mapping layer (mappers.py) for bidirectional domain model <-> Mem0 dict conversion with _synthorg_ metadata prefix - Mem0BackendConfig + config builder deriving from CompanyMemoryConfig - Factory wired up with deferred import (no longer raises MemoryConfigError for mem0 backend) - Shared knowledge via reserved __synthorg_shared__ namespace with publisher ownership tracking - 95 unit tests (adapter, mappers, config) + 6 integration tests (retrieval pipeline, shared knowledge flow) Closes #206 --- pyproject.toml | 5 + src/ai_company/memory/__init__.py | 9 +- src/ai_company/memory/backends/__init__.py | 5 + .../memory/backends/mem0/__init__.py | 6 + .../memory/backends/mem0/adapter.py | 700 ++++++++++++++++ src/ai_company/memory/backends/mem0/config.py | 109 +++ .../memory/backends/mem0/mappers.py | 261 ++++++ src/ai_company/memory/factory.py | 37 +- tests/integration/memory/test_mem0_backend.py | 271 +++++++ tests/unit/memory/backends/__init__.py | 0 tests/unit/memory/backends/mem0/__init__.py | 0 .../unit/memory/backends/mem0/test_adapter.py | 762 ++++++++++++++++++ .../unit/memory/backends/mem0/test_config.py | 125 +++ .../unit/memory/backends/mem0/test_mappers.py | 368 +++++++++ tests/unit/memory/test_factory.py | 19 +- tests/unit/memory/test_init.py | 1 + uv.lock | 236 ++++++ 17 files changed, 2891 insertions(+), 23 deletions(-) create mode 100644 src/ai_company/memory/backends/__init__.py create mode 100644 src/ai_company/memory/backends/mem0/__init__.py create mode 100644 src/ai_company/memory/backends/mem0/adapter.py create mode 100644 src/ai_company/memory/backends/mem0/config.py create mode 100644 src/ai_company/memory/backends/mem0/mappers.py create mode 100644 tests/integration/memory/test_mem0_backend.py create mode 100644 tests/unit/memory/backends/__init__.py create mode 100644 tests/unit/memory/backends/mem0/__init__.py create mode 100644 tests/unit/memory/backends/mem0/test_adapter.py create mode 100644 tests/unit/memory/backends/mem0/test_config.py create mode 100644 tests/unit/memory/backends/mem0/test_mappers.py diff --git a/pyproject.toml b/pyproject.toml index 8d0a051510..a54f4cea27 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,6 +21,7 @@ dependencies = [ "litellm==1.82.1", "litestar[standard,structlog,pydantic,brotli,prometheus]==2.21.1", "mcp==1.26.0", + "mem0ai==1.0.5", "pydantic==2.12.5", "pyjwt[crypto]==2.11.0", "pyyaml==6.0.3", @@ -184,6 +185,10 @@ ignore_missing_imports = true module = "mcp.*" ignore_missing_imports = true +[[tool.mypy.overrides]] +module = "mem0.*" +ignore_missing_imports = true + [[tool.mypy.overrides]] module = "litestar.*" ignore_missing_imports = true diff --git a/src/ai_company/memory/__init__.py b/src/ai_company/memory/__init__.py index 6ea6f8070a..d64975b521 100644 --- a/src/ai_company/memory/__init__.py +++ b/src/ai_company/memory/__init__.py @@ -3,11 +3,13 @@ Re-exports protocols (``MemoryBackend``, ``MemoryCapabilities``, ``SharedKnowledgeStore``, ``MemoryInjectionStrategy``, ``OrgMemoryBackend``, ``ConsolidationStrategy``, ``ArchivalStore``), -domain models, config models, factory, retrieval pipeline, -consolidation, org memory, and error hierarchy so consumers can -import from ``ai_company.memory`` directly. +concrete backends (``Mem0MemoryBackend``), domain models, config +models, factory, retrieval pipeline, consolidation, org memory, and +error hierarchy so consumers can import from ``ai_company.memory`` +directly. """ +from ai_company.memory.backends.mem0 import Mem0MemoryBackend from ai_company.memory.capabilities import MemoryCapabilities from ai_company.memory.config import ( CompanyMemoryConfig, @@ -74,6 +76,7 @@ "DefaultTokenEstimator", "InjectionPoint", "InjectionStrategy", + "Mem0MemoryBackend", "MemoryBackend", "MemoryCapabilities", "MemoryCapabilityError", diff --git a/src/ai_company/memory/backends/__init__.py b/src/ai_company/memory/backends/__init__.py new file mode 100644 index 0000000000..dd0affabc7 --- /dev/null +++ b/src/ai_company/memory/backends/__init__.py @@ -0,0 +1,5 @@ +"""Concrete memory backend implementations.""" + +from ai_company.memory.backends.mem0 import Mem0MemoryBackend + +__all__ = ["Mem0MemoryBackend"] diff --git a/src/ai_company/memory/backends/mem0/__init__.py b/src/ai_company/memory/backends/mem0/__init__.py new file mode 100644 index 0000000000..3c103e22b5 --- /dev/null +++ b/src/ai_company/memory/backends/mem0/__init__.py @@ -0,0 +1,6 @@ +"""Mem0-backed agent memory — adapter, config, and mappers.""" + +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend +from ai_company.memory.backends.mem0.config import Mem0BackendConfig + +__all__ = ["Mem0BackendConfig", "Mem0MemoryBackend"] diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py new file mode 100644 index 0000000000..802887ac43 --- /dev/null +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -0,0 +1,700 @@ +"""Mem0 memory backend adapter. + +Implements ``MemoryBackend``, ``MemoryCapabilities``, and +``SharedKnowledgeStore`` protocols using Mem0 (embedded Qdrant + SQLite) +as the storage layer. + +All Mem0 calls are synchronous — they run in ``asyncio.to_thread()`` +to avoid blocking the event loop. +""" + +import asyncio +from typing import TYPE_CHECKING, Any + +from ai_company.core.enums import MemoryCategory +from ai_company.core.types import NotBlankStr +from ai_company.memory.backends.mem0.config import ( + Mem0BackendConfig, + build_mem0_config_dict, +) +from ai_company.memory.backends.mem0.mappers import ( + apply_post_filters, + build_mem0_metadata, + mem0_result_to_entry, + query_to_mem0_getall_args, + query_to_mem0_search_args, +) +from ai_company.memory.errors import ( + MemoryConnectionError, + MemoryRetrievalError, + MemoryStoreError, +) +from ai_company.observability import get_logger + +if TYPE_CHECKING: + from ai_company.memory.models import ( + MemoryEntry, + MemoryQuery, + MemoryStoreRequest, + ) +from ai_company.observability.events.memory import ( + MEMORY_BACKEND_CONNECTED, + MEMORY_BACKEND_CONNECTING, + MEMORY_BACKEND_CONNECTION_FAILED, + MEMORY_BACKEND_CREATED, + MEMORY_BACKEND_DISCONNECTED, + MEMORY_BACKEND_DISCONNECTING, + MEMORY_BACKEND_HEALTH_CHECK, + MEMORY_BACKEND_NOT_CONNECTED, + MEMORY_ENTRY_COUNT_FAILED, + MEMORY_ENTRY_COUNTED, + MEMORY_ENTRY_DELETE_FAILED, + MEMORY_ENTRY_DELETED, + MEMORY_ENTRY_FETCH_FAILED, + MEMORY_ENTRY_FETCHED, + MEMORY_ENTRY_RETRIEVAL_FAILED, + MEMORY_ENTRY_RETRIEVED, + MEMORY_ENTRY_STORE_FAILED, + MEMORY_ENTRY_STORED, + MEMORY_SHARED_PUBLISH_FAILED, + MEMORY_SHARED_PUBLISHED, + MEMORY_SHARED_RETRACT_FAILED, + MEMORY_SHARED_RETRACTED, + MEMORY_SHARED_SEARCH_FAILED, + MEMORY_SHARED_SEARCHED, +) + +logger = get_logger(__name__) + +# Reserved user_id for the shared knowledge namespace. +_SHARED_NAMESPACE: str = "__synthorg_shared__" + +# Metadata key to track who published a shared memory. +_PUBLISHER_KEY: str = "_synthorg_publisher" + + +class Mem0MemoryBackend: + """Mem0-backed agent memory backend. + + Implements the ``MemoryBackend``, ``MemoryCapabilities``, and + ``SharedKnowledgeStore`` protocols. + + Args: + mem0_config: Mem0-specific backend configuration. + max_memories_per_agent: Per-agent memory limit (from company config). + """ + + def __init__( + self, + *, + mem0_config: Mem0BackendConfig, + max_memories_per_agent: int = 10_000, + ) -> None: + self._mem0_config = mem0_config + self._max_memories_per_agent = max_memories_per_agent + self._client: Any = None + self._connected = False + logger.debug( + MEMORY_BACKEND_CREATED, + backend="mem0", + data_dir=mem0_config.data_dir, + collection_name=mem0_config.collection_name, + ) + + # ── Lifecycle ───────────────────────────────────────────────── + + async def connect(self) -> None: + """Establish connection to Mem0. + + Creates the Mem0 ``Memory`` client with embedded Qdrant. + + Raises: + MemoryConnectionError: If Mem0 initialization fails. + """ + logger.info(MEMORY_BACKEND_CONNECTING, backend="mem0") + try: + from mem0 import Memory # noqa: PLC0415 + + config_dict = build_mem0_config_dict(self._mem0_config) + client = await asyncio.to_thread(Memory.from_config, config_dict) + except Exception as exc: + logger.exception( + MEMORY_BACKEND_CONNECTION_FAILED, + backend="mem0", + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to connect to Mem0: {exc}" + raise MemoryConnectionError(msg) from exc + self._client = client + self._connected = True + logger.info(MEMORY_BACKEND_CONNECTED, backend="mem0") + + async def disconnect(self) -> None: + """Close the Mem0 connection. + + Safe to call even if not connected. + """ + logger.info(MEMORY_BACKEND_DISCONNECTING, backend="mem0") + self._client = None + self._connected = False + logger.info(MEMORY_BACKEND_DISCONNECTED, backend="mem0") + + async def health_check(self) -> bool: + """Check whether the Mem0 backend is healthy. + + Returns: + ``True`` if connected, ``False`` otherwise. + """ + healthy = self._connected and self._client is not None + logger.debug( + MEMORY_BACKEND_HEALTH_CHECK, + backend="mem0", + healthy=healthy, + ) + return healthy + + @property + def is_connected(self) -> bool: + """Whether the backend has an active connection.""" + return self._connected + + @property + def backend_name(self) -> NotBlankStr: + """Human-readable backend identifier.""" + return NotBlankStr("mem0") + + # ── Capabilities ────────────────────────────────────────────── + + @property + def supported_categories(self) -> frozenset[MemoryCategory]: + """All memory categories are supported.""" + return frozenset(MemoryCategory) + + @property + def supports_graph(self) -> bool: + """Graph memory is not available in embedded mode.""" + return False + + @property + def supports_temporal(self) -> bool: + """Temporal tracking is available via timestamps.""" + return True + + @property + def supports_vector_search(self) -> bool: + """Vector search is available via embedded Qdrant.""" + return True + + @property + def supports_shared_access(self) -> bool: + """Cross-agent shared memory is available.""" + return True + + @property + def max_memories_per_agent(self) -> int | None: + """Maximum memories per agent from configuration.""" + return self._max_memories_per_agent + + # ── Connection guard ────────────────────────────────────────── + + def _require_connected(self) -> None: + """Raise ``MemoryConnectionError`` if not connected.""" + if not self._connected or self._client is None: + logger.warning( + MEMORY_BACKEND_NOT_CONNECTED, + backend="mem0", + ) + msg = "Not connected — call connect() first" + raise MemoryConnectionError(msg) + + # ── CRUD Operations ─────────────────────────────────────────── + + async def store( + self, + agent_id: NotBlankStr, + request: MemoryStoreRequest, + ) -> NotBlankStr: + """Store a memory entry for an agent. + + Args: + agent_id: Owning agent identifier. + request: Memory content and metadata. + + Returns: + The backend-assigned memory ID. + + Raises: + MemoryConnectionError: If the backend is not connected. + MemoryStoreError: If the store operation fails. + """ + self._require_connected() + try: + kwargs = { + "messages": [ + {"role": "user", "content": request.content}, + ], + "user_id": str(agent_id), + "metadata": build_mem0_metadata(request), + "infer": False, + } + result = await asyncio.to_thread(self._client.add, **kwargs) + results_list = result.get("results", []) + if not results_list: + msg = "Mem0 add returned no results" + raise MemoryStoreError(msg) # noqa: TRY301 + memory_id = NotBlankStr(str(results_list[0]["id"])) + except MemoryStoreError: + logger.exception( + MEMORY_ENTRY_STORE_FAILED, + agent_id=agent_id, + ) + raise + except Exception as exc: + logger.exception( + MEMORY_ENTRY_STORE_FAILED, + agent_id=agent_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to store memory: {exc}" + raise MemoryStoreError(msg) from exc + else: + logger.info( + MEMORY_ENTRY_STORED, + agent_id=agent_id, + memory_id=memory_id, + category=request.category.value, + ) + return memory_id + + async def retrieve( + self, + agent_id: NotBlankStr, + query: MemoryQuery, + ) -> tuple[MemoryEntry, ...]: + """Retrieve memories for an agent, ordered by relevance. + + Args: + agent_id: Owning agent identifier. + query: Retrieval parameters. + + Returns: + Matching memory entries ordered by relevance. + + Raises: + MemoryConnectionError: If the backend is not connected. + MemoryRetrievalError: If the retrieval fails. + """ + self._require_connected() + try: + if query.text is not None: + kwargs = query_to_mem0_search_args(str(agent_id), query) + raw_result = await asyncio.to_thread(self._client.search, **kwargs) + else: + kwargs = query_to_mem0_getall_args(str(agent_id), query) + raw_result = await asyncio.to_thread(self._client.get_all, **kwargs) + raw_list = raw_result.get("results", []) + entries = tuple( + mem0_result_to_entry(item, str(agent_id)) for item in raw_list + ) + entries = apply_post_filters(entries, query) + except MemoryRetrievalError: + raise + except Exception as exc: + logger.exception( + MEMORY_ENTRY_RETRIEVAL_FAILED, + agent_id=agent_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to retrieve memories: {exc}" + raise MemoryRetrievalError(msg) from exc + else: + logger.info( + MEMORY_ENTRY_RETRIEVED, + agent_id=agent_id, + count=len(entries), + ) + return entries + + async def get( + self, + agent_id: NotBlankStr, + memory_id: NotBlankStr, + ) -> MemoryEntry | None: + """Get a specific memory entry by ID. + + Args: + agent_id: Owning agent identifier. + memory_id: Memory identifier. + + Returns: + The memory entry, or ``None`` if not found. + + Raises: + MemoryConnectionError: If the backend is not connected. + MemoryRetrievalError: If the backend query fails. + """ + self._require_connected() + try: + raw = await asyncio.to_thread(self._client.get, str(memory_id)) + if raw is None: + logger.debug( + MEMORY_ENTRY_FETCHED, + agent_id=agent_id, + memory_id=memory_id, + found=False, + ) + return None + entry = mem0_result_to_entry(raw, str(agent_id)) + except MemoryRetrievalError: + raise + except Exception as exc: + logger.exception( + MEMORY_ENTRY_FETCH_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to get memory {memory_id}: {exc}" + raise MemoryRetrievalError(msg) from exc + else: + logger.debug( + MEMORY_ENTRY_FETCHED, + agent_id=agent_id, + memory_id=memory_id, + found=True, + ) + return entry + + async def delete( + self, + agent_id: NotBlankStr, + memory_id: NotBlankStr, + ) -> bool: + """Delete a specific memory entry. + + Args: + agent_id: Owning agent identifier. + memory_id: Memory identifier. + + Returns: + ``True`` if the entry was deleted, ``False`` if not found. + + Raises: + MemoryConnectionError: If the backend is not connected. + MemoryStoreError: If the delete operation fails. + """ + self._require_connected() + try: + # Check existence first — Mem0 delete doesn't indicate + # whether the entry existed. + existing = await asyncio.to_thread(self._client.get, str(memory_id)) + if existing is None: + logger.debug( + MEMORY_ENTRY_DELETED, + agent_id=agent_id, + memory_id=memory_id, + found=False, + ) + return False + await asyncio.to_thread(self._client.delete, str(memory_id)) + except MemoryStoreError: + raise + except Exception as exc: + logger.exception( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to delete memory {memory_id}: {exc}" + raise MemoryStoreError(msg) from exc + else: + logger.info( + MEMORY_ENTRY_DELETED, + agent_id=agent_id, + memory_id=memory_id, + found=True, + ) + return True + + async def count( + self, + agent_id: NotBlankStr, + *, + category: MemoryCategory | None = None, + ) -> int: + """Count memory entries for an agent. + + Note: This uses ``get_all()`` internally, which is O(n). + Acceptable because ``count()`` is not on the hot path. + + Args: + agent_id: Owning agent identifier. + category: Optional category filter. + + Returns: + Number of matching entries. + + Raises: + MemoryConnectionError: If the backend is not connected. + MemoryRetrievalError: If the count query fails. + """ + self._require_connected() + try: + raw_result = await asyncio.to_thread( + self._client.get_all, + user_id=str(agent_id), + limit=self._max_memories_per_agent, + ) + raw_list = raw_result.get("results", []) + if category is None: + count = len(raw_list) + else: + count = sum( + 1 for item in raw_list if _extract_category(item) == category + ) + except MemoryRetrievalError: + raise + except Exception as exc: + logger.exception( + MEMORY_ENTRY_COUNT_FAILED, + agent_id=agent_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to count memories: {exc}" + raise MemoryRetrievalError(msg) from exc + else: + logger.info( + MEMORY_ENTRY_COUNTED, + agent_id=agent_id, + count=count, + category=category.value if category else None, + ) + return count + + # ── SharedKnowledgeStore ────────────────────────────────────── + + async def publish( + self, + agent_id: NotBlankStr, + request: MemoryStoreRequest, + ) -> NotBlankStr: + """Publish a memory to the shared knowledge store. + + Uses a reserved namespace (``__synthorg_shared__``) and + records the publisher in metadata for ownership tracking. + + Args: + agent_id: Publishing agent identifier. + request: Memory content and metadata. + + Returns: + The backend-assigned shared memory ID. + + Raises: + MemoryStoreError: If the publish operation fails. + """ + self._require_connected() + try: + metadata = build_mem0_metadata(request) + metadata[_PUBLISHER_KEY] = str(agent_id) + kwargs = { + "messages": [ + {"role": "user", "content": request.content}, + ], + "user_id": _SHARED_NAMESPACE, + "metadata": metadata, + "infer": False, + } + result = await asyncio.to_thread(self._client.add, **kwargs) + results_list = result.get("results", []) + if not results_list: + msg = "Mem0 add returned no results for shared publish" + raise MemoryStoreError(msg) # noqa: TRY301 + memory_id = NotBlankStr(str(results_list[0]["id"])) + except MemoryStoreError: + logger.exception( + MEMORY_SHARED_PUBLISH_FAILED, + agent_id=agent_id, + ) + raise + except Exception as exc: + logger.exception( + MEMORY_SHARED_PUBLISH_FAILED, + agent_id=agent_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to publish shared memory: {exc}" + raise MemoryStoreError(msg) from exc + else: + logger.info( + MEMORY_SHARED_PUBLISHED, + agent_id=agent_id, + memory_id=memory_id, + ) + return memory_id + + async def search_shared( + self, + query: MemoryQuery, + *, + exclude_agent: NotBlankStr | None = None, + ) -> tuple[MemoryEntry, ...]: + """Search the shared knowledge store across agents. + + Args: + query: Search parameters. + exclude_agent: Optional agent ID to exclude from results. + + Returns: + Matching shared memory entries ordered by relevance. + + Raises: + MemoryRetrievalError: If the search fails. + """ + self._require_connected() + try: + if query.text is not None: + raw_result = await asyncio.to_thread( + self._client.search, + query=str(query.text), + user_id=_SHARED_NAMESPACE, + limit=query.limit, + ) + else: + raw_result = await asyncio.to_thread( + self._client.get_all, + user_id=_SHARED_NAMESPACE, + limit=query.limit, + ) + raw_list = raw_result.get("results", []) + + entries: list[MemoryEntry] = [] + for item in raw_list: + publisher = _extract_publisher(item) + entry = mem0_result_to_entry( + item, + publisher or _SHARED_NAMESPACE, + ) + entries.append(entry) + + result = tuple(entries) + result = apply_post_filters(result, query) + + if exclude_agent is not None: + result = tuple(e for e in result if e.agent_id != exclude_agent) + except MemoryRetrievalError: + raise + except Exception as exc: + logger.exception( + MEMORY_SHARED_SEARCH_FAILED, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to search shared knowledge: {exc}" + raise MemoryRetrievalError(msg) from exc + else: + logger.info( + MEMORY_SHARED_SEARCHED, + count=len(result), + exclude_agent=exclude_agent, + ) + return result + + async def retract( + self, + agent_id: NotBlankStr, + memory_id: NotBlankStr, + ) -> bool: + """Remove a memory from the shared knowledge store. + + Verifies publisher ownership before deletion. + + Args: + agent_id: Retracting agent identifier. + memory_id: Shared memory identifier. + + Returns: + ``True`` if retracted, ``False`` if not found. + + Raises: + MemoryStoreError: If the retraction operation fails. + """ + self._require_connected() + try: + raw = await asyncio.to_thread(self._client.get, str(memory_id)) + if raw is None: + logger.debug( + MEMORY_SHARED_RETRACTED, + agent_id=agent_id, + memory_id=memory_id, + found=False, + ) + return False + + publisher = _extract_publisher(raw) + if publisher != str(agent_id): + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="ownership mismatch", + publisher=publisher, + ) + msg = ( + f"Agent {agent_id} cannot retract memory " + f"{memory_id} published by {publisher}" + ) + raise MemoryStoreError(msg) # noqa: TRY301 + + await asyncio.to_thread(self._client.delete, str(memory_id)) + except MemoryStoreError: + raise + except Exception as exc: + logger.exception( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to retract shared memory {memory_id}: {exc}" + raise MemoryStoreError(msg) from exc + else: + logger.info( + MEMORY_SHARED_RETRACTED, + agent_id=agent_id, + memory_id=memory_id, + found=True, + ) + return True + + +# ── Module-level helpers ────────────────────────────────────────── + + +def _extract_category(raw: dict[str, Any]) -> MemoryCategory: + """Extract the memory category from a Mem0 result dict.""" + metadata = raw.get("metadata", {}) + if not metadata: + return MemoryCategory.WORKING + cat_str = metadata.get("_synthorg_category") + if cat_str: + return MemoryCategory(cat_str) + return MemoryCategory.WORKING + + +def _extract_publisher(raw: dict[str, Any]) -> str | None: + """Extract the publisher agent ID from a shared memory dict.""" + metadata = raw.get("metadata", {}) + if not metadata: + return None + value: str | None = metadata.get(_PUBLISHER_KEY) + return value diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py new file mode 100644 index 0000000000..43ad2b671a --- /dev/null +++ b/src/ai_company/memory/backends/mem0/config.py @@ -0,0 +1,109 @@ +"""Mem0 backend configuration and config builder. + +Isolates Mem0-specific settings from the core ``CompanyMemoryConfig``. +The ``build_mem0_config_dict`` function produces the dict that Mem0's +``Memory.from_config()`` expects. +""" + +from typing import Any + +from pydantic import BaseModel, ConfigDict, Field + +from ai_company.core.types import NotBlankStr # noqa: TC001 +from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 + + +class Mem0EmbedderConfig(BaseModel): + """Embedder settings for Mem0. + + Attributes: + provider: Embedding provider name. + model: Embedding model identifier. + dims: Embedding vector dimensions. + """ + + model_config = ConfigDict(frozen=True, allow_inf_nan=False) + + provider: NotBlankStr = Field( + default="openai", + description="Embedding provider name", + ) + model: NotBlankStr = Field( + default="text-embedding-3-small", + description="Embedding model identifier", + ) + dims: int = Field( + default=1536, + ge=1, + description="Embedding vector dimensions", + ) + + +class Mem0BackendConfig(BaseModel): + """Mem0-specific backend configuration. + + Attributes: + data_dir: Directory for Mem0 data persistence. + collection_name: Qdrant collection name. + embedder: Embedder settings. + """ + + model_config = ConfigDict(frozen=True, allow_inf_nan=False) + + data_dir: NotBlankStr = Field( + default="/data/memory", + description="Directory for Mem0 data persistence", + ) + collection_name: NotBlankStr = Field( + default="synthorg_memories", + description="Qdrant collection name", + ) + embedder: Mem0EmbedderConfig = Field( + default_factory=Mem0EmbedderConfig, + description="Embedder settings", + ) + + +def build_mem0_config_dict(config: Mem0BackendConfig) -> dict[str, Any]: + """Build the dict that ``Memory.from_config()`` expects. + + Args: + config: Mem0 backend configuration. + + Returns: + Configuration dict suitable for ``Memory.from_config()``. + """ + return { + "vector_store": { + "provider": "qdrant", + "config": { + "collection_name": config.collection_name, + "embedding_model_dims": config.embedder.dims, + "path": f"{config.data_dir}/qdrant", + }, + }, + "embedder": { + "provider": config.embedder.provider, + "config": { + "model": config.embedder.model, + }, + }, + "history_db_path": f"{config.data_dir}/history.db", + "version": "v1.1", + } + + +def build_config_from_company_config( + config: CompanyMemoryConfig, +) -> Mem0BackendConfig: + """Derive a ``Mem0BackendConfig`` from the top-level memory config. + + Args: + config: Company-wide memory configuration. + + Returns: + Mem0-specific backend configuration. + """ + return Mem0BackendConfig( + data_dir=config.storage.data_dir, + ) diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py new file mode 100644 index 0000000000..f8d6c52671 --- /dev/null +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -0,0 +1,261 @@ +"""Bidirectional mapping between SynthOrg domain models and Mem0 dicts. + +Pure functions — no I/O, no side effects. Each mapper handles one +direction of the conversion so the adapter stays thin. +""" + +from datetime import UTC, datetime +from typing import TYPE_CHECKING, Any + +from ai_company.core.enums import MemoryCategory +from ai_company.core.types import NotBlankStr +from ai_company.memory.models import ( + MemoryEntry, + MemoryMetadata, + MemoryQuery, + MemoryStoreRequest, +) + +if TYPE_CHECKING: + from pydantic import AwareDatetime + +# Metadata prefix avoids collisions with Mem0's own keys. +_PREFIX = "_synthorg_" + + +def build_mem0_metadata(request: MemoryStoreRequest) -> dict[str, Any]: + """Serialize a store request's metadata into Mem0-compatible dict. + + Args: + request: Memory store request with category and metadata. + + Returns: + Dict of prefixed metadata fields for Mem0. + """ + meta: dict[str, Any] = { + f"{_PREFIX}category": request.category.value, + f"{_PREFIX}confidence": request.metadata.confidence, + } + if request.metadata.source is not None: + meta[f"{_PREFIX}source"] = request.metadata.source + if request.metadata.tags: + meta[f"{_PREFIX}tags"] = list(request.metadata.tags) + if request.expires_at is not None: + meta[f"{_PREFIX}expires_at"] = request.expires_at.isoformat() + return meta + + +def store_request_to_mem0_args( + agent_id: str, + request: MemoryStoreRequest, +) -> dict[str, Any]: + """Convert a store request to ``Memory.add()`` keyword arguments. + + Args: + agent_id: Owning agent identifier. + request: Memory store request. + + Returns: + Dict of kwargs for ``Memory.add()``. + """ + messages = [{"role": "user", "content": request.content}] + metadata = build_mem0_metadata(request) + return { + "messages": messages, + "user_id": agent_id, + "metadata": metadata, + "infer": False, + } + + +def parse_mem0_datetime(raw: str | None) -> AwareDatetime | None: + """Parse a datetime string from Mem0 into an aware datetime. + + Mem0 stores timestamps as ISO 8601 strings. Naive datetimes + are assumed UTC. + + Args: + raw: ISO 8601 datetime string, or ``None``. + + Returns: + Aware datetime or ``None`` if input is ``None`` or empty. + """ + if not raw: + return None + dt = datetime.fromisoformat(raw) + if dt.tzinfo is None: + dt = dt.replace(tzinfo=UTC) + return dt + + +def normalize_relevance_score(score: float | None) -> float | None: + """Clamp a relevance score to [0.0, 1.0]. + + Args: + score: Raw score from Mem0 (may exceed bounds). + + Returns: + Clamped score, or ``None`` if input is ``None``. + """ + if score is None: + return None + return max(0.0, min(1.0, score)) + + +def parse_mem0_metadata( + raw_metadata: dict[str, Any] | None, +) -> tuple[MemoryCategory, MemoryMetadata, AwareDatetime | None]: + """Deserialize Mem0 metadata dict into domain objects. + + Args: + raw_metadata: Metadata dict from Mem0 result (may be ``None``). + + Returns: + Tuple of (category, metadata, expires_at). + """ + if not raw_metadata: + return ( + MemoryCategory.WORKING, + MemoryMetadata(), + None, + ) + + category_str = raw_metadata.get(f"{_PREFIX}category") + category = MemoryCategory(category_str) if category_str else MemoryCategory.WORKING + + confidence = raw_metadata.get(f"{_PREFIX}confidence", 1.0) + source = raw_metadata.get(f"{_PREFIX}source") + raw_tags = raw_metadata.get(f"{_PREFIX}tags", ()) + tags = tuple(NotBlankStr(t) for t in raw_tags if t and str(t).strip()) + + expires_at = parse_mem0_datetime( + raw_metadata.get(f"{_PREFIX}expires_at"), + ) + + metadata = MemoryMetadata( + source=source, + confidence=confidence, + tags=tags, + ) + return category, metadata, expires_at + + +def mem0_result_to_entry( + raw: dict[str, Any], + agent_id: str, +) -> MemoryEntry: + """Convert a single Mem0 result dict to a ``MemoryEntry``. + + Args: + raw: Single result dict from Mem0 (``search``, ``get``, or + ``get_all``). + agent_id: Owning agent identifier. + + Returns: + Domain ``MemoryEntry``. + """ + memory_id = NotBlankStr(str(raw["id"])) + content = NotBlankStr(str(raw.get("memory", raw.get("data", "")))) + + created_at = parse_mem0_datetime(raw.get("created_at")) + if created_at is None: + created_at = datetime.now(UTC) + updated_at = parse_mem0_datetime(raw.get("updated_at")) + + raw_metadata = raw.get("metadata") + category, metadata, expires_at = parse_mem0_metadata(raw_metadata) + + raw_score = raw.get("score") + relevance_score = normalize_relevance_score(raw_score) + + return MemoryEntry( + id=memory_id, + agent_id=NotBlankStr(agent_id), + category=category, + content=content, + metadata=metadata, + created_at=created_at, + updated_at=updated_at, + expires_at=expires_at, + relevance_score=relevance_score, + ) + + +def query_to_mem0_search_args( + agent_id: str, + query: MemoryQuery, +) -> dict[str, Any]: + """Convert a ``MemoryQuery`` to ``Memory.search()`` kwargs. + + Args: + agent_id: Owning agent identifier. + query: Retrieval query. + + Returns: + Dict of kwargs for ``Memory.search()``. + + Raises: + ValueError: If ``query.text`` is ``None`` (search requires text). + """ + if query.text is None: + msg = "search requires query.text to be set" + raise ValueError(msg) + return { + "query": query.text, + "user_id": agent_id, + "limit": query.limit, + } + + +def query_to_mem0_getall_args( + agent_id: str, + query: MemoryQuery, +) -> dict[str, Any]: + """Convert a ``MemoryQuery`` to ``Memory.get_all()`` kwargs. + + Args: + agent_id: Owning agent identifier. + query: Retrieval query. + + Returns: + Dict of kwargs for ``Memory.get_all()``. + """ + return { + "user_id": agent_id, + "limit": query.limit, + } + + +def apply_post_filters( + entries: tuple[MemoryEntry, ...], + query: MemoryQuery, +) -> tuple[MemoryEntry, ...]: + """Apply post-retrieval filters that Mem0 cannot handle natively. + + Filters by category, tags, time range, and minimum relevance. + + Args: + entries: Raw entries from Mem0. + query: Original query with filter criteria. + + Returns: + Filtered entries (order preserved). + """ + result: list[MemoryEntry] = [] + for entry in entries: + if query.categories and entry.category not in query.categories: + continue + if query.tags and not all(tag in entry.metadata.tags for tag in query.tags): + continue + if query.since and entry.created_at < query.since: + continue + if query.until and entry.created_at >= query.until: + continue + if ( + query.min_relevance > 0.0 + and entry.relevance_score is not None + and entry.relevance_score < query.min_relevance + ): + continue + result.append(entry) + return tuple(result) diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index c1f8f89fb6..f6a88d99fc 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -1,7 +1,8 @@ """Factory for creating memory backends from configuration. -Each company gets its own ``MemoryBackend`` instance. Concrete -backend registration happens in issue #41 (Mem0 adapter). +Each company gets its own ``MemoryBackend`` instance. The factory +dispatches to concrete backend implementations based on +``config.backend``. """ from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 @@ -9,7 +10,7 @@ from ai_company.memory.protocol import MemoryBackend # noqa: TC001 from ai_company.observability import get_logger from ai_company.observability.events.memory import ( - MEMORY_BACKEND_NOT_IMPLEMENTED, + MEMORY_BACKEND_CREATED, MEMORY_BACKEND_UNKNOWN, ) @@ -19,30 +20,34 @@ def create_memory_backend(config: CompanyMemoryConfig) -> MemoryBackend: """Create a memory backend from configuration. - Currently a placeholder — raises ``MemoryConfigError`` for all - backends. Concrete registration happens in #41. - Args: config: Memory configuration (includes backend selection and backend-specific settings). Returns: - A new, disconnected backend instance. Currently unreachable - — the function always raises while the Mem0 adapter (#41) - is pending. + A new, disconnected backend instance. The caller must call + ``connect()`` before use. Raises: - MemoryConfigError: If the backend is not yet implemented or - not recognized. + MemoryConfigError: If the backend is not recognized. """ if config.backend == "mem0": - msg = "mem0 backend not yet implemented" - logger.warning( - MEMORY_BACKEND_NOT_IMPLEMENTED, + from ai_company.memory.backends.mem0 import Mem0MemoryBackend # noqa: PLC0415 + from ai_company.memory.backends.mem0.config import ( # noqa: PLC0415 + build_config_from_company_config, + ) + + mem0_config = build_config_from_company_config(config) + backend = Mem0MemoryBackend( + mem0_config=mem0_config, + max_memories_per_agent=config.options.max_memories_per_agent, + ) + logger.info( + MEMORY_BACKEND_CREATED, backend="mem0", - reason=msg, + data_dir=mem0_config.data_dir, ) - raise MemoryConfigError(msg) + return backend # Defensive guard: config validation rejects unknown backends, so # this branch is unreachable under normal construction. It exists # as a safety net for callers that bypass validation (e.g. via diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py new file mode 100644 index 0000000000..d07fd83dd0 --- /dev/null +++ b/tests/integration/memory/test_mem0_backend.py @@ -0,0 +1,271 @@ +"""Integration tests for Mem0 backend with retrieval pipeline. + +Tests the adapter plugged into the retrieval pipeline (ranking + +context injection) using a mocked Mem0 client — validates the full +store -> retrieve -> rank -> format flow. +""" + +from datetime import UTC, datetime, timedelta +from unittest.mock import MagicMock + +import pytest + +from ai_company.core.enums import MemoryCategory +from ai_company.memory.backends.mem0.adapter import ( + _PUBLISHER_KEY, + Mem0MemoryBackend, +) +from ai_company.memory.backends.mem0.config import Mem0BackendConfig +from ai_company.memory.models import MemoryQuery, MemoryStoreRequest +from ai_company.memory.retrieval_config import MemoryRetrievalConfig +from ai_company.memory.retriever import ContextInjectionStrategy + +pytestmark = pytest.mark.timeout(30) + + +@pytest.fixture +def mock_client() -> MagicMock: + """Mock Mem0 Memory client.""" + return MagicMock() + + +@pytest.fixture +def backend(mock_client: MagicMock) -> Mem0MemoryBackend: + """Connected Mem0 backend with mocked client.""" + config = Mem0BackendConfig(data_dir="/tmp/test-integration") # noqa: S108 + b = Mem0MemoryBackend(mem0_config=config, max_memories_per_agent=100) + b._client = mock_client + b._connected = True + return b + + +@pytest.mark.integration +class TestMem0RetrievalPipeline: + """Test adapter integrated with the retrieval pipeline.""" + + async def test_store_then_retrieve( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store a memory, then retrieve it via semantic search.""" + mock_client.add.return_value = { + "results": [ + { + "id": "mem-int-001", + "memory": "project uses Litestar", + "event": "ADD", + }, + ], + } + mock_client.search.return_value = { + "results": [ + { + "id": "mem-int-001", + "memory": "project uses Litestar", + "score": 0.92, + "created_at": datetime.now(UTC).isoformat(), + "metadata": { + "_synthorg_category": "semantic", + "_synthorg_confidence": 1.0, + }, + }, + ], + } + + # Store + memory_id = await backend.store( + "test-agent-001", + MemoryStoreRequest( + category=MemoryCategory.SEMANTIC, + content="project uses Litestar", + ), + ) + assert memory_id == "mem-int-001" + + # Retrieve + entries = await backend.retrieve( + "test-agent-001", + MemoryQuery(text="what framework", limit=5), + ) + assert len(entries) == 1 + assert entries[0].content == "project uses Litestar" + assert entries[0].relevance_score == 0.92 + + async def test_pipeline_prepare_messages( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Full pipeline: retrieve -> rank -> format via ContextInjectionStrategy.""" + now = datetime.now(UTC) + mock_client.search.return_value = { + "results": [ + { + "id": "m1", + "memory": "agent prefers concise responses", + "score": 0.88, + "created_at": (now - timedelta(hours=1)).isoformat(), + "metadata": { + "_synthorg_category": "procedural", + "_synthorg_confidence": 0.9, + }, + }, + { + "id": "m2", + "memory": "last task was code review", + "score": 0.75, + "created_at": (now - timedelta(hours=24)).isoformat(), + "metadata": { + "_synthorg_category": "episodic", + "_synthorg_confidence": 0.8, + }, + }, + ], + } + + config = MemoryRetrievalConfig( + relevance_weight=0.7, + recency_weight=0.3, + min_relevance=0.1, + max_memories=10, + ) + strategy = ContextInjectionStrategy( + backend=backend, + config=config, + ) + + messages = await strategy.prepare_messages( + agent_id="test-agent-001", + query_text="what should I remember", + token_budget=500, + ) + + # Should produce at least one message with memory context + assert len(messages) >= 1 + # Content should include both memories (they pass min_relevance) + combined = " ".join(m.content for m in messages) + assert "concise responses" in combined + + async def test_shared_knowledge_flow( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Publish -> search_shared -> retract flow.""" + # Publish + mock_client.add.return_value = { + "results": [ + {"id": "shared-001", "memory": "company policy", "event": "ADD"}, + ], + } + + shared_id = await backend.publish( + "test-agent-001", + MemoryStoreRequest( + category=MemoryCategory.SEMANTIC, + content="company policy: always test code", + ), + ) + assert shared_id == "shared-001" + + # Search shared + mock_client.search.return_value = { + "results": [ + { + "id": "shared-001", + "memory": "company policy: always test code", + "score": 0.95, + "created_at": datetime.now(UTC).isoformat(), + "metadata": { + "_synthorg_category": "semantic", + _PUBLISHER_KEY: "test-agent-001", + }, + }, + ], + } + + entries = await backend.search_shared( + MemoryQuery(text="company policy"), + ) + assert len(entries) == 1 + assert entries[0].agent_id == "test-agent-001" + + # Retract + mock_client.get.return_value = { + "id": "shared-001", + "memory": "company policy: always test code", + "created_at": datetime.now(UTC).isoformat(), + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + } + mock_client.delete.return_value = None + + retracted = await backend.retract("test-agent-001", "shared-001") + assert retracted is True + + async def test_shared_search_excludes_agent( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Search shared knowledge excluding the requesting agent.""" + mock_client.search.return_value = { + "results": [ + { + "id": "s1", + "memory": "from agent 1", + "score": 0.9, + "created_at": datetime.now(UTC).isoformat(), + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + }, + { + "id": "s2", + "memory": "from agent 2", + "score": 0.85, + "created_at": datetime.now(UTC).isoformat(), + "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + }, + ], + } + + entries = await backend.search_shared( + MemoryQuery(text="knowledge"), + exclude_agent="test-agent-001", + ) + assert len(entries) == 1 + assert entries[0].agent_id == "test-agent-002" + + async def test_count_after_store( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Count memories after storing several entries.""" + mock_client.get_all.return_value = { + "results": [ + { + "id": "m1", + "memory": "first", + "metadata": {"_synthorg_category": "episodic"}, + }, + { + "id": "m2", + "memory": "second", + "metadata": {"_synthorg_category": "semantic"}, + }, + { + "id": "m3", + "memory": "third", + "metadata": {"_synthorg_category": "episodic"}, + }, + ], + } + + total = await backend.count("test-agent-001") + assert total == 3 + + episodic_count = await backend.count( + "test-agent-001", + category=MemoryCategory.EPISODIC, + ) + assert episodic_count == 2 diff --git a/tests/unit/memory/backends/__init__.py b/tests/unit/memory/backends/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/unit/memory/backends/mem0/__init__.py b/tests/unit/memory/backends/mem0/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py new file mode 100644 index 0000000000..01c142717b --- /dev/null +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -0,0 +1,762 @@ +"""Tests for the Mem0 memory backend adapter.""" + +from unittest.mock import MagicMock, patch + +import pytest + +from ai_company.core.enums import MemoryCategory +from ai_company.memory.backends.mem0.adapter import ( + _PUBLISHER_KEY, + _SHARED_NAMESPACE, + Mem0MemoryBackend, +) +from ai_company.memory.backends.mem0.config import Mem0BackendConfig +from ai_company.memory.errors import ( + MemoryConnectionError, + MemoryRetrievalError, + MemoryStoreError, +) +from ai_company.memory.models import ( + MemoryQuery, + MemoryStoreRequest, +) + +pytestmark = pytest.mark.timeout(30) + + +@pytest.fixture +def mem0_config() -> Mem0BackendConfig: + """Default Mem0 config for tests.""" + return Mem0BackendConfig(data_dir="/tmp/test-memory") # noqa: S108 + + +@pytest.fixture +def mock_client() -> MagicMock: + """Mock Mem0 Memory client.""" + return MagicMock() + + +@pytest.fixture +def backend( + mem0_config: Mem0BackendConfig, + mock_client: MagicMock, +) -> Mem0MemoryBackend: + """Connected backend with mocked Mem0 client.""" + b = Mem0MemoryBackend(mem0_config=mem0_config, max_memories_per_agent=100) + b._client = mock_client + b._connected = True + return b + + +def _mem0_add_result(memory_id: str = "mem-001") -> dict: + """Build a typical Mem0 add() return value.""" + return { + "results": [ + { + "id": memory_id, + "memory": "test content", + "event": "ADD", + }, + ], + } + + +def _mem0_search_result( + items: list[dict] | None = None, +) -> dict: + """Build a typical Mem0 search() return value.""" + if items is None: + items = [ + { + "id": "mem-001", + "memory": "found content", + "score": 0.85, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + "_synthorg_category": "episodic", + "_synthorg_confidence": 0.9, + }, + }, + ] + return {"results": items} + + +def _mem0_get_result(memory_id: str = "mem-001") -> dict: + """Build a typical Mem0 get() return value.""" + return { + "id": memory_id, + "memory": "stored content", + "created_at": "2026-03-12T10:00:00+00:00", + "updated_at": None, + "metadata": { + "_synthorg_category": "episodic", + "_synthorg_confidence": 1.0, + }, + } + + +def _make_store_request( + *, + category: MemoryCategory = MemoryCategory.EPISODIC, + content: str = "test content", +) -> MemoryStoreRequest: + """Helper to build a store request.""" + return MemoryStoreRequest(category=category, content=content) + + +# ── Properties ──────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestProperties: + def test_backend_name(self, backend: Mem0MemoryBackend) -> None: + assert backend.backend_name == "mem0" + + def test_is_connected_true(self, backend: Mem0MemoryBackend) -> None: + assert backend.is_connected is True + + def test_is_connected_false(self, mem0_config: Mem0BackendConfig) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + assert b.is_connected is False + + +# ── Capabilities ────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestCapabilities: + def test_supported_categories(self, backend: Mem0MemoryBackend) -> None: + assert backend.supported_categories == frozenset(MemoryCategory) + + def test_supports_graph_false(self, backend: Mem0MemoryBackend) -> None: + assert backend.supports_graph is False + + def test_supports_temporal_true(self, backend: Mem0MemoryBackend) -> None: + assert backend.supports_temporal is True + + def test_supports_vector_search_true( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert backend.supports_vector_search is True + + def test_supports_shared_access_true( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert backend.supports_shared_access is True + + def test_max_memories_per_agent( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert backend.max_memories_per_agent == 100 + + +# ── Lifecycle ───────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestLifecycle: + async def test_connect_success( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + mock_memory = MagicMock() + with patch( + "ai_company.memory.backends.mem0.adapter.asyncio.to_thread", + return_value=mock_memory, + ): + await b.connect() + + assert b.is_connected is True + + async def test_connect_failure_raises( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with ( + patch( + "ai_company.memory.backends.mem0.adapter.asyncio.to_thread", + side_effect=RuntimeError("connection failed"), + ), + pytest.raises(MemoryConnectionError, match="Failed to connect"), + ): + await b.connect() + assert b.is_connected is False + + async def test_disconnect(self, backend: Mem0MemoryBackend) -> None: + await backend.disconnect() + assert backend.is_connected is False + assert backend._client is None + + async def test_disconnect_when_not_connected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + await b.disconnect() # Should not raise + assert b.is_connected is False + + async def test_health_check_connected( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert await backend.health_check() is True + + async def test_health_check_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + assert await b.health_check() is False + + +# ── Connection guard ────────────────────────────────────────────── + + +@pytest.mark.unit +class TestConnectionGuard: + async def test_store_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.store("test-agent-001", _make_store_request()) + + async def test_retrieve_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.retrieve("test-agent-001", MemoryQuery(text="test")) + + async def test_get_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.get("test-agent-001", "mem-001") + + async def test_delete_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.delete("test-agent-001", "mem-001") + + async def test_count_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.count("test-agent-001") + + async def test_publish_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.publish("test-agent-001", _make_store_request()) + + async def test_search_shared_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.search_shared(MemoryQuery(text="test")) + + async def test_retract_raises_when_disconnected( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + b = Mem0MemoryBackend(mem0_config=mem0_config) + with pytest.raises(MemoryConnectionError, match="Not connected"): + await b.retract("test-agent-001", "mem-001") + + +# ── Store ───────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestStore: + async def test_store_success( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = _mem0_add_result("new-mem-id") + + memory_id = await backend.store( + "test-agent-001", + _make_store_request(), + ) + + assert memory_id == "new-mem-id" + mock_client.add.assert_called_once() + call_kwargs = mock_client.add.call_args[1] + assert call_kwargs["user_id"] == "test-agent-001" + assert call_kwargs["infer"] is False + + async def test_store_empty_results_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = {"results": []} + + with pytest.raises(MemoryStoreError, match="no results"): + await backend.store("test-agent-001", _make_store_request()) + + async def test_store_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.side_effect = RuntimeError("disk full") + + with pytest.raises(MemoryStoreError, match="Failed to store") as exc_info: + await backend.store("test-agent-001", _make_store_request()) + + assert exc_info.value.__cause__ is not None + + +# ── Retrieve ────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestRetrieve: + async def test_retrieve_with_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = _mem0_search_result() + + query = MemoryQuery(text="find relevant", limit=5) + entries = await backend.retrieve("test-agent-001", query) + + assert len(entries) == 1 + assert entries[0].content == "found content" + assert entries[0].relevance_score == 0.85 + mock_client.search.assert_called_once() + + async def test_retrieve_without_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = _mem0_search_result( + [ + { + "id": "mem-001", + "memory": "all content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {}, + }, + ], + ) + + query = MemoryQuery(text=None, limit=10) + entries = await backend.retrieve("test-agent-001", query) + + assert len(entries) == 1 + mock_client.get_all.assert_called_once() + + async def test_retrieve_applies_post_filters( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = _mem0_search_result( + [ + { + "id": "m1", + "memory": "episodic", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {"_synthorg_category": "episodic"}, + }, + { + "id": "m2", + "memory": "semantic", + "score": 0.8, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {"_synthorg_category": "semantic"}, + }, + ], + ) + + query = MemoryQuery( + text="test", + categories=frozenset({MemoryCategory.EPISODIC}), + ) + entries = await backend.retrieve("test-agent-001", query) + + assert len(entries) == 1 + assert entries[0].category == MemoryCategory.EPISODIC + + async def test_retrieve_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.side_effect = RuntimeError("search failed") + + with pytest.raises(MemoryRetrievalError, match="Failed to retrieve"): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + + +# ── Get ─────────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestGet: + async def test_get_existing( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = _mem0_get_result("mem-001") + + entry = await backend.get("test-agent-001", "mem-001") + + assert entry is not None + assert entry.id == "mem-001" + assert entry.agent_id == "test-agent-001" + + async def test_get_not_found( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = None + + entry = await backend.get("test-agent-001", "nonexistent") + + assert entry is None + + async def test_get_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.side_effect = RuntimeError("backend error") + + with pytest.raises(MemoryRetrievalError, match="Failed to get"): + await backend.get("test-agent-001", "mem-001") + + +# ── Delete ──────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestDelete: + async def test_delete_existing( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = _mem0_get_result("mem-001") + mock_client.delete.return_value = None + + result = await backend.delete("test-agent-001", "mem-001") + + assert result is True + mock_client.delete.assert_called_once_with("mem-001") + + async def test_delete_not_found( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = None + + result = await backend.delete("test-agent-001", "nonexistent") + + assert result is False + mock_client.delete.assert_not_called() + + async def test_delete_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.side_effect = RuntimeError("backend error") + + with pytest.raises(MemoryStoreError, match="Failed to delete"): + await backend.delete("test-agent-001", "mem-001") + + +# ── Count ───────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestCount: + async def test_count_all( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = { + "results": [ + {"id": "m1", "memory": "a", "metadata": {}}, + {"id": "m2", "memory": "b", "metadata": {}}, + ], + } + + count = await backend.count("test-agent-001") + assert count == 2 + + async def test_count_by_category( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = { + "results": [ + { + "id": "m1", + "memory": "a", + "metadata": {"_synthorg_category": "episodic"}, + }, + { + "id": "m2", + "memory": "b", + "metadata": {"_synthorg_category": "semantic"}, + }, + { + "id": "m3", + "memory": "c", + "metadata": {"_synthorg_category": "episodic"}, + }, + ], + } + + count = await backend.count( + "test-agent-001", + category=MemoryCategory.EPISODIC, + ) + assert count == 2 + + async def test_count_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.side_effect = RuntimeError("fail") + + with pytest.raises(MemoryRetrievalError, match="Failed to count"): + await backend.count("test-agent-001") + + +# ── Shared Knowledge Store ──────────────────────────────────────── + + +@pytest.mark.unit +class TestPublish: + async def test_publish_success( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = _mem0_add_result("shared-mem-001") + + memory_id = await backend.publish( + "test-agent-001", + _make_store_request(), + ) + + assert memory_id == "shared-mem-001" + call_kwargs = mock_client.add.call_args[1] + assert call_kwargs["user_id"] == _SHARED_NAMESPACE + assert _PUBLISHER_KEY in call_kwargs["metadata"] + assert call_kwargs["metadata"][_PUBLISHER_KEY] == "test-agent-001" + + async def test_publish_empty_results_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = {"results": []} + + with pytest.raises(MemoryStoreError, match="no results"): + await backend.publish("test-agent-001", _make_store_request()) + + async def test_publish_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.side_effect = RuntimeError("network error") + + with pytest.raises(MemoryStoreError, match="Failed to publish"): + await backend.publish("test-agent-001", _make_store_request()) + + +@pytest.mark.unit +class TestSearchShared: + async def test_search_shared_with_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = _mem0_search_result( + [ + { + "id": "shared-1", + "memory": "shared fact", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + "_synthorg_category": "semantic", + _PUBLISHER_KEY: "test-agent-002", + }, + }, + ], + ) + + query = MemoryQuery(text="find shared", limit=5) + entries = await backend.search_shared(query) + + assert len(entries) == 1 + assert entries[0].agent_id == "test-agent-002" + mock_client.search.assert_called_once() + call_kwargs = mock_client.search.call_args[1] + assert call_kwargs["user_id"] == _SHARED_NAMESPACE + + async def test_search_shared_without_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = _mem0_search_result( + [ + { + "id": "shared-1", + "memory": "shared fact", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + _PUBLISHER_KEY: "test-agent-002", + }, + }, + ], + ) + + query = MemoryQuery(text=None) + entries = await backend.search_shared(query) + + assert len(entries) == 1 + mock_client.get_all.assert_called_once() + + async def test_search_shared_exclude_agent( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = _mem0_search_result( + [ + { + "id": "s1", + "memory": "from agent 1", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + }, + { + "id": "s2", + "memory": "from agent 2", + "score": 0.8, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + }, + ], + ) + + query = MemoryQuery(text="test") + entries = await backend.search_shared( + query, + exclude_agent="test-agent-001", + ) + + assert len(entries) == 1 + assert entries[0].agent_id == "test-agent-002" + + async def test_search_shared_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.side_effect = RuntimeError("search error") + + with pytest.raises(MemoryRetrievalError, match="Failed to search"): + await backend.search_shared(MemoryQuery(text="test")) + + +@pytest.mark.unit +class TestRetract: + async def test_retract_success( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = { + "id": "shared-001", + "memory": "shared content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + } + mock_client.delete.return_value = None + + result = await backend.retract("test-agent-001", "shared-001") + + assert result is True + mock_client.delete.assert_called_once_with("shared-001") + + async def test_retract_not_found( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = None + + result = await backend.retract("test-agent-001", "nonexistent") + + assert result is False + + async def test_retract_ownership_mismatch( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = { + "id": "shared-001", + "memory": "content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + } + + with pytest.raises(MemoryStoreError, match="cannot retract"): + await backend.retract("test-agent-001", "shared-001") + + async def test_retract_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.side_effect = RuntimeError("backend error") + + with pytest.raises(MemoryStoreError, match="Failed to retract"): + await backend.retract("test-agent-001", "shared-001") diff --git a/tests/unit/memory/backends/mem0/test_config.py b/tests/unit/memory/backends/mem0/test_config.py new file mode 100644 index 0000000000..1025ca3630 --- /dev/null +++ b/tests/unit/memory/backends/mem0/test_config.py @@ -0,0 +1,125 @@ +"""Tests for Mem0 backend configuration.""" + +import pytest +from pydantic import ValidationError + +from ai_company.memory.backends.mem0.config import ( + Mem0BackendConfig, + Mem0EmbedderConfig, + build_config_from_company_config, + build_mem0_config_dict, +) +from ai_company.memory.config import CompanyMemoryConfig + +pytestmark = pytest.mark.timeout(30) + + +@pytest.mark.unit +class TestMem0EmbedderConfig: + def test_defaults(self) -> None: + config = Mem0EmbedderConfig() + assert config.provider == "openai" + assert config.model == "text-embedding-3-small" + assert config.dims == 1536 + + def test_custom_values(self) -> None: + config = Mem0EmbedderConfig( + provider="test-provider", + model="test-embedding-001", + dims=768, + ) + assert config.provider == "test-provider" + assert config.model == "test-embedding-001" + assert config.dims == 768 + + def test_frozen(self) -> None: + config = Mem0EmbedderConfig() + with pytest.raises(ValidationError): + config.dims = 512 # type: ignore[misc] + + def test_rejects_zero_dims(self) -> None: + with pytest.raises(ValidationError, match="dims"): + Mem0EmbedderConfig(dims=0) + + def test_rejects_blank_provider(self) -> None: + with pytest.raises(ValidationError): + Mem0EmbedderConfig(provider=" ") + + +@pytest.mark.unit +class TestMem0BackendConfig: + def test_defaults(self) -> None: + config = Mem0BackendConfig() + assert config.data_dir == "/data/memory" + assert config.collection_name == "synthorg_memories" + assert config.embedder.provider == "openai" + + def test_custom_data_dir(self) -> None: + config = Mem0BackendConfig(data_dir="/tmp/test-memory") # noqa: S108 + assert config.data_dir == "/tmp/test-memory" # noqa: S108 + + def test_custom_collection(self) -> None: + config = Mem0BackendConfig(collection_name="test-collection") + assert config.collection_name == "test-collection" + + def test_frozen(self) -> None: + config = Mem0BackendConfig() + with pytest.raises(ValidationError): + config.data_dir = "/other" # type: ignore[misc] + + +@pytest.mark.unit +class TestBuildMem0ConfigDict: + def test_default_config(self) -> None: + config = Mem0BackendConfig() + result = build_mem0_config_dict(config) + + assert result["vector_store"]["provider"] == "qdrant" + assert ( + result["vector_store"]["config"]["collection_name"] == "synthorg_memories" + ) + assert result["vector_store"]["config"]["embedding_model_dims"] == 1536 + assert result["vector_store"]["config"]["path"] == "/data/memory/qdrant" + assert result["embedder"]["provider"] == "openai" + assert result["embedder"]["config"]["model"] == "text-embedding-3-small" + assert result["history_db_path"] == "/data/memory/history.db" + assert result["version"] == "v1.1" + + def test_custom_config(self) -> None: + config = Mem0BackendConfig( + data_dir="/custom/path", + collection_name="custom-col", + embedder=Mem0EmbedderConfig( + provider="test-provider", + model="test-model", + dims=384, + ), + ) + result = build_mem0_config_dict(config) + + assert result["vector_store"]["config"]["path"] == "/custom/path/qdrant" + assert result["vector_store"]["config"]["collection_name"] == "custom-col" + assert result["vector_store"]["config"]["embedding_model_dims"] == 384 + assert result["embedder"]["provider"] == "test-provider" + assert result["embedder"]["config"]["model"] == "test-model" + assert result["history_db_path"] == "/custom/path/history.db" + + +@pytest.mark.unit +class TestBuildConfigFromCompanyConfig: + def test_derives_data_dir(self) -> None: + company_config = CompanyMemoryConfig( + backend="mem0", + ) + mem0_config = build_config_from_company_config(company_config) + + assert mem0_config.data_dir == company_config.storage.data_dir + + def test_custom_data_dir(self) -> None: + company_config = CompanyMemoryConfig( + backend="mem0", + storage={"data_dir": "/custom/data"}, + ) + mem0_config = build_config_from_company_config(company_config) + + assert mem0_config.data_dir == "/custom/data" diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py new file mode 100644 index 0000000000..307403242f --- /dev/null +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -0,0 +1,368 @@ +"""Tests for Mem0 mapping functions.""" + +from datetime import UTC, datetime, timedelta + +import pytest + +from ai_company.core.enums import MemoryCategory +from ai_company.memory.backends.mem0.mappers import ( + _PREFIX, + apply_post_filters, + build_mem0_metadata, + mem0_result_to_entry, + normalize_relevance_score, + parse_mem0_datetime, + parse_mem0_metadata, + query_to_mem0_getall_args, + query_to_mem0_search_args, + store_request_to_mem0_args, +) +from ai_company.memory.models import ( + MemoryEntry, + MemoryMetadata, + MemoryQuery, + MemoryStoreRequest, +) + +pytestmark = pytest.mark.timeout(30) + + +def _make_entry( # noqa: PLR0913 + *, + memory_id: str = "mem-1", + agent_id: str = "test-agent-001", + category: MemoryCategory = MemoryCategory.EPISODIC, + content: str = "test content", + tags: tuple[str, ...] = (), + relevance_score: float | None = None, + created_at: datetime | None = None, +) -> MemoryEntry: + """Helper to build a MemoryEntry for tests.""" + now = created_at or datetime.now(UTC) + return MemoryEntry( + id=memory_id, + agent_id=agent_id, + category=category, + content=content, + metadata=MemoryMetadata(tags=tags), + created_at=now, + relevance_score=relevance_score, + ) + + +@pytest.mark.unit +class TestBuildMem0Metadata: + def test_basic_request(self) -> None: + request = MemoryStoreRequest( + category=MemoryCategory.EPISODIC, + content="test content", + ) + meta = build_mem0_metadata(request) + + assert meta[f"{_PREFIX}category"] == "episodic" + assert meta[f"{_PREFIX}confidence"] == 1.0 + assert f"{_PREFIX}source" not in meta + assert f"{_PREFIX}tags" not in meta + assert f"{_PREFIX}expires_at" not in meta + + def test_full_metadata(self) -> None: + expires = datetime.now(UTC) + timedelta(days=7) + request = MemoryStoreRequest( + category=MemoryCategory.SEMANTIC, + content="important fact", + metadata=MemoryMetadata( + source="task-123", + confidence=0.85, + tags=("tag-a", "tag-b"), + ), + expires_at=expires, + ) + meta = build_mem0_metadata(request) + + assert meta[f"{_PREFIX}category"] == "semantic" + assert meta[f"{_PREFIX}confidence"] == 0.85 + assert meta[f"{_PREFIX}source"] == "task-123" + assert meta[f"{_PREFIX}tags"] == ["tag-a", "tag-b"] + assert meta[f"{_PREFIX}expires_at"] == expires.isoformat() + + +@pytest.mark.unit +class TestStoreRequestToMem0Args: + def test_basic_conversion(self) -> None: + request = MemoryStoreRequest( + category=MemoryCategory.WORKING, + content="remember this", + ) + args = store_request_to_mem0_args("test-agent-001", request) + + assert args["messages"] == [ + {"role": "user", "content": "remember this"}, + ] + assert args["user_id"] == "test-agent-001" + assert args["infer"] is False + assert f"{_PREFIX}category" in args["metadata"] + + +@pytest.mark.unit +class TestParseMem0Datetime: + def test_none_returns_none(self) -> None: + assert parse_mem0_datetime(None) is None + + def test_empty_string_returns_none(self) -> None: + assert parse_mem0_datetime("") is None + + def test_aware_iso_string(self) -> None: + dt = parse_mem0_datetime("2026-03-12T10:30:00+00:00") + assert dt is not None + assert dt.tzinfo is not None + assert dt.year == 2026 + + def test_naive_gets_utc(self) -> None: + dt = parse_mem0_datetime("2026-03-12T10:30:00") + assert dt is not None + assert dt.tzinfo == UTC + + def test_non_utc_timezone(self) -> None: + dt = parse_mem0_datetime("2026-03-12T10:30:00+05:30") + assert dt is not None + assert dt.utcoffset() == timedelta(hours=5, minutes=30) + + +@pytest.mark.unit +class TestNormalizeRelevanceScore: + def test_none_returns_none(self) -> None: + assert normalize_relevance_score(None) is None + + def test_in_range(self) -> None: + assert normalize_relevance_score(0.75) == 0.75 + + def test_below_zero_clamped(self) -> None: + assert normalize_relevance_score(-0.5) == 0.0 + + def test_above_one_clamped(self) -> None: + assert normalize_relevance_score(1.5) == 1.0 + + def test_boundaries(self) -> None: + assert normalize_relevance_score(0.0) == 0.0 + assert normalize_relevance_score(1.0) == 1.0 + + +@pytest.mark.unit +class TestParseMem0Metadata: + def test_none_metadata(self) -> None: + category, metadata, expires_at = parse_mem0_metadata(None) + assert category == MemoryCategory.WORKING + assert metadata.confidence == 1.0 + assert expires_at is None + + def test_empty_metadata(self) -> None: + category, metadata, _expires_at = parse_mem0_metadata({}) + assert category == MemoryCategory.WORKING + assert metadata.confidence == 1.0 + + def test_full_metadata(self) -> None: + raw = { + f"{_PREFIX}category": "semantic", + f"{_PREFIX}confidence": 0.9, + f"{_PREFIX}source": "task-456", + f"{_PREFIX}tags": ["important", "verified"], + f"{_PREFIX}expires_at": "2026-12-31T23:59:59+00:00", + } + category, metadata, expires_at = parse_mem0_metadata(raw) + + assert category == MemoryCategory.SEMANTIC + assert metadata.confidence == 0.9 + assert metadata.source == "task-456" + assert metadata.tags == ("important", "verified") + assert expires_at is not None + assert expires_at.year == 2026 + + def test_missing_category_defaults_to_working(self) -> None: + raw = {f"{_PREFIX}confidence": 0.5} + category, _metadata, _expires = parse_mem0_metadata(raw) + assert category == MemoryCategory.WORKING + + def test_empty_tags_filtered(self) -> None: + raw = {f"{_PREFIX}tags": ["valid", "", " ", "also-valid"]} + _category, metadata, _expires = parse_mem0_metadata(raw) + assert metadata.tags == ("valid", "also-valid") + + +@pytest.mark.unit +class TestMem0ResultToEntry: + def test_basic_result(self) -> None: + raw = { + "id": "abc-123", + "memory": "test content", + "created_at": "2026-03-12T10:00:00+00:00", + "updated_at": None, + "metadata": { + f"{_PREFIX}category": "episodic", + f"{_PREFIX}confidence": 0.8, + }, + } + entry = mem0_result_to_entry(raw, "test-agent-001") + + assert entry.id == "abc-123" + assert entry.agent_id == "test-agent-001" + assert entry.category == MemoryCategory.EPISODIC + assert entry.content == "test content" + assert entry.metadata.confidence == 0.8 + assert entry.relevance_score is None + + def test_with_score(self) -> None: + raw = { + "id": "def-456", + "memory": "scored content", + "created_at": "2026-03-12T10:00:00+00:00", + "score": 0.95, + "metadata": {}, + } + entry = mem0_result_to_entry(raw, "test-agent-001") + + assert entry.relevance_score == 0.95 + + def test_missing_created_at_uses_now(self) -> None: + raw = { + "id": "no-time", + "memory": "timeless", + "metadata": {}, + } + before = datetime.now(UTC) + entry = mem0_result_to_entry(raw, "test-agent-001") + after = datetime.now(UTC) + + assert before <= entry.created_at <= after + + def test_no_metadata(self) -> None: + raw = { + "id": "no-meta", + "memory": "bare content", + "created_at": "2026-03-12T10:00:00+00:00", + } + entry = mem0_result_to_entry(raw, "test-agent-001") + + assert entry.category == MemoryCategory.WORKING + assert entry.metadata.confidence == 1.0 + + +@pytest.mark.unit +class TestQueryToMem0SearchArgs: + def test_basic_search(self) -> None: + query = MemoryQuery(text="find this", limit=5) + args = query_to_mem0_search_args("test-agent-001", query) + + assert args["query"] == "find this" + assert args["user_id"] == "test-agent-001" + assert args["limit"] == 5 + + def test_raises_on_none_text(self) -> None: + query = MemoryQuery(text=None) + with pytest.raises(ValueError, match=r"search requires query\.text"): + query_to_mem0_search_args("test-agent-001", query) + + +@pytest.mark.unit +class TestQueryToMem0GetallArgs: + def test_basic_getall(self) -> None: + query = MemoryQuery(limit=20) + args = query_to_mem0_getall_args("test-agent-001", query) + + assert args["user_id"] == "test-agent-001" + assert args["limit"] == 20 + + +@pytest.mark.unit +class TestApplyPostFilters: + def test_no_filters_passes_all(self) -> None: + entries = ( + _make_entry(memory_id="m1"), + _make_entry(memory_id="m2"), + ) + query = MemoryQuery() + result = apply_post_filters(entries, query) + assert len(result) == 2 + + def test_category_filter(self) -> None: + entries = ( + _make_entry(memory_id="m1", category=MemoryCategory.EPISODIC), + _make_entry(memory_id="m2", category=MemoryCategory.SEMANTIC), + _make_entry(memory_id="m3", category=MemoryCategory.EPISODIC), + ) + query = MemoryQuery( + categories=frozenset({MemoryCategory.EPISODIC}), + ) + result = apply_post_filters(entries, query) + assert len(result) == 2 + assert all(e.category == MemoryCategory.EPISODIC for e in result) + + def test_tag_filter(self) -> None: + entries = ( + _make_entry(memory_id="m1", tags=("important",)), + _make_entry(memory_id="m2", tags=("trivial",)), + _make_entry(memory_id="m3", tags=("important", "verified")), + ) + query = MemoryQuery(tags=("important",)) + result = apply_post_filters(entries, query) + assert len(result) == 2 + + def test_time_range_filter(self) -> None: + now = datetime.now(UTC) + old = now - timedelta(hours=48) + recent = now - timedelta(hours=1) + + entries = ( + _make_entry(memory_id="m1", created_at=old), + _make_entry(memory_id="m2", created_at=recent), + ) + query = MemoryQuery(since=now - timedelta(hours=24)) + result = apply_post_filters(entries, query) + assert len(result) == 1 + assert result[0].id == "m2" + + def test_min_relevance_filter(self) -> None: + entries = ( + _make_entry(memory_id="m1", relevance_score=0.9), + _make_entry(memory_id="m2", relevance_score=0.3), + _make_entry(memory_id="m3", relevance_score=None), + ) + query = MemoryQuery(min_relevance=0.5) + result = apply_post_filters(entries, query) + # m1 passes (0.9 >= 0.5), m2 fails (0.3 < 0.5), m3 passes (None skips check) + assert len(result) == 2 + + def test_until_filter_exclusive(self) -> None: + now = datetime.now(UTC) + entries = ( + _make_entry(memory_id="m1", created_at=now - timedelta(hours=2)), + _make_entry(memory_id="m2", created_at=now), + ) + query = MemoryQuery(until=now) + result = apply_post_filters(entries, query) + assert len(result) == 1 + assert result[0].id == "m1" + + def test_combined_filters(self) -> None: + now = datetime.now(UTC) + entries = ( + _make_entry( + memory_id="m1", + category=MemoryCategory.EPISODIC, + tags=("important",), + created_at=now - timedelta(hours=1), + ), + _make_entry( + memory_id="m2", + category=MemoryCategory.SEMANTIC, + tags=("important",), + created_at=now - timedelta(hours=1), + ), + ) + query = MemoryQuery( + categories=frozenset({MemoryCategory.EPISODIC}), + tags=("important",), + since=now - timedelta(hours=2), + ) + result = apply_post_filters(entries, query) + assert len(result) == 1 + assert result[0].id == "m1" diff --git a/tests/unit/memory/test_factory.py b/tests/unit/memory/test_factory.py index 00235d7d90..26a9b39918 100644 --- a/tests/unit/memory/test_factory.py +++ b/tests/unit/memory/test_factory.py @@ -3,8 +3,8 @@ import pytest from pydantic import ValidationError +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend from ai_company.memory.config import CompanyMemoryConfig -from ai_company.memory.errors import MemoryConfigError from ai_company.memory.factory import create_memory_backend pytestmark = pytest.mark.timeout(30) @@ -12,10 +12,21 @@ @pytest.mark.unit class TestCreateMemoryBackend: - def test_mem0_raises_not_yet_implemented(self) -> None: + def test_mem0_creates_backend(self) -> None: config = CompanyMemoryConfig(backend="mem0") - with pytest.raises(MemoryConfigError, match="not yet implemented"): - create_memory_backend(config) + backend = create_memory_backend(config) + assert isinstance(backend, Mem0MemoryBackend) + assert backend.is_connected is False + assert backend.backend_name == "mem0" + + def test_mem0_passes_max_memories(self) -> None: + config = CompanyMemoryConfig( + backend="mem0", + options={"max_memories_per_agent": 500}, + ) + backend = create_memory_backend(config) + assert isinstance(backend, Mem0MemoryBackend) + assert backend.max_memories_per_agent == 500 def test_unknown_backend_rejected_by_config_validation(self) -> None: """Unknown backends are rejected by config validation.""" diff --git a/tests/unit/memory/test_init.py b/tests/unit/memory/test_init.py index 57a8ef229d..4ec452c863 100644 --- a/tests/unit/memory/test_init.py +++ b/tests/unit/memory/test_init.py @@ -16,6 +16,7 @@ def test_all_exports_importable(self) -> None: def test_all_has_expected_names(self) -> None: expected = { "ArchivalStore", + "Mem0MemoryBackend", "CompanyMemoryConfig", "ConsolidationConfig", "ConsolidationResult", diff --git a/uv.lock b/uv.lock index 7aba7cfde7..91de08fd1b 100644 --- a/uv.lock +++ b/uv.lock @@ -186,6 +186,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/3a/2a/7cc015f5b9f5db42b7d48157e23356022889fc354a2813c15934b7cb5c0e/attrs-25.4.0-py3-none-any.whl", hash = "sha256:adcf7e2a1fb3b36ac48d97835bb6d8ade15b8dcce26aba8bf1d14847b57a3373", size = 67615, upload-time = "2025-10-06T13:54:43.17Z" }, ] +[[package]] +name = "backoff" +version = "2.2.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/47/d7/5bbeb12c44d7c4f2fb5b56abce497eb5ed9f34d85701de869acedd602619/backoff-2.2.1.tar.gz", hash = "sha256:03f829f5bb1923180821643f8753b0502c3b682293992485b0eef2807afa5cba", size = 17001, upload-time = "2022-10-05T19:19:32.061Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/df/73/b6e24bd22e6720ca8ee9a85a0c4a2971af8497d8f3193fa05390cbd46e09/backoff-2.2.1-py3-none-any.whl", hash = "sha256:63579f9a0628e06278f7e47b7d7d5b6ce20dc65c5e96a6f3ca99a6adca0396e8", size = 15148, upload-time = "2022-10-05T19:19:30.546Z" }, +] + [[package]] name = "brotli" version = "1.2.0" @@ -606,6 +615,31 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f7/ec/67fbef5d497f86283db54c22eec6f6140243aae73265799baaaa19cd17fb/ghp_import-2.1.0-py3-none-any.whl", hash = "sha256:8337dd7b50877f163d4c0289bc1f1c7f127550241988d568c1db512c4324a619", size = 11034, upload-time = "2022-05-02T15:47:14.552Z" }, ] +[[package]] +name = "greenlet" +version = "3.3.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/a3/51/1664f6b78fc6ebbd98019a1fd730e83fa78f2db7058f72b1463d3612b8db/greenlet-3.3.2.tar.gz", hash = "sha256:2eaf067fc6d886931c7962e8c6bede15d2f01965560f3359b27c80bde2d151f2", size = 188267, upload-time = "2026-02-20T20:54:15.531Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/3f/ae/8bffcbd373b57a5992cd077cbe8858fff39110480a9d50697091faea6f39/greenlet-3.3.2-cp314-cp314-macosx_11_0_universal2.whl", hash = "sha256:8d1658d7291f9859beed69a776c10822a0a799bc4bfe1bd4272bb60e62507dab", size = 279650, upload-time = "2026-02-20T20:18:00.783Z" }, + { url = "https://files.pythonhosted.org/packages/d1/c0/45f93f348fa49abf32ac8439938726c480bd96b2a3c6f4d949ec0124b69f/greenlet-3.3.2-cp314-cp314-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:18cb1b7337bca281915b3c5d5ae19f4e76d35e1df80f4ad3c1a7be91fadf1082", size = 650295, upload-time = "2026-02-20T20:47:34.036Z" }, + { url = "https://files.pythonhosted.org/packages/b3/de/dd7589b3f2b8372069ab3e4763ea5329940fc7ad9dcd3e272a37516d7c9b/greenlet-3.3.2-cp314-cp314-manylinux_2_24_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:c2e47408e8ce1c6f1ceea0dffcdf6ebb85cc09e55c7af407c99f1112016e45e9", size = 662163, upload-time = "2026-02-20T20:56:01.295Z" }, + { url = "https://files.pythonhosted.org/packages/cd/ac/85804f74f1ccea31ba518dcc8ee6f14c79f73fe36fa1beba38930806df09/greenlet-3.3.2-cp314-cp314-manylinux_2_24_s390x.manylinux_2_28_s390x.whl", hash = "sha256:e3cb43ce200f59483eb82949bf1835a99cf43d7571e900d7c8d5c62cdf25d2f9", size = 675371, upload-time = "2026-02-20T21:02:49.664Z" }, + { url = "https://files.pythonhosted.org/packages/d2/d8/09bfa816572a4d83bccd6750df1926f79158b1c36c5f73786e26dbe4ee38/greenlet-3.3.2-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:63d10328839d1973e5ba35e98cccbca71b232b14051fd957b6f8b6e8e80d0506", size = 664160, upload-time = "2026-02-20T20:21:04.015Z" }, + { url = "https://files.pythonhosted.org/packages/48/cf/56832f0c8255d27f6c35d41b5ec91168d74ec721d85f01a12131eec6b93c/greenlet-3.3.2-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:8e4ab3cfb02993c8cc248ea73d7dae6cec0253e9afa311c9b37e603ca9fad2ce", size = 1619181, upload-time = "2026-02-20T20:49:36.052Z" }, + { url = "https://files.pythonhosted.org/packages/0a/23/b90b60a4aabb4cec0796e55f25ffbfb579a907c3898cd2905c8918acaa16/greenlet-3.3.2-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:94ad81f0fd3c0c0681a018a976e5c2bd2ca2d9d94895f23e7bb1af4e8af4e2d5", size = 1687713, upload-time = "2026-02-20T20:21:11.684Z" }, + { url = "https://files.pythonhosted.org/packages/f3/ca/2101ca3d9223a1dc125140dbc063644dca76df6ff356531eb27bc267b446/greenlet-3.3.2-cp314-cp314-win_amd64.whl", hash = "sha256:8c4dd0f3997cf2512f7601563cc90dfb8957c0cff1e3a1b23991d4ea1776c492", size = 232034, upload-time = "2026-02-20T20:20:08.186Z" }, + { url = "https://files.pythonhosted.org/packages/f6/4a/ecf894e962a59dea60f04877eea0fd5724618da89f1867b28ee8b91e811f/greenlet-3.3.2-cp314-cp314-win_arm64.whl", hash = "sha256:cd6f9e2bbd46321ba3bbb4c8a15794d32960e3b0ae2cc4d49a1a53d314805d71", size = 231437, upload-time = "2026-02-20T20:18:59.722Z" }, + { url = "https://files.pythonhosted.org/packages/98/6d/8f2ef704e614bcf58ed43cfb8d87afa1c285e98194ab2cfad351bf04f81e/greenlet-3.3.2-cp314-cp314t-macosx_11_0_universal2.whl", hash = "sha256:e26e72bec7ab387ac80caa7496e0f908ff954f31065b0ffc1f8ecb1338b11b54", size = 286617, upload-time = "2026-02-20T20:19:29.856Z" }, + { url = "https://files.pythonhosted.org/packages/5e/0d/93894161d307c6ea237a43988f27eba0947b360b99ac5239ad3fe09f0b47/greenlet-3.3.2-cp314-cp314t-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8b466dff7a4ffda6ca975979bab80bdadde979e29fc947ac3be4451428d8b0e4", size = 655189, upload-time = "2026-02-20T20:47:35.742Z" }, + { url = "https://files.pythonhosted.org/packages/f5/2c/d2d506ebd8abcb57386ec4f7ba20f4030cbe56eae541bc6fd6ef399c0b41/greenlet-3.3.2-cp314-cp314t-manylinux_2_24_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:b8bddc5b73c9720bea487b3bffdb1840fe4e3656fba3bd40aa1489e9f37877ff", size = 658225, upload-time = "2026-02-20T20:56:02.527Z" }, + { url = "https://files.pythonhosted.org/packages/d1/67/8197b7e7e602150938049d8e7f30de1660cfb87e4c8ee349b42b67bdb2e1/greenlet-3.3.2-cp314-cp314t-manylinux_2_24_s390x.manylinux_2_28_s390x.whl", hash = "sha256:59b3e2c40f6706b05a9cd299c836c6aa2378cabe25d021acd80f13abf81181cf", size = 666581, upload-time = "2026-02-20T21:02:51.526Z" }, + { url = "https://files.pythonhosted.org/packages/8e/30/3a09155fbf728673a1dea713572d2d31159f824a37c22da82127056c44e4/greenlet-3.3.2-cp314-cp314t-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b26b0f4428b871a751968285a1ac9648944cea09807177ac639b030bddebcea4", size = 657907, upload-time = "2026-02-20T20:21:05.259Z" }, + { url = "https://files.pythonhosted.org/packages/f3/fd/d05a4b7acd0154ed758797f0a43b4c0962a843bedfe980115e842c5b2d08/greenlet-3.3.2-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:1fb39a11ee2e4d94be9a76671482be9398560955c9e568550de0224e41104727", size = 1618857, upload-time = "2026-02-20T20:49:37.309Z" }, + { url = "https://files.pythonhosted.org/packages/6f/e1/50ee92a5db521de8f35075b5eff060dd43d39ebd46c2181a2042f7070385/greenlet-3.3.2-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:20154044d9085151bc309e7689d6f7ba10027f8f5a8c0676ad398b951913d89e", size = 1680010, upload-time = "2026-02-20T20:21:13.427Z" }, + { url = "https://files.pythonhosted.org/packages/29/4b/45d90626aef8e65336bed690106d1382f7a43665e2249017e9527df8823b/greenlet-3.3.2-cp314-cp314t-win_amd64.whl", hash = "sha256:c04c5e06ec3e022cbfe2cd4a846e1d4e50087444f875ff6d2c2ad8445495cf1a", size = 237086, upload-time = "2026-02-20T20:20:45.786Z" }, +] + [[package]] name = "griffe" version = "2.0.0" @@ -650,6 +684,27 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/4d/51/c936033e16d12b627ea334aaaaf42229c37620d0f15593456ab69ab48161/griffelib-2.0.0-py3-none-any.whl", hash = "sha256:01284878c966508b6d6f1dbff9b6fa607bc062d8261c5c7253cb285b06422a7f", size = 142004, upload-time = "2026-02-09T19:09:40.561Z" }, ] +[[package]] +name = "grpcio" +version = "1.78.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/06/8a/3d098f35c143a89520e568e6539cc098fcd294495910e359889ce8741c84/grpcio-1.78.0.tar.gz", hash = "sha256:7382b95189546f375c174f53a5fa873cef91c4b8005faa05cc5b3beea9c4f1c5", size = 12852416, upload-time = "2026-02-06T09:57:18.093Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/29/f2/b56e43e3c968bfe822fa6ce5bca10d5c723aa40875b48791ce1029bb78c7/grpcio-1.78.0-cp314-cp314-linux_armv7l.whl", hash = "sha256:e87cbc002b6f440482b3519e36e1313eb5443e9e9e73d6a52d43bd2004fcfd8e", size = 5920591, upload-time = "2026-02-06T09:56:20.758Z" }, + { url = "https://files.pythonhosted.org/packages/5d/81/1f3b65bd30c334167bfa8b0d23300a44e2725ce39bba5b76a2460d85f745/grpcio-1.78.0-cp314-cp314-macosx_11_0_universal2.whl", hash = "sha256:c41bc64626db62e72afec66b0c8a0da76491510015417c127bfc53b2fe6d7f7f", size = 11813685, upload-time = "2026-02-06T09:56:24.315Z" }, + { url = "https://files.pythonhosted.org/packages/0e/1c/bbe2f8216a5bd3036119c544d63c2e592bdf4a8ec6e4a1867592f4586b26/grpcio-1.78.0-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:8dfffba826efcf366b1e3ccc37e67afe676f290e13a3b48d31a46739f80a8724", size = 6487803, upload-time = "2026-02-06T09:56:27.367Z" }, + { url = "https://files.pythonhosted.org/packages/16/5c/a6b2419723ea7ddce6308259a55e8e7593d88464ce8db9f4aa857aba96fa/grpcio-1.78.0-cp314-cp314-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:74be1268d1439eaaf552c698cdb11cd594f0c49295ae6bb72c34ee31abbe611b", size = 7173206, upload-time = "2026-02-06T09:56:29.876Z" }, + { url = "https://files.pythonhosted.org/packages/df/1e/b8801345629a415ea7e26c83d75eb5dbe91b07ffe5210cc517348a8d4218/grpcio-1.78.0-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:be63c88b32e6c0f1429f1398ca5c09bc64b0d80950c8bb7807d7d7fb36fb84c7", size = 6693826, upload-time = "2026-02-06T09:56:32.305Z" }, + { url = "https://files.pythonhosted.org/packages/34/84/0de28eac0377742679a510784f049738a80424b17287739fc47d63c2439e/grpcio-1.78.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:3c586ac70e855c721bda8f548d38c3ca66ac791dc49b66a8281a1f99db85e452", size = 7277897, upload-time = "2026-02-06T09:56:34.915Z" }, + { url = "https://files.pythonhosted.org/packages/ca/9c/ad8685cfe20559a9edb66f735afdcb2b7d3de69b13666fdfc542e1916ebd/grpcio-1.78.0-cp314-cp314-musllinux_1_2_i686.whl", hash = "sha256:35eb275bf1751d2ffbd8f57cdbc46058e857cf3971041521b78b7db94bdaf127", size = 8252404, upload-time = "2026-02-06T09:56:37.553Z" }, + { url = "https://files.pythonhosted.org/packages/3c/05/33a7a4985586f27e1de4803887c417ec7ced145ebd069bc38a9607059e2b/grpcio-1.78.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:207db540302c884b8848036b80db352a832b99dfdf41db1eb554c2c2c7800f65", size = 7696837, upload-time = "2026-02-06T09:56:40.173Z" }, + { url = "https://files.pythonhosted.org/packages/73/77/7382241caf88729b106e49e7d18e3116216c778e6a7e833826eb96de22f7/grpcio-1.78.0-cp314-cp314-win32.whl", hash = "sha256:57bab6deef2f4f1ca76cc04565df38dc5713ae6c17de690721bdf30cb1e0545c", size = 4142439, upload-time = "2026-02-06T09:56:43.258Z" }, + { url = "https://files.pythonhosted.org/packages/48/b2/b096ccce418882fbfda4f7496f9357aaa9a5af1896a9a7f60d9f2b275a06/grpcio-1.78.0-cp314-cp314-win_amd64.whl", hash = "sha256:dce09d6116df20a96acfdbf85e4866258c3758180e8c49845d6ba8248b6d0bbb", size = 4929852, upload-time = "2026-02-06T09:56:45.885Z" }, +] + [[package]] name = "h11" version = "0.16.0" @@ -659,6 +714,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" }, ] +[[package]] +name = "h2" +version = "4.3.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "hpack" }, + { name = "hyperframe" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/1d/17/afa56379f94ad0fe8defd37d6eb3f89a25404ffc71d4d848893d270325fc/h2-4.3.0.tar.gz", hash = "sha256:6c59efe4323fa18b47a632221a1888bd7fde6249819beda254aeca909f221bf1", size = 2152026, upload-time = "2025-08-23T18:12:19.778Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/69/b2/119f6e6dcbd96f9069ce9a2665e0146588dc9f88f29549711853645e736a/h2-4.3.0-py3-none-any.whl", hash = "sha256:c438f029a25f7945c69e0ccf0fb951dc3f73a5f6412981daee861431b70e2bdd", size = 61779, upload-time = "2025-08-23T18:12:17.779Z" }, +] + [[package]] name = "hf-xet" version = "1.3.2" @@ -683,6 +751,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/cc/02/9a6e4ca1f3f73a164c0cd48e41b3cc56585dcc37e809250de443d673266f/hf_xet-1.3.2-cp37-abi3-win_arm64.whl", hash = "sha256:83d8ec273136171431833a6957e8f3af496bee227a0fe47c7b8b39c106d1749a", size = 3503976, upload-time = "2026-02-27T17:26:12.123Z" }, ] +[[package]] +name = "hpack" +version = "4.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/2c/48/71de9ed269fdae9c8057e5a4c0aa7402e8bb16f2c6e90b3aa53327b113f8/hpack-4.1.0.tar.gz", hash = "sha256:ec5eca154f7056aa06f196a557655c5b009b382873ac8d1e66e79e87535f1dca", size = 51276, upload-time = "2025-01-22T21:44:58.347Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/07/c6/80c95b1b2b94682a72cbdbfb85b81ae2daffa4291fbfa1b1464502ede10d/hpack-4.1.0-py3-none-any.whl", hash = "sha256:157ac792668d995c657d93111f46b4535ed114f0c9c8d672271bbec7eae1b496", size = 34357, upload-time = "2025-01-22T21:44:56.92Z" }, +] + [[package]] name = "httpcore" version = "1.0.9" @@ -726,6 +803,11 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517, upload-time = "2024-12-06T15:37:21.509Z" }, ] +[package.optional-dependencies] +http2 = [ + { name = "h2" }, +] + [[package]] name = "httpx-sse" version = "0.4.3" @@ -755,6 +837,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ec/74/2bc951622e2dbba1af9a460d93c51d15e458becd486e62c29cc0ccb08178/huggingface_hub-1.5.0-py3-none-any.whl", hash = "sha256:c9c0b3ab95a777fc91666111f3b3ede71c0cdced3614c553a64e98920585c4ee", size = 596261, upload-time = "2026-02-26T15:35:31.1Z" }, ] +[[package]] +name = "hyperframe" +version = "6.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/02/e7/94f8232d4a74cc99514c13a9f995811485a6903d48e5d952771ef6322e30/hyperframe-6.1.0.tar.gz", hash = "sha256:f630908a00854a7adeabd6382b43923a4c4cd4b821fcb527e6ab9e15382a3b08", size = 26566, upload-time = "2025-01-22T21:41:49.302Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/48/30/47d0bf6072f7252e6521f3447ccfa40b421b6824517f82854703d0f5a98b/hyperframe-6.1.0-py3-none-any.whl", hash = "sha256:b03380493a519fce58ea5af42e4a42317bf9bd425596f7a0835ffce80f1a42e5", size = 13007, upload-time = "2025-01-22T21:41:47.295Z" }, +] + [[package]] name = "identify" version = "2.6.16" @@ -1075,6 +1166,24 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979, upload-time = "2022-08-14T12:40:09.779Z" }, ] +[[package]] +name = "mem0ai" +version = "1.0.5" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "openai" }, + { name = "posthog" }, + { name = "protobuf" }, + { name = "pydantic" }, + { name = "pytz" }, + { name = "qdrant-client" }, + { name = "sqlalchemy" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/f3/79/2307e5fe1610d2ad0d08688af10cd5163861390deeb070f83449c0b65417/mem0ai-1.0.5.tar.gz", hash = "sha256:0835a0001ecac40ba2667bbf17629329c1b2f33eaa585e93a6be54d868a82f79", size = 182982, upload-time = "2026-03-03T22:27:09.488Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/78/0e/43ec9f125ebe6e8390805aa56237ee7165fc4f2b796122644cb0043e6631/mem0ai-1.0.5-py3-none-any.whl", hash = "sha256:0526814d2ec9134e21a628cc04ae0e6dc1779a579af92c481cb9fd7f7b8d17aa", size = 275991, upload-time = "2026-03-03T22:27:07.73Z" }, +] + [[package]] name = "mergedeep" version = "1.3.4" @@ -1289,6 +1398,35 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/88/b2/d0896bdcdc8d28a7fc5717c305f1a861c26e18c05047949fb371034d98bd/nodeenv-1.10.0-py2.py3-none-any.whl", hash = "sha256:5bb13e3eed2923615535339b3c620e76779af4cb4c6a90deccc9e36b274d3827", size = 23438, upload-time = "2025-12-20T14:08:52.782Z" }, ] +[[package]] +name = "numpy" +version = "2.4.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/10/8b/c265f4823726ab832de836cdd184d0986dcf94480f81e8739692a7ac7af2/numpy-2.4.3.tar.gz", hash = "sha256:483a201202b73495f00dbc83796c6ae63137a9bdade074f7648b3e32613412dd", size = 20727743, upload-time = "2026-03-09T07:58:53.426Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/70/ae/3936f79adebf8caf81bd7a599b90a561334a658be4dcc7b6329ebf4ee8de/numpy-2.4.3-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:5884ce5c7acfae1e4e1b6fde43797d10aa506074d25b531b4f54bde33c0c31d4", size = 16664563, upload-time = "2026-03-09T07:57:43.817Z" }, + { url = "https://files.pythonhosted.org/packages/9b/62/760f2b55866b496bb1fa7da2a6db076bef908110e568b02fcfc1422e2a3a/numpy-2.4.3-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:297837823f5bc572c5f9379b0c9f3a3365f08492cbdc33bcc3af174372ebb168", size = 14702161, upload-time = "2026-03-09T07:57:46.169Z" }, + { url = "https://files.pythonhosted.org/packages/32/af/a7a39464e2c0a21526fb4fb76e346fb172ebc92f6d1c7a07c2c139cc17b1/numpy-2.4.3-cp314-cp314-macosx_14_0_arm64.whl", hash = "sha256:a111698b4a3f8dcbe54c64a7708f049355abd603e619013c346553c1fd4ca90b", size = 5208738, upload-time = "2026-03-09T07:57:48.506Z" }, + { url = "https://files.pythonhosted.org/packages/29/8c/2a0cf86a59558fa078d83805589c2de490f29ed4fb336c14313a161d358a/numpy-2.4.3-cp314-cp314-macosx_14_0_x86_64.whl", hash = "sha256:4bd4741a6a676770e0e97fe9ab2e51de01183df3dcbcec591d26d331a40de950", size = 6543618, upload-time = "2026-03-09T07:57:50.591Z" }, + { url = "https://files.pythonhosted.org/packages/aa/b8/612ce010c0728b1c363fa4ea3aa4c22fe1c5da1de008486f8c2f5cb92fae/numpy-2.4.3-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:54f29b877279d51e210e0c80709ee14ccbbad647810e8f3d375561c45ef613dd", size = 15680676, upload-time = "2026-03-09T07:57:52.34Z" }, + { url = "https://files.pythonhosted.org/packages/a9/7e/4f120ecc54ba26ddf3dc348eeb9eb063f421de65c05fc961941798feea18/numpy-2.4.3-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:679f2a834bae9020f81534671c56fd0cc76dd7e5182f57131478e23d0dc59e24", size = 16613492, upload-time = "2026-03-09T07:57:54.91Z" }, + { url = "https://files.pythonhosted.org/packages/2c/86/1b6020db73be330c4b45d5c6ee4295d59cfeef0e3ea323959d053e5a6909/numpy-2.4.3-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:d84f0f881cb2225c2dfd7f78a10a5645d487a496c6668d6cc39f0f114164f3d0", size = 17031789, upload-time = "2026-03-09T07:57:57.641Z" }, + { url = "https://files.pythonhosted.org/packages/07/3a/3b90463bf41ebc21d1b7e06079f03070334374208c0f9a1f05e4ae8455e7/numpy-2.4.3-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:d213c7e6e8d211888cc359bab7199670a00f5b82c0978b9d1c75baf1eddbeac0", size = 18339941, upload-time = "2026-03-09T07:58:00.577Z" }, + { url = "https://files.pythonhosted.org/packages/a8/74/6d736c4cd962259fd8bae9be27363eb4883a2f9069763747347544c2a487/numpy-2.4.3-cp314-cp314-win32.whl", hash = "sha256:52077feedeff7c76ed7c9f1a0428558e50825347b7545bbb8523da2cd55c547a", size = 6007503, upload-time = "2026-03-09T07:58:03.331Z" }, + { url = "https://files.pythonhosted.org/packages/48/39/c56ef87af669364356bb011922ef0734fc49dad51964568634c72a009488/numpy-2.4.3-cp314-cp314-win_amd64.whl", hash = "sha256:0448e7f9caefb34b4b7dd2b77f21e8906e5d6f0365ad525f9f4f530b13df2afc", size = 12444915, upload-time = "2026-03-09T07:58:06.353Z" }, + { url = "https://files.pythonhosted.org/packages/9d/1f/ab8528e38d295fd349310807496fabb7cf9fe2e1f70b97bc20a483ea9d4a/numpy-2.4.3-cp314-cp314-win_arm64.whl", hash = "sha256:b44fd60341c4d9783039598efadd03617fa28d041fc37d22b62d08f2027fa0e7", size = 10494875, upload-time = "2026-03-09T07:58:08.734Z" }, + { url = "https://files.pythonhosted.org/packages/e6/ef/b7c35e4d5ef141b836658ab21a66d1a573e15b335b1d111d31f26c8ef80f/numpy-2.4.3-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:0a195f4216be9305a73c0e91c9b026a35f2161237cf1c6de9b681637772ea657", size = 14822225, upload-time = "2026-03-09T07:58:11.034Z" }, + { url = "https://files.pythonhosted.org/packages/cd/8d/7730fa9278cf6648639946cc816e7cc89f0d891602584697923375f801ed/numpy-2.4.3-cp314-cp314t-macosx_14_0_arm64.whl", hash = "sha256:cd32fbacb9fd1bf041bf8e89e4576b6f00b895f06d00914820ae06a616bdfef7", size = 5328769, upload-time = "2026-03-09T07:58:13.67Z" }, + { url = "https://files.pythonhosted.org/packages/47/01/d2a137317c958b074d338807c1b6a383406cdf8b8e53b075d804cc3d211d/numpy-2.4.3-cp314-cp314t-macosx_14_0_x86_64.whl", hash = "sha256:2e03c05abaee1f672e9d67bc858f300b5ccba1c21397211e8d77d98350972093", size = 6649461, upload-time = "2026-03-09T07:58:15.912Z" }, + { url = "https://files.pythonhosted.org/packages/5c/34/812ce12bc0f00272a4b0ec0d713cd237cb390666eb6206323d1cc9cedbb2/numpy-2.4.3-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:7d1ce23cce91fcea443320a9d0ece9b9305d4368875bab09538f7a5b4131938a", size = 15725809, upload-time = "2026-03-09T07:58:17.787Z" }, + { url = "https://files.pythonhosted.org/packages/25/c0/2aed473a4823e905e765fee3dc2cbf504bd3e68ccb1150fbdabd5c39f527/numpy-2.4.3-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c59020932feb24ed49ffd03704fbab89f22aa9c0d4b180ff45542fe8918f5611", size = 16655242, upload-time = "2026-03-09T07:58:20.476Z" }, + { url = "https://files.pythonhosted.org/packages/f2/c8/7e052b2fc87aa0e86de23f20e2c42bd261c624748aa8efd2c78f7bb8d8c6/numpy-2.4.3-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:9684823a78a6cd6ad7511fc5e25b07947d1d5b5e2812c93fe99d7d4195130720", size = 17080660, upload-time = "2026-03-09T07:58:23.067Z" }, + { url = "https://files.pythonhosted.org/packages/f3/3d/0876746044db2adcb11549f214d104f2e1be00f07a67edbb4e2812094847/numpy-2.4.3-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:0200b25c687033316fb39f0ff4e3e690e8957a2c3c8d22499891ec58c37a3eb5", size = 18380384, upload-time = "2026-03-09T07:58:25.839Z" }, + { url = "https://files.pythonhosted.org/packages/07/12/8160bea39da3335737b10308df4f484235fd297f556745f13092aa039d3b/numpy-2.4.3-cp314-cp314t-win32.whl", hash = "sha256:5e10da9e93247e554bb1d22f8edc51847ddd7dde52d85ce31024c1b4312bfba0", size = 6154547, upload-time = "2026-03-09T07:58:28.289Z" }, + { url = "https://files.pythonhosted.org/packages/42/f3/76534f61f80d74cc9cdf2e570d3d4eeb92c2280a27c39b0aaf471eda7b48/numpy-2.4.3-cp314-cp314t-win_amd64.whl", hash = "sha256:45f003dbdffb997a03da2d1d0cb41fbd24a87507fb41605c0420a3db5bd4667b", size = 12633645, upload-time = "2026-03-09T07:58:30.384Z" }, + { url = "https://files.pythonhosted.org/packages/1f/b6/7c0d4334c15983cec7f92a69e8ce9b1e6f31857e5ee3a413ac424e6bd63d/numpy-2.4.3-cp314-cp314t-win_arm64.whl", hash = "sha256:4d382735cecd7bcf090172489a525cd7d4087bc331f7df9f60ddc9a296cf208e", size = 10565454, upload-time = "2026-03-09T07:58:33.031Z" }, +] + [[package]] name = "openai" version = "2.24.0" @@ -1357,6 +1495,35 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/dd/34/b6f19941adcdaf415b5e8a8d577499f5b6a76b59cbae37f9b125a9ffe9f2/polyfactory-3.3.0-py3-none-any.whl", hash = "sha256:686abcaa761930d3df87b91e95b26b8d8cb9fdbbbe0b03d5f918acff5c72606e", size = 62707, upload-time = "2026-02-22T09:46:25.985Z" }, ] +[[package]] +name = "portalocker" +version = "3.2.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pywin32", marker = "sys_platform == 'win32'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/5e/77/65b857a69ed876e1951e88aaba60f5ce6120c33703f7cb61a3c894b8c1b6/portalocker-3.2.0.tar.gz", hash = "sha256:1f3002956a54a8c3730586c5c77bf18fae4149e07eaf1c29fc3faf4d5a3f89ac", size = 95644, upload-time = "2025-06-14T13:20:40.03Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/4b/a6/38c8e2f318bf67d338f4d629e93b0b4b9af331f455f0390ea8ce4a099b26/portalocker-3.2.0-py3-none-any.whl", hash = "sha256:3cdc5f565312224bc570c49337bd21428bba0ef363bbcf58b9ef4a9f11779968", size = 22424, upload-time = "2025-06-14T13:20:38.083Z" }, +] + +[[package]] +name = "posthog" +version = "7.9.12" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "backoff" }, + { name = "distro" }, + { name = "python-dateutil" }, + { name = "requests" }, + { name = "six" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/1c/a7/2865487853061fbd62383492237b546d2d8f7c1846272350d2b9e14138cd/posthog-7.9.12.tar.gz", hash = "sha256:ebabf2eb2e1c1fbf22b0759df4644623fa43cc6c9dcbe9fd429b7937d14251ec", size = 176828, upload-time = "2026-03-12T09:01:15.184Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/65/a9/7a803aed5a5649cf78ea7b31e90d0080181ba21f739243e1741a1e607f1f/posthog-7.9.12-py3-none-any.whl", hash = "sha256:7175bd1698a566bfea98a016c64e3456399f8046aeeca8f1d04ae5bf6c5a38d0", size = 202469, upload-time = "2026-03-12T09:01:13.38Z" }, +] + [[package]] name = "pre-commit" version = "4.5.1" @@ -1446,6 +1613,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5b/5a/bc7b4a4ef808fa59a816c17b20c4bef6884daebbdf627ff2a161da67da19/propcache-0.4.1-py3-none-any.whl", hash = "sha256:af2a6052aeb6cf17d3e46ee169099044fd8224cbaf75c76a2ef596e8163e2237", size = 13305, upload-time = "2025-10-08T19:49:00.792Z" }, ] +[[package]] +name = "protobuf" +version = "5.29.6" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/7e/57/394a763c103e0edf87f0938dafcd918d53b4c011dfc5c8ae80f3b0452dbb/protobuf-5.29.6.tar.gz", hash = "sha256:da9ee6a5424b6b30fd5e45c5ea663aef540ca95f9ad99d1e887e819cdf9b8723", size = 425623, upload-time = "2026-02-04T22:54:40.584Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d4/88/9ee58ff7863c479d6f8346686d4636dd4c415b0cbeed7a6a7d0617639c2a/protobuf-5.29.6-cp310-abi3-win32.whl", hash = "sha256:62e8a3114992c7c647bce37dcc93647575fc52d50e48de30c6fcb28a6a291eb1", size = 423357, upload-time = "2026-02-04T22:54:25.805Z" }, + { url = "https://files.pythonhosted.org/packages/1c/66/2dc736a4d576847134fb6d80bd995c569b13cdc7b815d669050bf0ce2d2c/protobuf-5.29.6-cp310-abi3-win_amd64.whl", hash = "sha256:7e6ad413275be172f67fdee0f43484b6de5a904cc1c3ea9804cb6fe2ff366eda", size = 435175, upload-time = "2026-02-04T22:54:28.592Z" }, + { url = "https://files.pythonhosted.org/packages/06/db/49b05966fd208ae3f44dcd33837b6243b4915c57561d730a43f881f24dea/protobuf-5.29.6-cp38-abi3-macosx_10_9_universal2.whl", hash = "sha256:b5a169e664b4057183a34bdc424540e86eea47560f3c123a0d64de4e137f9269", size = 418619, upload-time = "2026-02-04T22:54:30.266Z" }, + { url = "https://files.pythonhosted.org/packages/b7/d7/48cbf6b0c3c39761e47a99cb483405f0fde2be22cf00d71ef316ce52b458/protobuf-5.29.6-cp38-abi3-manylinux2014_aarch64.whl", hash = "sha256:a8866b2cff111f0f863c1b3b9e7572dc7eaea23a7fae27f6fc613304046483e6", size = 320284, upload-time = "2026-02-04T22:54:31.782Z" }, + { url = "https://files.pythonhosted.org/packages/e3/dd/cadd6ec43069247d91f6345fa7a0d2858bef6af366dbd7ba8f05d2c77d3b/protobuf-5.29.6-cp38-abi3-manylinux2014_x86_64.whl", hash = "sha256:e3387f44798ac1106af0233c04fb8abf543772ff241169946f698b3a9a3d3ab9", size = 320478, upload-time = "2026-02-04T22:54:32.909Z" }, + { url = "https://files.pythonhosted.org/packages/5a/cb/e3065b447186cb70aa65acc70c86baf482d82bf75625bf5a2c4f6919c6a3/protobuf-5.29.6-py3-none-any.whl", hash = "sha256:6b9edb641441b2da9fa8f428760fc136a49cf97a52076010cf22a2ff73438a86", size = 173126, upload-time = "2026-02-04T22:54:39.462Z" }, +] + [[package]] name = "pycparser" version = "3.0" @@ -1694,6 +1875,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/1b/d0/397f9626e711ff749a95d96b7af99b9c566a9bb5129b8e4c10fc4d100304/python_multipart-0.0.22-py3-none-any.whl", hash = "sha256:2b2cd894c83d21bf49d702499531c7bafd057d730c201782048f7945d82de155", size = 24579, upload-time = "2026-01-25T10:15:54.811Z" }, ] +[[package]] +name = "pytz" +version = "2026.1.post1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/56/db/b8721d71d945e6a8ac63c0fc900b2067181dbb50805958d4d4661cf7d277/pytz-2026.1.post1.tar.gz", hash = "sha256:3378dde6a0c3d26719182142c56e60c7f9af7e968076f31aae569d72a0358ee1", size = 321088, upload-time = "2026-03-03T07:47:50.683Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/10/99/781fe0c827be2742bcc775efefccb3b048a3a9c6ce9aec0cbf4a101677e5/pytz-2026.1.post1-py2.py3-none-any.whl", hash = "sha256:f2fd16142fda348286a75e1a524be810bb05d444e5a081f37f7affc635035f7a", size = 510489, upload-time = "2026-03-03T07:47:49.167Z" }, +] + [[package]] name = "pywin32" version = "311" @@ -1742,6 +1932,24 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/04/11/432f32f8097b03e3cd5fe57e88efb685d964e2e5178a48ed61e841f7fdce/pyyaml_env_tag-1.1-py3-none-any.whl", hash = "sha256:17109e1a528561e32f026364712fee1264bc2ea6715120891174ed1b980d2e04", size = 4722, upload-time = "2025-05-13T15:23:59.629Z" }, ] +[[package]] +name = "qdrant-client" +version = "1.17.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "grpcio" }, + { name = "httpx", extra = ["http2"] }, + { name = "numpy" }, + { name = "portalocker" }, + { name = "protobuf" }, + { name = "pydantic" }, + { name = "urllib3" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/20/fb/c9c4cecf6e7fdff2dbaeee0de40e93fe495379eb5fe2775b184ea45315da/qdrant_client-1.17.0.tar.gz", hash = "sha256:47eb033edb9be33a4babb4d87b0d8d5eaf03d52112dca0218db7f2030bf41ba9", size = 344839, upload-time = "2026-02-19T16:03:17.069Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c1/15/dfadbc9d8c9872e8ac45fa96f5099bb2855f23426bfea1bbcdc85e64ef6e/qdrant_client-1.17.0-py3-none-any.whl", hash = "sha256:f5b452c68c42b3580d3d266446fb00d3c6e3aae89c916e16585b3c704e108438", size = 390381, upload-time = "2026-02-19T16:03:15.486Z" }, +] + [[package]] name = "questionary" version = "2.1.1" @@ -1950,6 +2158,32 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e9/44/75a9c9421471a6c4805dbf2356f7c181a29c1879239abab1ea2cc8f38b40/sniffio-1.3.1-py3-none-any.whl", hash = "sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2", size = 10235, upload-time = "2024-02-25T23:20:01.196Z" }, ] +[[package]] +name = "sqlalchemy" +version = "2.0.48" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "greenlet", marker = "platform_machine == 'AMD64' or platform_machine == 'WIN32' or platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'ppc64le' or platform_machine == 'win32' or platform_machine == 'x86_64'" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/1f/73/b4a9737255583b5fa858e0bb8e116eb94b88c910164ed2ed719147bde3de/sqlalchemy-2.0.48.tar.gz", hash = "sha256:5ca74f37f3369b45e1f6b7b06afb182af1fd5dde009e4ffd831830d98cbe5fe7", size = 9886075, upload-time = "2026-03-02T15:28:51.474Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f7/b3/f437eaa1cf028bb3c927172c7272366393e73ccd104dcf5b6963f4ab5318/sqlalchemy-2.0.48-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:e2d0d88686e3d35a76f3e15a34e8c12d73fc94c1dea1cd55782e695cc14086dd", size = 2154401, upload-time = "2026-03-02T15:49:17.24Z" }, + { url = "https://files.pythonhosted.org/packages/6c/1c/b3abdf0f402aa3f60f0df6ea53d92a162b458fca2321d8f1f00278506402/sqlalchemy-2.0.48-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:49b7bddc1eebf011ea5ab722fdbe67a401caa34a350d278cc7733c0e88fecb1f", size = 3274528, upload-time = "2026-03-02T15:50:41.489Z" }, + { url = "https://files.pythonhosted.org/packages/f2/5e/327428a034407651a048f5e624361adf3f9fbac9d0fa98e981e9c6ff2f5e/sqlalchemy-2.0.48-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:426c5ca86415d9b8945c7073597e10de9644802e2ff502b8e1f11a7a2642856b", size = 3279523, upload-time = "2026-03-02T15:53:32.962Z" }, + { url = "https://files.pythonhosted.org/packages/2a/ca/ece73c81a918add0965b76b868b7b5359e068380b90ef1656ee995940c02/sqlalchemy-2.0.48-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:288937433bd44e3990e7da2402fabc44a3c6c25d3704da066b85b89a85474ae0", size = 3224312, upload-time = "2026-03-02T15:50:42.996Z" }, + { url = "https://files.pythonhosted.org/packages/88/11/fbaf1ae91fa4ee43f4fe79661cead6358644824419c26adb004941bdce7c/sqlalchemy-2.0.48-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:8183dc57ae7d9edc1346e007e840a9f3d6aa7b7f165203a99e16f447150140d2", size = 3246304, upload-time = "2026-03-02T15:53:34.937Z" }, + { url = "https://files.pythonhosted.org/packages/fa/a8/5fb0deb13930b4f2f698c5541ae076c18981173e27dd00376dbaea7a9c82/sqlalchemy-2.0.48-cp314-cp314-win32.whl", hash = "sha256:1182437cb2d97988cfea04cf6cdc0b0bb9c74f4d56ec3d08b81e23d621a28cc6", size = 2116565, upload-time = "2026-03-02T15:54:38.321Z" }, + { url = "https://files.pythonhosted.org/packages/95/7e/e83615cb63f80047f18e61e31e8e32257d39458426c23006deeaf48f463b/sqlalchemy-2.0.48-cp314-cp314-win_amd64.whl", hash = "sha256:144921da96c08feb9e2b052c5c5c1d0d151a292c6135623c6b2c041f2a45f9e0", size = 2142205, upload-time = "2026-03-02T15:54:39.831Z" }, + { url = "https://files.pythonhosted.org/packages/83/e3/69d8711b3f2c5135e9cde5f063bc1605860f0b2c53086d40c04017eb1f77/sqlalchemy-2.0.48-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5aee45fd2c6c0f2b9cdddf48c48535e7471e42d6fb81adfde801da0bd5b93241", size = 3563519, upload-time = "2026-03-02T15:57:52.387Z" }, + { url = "https://files.pythonhosted.org/packages/f8/4f/a7cce98facca73c149ea4578981594aaa5fd841e956834931de503359336/sqlalchemy-2.0.48-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:7cddca31edf8b0653090cbb54562ca027c421c58ddde2c0685f49ff56a1690e0", size = 3528611, upload-time = "2026-03-02T16:04:42.097Z" }, + { url = "https://files.pythonhosted.org/packages/cd/7d/5936c7a03a0b0cb0fa0cc425998821c6029756b0855a8f7ee70fba1de955/sqlalchemy-2.0.48-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:7a936f1bb23d370b7c8cc079d5fce4c7d18da87a33c6744e51a93b0f9e97e9b3", size = 3472326, upload-time = "2026-03-02T15:57:54.423Z" }, + { url = "https://files.pythonhosted.org/packages/f4/33/cea7dfc31b52904efe3dcdc169eb4514078887dff1f5ae28a7f4c5d54b3c/sqlalchemy-2.0.48-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:e004aa9248e8cb0a5f9b96d003ca7c1c0a5da8decd1066e7b53f59eb8ce7c62b", size = 3478453, upload-time = "2026-03-02T16:04:44.584Z" }, + { url = "https://files.pythonhosted.org/packages/c8/95/32107c4d13be077a9cae61e9ae49966a35dc4bf442a8852dd871db31f62e/sqlalchemy-2.0.48-cp314-cp314t-win32.whl", hash = "sha256:b8438ec5594980d405251451c5b7ea9aa58dda38eb7ac35fb7e4c696712ee24f", size = 2147209, upload-time = "2026-03-02T15:52:54.274Z" }, + { url = "https://files.pythonhosted.org/packages/d2/d7/1e073da7a4bc645eb83c76067284a0374e643bc4be57f14cc6414656f92c/sqlalchemy-2.0.48-cp314-cp314t-win_amd64.whl", hash = "sha256:d854b3970067297f3a7fbd7a4683587134aa9b3877ee15aa29eea478dc68f933", size = 2182198, upload-time = "2026-03-02T15:52:55.606Z" }, + { url = "https://files.pythonhosted.org/packages/46/2c/9664130905f03db57961b8980b05cab624afd114bf2be2576628a9f22da4/sqlalchemy-2.0.48-py3-none-any.whl", hash = "sha256:a66fe406437dd65cacd96a72689a3aaaecaebbcd62d81c5ac1c0fdbeac835096", size = 1940202, upload-time = "2026-03-02T15:52:43.285Z" }, +] + [[package]] name = "sse-starlette" version = "3.3.2" @@ -1996,6 +2230,7 @@ dependencies = [ { name = "litellm" }, { name = "litestar", extra = ["brotli", "prometheus", "pydantic", "standard", "structlog"] }, { name = "mcp" }, + { name = "mem0ai" }, { name = "pydantic" }, { name = "pyjwt", extra = ["crypto"] }, { name = "pyyaml" }, @@ -2048,6 +2283,7 @@ requires-dist = [ { name = "litellm", specifier = "==1.82.1" }, { name = "litestar", extras = ["brotli", "prometheus", "pydantic", "standard", "structlog"], specifier = "==2.21.1" }, { name = "mcp", specifier = "==1.26.0" }, + { name = "mem0ai", specifier = "==1.0.5" }, { name = "pydantic", specifier = "==2.12.5" }, { name = "pyjwt", extras = ["crypto"], specifier = "==2.11.0" }, { name = "pyyaml", specifier = "==6.0.3" }, From b51c09aa59d9c61690aa23e2b6bab985c4725ce0 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 00:24:56 +0100 Subject: [PATCH 02/17] refactor: harden Mem0 adapter with review findings from 9 agents Pre-reviewed by 9 agents, 33 findings addressed: - Remove vendor-specific defaults from Mem0EmbedderConfig (now required) - Fix metadata dict mutation in publish() (use spread instead) - Add path traversal validation on Mem0BackendConfig.data_dir - Extract _validate_add_result helper for store/publish result validation - Add explicit ImportError handling in connect() - Make health_check() probe backend with lightweight get_all call - Add defensive category parsing in _extract_category (handle invalid enums) - Add missing publisher check in retract() (not a shared memory entry) - Replace C-style loop in search_shared with generator expression - Use meaningful variable names (raw_entries/filtered vs triple result) - Change logger.exception to logger.warning for self-raised errors - Add structured error/error_type kwargs to all re-raise logs - Add MemoryConnectionError to Raises docstrings for shared methods - Update factory to require embedder config (no vendor defaults) - Remove dead store_request_to_mem0_args function and its tests - Add tests: path traversal, missing id, empty content, malformed datetime, invalid category, no-publisher retract, delete-after-get failure, protocol conformance, health_check probe failure --- src/ai_company/memory/__init__.py | 6 +- .../memory/backends/mem0/__init__.py | 4 +- .../memory/backends/mem0/adapter.py | 223 +++++++++++++----- src/ai_company/memory/backends/mem0/config.py | 39 ++- .../memory/backends/mem0/mappers.py | 68 +++--- src/ai_company/memory/factory.py | 32 ++- tests/integration/memory/test_mem0_backend.py | 26 +- .../unit/memory/backends/mem0/test_adapter.py | 143 ++++++++++- .../unit/memory/backends/mem0/test_config.py | 114 +++++++-- .../unit/memory/backends/mem0/test_mappers.py | 55 +++-- tests/unit/memory/test_factory.py | 28 ++- tests/unit/memory/test_init.py | 1 + 12 files changed, 581 insertions(+), 158 deletions(-) diff --git a/src/ai_company/memory/__init__.py b/src/ai_company/memory/__init__.py index d64975b521..b2990a46db 100644 --- a/src/ai_company/memory/__init__.py +++ b/src/ai_company/memory/__init__.py @@ -9,7 +9,10 @@ directly. """ -from ai_company.memory.backends.mem0 import Mem0MemoryBackend +from ai_company.memory.backends.mem0 import ( + Mem0EmbedderConfig, + Mem0MemoryBackend, +) from ai_company.memory.capabilities import MemoryCapabilities from ai_company.memory.config import ( CompanyMemoryConfig, @@ -76,6 +79,7 @@ "DefaultTokenEstimator", "InjectionPoint", "InjectionStrategy", + "Mem0EmbedderConfig", "Mem0MemoryBackend", "MemoryBackend", "MemoryCapabilities", diff --git a/src/ai_company/memory/backends/mem0/__init__.py b/src/ai_company/memory/backends/mem0/__init__.py index 3c103e22b5..00e1235bf9 100644 --- a/src/ai_company/memory/backends/mem0/__init__.py +++ b/src/ai_company/memory/backends/mem0/__init__.py @@ -1,6 +1,6 @@ """Mem0-backed agent memory — adapter, config, and mappers.""" from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend -from ai_company.memory.backends.mem0.config import Mem0BackendConfig +from ai_company.memory.backends.mem0.config import Mem0BackendConfig, Mem0EmbedderConfig -__all__ = ["Mem0BackendConfig", "Mem0MemoryBackend"] +__all__ = ["Mem0BackendConfig", "Mem0EmbedderConfig", "Mem0MemoryBackend"] diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 802887ac43..2333c123ad 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -41,7 +41,6 @@ MEMORY_BACKEND_CONNECTED, MEMORY_BACKEND_CONNECTING, MEMORY_BACKEND_CONNECTION_FAILED, - MEMORY_BACKEND_CREATED, MEMORY_BACKEND_DISCONNECTED, MEMORY_BACKEND_DISCONNECTING, MEMORY_BACKEND_HEALTH_CHECK, @@ -56,6 +55,7 @@ MEMORY_ENTRY_RETRIEVED, MEMORY_ENTRY_STORE_FAILED, MEMORY_ENTRY_STORED, + MEMORY_MODEL_INVALID, MEMORY_SHARED_PUBLISH_FAILED, MEMORY_SHARED_PUBLISHED, MEMORY_SHARED_RETRACT_FAILED, @@ -73,6 +73,31 @@ _PUBLISHER_KEY: str = "_synthorg_publisher" +def _validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: + """Extract and validate the memory ID from a Mem0 ``add`` result. + + Args: + result: Raw result dict from ``Memory.add()``. + context: Human-readable context for error messages + (e.g. ``"store"`` or ``"shared publish"``). + + Returns: + The backend-assigned memory ID. + + Raises: + MemoryStoreError: If the result is missing or malformed. + """ + results_list = result.get("results", []) + if not results_list: + msg = f"Mem0 add returned no results for {context}" + raise MemoryStoreError(msg) + first = results_list[0] + if "id" not in first: + msg = f"Mem0 add result missing 'id' for {context}: keys={list(first.keys())}" + raise MemoryStoreError(msg) + return NotBlankStr(str(first["id"])) + + class Mem0MemoryBackend: """Mem0-backed agent memory backend. @@ -94,12 +119,6 @@ def __init__( self._max_memories_per_agent = max_memories_per_agent self._client: Any = None self._connected = False - logger.debug( - MEMORY_BACKEND_CREATED, - backend="mem0", - data_dir=mem0_config.data_dir, - collection_name=mem0_config.collection_name, - ) # ── Lifecycle ───────────────────────────────────────────────── @@ -109,16 +128,26 @@ async def connect(self) -> None: Creates the Mem0 ``Memory`` client with embedded Qdrant. Raises: - MemoryConnectionError: If Mem0 initialization fails. + MemoryConnectionError: If Mem0 is not installed or + initialization fails. """ logger.info(MEMORY_BACKEND_CONNECTING, backend="mem0") try: from mem0 import Memory # noqa: PLC0415 - + except ImportError as exc: + logger.warning( + MEMORY_BACKEND_CONNECTION_FAILED, + backend="mem0", + error=str(exc), + error_type="ImportError", + ) + msg = "mem0 package is not installed" + raise MemoryConnectionError(msg) from exc + try: config_dict = build_mem0_config_dict(self._mem0_config) client = await asyncio.to_thread(Memory.from_config, config_dict) except Exception as exc: - logger.exception( + logger.warning( MEMORY_BACKEND_CONNECTION_FAILED, backend="mem0", error=str(exc), @@ -143,16 +172,39 @@ async def disconnect(self) -> None: async def health_check(self) -> bool: """Check whether the Mem0 backend is healthy. + Probes the backend with a lightweight ``get_all`` call to + verify the connection is functional, not just flagged as + connected. + Returns: - ``True`` if connected, ``False`` otherwise. + ``True`` if the backend responds, ``False`` otherwise. """ - healthy = self._connected and self._client is not None + if not self._connected or self._client is None: + logger.debug( + MEMORY_BACKEND_HEALTH_CHECK, + backend="mem0", + healthy=False, + ) + return False + try: + await asyncio.to_thread( + self._client.get_all, + user_id=_SHARED_NAMESPACE, + limit=1, + ) + except Exception: + logger.debug( + MEMORY_BACKEND_HEALTH_CHECK, + backend="mem0", + healthy=False, + ) + return False logger.debug( MEMORY_BACKEND_HEALTH_CHECK, backend="mem0", - healthy=healthy, + healthy=True, ) - return healthy + return True @property def is_connected(self) -> bool: @@ -239,19 +291,17 @@ async def store( "infer": False, } result = await asyncio.to_thread(self._client.add, **kwargs) - results_list = result.get("results", []) - if not results_list: - msg = "Mem0 add returned no results" - raise MemoryStoreError(msg) # noqa: TRY301 - memory_id = NotBlankStr(str(results_list[0]["id"])) + memory_id = _validate_add_result(result, context="store") except MemoryStoreError: - logger.exception( + logger.warning( MEMORY_ENTRY_STORE_FAILED, agent_id=agent_id, + error="Mem0 add returned no results or missing id", + error_type="MemoryStoreError", ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_ENTRY_STORE_FAILED, agent_id=agent_id, error=str(exc), @@ -300,9 +350,15 @@ async def retrieve( ) entries = apply_post_filters(entries, query) except MemoryRetrievalError: + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + agent_id=agent_id, + error="retrieval error during result mapping", + error_type="MemoryRetrievalError", + ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_ENTRY_RETRIEVAL_FAILED, agent_id=agent_id, error=str(exc), @@ -349,9 +405,16 @@ async def get( return None entry = mem0_result_to_entry(raw, str(agent_id)) except MemoryRetrievalError: + logger.warning( + MEMORY_ENTRY_FETCH_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error="retrieval error during result mapping", + error_type="MemoryRetrievalError", + ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_ENTRY_FETCH_FAILED, agent_id=agent_id, memory_id=memory_id, @@ -402,9 +465,16 @@ async def delete( return False await asyncio.to_thread(self._client.delete, str(memory_id)) except MemoryStoreError: + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error="delete operation failed", + error_type="MemoryStoreError", + ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_ENTRY_DELETE_FAILED, agent_id=agent_id, memory_id=memory_id, @@ -430,8 +500,8 @@ async def count( ) -> int: """Count memory entries for an agent. - Note: This uses ``get_all()`` internally, which is O(n). - Acceptable because ``count()`` is not on the hot path. + Uses ``get_all()`` internally — O(n) in the agent's memory + count. Acceptable because ``count()`` is not on the hot path. Args: agent_id: Owning agent identifier. @@ -453,15 +523,21 @@ async def count( ) raw_list = raw_result.get("results", []) if category is None: - count = len(raw_list) + total = len(raw_list) else: - count = sum( + total = sum( 1 for item in raw_list if _extract_category(item) == category ) except MemoryRetrievalError: + logger.warning( + MEMORY_ENTRY_COUNT_FAILED, + agent_id=agent_id, + error="count query failed", + error_type="MemoryRetrievalError", + ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_ENTRY_COUNT_FAILED, agent_id=agent_id, error=str(exc), @@ -473,10 +549,10 @@ async def count( logger.info( MEMORY_ENTRY_COUNTED, agent_id=agent_id, - count=count, + count=total, category=category.value if category else None, ) - return count + return total # ── SharedKnowledgeStore ────────────────────────────────────── @@ -498,12 +574,15 @@ async def publish( The backend-assigned shared memory ID. Raises: + MemoryConnectionError: If the backend is not connected. MemoryStoreError: If the publish operation fails. """ self._require_connected() try: - metadata = build_mem0_metadata(request) - metadata[_PUBLISHER_KEY] = str(agent_id) + metadata = { + **build_mem0_metadata(request), + _PUBLISHER_KEY: str(agent_id), + } kwargs = { "messages": [ {"role": "user", "content": request.content}, @@ -513,19 +592,17 @@ async def publish( "infer": False, } result = await asyncio.to_thread(self._client.add, **kwargs) - results_list = result.get("results", []) - if not results_list: - msg = "Mem0 add returned no results for shared publish" - raise MemoryStoreError(msg) # noqa: TRY301 - memory_id = NotBlankStr(str(results_list[0]["id"])) + memory_id = _validate_add_result(result, context="shared publish") except MemoryStoreError: - logger.exception( + logger.warning( MEMORY_SHARED_PUBLISH_FAILED, agent_id=agent_id, + error="publish operation returned no results or missing id", + error_type="MemoryStoreError", ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_SHARED_PUBLISH_FAILED, agent_id=agent_id, error=str(exc), @@ -557,6 +634,7 @@ async def search_shared( Matching shared memory entries ordered by relevance. Raises: + MemoryConnectionError: If the backend is not connected. MemoryRetrievalError: If the search fails. """ self._require_connected() @@ -576,24 +654,26 @@ async def search_shared( ) raw_list = raw_result.get("results", []) - entries: list[MemoryEntry] = [] - for item in raw_list: - publisher = _extract_publisher(item) - entry = mem0_result_to_entry( + raw_entries = tuple( + mem0_result_to_entry( item, - publisher or _SHARED_NAMESPACE, + _extract_publisher(item) or _SHARED_NAMESPACE, ) - entries.append(entry) - - result = tuple(entries) - result = apply_post_filters(result, query) + for item in raw_list + ) + filtered = apply_post_filters(raw_entries, query) if exclude_agent is not None: - result = tuple(e for e in result if e.agent_id != exclude_agent) + filtered = tuple(e for e in filtered if e.agent_id != exclude_agent) except MemoryRetrievalError: + logger.warning( + MEMORY_SHARED_SEARCH_FAILED, + error="search failed during result mapping", + error_type="MemoryRetrievalError", + ) raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_SHARED_SEARCH_FAILED, error=str(exc), error_type=type(exc).__name__, @@ -603,10 +683,10 @@ async def search_shared( else: logger.info( MEMORY_SHARED_SEARCHED, - count=len(result), + count=len(filtered), exclude_agent=exclude_agent, ) - return result + return filtered async def retract( self, @@ -625,7 +705,9 @@ async def retract( ``True`` if retracted, ``False`` if not found. Raises: - MemoryStoreError: If the retraction operation fails. + MemoryConnectionError: If the backend is not connected. + MemoryStoreError: If the retraction operation fails or + ownership verification fails. """ self._require_connected() try: @@ -640,6 +722,19 @@ async def retract( return False publisher = _extract_publisher(raw) + if publisher is None: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="not a shared memory entry (no publisher)", + ) + msg = ( + f"Memory {memory_id} is not a shared memory entry " + f"(no publisher metadata)" + ) + raise MemoryStoreError(msg) # noqa: TRY301 + if publisher != str(agent_id): logger.warning( MEMORY_SHARED_RETRACT_FAILED, @@ -658,7 +753,7 @@ async def retract( except MemoryStoreError: raise except Exception as exc: - logger.exception( + logger.warning( MEMORY_SHARED_RETRACT_FAILED, agent_id=agent_id, memory_id=memory_id, @@ -681,13 +776,27 @@ async def retract( def _extract_category(raw: dict[str, Any]) -> MemoryCategory: - """Extract the memory category from a Mem0 result dict.""" + """Extract the memory category from a Mem0 result dict. + + Returns ``MemoryCategory.WORKING`` if the category is missing + or unrecognised. + """ metadata = raw.get("metadata", {}) if not metadata: return MemoryCategory.WORKING cat_str = metadata.get("_synthorg_category") if cat_str: - return MemoryCategory(cat_str) + try: + return MemoryCategory(cat_str) + except ValueError: + logger.warning( + MEMORY_MODEL_INVALID, + field="category", + raw_value=cat_str, + reason="unrecognized category in _extract_category, " + "defaulting to WORKING", + ) + return MemoryCategory.WORKING return MemoryCategory.WORKING diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 43ad2b671a..28665ecd5d 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -5,9 +5,10 @@ ``Memory.from_config()`` expects. """ -from typing import Any +from pathlib import PurePosixPath, PureWindowsPath +from typing import Any, Self -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator from ai_company.core.types import NotBlankStr # noqa: TC001 from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 @@ -16,21 +17,24 @@ class Mem0EmbedderConfig(BaseModel): """Embedder settings for Mem0. + ``provider`` and ``model`` are required — callers must supply them + explicitly so that vendor-specific identifiers stay out of source + defaults. Pass values that the Mem0 SDK recognises (e.g. via + company YAML config). + Attributes: - provider: Embedding provider name. - model: Embedding model identifier. + provider: Embedding provider name (Mem0 SDK identifier). + model: Embedding model identifier (Mem0 SDK identifier). dims: Embedding vector dimensions. """ model_config = ConfigDict(frozen=True, allow_inf_nan=False) provider: NotBlankStr = Field( - default="openai", - description="Embedding provider name", + description="Embedding provider name (Mem0 SDK identifier)", ) model: NotBlankStr = Field( - default="text-embedding-3-small", - description="Embedding model identifier", + description="Embedding model identifier (Mem0 SDK identifier)", ) dims: int = Field( default=1536, @@ -45,7 +49,7 @@ class Mem0BackendConfig(BaseModel): Attributes: data_dir: Directory for Mem0 data persistence. collection_name: Qdrant collection name. - embedder: Embedder settings. + embedder: Embedder settings (required — no defaults). """ model_config = ConfigDict(frozen=True, allow_inf_nan=False) @@ -59,10 +63,20 @@ class Mem0BackendConfig(BaseModel): description="Qdrant collection name", ) embedder: Mem0EmbedderConfig = Field( - default_factory=Mem0EmbedderConfig, description="Embedder settings", ) + @model_validator(mode="after") + def _reject_traversal(self) -> Self: + """Reject parent-directory traversal to prevent path escapes.""" + parts = ( + PureWindowsPath(self.data_dir).parts + PurePosixPath(self.data_dir).parts + ) + if ".." in parts: + msg = "data_dir must not contain parent-directory traversal (..)" + raise ValueError(msg) + return self + def build_mem0_config_dict(config: Mem0BackendConfig) -> dict[str, Any]: """Build the dict that ``Memory.from_config()`` expects. @@ -95,15 +109,20 @@ def build_mem0_config_dict(config: Mem0BackendConfig) -> dict[str, Any]: def build_config_from_company_config( config: CompanyMemoryConfig, + *, + embedder: Mem0EmbedderConfig, ) -> Mem0BackendConfig: """Derive a ``Mem0BackendConfig`` from the top-level memory config. Args: config: Company-wide memory configuration. + embedder: Embedder settings (provider and model must be + supplied explicitly to avoid vendor names in defaults). Returns: Mem0-specific backend configuration. """ return Mem0BackendConfig( data_dir=config.storage.data_dir, + embedder=embedder, ) diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index f8d6c52671..414465d0e2 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -1,7 +1,8 @@ """Bidirectional mapping between SynthOrg domain models and Mem0 dicts. -Pure functions — no I/O, no side effects. Each mapper handles one -direction of the conversion so the adapter stays thin. +Stateless mapping functions — no I/O, no persistent side effects. +Each mapper handles one direction of the conversion so the adapter +stays thin. """ from datetime import UTC, datetime @@ -9,16 +10,21 @@ from ai_company.core.enums import MemoryCategory from ai_company.core.types import NotBlankStr +from ai_company.memory.errors import MemoryRetrievalError from ai_company.memory.models import ( MemoryEntry, MemoryMetadata, MemoryQuery, MemoryStoreRequest, ) +from ai_company.observability import get_logger +from ai_company.observability.events.memory import MEMORY_MODEL_INVALID if TYPE_CHECKING: from pydantic import AwareDatetime +logger = get_logger(__name__) + # Metadata prefix avoids collisions with Mem0's own keys. _PREFIX = "_synthorg_" @@ -45,29 +51,6 @@ def build_mem0_metadata(request: MemoryStoreRequest) -> dict[str, Any]: return meta -def store_request_to_mem0_args( - agent_id: str, - request: MemoryStoreRequest, -) -> dict[str, Any]: - """Convert a store request to ``Memory.add()`` keyword arguments. - - Args: - agent_id: Owning agent identifier. - request: Memory store request. - - Returns: - Dict of kwargs for ``Memory.add()``. - """ - messages = [{"role": "user", "content": request.content}] - metadata = build_mem0_metadata(request) - return { - "messages": messages, - "user_id": agent_id, - "metadata": metadata, - "infer": False, - } - - def parse_mem0_datetime(raw: str | None) -> AwareDatetime | None: """Parse a datetime string from Mem0 into an aware datetime. @@ -82,7 +65,16 @@ def parse_mem0_datetime(raw: str | None) -> AwareDatetime | None: """ if not raw: return None - dt = datetime.fromisoformat(raw) + try: + dt = datetime.fromisoformat(raw) + except ValueError: + logger.warning( + MEMORY_MODEL_INVALID, + field="datetime", + raw_value=raw, + reason="malformed ISO 8601 datetime, returning None", + ) + return None if dt.tzinfo is None: dt = dt.replace(tzinfo=UTC) return dt @@ -121,7 +113,19 @@ def parse_mem0_metadata( ) category_str = raw_metadata.get(f"{_PREFIX}category") - category = MemoryCategory(category_str) if category_str else MemoryCategory.WORKING + if category_str: + try: + category = MemoryCategory(category_str) + except ValueError: + logger.warning( + MEMORY_MODEL_INVALID, + field="category", + raw_value=category_str, + reason="unrecognized category, defaulting to WORKING", + ) + category = MemoryCategory.WORKING + else: + category = MemoryCategory.WORKING confidence = raw_metadata.get(f"{_PREFIX}confidence", 1.0) source = raw_metadata.get(f"{_PREFIX}source") @@ -154,8 +158,16 @@ def mem0_result_to_entry( Returns: Domain ``MemoryEntry``. """ + if "id" not in raw: + msg = f"Mem0 result missing required 'id' field: keys={list(raw.keys())}" + raise MemoryRetrievalError(msg) memory_id = NotBlankStr(str(raw["id"])) - content = NotBlankStr(str(raw.get("memory", raw.get("data", "")))) + + raw_content = raw.get("memory") or raw.get("data") + if not raw_content or not str(raw_content).strip(): + msg = f"Mem0 result {raw.get('id', '?')} has empty content" + raise MemoryRetrievalError(msg) + content = NotBlankStr(str(raw_content)) created_at = parse_mem0_datetime(raw.get("created_at")) if created_at is None: diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index f6a88d99fc..36de270c1a 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -5,6 +5,8 @@ ``config.backend``. """ +from typing import Any + from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 from ai_company.memory.errors import MemoryConfigError from ai_company.memory.protocol import MemoryBackend # noqa: TC001 @@ -17,27 +19,51 @@ logger = get_logger(__name__) -def create_memory_backend(config: CompanyMemoryConfig) -> MemoryBackend: +def create_memory_backend( + config: CompanyMemoryConfig, + *, + embedder: Any = None, +) -> MemoryBackend: """Create a memory backend from configuration. Args: config: Memory configuration (includes backend selection and backend-specific settings). + embedder: Backend-specific embedder configuration. Required + for the ``"mem0"`` backend (must be a + ``Mem0EmbedderConfig`` instance). Returns: A new, disconnected backend instance. The caller must call ``connect()`` before use. Raises: - MemoryConfigError: If the backend is not recognized. + MemoryConfigError: If the backend is not recognized or + required configuration is missing. """ if config.backend == "mem0": from ai_company.memory.backends.mem0 import Mem0MemoryBackend # noqa: PLC0415 from ai_company.memory.backends.mem0.config import ( # noqa: PLC0415 + Mem0EmbedderConfig, build_config_from_company_config, ) - mem0_config = build_config_from_company_config(config) + if embedder is None: + msg = ( + "Mem0 backend requires an embedder configuration — " + "pass a Mem0EmbedderConfig instance" + ) + raise MemoryConfigError(msg) + if not isinstance(embedder, Mem0EmbedderConfig): + msg = ( + f"embedder must be a Mem0EmbedderConfig, got {type(embedder).__name__}" + ) + raise MemoryConfigError(msg) + + mem0_config = build_config_from_company_config( + config, + embedder=embedder, + ) backend = Mem0MemoryBackend( mem0_config=mem0_config, max_memories_per_agent=config.options.max_memories_per_agent, diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py index d07fd83dd0..49e334ab4e 100644 --- a/tests/integration/memory/test_mem0_backend.py +++ b/tests/integration/memory/test_mem0_backend.py @@ -15,7 +15,10 @@ _PUBLISHER_KEY, Mem0MemoryBackend, ) -from ai_company.memory.backends.mem0.config import Mem0BackendConfig +from ai_company.memory.backends.mem0.config import ( + Mem0BackendConfig, + Mem0EmbedderConfig, +) from ai_company.memory.models import MemoryQuery, MemoryStoreRequest from ai_company.memory.retrieval_config import MemoryRetrievalConfig from ai_company.memory.retriever import ContextInjectionStrategy @@ -23,6 +26,14 @@ pytestmark = pytest.mark.timeout(30) +def _test_embedder() -> Mem0EmbedderConfig: + """Vendor-agnostic embedder config for tests.""" + return Mem0EmbedderConfig( + provider="test-provider", + model="test-embedding-001", + ) + + @pytest.fixture def mock_client() -> MagicMock: """Mock Mem0 Memory client.""" @@ -32,7 +43,10 @@ def mock_client() -> MagicMock: @pytest.fixture def backend(mock_client: MagicMock) -> Mem0MemoryBackend: """Connected Mem0 backend with mocked client.""" - config = Mem0BackendConfig(data_dir="/tmp/test-integration") # noqa: S108 + config = Mem0BackendConfig( + data_dir="/tmp/test-integration", # noqa: S108 + embedder=_test_embedder(), + ) b = Mem0MemoryBackend(mem0_config=config, max_memories_per_agent=100) b._client = mock_client b._connected = True @@ -144,7 +158,7 @@ async def test_pipeline_prepare_messages( # Should produce at least one message with memory context assert len(messages) >= 1 # Content should include both memories (they pass min_relevance) - combined = " ".join(m.content for m in messages) + combined = " ".join(m.content for m in messages if m.content) assert "concise responses" in combined async def test_shared_knowledge_flow( @@ -156,7 +170,11 @@ async def test_shared_knowledge_flow( # Publish mock_client.add.return_value = { "results": [ - {"id": "shared-001", "memory": "company policy", "event": "ADD"}, + { + "id": "shared-001", + "memory": "company policy", + "event": "ADD", + }, ], } diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index 01c142717b..5dc3c7f927 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -1,5 +1,6 @@ """Tests for the Mem0 memory backend adapter.""" +from typing import Any from unittest.mock import MagicMock, patch import pytest @@ -10,7 +11,10 @@ _SHARED_NAMESPACE, Mem0MemoryBackend, ) -from ai_company.memory.backends.mem0.config import Mem0BackendConfig +from ai_company.memory.backends.mem0.config import ( + Mem0BackendConfig, + Mem0EmbedderConfig, +) from ai_company.memory.errors import ( MemoryConnectionError, MemoryRetrievalError, @@ -24,10 +28,21 @@ pytestmark = pytest.mark.timeout(30) +def _test_embedder() -> Mem0EmbedderConfig: + """Vendor-agnostic embedder config for tests.""" + return Mem0EmbedderConfig( + provider="test-provider", + model="test-embedding-001", + ) + + @pytest.fixture def mem0_config() -> Mem0BackendConfig: """Default Mem0 config for tests.""" - return Mem0BackendConfig(data_dir="/tmp/test-memory") # noqa: S108 + return Mem0BackendConfig( + data_dir="/tmp/test-memory", # noqa: S108 + embedder=_test_embedder(), + ) @pytest.fixture @@ -48,7 +63,7 @@ def backend( return b -def _mem0_add_result(memory_id: str = "mem-001") -> dict: +def _mem0_add_result(memory_id: str = "mem-001") -> dict[str, Any]: """Build a typical Mem0 add() return value.""" return { "results": [ @@ -62,8 +77,8 @@ def _mem0_add_result(memory_id: str = "mem-001") -> dict: def _mem0_search_result( - items: list[dict] | None = None, -) -> dict: + items: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: """Build a typical Mem0 search() return value.""" if items is None: items = [ @@ -81,7 +96,7 @@ def _mem0_search_result( return {"results": items} -def _mem0_get_result(memory_id: str = "mem-001") -> dict: +def _mem0_get_result(memory_id: str = "mem-001") -> dict[str, Any]: """Build a typical Mem0 get() return value.""" return { "id": memory_id, @@ -153,6 +168,46 @@ def test_max_memories_per_agent( assert backend.max_memories_per_agent == 100 +# ── Protocol Conformance ───────────────────────────────────────── + + +@pytest.mark.unit +class TestProtocolConformance: + """Verify Mem0MemoryBackend conforms to protocol interfaces.""" + + def test_has_memory_backend_methods( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert hasattr(backend, "connect") + assert hasattr(backend, "disconnect") + assert hasattr(backend, "health_check") + assert hasattr(backend, "store") + assert hasattr(backend, "retrieve") + assert hasattr(backend, "get") + assert hasattr(backend, "delete") + assert hasattr(backend, "count") + + def test_has_capabilities_properties( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert hasattr(backend, "supported_categories") + assert hasattr(backend, "supports_graph") + assert hasattr(backend, "supports_temporal") + assert hasattr(backend, "supports_vector_search") + assert hasattr(backend, "supports_shared_access") + assert hasattr(backend, "max_memories_per_agent") + + def test_has_shared_knowledge_methods( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert hasattr(backend, "publish") + assert hasattr(backend, "search_shared") + assert hasattr(backend, "retract") + + # ── Lifecycle ───────────────────────────────────────────────────── @@ -203,7 +258,9 @@ async def test_disconnect_when_not_connected( async def test_health_check_connected( self, backend: Mem0MemoryBackend, + mock_client: MagicMock, ) -> None: + mock_client.get_all.return_value = {"results": []} assert await backend.health_check() is True async def test_health_check_disconnected( @@ -213,6 +270,14 @@ async def test_health_check_disconnected( b = Mem0MemoryBackend(mem0_config=mem0_config) assert await b.health_check() is False + async def test_health_check_probe_failure( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.side_effect = RuntimeError("backend down") + assert await backend.health_check() is False + # ── Connection guard ────────────────────────────────────────────── @@ -317,6 +382,18 @@ async def test_store_empty_results_raises( with pytest.raises(MemoryStoreError, match="no results"): await backend.store("test-agent-001", _make_store_request()) + async def test_store_missing_id_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = { + "results": [{"memory": "no id", "event": "ADD"}], + } + + with pytest.raises(MemoryStoreError, match="missing 'id'"): + await backend.store("test-agent-001", _make_store_request()) + async def test_store_exception_wraps( self, backend: Mem0MemoryBackend, @@ -499,6 +576,17 @@ async def test_delete_exception_wraps( with pytest.raises(MemoryStoreError, match="Failed to delete"): await backend.delete("test-agent-001", "mem-001") + async def test_delete_get_ok_but_delete_fails( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = _mem0_get_result("mem-001") + mock_client.delete.side_effect = RuntimeError("delete failed") + + with pytest.raises(MemoryStoreError, match="Failed to delete"): + await backend.delete("test-agent-001", "mem-001") + # ── Count ───────────────────────────────────────────────────────── @@ -551,6 +639,34 @@ async def test_count_by_category( ) assert count == 2 + async def test_count_with_invalid_category_in_data( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Invalid category in stored data defaults to WORKING.""" + mock_client.get_all.return_value = { + "results": [ + { + "id": "m1", + "memory": "a", + "metadata": {"_synthorg_category": "bogus_category"}, + }, + { + "id": "m2", + "memory": "b", + "metadata": {"_synthorg_category": "episodic"}, + }, + ], + } + + count = await backend.count( + "test-agent-001", + category=MemoryCategory.WORKING, + ) + # "bogus_category" defaults to WORKING + assert count == 1 + async def test_count_exception_wraps( self, backend: Mem0MemoryBackend, @@ -751,6 +867,21 @@ async def test_retract_ownership_mismatch( with pytest.raises(MemoryStoreError, match="cannot retract"): await backend.retract("test-agent-001", "shared-001") + async def test_retract_no_publisher_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = { + "id": "not-shared-001", + "memory": "private content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {}, + } + + with pytest.raises(MemoryStoreError, match="not a shared memory"): + await backend.retract("test-agent-001", "not-shared-001") + async def test_retract_exception_wraps( self, backend: Mem0MemoryBackend, diff --git a/tests/unit/memory/backends/mem0/test_config.py b/tests/unit/memory/backends/mem0/test_config.py index 1025ca3630..419cf2e0e7 100644 --- a/tests/unit/memory/backends/mem0/test_config.py +++ b/tests/unit/memory/backends/mem0/test_config.py @@ -9,18 +9,27 @@ build_config_from_company_config, build_mem0_config_dict, ) -from ai_company.memory.config import CompanyMemoryConfig +from ai_company.memory.config import CompanyMemoryConfig, MemoryStorageConfig pytestmark = pytest.mark.timeout(30) +def _embedder( + *, + provider: str = "test-provider", + model: str = "test-embedding-001", + dims: int = 1536, +) -> Mem0EmbedderConfig: + """Build a test embedder config with vendor-agnostic defaults.""" + return Mem0EmbedderConfig(provider=provider, model=model, dims=dims) + + @pytest.mark.unit class TestMem0EmbedderConfig: - def test_defaults(self) -> None: - config = Mem0EmbedderConfig() - assert config.provider == "openai" - assert config.model == "text-embedding-3-small" - assert config.dims == 1536 + def test_requires_provider_and_model(self) -> None: + """Provider and model are required — no vendor-specific defaults.""" + with pytest.raises(ValidationError): + Mem0EmbedderConfig() # type: ignore[call-arg] def test_custom_values(self) -> None: config = Mem0EmbedderConfig( @@ -32,46 +41,84 @@ def test_custom_values(self) -> None: assert config.model == "test-embedding-001" assert config.dims == 768 + def test_default_dims(self) -> None: + config = _embedder() + assert config.dims == 1536 + def test_frozen(self) -> None: - config = Mem0EmbedderConfig() + config = _embedder() with pytest.raises(ValidationError): config.dims = 512 # type: ignore[misc] def test_rejects_zero_dims(self) -> None: with pytest.raises(ValidationError, match="dims"): - Mem0EmbedderConfig(dims=0) + _embedder(dims=0) def test_rejects_blank_provider(self) -> None: with pytest.raises(ValidationError): - Mem0EmbedderConfig(provider=" ") + Mem0EmbedderConfig(provider=" ", model="test-model") + + def test_rejects_blank_model(self) -> None: + with pytest.raises(ValidationError): + Mem0EmbedderConfig(provider="test-provider", model=" ") @pytest.mark.unit class TestMem0BackendConfig: - def test_defaults(self) -> None: - config = Mem0BackendConfig() + def test_defaults_with_embedder(self) -> None: + config = Mem0BackendConfig(embedder=_embedder()) assert config.data_dir == "/data/memory" assert config.collection_name == "synthorg_memories" - assert config.embedder.provider == "openai" + + def test_requires_embedder(self) -> None: + with pytest.raises(ValidationError): + Mem0BackendConfig() # type: ignore[call-arg] def test_custom_data_dir(self) -> None: - config = Mem0BackendConfig(data_dir="/tmp/test-memory") # noqa: S108 + config = Mem0BackendConfig( + data_dir="/tmp/test-memory", # noqa: S108 + embedder=_embedder(), + ) assert config.data_dir == "/tmp/test-memory" # noqa: S108 def test_custom_collection(self) -> None: - config = Mem0BackendConfig(collection_name="test-collection") + config = Mem0BackendConfig( + collection_name="test-collection", + embedder=_embedder(), + ) assert config.collection_name == "test-collection" def test_frozen(self) -> None: - config = Mem0BackendConfig() + config = Mem0BackendConfig(embedder=_embedder()) with pytest.raises(ValidationError): config.data_dir = "/other" # type: ignore[misc] + def test_rejects_parent_traversal_unix(self) -> None: + with pytest.raises(ValidationError, match="parent-directory traversal"): + Mem0BackendConfig( + data_dir="/data/../escape", + embedder=_embedder(), + ) + + def test_rejects_parent_traversal_windows(self) -> None: + with pytest.raises(ValidationError, match="parent-directory traversal"): + Mem0BackendConfig( + data_dir="C:\\data\\..\\escape", + embedder=_embedder(), + ) + + def test_accepts_valid_nested_path(self) -> None: + config = Mem0BackendConfig( + data_dir="/data/sub/dir", + embedder=_embedder(), + ) + assert config.data_dir == "/data/sub/dir" + @pytest.mark.unit class TestBuildMem0ConfigDict: def test_default_config(self) -> None: - config = Mem0BackendConfig() + config = Mem0BackendConfig(embedder=_embedder()) result = build_mem0_config_dict(config) assert result["vector_store"]["provider"] == "qdrant" @@ -80,8 +127,8 @@ def test_default_config(self) -> None: ) assert result["vector_store"]["config"]["embedding_model_dims"] == 1536 assert result["vector_store"]["config"]["path"] == "/data/memory/qdrant" - assert result["embedder"]["provider"] == "openai" - assert result["embedder"]["config"]["model"] == "text-embedding-3-small" + assert result["embedder"]["provider"] == "test-provider" + assert result["embedder"]["config"]["model"] == "test-embedding-001" assert result["history_db_path"] == "/data/memory/history.db" assert result["version"] == "v1.1" @@ -108,18 +155,35 @@ def test_custom_config(self) -> None: @pytest.mark.unit class TestBuildConfigFromCompanyConfig: def test_derives_data_dir(self) -> None: - company_config = CompanyMemoryConfig( - backend="mem0", + company_config = CompanyMemoryConfig(backend="mem0") + mem0_config = build_config_from_company_config( + company_config, + embedder=_embedder(), ) - mem0_config = build_config_from_company_config(company_config) - assert mem0_config.data_dir == company_config.storage.data_dir def test_custom_data_dir(self) -> None: company_config = CompanyMemoryConfig( backend="mem0", - storage={"data_dir": "/custom/data"}, + storage=MemoryStorageConfig(data_dir="/custom/data"), + ) + mem0_config = build_config_from_company_config( + company_config, + embedder=_embedder(), ) - mem0_config = build_config_from_company_config(company_config) - assert mem0_config.data_dir == "/custom/data" + + def test_passes_embedder_through(self) -> None: + embedder = Mem0EmbedderConfig( + provider="test-provider", + model="test-model-xl", + dims=4096, + ) + company_config = CompanyMemoryConfig(backend="mem0") + mem0_config = build_config_from_company_config( + company_config, + embedder=embedder, + ) + assert mem0_config.embedder.provider == "test-provider" + assert mem0_config.embedder.model == "test-model-xl" + assert mem0_config.embedder.dims == 4096 diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index 307403242f..ba30f5e1fc 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -15,8 +15,8 @@ parse_mem0_metadata, query_to_mem0_getall_args, query_to_mem0_search_args, - store_request_to_mem0_args, ) +from ai_company.memory.errors import MemoryRetrievalError from ai_company.memory.models import ( MemoryEntry, MemoryMetadata, @@ -86,23 +86,6 @@ def test_full_metadata(self) -> None: assert meta[f"{_PREFIX}expires_at"] == expires.isoformat() -@pytest.mark.unit -class TestStoreRequestToMem0Args: - def test_basic_conversion(self) -> None: - request = MemoryStoreRequest( - category=MemoryCategory.WORKING, - content="remember this", - ) - args = store_request_to_mem0_args("test-agent-001", request) - - assert args["messages"] == [ - {"role": "user", "content": "remember this"}, - ] - assert args["user_id"] == "test-agent-001" - assert args["infer"] is False - assert f"{_PREFIX}category" in args["metadata"] - - @pytest.mark.unit class TestParseMem0Datetime: def test_none_returns_none(self) -> None: @@ -127,6 +110,12 @@ def test_non_utc_timezone(self) -> None: assert dt is not None assert dt.utcoffset() == timedelta(hours=5, minutes=30) + def test_malformed_datetime_returns_none(self) -> None: + assert parse_mem0_datetime("not-a-date") is None + + def test_partial_datetime_returns_none(self) -> None: + assert parse_mem0_datetime("2026-13-45T99:99:99") is None + @pytest.mark.unit class TestNormalizeRelevanceScore: @@ -182,6 +171,11 @@ def test_missing_category_defaults_to_working(self) -> None: category, _metadata, _expires = parse_mem0_metadata(raw) assert category == MemoryCategory.WORKING + def test_invalid_category_defaults_to_working(self) -> None: + raw = {f"{_PREFIX}category": "nonexistent_category"} + category, _metadata, _expires = parse_mem0_metadata(raw) + assert category == MemoryCategory.WORKING + def test_empty_tags_filtered(self) -> None: raw = {f"{_PREFIX}tags": ["valid", "", " ", "also-valid"]} _category, metadata, _expires = parse_mem0_metadata(raw) @@ -245,6 +239,31 @@ def test_no_metadata(self) -> None: assert entry.category == MemoryCategory.WORKING assert entry.metadata.confidence == 1.0 + def test_missing_id_raises(self) -> None: + raw = {"memory": "no id here", "metadata": {}} + with pytest.raises(MemoryRetrievalError, match="missing required 'id'"): + mem0_result_to_entry(raw, "test-agent-001") + + def test_empty_content_raises(self) -> None: + raw = {"id": "empty-content", "memory": "", "metadata": {}} + with pytest.raises(MemoryRetrievalError, match="empty content"): + mem0_result_to_entry(raw, "test-agent-001") + + def test_whitespace_only_content_raises(self) -> None: + raw = {"id": "ws-only", "memory": " ", "metadata": {}} + with pytest.raises(MemoryRetrievalError, match="empty content"): + mem0_result_to_entry(raw, "test-agent-001") + + def test_data_key_fallback(self) -> None: + raw = { + "id": "data-key", + "data": "content via data key", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {}, + } + entry = mem0_result_to_entry(raw, "test-agent-001") + assert entry.content == "content via data key" + @pytest.mark.unit class TestQueryToMem0SearchArgs: diff --git a/tests/unit/memory/test_factory.py b/tests/unit/memory/test_factory.py index 26a9b39918..ccd3f6dfac 100644 --- a/tests/unit/memory/test_factory.py +++ b/tests/unit/memory/test_factory.py @@ -4,17 +4,27 @@ from pydantic import ValidationError from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend -from ai_company.memory.config import CompanyMemoryConfig +from ai_company.memory.backends.mem0.config import Mem0EmbedderConfig +from ai_company.memory.config import CompanyMemoryConfig, MemoryOptionsConfig +from ai_company.memory.errors import MemoryConfigError from ai_company.memory.factory import create_memory_backend pytestmark = pytest.mark.timeout(30) +def _test_embedder() -> Mem0EmbedderConfig: + """Vendor-agnostic embedder config for tests.""" + return Mem0EmbedderConfig( + provider="test-provider", + model="test-embedding-001", + ) + + @pytest.mark.unit class TestCreateMemoryBackend: def test_mem0_creates_backend(self) -> None: config = CompanyMemoryConfig(backend="mem0") - backend = create_memory_backend(config) + backend = create_memory_backend(config, embedder=_test_embedder()) assert isinstance(backend, Mem0MemoryBackend) assert backend.is_connected is False assert backend.backend_name == "mem0" @@ -22,9 +32,9 @@ def test_mem0_creates_backend(self) -> None: def test_mem0_passes_max_memories(self) -> None: config = CompanyMemoryConfig( backend="mem0", - options={"max_memories_per_agent": 500}, + options=MemoryOptionsConfig(max_memories_per_agent=500), ) - backend = create_memory_backend(config) + backend = create_memory_backend(config, embedder=_test_embedder()) assert isinstance(backend, Mem0MemoryBackend) assert backend.max_memories_per_agent == 500 @@ -32,3 +42,13 @@ def test_unknown_backend_rejected_by_config_validation(self) -> None: """Unknown backends are rejected by config validation.""" with pytest.raises(ValidationError, match="Unknown memory backend"): CompanyMemoryConfig(backend="nonexistent") + + def test_mem0_without_embedder_raises(self) -> None: + config = CompanyMemoryConfig(backend="mem0") + with pytest.raises(MemoryConfigError, match="requires an embedder"): + create_memory_backend(config) + + def test_mem0_wrong_embedder_type_raises(self) -> None: + config = CompanyMemoryConfig(backend="mem0") + with pytest.raises(MemoryConfigError, match="must be a Mem0EmbedderConfig"): + create_memory_backend(config, embedder="not-a-config") diff --git a/tests/unit/memory/test_init.py b/tests/unit/memory/test_init.py index 4ec452c863..181af21f29 100644 --- a/tests/unit/memory/test_init.py +++ b/tests/unit/memory/test_init.py @@ -16,6 +16,7 @@ def test_all_exports_importable(self) -> None: def test_all_has_expected_names(self) -> None: expected = { "ArchivalStore", + "Mem0EmbedderConfig", "Mem0MemoryBackend", "CompanyMemoryConfig", "ConsolidationConfig", From d1a552aa69289853dc0ef92bba7f8fef6422f825 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 06:55:06 +0100 Subject: [PATCH 03/17] fix: address 29 PR review findings from local agents and external reviewers - Add MemoryError/RecursionError guards before all except Exception blocks - Capture exceptions in re-raise blocks for structured logging - Move helper functions (validate_add_result, extract_category, extract_publisher) from adapter.py to mappers.py to keep adapter under 800-line limit - Add type/tag/confidence validation in mappers (items 9, 11, 23) - Add logging before all raise statements (items 7, 20, 21) - Fix health_check exception log level from debug to warning - Add search_shared query context to error logs - Strengthen factory.py type annotation (embedder: Mem0EmbedderConfig | None) - Parametrize config traversal tests, add ImportError/MemoryError/publish tests - Update CLAUDE.md, README, roadmap, and design spec for Mem0 adapter status --- CLAUDE.md | 3 +- README.md | 2 +- docs/design/memory.md | 2 +- docs/roadmap/index.md | 1 - .../memory/backends/mem0/adapter.py | 138 +++++++----------- src/ai_company/memory/backends/mem0/config.py | 1 + .../memory/backends/mem0/mappers.py | 117 ++++++++++++++- src/ai_company/memory/factory.py | 20 ++- tests/integration/memory/test_mem0_backend.py | 6 +- .../unit/memory/backends/mem0/test_adapter.py | 61 +++++++- .../unit/memory/backends/mem0/test_config.py | 20 +-- .../unit/memory/backends/mem0/test_mappers.py | 2 +- tests/unit/memory/test_factory.py | 2 +- 13 files changed, 259 insertions(+), 116 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 4f913153dd..0a01986a81 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -94,7 +94,7 @@ src/ai_company/ core/ # Shared domain models, base classes, and resilience config (RetryConfig, RateLimiterConfig) engine/ # Agent orchestration, execution loops, parallel execution, task decomposition, routing, task assignment, centralized single-writer task state engine (TaskEngine), task lifecycle, recovery, shutdown, workspace isolation, coordination (multi-agent pipeline: TopologyDispatcher protocol, 4 dispatchers — SAS/centralized/decentralized/context-dependent, wave execution, workspace lifecycle integration), coordination error classification, and prompt policy validation hr/ # HR engine: hiring, firing, onboarding, offboarding, agent registry, performance tracking (task metrics, collaboration scoring, trend detection), promotion/demotion (criteria evaluation, approval strategies, model mapping) - memory/ # Persistent agent memory (Mem0 initial, custom stack future — see Decision Log), retrieval pipeline (ranking, injection, context formatting, non-inferable filtering), shared org memory (org/), consolidation/archival (consolidation/) + memory/ # Persistent agent memory (pluggable MemoryBackend protocol), backends/ (Mem0 adapter: backends/mem0/), retrieval pipeline (ranking, injection, context formatting, non-inferable filtering), shared org memory (org/), consolidation/archival (consolidation/) persistence/ # Operational data persistence — pluggable PersistenceBackend protocol, SQLite initial (see Memory & Persistence design page) observability/ # Structured logging, correlation tracking, log sinks providers/ # LLM provider abstraction (LiteLLM adapter) @@ -205,4 +205,5 @@ src/ai_company/ - **Pinned**: all versions use `==` in `pyproject.toml` - **Groups**: `test` (pytest + plugins), `dev` (includes test + ruff, mypy, pre-commit, commitizen) +- **Optional**: `mem0ai` (Mem0 memory backend — only needed when `backend: "mem0"` is configured) - **Install**: `uv sync` installs everything (dev group is default) diff --git a/README.md b/README.md index a59c08b8ee..821d9cb412 100644 --- a/README.md +++ b/README.md @@ -127,7 +127,7 @@ graph TB ## Status -Core framework complete — agent engine, multi-agent coordination, API, security, HR, memory, and budget systems are implemented. Remaining: Mem0 adapter backend, approval workflow gates, CLI, web dashboard. See the [roadmap](docs/roadmap/index.md) for details. +Core framework complete — agent engine, multi-agent coordination, API, security, HR, memory (including Mem0 backend adapter), and budget systems are implemented. Remaining: approval workflow gates, CLI, web dashboard. See the [roadmap](docs/roadmap/index.md) for details. ## License diff --git a/docs/design/memory.md b/docs/design/memory.md index 4e676aea98..25db5ca432 100644 --- a/docs/design/memory.md +++ b/docs/design/memory.md @@ -30,7 +30,7 @@ configuration without modifying application code. +----------+----------+-----------+---------------+ | Storage Backend | | SQLite / PostgreSQL / File-based | -| + Mem0 (initial) / Custom Stack (future) | +| + Mem0 (initial, implemented) / Custom (future)| | See Decision Log | +-------------------------------------------------+ ``` diff --git a/docs/roadmap/index.md b/docs/roadmap/index.md index 39d0419fca..033f1bf047 100644 --- a/docs/roadmap/index.md +++ b/docs/roadmap/index.md @@ -22,7 +22,6 @@ The SynthOrg core framework is complete. The following subsystems are built and | Area | Description | |------|-------------| -| **Mem0 adapter** | Concrete `MemoryBackend` implementation using the Mem0 library | | **Approval workflow gates** | Runtime wiring for human-in-the-loop approval queues | | **CLI** | Terminal interface wrapping the REST API (may not be needed) | | **Web dashboard** | Vue 3 frontend for monitoring and managing the synthetic organization | diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 2333c123ad..4d8f487bc5 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -18,11 +18,15 @@ build_mem0_config_dict, ) from ai_company.memory.backends.mem0.mappers import ( + _PUBLISHER_KEY, apply_post_filters, build_mem0_metadata, + extract_category, + extract_publisher, mem0_result_to_entry, query_to_mem0_getall_args, query_to_mem0_search_args, + validate_add_result, ) from ai_company.memory.errors import ( MemoryConnectionError, @@ -37,6 +41,7 @@ MemoryQuery, MemoryStoreRequest, ) + from ai_company.observability.events.memory import ( MEMORY_BACKEND_CONNECTED, MEMORY_BACKEND_CONNECTING, @@ -55,7 +60,6 @@ MEMORY_ENTRY_RETRIEVED, MEMORY_ENTRY_STORE_FAILED, MEMORY_ENTRY_STORED, - MEMORY_MODEL_INVALID, MEMORY_SHARED_PUBLISH_FAILED, MEMORY_SHARED_PUBLISHED, MEMORY_SHARED_RETRACT_FAILED, @@ -69,34 +73,6 @@ # Reserved user_id for the shared knowledge namespace. _SHARED_NAMESPACE: str = "__synthorg_shared__" -# Metadata key to track who published a shared memory. -_PUBLISHER_KEY: str = "_synthorg_publisher" - - -def _validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: - """Extract and validate the memory ID from a Mem0 ``add`` result. - - Args: - result: Raw result dict from ``Memory.add()``. - context: Human-readable context for error messages - (e.g. ``"store"`` or ``"shared publish"``). - - Returns: - The backend-assigned memory ID. - - Raises: - MemoryStoreError: If the result is missing or malformed. - """ - results_list = result.get("results", []) - if not results_list: - msg = f"Mem0 add returned no results for {context}" - raise MemoryStoreError(msg) - first = results_list[0] - if "id" not in first: - msg = f"Mem0 add result missing 'id' for {context}: keys={list(first.keys())}" - raise MemoryStoreError(msg) - return NotBlankStr(str(first["id"])) - class Mem0MemoryBackend: """Mem0-backed agent memory backend. @@ -146,6 +122,8 @@ async def connect(self) -> None: try: config_dict = build_mem0_config_dict(self._mem0_config) client = await asyncio.to_thread(Memory.from_config, config_dict) + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_BACKEND_CONNECTION_FAILED, @@ -192,8 +170,10 @@ async def health_check(self) -> bool: user_id=_SHARED_NAMESPACE, limit=1, ) + except MemoryError, RecursionError: + raise except Exception: - logger.debug( + logger.warning( MEMORY_BACKEND_HEALTH_CHECK, backend="mem0", healthy=False, @@ -291,15 +271,17 @@ async def store( "infer": False, } result = await asyncio.to_thread(self._client.add, **kwargs) - memory_id = _validate_add_result(result, context="store") - except MemoryStoreError: + memory_id = validate_add_result(result, context="store") + except MemoryStoreError as exc: logger.warning( MEMORY_ENTRY_STORE_FAILED, agent_id=agent_id, - error="Mem0 add returned no results or missing id", + error=str(exc), error_type="MemoryStoreError", ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_ENTRY_STORE_FAILED, @@ -325,6 +307,9 @@ async def retrieve( ) -> tuple[MemoryEntry, ...]: """Retrieve memories for an agent, ordered by relevance. + Uses ``search()`` when ``query.text`` is set, otherwise falls + back to ``get_all()`` for unfiltered retrieval. + Args: agent_id: Owning agent identifier. query: Retrieval parameters. @@ -349,14 +334,16 @@ async def retrieve( mem0_result_to_entry(item, str(agent_id)) for item in raw_list ) entries = apply_post_filters(entries, query) - except MemoryRetrievalError: + except MemoryRetrievalError as exc: logger.warning( MEMORY_ENTRY_RETRIEVAL_FAILED, agent_id=agent_id, - error="retrieval error during result mapping", + error=str(exc), error_type="MemoryRetrievalError", ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_ENTRY_RETRIEVAL_FAILED, @@ -404,15 +391,17 @@ async def get( ) return None entry = mem0_result_to_entry(raw, str(agent_id)) - except MemoryRetrievalError: + except MemoryRetrievalError as exc: logger.warning( MEMORY_ENTRY_FETCH_FAILED, agent_id=agent_id, memory_id=memory_id, - error="retrieval error during result mapping", + error=str(exc), error_type="MemoryRetrievalError", ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_ENTRY_FETCH_FAILED, @@ -464,15 +453,17 @@ async def delete( ) return False await asyncio.to_thread(self._client.delete, str(memory_id)) - except MemoryStoreError: + except MemoryStoreError as exc: logger.warning( MEMORY_ENTRY_DELETE_FAILED, agent_id=agent_id, memory_id=memory_id, - error="delete operation failed", + error=str(exc), error_type="MemoryStoreError", ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_ENTRY_DELETE_FAILED, @@ -526,16 +517,18 @@ async def count( total = len(raw_list) else: total = sum( - 1 for item in raw_list if _extract_category(item) == category + 1 for item in raw_list if extract_category(item) == category ) - except MemoryRetrievalError: + except MemoryRetrievalError as exc: logger.warning( MEMORY_ENTRY_COUNT_FAILED, agent_id=agent_id, - error="count query failed", + error=str(exc), error_type="MemoryRetrievalError", ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_ENTRY_COUNT_FAILED, @@ -592,15 +585,17 @@ async def publish( "infer": False, } result = await asyncio.to_thread(self._client.add, **kwargs) - memory_id = _validate_add_result(result, context="shared publish") - except MemoryStoreError: + memory_id = validate_add_result(result, context="shared publish") + except MemoryStoreError as exc: logger.warning( MEMORY_SHARED_PUBLISH_FAILED, agent_id=agent_id, - error="publish operation returned no results or missing id", + error=str(exc), error_type="MemoryStoreError", ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_SHARED_PUBLISH_FAILED, @@ -657,7 +652,7 @@ async def search_shared( raw_entries = tuple( mem0_result_to_entry( item, - _extract_publisher(item) or _SHARED_NAMESPACE, + extract_publisher(item) or _SHARED_NAMESPACE, ) for item in raw_list ) @@ -665,18 +660,24 @@ async def search_shared( if exclude_agent is not None: filtered = tuple(e for e in filtered if e.agent_id != exclude_agent) - except MemoryRetrievalError: + except MemoryRetrievalError as exc: logger.warning( MEMORY_SHARED_SEARCH_FAILED, - error="search failed during result mapping", + error=str(exc), error_type="MemoryRetrievalError", + query_text=query.text, + exclude_agent=exclude_agent, ) raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_SHARED_SEARCH_FAILED, error=str(exc), error_type=type(exc).__name__, + query_text=query.text, + exclude_agent=exclude_agent, ) msg = f"Failed to search shared knowledge: {exc}" raise MemoryRetrievalError(msg) from exc @@ -721,7 +722,7 @@ async def retract( ) return False - publisher = _extract_publisher(raw) + publisher = extract_publisher(raw) if publisher is None: logger.warning( MEMORY_SHARED_RETRACT_FAILED, @@ -752,6 +753,8 @@ async def retract( await asyncio.to_thread(self._client.delete, str(memory_id)) except MemoryStoreError: raise + except MemoryError, RecursionError: + raise except Exception as exc: logger.warning( MEMORY_SHARED_RETRACT_FAILED, @@ -770,40 +773,3 @@ async def retract( found=True, ) return True - - -# ── Module-level helpers ────────────────────────────────────────── - - -def _extract_category(raw: dict[str, Any]) -> MemoryCategory: - """Extract the memory category from a Mem0 result dict. - - Returns ``MemoryCategory.WORKING`` if the category is missing - or unrecognised. - """ - metadata = raw.get("metadata", {}) - if not metadata: - return MemoryCategory.WORKING - cat_str = metadata.get("_synthorg_category") - if cat_str: - try: - return MemoryCategory(cat_str) - except ValueError: - logger.warning( - MEMORY_MODEL_INVALID, - field="category", - raw_value=cat_str, - reason="unrecognized category in _extract_category, " - "defaulting to WORKING", - ) - return MemoryCategory.WORKING - return MemoryCategory.WORKING - - -def _extract_publisher(raw: dict[str, Any]) -> str | None: - """Extract the publisher agent ID from a shared memory dict.""" - metadata = raw.get("metadata", {}) - if not metadata: - return None - value: str | None = metadata.get(_PUBLISHER_KEY) - return value diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 28665ecd5d..965898f6ce 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -103,6 +103,7 @@ def build_mem0_config_dict(config: Mem0BackendConfig) -> dict[str, Any]: }, }, "history_db_path": f"{config.data_dir}/history.db", + # Mem0 config schema version — required by Memory.from_config(). "version": "v1.1", } diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 414465d0e2..f04188be70 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -10,7 +10,7 @@ from ai_company.core.enums import MemoryCategory from ai_company.core.types import NotBlankStr -from ai_company.memory.errors import MemoryRetrievalError +from ai_company.memory.errors import MemoryRetrievalError, MemoryStoreError from ai_company.memory.models import ( MemoryEntry, MemoryMetadata, @@ -18,7 +18,10 @@ MemoryStoreRequest, ) from ai_company.observability import get_logger -from ai_company.observability.events.memory import MEMORY_MODEL_INVALID +from ai_company.observability.events.memory import ( + MEMORY_ENTRY_STORE_FAILED, + MEMORY_MODEL_INVALID, +) if TYPE_CHECKING: from pydantic import AwareDatetime @@ -28,6 +31,9 @@ # Metadata prefix avoids collisions with Mem0's own keys. _PREFIX = "_synthorg_" +# Metadata key to track who published a shared memory. +_PUBLISHER_KEY: str = "_synthorg_publisher" + def build_mem0_metadata(request: MemoryStoreRequest) -> dict[str, Any]: """Serialize a store request's metadata into Mem0-compatible dict. @@ -67,7 +73,7 @@ def parse_mem0_datetime(raw: str | None) -> AwareDatetime | None: return None try: dt = datetime.fromisoformat(raw) - except ValueError: + except ValueError, TypeError: logger.warning( MEMORY_MODEL_INVALID, field="datetime", @@ -127,9 +133,23 @@ def parse_mem0_metadata( else: category = MemoryCategory.WORKING - confidence = raw_metadata.get(f"{_PREFIX}confidence", 1.0) + raw_confidence = raw_metadata.get(f"{_PREFIX}confidence", 1.0) + try: + confidence = float(raw_confidence) + except ValueError, TypeError: + logger.warning( + MEMORY_MODEL_INVALID, + field="confidence", + raw_value=raw_confidence, + reason="non-numeric confidence, defaulting to 1.0", + ) + confidence = 1.0 source = raw_metadata.get(f"{_PREFIX}source") raw_tags = raw_metadata.get(f"{_PREFIX}tags", ()) + if isinstance(raw_tags, str): + raw_tags = [raw_tags] + elif not isinstance(raw_tags, (list, tuple)): + raw_tags = () tags = tuple(NotBlankStr(t) for t in raw_tags if t and str(t).strip()) expires_at = parse_mem0_datetime( @@ -158,14 +178,27 @@ def mem0_result_to_entry( Returns: Domain ``MemoryEntry``. """ - if "id" not in raw: - msg = f"Mem0 result missing required 'id' field: keys={list(raw.keys())}" + raw_id = raw.get("id") + if raw_id is None or not str(raw_id).strip(): + msg = f"Mem0 result has missing or blank 'id': keys={list(raw.keys())}" + logger.warning( + MEMORY_MODEL_INVALID, + field="id", + raw_value=raw_id, + reason=msg, + ) raise MemoryRetrievalError(msg) - memory_id = NotBlankStr(str(raw["id"])) + memory_id = NotBlankStr(str(raw_id)) raw_content = raw.get("memory") or raw.get("data") if not raw_content or not str(raw_content).strip(): msg = f"Mem0 result {raw.get('id', '?')} has empty content" + logger.warning( + MEMORY_MODEL_INVALID, + field="content", + raw_value=raw_content, + reason=msg, + ) raise MemoryRetrievalError(msg) content = NotBlankStr(str(raw_content)) @@ -211,6 +244,12 @@ def query_to_mem0_search_args( """ if query.text is None: msg = "search requires query.text to be set" + logger.warning( + MEMORY_MODEL_INVALID, + field="query.text", + raw_value=None, + reason=msg, + ) raise ValueError(msg) return { "query": query.text, @@ -271,3 +310,67 @@ def apply_post_filters( continue result.append(entry) return tuple(result) + + +# ── Adapter helpers (moved here to keep adapter.py under 800 lines) ── + + +def validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: + """Extract and validate the memory ID from a Mem0 ``add`` result. + + Args: + result: Raw result dict from ``Memory.add()``. + context: Human-readable context for error messages + (e.g. ``"store"`` or ``"shared publish"``). + + Returns: + The backend-assigned memory ID. + + Raises: + MemoryStoreError: If the result is missing or malformed. + """ + results_list = result.get("results") + if not isinstance(results_list, list) or not results_list: + msg = f"Mem0 add returned no results for {context}" + logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) + raise MemoryStoreError(msg) + first = results_list[0] + if "id" not in first: + msg = f"Mem0 add result missing 'id' for {context}: keys={list(first.keys())}" + logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) + raise MemoryStoreError(msg) + return NotBlankStr(str(first["id"])) + + +def extract_category(raw: dict[str, Any]) -> MemoryCategory: + """Extract the memory category from a Mem0 result dict. + + Returns ``MemoryCategory.WORKING`` if the category is missing + or unrecognised. + """ + metadata = raw.get("metadata", {}) + if not metadata: + return MemoryCategory.WORKING + cat_str = metadata.get(f"{_PREFIX}category") + if cat_str: + try: + return MemoryCategory(cat_str) + except ValueError: + logger.warning( + MEMORY_MODEL_INVALID, + field="category", + raw_value=cat_str, + reason="unrecognized category in extract_category, " + "defaulting to WORKING", + ) + return MemoryCategory.WORKING + return MemoryCategory.WORKING + + +def extract_publisher(raw: dict[str, Any]) -> str | None: + """Extract the publisher agent ID from a shared memory dict.""" + metadata = raw.get("metadata", {}) + if not metadata: + return None + value: str | None = metadata.get(_PUBLISHER_KEY) + return value diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 36de270c1a..151668b09e 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -5,9 +5,12 @@ ``config.backend``. """ -from typing import Any +from typing import TYPE_CHECKING from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 + +if TYPE_CHECKING: + from ai_company.memory.backends.mem0.config import Mem0EmbedderConfig from ai_company.memory.errors import MemoryConfigError from ai_company.memory.protocol import MemoryBackend # noqa: TC001 from ai_company.observability import get_logger @@ -22,7 +25,7 @@ def create_memory_backend( config: CompanyMemoryConfig, *, - embedder: Any = None, + embedder: Mem0EmbedderConfig | None = None, ) -> MemoryBackend: """Create a memory backend from configuration. @@ -53,11 +56,22 @@ def create_memory_backend( "Mem0 backend requires an embedder configuration — " "pass a Mem0EmbedderConfig instance" ) + logger.warning( + MEMORY_BACKEND_UNKNOWN, + backend="mem0", + error=msg, + ) raise MemoryConfigError(msg) if not isinstance(embedder, Mem0EmbedderConfig): - msg = ( + msg = ( # type: ignore[unreachable] f"embedder must be a Mem0EmbedderConfig, got {type(embedder).__name__}" ) + logger.warning( + MEMORY_BACKEND_UNKNOWN, + backend="mem0", + error=msg, + embedder_type=type(embedder).__name__, + ) raise MemoryConfigError(msg) mem0_config = build_config_from_company_config( diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py index 49e334ab4e..6372d7e49f 100644 --- a/tests/integration/memory/test_mem0_backend.py +++ b/tests/integration/memory/test_mem0_backend.py @@ -11,14 +11,12 @@ import pytest from ai_company.core.enums import MemoryCategory -from ai_company.memory.backends.mem0.adapter import ( - _PUBLISHER_KEY, - Mem0MemoryBackend, -) +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend from ai_company.memory.backends.mem0.config import ( Mem0BackendConfig, Mem0EmbedderConfig, ) +from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY from ai_company.memory.models import MemoryQuery, MemoryStoreRequest from ai_company.memory.retrieval_config import MemoryRetrievalConfig from ai_company.memory.retriever import ContextInjectionStrategy diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index 5dc3c7f927..ab4360c357 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -1,5 +1,6 @@ """Tests for the Mem0 memory backend adapter.""" +import sys from typing import Any from unittest.mock import MagicMock, patch @@ -7,7 +8,6 @@ from ai_company.core.enums import MemoryCategory from ai_company.memory.backends.mem0.adapter import ( - _PUBLISHER_KEY, _SHARED_NAMESPACE, Mem0MemoryBackend, ) @@ -15,6 +15,7 @@ Mem0BackendConfig, Mem0EmbedderConfig, ) +from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY from ai_company.memory.errors import ( MemoryConnectionError, MemoryRetrievalError, @@ -278,6 +279,28 @@ async def test_health_check_probe_failure( mock_client.get_all.side_effect = RuntimeError("backend down") assert await backend.health_check() is False + async def test_connect_import_error_raises( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + """ImportError when mem0 package is not installed.""" + b = Mem0MemoryBackend(mem0_config=mem0_config) + with ( + patch.dict(sys.modules, {"mem0": None}), + pytest.raises(MemoryConnectionError, match="not installed"), + ): + await b.connect() + + async def test_health_check_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError propagates through health_check.""" + mock_client.get_all.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.health_check() + # ── Connection guard ────────────────────────────────────────────── @@ -406,6 +429,16 @@ async def test_store_exception_wraps( assert exc_info.value.__cause__ is not None + async def test_store_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping.""" + mock_client.add.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.store("test-agent-001", _make_store_request()) + # ── Retrieve ────────────────────────────────────────────────────── @@ -495,6 +528,19 @@ async def test_retrieve_exception_wraps( MemoryQuery(text="test"), ) + async def test_retrieve_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping.""" + mock_client.search.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + # ── Get ─────────────────────────────────────────────────────────── @@ -721,6 +767,19 @@ async def test_publish_exception_wraps( with pytest.raises(MemoryStoreError, match="Failed to publish"): await backend.publish("test-agent-001", _make_store_request()) + async def test_publish_missing_id_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Publish result missing 'id' raises MemoryStoreError.""" + mock_client.add.return_value = { + "results": [{"memory": "no id", "event": "ADD"}], + } + + with pytest.raises(MemoryStoreError, match="missing 'id'"): + await backend.publish("test-agent-001", _make_store_request()) + @pytest.mark.unit class TestSearchShared: diff --git a/tests/unit/memory/backends/mem0/test_config.py b/tests/unit/memory/backends/mem0/test_config.py index 419cf2e0e7..1a4e309f46 100644 --- a/tests/unit/memory/backends/mem0/test_config.py +++ b/tests/unit/memory/backends/mem0/test_config.py @@ -93,17 +93,19 @@ def test_frozen(self) -> None: with pytest.raises(ValidationError): config.data_dir = "/other" # type: ignore[misc] - def test_rejects_parent_traversal_unix(self) -> None: + @pytest.mark.parametrize( + "data_dir", + [ + "/data/../escape", + "C:\\data\\..\\escape", + "data/../../escape", + ], + ids=["unix", "windows", "relative"], + ) + def test_rejects_parent_traversal(self, data_dir: str) -> None: with pytest.raises(ValidationError, match="parent-directory traversal"): Mem0BackendConfig( - data_dir="/data/../escape", - embedder=_embedder(), - ) - - def test_rejects_parent_traversal_windows(self) -> None: - with pytest.raises(ValidationError, match="parent-directory traversal"): - Mem0BackendConfig( - data_dir="C:\\data\\..\\escape", + data_dir=data_dir, embedder=_embedder(), ) diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index ba30f5e1fc..34c6d461f8 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -241,7 +241,7 @@ def test_no_metadata(self) -> None: def test_missing_id_raises(self) -> None: raw = {"memory": "no id here", "metadata": {}} - with pytest.raises(MemoryRetrievalError, match="missing required 'id'"): + with pytest.raises(MemoryRetrievalError, match="missing or blank 'id'"): mem0_result_to_entry(raw, "test-agent-001") def test_empty_content_raises(self) -> None: diff --git a/tests/unit/memory/test_factory.py b/tests/unit/memory/test_factory.py index ccd3f6dfac..147f514153 100644 --- a/tests/unit/memory/test_factory.py +++ b/tests/unit/memory/test_factory.py @@ -51,4 +51,4 @@ def test_mem0_without_embedder_raises(self) -> None: def test_mem0_wrong_embedder_type_raises(self) -> None: config = CompanyMemoryConfig(backend="mem0") with pytest.raises(MemoryConfigError, match="must be a Mem0EmbedderConfig"): - create_memory_backend(config, embedder="not-a-config") + create_memory_backend(config, embedder="not-a-config") # type: ignore[arg-type] From 258d1a0ed4916454e6186a530dc969345e7c3c74 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 07:13:32 +0100 Subject: [PATCH 04/17] fix: address round-2 PR review findings and fix dependency review CI - Fix dependency review CI: add LicenseRef-scancode-protobuf, ZPL-2.1 to license allow-list and allow-dependencies-licenses for packages with null SPDX metadata (mem0ai, numpy, qdrant-client, posthog) - Fix validate_add_result blank-ID gap: check for None and whitespace IDs, not just missing key (Greptile finding) - Fix factory.py wrong event constant: use MEMORY_BACKEND_CONFIG_INVALID instead of MEMORY_BACKEND_UNKNOWN for embedder config errors (Greptile) - Fix retract() bare except: capture as exc and log error detail - Document count() limitation: capped at max_memories_per_agent - Add 37 new tests: validate_add_result (blank/None/whitespace/numeric ID, non-list results), extract_category, extract_publisher, MemoryError re-raise for all operations, blank-ID through store/retrieve, shared namespace fallback, count empty results, retract delete failure --- .github/workflows/dependency-review.yml | 15 +- .../memory/backends/mem0/adapter.py | 18 +- .../memory/backends/mem0/mappers.py | 10 +- src/ai_company/memory/factory.py | 7 +- src/ai_company/observability/events/memory.py | 1 + .../unit/memory/backends/mem0/test_adapter.py | 169 +++++++++++++++++- .../unit/memory/backends/mem0/test_mappers.py | 106 ++++++++++- 7 files changed, 315 insertions(+), 11 deletions(-) diff --git a/.github/workflows/dependency-review.yml b/.github/workflows/dependency-review.yml index 562052bfce..3859ba196a 100644 --- a/.github/workflows/dependency-review.yml +++ b/.github/workflows/dependency-review.yml @@ -29,11 +29,24 @@ jobs: # pycparser 3.0, sse-starlette 3.3.2 — MIT per classifiers, scancode misdetects # LGPL-*: @img/sharp-libvips-* (Astro image optimization, build-time only) # BlueOak-1.0.0: lru-cache, sax (Astro transitive deps, permissive) + # LicenseRef-scancode-protobuf: protobuf 5.29.6 (BSD-3-Clause AND scancode-protobuf, permissive) + # ZPL-2.1: pytz 2026.1 (MIT AND ZPL-2.1, Zope Public License — permissive) + # Null: mem0ai, numpy, qdrant-client, posthog — license metadata missing + # from lockfile but all are permissive (Apache-2.0 / BSD-3-Clause / MIT) allow-licenses: >- MIT, MIT-0, Apache-2.0, BSD-2-Clause, BSD-3-Clause, ISC, MPL-2.0, PSF-2.0, Unlicense, 0BSD, CC0-1.0, Python-2.0, Python-2.0.1, - LicenseRef-scancode-free-unknown, + LicenseRef-scancode-free-unknown, LicenseRef-scancode-protobuf, + ZPL-2.1, LGPL-2.0-only, LGPL-2.1-only, LGPL-3.0-only, LGPL-3.0-or-later, BlueOak-1.0.0 + # Packages with null/missing SPDX license metadata in lockfile. + # Verified manually: mem0ai (Apache-2.0), numpy (BSD-3-Clause), + # qdrant-client (Apache-2.0), posthog (MIT). + allow-dependencies-licenses: >- + pkg:pypi/mem0ai, + pkg:pypi/numpy, + pkg:pypi/qdrant-client, + pkg:pypi/posthog comment-summary-in-pr: always diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 4d8f487bc5..008ee23589 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -494,12 +494,19 @@ async def count( Uses ``get_all()`` internally — O(n) in the agent's memory count. Acceptable because ``count()`` is not on the hot path. + .. note:: + Results are capped at ``max_memories_per_agent``. If an + agent has more memories than this limit the count will be + an underestimate. This is consistent with the adapter's + store/retrieve semantics which also respect the cap. + Args: agent_id: Owning agent identifier. category: Optional category filter. Returns: - Number of matching entries. + Number of matching entries (capped at + ``max_memories_per_agent``). Raises: MemoryConnectionError: If the backend is not connected. @@ -751,7 +758,14 @@ async def retract( raise MemoryStoreError(msg) # noqa: TRY301 await asyncio.to_thread(self._client.delete, str(memory_id)) - except MemoryStoreError: + except MemoryStoreError as exc: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error=str(exc), + error_type="MemoryStoreError", + ) raise except MemoryError, RecursionError: raise diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index f04188be70..207ae52742 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -335,11 +335,15 @@ def validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) raise MemoryStoreError(msg) first = results_list[0] - if "id" not in first: - msg = f"Mem0 add result missing 'id' for {context}: keys={list(first.keys())}" + raw_id = first.get("id") + if raw_id is None or not str(raw_id).strip(): + msg = ( + f"Mem0 add result has missing or blank 'id' for {context}: " + f"keys={list(first.keys())}" + ) logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) raise MemoryStoreError(msg) - return NotBlankStr(str(first["id"])) + return NotBlankStr(str(raw_id)) def extract_category(raw: dict[str, Any]) -> MemoryCategory: diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 151668b09e..53b06f37a4 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -15,6 +15,7 @@ from ai_company.memory.protocol import MemoryBackend # noqa: TC001 from ai_company.observability import get_logger from ai_company.observability.events.memory import ( + MEMORY_BACKEND_CONFIG_INVALID, MEMORY_BACKEND_CREATED, MEMORY_BACKEND_UNKNOWN, ) @@ -57,8 +58,9 @@ def create_memory_backend( "pass a Mem0EmbedderConfig instance" ) logger.warning( - MEMORY_BACKEND_UNKNOWN, + MEMORY_BACKEND_CONFIG_INVALID, backend="mem0", + reason="missing_embedder", error=msg, ) raise MemoryConfigError(msg) @@ -67,8 +69,9 @@ def create_memory_backend( f"embedder must be a Mem0EmbedderConfig, got {type(embedder).__name__}" ) logger.warning( - MEMORY_BACKEND_UNKNOWN, + MEMORY_BACKEND_CONFIG_INVALID, backend="mem0", + reason="invalid_embedder_type", error=msg, embedder_type=type(embedder).__name__, ) diff --git a/src/ai_company/observability/events/memory.py b/src/ai_company/observability/events/memory.py index 203a3d7c91..6d22ddc90b 100644 --- a/src/ai_company/observability/events/memory.py +++ b/src/ai_company/observability/events/memory.py @@ -19,6 +19,7 @@ MEMORY_BACKEND_CREATED: Final[str] = "memory.backend.created" MEMORY_BACKEND_NOT_IMPLEMENTED: Final[str] = "memory.backend.not_implemented" MEMORY_BACKEND_UNKNOWN: Final[str] = "memory.backend.unknown" +MEMORY_BACKEND_CONFIG_INVALID: Final[str] = "memory.backend.config_invalid" MEMORY_BACKEND_NOT_CONNECTED: Final[str] = "memory.backend.not_connected" # ── Entry operations ────────────────────────────────────────────── diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index ab4360c357..6c8508cfe9 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -414,7 +414,7 @@ async def test_store_missing_id_raises( "results": [{"memory": "no id", "event": "ADD"}], } - with pytest.raises(MemoryStoreError, match="missing 'id'"): + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): await backend.store("test-agent-001", _make_store_request()) async def test_store_exception_wraps( @@ -777,7 +777,7 @@ async def test_publish_missing_id_raises( "results": [{"memory": "no id", "event": "ADD"}], } - with pytest.raises(MemoryStoreError, match="missing 'id'"): + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): await backend.publish("test-agent-001", _make_store_request()) @@ -950,3 +950,168 @@ async def test_retract_exception_wraps( with pytest.raises(MemoryStoreError, match="Failed to retract"): await backend.retract("test-agent-001", "shared-001") + + async def test_retract_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping.""" + mock_client.get.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.retract("test-agent-001", "shared-001") + + async def test_retract_delete_failure_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Exception during delete phase wraps in MemoryStoreError.""" + mock_client.get.return_value = { + "id": "shared-001", + "memory": "content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + } + mock_client.delete.side_effect = RuntimeError("delete failed") + + with pytest.raises(MemoryStoreError, match="Failed to retract"): + await backend.retract("test-agent-001", "shared-001") + + +@pytest.mark.unit +class TestAdditionalEdgeCases: + """Edge cases for improved coverage.""" + + async def test_store_blank_id_from_add_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store result with blank ID raises MemoryStoreError.""" + mock_client.add.return_value = { + "results": [{"id": "", "event": "ADD"}], + } + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + await backend.store("test-agent-001", _make_store_request()) + + async def test_store_whitespace_id_from_add_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store result with whitespace-only ID raises MemoryStoreError.""" + mock_client.add.return_value = { + "results": [{"id": " ", "event": "ADD"}], + } + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + await backend.store("test-agent-001", _make_store_request()) + + async def test_get_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping in get().""" + mock_client.get.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.get("test-agent-001", "mem-001") + + async def test_delete_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping in delete().""" + mock_client.get.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.delete("test-agent-001", "mem-001") + + async def test_count_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping in count().""" + mock_client.get_all.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.count("test-agent-001") + + async def test_publish_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping in publish().""" + mock_client.add.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.publish("test-agent-001", _make_store_request()) + + async def test_search_shared_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """MemoryError is re-raised without wrapping in search_shared().""" + mock_client.search.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.search_shared(MemoryQuery(text="test")) + + async def test_store_non_list_results_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store result with non-list 'results' raises MemoryStoreError.""" + mock_client.add.return_value = {"results": "not-a-list"} + with pytest.raises(MemoryStoreError, match="no results"): + await backend.store("test-agent-001", _make_store_request()) + + async def test_retrieve_invalid_entry_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Invalid entry in search results wraps as MemoryRetrievalError.""" + mock_client.search.return_value = { + "results": [ + {"id": "", "memory": "blank id", "metadata": {}}, + ], + } + with pytest.raises(MemoryRetrievalError, match="missing or blank"): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + + async def test_search_shared_no_publisher_uses_namespace( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Entries without publisher metadata use the shared namespace.""" + mock_client.search.return_value = _mem0_search_result( + [ + { + "id": "shared-1", + "memory": "orphan fact", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {"_synthorg_category": "semantic"}, + }, + ], + ) + + entries = await backend.search_shared(MemoryQuery(text="test")) + assert len(entries) == 1 + assert entries[0].agent_id == _SHARED_NAMESPACE + + async def test_count_empty_results( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Count returns 0 for empty results.""" + mock_client.get_all.return_value = {"results": []} + count = await backend.count("test-agent-001") + assert count == 0 diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index 34c6d461f8..c049bba835 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -1,6 +1,7 @@ """Tests for Mem0 mapping functions.""" from datetime import UTC, datetime, timedelta +from typing import Any import pytest @@ -9,14 +10,17 @@ _PREFIX, apply_post_filters, build_mem0_metadata, + extract_category, + extract_publisher, mem0_result_to_entry, normalize_relevance_score, parse_mem0_datetime, parse_mem0_metadata, query_to_mem0_getall_args, query_to_mem0_search_args, + validate_add_result, ) -from ai_company.memory.errors import MemoryRetrievalError +from ai_company.memory.errors import MemoryRetrievalError, MemoryStoreError from ai_company.memory.models import ( MemoryEntry, MemoryMetadata, @@ -385,3 +389,103 @@ def test_combined_filters(self) -> None: result = apply_post_filters(entries, query) assert len(result) == 1 assert result[0].id == "m1" + + +@pytest.mark.unit +class TestValidateAddResult: + def test_valid_result(self) -> None: + result = {"results": [{"id": "mem-001", "event": "ADD"}]} + memory_id = validate_add_result(result, context="test") + assert memory_id == "mem-001" + + def test_empty_results_raises(self) -> None: + result: dict[str, Any] = {"results": []} + with pytest.raises(MemoryStoreError, match="no results"): + validate_add_result(result, context="test") + + def test_missing_results_key_raises(self) -> None: + result = {"data": "something"} + with pytest.raises(MemoryStoreError, match="no results"): + validate_add_result(result, context="test") + + def test_non_list_results_raises(self) -> None: + result = {"results": "not-a-list"} + with pytest.raises(MemoryStoreError, match="no results"): + validate_add_result(result, context="test") + + def test_missing_id_raises(self) -> None: + result = {"results": [{"memory": "no id"}]} + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + validate_add_result(result, context="test") + + def test_none_id_raises(self) -> None: + result = {"results": [{"id": None, "event": "ADD"}]} + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + validate_add_result(result, context="test") + + def test_blank_id_raises(self) -> None: + result = {"results": [{"id": "", "event": "ADD"}]} + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + validate_add_result(result, context="test") + + def test_whitespace_only_id_raises(self) -> None: + result = {"results": [{"id": " ", "event": "ADD"}]} + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + validate_add_result(result, context="test") + + def test_numeric_id_coerced_to_string(self) -> None: + result = {"results": [{"id": 42, "event": "ADD"}]} + memory_id = validate_add_result(result, context="test") + assert memory_id == "42" + + +@pytest.mark.unit +class TestExtractCategory: + def test_valid_category(self) -> None: + raw = {"metadata": {f"{_PREFIX}category": "episodic"}} + assert extract_category(raw) == MemoryCategory.EPISODIC + + def test_missing_metadata(self) -> None: + raw = {"id": "m1", "memory": "content"} + assert extract_category(raw) == MemoryCategory.WORKING + + def test_empty_metadata(self) -> None: + raw: dict[str, Any] = {"metadata": {}} + assert extract_category(raw) == MemoryCategory.WORKING + + def test_none_metadata(self) -> None: + raw: dict[str, Any] = {"metadata": None} + assert extract_category(raw) == MemoryCategory.WORKING + + def test_invalid_category_defaults(self) -> None: + raw = {"metadata": {f"{_PREFIX}category": "nonexistent"}} + assert extract_category(raw) == MemoryCategory.WORKING + + def test_missing_category_key(self) -> None: + raw = {"metadata": {f"{_PREFIX}confidence": 0.9}} + assert extract_category(raw) == MemoryCategory.WORKING + + +@pytest.mark.unit +class TestExtractPublisher: + def test_valid_publisher(self) -> None: + from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY + + raw = {"metadata": {_PUBLISHER_KEY: "test-agent-001"}} + assert extract_publisher(raw) == "test-agent-001" + + def test_missing_metadata(self) -> None: + raw = {"id": "m1"} + assert extract_publisher(raw) is None + + def test_empty_metadata(self) -> None: + raw: dict[str, Any] = {"metadata": {}} + assert extract_publisher(raw) is None + + def test_none_metadata(self) -> None: + raw = {"metadata": None} + assert extract_publisher(raw) is None + + def test_no_publisher_key(self) -> None: + raw = {"metadata": {"_synthorg_category": "semantic"}} + assert extract_publisher(raw) is None From a7ee906965a4352539466c7fdcd30bec1ef1d4b6 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 07:19:04 +0100 Subject: [PATCH 05/17] fix: add non-dict metadata guards and created_at fallback logging --- src/ai_company/memory/backends/mem0/mappers.py | 12 +++++++++--- tests/integration/memory/test_mem0_backend.py | 1 + tests/unit/memory/backends/mem0/test_mappers.py | 16 ++++++++++++++++ 3 files changed, 26 insertions(+), 3 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 207ae52742..c9b00b17ea 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -111,7 +111,7 @@ def parse_mem0_metadata( Returns: Tuple of (category, metadata, expires_at). """ - if not raw_metadata: + if not raw_metadata or not isinstance(raw_metadata, dict): return ( MemoryCategory.WORKING, MemoryMetadata(), @@ -204,6 +204,12 @@ def mem0_result_to_entry( created_at = parse_mem0_datetime(raw.get("created_at")) if created_at is None: + logger.debug( + MEMORY_MODEL_INVALID, + field="created_at", + memory_id=str(raw.get("id", "?")), + reason="missing or unparseable created_at, defaulting to now()", + ) created_at = datetime.now(UTC) updated_at = parse_mem0_datetime(raw.get("updated_at")) @@ -353,7 +359,7 @@ def extract_category(raw: dict[str, Any]) -> MemoryCategory: or unrecognised. """ metadata = raw.get("metadata", {}) - if not metadata: + if not metadata or not isinstance(metadata, dict): return MemoryCategory.WORKING cat_str = metadata.get(f"{_PREFIX}category") if cat_str: @@ -374,7 +380,7 @@ def extract_category(raw: dict[str, Any]) -> MemoryCategory: def extract_publisher(raw: dict[str, Any]) -> str | None: """Extract the publisher agent ID from a shared memory dict.""" metadata = raw.get("metadata", {}) - if not metadata: + if not metadata or not isinstance(metadata, dict): return None value: str | None = metadata.get(_PUBLISHER_KEY) return value diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py index 6372d7e49f..d95affa1d0 100644 --- a/tests/integration/memory/test_mem0_backend.py +++ b/tests/integration/memory/test_mem0_backend.py @@ -158,6 +158,7 @@ async def test_pipeline_prepare_messages( # Content should include both memories (they pass min_relevance) combined = " ".join(m.content for m in messages if m.content) assert "concise responses" in combined + assert "code review" in combined async def test_shared_knowledge_flow( self, diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index c049bba835..00c82864c4 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -457,6 +457,14 @@ def test_none_metadata(self) -> None: raw: dict[str, Any] = {"metadata": None} assert extract_category(raw) == MemoryCategory.WORKING + def test_list_metadata_defaults(self) -> None: + raw: dict[str, Any] = {"metadata": ["not", "a", "dict"]} + assert extract_category(raw) == MemoryCategory.WORKING + + def test_string_metadata_defaults(self) -> None: + raw: dict[str, Any] = {"metadata": "oops"} + assert extract_category(raw) == MemoryCategory.WORKING + def test_invalid_category_defaults(self) -> None: raw = {"metadata": {f"{_PREFIX}category": "nonexistent"}} assert extract_category(raw) == MemoryCategory.WORKING @@ -489,3 +497,11 @@ def test_none_metadata(self) -> None: def test_no_publisher_key(self) -> None: raw = {"metadata": {"_synthorg_category": "semantic"}} assert extract_publisher(raw) is None + + def test_list_metadata_returns_none(self) -> None: + raw: dict[str, Any] = {"metadata": ["not", "a", "dict"]} + assert extract_publisher(raw) is None + + def test_string_metadata_returns_none(self) -> None: + raw: dict[str, Any] = {"metadata": "oops"} + assert extract_publisher(raw) is None From 5aedbc44ab8de6815730da044dc155ba43c93c9d Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 07:25:10 +0100 Subject: [PATCH 06/17] fix: remove double-logging in retract() ownership check path --- src/ai_company/memory/backends/mem0/adapter.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 008ee23589..5850069cd0 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -758,14 +758,10 @@ async def retract( raise MemoryStoreError(msg) # noqa: TRY301 await asyncio.to_thread(self._client.delete, str(memory_id)) - except MemoryStoreError as exc: - logger.warning( - MEMORY_SHARED_RETRACT_FAILED, - agent_id=agent_id, - memory_id=memory_id, - error=str(exc), - error_type="MemoryStoreError", - ) + except MemoryStoreError: + # Ownership-check MemoryStoreErrors are already logged + # with context (reason, publisher) above — re-raise + # without duplicate logging. raise except MemoryError, RecursionError: raise From ab58ccf834f58ca3419fca6597c05fcc0064f9d1 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 08:04:01 +0100 Subject: [PATCH 07/17] fix: address round-3 PR review findings across adapter, mappers, config, and tests - mappers.py: add str() conversion before NotBlankStr for tags, clamp confidence to [0.0, 1.0], log missing/non-dict metadata and unexpected tag types, guard blank publisher strings, document None relevance score behavior in apply_post_filters, clean up historical comment - config.py: add structured logging to _reject_traversal validator, validate unsupported storage overrides in build_config_from_company_config - factory.py: wrap ValueError from config construction as MemoryConfigError - pyproject.toml: move mem0ai to optional-dependencies, add TC to test per-file-ignores - __init__.py: guard Mem0 imports with contextlib.suppress(ImportError) - adapter.py: log during disconnect reset failure instead of bare pass - dependency-review.yml: version-pin allowed PURLs - docs: add embedder config to memory.md YAML example, update roadmap Current Status to mention Mem0 adapter - tests: split test_adapter.py (1118 lines) into conftest.py + test_adapter.py + test_adapter_crud.py + test_adapter_shared.py; add protocol isinstance checks, connect MemoryError/RecursionError propagation, agent_id validation, ownership mismatch, and shared namespace guard tests - Clean up 78 now-unused noqa: TC* directives across test files --- .github/workflows/dependency-review.yml | 8 +- docs/design/memory.md | 7 + docs/roadmap/index.md | 2 +- pyproject.toml | 5 +- src/ai_company/memory/__init__.py | 11 +- .../memory/backends/mem0/adapter.py | 215 +++- src/ai_company/memory/backends/mem0/config.py | 43 + .../memory/backends/mem0/mappers.py | 28 +- src/ai_company/memory/factory.py | 18 +- src/ai_company/observability/events/memory.py | 1 + .../communication/test_meeting_integration.py | 4 +- tests/integration/tools/conftest.py | 2 +- .../tools/test_sandbox_integration.py | 2 +- tests/unit/api/auth/test_controller.py | 4 +- tests/unit/api/conftest.py | 18 +- tests/unit/api/controllers/test_agents.py | 2 +- tests/unit/api/controllers/test_analytics.py | 2 +- tests/unit/api/controllers/test_approvals.py | 4 +- tests/unit/api/controllers/test_artifacts.py | 2 +- tests/unit/api/controllers/test_autonomy.py | 2 +- tests/unit/api/controllers/test_budget.py | 4 +- tests/unit/api/controllers/test_company.py | 2 +- .../unit/api/controllers/test_departments.py | 2 +- tests/unit/api/controllers/test_meetings.py | 2 +- tests/unit/api/controllers/test_messages.py | 2 +- tests/unit/api/controllers/test_projects.py | 2 +- tests/unit/api/controllers/test_providers.py | 2 +- tests/unit/api/controllers/test_tasks.py | 2 +- tests/unit/api/test_app.py | 2 +- tests/unit/api/test_guards.py | 2 +- tests/unit/api/test_health.py | 2 +- tests/unit/api/test_middleware.py | 2 +- tests/unit/budget/test_category_analytics.py | 2 +- .../test_authority_strategy.py | 2 +- .../test_debate_strategy.py | 2 +- .../conflict_resolution/test_helpers.py | 2 +- .../test_hybrid_strategy.py | 2 +- .../conflict_resolution/test_service.py | 2 +- tests/unit/communication/meeting/conftest.py | 2 +- .../meeting/test_orchestrator.py | 2 +- .../meeting/test_position_papers.py | 2 +- .../communication/meeting/test_protocol.py | 2 +- .../communication/meeting/test_round_robin.py | 2 +- .../meeting/test_structured_phases.py | 2 +- tests/unit/engine/task_engine_helpers.py | 2 +- tests/unit/engine/test_agent_engine.py | 2 +- tests/unit/engine/test_agent_engine_errors.py | 4 +- .../engine/test_agent_engine_lifecycle.py | 4 +- tests/unit/engine/test_context.py | 4 +- tests/unit/engine/test_loop_protocol.py | 2 +- tests/unit/engine/test_metrics.py | 2 +- tests/unit/engine/test_plan_execute_loop.py | 2 +- tests/unit/engine/test_react_loop.py | 2 +- tests/unit/engine/test_routing_models.py | 2 +- .../unit/engine/test_task_engine_mutations.py | 2 +- tests/unit/hr/test_full_snapshot_strategy.py | 2 +- tests/unit/hr/test_hiring_service.py | 4 +- tests/unit/hr/test_offboarding_service.py | 8 +- tests/unit/hr/test_onboarding_service.py | 4 +- tests/unit/hr/test_registry.py | 2 +- tests/unit/memory/backends/mem0/conftest.py | 105 ++ .../unit/memory/backends/mem0/test_adapter.py | 925 ++---------------- .../memory/backends/mem0/test_adapter_crud.py | 515 ++++++++++ .../backends/mem0/test_adapter_shared.py | 327 +++++++ tests/unit/providers/conftest.py | 2 +- tests/unit/providers/test_protocol.py | 2 +- tests/unit/tools/git/conftest.py | 2 +- .../tools/git/test_git_sandbox_integration.py | 2 +- tests/unit/tools/sandbox/conftest.py | 2 +- tests/unit/tools/sandbox/test_protocol.py | 6 +- uv.lock | 9 +- 71 files changed, 1376 insertions(+), 999 deletions(-) create mode 100644 tests/unit/memory/backends/mem0/conftest.py create mode 100644 tests/unit/memory/backends/mem0/test_adapter_crud.py create mode 100644 tests/unit/memory/backends/mem0/test_adapter_shared.py diff --git a/.github/workflows/dependency-review.yml b/.github/workflows/dependency-review.yml index 3859ba196a..68bbc3d61c 100644 --- a/.github/workflows/dependency-review.yml +++ b/.github/workflows/dependency-review.yml @@ -45,8 +45,8 @@ jobs: # Verified manually: mem0ai (Apache-2.0), numpy (BSD-3-Clause), # qdrant-client (Apache-2.0), posthog (MIT). allow-dependencies-licenses: >- - pkg:pypi/mem0ai, - pkg:pypi/numpy, - pkg:pypi/qdrant-client, - pkg:pypi/posthog + pkg:pypi/mem0ai@1.0.5, + pkg:pypi/numpy@2.4.3, + pkg:pypi/qdrant-client@1.17.0, + pkg:pypi/posthog@7.9.12 comment-summary-in-pr: always diff --git a/docs/design/memory.md b/docs/design/memory.md index 25db5ca432..1d494fbe7a 100644 --- a/docs/design/memory.md +++ b/docs/design/memory.md @@ -283,6 +283,13 @@ memory: max_memories_per_agent: 10000 consolidation_interval: "daily" shared_knowledge_base: true + +# Embedder config is passed programmatically via the factory: +# create_memory_backend(config, embedder=Mem0EmbedderConfig( +# provider="", +# model="", +# dims=1536, +# )) ``` Configuration is modeled by `CompanyMemoryConfig` (top-level), `MemoryStorageConfig` diff --git a/docs/roadmap/index.md b/docs/roadmap/index.md index 033f1bf047..ef10fae0ee 100644 --- a/docs/roadmap/index.md +++ b/docs/roadmap/index.md @@ -8,7 +8,7 @@ The SynthOrg core framework is complete. The following subsystems are built and - Budget and cost management (tracking, enforcement, CFO optimization, quotas) - Agent engine (execution loops, parallel execution, task decomposition, routing, assignment, recovery, shutdown) - Communication layer (message bus, delegation, loop prevention, conflict resolution, meeting protocol) -- Memory system (pluggable backend protocol, retrieval pipeline, shared org memory, consolidation) +- Memory system (pluggable backend protocol, Mem0 adapter, retrieval pipeline, shared org memory, consolidation) - Security and approval system (rule engine, output scanning, progressive trust, autonomy levels, timeout policies) - Tool system (file system, git, code runner, MCP bridge, sandboxing, permissions) - HR engine (hiring, firing, onboarding, offboarding, registry, performance tracking, promotions) diff --git a/pyproject.toml b/pyproject.toml index a54f4cea27..a93f1421d9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,13 +21,15 @@ dependencies = [ "litellm==1.82.1", "litestar[standard,structlog,pydantic,brotli,prometheus]==2.21.1", "mcp==1.26.0", - "mem0ai==1.0.5", "pydantic==2.12.5", "pyjwt[crypto]==2.11.0", "pyyaml==6.0.3", "structlog==25.5.0", ] +[project.optional-dependencies] +mem0 = ["mem0ai==1.0.5"] + [build-system] requires = ["hatchling==1.29.0"] build-backend = "hatchling.build" @@ -135,6 +137,7 @@ convention = "google" "PLR2004", # magic values in tests "SLF001", # private member access in tests "PLC0415", # local imports in test functions + "TC", # type-checking imports (fixture type hints) ] "__init__.py" = ["F401"] "scripts/**/*.py" = [ diff --git a/src/ai_company/memory/__init__.py b/src/ai_company/memory/__init__.py index b2990a46db..172f4c58be 100644 --- a/src/ai_company/memory/__init__.py +++ b/src/ai_company/memory/__init__.py @@ -9,10 +9,13 @@ directly. """ -from ai_company.memory.backends.mem0 import ( - Mem0EmbedderConfig, - Mem0MemoryBackend, -) +import contextlib + +with contextlib.suppress(ImportError): # mem0ai is optional + from ai_company.memory.backends.mem0 import ( + Mem0EmbedderConfig, + Mem0MemoryBackend, + ) from ai_company.memory.capabilities import MemoryCapabilities from ai_company.memory.config import ( CompanyMemoryConfig, diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 5850069cd0..e59a15abc6 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -1,14 +1,18 @@ """Mem0 memory backend adapter. Implements ``MemoryBackend``, ``MemoryCapabilities``, and -``SharedKnowledgeStore`` protocols using Mem0 (embedded Qdrant + SQLite) -as the storage layer. +``SharedKnowledgeStore`` protocols using Mem0 as the storage layer +(default: embedded Qdrant + SQLite). All Mem0 calls are synchronous — they run in ``asyncio.to_thread()`` to avoid blocking the event loop. + +All methods re-raise ``builtins.MemoryError`` and ``RecursionError`` +immediately without wrapping, to avoid masking system-level failures. """ import asyncio +import builtins from typing import TYPE_CHECKING, Any from ai_company.core.enums import MemoryCategory @@ -43,6 +47,7 @@ ) from ai_company.observability.events.memory import ( + MEMORY_BACKEND_AGENT_ID_REJECTED, MEMORY_BACKEND_CONNECTED, MEMORY_BACKEND_CONNECTING, MEMORY_BACKEND_CONNECTION_FAILED, @@ -74,6 +79,40 @@ _SHARED_NAMESPACE: str = "__synthorg_shared__" +def _validate_mem0_result( + raw_result: Any, + *, + context: str, +) -> list[dict[str, Any]]: + """Validate and extract the results list from a Mem0 response. + + Args: + raw_result: Raw return value from a Mem0 SDK call. + context: Human-readable context for error messages. + + Returns: + The ``"results"`` list from the response. + + Raises: + MemoryRetrievalError: If the response is not a dict or + ``"results"`` is not a list. + """ + if not isinstance(raw_result, dict): + msg = ( + f"Unexpected Mem0 response type for {context}: " + f"{type(raw_result).__name__}, expected dict" + ) + raise MemoryRetrievalError(msg) + raw_list = raw_result.get("results", []) + if not isinstance(raw_list, list): + msg = ( + f"Unexpected Mem0 results type for {context}: " + f"{type(raw_list).__name__}, expected list" + ) + raise MemoryRetrievalError(msg) + return raw_list + + class Mem0MemoryBackend: """Mem0-backed agent memory backend. @@ -122,7 +161,7 @@ async def connect(self) -> None: try: config_dict = build_mem0_config_dict(self._mem0_config) client = await asyncio.to_thread(Memory.from_config, config_dict) - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -140,9 +179,19 @@ async def connect(self) -> None: async def disconnect(self) -> None: """Close the Mem0 connection. - Safe to call even if not connected. + Attempts to close the underlying client resources before + releasing the reference. Safe to call even if not connected. """ logger.info(MEMORY_BACKEND_DISCONNECTING, backend="mem0") + if self._client is not None: + try: + await asyncio.to_thread(self._client.reset) + except Exception: + logger.debug( + MEMORY_BACKEND_DISCONNECTING, + backend="mem0", + note="reset failed during disconnect, ignoring", + ) self._client = None self._connected = False logger.info(MEMORY_BACKEND_DISCONNECTED, backend="mem0") @@ -170,13 +219,15 @@ async def health_check(self) -> bool: user_id=_SHARED_NAMESPACE, limit=1, ) - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise - except Exception: + except Exception as exc: logger.warning( MEMORY_BACKEND_HEALTH_CHECK, backend="mem0", healthy=False, + error=str(exc), + error_type=type(exc).__name__, ) return False logger.debug( @@ -228,7 +279,7 @@ def max_memories_per_agent(self) -> int | None: """Maximum memories per agent from configuration.""" return self._max_memories_per_agent - # ── Connection guard ────────────────────────────────────────── + # ── Guards ──────────────────────────────────────────────────── def _require_connected(self) -> None: """Raise ``MemoryConnectionError`` if not connected.""" @@ -240,6 +291,25 @@ def _require_connected(self) -> None: msg = "Not connected — call connect() first" raise MemoryConnectionError(msg) + def _validate_agent_id(self, agent_id: NotBlankStr) -> None: + """Reject the reserved shared namespace as an agent ID. + + Raises: + MemoryStoreError: If ``agent_id`` collides with the + reserved ``_SHARED_NAMESPACE``. + """ + if str(agent_id) == _SHARED_NAMESPACE: + logger.warning( + MEMORY_BACKEND_AGENT_ID_REJECTED, + agent_id=agent_id, + reason="reserved shared namespace", + ) + msg = ( + f"agent_id must not be the reserved shared namespace: " + f"{_SHARED_NAMESPACE!r}" + ) + raise MemoryStoreError(msg) + # ── CRUD Operations ─────────────────────────────────────────── async def store( @@ -261,6 +331,7 @@ async def store( MemoryStoreError: If the store operation fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: kwargs = { "messages": [ @@ -280,7 +351,7 @@ async def store( error_type="MemoryStoreError", ) raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -308,7 +379,8 @@ async def retrieve( """Retrieve memories for an agent, ordered by relevance. Uses ``search()`` when ``query.text`` is set, otherwise falls - back to ``get_all()`` for unfiltered retrieval. + back to ``get_all()`` for non-semantic retrieval (post-filters + still apply). Args: agent_id: Owning agent identifier. @@ -322,6 +394,7 @@ async def retrieve( MemoryRetrievalError: If the retrieval fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: if query.text is not None: kwargs = query_to_mem0_search_args(str(agent_id), query) @@ -329,7 +402,7 @@ async def retrieve( else: kwargs = query_to_mem0_getall_args(str(agent_id), query) raw_result = await asyncio.to_thread(self._client.get_all, **kwargs) - raw_list = raw_result.get("results", []) + raw_list = _validate_mem0_result(raw_result, context="retrieve") entries = tuple( mem0_result_to_entry(item, str(agent_id)) for item in raw_list ) @@ -342,7 +415,7 @@ async def retrieve( error_type="MemoryRetrievalError", ) raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -368,18 +441,22 @@ async def get( ) -> MemoryEntry | None: """Get a specific memory entry by ID. + Verifies ownership: if the retrieved memory belongs to a + different agent the method returns ``None``. + Args: agent_id: Owning agent identifier. memory_id: Memory identifier. Returns: - The memory entry, or ``None`` if not found. + The memory entry, or ``None`` if not found or not owned. Raises: MemoryConnectionError: If the backend is not connected. MemoryRetrievalError: If the backend query fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: raw = await asyncio.to_thread(self._client.get, str(memory_id)) if raw is None: @@ -390,6 +467,17 @@ async def get( found=False, ) return None + owner = raw.get("user_id") + if owner is not None and str(owner) != str(agent_id): + logger.debug( + MEMORY_ENTRY_FETCHED, + agent_id=agent_id, + memory_id=memory_id, + found=False, + reason="ownership mismatch", + actual_owner=str(owner), + ) + return None entry = mem0_result_to_entry(raw, str(agent_id)) except MemoryRetrievalError as exc: logger.warning( @@ -400,7 +488,7 @@ async def get( error_type="MemoryRetrievalError", ) raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -428,6 +516,9 @@ async def delete( ) -> bool: """Delete a specific memory entry. + Verifies ownership before deletion. Shared-namespace entries + must be removed through ``retract()`` instead. + Args: agent_id: Owning agent identifier. memory_id: Memory identifier. @@ -437,9 +528,11 @@ async def delete( Raises: MemoryConnectionError: If the backend is not connected. - MemoryStoreError: If the delete operation fails. + MemoryStoreError: If the delete operation fails or + ownership verification fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: # Check existence first — Mem0 delete doesn't indicate # whether the entry existed. @@ -452,17 +545,38 @@ async def delete( found=False, ) return False + # Block deletion of shared-namespace entries — use retract(). + owner = existing.get("user_id") + if owner is not None and str(owner) == _SHARED_NAMESPACE: + msg = ( + f"Memory {memory_id} belongs to the shared namespace — " + f"use retract() to remove shared entries" + ) + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="shared namespace entry", + ) + raise MemoryStoreError(msg) # noqa: TRY301 + # Verify ownership — reject cross-agent deletion. + if owner is not None and str(owner) != str(agent_id): + msg = ( + f"Agent {agent_id} cannot delete memory " + f"{memory_id} owned by {owner}" + ) + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="ownership mismatch", + actual_owner=str(owner), + ) + raise MemoryStoreError(msg) # noqa: TRY301 await asyncio.to_thread(self._client.delete, str(memory_id)) - except MemoryStoreError as exc: - logger.warning( - MEMORY_ENTRY_DELETE_FAILED, - agent_id=agent_id, - memory_id=memory_id, - error=str(exc), - error_type="MemoryStoreError", - ) + except MemoryStoreError: raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -491,14 +605,15 @@ async def count( ) -> int: """Count memory entries for an agent. - Uses ``get_all()`` internally — O(n) in the agent's memory - count. Acceptable because ``count()`` is not on the hot path. + Uses ``get_all()`` internally — retrieves all of the agent's + memories, so cost scales linearly with the agent's memory count. + Acceptable because ``count()`` is not on the hot path. - .. note:: - Results are capped at ``max_memories_per_agent``. If an - agent has more memories than this limit the count will be - an underestimate. This is consistent with the adapter's - store/retrieve semantics which also respect the cap. + Note: + Results are capped at ``max_memories_per_agent``. If an + agent has more memories than this limit the count will be + an underestimate. This is consistent with the adapter's + store/retrieve semantics which also respect the cap. Args: agent_id: Owning agent identifier. @@ -513,13 +628,14 @@ async def count( MemoryRetrievalError: If the count query fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: raw_result = await asyncio.to_thread( self._client.get_all, user_id=str(agent_id), limit=self._max_memories_per_agent, ) - raw_list = raw_result.get("results", []) + raw_list = _validate_mem0_result(raw_result, context="count") if category is None: total = len(raw_list) else: @@ -534,7 +650,7 @@ async def count( error_type="MemoryRetrievalError", ) raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -546,12 +662,24 @@ async def count( msg = f"Failed to count memories: {exc}" raise MemoryRetrievalError(msg) from exc else: - logger.info( - MEMORY_ENTRY_COUNTED, - agent_id=agent_id, - count=total, - category=category.value if category else None, - ) + truncated = total == self._max_memories_per_agent + if truncated: + logger.warning( + MEMORY_ENTRY_COUNTED, + agent_id=agent_id, + count=total, + category=category.value if category else None, + truncated=True, + reason="count equals max_memories_per_agent, " + "actual count may be higher", + ) + else: + logger.info( + MEMORY_ENTRY_COUNTED, + agent_id=agent_id, + count=total, + category=category.value if category else None, + ) return total # ── SharedKnowledgeStore ────────────────────────────────────── @@ -601,7 +729,7 @@ async def publish( error_type="MemoryStoreError", ) raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -654,7 +782,10 @@ async def search_shared( user_id=_SHARED_NAMESPACE, limit=query.limit, ) - raw_list = raw_result.get("results", []) + raw_list = _validate_mem0_result( + raw_result, + context="search_shared", + ) raw_entries = tuple( mem0_result_to_entry( @@ -676,7 +807,7 @@ async def search_shared( exclude_agent=exclude_agent, ) raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( @@ -763,7 +894,7 @@ async def retract( # with context (reason, publisher) above — re-raise # without duplicate logging. raise - except MemoryError, RecursionError: + except builtins.MemoryError, RecursionError: raise except Exception as exc: logger.warning( diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 965898f6ce..7b78fe6369 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -12,6 +12,12 @@ from ai_company.core.types import NotBlankStr # noqa: TC001 from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 +from ai_company.observability import get_logger +from ai_company.observability.events.memory import ( + MEMORY_BACKEND_CONFIG_INVALID, +) + +logger = get_logger(__name__) class Mem0EmbedderConfig(BaseModel): @@ -74,6 +80,13 @@ def _reject_traversal(self) -> Self: ) if ".." in parts: msg = "data_dir must not contain parent-directory traversal (..)" + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + field="data_dir", + value=self.data_dir, + reason=msg, + ) raise ValueError(msg) return self @@ -122,7 +135,37 @@ def build_config_from_company_config( Returns: Mem0-specific backend configuration. + + Raises: + ValueError: If the storage config specifies a vector or + history store that the Mem0 backend does not support. """ + if config.storage.vector_store not in ("qdrant", "qdrant-external"): + msg = ( + f"Mem0 backend only supports qdrant vector stores, " + f"got {config.storage.vector_store!r}" + ) + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + field="vector_store", + value=config.storage.vector_store, + reason=msg, + ) + raise ValueError(msg) + if config.storage.history_store != "sqlite": + msg = ( + f"Mem0 backend only supports sqlite history store, " + f"got {config.storage.history_store!r}" + ) + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + field="history_store", + value=config.storage.history_store, + reason=msg, + ) + raise ValueError(msg) return Mem0BackendConfig( data_dir=config.storage.data_dir, embedder=embedder, diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index c9b00b17ea..b25158dc7c 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -112,6 +112,12 @@ def parse_mem0_metadata( Tuple of (category, metadata, expires_at). """ if not raw_metadata or not isinstance(raw_metadata, dict): + logger.debug( + MEMORY_MODEL_INVALID, + field="metadata", + raw_value=type(raw_metadata).__name__ if raw_metadata else None, + reason="missing or non-dict metadata, using defaults", + ) return ( MemoryCategory.WORKING, MemoryMetadata(), @@ -144,13 +150,20 @@ def parse_mem0_metadata( reason="non-numeric confidence, defaulting to 1.0", ) confidence = 1.0 + confidence = max(0.0, min(1.0, confidence)) source = raw_metadata.get(f"{_PREFIX}source") raw_tags = raw_metadata.get(f"{_PREFIX}tags", ()) if isinstance(raw_tags, str): raw_tags = [raw_tags] elif not isinstance(raw_tags, (list, tuple)): + logger.debug( + MEMORY_MODEL_INVALID, + field="tags", + raw_value=type(raw_tags).__name__, + reason="unexpected tags type, ignoring", + ) raw_tags = () - tags = tuple(NotBlankStr(t) for t in raw_tags if t and str(t).strip()) + tags = tuple(NotBlankStr(str(t)) for t in raw_tags if t and str(t).strip()) expires_at = parse_mem0_datetime( raw_metadata.get(f"{_PREFIX}expires_at"), @@ -290,6 +303,9 @@ def apply_post_filters( """Apply post-retrieval filters that Mem0 cannot handle natively. Filters by category, tags, time range, and minimum relevance. + Entries with ``relevance_score=None`` (e.g. from ``get_all``) + are never excluded by ``min_relevance`` — the filter only + applies when a score is present. Args: entries: Raw entries from Mem0. @@ -318,7 +334,7 @@ def apply_post_filters( return tuple(result) -# ── Adapter helpers (moved here to keep adapter.py under 800 lines) ── +# ── Adapter helpers ────────────────────────────────────────────────── def validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: @@ -378,9 +394,15 @@ def extract_category(raw: dict[str, Any]) -> MemoryCategory: def extract_publisher(raw: dict[str, Any]) -> str | None: - """Extract the publisher agent ID from a shared memory dict.""" + """Extract the publisher agent ID from a shared memory dict. + + Returns ``None`` if the publisher key is missing, non-dict + metadata, or the value is blank after stripping. + """ metadata = raw.get("metadata", {}) if not metadata or not isinstance(metadata, dict): return None value: str | None = metadata.get(_PUBLISHER_KEY) + if value is not None and not str(value).strip(): + return None return value diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 53b06f37a4..53bc44d61c 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -77,10 +77,20 @@ def create_memory_backend( ) raise MemoryConfigError(msg) - mem0_config = build_config_from_company_config( - config, - embedder=embedder, - ) + try: + mem0_config = build_config_from_company_config( + config, + embedder=embedder, + ) + except ValueError as exc: + msg = f"Invalid Mem0 configuration: {exc}" + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="config_build_failed", + error=msg, + ) + raise MemoryConfigError(msg) from exc backend = Mem0MemoryBackend( mem0_config=mem0_config, max_memories_per_agent=config.options.max_memories_per_agent, diff --git a/src/ai_company/observability/events/memory.py b/src/ai_company/observability/events/memory.py index 6d22ddc90b..b4e259031a 100644 --- a/src/ai_company/observability/events/memory.py +++ b/src/ai_company/observability/events/memory.py @@ -21,6 +21,7 @@ MEMORY_BACKEND_UNKNOWN: Final[str] = "memory.backend.unknown" MEMORY_BACKEND_CONFIG_INVALID: Final[str] = "memory.backend.config_invalid" MEMORY_BACKEND_NOT_CONNECTED: Final[str] = "memory.backend.not_connected" +MEMORY_BACKEND_AGENT_ID_REJECTED: Final[str] = "memory.backend.agent_id_rejected" # ── Entry operations ────────────────────────────────────────────── diff --git a/tests/integration/communication/test_meeting_integration.py b/tests/integration/communication/test_meeting_integration.py index bd0ceb8147..57b4c3b356 100644 --- a/tests/integration/communication/test_meeting_integration.py +++ b/tests/integration/communication/test_meeting_integration.py @@ -26,8 +26,8 @@ PositionPapersProtocol, ) from ai_company.communication.meeting.protocol import ( - AgentCaller, # noqa: TC001 - MeetingProtocol, # noqa: TC001 + AgentCaller, + MeetingProtocol, ) from ai_company.communication.meeting.round_robin import RoundRobinProtocol from ai_company.communication.meeting.structured_phases import ( diff --git a/tests/integration/tools/conftest.py b/tests/integration/tools/conftest.py index 7d33626b60..aee22e47d5 100644 --- a/tests/integration/tools/conftest.py +++ b/tests/integration/tools/conftest.py @@ -2,7 +2,7 @@ import os import subprocess -from pathlib import Path # noqa: TC003 — pytest evaluates annotations +from pathlib import Path import pytest diff --git a/tests/integration/tools/test_sandbox_integration.py b/tests/integration/tools/test_sandbox_integration.py index fdc411a070..d012713052 100644 --- a/tests/integration/tools/test_sandbox_integration.py +++ b/tests/integration/tools/test_sandbox_integration.py @@ -1,7 +1,7 @@ """Integration tests for subprocess sandbox with real git.""" import os -from pathlib import Path # noqa: TC003 — used at runtime +from pathlib import Path import pytest diff --git a/tests/unit/api/auth/test_controller.py b/tests/unit/api/auth/test_controller.py index abfc63d0d3..11c6f3986f 100644 --- a/tests/unit/api/auth/test_controller.py +++ b/tests/unit/api/auth/test_controller.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from ai_company.api.guards import HumanRole from tests.unit.api.conftest import make_auth_headers @@ -44,7 +44,7 @@ def test_setup_409_when_users_exist( from datetime import UTC, datetime from ai_company.api.auth.models import User - from ai_company.api.auth.service import AuthService # noqa: TC001 + from ai_company.api.auth.service import AuthService from ai_company.api.guards import HumanRole app_state = bare_client.app.state["app_state"] diff --git a/tests/unit/api/conftest.py b/tests/unit/api/conftest.py index c0f0a014dc..feed2d053a 100644 --- a/tests/unit/api/conftest.py +++ b/tests/unit/api/conftest.py @@ -2,7 +2,7 @@ import asyncio import uuid -from collections.abc import Generator # noqa: TC003 +from collections.abc import Generator from datetime import UTC, datetime, timedelta from typing import Any @@ -15,10 +15,10 @@ from ai_company.api.auth.models import ApiKey, User from ai_company.api.auth.service import AuthService from ai_company.api.guards import HumanRole -from ai_company.budget.cost_record import CostRecord # noqa: TC001 +from ai_company.budget.cost_record import CostRecord from ai_company.budget.tracker import CostTracker -from ai_company.communication.channel import Channel # noqa: TC001 -from ai_company.communication.message import Message # noqa: TC001 +from ai_company.communication.channel import Channel +from ai_company.communication.message import Message from ai_company.config.schema import RootConfig from ai_company.core.approval import ApprovalItem from ai_company.core.enums import ( @@ -28,15 +28,15 @@ ) from ai_company.core.task import Task from ai_company.engine.task_engine import TaskEngine -from ai_company.hr.enums import LifecycleEventType # noqa: TC001 -from ai_company.hr.models import AgentLifecycleEvent # noqa: TC001 -from ai_company.hr.performance.models import ( # noqa: TC001 +from ai_company.hr.enums import LifecycleEventType +from ai_company.hr.models import AgentLifecycleEvent +from ai_company.hr.performance.models import ( CollaborationMetricRecord, TaskMetricRecord, ) from ai_company.persistence.errors import DuplicateRecordError, QueryError -from ai_company.security.models import AuditEntry, AuditVerdictStr # noqa: TC001 -from ai_company.security.timeout.parked_context import ParkedContext # noqa: TC001 +from ai_company.security.models import AuditEntry, AuditVerdictStr +from ai_company.security.timeout.parked_context import ParkedContext # ── Test auth constants ─────────────────────────────────────── diff --git a/tests/unit/api/controllers/test_agents.py b/tests/unit/api/controllers/test_agents.py index c6632b2282..4ebb33f665 100644 --- a/tests/unit/api/controllers/test_agents.py +++ b/tests/unit/api/controllers/test_agents.py @@ -28,7 +28,7 @@ def test_list_agents_with_data( fake_message_bus: FakeMessageBus, ) -> None: from ai_company.api.app import create_app - from ai_company.api.auth.service import AuthService # noqa: TC001 + from ai_company.api.auth.service import AuthService from ai_company.budget.tracker import CostTracker from tests.unit.api.conftest import _make_test_auth_service, _seed_test_users diff --git a/tests/unit/api/controllers/test_analytics.py b/tests/unit/api/controllers/test_analytics.py index 49fc03f29d..87dcf5b7e5 100644 --- a/tests/unit/api/controllers/test_analytics.py +++ b/tests/unit/api/controllers/test_analytics.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from ai_company.core.enums import TaskStatus from tests.unit.api.conftest import make_auth_headers diff --git a/tests/unit/api/controllers/test_approvals.py b/tests/unit/api/controllers/test_approvals.py index 29a986353e..68c2b3ca68 100644 --- a/tests/unit/api/controllers/test_approvals.py +++ b/tests/unit/api/controllers/test_approvals.py @@ -4,9 +4,9 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient -from ai_company.api.approval_store import ApprovalStore # noqa: TC001 +from ai_company.api.approval_store import ApprovalStore from ai_company.core.approval import ApprovalItem from ai_company.core.enums import ApprovalRiskLevel, ApprovalStatus from tests.unit.api.conftest import make_approval, make_auth_headers diff --git a/tests/unit/api/controllers/test_artifacts.py b/tests/unit/api/controllers/test_artifacts.py index f53093e7ef..baf2638cab 100644 --- a/tests/unit/api/controllers/test_artifacts.py +++ b/tests/unit/api/controllers/test_artifacts.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/controllers/test_autonomy.py b/tests/unit/api/controllers/test_autonomy.py index c84cde81d7..0a92b4af49 100644 --- a/tests/unit/api/controllers/test_autonomy.py +++ b/tests/unit/api/controllers/test_autonomy.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from tests.unit.api.conftest import make_auth_headers diff --git a/tests/unit/api/controllers/test_budget.py b/tests/unit/api/controllers/test_budget.py index 28c8e47a68..e2e41cc9c2 100644 --- a/tests/unit/api/controllers/test_budget.py +++ b/tests/unit/api/controllers/test_budget.py @@ -4,10 +4,10 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from ai_company.budget.cost_record import CostRecord -from ai_company.budget.tracker import CostTracker # noqa: TC001 +from ai_company.budget.tracker import CostTracker from tests.unit.api.conftest import make_auth_headers _HEADERS = make_auth_headers("ceo") diff --git a/tests/unit/api/controllers/test_company.py b/tests/unit/api/controllers/test_company.py index 02a554f3fa..66295dcc81 100644 --- a/tests/unit/api/controllers/test_company.py +++ b/tests/unit/api/controllers/test_company.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from tests.unit.api.conftest import make_auth_headers diff --git a/tests/unit/api/controllers/test_departments.py b/tests/unit/api/controllers/test_departments.py index c9afa32996..21fa024b01 100644 --- a/tests/unit/api/controllers/test_departments.py +++ b/tests/unit/api/controllers/test_departments.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/controllers/test_meetings.py b/tests/unit/api/controllers/test_meetings.py index 49d6a26d0a..774d6cfbb6 100644 --- a/tests/unit/api/controllers/test_meetings.py +++ b/tests/unit/api/controllers/test_meetings.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/controllers/test_messages.py b/tests/unit/api/controllers/test_messages.py index c5f895b675..09ec10de45 100644 --- a/tests/unit/api/controllers/test_messages.py +++ b/tests/unit/api/controllers/test_messages.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/controllers/test_projects.py b/tests/unit/api/controllers/test_projects.py index 086711203e..cf7b8c7a36 100644 --- a/tests/unit/api/controllers/test_projects.py +++ b/tests/unit/api/controllers/test_projects.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/controllers/test_providers.py b/tests/unit/api/controllers/test_providers.py index 3fb729b119..724ca24b82 100644 --- a/tests/unit/api/controllers/test_providers.py +++ b/tests/unit/api/controllers/test_providers.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/controllers/test_tasks.py b/tests/unit/api/controllers/test_tasks.py index 04755227d9..55e4801c79 100644 --- a/tests/unit/api/controllers/test_tasks.py +++ b/tests/unit/api/controllers/test_tasks.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from tests.unit.api.conftest import FakePersistenceBackend, make_auth_headers, make_task diff --git a/tests/unit/api/test_app.py b/tests/unit/api/test_app.py index 764000614c..8fb34a63db 100644 --- a/tests/unit/api/test_app.py +++ b/tests/unit/api/test_app.py @@ -4,7 +4,7 @@ import pytest from litestar import Litestar -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from ai_company.api.app import create_app diff --git a/tests/unit/api/test_guards.py b/tests/unit/api/test_guards.py index 4a1a543ae7..074d766837 100644 --- a/tests/unit/api/test_guards.py +++ b/tests/unit/api/test_guards.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from tests.unit.api.conftest import make_auth_headers diff --git a/tests/unit/api/test_health.py b/tests/unit/api/test_health.py index e08fff1e46..ba153b21ea 100644 --- a/tests/unit/api/test_health.py +++ b/tests/unit/api/test_health.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient @pytest.mark.unit diff --git a/tests/unit/api/test_middleware.py b/tests/unit/api/test_middleware.py index 418a090f88..16563759b4 100644 --- a/tests/unit/api/test_middleware.py +++ b/tests/unit/api/test_middleware.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from litestar.testing import TestClient # noqa: TC002 +from litestar.testing import TestClient from ai_company.api.middleware import ( _API_CSP, diff --git a/tests/unit/budget/test_category_analytics.py b/tests/unit/budget/test_category_analytics.py index 808cb59c46..0e0808fd3f 100644 --- a/tests/unit/budget/test_category_analytics.py +++ b/tests/unit/budget/test_category_analytics.py @@ -12,7 +12,7 @@ ) from ai_company.budget.coordination_config import OrchestrationAlertThresholds from ai_company.budget.cost_record import CostRecord -from ai_company.budget.tracker import CostTracker # noqa: TC001 +from ai_company.budget.tracker import CostTracker pytestmark = pytest.mark.timeout(30) diff --git a/tests/unit/communication/conflict_resolution/test_authority_strategy.py b/tests/unit/communication/conflict_resolution/test_authority_strategy.py index d2a590426b..bafea4d295 100644 --- a/tests/unit/communication/conflict_resolution/test_authority_strategy.py +++ b/tests/unit/communication/conflict_resolution/test_authority_strategy.py @@ -9,7 +9,7 @@ ConflictResolutionOutcome, ) from ai_company.communication.delegation.hierarchy import ( - HierarchyResolver, # noqa: TC001 + HierarchyResolver, ) from ai_company.communication.enums import ConflictResolutionStrategy from ai_company.communication.errors import ConflictHierarchyError diff --git a/tests/unit/communication/conflict_resolution/test_debate_strategy.py b/tests/unit/communication/conflict_resolution/test_debate_strategy.py index d8fe59666e..a652552115 100644 --- a/tests/unit/communication/conflict_resolution/test_debate_strategy.py +++ b/tests/unit/communication/conflict_resolution/test_debate_strategy.py @@ -12,7 +12,7 @@ ) from ai_company.communication.conflict_resolution.protocol import JudgeDecision from ai_company.communication.delegation.hierarchy import ( - HierarchyResolver, # noqa: TC001 + HierarchyResolver, ) from ai_company.communication.enums import ConflictResolutionStrategy from ai_company.communication.errors import ( diff --git a/tests/unit/communication/conflict_resolution/test_helpers.py b/tests/unit/communication/conflict_resolution/test_helpers.py index 079a287934..43e5ab8a22 100644 --- a/tests/unit/communication/conflict_resolution/test_helpers.py +++ b/tests/unit/communication/conflict_resolution/test_helpers.py @@ -9,7 +9,7 @@ pick_highest_seniority, ) from ai_company.communication.delegation.hierarchy import ( - HierarchyResolver, # noqa: TC001 + HierarchyResolver, ) from ai_company.communication.errors import ConflictStrategyError from ai_company.core.enums import SeniorityLevel diff --git a/tests/unit/communication/conflict_resolution/test_hybrid_strategy.py b/tests/unit/communication/conflict_resolution/test_hybrid_strategy.py index c223168d0d..184d1bb20d 100644 --- a/tests/unit/communication/conflict_resolution/test_hybrid_strategy.py +++ b/tests/unit/communication/conflict_resolution/test_hybrid_strategy.py @@ -15,7 +15,7 @@ ) from ai_company.communication.conflict_resolution.protocol import JudgeDecision from ai_company.communication.delegation.hierarchy import ( - HierarchyResolver, # noqa: TC001 + HierarchyResolver, ) from ai_company.communication.enums import ConflictResolutionStrategy from ai_company.core.enums import SeniorityLevel diff --git a/tests/unit/communication/conflict_resolution/test_service.py b/tests/unit/communication/conflict_resolution/test_service.py index 8ab9410c2a..287315053b 100644 --- a/tests/unit/communication/conflict_resolution/test_service.py +++ b/tests/unit/communication/conflict_resolution/test_service.py @@ -20,7 +20,7 @@ ConflictResolutionService, ) from ai_company.communication.delegation.hierarchy import ( - HierarchyResolver, # noqa: TC001 + HierarchyResolver, ) from ai_company.communication.enums import ( ConflictResolutionStrategy, diff --git a/tests/unit/communication/meeting/conftest.py b/tests/unit/communication/meeting/conftest.py index 52110dc583..c83bb424b3 100644 --- a/tests/unit/communication/meeting/conftest.py +++ b/tests/unit/communication/meeting/conftest.py @@ -13,7 +13,7 @@ MeetingAgenda, MeetingAgendaItem, ) -from ai_company.communication.meeting.protocol import AgentCaller # noqa: TC001 +from ai_company.communication.meeting.protocol import AgentCaller def make_agent_response( diff --git a/tests/unit/communication/meeting/test_orchestrator.py b/tests/unit/communication/meeting/test_orchestrator.py index 6e03b3effd..4d14549da0 100644 --- a/tests/unit/communication/meeting/test_orchestrator.py +++ b/tests/unit/communication/meeting/test_orchestrator.py @@ -28,7 +28,7 @@ from ai_company.communication.meeting.position_papers import ( PositionPapersProtocol, ) -from ai_company.communication.meeting.protocol import MeetingProtocol # noqa: TC001 +from ai_company.communication.meeting.protocol import MeetingProtocol from ai_company.communication.meeting.round_robin import RoundRobinProtocol from ai_company.core.enums import Priority from tests.unit.communication.meeting.conftest import ( diff --git a/tests/unit/communication/meeting/test_position_papers.py b/tests/unit/communication/meeting/test_position_papers.py index 9ba70531d0..2c84383879 100644 --- a/tests/unit/communication/meeting/test_position_papers.py +++ b/tests/unit/communication/meeting/test_position_papers.py @@ -10,7 +10,7 @@ from ai_company.communication.meeting.errors import ( MeetingBudgetExhaustedError, ) -from ai_company.communication.meeting.models import MeetingAgenda # noqa: TC001 +from ai_company.communication.meeting.models import MeetingAgenda from ai_company.communication.meeting.position_papers import ( PositionPapersProtocol, ) diff --git a/tests/unit/communication/meeting/test_protocol.py b/tests/unit/communication/meeting/test_protocol.py index 9348eeca1b..40e7eeea1e 100644 --- a/tests/unit/communication/meeting/test_protocol.py +++ b/tests/unit/communication/meeting/test_protocol.py @@ -3,7 +3,7 @@ import pytest from ai_company.communication.meeting.enums import MeetingProtocolType -from ai_company.communication.meeting.models import ( # noqa: TC001 +from ai_company.communication.meeting.models import ( MeetingAgenda, MeetingMinutes, ) diff --git a/tests/unit/communication/meeting/test_round_robin.py b/tests/unit/communication/meeting/test_round_robin.py index c8c37d02f1..9a8c60b4d3 100644 --- a/tests/unit/communication/meeting/test_round_robin.py +++ b/tests/unit/communication/meeting/test_round_robin.py @@ -10,7 +10,7 @@ from ai_company.communication.meeting.errors import ( MeetingBudgetExhaustedError, ) -from ai_company.communication.meeting.models import MeetingAgenda # noqa: TC001 +from ai_company.communication.meeting.models import MeetingAgenda from ai_company.communication.meeting.protocol import MeetingProtocol from ai_company.communication.meeting.round_robin import RoundRobinProtocol from tests.unit.communication.meeting.conftest import ( diff --git a/tests/unit/communication/meeting/test_structured_phases.py b/tests/unit/communication/meeting/test_structured_phases.py index 810d174294..cda3b85f08 100644 --- a/tests/unit/communication/meeting/test_structured_phases.py +++ b/tests/unit/communication/meeting/test_structured_phases.py @@ -10,7 +10,7 @@ from ai_company.communication.meeting.errors import ( MeetingBudgetExhaustedError, ) -from ai_company.communication.meeting.models import MeetingAgenda # noqa: TC001 +from ai_company.communication.meeting.models import MeetingAgenda from ai_company.communication.meeting.protocol import ConflictDetector, MeetingProtocol from ai_company.communication.meeting.structured_phases import ( KeywordConflictDetector, diff --git a/tests/unit/engine/task_engine_helpers.py b/tests/unit/engine/task_engine_helpers.py index 2a6767e165..23e7f1a30e 100644 --- a/tests/unit/engine/task_engine_helpers.py +++ b/tests/unit/engine/task_engine_helpers.py @@ -3,7 +3,7 @@ import copy from typing import TYPE_CHECKING -from ai_company.core.task import Task # noqa: TC001 +from ai_company.core.task import Task from ai_company.engine.task_engine_models import CreateTaskData if TYPE_CHECKING: diff --git a/tests/unit/engine/test_agent_engine.py b/tests/unit/engine/test_agent_engine.py index ee15b23adf..bcda021ad9 100644 --- a/tests/unit/engine/test_agent_engine.py +++ b/tests/unit/engine/test_agent_engine.py @@ -9,7 +9,7 @@ from ai_company.budget.coordination_config import ErrorTaxonomyConfig from ai_company.budget.tracker import CostTracker -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import AgentStatus, Priority, TaskStatus, TaskType from ai_company.core.task import Task from ai_company.engine.agent_engine import AgentEngine diff --git a/tests/unit/engine/test_agent_engine_errors.py b/tests/unit/engine/test_agent_engine_errors.py index 57c9d7ac35..9f9aa2cd68 100644 --- a/tests/unit/engine/test_agent_engine_errors.py +++ b/tests/unit/engine/test_agent_engine_errors.py @@ -5,9 +5,9 @@ import pytest -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import TaskStatus -from ai_company.core.task import Task # noqa: TC001 +from ai_company.core.task import Task from ai_company.engine.agent_engine import AgentEngine from ai_company.engine.context import AgentContext from ai_company.engine.loop_protocol import ( diff --git a/tests/unit/engine/test_agent_engine_lifecycle.py b/tests/unit/engine/test_agent_engine_lifecycle.py index 725f79662f..71f7037337 100644 --- a/tests/unit/engine/test_agent_engine_lifecycle.py +++ b/tests/unit/engine/test_agent_engine_lifecycle.py @@ -6,9 +6,9 @@ import pytest -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import TaskStatus -from ai_company.core.task import Task # noqa: TC001 +from ai_company.core.task import Task from ai_company.engine.agent_engine import AgentEngine from ai_company.engine.context import AgentContext from ai_company.engine.loop_protocol import ( diff --git a/tests/unit/engine/test_context.py b/tests/unit/engine/test_context.py index 9a09f5cb71..4149ecb7b5 100644 --- a/tests/unit/engine/test_context.py +++ b/tests/unit/engine/test_context.py @@ -6,9 +6,9 @@ import structlog.testing from pydantic import ValidationError -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import TaskStatus -from ai_company.core.task import Task # noqa: TC001 +from ai_company.core.task import Task from ai_company.engine.context import ( DEFAULT_MAX_TURNS, AgentContext, diff --git a/tests/unit/engine/test_loop_protocol.py b/tests/unit/engine/test_loop_protocol.py index ba5741ce1e..56360f45e0 100644 --- a/tests/unit/engine/test_loop_protocol.py +++ b/tests/unit/engine/test_loop_protocol.py @@ -6,7 +6,7 @@ from ai_company.budget.call_category import LLMCallCategory from ai_company.core.enums import Complexity, Priority, TaskStatus, TaskType from ai_company.core.task import Task -from ai_company.engine.context import AgentContext # noqa: TC001 +from ai_company.engine.context import AgentContext from ai_company.engine.loop_protocol import ( ExecutionLoop, ExecutionResult, diff --git a/tests/unit/engine/test_metrics.py b/tests/unit/engine/test_metrics.py index f3490baf2e..3704d90bc9 100644 --- a/tests/unit/engine/test_metrics.py +++ b/tests/unit/engine/test_metrics.py @@ -3,7 +3,7 @@ import pytest from pydantic import ValidationError -from ai_company.engine.context import AgentContext # noqa: TC001 +from ai_company.engine.context import AgentContext from ai_company.engine.loop_protocol import ( ExecutionResult, TerminationReason, diff --git a/tests/unit/engine/test_plan_execute_loop.py b/tests/unit/engine/test_plan_execute_loop.py index 447f463951..8b15dd4de2 100644 --- a/tests/unit/engine/test_plan_execute_loop.py +++ b/tests/unit/engine/test_plan_execute_loop.py @@ -6,7 +6,7 @@ import pytest from ai_company.budget.call_category import LLMCallCategory -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import ToolCategory from ai_company.engine.context import AgentContext from ai_company.engine.loop_protocol import TerminationReason diff --git a/tests/unit/engine/test_react_loop.py b/tests/unit/engine/test_react_loop.py index 026bf90d94..f60b14c914 100644 --- a/tests/unit/engine/test_react_loop.py +++ b/tests/unit/engine/test_react_loop.py @@ -5,7 +5,7 @@ import pytest -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import ToolCategory from ai_company.engine.context import AgentContext from ai_company.engine.loop_protocol import TerminationReason diff --git a/tests/unit/engine/test_routing_models.py b/tests/unit/engine/test_routing_models.py index 3a12c9d4bc..eff747f416 100644 --- a/tests/unit/engine/test_routing_models.py +++ b/tests/unit/engine/test_routing_models.py @@ -2,7 +2,7 @@ import pytest -from ai_company.core.agent import AgentIdentity # noqa: TC001 +from ai_company.core.agent import AgentIdentity from ai_company.core.enums import CoordinationTopology from ai_company.engine.routing.models import ( AutoTopologyConfig, diff --git a/tests/unit/engine/test_task_engine_mutations.py b/tests/unit/engine/test_task_engine_mutations.py index a9c1df9a8d..8724bba2de 100644 --- a/tests/unit/engine/test_task_engine_mutations.py +++ b/tests/unit/engine/test_task_engine_mutations.py @@ -5,7 +5,7 @@ import pytest from ai_company.core.enums import TaskStatus -from ai_company.core.task import Task # noqa: TC001 +from ai_company.core.task import Task from ai_company.engine.errors import ( TaskMutationError, TaskNotFoundError, diff --git a/tests/unit/hr/test_full_snapshot_strategy.py b/tests/unit/hr/test_full_snapshot_strategy.py index fc4e494a53..1d9482120c 100644 --- a/tests/unit/hr/test_full_snapshot_strategy.py +++ b/tests/unit/hr/test_full_snapshot_strategy.py @@ -9,7 +9,7 @@ from ai_company.core.types import NotBlankStr from ai_company.hr.errors import MemoryArchivalError from ai_company.hr.full_snapshot_strategy import FullSnapshotStrategy -from ai_company.memory.consolidation.models import ArchivalEntry # noqa: TC001 +from ai_company.memory.consolidation.models import ArchivalEntry from ai_company.memory.models import MemoryEntry, MemoryMetadata, MemoryQuery from ai_company.memory.org.models import OrgFactAuthor, OrgFactWriteRequest diff --git a/tests/unit/hr/test_hiring_service.py b/tests/unit/hr/test_hiring_service.py index 3a21f79358..dc6b6f9950 100644 --- a/tests/unit/hr/test_hiring_service.py +++ b/tests/unit/hr/test_hiring_service.py @@ -12,8 +12,8 @@ InvalidCandidateError, ) from ai_company.hr.hiring_service import HiringService -from ai_company.hr.onboarding_service import OnboardingService # noqa: TC001 -from ai_company.hr.registry import AgentRegistryService # noqa: TC001 +from ai_company.hr.onboarding_service import OnboardingService +from ai_company.hr.registry import AgentRegistryService from tests.unit.hr.conftest import make_candidate_card, make_hiring_request diff --git a/tests/unit/hr/test_offboarding_service.py b/tests/unit/hr/test_offboarding_service.py index cd7e7d9b0a..33582d37d9 100644 --- a/tests/unit/hr/test_offboarding_service.py +++ b/tests/unit/hr/test_offboarding_service.py @@ -5,10 +5,10 @@ import pytest -from ai_company.communication.channel import Channel # noqa: TC001 -from ai_company.communication.message import Message # noqa: TC001 +from ai_company.communication.channel import Channel +from ai_company.communication.message import Message from ai_company.core.enums import AgentStatus, TaskStatus -from ai_company.core.task import Task # noqa: TC001 +from ai_company.core.task import Task from ai_company.core.types import NotBlankStr from ai_company.hr.archival_protocol import ArchivalResult from ai_company.hr.errors import ( @@ -19,7 +19,7 @@ ) from ai_company.hr.models import OffboardingRecord from ai_company.hr.offboarding_service import OffboardingService -from ai_company.hr.registry import AgentRegistryService # noqa: TC001 +from ai_company.hr.registry import AgentRegistryService from tests.unit.hr.conftest import ( make_agent_identity, make_firing_request, diff --git a/tests/unit/hr/test_onboarding_service.py b/tests/unit/hr/test_onboarding_service.py index 170ea50ad4..8dfe59eb12 100644 --- a/tests/unit/hr/test_onboarding_service.py +++ b/tests/unit/hr/test_onboarding_service.py @@ -5,8 +5,8 @@ from ai_company.core.enums import AgentStatus from ai_company.hr.enums import OnboardingStep from ai_company.hr.errors import OnboardingError -from ai_company.hr.onboarding_service import OnboardingService # noqa: TC001 -from ai_company.hr.registry import AgentRegistryService # noqa: TC001 +from ai_company.hr.onboarding_service import OnboardingService +from ai_company.hr.registry import AgentRegistryService from tests.unit.hr.conftest import make_agent_identity diff --git a/tests/unit/hr/test_registry.py b/tests/unit/hr/test_registry.py index 38dcb452c1..4b5dfb535e 100644 --- a/tests/unit/hr/test_registry.py +++ b/tests/unit/hr/test_registry.py @@ -4,7 +4,7 @@ from ai_company.core.enums import AgentStatus, SeniorityLevel from ai_company.hr.errors import AgentAlreadyRegisteredError, AgentNotFoundError -from ai_company.hr.registry import AgentRegistryService # noqa: TC001 +from ai_company.hr.registry import AgentRegistryService from tests.unit.hr.conftest import make_agent_identity diff --git a/tests/unit/memory/backends/mem0/conftest.py b/tests/unit/memory/backends/mem0/conftest.py new file mode 100644 index 0000000000..30b2c20171 --- /dev/null +++ b/tests/unit/memory/backends/mem0/conftest.py @@ -0,0 +1,105 @@ +"""Shared fixtures for Mem0 adapter tests.""" + +from typing import Any +from unittest.mock import MagicMock + +import pytest + +from ai_company.core.enums import MemoryCategory +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend +from ai_company.memory.backends.mem0.config import ( + Mem0BackendConfig, + Mem0EmbedderConfig, +) +from ai_company.memory.models import MemoryStoreRequest + + +def _test_embedder() -> Mem0EmbedderConfig: + """Vendor-agnostic embedder config for tests.""" + return Mem0EmbedderConfig( + provider="test-provider", + model="test-embedding-001", + ) + + +@pytest.fixture +def mem0_config() -> Mem0BackendConfig: + """Default Mem0 config for tests.""" + return Mem0BackendConfig( + data_dir="/tmp/test-memory", # noqa: S108 + embedder=_test_embedder(), + ) + + +@pytest.fixture +def mock_client() -> MagicMock: + """Mock Mem0 Memory client.""" + return MagicMock() + + +@pytest.fixture +def backend( + mem0_config: Mem0BackendConfig, + mock_client: MagicMock, +) -> Mem0MemoryBackend: + """Connected backend with mocked Mem0 client.""" + b = Mem0MemoryBackend(mem0_config=mem0_config, max_memories_per_agent=100) + b._client = mock_client + b._connected = True + return b + + +def mem0_add_result(memory_id: str = "mem-001") -> dict[str, Any]: + """Build a typical Mem0 add() return value.""" + return { + "results": [ + { + "id": memory_id, + "memory": "test content", + "event": "ADD", + }, + ], + } + + +def mem0_search_result( + items: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + """Build a typical Mem0 search() return value.""" + if items is None: + items = [ + { + "id": "mem-001", + "memory": "found content", + "score": 0.85, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + "_synthorg_category": "episodic", + "_synthorg_confidence": 0.9, + }, + }, + ] + return {"results": items} + + +def mem0_get_result(memory_id: str = "mem-001") -> dict[str, Any]: + """Build a typical Mem0 get() return value.""" + return { + "id": memory_id, + "memory": "stored content", + "created_at": "2026-03-12T10:00:00+00:00", + "updated_at": None, + "metadata": { + "_synthorg_category": "episodic", + "_synthorg_confidence": 1.0, + }, + } + + +def make_store_request( + *, + category: MemoryCategory = MemoryCategory.EPISODIC, + content: str = "test content", +) -> MemoryStoreRequest: + """Helper to build a store request.""" + return MemoryStoreRequest(category=category, content=content) diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index 6c8508cfe9..b65b164baa 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -1,123 +1,22 @@ -"""Tests for the Mem0 memory backend adapter.""" +"""Tests for Mem0 adapter — properties, capabilities, protocol, lifecycle.""" import sys -from typing import Any from unittest.mock import MagicMock, patch import pytest from ai_company.core.enums import MemoryCategory -from ai_company.memory.backends.mem0.adapter import ( - _SHARED_NAMESPACE, - Mem0MemoryBackend, -) -from ai_company.memory.backends.mem0.config import ( - Mem0BackendConfig, - Mem0EmbedderConfig, -) -from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY -from ai_company.memory.errors import ( - MemoryConnectionError, - MemoryRetrievalError, - MemoryStoreError, -) -from ai_company.memory.models import ( - MemoryQuery, - MemoryStoreRequest, -) - -pytestmark = pytest.mark.timeout(30) - - -def _test_embedder() -> Mem0EmbedderConfig: - """Vendor-agnostic embedder config for tests.""" - return Mem0EmbedderConfig( - provider="test-provider", - model="test-embedding-001", - ) - - -@pytest.fixture -def mem0_config() -> Mem0BackendConfig: - """Default Mem0 config for tests.""" - return Mem0BackendConfig( - data_dir="/tmp/test-memory", # noqa: S108 - embedder=_test_embedder(), - ) - - -@pytest.fixture -def mock_client() -> MagicMock: - """Mock Mem0 Memory client.""" - return MagicMock() - - -@pytest.fixture -def backend( - mem0_config: Mem0BackendConfig, - mock_client: MagicMock, -) -> Mem0MemoryBackend: - """Connected backend with mocked Mem0 client.""" - b = Mem0MemoryBackend(mem0_config=mem0_config, max_memories_per_agent=100) - b._client = mock_client - b._connected = True - return b - - -def _mem0_add_result(memory_id: str = "mem-001") -> dict[str, Any]: - """Build a typical Mem0 add() return value.""" - return { - "results": [ - { - "id": memory_id, - "memory": "test content", - "event": "ADD", - }, - ], - } +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend +from ai_company.memory.backends.mem0.config import Mem0BackendConfig +from ai_company.memory.capabilities import MemoryCapabilities +from ai_company.memory.errors import MemoryConnectionError +from ai_company.memory.models import MemoryQuery +from ai_company.memory.protocol import MemoryBackend +from ai_company.memory.shared import SharedKnowledgeStore +from .conftest import make_store_request -def _mem0_search_result( - items: list[dict[str, Any]] | None = None, -) -> dict[str, Any]: - """Build a typical Mem0 search() return value.""" - if items is None: - items = [ - { - "id": "mem-001", - "memory": "found content", - "score": 0.85, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": { - "_synthorg_category": "episodic", - "_synthorg_confidence": 0.9, - }, - }, - ] - return {"results": items} - - -def _mem0_get_result(memory_id: str = "mem-001") -> dict[str, Any]: - """Build a typical Mem0 get() return value.""" - return { - "id": memory_id, - "memory": "stored content", - "created_at": "2026-03-12T10:00:00+00:00", - "updated_at": None, - "metadata": { - "_synthorg_category": "episodic", - "_synthorg_confidence": 1.0, - }, - } - - -def _make_store_request( - *, - category: MemoryCategory = MemoryCategory.EPISODIC, - content: str = "test content", -) -> MemoryStoreRequest: - """Helper to build a store request.""" - return MemoryStoreRequest(category=category, content=content) +pytestmark = pytest.mark.timeout(30) # ── Properties ──────────────────────────────────────────────────── @@ -147,7 +46,10 @@ def test_supported_categories(self, backend: Mem0MemoryBackend) -> None: def test_supports_graph_false(self, backend: Mem0MemoryBackend) -> None: assert backend.supports_graph is False - def test_supports_temporal_true(self, backend: Mem0MemoryBackend) -> None: + def test_supports_temporal_true( + self, + backend: Mem0MemoryBackend, + ) -> None: assert backend.supports_temporal is True def test_supports_vector_search_true( @@ -208,6 +110,24 @@ def test_has_shared_knowledge_methods( assert hasattr(backend, "search_shared") assert hasattr(backend, "retract") + def test_isinstance_memory_backend( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert isinstance(backend, MemoryBackend) + + def test_isinstance_memory_capabilities( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert isinstance(backend, MemoryCapabilities) + + def test_isinstance_shared_knowledge_store( + self, + backend: Mem0MemoryBackend, + ) -> None: + assert isinstance(backend, SharedKnowledgeStore) + # ── Lifecycle ───────────────────────────────────────────────────── @@ -296,11 +216,41 @@ async def test_health_check_reraises_memory_error( backend: Mem0MemoryBackend, mock_client: MagicMock, ) -> None: - """MemoryError propagates through health_check.""" + """builtins.MemoryError propagates through health_check.""" mock_client.get_all.side_effect = MemoryError("out of memory") with pytest.raises(MemoryError): await backend.health_check() + async def test_connect_memory_error_propagates( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + """builtins.MemoryError from connect is not wrapped.""" + b = Mem0MemoryBackend(mem0_config=mem0_config) + with ( + patch( + "ai_company.memory.backends.mem0.adapter.asyncio.to_thread", + side_effect=MemoryError("out of memory"), + ), + pytest.raises(MemoryError), + ): + await b.connect() + + async def test_connect_recursion_error_propagates( + self, + mem0_config: Mem0BackendConfig, + ) -> None: + """RecursionError from connect is not wrapped.""" + b = Mem0MemoryBackend(mem0_config=mem0_config) + with ( + patch( + "ai_company.memory.backends.mem0.adapter.asyncio.to_thread", + side_effect=RecursionError("infinite loop"), + ), + pytest.raises(RecursionError), + ): + await b.connect() + # ── Connection guard ────────────────────────────────────────────── @@ -313,7 +263,7 @@ async def test_store_raises_when_disconnected( ) -> None: b = Mem0MemoryBackend(mem0_config=mem0_config) with pytest.raises(MemoryConnectionError, match="Not connected"): - await b.store("test-agent-001", _make_store_request()) + await b.store("test-agent-001", make_store_request()) async def test_retrieve_raises_when_disconnected( self, @@ -353,7 +303,7 @@ async def test_publish_raises_when_disconnected( ) -> None: b = Mem0MemoryBackend(mem0_config=mem0_config) with pytest.raises(MemoryConnectionError, match="Not connected"): - await b.publish("test-agent-001", _make_store_request()) + await b.publish("test-agent-001", make_store_request()) async def test_search_shared_raises_when_disconnected( self, @@ -370,748 +320,3 @@ async def test_retract_raises_when_disconnected( b = Mem0MemoryBackend(mem0_config=mem0_config) with pytest.raises(MemoryConnectionError, match="Not connected"): await b.retract("test-agent-001", "mem-001") - - -# ── Store ───────────────────────────────────────────────────────── - - -@pytest.mark.unit -class TestStore: - async def test_store_success( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.return_value = _mem0_add_result("new-mem-id") - - memory_id = await backend.store( - "test-agent-001", - _make_store_request(), - ) - - assert memory_id == "new-mem-id" - mock_client.add.assert_called_once() - call_kwargs = mock_client.add.call_args[1] - assert call_kwargs["user_id"] == "test-agent-001" - assert call_kwargs["infer"] is False - - async def test_store_empty_results_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.return_value = {"results": []} - - with pytest.raises(MemoryStoreError, match="no results"): - await backend.store("test-agent-001", _make_store_request()) - - async def test_store_missing_id_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.return_value = { - "results": [{"memory": "no id", "event": "ADD"}], - } - - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - await backend.store("test-agent-001", _make_store_request()) - - async def test_store_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.side_effect = RuntimeError("disk full") - - with pytest.raises(MemoryStoreError, match="Failed to store") as exc_info: - await backend.store("test-agent-001", _make_store_request()) - - assert exc_info.value.__cause__ is not None - - async def test_store_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping.""" - mock_client.add.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.store("test-agent-001", _make_store_request()) - - -# ── Retrieve ────────────────────────────────────────────────────── - - -@pytest.mark.unit -class TestRetrieve: - async def test_retrieve_with_text( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.search.return_value = _mem0_search_result() - - query = MemoryQuery(text="find relevant", limit=5) - entries = await backend.retrieve("test-agent-001", query) - - assert len(entries) == 1 - assert entries[0].content == "found content" - assert entries[0].relevance_score == 0.85 - mock_client.search.assert_called_once() - - async def test_retrieve_without_text( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get_all.return_value = _mem0_search_result( - [ - { - "id": "mem-001", - "memory": "all content", - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {}, - }, - ], - ) - - query = MemoryQuery(text=None, limit=10) - entries = await backend.retrieve("test-agent-001", query) - - assert len(entries) == 1 - mock_client.get_all.assert_called_once() - - async def test_retrieve_applies_post_filters( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.search.return_value = _mem0_search_result( - [ - { - "id": "m1", - "memory": "episodic", - "score": 0.9, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {"_synthorg_category": "episodic"}, - }, - { - "id": "m2", - "memory": "semantic", - "score": 0.8, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {"_synthorg_category": "semantic"}, - }, - ], - ) - - query = MemoryQuery( - text="test", - categories=frozenset({MemoryCategory.EPISODIC}), - ) - entries = await backend.retrieve("test-agent-001", query) - - assert len(entries) == 1 - assert entries[0].category == MemoryCategory.EPISODIC - - async def test_retrieve_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.search.side_effect = RuntimeError("search failed") - - with pytest.raises(MemoryRetrievalError, match="Failed to retrieve"): - await backend.retrieve( - "test-agent-001", - MemoryQuery(text="test"), - ) - - async def test_retrieve_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping.""" - mock_client.search.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.retrieve( - "test-agent-001", - MemoryQuery(text="test"), - ) - - -# ── Get ─────────────────────────────────────────────────────────── - - -@pytest.mark.unit -class TestGet: - async def test_get_existing( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = _mem0_get_result("mem-001") - - entry = await backend.get("test-agent-001", "mem-001") - - assert entry is not None - assert entry.id == "mem-001" - assert entry.agent_id == "test-agent-001" - - async def test_get_not_found( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = None - - entry = await backend.get("test-agent-001", "nonexistent") - - assert entry is None - - async def test_get_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.side_effect = RuntimeError("backend error") - - with pytest.raises(MemoryRetrievalError, match="Failed to get"): - await backend.get("test-agent-001", "mem-001") - - -# ── Delete ──────────────────────────────────────────────────────── - - -@pytest.mark.unit -class TestDelete: - async def test_delete_existing( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = _mem0_get_result("mem-001") - mock_client.delete.return_value = None - - result = await backend.delete("test-agent-001", "mem-001") - - assert result is True - mock_client.delete.assert_called_once_with("mem-001") - - async def test_delete_not_found( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = None - - result = await backend.delete("test-agent-001", "nonexistent") - - assert result is False - mock_client.delete.assert_not_called() - - async def test_delete_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.side_effect = RuntimeError("backend error") - - with pytest.raises(MemoryStoreError, match="Failed to delete"): - await backend.delete("test-agent-001", "mem-001") - - async def test_delete_get_ok_but_delete_fails( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = _mem0_get_result("mem-001") - mock_client.delete.side_effect = RuntimeError("delete failed") - - with pytest.raises(MemoryStoreError, match="Failed to delete"): - await backend.delete("test-agent-001", "mem-001") - - -# ── Count ───────────────────────────────────────────────────────── - - -@pytest.mark.unit -class TestCount: - async def test_count_all( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get_all.return_value = { - "results": [ - {"id": "m1", "memory": "a", "metadata": {}}, - {"id": "m2", "memory": "b", "metadata": {}}, - ], - } - - count = await backend.count("test-agent-001") - assert count == 2 - - async def test_count_by_category( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get_all.return_value = { - "results": [ - { - "id": "m1", - "memory": "a", - "metadata": {"_synthorg_category": "episodic"}, - }, - { - "id": "m2", - "memory": "b", - "metadata": {"_synthorg_category": "semantic"}, - }, - { - "id": "m3", - "memory": "c", - "metadata": {"_synthorg_category": "episodic"}, - }, - ], - } - - count = await backend.count( - "test-agent-001", - category=MemoryCategory.EPISODIC, - ) - assert count == 2 - - async def test_count_with_invalid_category_in_data( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Invalid category in stored data defaults to WORKING.""" - mock_client.get_all.return_value = { - "results": [ - { - "id": "m1", - "memory": "a", - "metadata": {"_synthorg_category": "bogus_category"}, - }, - { - "id": "m2", - "memory": "b", - "metadata": {"_synthorg_category": "episodic"}, - }, - ], - } - - count = await backend.count( - "test-agent-001", - category=MemoryCategory.WORKING, - ) - # "bogus_category" defaults to WORKING - assert count == 1 - - async def test_count_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get_all.side_effect = RuntimeError("fail") - - with pytest.raises(MemoryRetrievalError, match="Failed to count"): - await backend.count("test-agent-001") - - -# ── Shared Knowledge Store ──────────────────────────────────────── - - -@pytest.mark.unit -class TestPublish: - async def test_publish_success( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.return_value = _mem0_add_result("shared-mem-001") - - memory_id = await backend.publish( - "test-agent-001", - _make_store_request(), - ) - - assert memory_id == "shared-mem-001" - call_kwargs = mock_client.add.call_args[1] - assert call_kwargs["user_id"] == _SHARED_NAMESPACE - assert _PUBLISHER_KEY in call_kwargs["metadata"] - assert call_kwargs["metadata"][_PUBLISHER_KEY] == "test-agent-001" - - async def test_publish_empty_results_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.return_value = {"results": []} - - with pytest.raises(MemoryStoreError, match="no results"): - await backend.publish("test-agent-001", _make_store_request()) - - async def test_publish_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.add.side_effect = RuntimeError("network error") - - with pytest.raises(MemoryStoreError, match="Failed to publish"): - await backend.publish("test-agent-001", _make_store_request()) - - async def test_publish_missing_id_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Publish result missing 'id' raises MemoryStoreError.""" - mock_client.add.return_value = { - "results": [{"memory": "no id", "event": "ADD"}], - } - - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - await backend.publish("test-agent-001", _make_store_request()) - - -@pytest.mark.unit -class TestSearchShared: - async def test_search_shared_with_text( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.search.return_value = _mem0_search_result( - [ - { - "id": "shared-1", - "memory": "shared fact", - "score": 0.9, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": { - "_synthorg_category": "semantic", - _PUBLISHER_KEY: "test-agent-002", - }, - }, - ], - ) - - query = MemoryQuery(text="find shared", limit=5) - entries = await backend.search_shared(query) - - assert len(entries) == 1 - assert entries[0].agent_id == "test-agent-002" - mock_client.search.assert_called_once() - call_kwargs = mock_client.search.call_args[1] - assert call_kwargs["user_id"] == _SHARED_NAMESPACE - - async def test_search_shared_without_text( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get_all.return_value = _mem0_search_result( - [ - { - "id": "shared-1", - "memory": "shared fact", - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": { - _PUBLISHER_KEY: "test-agent-002", - }, - }, - ], - ) - - query = MemoryQuery(text=None) - entries = await backend.search_shared(query) - - assert len(entries) == 1 - mock_client.get_all.assert_called_once() - - async def test_search_shared_exclude_agent( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.search.return_value = _mem0_search_result( - [ - { - "id": "s1", - "memory": "from agent 1", - "score": 0.9, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, - }, - { - "id": "s2", - "memory": "from agent 2", - "score": 0.8, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-002"}, - }, - ], - ) - - query = MemoryQuery(text="test") - entries = await backend.search_shared( - query, - exclude_agent="test-agent-001", - ) - - assert len(entries) == 1 - assert entries[0].agent_id == "test-agent-002" - - async def test_search_shared_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.search.side_effect = RuntimeError("search error") - - with pytest.raises(MemoryRetrievalError, match="Failed to search"): - await backend.search_shared(MemoryQuery(text="test")) - - -@pytest.mark.unit -class TestRetract: - async def test_retract_success( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = { - "id": "shared-001", - "memory": "shared content", - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, - } - mock_client.delete.return_value = None - - result = await backend.retract("test-agent-001", "shared-001") - - assert result is True - mock_client.delete.assert_called_once_with("shared-001") - - async def test_retract_not_found( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = None - - result = await backend.retract("test-agent-001", "nonexistent") - - assert result is False - - async def test_retract_ownership_mismatch( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = { - "id": "shared-001", - "memory": "content", - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-002"}, - } - - with pytest.raises(MemoryStoreError, match="cannot retract"): - await backend.retract("test-agent-001", "shared-001") - - async def test_retract_no_publisher_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.return_value = { - "id": "not-shared-001", - "memory": "private content", - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {}, - } - - with pytest.raises(MemoryStoreError, match="not a shared memory"): - await backend.retract("test-agent-001", "not-shared-001") - - async def test_retract_exception_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - mock_client.get.side_effect = RuntimeError("backend error") - - with pytest.raises(MemoryStoreError, match="Failed to retract"): - await backend.retract("test-agent-001", "shared-001") - - async def test_retract_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping.""" - mock_client.get.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.retract("test-agent-001", "shared-001") - - async def test_retract_delete_failure_wraps( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Exception during delete phase wraps in MemoryStoreError.""" - mock_client.get.return_value = { - "id": "shared-001", - "memory": "content", - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, - } - mock_client.delete.side_effect = RuntimeError("delete failed") - - with pytest.raises(MemoryStoreError, match="Failed to retract"): - await backend.retract("test-agent-001", "shared-001") - - -@pytest.mark.unit -class TestAdditionalEdgeCases: - """Edge cases for improved coverage.""" - - async def test_store_blank_id_from_add_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Store result with blank ID raises MemoryStoreError.""" - mock_client.add.return_value = { - "results": [{"id": "", "event": "ADD"}], - } - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - await backend.store("test-agent-001", _make_store_request()) - - async def test_store_whitespace_id_from_add_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Store result with whitespace-only ID raises MemoryStoreError.""" - mock_client.add.return_value = { - "results": [{"id": " ", "event": "ADD"}], - } - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - await backend.store("test-agent-001", _make_store_request()) - - async def test_get_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping in get().""" - mock_client.get.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.get("test-agent-001", "mem-001") - - async def test_delete_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping in delete().""" - mock_client.get.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.delete("test-agent-001", "mem-001") - - async def test_count_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping in count().""" - mock_client.get_all.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.count("test-agent-001") - - async def test_publish_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping in publish().""" - mock_client.add.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.publish("test-agent-001", _make_store_request()) - - async def test_search_shared_reraises_memory_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """MemoryError is re-raised without wrapping in search_shared().""" - mock_client.search.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): - await backend.search_shared(MemoryQuery(text="test")) - - async def test_store_non_list_results_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Store result with non-list 'results' raises MemoryStoreError.""" - mock_client.add.return_value = {"results": "not-a-list"} - with pytest.raises(MemoryStoreError, match="no results"): - await backend.store("test-agent-001", _make_store_request()) - - async def test_retrieve_invalid_entry_raises( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Invalid entry in search results wraps as MemoryRetrievalError.""" - mock_client.search.return_value = { - "results": [ - {"id": "", "memory": "blank id", "metadata": {}}, - ], - } - with pytest.raises(MemoryRetrievalError, match="missing or blank"): - await backend.retrieve( - "test-agent-001", - MemoryQuery(text="test"), - ) - - async def test_search_shared_no_publisher_uses_namespace( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Entries without publisher metadata use the shared namespace.""" - mock_client.search.return_value = _mem0_search_result( - [ - { - "id": "shared-1", - "memory": "orphan fact", - "score": 0.9, - "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {"_synthorg_category": "semantic"}, - }, - ], - ) - - entries = await backend.search_shared(MemoryQuery(text="test")) - assert len(entries) == 1 - assert entries[0].agent_id == _SHARED_NAMESPACE - - async def test_count_empty_results( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """Count returns 0 for empty results.""" - mock_client.get_all.return_value = {"results": []} - count = await backend.count("test-agent-001") - assert count == 0 diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py new file mode 100644 index 0000000000..3fbcccb859 --- /dev/null +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -0,0 +1,515 @@ +"""Tests for Mem0 adapter — store, retrieve, get, delete, count.""" + +from unittest.mock import MagicMock + +import pytest + +from ai_company.core.enums import MemoryCategory +from ai_company.memory.backends.mem0.adapter import ( + _SHARED_NAMESPACE, + Mem0MemoryBackend, +) +from ai_company.memory.errors import ( + MemoryRetrievalError, + MemoryStoreError, +) +from ai_company.memory.models import MemoryQuery + +from .conftest import ( + make_store_request, + mem0_add_result, + mem0_get_result, + mem0_search_result, +) + +pytestmark = pytest.mark.timeout(30) + + +# ── Store ───────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestStore: + async def test_store_success( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = mem0_add_result("new-mem-id") + + memory_id = await backend.store( + "test-agent-001", + make_store_request(), + ) + + assert memory_id == "new-mem-id" + mock_client.add.assert_called_once() + call_kwargs = mock_client.add.call_args[1] + assert call_kwargs["user_id"] == "test-agent-001" + assert call_kwargs["infer"] is False + + async def test_store_empty_results_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = {"results": []} + + with pytest.raises(MemoryStoreError, match="no results"): + await backend.store("test-agent-001", make_store_request()) + + async def test_store_missing_id_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = { + "results": [{"memory": "no id", "event": "ADD"}], + } + + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + await backend.store("test-agent-001", make_store_request()) + + async def test_store_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.side_effect = RuntimeError("disk full") + + with pytest.raises(MemoryStoreError, match="Failed to store") as exc_info: + await backend.store("test-agent-001", make_store_request()) + + assert exc_info.value.__cause__ is not None + + async def test_store_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping.""" + mock_client.add.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.store("test-agent-001", make_store_request()) + + async def test_store_blank_id_from_add_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store result with blank ID raises MemoryStoreError.""" + mock_client.add.return_value = { + "results": [{"id": "", "event": "ADD"}], + } + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + await backend.store("test-agent-001", make_store_request()) + + async def test_store_whitespace_id_from_add_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store result with whitespace-only ID raises MemoryStoreError.""" + mock_client.add.return_value = { + "results": [{"id": " ", "event": "ADD"}], + } + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + await backend.store("test-agent-001", make_store_request()) + + async def test_store_non_list_results_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Store result with non-list 'results' raises MemoryStoreError.""" + mock_client.add.return_value = {"results": "not-a-list"} + with pytest.raises(MemoryStoreError, match="no results"): + await backend.store("test-agent-001", make_store_request()) + + async def test_store_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """Storing with the shared namespace agent ID is rejected.""" + with pytest.raises(MemoryStoreError, match="reserved shared namespace"): + await backend.store(_SHARED_NAMESPACE, make_store_request()) + + +# ── Retrieve ────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestRetrieve: + async def test_retrieve_with_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = mem0_search_result() + + query = MemoryQuery(text="find relevant", limit=5) + entries = await backend.retrieve("test-agent-001", query) + + assert len(entries) == 1 + assert entries[0].content == "found content" + assert entries[0].relevance_score == 0.85 + mock_client.search.assert_called_once() + + async def test_retrieve_without_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = mem0_search_result( + [ + { + "id": "mem-001", + "memory": "all content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {}, + }, + ], + ) + + query = MemoryQuery(text=None, limit=10) + entries = await backend.retrieve("test-agent-001", query) + + assert len(entries) == 1 + mock_client.get_all.assert_called_once() + + async def test_retrieve_applies_post_filters( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = mem0_search_result( + [ + { + "id": "m1", + "memory": "episodic", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {"_synthorg_category": "episodic"}, + }, + { + "id": "m2", + "memory": "semantic", + "score": 0.8, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {"_synthorg_category": "semantic"}, + }, + ], + ) + + query = MemoryQuery( + text="test", + categories=frozenset({MemoryCategory.EPISODIC}), + ) + entries = await backend.retrieve("test-agent-001", query) + + assert len(entries) == 1 + assert entries[0].category == MemoryCategory.EPISODIC + + async def test_retrieve_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.side_effect = RuntimeError("search failed") + + with pytest.raises(MemoryRetrievalError, match="Failed to retrieve"): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + + async def test_retrieve_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping.""" + mock_client.search.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + + async def test_retrieve_invalid_entry_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Invalid entry in search results wraps as MemoryRetrievalError.""" + mock_client.search.return_value = { + "results": [ + {"id": "", "memory": "blank id", "metadata": {}}, + ], + } + with pytest.raises(MemoryRetrievalError, match="missing or blank"): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + + +# ── Get ─────────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestGet: + async def test_get_existing( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = mem0_get_result("mem-001") + + entry = await backend.get("test-agent-001", "mem-001") + + assert entry is not None + assert entry.id == "mem-001" + assert entry.agent_id == "test-agent-001" + + async def test_get_not_found( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = None + + entry = await backend.get("test-agent-001", "nonexistent") + + assert entry is None + + async def test_get_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.side_effect = RuntimeError("backend error") + + with pytest.raises(MemoryRetrievalError, match="Failed to get"): + await backend.get("test-agent-001", "mem-001") + + async def test_get_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping in get().""" + mock_client.get.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.get("test-agent-001", "mem-001") + + async def test_get_ownership_mismatch_returns_none( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """get() returns None when user_id doesn't match agent_id.""" + result = mem0_get_result("mem-001") + result["user_id"] = "other-agent" + mock_client.get.return_value = result + + entry = await backend.get("test-agent-001", "mem-001") + assert entry is None + + +# ── Delete ──────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestDelete: + async def test_delete_existing( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = mem0_get_result("mem-001") + mock_client.delete.return_value = None + + result = await backend.delete("test-agent-001", "mem-001") + + assert result is True + mock_client.delete.assert_called_once_with("mem-001") + + async def test_delete_not_found( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = None + + result = await backend.delete("test-agent-001", "nonexistent") + + assert result is False + mock_client.delete.assert_not_called() + + async def test_delete_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.side_effect = RuntimeError("backend error") + + with pytest.raises(MemoryStoreError, match="Failed to delete"): + await backend.delete("test-agent-001", "mem-001") + + async def test_delete_get_ok_but_delete_fails( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = mem0_get_result("mem-001") + mock_client.delete.side_effect = RuntimeError("delete failed") + + with pytest.raises(MemoryStoreError, match="Failed to delete"): + await backend.delete("test-agent-001", "mem-001") + + async def test_delete_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping in delete().""" + mock_client.get.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.delete("test-agent-001", "mem-001") + + async def test_delete_shared_namespace_entry_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """delete() rejects entries belonging to the shared namespace.""" + result = mem0_get_result("mem-001") + result["user_id"] = _SHARED_NAMESPACE + mock_client.get.return_value = result + + with pytest.raises(MemoryStoreError, match="shared namespace"): + await backend.delete("test-agent-001", "mem-001") + + async def test_delete_ownership_mismatch_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """delete() rejects when user_id doesn't match agent_id.""" + result = mem0_get_result("mem-001") + result["user_id"] = "other-agent" + mock_client.get.return_value = result + + with pytest.raises(MemoryStoreError, match="cannot delete"): + await backend.delete("test-agent-001", "mem-001") + + +# ── Count ───────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestCount: + async def test_count_all( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = { + "results": [ + {"id": "m1", "memory": "a", "metadata": {}}, + {"id": "m2", "memory": "b", "metadata": {}}, + ], + } + + count = await backend.count("test-agent-001") + assert count == 2 + + async def test_count_by_category( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = { + "results": [ + { + "id": "m1", + "memory": "a", + "metadata": {"_synthorg_category": "episodic"}, + }, + { + "id": "m2", + "memory": "b", + "metadata": {"_synthorg_category": "semantic"}, + }, + { + "id": "m3", + "memory": "c", + "metadata": {"_synthorg_category": "episodic"}, + }, + ], + } + + count = await backend.count( + "test-agent-001", + category=MemoryCategory.EPISODIC, + ) + assert count == 2 + + async def test_count_with_invalid_category_in_data( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Invalid category in stored data defaults to WORKING.""" + mock_client.get_all.return_value = { + "results": [ + { + "id": "m1", + "memory": "a", + "metadata": {"_synthorg_category": "bogus_category"}, + }, + { + "id": "m2", + "memory": "b", + "metadata": {"_synthorg_category": "episodic"}, + }, + ], + } + + count = await backend.count( + "test-agent-001", + category=MemoryCategory.WORKING, + ) + # "bogus_category" defaults to WORKING + assert count == 1 + + async def test_count_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.side_effect = RuntimeError("fail") + + with pytest.raises(MemoryRetrievalError, match="Failed to count"): + await backend.count("test-agent-001") + + async def test_count_empty_results( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Count returns 0 for empty results.""" + mock_client.get_all.return_value = {"results": []} + count = await backend.count("test-agent-001") + assert count == 0 + + async def test_count_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping in count().""" + mock_client.get_all.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.count("test-agent-001") diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py new file mode 100644 index 0000000000..4711ce5ecb --- /dev/null +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -0,0 +1,327 @@ +"""Tests for Mem0 adapter — shared knowledge store (publish, search, retract).""" + +from unittest.mock import MagicMock + +import pytest + +from ai_company.memory.backends.mem0.adapter import ( + _SHARED_NAMESPACE, + Mem0MemoryBackend, +) +from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY +from ai_company.memory.errors import ( + MemoryRetrievalError, + MemoryStoreError, +) +from ai_company.memory.models import MemoryQuery + +from .conftest import ( + make_store_request, + mem0_add_result, + mem0_search_result, +) + +pytestmark = pytest.mark.timeout(30) + + +# ── Publish ────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestPublish: + async def test_publish_success( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = mem0_add_result("shared-mem-001") + + memory_id = await backend.publish( + "test-agent-001", + make_store_request(), + ) + + assert memory_id == "shared-mem-001" + call_kwargs = mock_client.add.call_args[1] + assert call_kwargs["user_id"] == _SHARED_NAMESPACE + assert _PUBLISHER_KEY in call_kwargs["metadata"] + assert call_kwargs["metadata"][_PUBLISHER_KEY] == "test-agent-001" + + async def test_publish_empty_results_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.return_value = {"results": []} + + with pytest.raises(MemoryStoreError, match="no results"): + await backend.publish("test-agent-001", make_store_request()) + + async def test_publish_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.add.side_effect = RuntimeError("network error") + + with pytest.raises(MemoryStoreError, match="Failed to publish"): + await backend.publish("test-agent-001", make_store_request()) + + async def test_publish_missing_id_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Publish result missing 'id' raises MemoryStoreError.""" + mock_client.add.return_value = { + "results": [{"memory": "no id", "event": "ADD"}], + } + + with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + await backend.publish("test-agent-001", make_store_request()) + + async def test_publish_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping.""" + mock_client.add.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.publish("test-agent-001", make_store_request()) + + +# ── SearchShared ───────────────────────────────────────────────── + + +@pytest.mark.unit +class TestSearchShared: + async def test_search_shared_with_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = mem0_search_result( + [ + { + "id": "shared-1", + "memory": "shared fact", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + "_synthorg_category": "semantic", + _PUBLISHER_KEY: "test-agent-002", + }, + }, + ], + ) + + query = MemoryQuery(text="find shared", limit=5) + entries = await backend.search_shared(query) + + assert len(entries) == 1 + assert entries[0].agent_id == "test-agent-002" + mock_client.search.assert_called_once() + call_kwargs = mock_client.search.call_args[1] + assert call_kwargs["user_id"] == _SHARED_NAMESPACE + + async def test_search_shared_without_text( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get_all.return_value = mem0_search_result( + [ + { + "id": "shared-1", + "memory": "shared fact", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + _PUBLISHER_KEY: "test-agent-002", + }, + }, + ], + ) + + query = MemoryQuery(text=None) + entries = await backend.search_shared(query) + + assert len(entries) == 1 + mock_client.get_all.assert_called_once() + + async def test_search_shared_exclude_agent( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.return_value = mem0_search_result( + [ + { + "id": "s1", + "memory": "from agent 1", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + }, + { + "id": "s2", + "memory": "from agent 2", + "score": 0.8, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + }, + ], + ) + + query = MemoryQuery(text="test") + entries = await backend.search_shared( + query, + exclude_agent="test-agent-001", + ) + + assert len(entries) == 1 + assert entries[0].agent_id == "test-agent-002" + + async def test_search_shared_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.search.side_effect = RuntimeError("search error") + + with pytest.raises(MemoryRetrievalError, match="Failed to search"): + await backend.search_shared(MemoryQuery(text="test")) + + async def test_search_shared_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping.""" + mock_client.search.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.search_shared(MemoryQuery(text="test")) + + async def test_search_shared_no_publisher_uses_namespace( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Entries without publisher metadata use the shared namespace.""" + mock_client.search.return_value = mem0_search_result( + [ + { + "id": "shared-1", + "memory": "orphan fact", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {"_synthorg_category": "semantic"}, + }, + ], + ) + + entries = await backend.search_shared(MemoryQuery(text="test")) + assert len(entries) == 1 + assert entries[0].agent_id == _SHARED_NAMESPACE + + +# ── Retract ────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestRetract: + async def test_retract_success( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = { + "id": "shared-001", + "memory": "shared content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + } + mock_client.delete.return_value = None + + result = await backend.retract("test-agent-001", "shared-001") + + assert result is True + mock_client.delete.assert_called_once_with("shared-001") + + async def test_retract_not_found( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = None + + result = await backend.retract("test-agent-001", "nonexistent") + + assert result is False + + async def test_retract_ownership_mismatch( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = { + "id": "shared-001", + "memory": "content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + } + + with pytest.raises(MemoryStoreError, match="cannot retract"): + await backend.retract("test-agent-001", "shared-001") + + async def test_retract_no_publisher_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.return_value = { + "id": "not-shared-001", + "memory": "private content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {}, + } + + with pytest.raises(MemoryStoreError, match="not a shared memory"): + await backend.retract("test-agent-001", "not-shared-001") + + async def test_retract_exception_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + mock_client.get.side_effect = RuntimeError("backend error") + + with pytest.raises(MemoryStoreError, match="Failed to retract"): + await backend.retract("test-agent-001", "shared-001") + + async def test_retract_reraises_memory_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """builtins.MemoryError is re-raised without wrapping.""" + mock_client.get.side_effect = MemoryError("out of memory") + with pytest.raises(MemoryError): + await backend.retract("test-agent-001", "shared-001") + + async def test_retract_delete_failure_wraps( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """Exception during delete phase wraps in MemoryStoreError.""" + mock_client.get.return_value = { + "id": "shared-001", + "memory": "content", + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + } + mock_client.delete.side_effect = RuntimeError("delete failed") + + with pytest.raises(MemoryStoreError, match="Failed to retract"): + await backend.retract("test-agent-001", "shared-001") diff --git a/tests/unit/providers/conftest.py b/tests/unit/providers/conftest.py index abcfe0eb63..c5c8c834cc 100644 --- a/tests/unit/providers/conftest.py +++ b/tests/unit/providers/conftest.py @@ -1,6 +1,6 @@ """Unit test configuration and fixtures for provider models.""" -from collections.abc import AsyncIterator # noqa: TC003 +from collections.abc import AsyncIterator import pytest from polyfactory.factories.pydantic_factory import ModelFactory diff --git a/tests/unit/providers/test_protocol.py b/tests/unit/providers/test_protocol.py index f66bdc41d6..c9aed4b697 100644 --- a/tests/unit/providers/test_protocol.py +++ b/tests/unit/providers/test_protocol.py @@ -1,6 +1,6 @@ """Tests for CompletionProvider protocol and BaseCompletionProvider ABC.""" -from collections.abc import AsyncIterator # noqa: TC003 +from collections.abc import AsyncIterator from unittest.mock import AsyncMock, MagicMock import pytest diff --git a/tests/unit/tools/git/conftest.py b/tests/unit/tools/git/conftest.py index 54ec344b0d..6ff2874392 100644 --- a/tests/unit/tools/git/conftest.py +++ b/tests/unit/tools/git/conftest.py @@ -2,7 +2,7 @@ import os import subprocess -from pathlib import Path # noqa: TC003 — pytest evaluates annotations +from pathlib import Path import pytest diff --git a/tests/unit/tools/git/test_git_sandbox_integration.py b/tests/unit/tools/git/test_git_sandbox_integration.py index 9c534ec382..f293f53375 100644 --- a/tests/unit/tools/git/test_git_sandbox_integration.py +++ b/tests/unit/tools/git/test_git_sandbox_integration.py @@ -1,6 +1,6 @@ """Tests for git tools with sandbox integration.""" -from pathlib import Path # noqa: TC003 — used at runtime +from pathlib import Path from unittest.mock import AsyncMock import pytest diff --git a/tests/unit/tools/sandbox/conftest.py b/tests/unit/tools/sandbox/conftest.py index b8d5d05bf5..81e9542370 100644 --- a/tests/unit/tools/sandbox/conftest.py +++ b/tests/unit/tools/sandbox/conftest.py @@ -1,6 +1,6 @@ """Fixtures for sandbox tests.""" -from pathlib import Path # noqa: TC003 — pytest evaluates annotations +from pathlib import Path import pytest diff --git a/tests/unit/tools/sandbox/test_protocol.py b/tests/unit/tools/sandbox/test_protocol.py index 534cf64dd1..62b8812477 100644 --- a/tests/unit/tools/sandbox/test_protocol.py +++ b/tests/unit/tools/sandbox/test_protocol.py @@ -1,7 +1,7 @@ """Tests for SandboxBackend protocol.""" -from collections.abc import Mapping # noqa: TC003 — used at runtime -from pathlib import Path # noqa: TC003 — used at runtime by DockerSandbox +from collections.abc import Mapping +from pathlib import Path import pytest @@ -9,7 +9,7 @@ from ai_company.tools.sandbox.docker_sandbox import DockerSandbox from ai_company.tools.sandbox.protocol import SandboxBackend from ai_company.tools.sandbox.result import SandboxResult -from ai_company.tools.sandbox.subprocess_sandbox import SubprocessSandbox # noqa: TC001 +from ai_company.tools.sandbox.subprocess_sandbox import SubprocessSandbox pytestmark = [pytest.mark.unit, pytest.mark.timeout(30)] diff --git a/uv.lock b/uv.lock index 91de08fd1b..aac3435471 100644 --- a/uv.lock +++ b/uv.lock @@ -2230,13 +2230,17 @@ dependencies = [ { name = "litellm" }, { name = "litestar", extra = ["brotli", "prometheus", "pydantic", "standard", "structlog"] }, { name = "mcp" }, - { name = "mem0ai" }, { name = "pydantic" }, { name = "pyjwt", extra = ["crypto"] }, { name = "pyyaml" }, { name = "structlog" }, ] +[package.optional-dependencies] +mem0 = [ + { name = "mem0ai" }, +] + [package.dev-dependencies] dev = [ { name = "commitizen" }, @@ -2283,12 +2287,13 @@ requires-dist = [ { name = "litellm", specifier = "==1.82.1" }, { name = "litestar", extras = ["brotli", "prometheus", "pydantic", "standard", "structlog"], specifier = "==2.21.1" }, { name = "mcp", specifier = "==1.26.0" }, - { name = "mem0ai", specifier = "==1.0.5" }, + { name = "mem0ai", marker = "extra == 'mem0'", specifier = "==1.0.5" }, { name = "pydantic", specifier = "==2.12.5" }, { name = "pyjwt", extras = ["crypto"], specifier = "==2.11.0" }, { name = "pyyaml", specifier = "==6.0.3" }, { name = "structlog", specifier = "==25.5.0" }, ] +provides-extras = ["mem0"] [package.metadata.requires-dev] dev = [ From 5a4bc0a38564819f59243de66622d1c8674fab76 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 08:07:40 +0100 Subject: [PATCH 08/17] fix: install mem0 optional extra in CI setup action MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit mem0ai was moved to optional-dependencies — CI needs --extra mem0 in the uv sync command to install it for tests. --- .github/actions/setup-python-uv/action.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/actions/setup-python-uv/action.yml b/.github/actions/setup-python-uv/action.yml index 22589d2d1c..4383eb24d2 100644 --- a/.github/actions/setup-python-uv/action.yml +++ b/.github/actions/setup-python-uv/action.yml @@ -25,4 +25,4 @@ runs: shell: bash env: PYTHON_VERSION: ${{ inputs.python-version }} - run: uv sync --frozen --python "$PYTHON_VERSION" + run: uv sync --frozen --python "$PYTHON_VERSION" --extra mem0 From d9e0dee981a851ea2e30132e6b2e8f09b465c180 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 08:10:40 +0100 Subject: [PATCH 09/17] fix: revert mem0ai to required dependency MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit mem0ai is the only memory backend — no reason to make it optional. Reverts the optional-dependencies split and CI extra flag. --- .github/actions/setup-python-uv/action.yml | 2 +- pyproject.toml | 4 +--- src/ai_company/memory/__init__.py | 11 ++++------- uv.lock | 9 ++------- 4 files changed, 8 insertions(+), 18 deletions(-) diff --git a/.github/actions/setup-python-uv/action.yml b/.github/actions/setup-python-uv/action.yml index 4383eb24d2..22589d2d1c 100644 --- a/.github/actions/setup-python-uv/action.yml +++ b/.github/actions/setup-python-uv/action.yml @@ -25,4 +25,4 @@ runs: shell: bash env: PYTHON_VERSION: ${{ inputs.python-version }} - run: uv sync --frozen --python "$PYTHON_VERSION" --extra mem0 + run: uv sync --frozen --python "$PYTHON_VERSION" diff --git a/pyproject.toml b/pyproject.toml index a93f1421d9..f9230763f6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,15 +21,13 @@ dependencies = [ "litellm==1.82.1", "litestar[standard,structlog,pydantic,brotli,prometheus]==2.21.1", "mcp==1.26.0", + "mem0ai==1.0.5", "pydantic==2.12.5", "pyjwt[crypto]==2.11.0", "pyyaml==6.0.3", "structlog==25.5.0", ] -[project.optional-dependencies] -mem0 = ["mem0ai==1.0.5"] - [build-system] requires = ["hatchling==1.29.0"] build-backend = "hatchling.build" diff --git a/src/ai_company/memory/__init__.py b/src/ai_company/memory/__init__.py index 172f4c58be..b2990a46db 100644 --- a/src/ai_company/memory/__init__.py +++ b/src/ai_company/memory/__init__.py @@ -9,13 +9,10 @@ directly. """ -import contextlib - -with contextlib.suppress(ImportError): # mem0ai is optional - from ai_company.memory.backends.mem0 import ( - Mem0EmbedderConfig, - Mem0MemoryBackend, - ) +from ai_company.memory.backends.mem0 import ( + Mem0EmbedderConfig, + Mem0MemoryBackend, +) from ai_company.memory.capabilities import MemoryCapabilities from ai_company.memory.config import ( CompanyMemoryConfig, diff --git a/uv.lock b/uv.lock index aac3435471..91de08fd1b 100644 --- a/uv.lock +++ b/uv.lock @@ -2230,17 +2230,13 @@ dependencies = [ { name = "litellm" }, { name = "litestar", extra = ["brotli", "prometheus", "pydantic", "standard", "structlog"] }, { name = "mcp" }, + { name = "mem0ai" }, { name = "pydantic" }, { name = "pyjwt", extra = ["crypto"] }, { name = "pyyaml" }, { name = "structlog" }, ] -[package.optional-dependencies] -mem0 = [ - { name = "mem0ai" }, -] - [package.dev-dependencies] dev = [ { name = "commitizen" }, @@ -2287,13 +2283,12 @@ requires-dist = [ { name = "litellm", specifier = "==1.82.1" }, { name = "litestar", extras = ["brotli", "prometheus", "pydantic", "standard", "structlog"], specifier = "==2.21.1" }, { name = "mcp", specifier = "==1.26.0" }, - { name = "mem0ai", marker = "extra == 'mem0'", specifier = "==1.0.5" }, + { name = "mem0ai", specifier = "==1.0.5" }, { name = "pydantic", specifier = "==2.12.5" }, { name = "pyjwt", extras = ["crypto"], specifier = "==2.11.0" }, { name = "pyyaml", specifier = "==6.0.3" }, { name = "structlog", specifier = "==25.5.0" }, ] -provides-extras = ["mem0"] [package.metadata.requires-dev] dev = [ From 3c0c85511dbc93efdd266d0a9c4a31e55385eba9 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 08:23:33 +0100 Subject: [PATCH 10/17] fix: address round-4 PR review findings from CodeRabbit, Greptile, and Gemini MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CRITICAL: remove Memory.reset() from disconnect() — it wiped all data - MAJOR: wrap Mem0MemoryBackend init errors in factory as MemoryConfigError - MEDIUM: coerce string scores in normalize_relevance_score - MEDIUM: validate dict structure in validate_add_result - MEDIUM: sanitize source metadata via _coerce_source helper - MEDIUM: update CLAUDE.md — mem0ai is a required dependency, not optional - MEDIUM: extract _coerce_confidence, _coerce_source, _normalize_tags helpers - MINOR: coerce non-string values in extract_publisher - MINOR: add user_id param to mem0_get_result fixture - MINOR: document adapter.py 800-line exemption - Tests: add coverage for string scores, non-dict results, source sanitization, publisher coercion, disconnect-no-reset --- CLAUDE.md | 2 +- .../memory/backends/mem0/adapter.py | 19 ++- .../memory/backends/mem0/mappers.py | 136 +++++++++++++----- src/ai_company/memory/factory.py | 18 ++- tests/unit/memory/backends/mem0/conftest.py | 18 ++- .../unit/memory/backends/mem0/test_adapter.py | 9 ++ .../memory/backends/mem0/test_adapter_crud.py | 21 +-- .../unit/memory/backends/mem0/test_mappers.py | 37 +++++ 8 files changed, 195 insertions(+), 65 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 0a01986a81..020de4beea 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -205,5 +205,5 @@ src/ai_company/ - **Pinned**: all versions use `==` in `pyproject.toml` - **Groups**: `test` (pytest + plugins), `dev` (includes test + ruff, mypy, pre-commit, commitizen) -- **Optional**: `mem0ai` (Mem0 memory backend — only needed when `backend: "mem0"` is configured) +- **Required**: `mem0ai` (Mem0 memory backend — the default and currently only backend) - **Install**: `uv sync` installs everything (dev group is default) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index e59a15abc6..de75f75346 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -9,6 +9,12 @@ All methods re-raise ``builtins.MemoryError`` and ``RecursionError`` immediately without wrapping, to avoid masking system-level failures. + +Note: This file exceeds the 800-line guideline because the single +``Mem0MemoryBackend`` class implements three protocols cohesively +(``MemoryBackend``, ``MemoryCapabilities``, ``SharedKnowledgeStore``). +Splitting would fragment the unified client lifecycle and connection +guard logic. """ import asyncio @@ -179,19 +185,10 @@ async def connect(self) -> None: async def disconnect(self) -> None: """Close the Mem0 connection. - Attempts to close the underlying client resources before - releasing the reference. Safe to call even if not connected. + Releases the client reference so the garbage collector can + reclaim resources. Safe to call even if not connected. """ logger.info(MEMORY_BACKEND_DISCONNECTING, backend="mem0") - if self._client is not None: - try: - await asyncio.to_thread(self._client.reset) - except Exception: - logger.debug( - MEMORY_BACKEND_DISCONNECTING, - backend="mem0", - note="reset failed during disconnect, ignoring", - ) self._client = None self._connected = False logger.info(MEMORY_BACKEND_DISCONNECTED, backend="mem0") diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index b25158dc7c..28b313d595 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -86,18 +86,90 @@ def parse_mem0_datetime(raw: str | None) -> AwareDatetime | None: return dt -def normalize_relevance_score(score: float | None) -> float | None: - """Clamp a relevance score to [0.0, 1.0]. +def normalize_relevance_score(score: Any) -> float | None: + """Coerce and clamp a relevance score to [0.0, 1.0]. Args: - score: Raw score from Mem0 (may exceed bounds). + score: Raw score from Mem0 (may be ``None``, numeric, + or a string representation of a number). Returns: - Clamped score, or ``None`` if input is ``None``. + Clamped score, or ``None`` if input is ``None`` or + cannot be converted to a float. """ if score is None: return None - return max(0.0, min(1.0, score)) + try: + numeric = float(score) + except ValueError, TypeError: + logger.warning( + MEMORY_MODEL_INVALID, + field="score", + raw_value=score, + reason="non-numeric relevance score, returning None", + ) + return None + return max(0.0, min(1.0, numeric)) + + +def _coerce_confidence(raw_metadata: dict[str, Any]) -> float: + """Extract and clamp confidence from Mem0 metadata. + + Returns a float in [0.0, 1.0], defaulting to 1.0 on failure. + """ + raw = raw_metadata.get(f"{_PREFIX}confidence", 1.0) + try: + value = float(raw) + except ValueError, TypeError: + logger.warning( + MEMORY_MODEL_INVALID, + field="confidence", + raw_value=raw, + reason="non-numeric confidence, defaulting to 1.0", + ) + return 1.0 + return max(0.0, min(1.0, value)) + + +def _coerce_source(raw_metadata: dict[str, Any]) -> str | None: + """Extract and sanitize the source field from Mem0 metadata. + + Returns ``None`` if the value is missing, non-string, or blank. + """ + raw = raw_metadata.get(f"{_PREFIX}source") + if raw is None: + return None + coerced = str(raw).strip() + if not coerced: + logger.debug( + MEMORY_MODEL_INVALID, + field="source", + raw_value=raw, + reason="blank source after coercion, returning None", + ) + return None + return coerced + + +def _normalize_tags( + raw_metadata: dict[str, Any], +) -> tuple[NotBlankStr, ...]: + """Extract and normalize tags from Mem0 metadata. + + Handles string, list, tuple, and unexpected types gracefully. + """ + raw_tags = raw_metadata.get(f"{_PREFIX}tags", ()) + if isinstance(raw_tags, str): + raw_tags = [raw_tags] + elif not isinstance(raw_tags, (list, tuple)): + logger.debug( + MEMORY_MODEL_INVALID, + field="tags", + raw_value=type(raw_tags).__name__, + reason="unexpected tags type, ignoring", + ) + raw_tags = () + return tuple(NotBlankStr(str(t)) for t in raw_tags if t and str(t).strip()) def parse_mem0_metadata( @@ -139,32 +211,9 @@ def parse_mem0_metadata( else: category = MemoryCategory.WORKING - raw_confidence = raw_metadata.get(f"{_PREFIX}confidence", 1.0) - try: - confidence = float(raw_confidence) - except ValueError, TypeError: - logger.warning( - MEMORY_MODEL_INVALID, - field="confidence", - raw_value=raw_confidence, - reason="non-numeric confidence, defaulting to 1.0", - ) - confidence = 1.0 - confidence = max(0.0, min(1.0, confidence)) - source = raw_metadata.get(f"{_PREFIX}source") - raw_tags = raw_metadata.get(f"{_PREFIX}tags", ()) - if isinstance(raw_tags, str): - raw_tags = [raw_tags] - elif not isinstance(raw_tags, (list, tuple)): - logger.debug( - MEMORY_MODEL_INVALID, - field="tags", - raw_value=type(raw_tags).__name__, - reason="unexpected tags type, ignoring", - ) - raw_tags = () - tags = tuple(NotBlankStr(str(t)) for t in raw_tags if t and str(t).strip()) - + confidence = _coerce_confidence(raw_metadata) + source = _coerce_source(raw_metadata) + tags = _normalize_tags(raw_metadata) expires_at = parse_mem0_datetime( raw_metadata.get(f"{_PREFIX}expires_at"), ) @@ -337,11 +386,11 @@ def apply_post_filters( # ── Adapter helpers ────────────────────────────────────────────────── -def validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: +def validate_add_result(result: Any, *, context: str) -> NotBlankStr: """Extract and validate the memory ID from a Mem0 ``add`` result. Args: - result: Raw result dict from ``Memory.add()``. + result: Raw result from ``Memory.add()`` (expected dict). context: Human-readable context for error messages (e.g. ``"store"`` or ``"shared publish"``). @@ -351,12 +400,24 @@ def validate_add_result(result: dict[str, Any], *, context: str) -> NotBlankStr: Raises: MemoryStoreError: If the result is missing or malformed. """ + if not isinstance(result, dict): + msg = ( + f"Mem0 add returned unexpected type for {context}: {type(result).__name__}" + ) + logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) + raise MemoryStoreError(msg) results_list = result.get("results") if not isinstance(results_list, list) or not results_list: msg = f"Mem0 add returned no results for {context}" logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) raise MemoryStoreError(msg) first = results_list[0] + if not isinstance(first, dict): + msg = ( + f"Mem0 add result item is not a dict for {context}: {type(first).__name__}" + ) + logger.warning(MEMORY_ENTRY_STORE_FAILED, context=context, error=msg) + raise MemoryStoreError(msg) raw_id = first.get("id") if raw_id is None or not str(raw_id).strip(): msg = ( @@ -397,12 +458,13 @@ def extract_publisher(raw: dict[str, Any]) -> str | None: """Extract the publisher agent ID from a shared memory dict. Returns ``None`` if the publisher key is missing, non-dict - metadata, or the value is blank after stripping. + metadata, or the value is blank after coercion and stripping. """ metadata = raw.get("metadata", {}) if not metadata or not isinstance(metadata, dict): return None - value: str | None = metadata.get(_PUBLISHER_KEY) - if value is not None and not str(value).strip(): + value = metadata.get(_PUBLISHER_KEY) + if value is None: return None - return value + coerced = str(value).strip() + return coerced or None diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 53bc44d61c..73359d7fac 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -91,10 +91,20 @@ def create_memory_backend( error=msg, ) raise MemoryConfigError(msg) from exc - backend = Mem0MemoryBackend( - mem0_config=mem0_config, - max_memories_per_agent=config.options.max_memories_per_agent, - ) + try: + backend = Mem0MemoryBackend( + mem0_config=mem0_config, + max_memories_per_agent=config.options.max_memories_per_agent, + ) + except Exception as exc: + msg = f"Failed to create Mem0 backend: {exc}" + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="backend_init_failed", + error=msg, + ) + raise MemoryConfigError(msg) from exc logger.info( MEMORY_BACKEND_CREATED, backend="mem0", diff --git a/tests/unit/memory/backends/mem0/conftest.py b/tests/unit/memory/backends/mem0/conftest.py index 30b2c20171..850b842804 100644 --- a/tests/unit/memory/backends/mem0/conftest.py +++ b/tests/unit/memory/backends/mem0/conftest.py @@ -82,9 +82,18 @@ def mem0_search_result( return {"results": items} -def mem0_get_result(memory_id: str = "mem-001") -> dict[str, Any]: - """Build a typical Mem0 get() return value.""" - return { +def mem0_get_result( + memory_id: str = "mem-001", + *, + user_id: str | None = None, +) -> dict[str, Any]: + """Build a typical Mem0 get() return value. + + Args: + memory_id: Memory identifier. + user_id: Optional owner ``user_id`` for ownership tests. + """ + result: dict[str, Any] = { "id": memory_id, "memory": "stored content", "created_at": "2026-03-12T10:00:00+00:00", @@ -94,6 +103,9 @@ def mem0_get_result(memory_id: str = "mem-001") -> dict[str, Any]: "_synthorg_confidence": 1.0, }, } + if user_id is not None: + result["user_id"] = user_id + return result def make_store_request( diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index b65b164baa..ed31b61021 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -168,6 +168,15 @@ async def test_disconnect(self, backend: Mem0MemoryBackend) -> None: assert backend.is_connected is False assert backend._client is None + async def test_disconnect_does_not_call_reset( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """disconnect() must never call reset() — that wipes all data.""" + await backend.disconnect() + mock_client.reset.assert_not_called() + async def test_disconnect_when_not_connected( self, mem0_config: Mem0BackendConfig, diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index 3fbcccb859..50ebeebf5f 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -309,9 +309,10 @@ async def test_get_ownership_mismatch_returns_none( mock_client: MagicMock, ) -> None: """get() returns None when user_id doesn't match agent_id.""" - result = mem0_get_result("mem-001") - result["user_id"] = "other-agent" - mock_client.get.return_value = result + mock_client.get.return_value = mem0_get_result( + "mem-001", + user_id="other-agent", + ) entry = await backend.get("test-agent-001", "mem-001") assert entry is None @@ -384,9 +385,10 @@ async def test_delete_shared_namespace_entry_raises( mock_client: MagicMock, ) -> None: """delete() rejects entries belonging to the shared namespace.""" - result = mem0_get_result("mem-001") - result["user_id"] = _SHARED_NAMESPACE - mock_client.get.return_value = result + mock_client.get.return_value = mem0_get_result( + "mem-001", + user_id=_SHARED_NAMESPACE, + ) with pytest.raises(MemoryStoreError, match="shared namespace"): await backend.delete("test-agent-001", "mem-001") @@ -397,9 +399,10 @@ async def test_delete_ownership_mismatch_raises( mock_client: MagicMock, ) -> None: """delete() rejects when user_id doesn't match agent_id.""" - result = mem0_get_result("mem-001") - result["user_id"] = "other-agent" - mock_client.get.return_value = result + mock_client.get.return_value = mem0_get_result( + "mem-001", + user_id="other-agent", + ) with pytest.raises(MemoryStoreError, match="cannot delete"): await backend.delete("test-agent-001", "mem-001") diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index 00c82864c4..b8c203ad81 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -8,6 +8,7 @@ from ai_company.core.enums import MemoryCategory from ai_company.memory.backends.mem0.mappers import ( _PREFIX, + _PUBLISHER_KEY, apply_post_filters, build_mem0_metadata, extract_category, @@ -139,6 +140,15 @@ def test_boundaries(self) -> None: assert normalize_relevance_score(0.0) == 0.0 assert normalize_relevance_score(1.0) == 1.0 + def test_string_score_coerced(self) -> None: + assert normalize_relevance_score("0.82") == 0.82 + + def test_non_numeric_string_returns_none(self) -> None: + assert normalize_relevance_score("not-a-number") is None + + def test_non_numeric_type_returns_none(self) -> None: + assert normalize_relevance_score(object()) is None + @pytest.mark.unit class TestParseMem0Metadata: @@ -185,6 +195,16 @@ def test_empty_tags_filtered(self) -> None: _category, metadata, _expires = parse_mem0_metadata(raw) assert metadata.tags == ("valid", "also-valid") + def test_blank_source_returns_none(self) -> None: + raw = {f"{_PREFIX}source": " "} + _category, metadata, _expires = parse_mem0_metadata(raw) + assert metadata.source is None + + def test_non_string_source_coerced(self) -> None: + raw = {f"{_PREFIX}source": 42} + _category, metadata, _expires = parse_mem0_metadata(raw) + assert metadata.source == "42" + @pytest.mark.unit class TestMem0ResultToEntry: @@ -438,6 +458,15 @@ def test_numeric_id_coerced_to_string(self) -> None: memory_id = validate_add_result(result, context="test") assert memory_id == "42" + def test_non_dict_result_raises(self) -> None: + with pytest.raises(MemoryStoreError, match="unexpected type"): + validate_add_result("not-a-dict", context="test") # type: ignore[arg-type] + + def test_non_dict_first_item_raises(self) -> None: + result: dict[str, Any] = {"results": ["not-a-dict"]} + with pytest.raises(MemoryStoreError, match="not a dict"): + validate_add_result(result, context="test") + @pytest.mark.unit class TestExtractCategory: @@ -505,3 +534,11 @@ def test_list_metadata_returns_none(self) -> None: def test_string_metadata_returns_none(self) -> None: raw: dict[str, Any] = {"metadata": "oops"} assert extract_publisher(raw) is None + + def test_numeric_publisher_coerced_to_string(self) -> None: + raw: dict[str, Any] = {"metadata": {_PUBLISHER_KEY: 42}} + assert extract_publisher(raw) == "42" + + def test_blank_publisher_returns_none(self) -> None: + raw: dict[str, Any] = {"metadata": {_PUBLISHER_KEY: " "}} + assert extract_publisher(raw) is None From 71469c2272449ea04fc3c8e00a6b29f9f192abc1 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 08:43:28 +0100 Subject: [PATCH 11/17] fix: address round-5 PR review findings and add missing test coverage - Remove unused noqa A004 directive (ruff auto-fix) - Fix import sorting in adapter.py (ruff auto-fix) - Add connect() idempotency test (verifies no-op when already connected) - Add _validate_agent_id error_cls tests for read ops (retrieve, get, count) - Add expires_at filtering tests in apply_post_filters - Add config validation tests using model_construct to bypass Pydantic - Fix factory init test patch target to match local import --- .../memory/backends/mem0/adapter.py | 31 ++++++++++---- .../memory/backends/mem0/mappers.py | 12 ++++-- .../unit/memory/backends/mem0/test_adapter.py | 15 +++++++ .../memory/backends/mem0/test_adapter_crud.py | 27 ++++++++++++ .../unit/memory/backends/mem0/test_config.py | 38 +++++++++++++++++ .../unit/memory/backends/mem0/test_mappers.py | 35 +++++++++++++++- tests/unit/memory/test_factory.py | 42 ++++++++++++++++++- 7 files changed, 187 insertions(+), 13 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index de75f75346..8c78529a62 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -43,6 +43,9 @@ MemoryRetrievalError, MemoryStoreError, ) +from ai_company.memory.errors import ( + MemoryError as DomainMemoryError, +) from ai_company.observability import get_logger if TYPE_CHECKING: @@ -147,11 +150,14 @@ async def connect(self) -> None: """Establish connection to Mem0. Creates the Mem0 ``Memory`` client with embedded Qdrant. + Idempotent — returns immediately if already connected. Raises: MemoryConnectionError: If Mem0 is not installed or initialization fails. """ + if self._connected: + return logger.info(MEMORY_BACKEND_CONNECTING, backend="mem0") try: from mem0 import Memory # noqa: PLC0415 @@ -288,12 +294,23 @@ def _require_connected(self) -> None: msg = "Not connected — call connect() first" raise MemoryConnectionError(msg) - def _validate_agent_id(self, agent_id: NotBlankStr) -> None: + def _validate_agent_id( + self, + agent_id: NotBlankStr, + *, + error_cls: type[DomainMemoryError] = MemoryStoreError, + ) -> None: """Reject the reserved shared namespace as an agent ID. + Args: + agent_id: Agent identifier to validate. + error_cls: Error class to raise on rejection — defaults to + ``MemoryStoreError`` for write ops, pass + ``MemoryRetrievalError`` for read ops. + Raises: - MemoryStoreError: If ``agent_id`` collides with the - reserved ``_SHARED_NAMESPACE``. + DomainMemoryError: (subclass per ``error_cls``) If + ``agent_id`` collides with ``_SHARED_NAMESPACE``. """ if str(agent_id) == _SHARED_NAMESPACE: logger.warning( @@ -305,7 +322,7 @@ def _validate_agent_id(self, agent_id: NotBlankStr) -> None: f"agent_id must not be the reserved shared namespace: " f"{_SHARED_NAMESPACE!r}" ) - raise MemoryStoreError(msg) + raise error_cls(msg) # ── CRUD Operations ─────────────────────────────────────────── @@ -391,7 +408,7 @@ async def retrieve( MemoryRetrievalError: If the retrieval fails. """ self._require_connected() - self._validate_agent_id(agent_id) + self._validate_agent_id(agent_id, error_cls=MemoryRetrievalError) try: if query.text is not None: kwargs = query_to_mem0_search_args(str(agent_id), query) @@ -453,7 +470,7 @@ async def get( MemoryRetrievalError: If the backend query fails. """ self._require_connected() - self._validate_agent_id(agent_id) + self._validate_agent_id(agent_id, error_cls=MemoryRetrievalError) try: raw = await asyncio.to_thread(self._client.get, str(memory_id)) if raw is None: @@ -625,7 +642,7 @@ async def count( MemoryRetrievalError: If the count query fails. """ self._require_connected() - self._validate_agent_id(agent_id) + self._validate_agent_id(agent_id, error_cls=MemoryRetrievalError) try: raw_result = await asyncio.to_thread( self._client.get_all, diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 28b313d595..83982af063 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -351,10 +351,11 @@ def apply_post_filters( ) -> tuple[MemoryEntry, ...]: """Apply post-retrieval filters that Mem0 cannot handle natively. - Filters by category, tags, time range, and minimum relevance. - Entries with ``relevance_score=None`` (e.g. from ``get_all``) - are never excluded by ``min_relevance`` — the filter only - applies when a score is present. + Filters expired entries, then applies category, tags, time range, + and minimum relevance filters. Entries with + ``relevance_score=None`` (e.g. from ``get_all``) are never + excluded by ``min_relevance`` — the filter only applies when a + score is present. Args: entries: Raw entries from Mem0. @@ -363,8 +364,11 @@ def apply_post_filters( Returns: Filtered entries (order preserved). """ + now = datetime.now(UTC) result: list[MemoryEntry] = [] for entry in entries: + if entry.expires_at is not None and entry.expires_at <= now: + continue if query.categories and entry.category not in query.categories: continue if query.tags and not all(tag in entry.metadata.tags for tag in query.tags): diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index ed31b61021..e345764005 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -148,6 +148,21 @@ async def test_connect_success( assert b.is_connected is True + async def test_connect_idempotent_when_already_connected( + self, + backend: Mem0MemoryBackend, + ) -> None: + """connect() is a no-op when already connected.""" + assert backend.is_connected is True + # Calling connect() again should not re-create client + original_client = backend._client + with patch( + "ai_company.memory.backends.mem0.adapter.asyncio.to_thread", + ) as mock_to_thread: + await backend.connect() + mock_to_thread.assert_not_called() + assert backend._client is original_client + async def test_connect_failure_raises( self, mem0_config: Mem0BackendConfig, diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index 50ebeebf5f..411bfa59bf 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -236,6 +236,17 @@ async def test_retrieve_reraises_memory_error( MemoryQuery(text="test"), ) + async def test_retrieve_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """retrieve() rejects the shared namespace with MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="reserved shared namespace"): + await backend.retrieve( + _SHARED_NAMESPACE, + MemoryQuery(text="test"), + ) + async def test_retrieve_invalid_entry_raises( self, backend: Mem0MemoryBackend, @@ -303,6 +314,14 @@ async def test_get_reraises_memory_error( with pytest.raises(MemoryError): await backend.get("test-agent-001", "mem-001") + async def test_get_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """get() rejects the shared namespace with MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="reserved shared namespace"): + await backend.get(_SHARED_NAMESPACE, "mem-001") + async def test_get_ownership_mismatch_returns_none( self, backend: Mem0MemoryBackend, @@ -487,6 +506,14 @@ async def test_count_with_invalid_category_in_data( # "bogus_category" defaults to WORKING assert count == 1 + async def test_count_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """count() rejects the shared namespace with MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="reserved shared namespace"): + await backend.count(_SHARED_NAMESPACE) + async def test_count_exception_wraps( self, backend: Mem0MemoryBackend, diff --git a/tests/unit/memory/backends/mem0/test_config.py b/tests/unit/memory/backends/mem0/test_config.py index 1a4e309f46..a85326a0ea 100644 --- a/tests/unit/memory/backends/mem0/test_config.py +++ b/tests/unit/memory/backends/mem0/test_config.py @@ -189,3 +189,41 @@ def test_passes_embedder_through(self) -> None: assert mem0_config.embedder.provider == "test-provider" assert mem0_config.embedder.model == "test-model-xl" assert mem0_config.embedder.dims == 4096 + + def test_rejects_unsupported_vector_store(self) -> None: + storage = MemoryStorageConfig.model_construct( + vector_store="chroma", + history_store="sqlite", + data_dir="/data/memory", + ) + company_config = CompanyMemoryConfig.model_construct( + backend="mem0", + storage=storage, + ) + with pytest.raises(ValueError, match="qdrant"): + build_config_from_company_config( + company_config, + embedder=_embedder(), + ) + + def test_rejects_unsupported_history_store(self) -> None: + company_config = CompanyMemoryConfig( + backend="mem0", + storage=MemoryStorageConfig(history_store="postgresql"), + ) + with pytest.raises(ValueError, match="sqlite"): + build_config_from_company_config( + company_config, + embedder=_embedder(), + ) + + def test_accepts_qdrant_external(self) -> None: + company_config = CompanyMemoryConfig( + backend="mem0", + storage=MemoryStorageConfig(vector_store="qdrant-external"), + ) + mem0_config = build_config_from_company_config( + company_config, + embedder=_embedder(), + ) + assert mem0_config is not None diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index b8c203ad81..3cfbd48c8a 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -41,6 +41,7 @@ def _make_entry( # noqa: PLR0913 tags: tuple[str, ...] = (), relevance_score: float | None = None, created_at: datetime | None = None, + expires_at: datetime | None = None, ) -> MemoryEntry: """Helper to build a MemoryEntry for tests.""" now = created_at or datetime.now(UTC) @@ -52,6 +53,7 @@ def _make_entry( # noqa: PLR0913 metadata=MemoryMetadata(tags=tags), created_at=now, relevance_score=relevance_score, + expires_at=expires_at, ) @@ -385,6 +387,37 @@ def test_until_filter_exclusive(self) -> None: assert len(result) == 1 assert result[0].id == "m1" + def test_expired_entries_excluded(self) -> None: + """Entries with expires_at in the past are filtered out.""" + now = datetime.now(UTC) + past = now - timedelta(days=7) + entries = ( + _make_entry( + memory_id="m1", + created_at=past, + expires_at=now - timedelta(hours=1), + ), + _make_entry( + memory_id="m2", + created_at=past, + expires_at=now + timedelta(hours=1), + ), + _make_entry(memory_id="m3"), # no expires_at + ) + query = MemoryQuery() + result = apply_post_filters(entries, query) + assert len(result) == 2 + assert {e.id for e in result} == {"m2", "m3"} + + def test_exactly_expired_entry_excluded(self) -> None: + """Entry with expires_at == now is excluded (<=).""" + now = datetime.now(UTC) + past = now - timedelta(days=7) + entries = (_make_entry(memory_id="m1", created_at=past, expires_at=now),) + query = MemoryQuery() + result = apply_post_filters(entries, query) + assert len(result) == 0 + def test_combined_filters(self) -> None: now = datetime.now(UTC) entries = ( @@ -460,7 +493,7 @@ def test_numeric_id_coerced_to_string(self) -> None: def test_non_dict_result_raises(self) -> None: with pytest.raises(MemoryStoreError, match="unexpected type"): - validate_add_result("not-a-dict", context="test") # type: ignore[arg-type] + validate_add_result("not-a-dict", context="test") def test_non_dict_first_item_raises(self) -> None: result: dict[str, Any] = {"results": ["not-a-dict"]} diff --git a/tests/unit/memory/test_factory.py b/tests/unit/memory/test_factory.py index 147f514153..0a9b35d744 100644 --- a/tests/unit/memory/test_factory.py +++ b/tests/unit/memory/test_factory.py @@ -1,11 +1,17 @@ """Tests for memory backend factory.""" +from unittest.mock import patch + import pytest from pydantic import ValidationError from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend from ai_company.memory.backends.mem0.config import Mem0EmbedderConfig -from ai_company.memory.config import CompanyMemoryConfig, MemoryOptionsConfig +from ai_company.memory.config import ( + CompanyMemoryConfig, + MemoryOptionsConfig, + MemoryStorageConfig, +) from ai_company.memory.errors import MemoryConfigError from ai_company.memory.factory import create_memory_backend @@ -52,3 +58,37 @@ def test_mem0_wrong_embedder_type_raises(self) -> None: config = CompanyMemoryConfig(backend="mem0") with pytest.raises(MemoryConfigError, match="must be a Mem0EmbedderConfig"): create_memory_backend(config, embedder="not-a-config") # type: ignore[arg-type] + + def test_config_build_error_wraps_as_memory_config_error(self) -> None: + """ValueError from build_config_from_company_config wraps.""" + storage = MemoryStorageConfig.model_construct( + vector_store="chroma", + history_store="sqlite", + data_dir="/data/memory", + ) + config = CompanyMemoryConfig.model_construct( + backend="mem0", + storage=storage, + ) + with pytest.raises(MemoryConfigError, match="Invalid Mem0 configuration"): + create_memory_backend(config, embedder=_test_embedder()) + + def test_backend_init_error_wraps_as_memory_config_error(self) -> None: + """Exception from Mem0MemoryBackend() constructor wraps.""" + config = CompanyMemoryConfig(backend="mem0") + with ( + patch( + "ai_company.memory.backends.mem0.Mem0MemoryBackend", + side_effect=RuntimeError("init boom"), + ), + pytest.raises(MemoryConfigError, match="Failed to create Mem0"), + ): + create_memory_backend(config, embedder=_test_embedder()) + + def test_unknown_backend_bypassing_validation_raises(self) -> None: + """Defensive guard when model_construct bypasses validation.""" + config = CompanyMemoryConfig.model_construct( + backend="nonexistent", + ) + with pytest.raises(MemoryConfigError, match="Unknown memory backend"): + create_memory_backend(config, embedder=_test_embedder()) From 3adf18badbb6f95191c2db6794497acab5ff19ad Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 09:24:45 +0100 Subject: [PATCH 12/17] fix: address round-6 PR review findings across adapter, docs, tests, and factory MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add Mem0Client protocol (TYPE_CHECKING) for structural typing documentation - Add asyncio.Lock + double-check locking to prevent concurrent connect() - Expand docstrings: adapter, config, design spec protocol snippets - Replace inline generator in search_shared with explicit loop + publisher fallback logging - Narrow factory exception catch to (ValueError, ValidationError) - Add YAML enum value annotations in design spec - Update CLAUDE.md logging event domains - Remove unused event constants (MEMORY_BACKEND_NOT_IMPLEMENTED, MEMORY_CAPABILITY_UNSUPPORTED) - Rename _PUBLISHER_KEY → PUBLISHER_KEY (public constant) - Add tests: RecursionError propagation, _validate_mem0_result, _coerce_confidence, _normalize_tags, count truncation warning, category post-filter, publisher fallback --- CLAUDE.md | 2 +- docs/design/memory.md | 45 ++-- .../memory/backends/mem0/adapter.py | 219 +++++++++++++----- src/ai_company/memory/backends/mem0/config.py | 23 +- .../memory/backends/mem0/mappers.py | 30 ++- src/ai_company/memory/factory.py | 4 +- src/ai_company/memory/models.py | 2 +- src/ai_company/observability/events/memory.py | 5 - tests/integration/memory/test_mem0_backend.py | 10 +- .../memory/backends/mem0/test_adapter_crud.py | 110 +++++++++ .../backends/mem0/test_adapter_shared.py | 91 +++++++- .../unit/memory/backends/mem0/test_mappers.py | 66 +++++- tests/unit/memory/test_factory.py | 28 ++- tests/unit/observability/test_events.py | 2 - 14 files changed, 509 insertions(+), 128 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 020de4beea..201443436a 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -127,7 +127,7 @@ src/ai_company/ - **Every module** with business logic MUST have: `from ai_company.observability import get_logger` then `logger = get_logger(__name__)` - **Never** use `import logging` / `logging.getLogger()` / `print()` in application code - **Variable name**: always `logger` (not `_logger`, not `log`) -- **Event names**: always use constants from the domain-specific module under `ai_company.observability.events` (e.g. `PROVIDER_CALL_START` from `events.provider`, `BUDGET_RECORD_ADDED` from `events.budget`, `CFO_ANOMALY_DETECTED` from `events.cfo`, `CONFLICT_DETECTED` from `events.conflict`, `MEETING_STARTED` from `events.meeting`, `CLASSIFICATION_START` from `events.classification`, `CONSOLIDATION_START` from `events.consolidation`, `ORG_MEMORY_QUERY_START` from `events.org_memory`, `API_REQUEST_STARTED` from `events.api`, `CODE_RUNNER_EXECUTE_START` from `events.code_runner`, `DOCKER_EXECUTE_START` from `events.docker`, `MCP_INVOKE_START` from `events.mcp`, `SECURITY_EVALUATE_START` from `events.security`, `HR_HIRING_REQUEST_CREATED` from `events.hr`, `PERF_METRIC_RECORDED` from `events.performance`, `TRUST_EVALUATE_START` from `events.trust`, `PROMOTION_EVALUATE_START` from `events.promotion`, `PROMPT_BUILD_START` from `events.prompt`, `MEMORY_RETRIEVAL_START` from `events.memory`, `AUTONOMY_ACTION_AUTO_APPROVED` from `events.autonomy`, `TIMEOUT_POLICY_EVALUATED` from `events.timeout`, `PERSISTENCE_AUDIT_ENTRY_SAVED` from `events.persistence`, `TASK_ENGINE_STARTED` from `events.task_engine`, `COORDINATION_STARTED` from `events.coordination`). Import directly: `from ai_company.observability.events. import EVENT_CONSTANT` +- **Event names**: always use constants from the domain-specific module under `ai_company.observability.events` (e.g. `PROVIDER_CALL_START` from `events.provider`, `BUDGET_RECORD_ADDED` from `events.budget`, `CFO_ANOMALY_DETECTED` from `events.cfo`, `CONFLICT_DETECTED` from `events.conflict`, `MEETING_STARTED` from `events.meeting`, `CLASSIFICATION_START` from `events.classification`, `CONSOLIDATION_START` from `events.consolidation`, `ORG_MEMORY_QUERY_START` from `events.org_memory`, `API_REQUEST_STARTED` from `events.api`, `CODE_RUNNER_EXECUTE_START` from `events.code_runner`, `DOCKER_EXECUTE_START` from `events.docker`, `MCP_INVOKE_START` from `events.mcp`, `SECURITY_EVALUATE_START` from `events.security`, `HR_HIRING_REQUEST_CREATED` from `events.hr`, `PERF_METRIC_RECORDED` from `events.performance`, `TRUST_EVALUATE_START` from `events.trust`, `PROMOTION_EVALUATE_START` from `events.promotion`, `PROMPT_BUILD_START` from `events.prompt`, `MEMORY_RETRIEVAL_START` from `events.memory`, `AUTONOMY_ACTION_AUTO_APPROVED` from `events.autonomy`, `TIMEOUT_POLICY_EVALUATED` from `events.timeout`, `PERSISTENCE_AUDIT_ENTRY_SAVED` from `events.persistence`, `TASK_ENGINE_STARTED` from `events.task_engine`, `COORDINATION_STARTED` from `events.coordination`, `COMMUNICATION_DISPATCH_START` from `events.communication`, `COMPANY_STARTED` from `events.company`, `CONFIG_LOADED` from `events.config`, `CORRELATION_ID_CREATED` from `events.correlation`, `DECOMPOSITION_STARTED` from `events.decomposition`, `DELEGATION_STARTED` from `events.delegation`, `EXECUTION_LOOP_STARTED` from `events.execution`, `GIT_OPERATION_START` from `events.git`, `PARALLEL_EXECUTION_STARTED` from `events.parallel`, `PERSONALITY_LOADED` from `events.personality`, `QUOTA_CHECKED` from `events.quota`, `ROLE_ASSIGNED` from `events.role`, `ROUTING_STARTED` from `events.routing`, `SANDBOX_EXECUTE_START` from `events.sandbox`, `TASK_CREATED` from `events.task`, `TASK_ASSIGNMENT_STARTED` from `events.task_assignment`, `TASK_ROUTING_STARTED` from `events.task_routing`, `TEMPLATE_LOADED` from `events.template`, `TOOL_INVOKE_START` from `events.tool`, `WORKSPACE_CREATED` from `events.workspace`). Import directly: `from ai_company.observability.events. import EVENT_CONSTANT` - **Structured kwargs**: always `logger.info(EVENT, key=value)` — never `logger.info("msg %s", val)` - **All error paths** must log at WARNING or ERROR with context before raising - **All state transitions** must log at INFO diff --git a/docs/design/memory.md b/docs/design/memory.md index 1d494fbe7a..7972a78396 100644 --- a/docs/design/memory.md +++ b/docs/design/memory.md @@ -29,8 +29,8 @@ configuration without modifying application code. | context | decisions| learned | | +----------+----------+-----------+---------------+ | Storage Backend | -| SQLite / PostgreSQL / File-based | -| + Mem0 (initial, implemented) / Custom (future)| +| Mem0 (initial, implemented) / Custom (future) | +| Qdrant (embedded) + SQLite history | | See Decision Log | +-------------------------------------------------+ ``` @@ -60,8 +60,8 @@ Memory persistence is configurable per agent, from no persistence to fully persi ```yaml memory: - level: "persistent" # none, session, project, persistent (default: session) - backend: "mem0" # mem0, custom, cognee, graphiti (future) -- see Decision Log + level: "persistent" # none | session | project | persistent (default: session) + backend: "mem0" # mem0 | custom | cognee | graphiti (future) -- see Decision Log storage: data_dir: "/data/memory" # mounted Docker volume path vector_store: "qdrant" # qdrant (embedded), qdrant-external, etc. @@ -206,11 +206,21 @@ class MemoryBackend(Protocol): @property def backend_name(self) -> NotBlankStr: ... - async def store(self, agent_id: NotBlankStr, request: MemoryStoreRequest) -> NotBlankStr: ... - async def retrieve(self, agent_id: NotBlankStr, query: MemoryQuery) -> tuple[MemoryEntry, ...]: ... - async def get(self, agent_id: NotBlankStr, memory_id: NotBlankStr) -> MemoryEntry | None: ... - async def delete(self, agent_id: NotBlankStr, memory_id: NotBlankStr) -> bool: ... - async def count(self, agent_id: NotBlankStr, *, category: MemoryCategory | None = None) -> int: ... + async def store(self, agent_id: NotBlankStr, request: MemoryStoreRequest) -> NotBlankStr: + """Raises: MemoryConnectionError, MemoryStoreError.""" + ... + async def retrieve(self, agent_id: NotBlankStr, query: MemoryQuery) -> tuple[MemoryEntry, ...]: + """Raises: MemoryConnectionError, MemoryRetrievalError.""" + ... + async def get(self, agent_id: NotBlankStr, memory_id: NotBlankStr) -> MemoryEntry | None: + """Raises: MemoryConnectionError, MemoryRetrievalError.""" + ... + async def delete(self, agent_id: NotBlankStr, memory_id: NotBlankStr) -> bool: + """Raises: MemoryConnectionError, MemoryStoreError.""" + ... + async def count(self, agent_id: NotBlankStr, *, category: MemoryCategory | None = None) -> int: + """Raises: MemoryConnectionError, MemoryRetrievalError.""" + ... ``` ### MemoryCapabilities Protocol @@ -248,9 +258,15 @@ clean. class SharedKnowledgeStore(Protocol): """Cross-agent shared knowledge operations.""" - async def publish(self, agent_id: NotBlankStr, request: MemoryStoreRequest) -> NotBlankStr: ... - async def search_shared(self, query: MemoryQuery, *, exclude_agent: NotBlankStr | None = None) -> tuple[MemoryEntry, ...]: ... - async def retract(self, agent_id: NotBlankStr, memory_id: NotBlankStr) -> bool: ... + async def publish(self, agent_id: NotBlankStr, request: MemoryStoreRequest) -> NotBlankStr: + """Raises: MemoryConnectionError, MemoryStoreError.""" + ... + async def search_shared(self, query: MemoryQuery, *, exclude_agent: NotBlankStr | None = None) -> tuple[MemoryEntry, ...]: + """Raises: MemoryConnectionError, MemoryRetrievalError.""" + ... + async def retract(self, agent_id: NotBlankStr, memory_id: NotBlankStr) -> bool: + """Raises: MemoryConnectionError, MemoryStoreError.""" + ... ``` ### Error Hierarchy @@ -294,8 +310,9 @@ memory: Configuration is modeled by `CompanyMemoryConfig` (top-level), `MemoryStorageConfig` (storage paths/backends), and `MemoryOptionsConfig` (behaviour tuning). All are frozen -Pydantic models. The `create_memory_backend(config)` factory returns an isolated -`MemoryBackend` instance per company. +Pydantic models. The `create_memory_backend(config, *, embedder=...)` factory returns an +isolated `MemoryBackend` instance per company. The `embedder` kwarg is required for the +Mem0 backend (must be a `Mem0EmbedderConfig`). ### Consolidation and Retention diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 8c78529a62..1133aba7fd 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -2,7 +2,7 @@ Implements ``MemoryBackend``, ``MemoryCapabilities``, and ``SharedKnowledgeStore`` protocols using Mem0 as the storage layer -(default: embedded Qdrant + SQLite). +(default: Qdrant (embedded by default) + SQLite). All Mem0 calls are synchronous — they run in ``asyncio.to_thread()`` to avoid blocking the event loop. @@ -19,7 +19,7 @@ import asyncio import builtins -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, runtime_checkable from ai_company.core.enums import MemoryCategory from ai_company.core.types import NotBlankStr @@ -28,7 +28,7 @@ build_mem0_config_dict, ) from ai_company.memory.backends.mem0.mappers import ( - _PUBLISHER_KEY, + PUBLISHER_KEY, apply_post_filters, build_mem0_metadata, extract_category, @@ -47,14 +47,6 @@ MemoryError as DomainMemoryError, ) from ai_company.observability import get_logger - -if TYPE_CHECKING: - from ai_company.memory.models import ( - MemoryEntry, - MemoryQuery, - MemoryStoreRequest, - ) - from ai_company.observability.events.memory import ( MEMORY_BACKEND_AGENT_ID_REJECTED, MEMORY_BACKEND_CONNECTED, @@ -82,9 +74,50 @@ MEMORY_SHARED_SEARCHED, ) +if TYPE_CHECKING: + from typing import Protocol + + from ai_company.memory.models import ( + MemoryEntry, + MemoryQuery, + MemoryStoreRequest, + ) + + @runtime_checkable + class Mem0Client(Protocol): + """Structural type for the Mem0 ``Memory`` client. + + Defines the subset of ``Memory`` methods that the adapter + uses, so the rest of the codebase does not depend on + ``Any``. + """ + + def add(self, **kwargs: Any) -> dict[str, Any]: + """Add a memory entry.""" + ... + + def search(self, **kwargs: Any) -> dict[str, Any]: + """Search memories.""" + ... + + def get_all(self, **kwargs: Any) -> dict[str, Any]: + """Get all memories for a user.""" + ... + + def get(self, memory_id: str) -> dict[str, Any] | None: + """Get a single memory by ID.""" + ... + + def delete(self, memory_id: str) -> None: + """Delete a memory by ID.""" + ... + + logger = get_logger(__name__) # Reserved user_id for the shared knowledge namespace. +# All shared memories are stored under this Mem0 ``user_id`` so they +# are isolated from per-agent memories and can be queried centrally. _SHARED_NAMESPACE: str = "__synthorg_shared__" @@ -111,13 +144,34 @@ def _validate_mem0_result( f"Unexpected Mem0 response type for {context}: " f"{type(raw_result).__name__}, expected dict" ) + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + context=context, + error=msg, + ) + raise MemoryRetrievalError(msg) + if "results" not in raw_result: + msg = ( + f"Mem0 response missing 'results' key for {context}: " + f"keys={list(raw_result.keys())}" + ) + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + context=context, + error=msg, + ) raise MemoryRetrievalError(msg) - raw_list = raw_result.get("results", []) + raw_list = raw_result["results"] if not isinstance(raw_list, list): msg = ( f"Unexpected Mem0 results type for {context}: " f"{type(raw_list).__name__}, expected list" ) + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + context=context, + error=msg, + ) raise MemoryRetrievalError(msg) return raw_list @@ -143,14 +197,17 @@ def __init__( self._max_memories_per_agent = max_memories_per_agent self._client: Any = None self._connected = False + self._connect_lock = asyncio.Lock() # ── Lifecycle ───────────────────────────────────────────────── async def connect(self) -> None: """Establish connection to Mem0. - Creates the Mem0 ``Memory`` client with embedded Qdrant. - Idempotent — returns immediately if already connected. + Creates the Mem0 ``Memory`` client with Qdrant + (embedded by default). Idempotent — returns immediately + if already connected. Uses a lock to prevent concurrent + ``connect()`` calls from creating duplicate clients. Raises: MemoryConnectionError: If Mem0 is not installed or @@ -158,35 +215,40 @@ async def connect(self) -> None: """ if self._connected: return - logger.info(MEMORY_BACKEND_CONNECTING, backend="mem0") - try: - from mem0 import Memory # noqa: PLC0415 - except ImportError as exc: - logger.warning( - MEMORY_BACKEND_CONNECTION_FAILED, - backend="mem0", - error=str(exc), - error_type="ImportError", - ) - msg = "mem0 package is not installed" - raise MemoryConnectionError(msg) from exc - try: - config_dict = build_mem0_config_dict(self._mem0_config) - client = await asyncio.to_thread(Memory.from_config, config_dict) - except builtins.MemoryError, RecursionError: - raise - except Exception as exc: - logger.warning( - MEMORY_BACKEND_CONNECTION_FAILED, - backend="mem0", - error=str(exc), - error_type=type(exc).__name__, - ) - msg = f"Failed to connect to Mem0: {exc}" - raise MemoryConnectionError(msg) from exc - self._client = client - self._connected = True - logger.info(MEMORY_BACKEND_CONNECTED, backend="mem0") + async with self._connect_lock: + # Double-check after acquiring the lock — another + # coroutine may have connected while we waited. + if self._connected: + return # type: ignore[unreachable] # concurrent state change + logger.info(MEMORY_BACKEND_CONNECTING, backend="mem0") + try: + from mem0 import Memory # noqa: PLC0415 + except ImportError as exc: + logger.warning( + MEMORY_BACKEND_CONNECTION_FAILED, + backend="mem0", + error=str(exc), + error_type="ImportError", + ) + msg = "mem0 package is not installed" + raise MemoryConnectionError(msg) from exc + try: + config_dict = build_mem0_config_dict(self._mem0_config) + client = await asyncio.to_thread(Memory.from_config, config_dict) + except builtins.MemoryError, RecursionError: + raise + except Exception as exc: + logger.warning( + MEMORY_BACKEND_CONNECTION_FAILED, + backend="mem0", + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to connect to Mem0: {exc}" + raise MemoryConnectionError(msg) from exc + self._client = client + self._connected = True + logger.info(MEMORY_BACKEND_CONNECTED, backend="mem0") async def disconnect(self) -> None: """Close the Mem0 connection. @@ -306,11 +368,14 @@ def _validate_agent_id( agent_id: Agent identifier to validate. error_cls: Error class to raise on rejection — defaults to ``MemoryStoreError`` for write ops, pass - ``MemoryRetrievalError`` for read ops. + ``MemoryRetrievalError`` or ``MemoryConnectionError`` + for read/connection ops. Raises: - DomainMemoryError: (subclass per ``error_cls``) If - ``agent_id`` collides with ``_SHARED_NAMESPACE``. + MemoryStoreError: If ``agent_id`` collides with + ``_SHARED_NAMESPACE`` (default). + MemoryRetrievalError: If ``error_cls`` was set to + ``MemoryRetrievalError``. """ if str(agent_id) == _SHARED_NAMESPACE: logger.warning( @@ -417,9 +482,7 @@ async def retrieve( kwargs = query_to_mem0_getall_args(str(agent_id), query) raw_result = await asyncio.to_thread(self._client.get_all, **kwargs) raw_list = _validate_mem0_result(raw_result, context="retrieve") - entries = tuple( - mem0_result_to_entry(item, str(agent_id)) for item in raw_list - ) + entries = tuple(mem0_result_to_entry(item, agent_id) for item in raw_list) entries = apply_post_filters(entries, query) except MemoryRetrievalError as exc: logger.warning( @@ -482,7 +545,16 @@ async def get( ) return None owner = raw.get("user_id") - if owner is not None and str(owner) != str(agent_id): + if owner is None: + logger.warning( + MEMORY_ENTRY_FETCHED, + agent_id=agent_id, + memory_id=memory_id, + found=True, + reason="memory has no user_id — ownership " + "unverifiable, returning to requesting agent", + ) + elif str(owner) != str(agent_id): logger.debug( MEMORY_ENTRY_FETCHED, agent_id=agent_id, @@ -492,7 +564,7 @@ async def get( actual_owner=str(owner), ) return None - entry = mem0_result_to_entry(raw, str(agent_id)) + entry = mem0_result_to_entry(raw, agent_id) except MemoryRetrievalError as exc: logger.warning( MEMORY_ENTRY_FETCH_FAILED, @@ -561,6 +633,14 @@ async def delete( return False # Block deletion of shared-namespace entries — use retract(). owner = existing.get("user_id") + if owner is None: + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="memory has no user_id — ownership " + "unverifiable, proceeding with delete", + ) if owner is not None and str(owner) == _SHARED_NAMESPACE: msg = ( f"Memory {memory_id} belongs to the shared namespace — " @@ -626,8 +706,11 @@ async def count( Note: Results are capped at ``max_memories_per_agent``. If an agent has more memories than this limit the count will be - an underestimate. This is consistent with the adapter's - store/retrieve semantics which also respect the cap. + an underestimate. Truncation is detected when the raw + result set (before any category filter) reaches + ``max_memories_per_agent``. This is consistent with the + adapter's store/retrieve semantics which also respect the + cap. Args: agent_id: Owning agent identifier. @@ -676,7 +759,7 @@ async def count( msg = f"Failed to count memories: {exc}" raise MemoryRetrievalError(msg) from exc else: - truncated = total == self._max_memories_per_agent + truncated = len(raw_list) == self._max_memories_per_agent if truncated: logger.warning( MEMORY_ENTRY_COUNTED, @@ -684,8 +767,9 @@ async def count( count=total, category=category.value if category else None, truncated=True, - reason="count equals max_memories_per_agent, " - "actual count may be higher", + reason="raw result set reached max_memories_per_agent " + "limit — actual count may be higher " + "(conservative estimate)", ) else: logger.info( @@ -720,10 +804,11 @@ async def publish( MemoryStoreError: If the publish operation fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: metadata = { **build_mem0_metadata(request), - _PUBLISHER_KEY: str(agent_id), + PUBLISHER_KEY: str(agent_id), } kwargs = { "messages": [ @@ -801,13 +886,21 @@ async def search_shared( context="search_shared", ) - raw_entries = tuple( - mem0_result_to_entry( - item, - extract_publisher(item) or _SHARED_NAMESPACE, + entries_list: list[MemoryEntry] = [] + for item in raw_list: + publisher = extract_publisher(item) + if publisher is None: + logger.debug( + MEMORY_SHARED_SEARCHED, + memory_id=item.get("id", "?"), + reason="no publisher metadata — " + "attributing to shared namespace", + ) + publisher = _SHARED_NAMESPACE + entries_list.append( + mem0_result_to_entry(item, NotBlankStr(publisher)), ) - for item in raw_list - ) + raw_entries = tuple(entries_list) filtered = apply_post_filters(raw_entries, query) if exclude_agent is not None: diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 7b78fe6369..612677ac6c 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -21,12 +21,13 @@ class Mem0EmbedderConfig(BaseModel): - """Embedder settings for Mem0. + """Embedder settings for the Mem0 memory backend. - ``provider`` and ``model`` are required — callers must supply them - explicitly so that vendor-specific identifiers stay out of source - defaults. Pass values that the Mem0 SDK recognises (e.g. via - company YAML config). + Both ``provider`` and ``model`` are required — callers must + supply them explicitly so that vendor-specific identifiers stay + out of source defaults. The values must be valid Mem0 SDK + identifiers (e.g. ``"openai"``, ``"text-embedding-ada-002"``); + see the Mem0 documentation for supported providers and models. Attributes: provider: Embedding provider name (Mem0 SDK identifier). @@ -74,7 +75,12 @@ class Mem0BackendConfig(BaseModel): @model_validator(mode="after") def _reject_traversal(self) -> Self: - """Reject parent-directory traversal to prevent path escapes.""" + """Reject parent-directory traversal to prevent path escapes. + + Note: ``build_config_from_company_config`` passes ``data_dir`` + from ``CompanyMemoryConfig``, so this check also protects + the factory path. + """ parts = ( PureWindowsPath(self.data_dir).parts + PurePosixPath(self.data_dir).parts ) @@ -138,7 +144,10 @@ def build_config_from_company_config( Raises: ValueError: If the storage config specifies a vector or - history store that the Mem0 backend does not support. + history store that the Mem0 backend does not support, + or if ``data_dir`` contains parent-directory traversal + (``..``) — propagated from ``Mem0BackendConfig`` + validation. """ if config.storage.vector_store not in ("qdrant", "qdrant-external"): msg = ( diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 83982af063..fdc9d60bda 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -32,7 +32,8 @@ _PREFIX = "_synthorg_" # Metadata key to track who published a shared memory. -_PUBLISHER_KEY: str = "_synthorg_publisher" +# Public because the adapter module needs it for ownership tracking. +PUBLISHER_KEY: str = "_synthorg_publisher" def build_mem0_metadata(request: MemoryStoreRequest) -> dict[str, Any]: @@ -115,7 +116,10 @@ def normalize_relevance_score(score: Any) -> float | None: def _coerce_confidence(raw_metadata: dict[str, Any]) -> float: """Extract and clamp confidence from Mem0 metadata. - Returns a float in [0.0, 1.0], defaulting to 1.0 on failure. + Returns a float in [0.0, 1.0]. Defaults to 1.0 when the key is + absent (newly stored entries always write it), or 0.5 when the + value is present but non-numeric (corrupt data gets a conservative + mid-range default rather than maximum confidence). """ raw = raw_metadata.get(f"{_PREFIX}confidence", 1.0) try: @@ -125,9 +129,9 @@ def _coerce_confidence(raw_metadata: dict[str, Any]) -> float: MEMORY_MODEL_INVALID, field="confidence", raw_value=raw, - reason="non-numeric confidence, defaulting to 1.0", + reason="non-numeric confidence, defaulting to 0.5", ) - return 1.0 + return 0.5 return max(0.0, min(1.0, value)) @@ -162,7 +166,7 @@ def _normalize_tags( if isinstance(raw_tags, str): raw_tags = [raw_tags] elif not isinstance(raw_tags, (list, tuple)): - logger.debug( + logger.warning( MEMORY_MODEL_INVALID, field="tags", raw_value=type(raw_tags).__name__, @@ -184,7 +188,8 @@ def parse_mem0_metadata( Tuple of (category, metadata, expires_at). """ if not raw_metadata or not isinstance(raw_metadata, dict): - logger.debug( + log_fn = logger.warning if raw_metadata is not None else logger.debug + log_fn( MEMORY_MODEL_INVALID, field="metadata", raw_value=type(raw_metadata).__name__ if raw_metadata else None, @@ -228,14 +233,14 @@ def parse_mem0_metadata( def mem0_result_to_entry( raw: dict[str, Any], - agent_id: str, + agent_id: NotBlankStr, ) -> MemoryEntry: """Convert a single Mem0 result dict to a ``MemoryEntry``. Args: raw: Single result dict from Mem0 (``search``, ``get``, or ``get_all``). - agent_id: Owning agent identifier. + agent_id: Owning agent identifier (must be ``NotBlankStr``). Returns: Domain ``MemoryEntry``. @@ -266,7 +271,7 @@ def mem0_result_to_entry( created_at = parse_mem0_datetime(raw.get("created_at")) if created_at is None: - logger.debug( + logger.warning( MEMORY_MODEL_INVALID, field="created_at", memory_id=str(raw.get("id", "?")), @@ -283,7 +288,7 @@ def mem0_result_to_entry( return MemoryEntry( id=memory_id, - agent_id=NotBlankStr(agent_id), + agent_id=agent_id, category=category, content=content, metadata=metadata, @@ -357,6 +362,9 @@ def apply_post_filters( excluded by ``min_relevance`` — the filter only applies when a score is present. + Time range uses a half-open interval: entries with + ``created_at >= since`` and ``created_at < until`` are included. + Args: entries: Raw entries from Mem0. query: Original query with filter criteria. @@ -467,7 +475,7 @@ def extract_publisher(raw: dict[str, Any]) -> str | None: metadata = raw.get("metadata", {}) if not metadata or not isinstance(metadata, dict): return None - value = metadata.get(_PUBLISHER_KEY) + value = metadata.get(PUBLISHER_KEY) if value is None: return None coerced = str(value).strip() diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 73359d7fac..271748f09b 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -7,6 +7,8 @@ from typing import TYPE_CHECKING +from pydantic import ValidationError + from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 if TYPE_CHECKING: @@ -96,7 +98,7 @@ def create_memory_backend( mem0_config=mem0_config, max_memories_per_agent=config.options.max_memories_per_agent, ) - except Exception as exc: + except (ValueError, ValidationError) as exc: msg = f"Failed to create Mem0 backend: {exc}" logger.warning( MEMORY_BACKEND_CONFIG_INVALID, diff --git a/src/ai_company/memory/models.py b/src/ai_company/memory/models.py index b8dd9703f0..23f9e25f0f 100644 --- a/src/ai_company/memory/models.py +++ b/src/ai_company/memory/models.py @@ -194,7 +194,7 @@ class MemoryQuery(BaseModel): ) since: AwareDatetime | None = Field( default=None, - description="Only memories created after this timestamp", + description="Only memories created at or after this timestamp", ) until: AwareDatetime | None = Field( default=None, diff --git a/src/ai_company/observability/events/memory.py b/src/ai_company/observability/events/memory.py index b4e259031a..1d7570a3ff 100644 --- a/src/ai_company/observability/events/memory.py +++ b/src/ai_company/observability/events/memory.py @@ -17,7 +17,6 @@ MEMORY_BACKEND_DISCONNECTED: Final[str] = "memory.backend.disconnected" MEMORY_BACKEND_HEALTH_CHECK: Final[str] = "memory.backend.health_check" MEMORY_BACKEND_CREATED: Final[str] = "memory.backend.created" -MEMORY_BACKEND_NOT_IMPLEMENTED: Final[str] = "memory.backend.not_implemented" MEMORY_BACKEND_UNKNOWN: Final[str] = "memory.backend.unknown" MEMORY_BACKEND_CONFIG_INVALID: Final[str] = "memory.backend.config_invalid" MEMORY_BACKEND_NOT_CONNECTED: Final[str] = "memory.backend.not_connected" @@ -49,10 +48,6 @@ MEMORY_MODEL_INVALID: Final[str] = "memory.model.invalid" -# ── Capability checks ──────────────────────────────────────────── - -MEMORY_CAPABILITY_UNSUPPORTED: Final[str] = "memory.capability.unsupported" - # ── Retrieval pipeline ────────────────────────────────────────── MEMORY_RETRIEVAL_START: Final[str] = "memory.retrieval.start" diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py index d95affa1d0..abc43c2431 100644 --- a/tests/integration/memory/test_mem0_backend.py +++ b/tests/integration/memory/test_mem0_backend.py @@ -16,7 +16,7 @@ Mem0BackendConfig, Mem0EmbedderConfig, ) -from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY +from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY from ai_company.memory.models import MemoryQuery, MemoryStoreRequest from ai_company.memory.retrieval_config import MemoryRetrievalConfig from ai_company.memory.retriever import ContextInjectionStrategy @@ -196,7 +196,7 @@ async def test_shared_knowledge_flow( "created_at": datetime.now(UTC).isoformat(), "metadata": { "_synthorg_category": "semantic", - _PUBLISHER_KEY: "test-agent-001", + PUBLISHER_KEY: "test-agent-001", }, }, ], @@ -213,7 +213,7 @@ async def test_shared_knowledge_flow( "id": "shared-001", "memory": "company policy: always test code", "created_at": datetime.now(UTC).isoformat(), - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.return_value = None @@ -233,14 +233,14 @@ async def test_shared_search_excludes_agent( "memory": "from agent 1", "score": 0.9, "created_at": datetime.now(UTC).isoformat(), - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + "metadata": {PUBLISHER_KEY: "test-agent-001"}, }, { "id": "s2", "memory": "from agent 2", "score": 0.85, "created_at": datetime.now(UTC).isoformat(), - "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + "metadata": {PUBLISHER_KEY: "test-agent-002"}, }, ], } diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index 411bfa59bf..4647ea2e0a 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -8,6 +8,7 @@ from ai_company.memory.backends.mem0.adapter import ( _SHARED_NAMESPACE, Mem0MemoryBackend, + _validate_mem0_result, ) from ai_company.memory.errors import ( MemoryRetrievalError, @@ -129,10 +130,22 @@ async def test_store_non_list_results_raises( async def test_store_rejects_shared_namespace_agent_id( self, backend: Mem0MemoryBackend, + mock_client: MagicMock, ) -> None: """Storing with the shared namespace agent ID is rejected.""" with pytest.raises(MemoryStoreError, match="reserved shared namespace"): await backend.store(_SHARED_NAMESPACE, make_store_request()) + mock_client.add.assert_not_called() + + async def test_store_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping.""" + mock_client.add.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.store("test-agent-001", make_store_request()) # ── Retrieve ────────────────────────────────────────────────────── @@ -236,6 +249,19 @@ async def test_retrieve_reraises_memory_error( MemoryQuery(text="test"), ) + async def test_retrieve_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping.""" + mock_client.search.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.retrieve( + "test-agent-001", + MemoryQuery(text="test"), + ) + async def test_retrieve_rejects_shared_namespace_agent_id( self, backend: Mem0MemoryBackend, @@ -314,6 +340,16 @@ async def test_get_reraises_memory_error( with pytest.raises(MemoryError): await backend.get("test-agent-001", "mem-001") + async def test_get_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping in get().""" + mock_client.get.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.get("test-agent-001", "mem-001") + async def test_get_rejects_shared_namespace_agent_id( self, backend: Mem0MemoryBackend, @@ -398,6 +434,16 @@ async def test_delete_reraises_memory_error( with pytest.raises(MemoryError): await backend.delete("test-agent-001", "mem-001") + async def test_delete_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping in delete().""" + mock_client.get.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.delete("test-agent-001", "mem-001") + async def test_delete_shared_namespace_entry_raises( self, backend: Mem0MemoryBackend, @@ -534,6 +580,23 @@ async def test_count_empty_results( count = await backend.count("test-agent-001") assert count == 0 + async def test_count_truncation_warning( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """count() returns the count even when truncated at max_memories_per_agent.""" + # backend fixture has max_memories_per_agent=100 + items = [ + {"id": f"m{i}", "memory": f"content-{i}", "metadata": {}} + for i in range(100) + ] + mock_client.get_all.return_value = {"results": items} + + count = await backend.count("test-agent-001") + # Truncation should still return a valid count + assert count == 100 + async def test_count_reraises_memory_error( self, backend: Mem0MemoryBackend, @@ -543,3 +606,50 @@ async def test_count_reraises_memory_error( mock_client.get_all.side_effect = MemoryError("out of memory") with pytest.raises(MemoryError): await backend.count("test-agent-001") + + async def test_count_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping in count().""" + mock_client.get_all.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.count("test-agent-001") + + +# ── _validate_mem0_result ──────────────────────────────────────── + + +@pytest.mark.unit +class TestValidateMem0Result: + def test_non_dict_raises(self) -> None: + """Non-dict response raises MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="Unexpected Mem0 response type"): + _validate_mem0_result("not-a-dict", context="test") + + def test_missing_results_key_raises(self) -> None: + """Dict without 'results' key raises MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="missing 'results' key"): + _validate_mem0_result({"data": []}, context="test") + + def test_non_list_results_raises(self) -> None: + """Non-list 'results' value raises MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="Unexpected Mem0 results type"): + _validate_mem0_result({"results": "not-a-list"}, context="test") + + def test_valid_response(self) -> None: + """Valid response returns the results list.""" + items = [{"id": "m1"}] + result = _validate_mem0_result({"results": items}, context="test") + assert result == items + + def test_empty_results(self) -> None: + """Empty results list is valid.""" + result = _validate_mem0_result({"results": []}, context="test") + assert result == [] + + def test_none_raises(self) -> None: + """None response raises MemoryRetrievalError.""" + with pytest.raises(MemoryRetrievalError, match="Unexpected Mem0 response type"): + _validate_mem0_result(None, context="test") diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py index 4711ce5ecb..7879e7ff64 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_shared.py +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -8,7 +8,7 @@ _SHARED_NAMESPACE, Mem0MemoryBackend, ) -from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY +from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY from ai_company.memory.errors import ( MemoryRetrievalError, MemoryStoreError, @@ -44,8 +44,8 @@ async def test_publish_success( assert memory_id == "shared-mem-001" call_kwargs = mock_client.add.call_args[1] assert call_kwargs["user_id"] == _SHARED_NAMESPACE - assert _PUBLISHER_KEY in call_kwargs["metadata"] - assert call_kwargs["metadata"][_PUBLISHER_KEY] == "test-agent-001" + assert PUBLISHER_KEY in call_kwargs["metadata"] + assert call_kwargs["metadata"][PUBLISHER_KEY] == "test-agent-001" async def test_publish_empty_results_raises( self, @@ -90,6 +90,16 @@ async def test_publish_reraises_memory_error( with pytest.raises(MemoryError): await backend.publish("test-agent-001", make_store_request()) + async def test_publish_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping.""" + mock_client.add.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.publish("test-agent-001", make_store_request()) + # ── SearchShared ───────────────────────────────────────────────── @@ -110,7 +120,7 @@ async def test_search_shared_with_text( "created_at": "2026-03-12T10:00:00+00:00", "metadata": { "_synthorg_category": "semantic", - _PUBLISHER_KEY: "test-agent-002", + PUBLISHER_KEY: "test-agent-002", }, }, ], @@ -137,7 +147,7 @@ async def test_search_shared_without_text( "memory": "shared fact", "created_at": "2026-03-12T10:00:00+00:00", "metadata": { - _PUBLISHER_KEY: "test-agent-002", + PUBLISHER_KEY: "test-agent-002", }, }, ], @@ -161,14 +171,14 @@ async def test_search_shared_exclude_agent( "memory": "from agent 1", "score": 0.9, "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + "metadata": {PUBLISHER_KEY: "test-agent-001"}, }, { "id": "s2", "memory": "from agent 2", "score": 0.8, "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + "metadata": {PUBLISHER_KEY: "test-agent-002"}, }, ], ) @@ -202,6 +212,57 @@ async def test_search_shared_reraises_memory_error( with pytest.raises(MemoryError): await backend.search_shared(MemoryQuery(text="test")) + async def test_search_shared_with_category_post_filter( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """search_shared applies post-filters (e.g. category filter).""" + mock_client.search.return_value = mem0_search_result( + [ + { + "id": "s1", + "memory": "episodic fact", + "score": 0.9, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + "_synthorg_category": "episodic", + PUBLISHER_KEY: "test-agent-001", + }, + }, + { + "id": "s2", + "memory": "semantic fact", + "score": 0.8, + "created_at": "2026-03-12T10:00:00+00:00", + "metadata": { + "_synthorg_category": "semantic", + PUBLISHER_KEY: "test-agent-001", + }, + }, + ], + ) + + from ai_company.core.enums import MemoryCategory + + query = MemoryQuery( + text="test", + categories=frozenset({MemoryCategory.SEMANTIC}), + ) + entries = await backend.search_shared(query) + assert len(entries) == 1 + assert entries[0].category == MemoryCategory.SEMANTIC + + async def test_search_shared_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping.""" + mock_client.search.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.search_shared(MemoryQuery(text="test")) + async def test_search_shared_no_publisher_uses_namespace( self, backend: Mem0MemoryBackend, @@ -239,7 +300,7 @@ async def test_retract_success( "id": "shared-001", "memory": "shared content", "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.return_value = None @@ -268,7 +329,7 @@ async def test_retract_ownership_mismatch( "id": "shared-001", "memory": "content", "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-002"}, + "metadata": {PUBLISHER_KEY: "test-agent-002"}, } with pytest.raises(MemoryStoreError, match="cannot retract"): @@ -309,6 +370,16 @@ async def test_retract_reraises_memory_error( with pytest.raises(MemoryError): await backend.retract("test-agent-001", "shared-001") + async def test_retract_reraises_recursion_error( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """RecursionError is re-raised without wrapping.""" + mock_client.get.side_effect = RecursionError("infinite loop") + with pytest.raises(RecursionError): + await backend.retract("test-agent-001", "shared-001") + async def test_retract_delete_failure_wraps( self, backend: Mem0MemoryBackend, @@ -319,7 +390,7 @@ async def test_retract_delete_failure_wraps( "id": "shared-001", "memory": "content", "created_at": "2026-03-12T10:00:00+00:00", - "metadata": {_PUBLISHER_KEY: "test-agent-001"}, + "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.side_effect = RuntimeError("delete failed") diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index 3cfbd48c8a..b28d6b5305 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -8,7 +8,9 @@ from ai_company.core.enums import MemoryCategory from ai_company.memory.backends.mem0.mappers import ( _PREFIX, - _PUBLISHER_KEY, + PUBLISHER_KEY, + _coerce_confidence, + _normalize_tags, apply_post_filters, build_mem0_metadata, extract_category, @@ -539,9 +541,9 @@ def test_missing_category_key(self) -> None: @pytest.mark.unit class TestExtractPublisher: def test_valid_publisher(self) -> None: - from ai_company.memory.backends.mem0.mappers import _PUBLISHER_KEY + from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY - raw = {"metadata": {_PUBLISHER_KEY: "test-agent-001"}} + raw = {"metadata": {PUBLISHER_KEY: "test-agent-001"}} assert extract_publisher(raw) == "test-agent-001" def test_missing_metadata(self) -> None: @@ -569,9 +571,63 @@ def test_string_metadata_returns_none(self) -> None: assert extract_publisher(raw) is None def test_numeric_publisher_coerced_to_string(self) -> None: - raw: dict[str, Any] = {"metadata": {_PUBLISHER_KEY: 42}} + raw: dict[str, Any] = {"metadata": {PUBLISHER_KEY: 42}} assert extract_publisher(raw) == "42" def test_blank_publisher_returns_none(self) -> None: - raw: dict[str, Any] = {"metadata": {_PUBLISHER_KEY: " "}} + raw: dict[str, Any] = {"metadata": {PUBLISHER_KEY: " "}} assert extract_publisher(raw) is None + + +@pytest.mark.unit +class TestCoerceConfidence: + def test_default_when_absent(self) -> None: + """Missing key returns 1.0.""" + assert _coerce_confidence({}) == 1.0 + + def test_numeric_value(self) -> None: + assert _coerce_confidence({f"{_PREFIX}confidence": 0.7}) == 0.7 + + def test_non_numeric_returns_half(self) -> None: + """Non-numeric confidence defaults to 0.5.""" + assert _coerce_confidence({f"{_PREFIX}confidence": "not-a-number"}) == 0.5 + + def test_above_one_clamped(self) -> None: + assert _coerce_confidence({f"{_PREFIX}confidence": 1.5}) == 1.0 + + def test_below_zero_clamped(self) -> None: + assert _coerce_confidence({f"{_PREFIX}confidence": -0.5}) == 0.0 + + def test_object_type_returns_half(self) -> None: + """Object type that can't convert to float defaults to 0.5.""" + assert _coerce_confidence({f"{_PREFIX}confidence": object()}) == 0.5 + + +@pytest.mark.unit +class TestNormalizeTags: + def test_single_string_wrapped_in_list(self) -> None: + """A single string tag is wrapped into a tuple.""" + result = _normalize_tags({f"{_PREFIX}tags": "solo-tag"}) + assert result == ("solo-tag",) + + def test_list_of_strings(self) -> None: + result = _normalize_tags({f"{_PREFIX}tags": ["a", "b", "c"]}) + assert result == ("a", "b", "c") + + def test_empty_strings_filtered(self) -> None: + result = _normalize_tags({f"{_PREFIX}tags": ["valid", "", " "]}) + assert result == ("valid",) + + def test_dict_type_ignored(self) -> None: + """dict type for tags is unexpected and returns empty tuple.""" + result = _normalize_tags({f"{_PREFIX}tags": {"key": "value"}}) + assert result == () + + def test_int_type_ignored(self) -> None: + """int type for tags is unexpected and returns empty tuple.""" + result = _normalize_tags({f"{_PREFIX}tags": 42}) + assert result == () + + def test_missing_key_returns_empty(self) -> None: + result = _normalize_tags({}) + assert result == () diff --git a/tests/unit/memory/test_factory.py b/tests/unit/memory/test_factory.py index 0a9b35d744..e22220bb82 100644 --- a/tests/unit/memory/test_factory.py +++ b/tests/unit/memory/test_factory.py @@ -73,13 +73,35 @@ def test_config_build_error_wraps_as_memory_config_error(self) -> None: with pytest.raises(MemoryConfigError, match="Invalid Mem0 configuration"): create_memory_backend(config, embedder=_test_embedder()) - def test_backend_init_error_wraps_as_memory_config_error(self) -> None: - """Exception from Mem0MemoryBackend() constructor wraps.""" + def test_backend_init_value_error_wraps(self) -> None: + """ValueError from Mem0MemoryBackend() constructor wraps.""" config = CompanyMemoryConfig(backend="mem0") with ( patch( "ai_company.memory.backends.mem0.Mem0MemoryBackend", - side_effect=RuntimeError("init boom"), + side_effect=ValueError("init boom"), + ), + pytest.raises(MemoryConfigError, match="Failed to create Mem0"), + ): + create_memory_backend(config, embedder=_test_embedder()) + + def test_backend_init_validation_error_wraps(self) -> None: + """ValidationError from Mem0MemoryBackend() constructor wraps.""" + from pydantic import BaseModel + + class _Dummy(BaseModel): + x: int + + try: + _Dummy(x="not-an-int") # type: ignore[arg-type] + except ValidationError as ve: + side_effect: ValidationError = ve + + config = CompanyMemoryConfig(backend="mem0") + with ( + patch( + "ai_company.memory.backends.mem0.Mem0MemoryBackend", + side_effect=side_effect, ), pytest.raises(MemoryConfigError, match="Failed to create Mem0"), ): diff --git a/tests/unit/observability/test_events.py b/tests/unit/observability/test_events.py index 62b229c749..3a3c52b425 100644 --- a/tests/unit/observability/test_events.py +++ b/tests/unit/observability/test_events.py @@ -476,7 +476,6 @@ def test_org_memory_events_exist( ("MEMORY_BACKEND_DISCONNECTED", "memory.backend.disconnected"), ("MEMORY_BACKEND_HEALTH_CHECK", "memory.backend.health_check"), ("MEMORY_BACKEND_CREATED", "memory.backend.created"), - ("MEMORY_BACKEND_NOT_IMPLEMENTED", "memory.backend.not_implemented"), ("MEMORY_BACKEND_UNKNOWN", "memory.backend.unknown"), ("MEMORY_BACKEND_NOT_CONNECTED", "memory.backend.not_connected"), ("MEMORY_ENTRY_STORED", "memory.entry.stored"), @@ -496,7 +495,6 @@ def test_org_memory_events_exist( ("MEMORY_SHARED_RETRACTED", "memory.shared.retracted"), ("MEMORY_SHARED_RETRACT_FAILED", "memory.shared.retract_failed"), ("MEMORY_MODEL_INVALID", "memory.model.invalid"), - ("MEMORY_CAPABILITY_UNSUPPORTED", "memory.capability.unsupported"), ], ) def test_memory_events_exist(self, constant_name: str, expected: str) -> None: From 16975cb8fbd537b887e7b8dd79108d89d6c7717b Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 09:46:45 +0100 Subject: [PATCH 13/17] fix: address round-7 PR review findings across adapter, mappers, config, factory, and tests MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Remove @runtime_checkable from TYPE_CHECKING-only Mem0Client protocol - Add _connect_lock to disconnect() to prevent race with connect() - Fix MEMORY_ENTRY_DELETE_FAILED event on non-failure path (unverifiable owner) - Add shared namespace verification in retract() (defense-in-depth) - Add _validate_agent_id() call in retract() - Add math.isfinite() guards in normalize_relevance_score() and _coerce_confidence() - Prefer updated_at over now() as created_at fallback in mem0_result_to_entry() - Reject qdrant-external in build_config_from_company_config() (not supported) - Extract _create_mem0_backend() helper from create_memory_backend() - Fix MemoryQuery class docstring: "after" → "at or after" for since field - Use explicit builtins.MemoryError in all test assertions --- .../memory/backends/mem0/adapter.py | 34 +++- src/ai_company/memory/backends/mem0/config.py | 4 +- .../memory/backends/mem0/mappers.py | 29 +++- src/ai_company/memory/factory.py | 151 ++++++++++-------- src/ai_company/memory/models.py | 2 +- tests/integration/memory/test_mem0_backend.py | 6 +- .../unit/memory/backends/mem0/test_adapter.py | 9 +- .../memory/backends/mem0/test_adapter_crud.py | 21 +-- .../backends/mem0/test_adapter_shared.py | 34 +++- .../unit/memory/backends/mem0/test_config.py | 13 +- .../unit/memory/backends/mem0/test_mappers.py | 20 +++ 11 files changed, 218 insertions(+), 105 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 1133aba7fd..e9677a65c8 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -19,7 +19,7 @@ import asyncio import builtins -from typing import TYPE_CHECKING, Any, runtime_checkable +from typing import TYPE_CHECKING, Any from ai_company.core.enums import MemoryCategory from ai_company.core.types import NotBlankStr @@ -83,7 +83,6 @@ MemoryStoreRequest, ) - @runtime_checkable class Mem0Client(Protocol): """Structural type for the Mem0 ``Memory`` client. @@ -255,11 +254,14 @@ async def disconnect(self) -> None: Releases the client reference so the garbage collector can reclaim resources. Safe to call even if not connected. + Acquires ``_connect_lock`` to prevent racing with an + in-progress ``connect()`` call. """ - logger.info(MEMORY_BACKEND_DISCONNECTING, backend="mem0") - self._client = None - self._connected = False - logger.info(MEMORY_BACKEND_DISCONNECTED, backend="mem0") + async with self._connect_lock: + logger.info(MEMORY_BACKEND_DISCONNECTING, backend="mem0") + self._client = None + self._connected = False + logger.info(MEMORY_BACKEND_DISCONNECTED, backend="mem0") async def health_check(self) -> bool: """Check whether the Mem0 backend is healthy. @@ -635,9 +637,10 @@ async def delete( owner = existing.get("user_id") if owner is None: logger.warning( - MEMORY_ENTRY_DELETE_FAILED, + MEMORY_ENTRY_DELETED, agent_id=agent_id, memory_id=memory_id, + unverifiable_ownership=True, reason="memory has no user_id — ownership " "unverifiable, proceeding with delete", ) @@ -956,6 +959,7 @@ async def retract( ownership verification fails. """ self._require_connected() + self._validate_agent_id(agent_id) try: raw = await asyncio.to_thread(self._client.get, str(memory_id)) if raw is None: @@ -967,6 +971,22 @@ async def retract( ) return False + # Verify this memory belongs to the shared namespace. + owner_ns = raw.get("user_id") + if owner_ns != _SHARED_NAMESPACE: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="not in shared namespace", + actual_namespace=str(owner_ns), + ) + msg = ( + f"Memory {memory_id} is not in the shared namespace — " + f"use delete() to remove private entries" + ) + raise MemoryStoreError(msg) # noqa: TRY301 + publisher = extract_publisher(raw) if publisher is None: logger.warning( diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 612677ac6c..6ff77c7aa6 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -149,9 +149,9 @@ def build_config_from_company_config( (``..``) — propagated from ``Mem0BackendConfig`` validation. """ - if config.storage.vector_store not in ("qdrant", "qdrant-external"): + if config.storage.vector_store != "qdrant": msg = ( - f"Mem0 backend only supports qdrant vector stores, " + f"Mem0 backend only supports embedded qdrant vector store, " f"got {config.storage.vector_store!r}" ) logger.warning( diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index fdc9d60bda..5fda6621d9 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -5,6 +5,7 @@ stays thin. """ +import math from datetime import UTC, datetime from typing import TYPE_CHECKING, Any @@ -110,6 +111,14 @@ def normalize_relevance_score(score: Any) -> float | None: reason="non-numeric relevance score, returning None", ) return None + if not math.isfinite(numeric): + logger.warning( + MEMORY_MODEL_INVALID, + field="score", + raw_value=score, + reason="non-finite relevance score, returning None", + ) + return None return max(0.0, min(1.0, numeric)) @@ -132,6 +141,14 @@ def _coerce_confidence(raw_metadata: dict[str, Any]) -> float: reason="non-numeric confidence, defaulting to 0.5", ) return 0.5 + if not math.isfinite(value): + logger.warning( + MEMORY_MODEL_INVALID, + field="confidence", + raw_value=raw, + reason="non-finite confidence, defaulting to 0.5", + ) + return 0.5 return max(0.0, min(1.0, value)) @@ -270,15 +287,21 @@ def mem0_result_to_entry( content = NotBlankStr(str(raw_content)) created_at = parse_mem0_datetime(raw.get("created_at")) + updated_at = parse_mem0_datetime(raw.get("updated_at")) if created_at is None: + # Prefer updated_at as a fallback — it is a closer + # approximation than now() and avoids violating the + # MemoryEntry invariant (updated_at >= created_at). + fallback = updated_at or datetime.now(UTC) + fallback_source = "updated_at" if updated_at else "now()" logger.warning( MEMORY_MODEL_INVALID, field="created_at", memory_id=str(raw.get("id", "?")), - reason="missing or unparseable created_at, defaulting to now()", + reason=f"missing or unparseable created_at, " + f"defaulting to {fallback_source}", ) - created_at = datetime.now(UTC) - updated_at = parse_mem0_datetime(raw.get("updated_at")) + created_at = fallback raw_metadata = raw.get("metadata") category, metadata, expires_at = parse_mem0_metadata(raw_metadata) diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 271748f09b..6c1d181f1a 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -25,6 +25,91 @@ logger = get_logger(__name__) +def _create_mem0_backend( + config: CompanyMemoryConfig, + *, + embedder: Mem0EmbedderConfig | None, +) -> MemoryBackend: + """Create a Mem0 memory backend from configuration. + + Args: + config: Company-wide memory configuration. + embedder: Mem0-specific embedder configuration (required). + + Returns: + A new, disconnected ``Mem0MemoryBackend`` instance. + + Raises: + MemoryConfigError: If embedder is missing/invalid or + backend construction fails. + """ + from ai_company.memory.backends.mem0 import Mem0MemoryBackend # noqa: PLC0415 + from ai_company.memory.backends.mem0.config import ( # noqa: PLC0415 + Mem0EmbedderConfig, + build_config_from_company_config, + ) + + if embedder is None: + msg = ( + "Mem0 backend requires an embedder configuration — " + "pass a Mem0EmbedderConfig instance" + ) + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="missing_embedder", + error=msg, + ) + raise MemoryConfigError(msg) + if not isinstance(embedder, Mem0EmbedderConfig): + msg = ( # type: ignore[unreachable] + f"embedder must be a Mem0EmbedderConfig, got {type(embedder).__name__}" + ) + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="invalid_embedder_type", + error=msg, + embedder_type=type(embedder).__name__, + ) + raise MemoryConfigError(msg) + + try: + mem0_config = build_config_from_company_config( + config, + embedder=embedder, + ) + except ValueError as exc: + msg = f"Invalid Mem0 configuration: {exc}" + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="config_build_failed", + error=msg, + ) + raise MemoryConfigError(msg) from exc + try: + backend = Mem0MemoryBackend( + mem0_config=mem0_config, + max_memories_per_agent=config.options.max_memories_per_agent, + ) + except (ValueError, ValidationError) as exc: + msg = f"Failed to create Mem0 backend: {exc}" + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="backend_init_failed", + error=msg, + ) + raise MemoryConfigError(msg) from exc + logger.info( + MEMORY_BACKEND_CREATED, + backend="mem0", + data_dir=mem0_config.data_dir, + ) + return backend + + def create_memory_backend( config: CompanyMemoryConfig, *, @@ -48,71 +133,7 @@ def create_memory_backend( required configuration is missing. """ if config.backend == "mem0": - from ai_company.memory.backends.mem0 import Mem0MemoryBackend # noqa: PLC0415 - from ai_company.memory.backends.mem0.config import ( # noqa: PLC0415 - Mem0EmbedderConfig, - build_config_from_company_config, - ) - - if embedder is None: - msg = ( - "Mem0 backend requires an embedder configuration — " - "pass a Mem0EmbedderConfig instance" - ) - logger.warning( - MEMORY_BACKEND_CONFIG_INVALID, - backend="mem0", - reason="missing_embedder", - error=msg, - ) - raise MemoryConfigError(msg) - if not isinstance(embedder, Mem0EmbedderConfig): - msg = ( # type: ignore[unreachable] - f"embedder must be a Mem0EmbedderConfig, got {type(embedder).__name__}" - ) - logger.warning( - MEMORY_BACKEND_CONFIG_INVALID, - backend="mem0", - reason="invalid_embedder_type", - error=msg, - embedder_type=type(embedder).__name__, - ) - raise MemoryConfigError(msg) - - try: - mem0_config = build_config_from_company_config( - config, - embedder=embedder, - ) - except ValueError as exc: - msg = f"Invalid Mem0 configuration: {exc}" - logger.warning( - MEMORY_BACKEND_CONFIG_INVALID, - backend="mem0", - reason="config_build_failed", - error=msg, - ) - raise MemoryConfigError(msg) from exc - try: - backend = Mem0MemoryBackend( - mem0_config=mem0_config, - max_memories_per_agent=config.options.max_memories_per_agent, - ) - except (ValueError, ValidationError) as exc: - msg = f"Failed to create Mem0 backend: {exc}" - logger.warning( - MEMORY_BACKEND_CONFIG_INVALID, - backend="mem0", - reason="backend_init_failed", - error=msg, - ) - raise MemoryConfigError(msg) from exc - logger.info( - MEMORY_BACKEND_CREATED, - backend="mem0", - data_dir=mem0_config.data_dir, - ) - return backend + return _create_mem0_backend(config, embedder=embedder) # Defensive guard: config validation rejects unknown backends, so # this branch is unreachable under normal construction. It exists # as a safety net for callers that bypass validation (e.g. via diff --git a/src/ai_company/memory/models.py b/src/ai_company/memory/models.py index 23f9e25f0f..1252b8b720 100644 --- a/src/ai_company/memory/models.py +++ b/src/ai_company/memory/models.py @@ -162,7 +162,7 @@ class MemoryQuery(BaseModel): tags: Filter by tags (AND semantics). min_relevance: Minimum relevance score threshold. limit: Maximum number of results. - since: Only memories created after this timestamp. + since: Only memories created at or after this timestamp. until: Only memories created before this timestamp. """ diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py index abc43c2431..c6a6317080 100644 --- a/tests/integration/memory/test_mem0_backend.py +++ b/tests/integration/memory/test_mem0_backend.py @@ -11,7 +11,10 @@ import pytest from ai_company.core.enums import MemoryCategory -from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend +from ai_company.memory.backends.mem0.adapter import ( + _SHARED_NAMESPACE, + Mem0MemoryBackend, +) from ai_company.memory.backends.mem0.config import ( Mem0BackendConfig, Mem0EmbedderConfig, @@ -213,6 +216,7 @@ async def test_shared_knowledge_flow( "id": "shared-001", "memory": "company policy: always test code", "created_at": datetime.now(UTC).isoformat(), + "user_id": _SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.return_value = None diff --git a/tests/unit/memory/backends/mem0/test_adapter.py b/tests/unit/memory/backends/mem0/test_adapter.py index e345764005..6c6e7e11e1 100644 --- a/tests/unit/memory/backends/mem0/test_adapter.py +++ b/tests/unit/memory/backends/mem0/test_adapter.py @@ -1,5 +1,6 @@ """Tests for Mem0 adapter — properties, capabilities, protocol, lifecycle.""" +import builtins import sys from unittest.mock import MagicMock, patch @@ -241,8 +242,8 @@ async def test_health_check_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError propagates through health_check.""" - mock_client.get_all.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.get_all.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.health_check() async def test_connect_memory_error_propagates( @@ -254,9 +255,9 @@ async def test_connect_memory_error_propagates( with ( patch( "ai_company.memory.backends.mem0.adapter.asyncio.to_thread", - side_effect=MemoryError("out of memory"), + side_effect=builtins.MemoryError("out of memory"), ), - pytest.raises(MemoryError), + pytest.raises(builtins.MemoryError), ): await b.connect() diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index 4647ea2e0a..0783421b0f 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -1,5 +1,6 @@ """Tests for Mem0 adapter — store, retrieve, get, delete, count.""" +import builtins from unittest.mock import MagicMock import pytest @@ -89,8 +90,8 @@ async def test_store_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping.""" - mock_client.add.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.add.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.store("test-agent-001", make_store_request()) async def test_store_blank_id_from_add_raises( @@ -242,8 +243,8 @@ async def test_retrieve_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping.""" - mock_client.search.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.search.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.retrieve( "test-agent-001", MemoryQuery(text="test"), @@ -336,8 +337,8 @@ async def test_get_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping in get().""" - mock_client.get.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.get.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.get("test-agent-001", "mem-001") async def test_get_reraises_recursion_error( @@ -430,8 +431,8 @@ async def test_delete_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping in delete().""" - mock_client.get.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.get.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.delete("test-agent-001", "mem-001") async def test_delete_reraises_recursion_error( @@ -603,8 +604,8 @@ async def test_count_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping in count().""" - mock_client.get_all.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.get_all.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.count("test-agent-001") async def test_count_reraises_recursion_error( diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py index 7879e7ff64..c8aa457a45 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_shared.py +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -1,5 +1,6 @@ """Tests for Mem0 adapter — shared knowledge store (publish, search, retract).""" +import builtins from unittest.mock import MagicMock import pytest @@ -86,8 +87,8 @@ async def test_publish_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping.""" - mock_client.add.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.add.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.publish("test-agent-001", make_store_request()) async def test_publish_reraises_recursion_error( @@ -208,8 +209,8 @@ async def test_search_shared_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping.""" - mock_client.search.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.search.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.search_shared(MemoryQuery(text="test")) async def test_search_shared_with_category_post_filter( @@ -300,6 +301,7 @@ async def test_retract_success( "id": "shared-001", "memory": "shared content", "created_at": "2026-03-12T10:00:00+00:00", + "user_id": _SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.return_value = None @@ -329,6 +331,7 @@ async def test_retract_ownership_mismatch( "id": "shared-001", "memory": "content", "created_at": "2026-03-12T10:00:00+00:00", + "user_id": _SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-002"}, } @@ -344,6 +347,7 @@ async def test_retract_no_publisher_raises( "id": "not-shared-001", "memory": "private content", "created_at": "2026-03-12T10:00:00+00:00", + "user_id": _SHARED_NAMESPACE, "metadata": {}, } @@ -366,8 +370,8 @@ async def test_retract_reraises_memory_error( mock_client: MagicMock, ) -> None: """builtins.MemoryError is re-raised without wrapping.""" - mock_client.get.side_effect = MemoryError("out of memory") - with pytest.raises(MemoryError): + mock_client.get.side_effect = builtins.MemoryError("out of memory") + with pytest.raises(builtins.MemoryError): await backend.retract("test-agent-001", "shared-001") async def test_retract_reraises_recursion_error( @@ -380,6 +384,23 @@ async def test_retract_reraises_recursion_error( with pytest.raises(RecursionError): await backend.retract("test-agent-001", "shared-001") + async def test_retract_not_shared_namespace_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """retract() rejects memories not in the shared namespace.""" + mock_client.get.return_value = { + "id": "private-001", + "memory": "private content", + "created_at": "2026-03-12T10:00:00+00:00", + "user_id": "test-agent-001", + "metadata": {PUBLISHER_KEY: "test-agent-001"}, + } + + with pytest.raises(MemoryStoreError, match="not in the shared namespace"): + await backend.retract("test-agent-001", "private-001") + async def test_retract_delete_failure_wraps( self, backend: Mem0MemoryBackend, @@ -390,6 +411,7 @@ async def test_retract_delete_failure_wraps( "id": "shared-001", "memory": "content", "created_at": "2026-03-12T10:00:00+00:00", + "user_id": _SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.side_effect = RuntimeError("delete failed") diff --git a/tests/unit/memory/backends/mem0/test_config.py b/tests/unit/memory/backends/mem0/test_config.py index a85326a0ea..3f0c1fd188 100644 --- a/tests/unit/memory/backends/mem0/test_config.py +++ b/tests/unit/memory/backends/mem0/test_config.py @@ -217,13 +217,14 @@ def test_rejects_unsupported_history_store(self) -> None: embedder=_embedder(), ) - def test_accepts_qdrant_external(self) -> None: + def test_rejects_qdrant_external(self) -> None: + """qdrant-external is not supported — only embedded qdrant.""" company_config = CompanyMemoryConfig( backend="mem0", storage=MemoryStorageConfig(vector_store="qdrant-external"), ) - mem0_config = build_config_from_company_config( - company_config, - embedder=_embedder(), - ) - assert mem0_config is not None + with pytest.raises(ValueError, match="embedded qdrant"): + build_config_from_company_config( + company_config, + embedder=_embedder(), + ) diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index b28d6b5305..5b774fe960 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -153,6 +153,18 @@ def test_non_numeric_string_returns_none(self) -> None: def test_non_numeric_type_returns_none(self) -> None: assert normalize_relevance_score(object()) is None + def test_nan_returns_none(self) -> None: + assert normalize_relevance_score(float("nan")) is None + + def test_nan_string_returns_none(self) -> None: + assert normalize_relevance_score("nan") is None + + def test_inf_returns_none(self) -> None: + assert normalize_relevance_score(float("inf")) is None + + def test_neg_inf_returns_none(self) -> None: + assert normalize_relevance_score(float("-inf")) is None + @pytest.mark.unit class TestParseMem0Metadata: @@ -602,6 +614,14 @@ def test_object_type_returns_half(self) -> None: """Object type that can't convert to float defaults to 0.5.""" assert _coerce_confidence({f"{_PREFIX}confidence": object()}) == 0.5 + def test_nan_returns_half(self) -> None: + """NaN confidence defaults to 0.5.""" + assert _coerce_confidence({f"{_PREFIX}confidence": float("nan")}) == 0.5 + + def test_inf_returns_half(self) -> None: + """Inf confidence defaults to 0.5.""" + assert _coerce_confidence({f"{_PREFIX}confidence": float("inf")}) == 0.5 + @pytest.mark.unit class TestNormalizeTags: From 6b69d770b6a8164cb864cf44ddb4f5740f22eebc Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 11:05:19 +0100 Subject: [PATCH 14/17] =?UTF-8?q?fix:=20address=20round-8=20PR=20review=20?= =?UTF-8?q?findings=20=E2=80=94=20orphan=20rejection,=20vendor=20names,=20?= =?UTF-8?q?pathlib?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - get(): return None for memories with no user_id (unverifiable ownership) - delete(): raise MemoryStoreError for orphan memories instead of allowing - search_shared(): reject exclude_agent matching reserved shared namespace - config: replace vendor-specific docstring examples with generic identifiers - config: use PurePosixPath for Docker-targeted path construction - factory: catch ValidationError alongside ValueError in config build - mappers: improved created_at fallback cascade (updated_at → expires_at → now) - tests: add orphan rejection tests (get/delete), namespace guard test - tests: freeze clock in expiry test, parameterize validate_add_result tests --- .../memory/backends/mem0/adapter.py | 24 ++++-- src/ai_company/memory/backends/mem0/config.py | 10 ++- .../memory/backends/mem0/mappers.py | 24 ++++-- src/ai_company/memory/factory.py | 2 +- .../memory/backends/mem0/test_adapter_crud.py | 38 +++++++++- .../backends/mem0/test_adapter_shared.py | 14 ++++ .../unit/memory/backends/mem0/test_mappers.py | 76 ++++++++----------- 7 files changed, 119 insertions(+), 69 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index e9677a65c8..0e3c3a2424 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -552,11 +552,12 @@ async def get( MEMORY_ENTRY_FETCHED, agent_id=agent_id, memory_id=memory_id, - found=True, + found=False, reason="memory has no user_id — ownership " - "unverifiable, returning to requesting agent", + "unverifiable, refusing to return", ) - elif str(owner) != str(agent_id): + return None + if str(owner) != str(agent_id): logger.debug( MEMORY_ENTRY_FETCHED, agent_id=agent_id, @@ -636,14 +637,17 @@ async def delete( # Block deletion of shared-namespace entries — use retract(). owner = existing.get("user_id") if owner is None: + msg = ( + f"Memory {memory_id} has no user_id — ownership " + f"unverifiable, refusing deletion" + ) logger.warning( - MEMORY_ENTRY_DELETED, + MEMORY_ENTRY_DELETE_FAILED, agent_id=agent_id, memory_id=memory_id, - unverifiable_ownership=True, - reason="memory has no user_id — ownership " - "unverifiable, proceeding with delete", + reason="unverifiable_ownership", ) + raise MemoryStoreError(msg) # noqa: TRY301 if owner is not None and str(owner) == _SHARED_NAMESPACE: msg = ( f"Memory {memory_id} belongs to the shared namespace — " @@ -870,6 +874,12 @@ async def search_shared( MemoryRetrievalError: If the search fails. """ self._require_connected() + if exclude_agent is not None and str(exclude_agent) == _SHARED_NAMESPACE: + msg = ( + "exclude_agent must not be the reserved shared namespace: " + f"{_SHARED_NAMESPACE!r}" + ) + raise MemoryRetrievalError(msg) try: if query.text is not None: raw_result = await asyncio.to_thread( diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 6ff77c7aa6..696a4119db 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -26,8 +26,9 @@ class Mem0EmbedderConfig(BaseModel): Both ``provider`` and ``model`` are required — callers must supply them explicitly so that vendor-specific identifiers stay out of source defaults. The values must be valid Mem0 SDK - identifiers (e.g. ``"openai"``, ``"text-embedding-ada-002"``); - see the Mem0 documentation for supported providers and models. + identifiers (e.g. ``"example-provider"``, + ``"example-medium-001"``); see the Mem0 documentation for + supported providers and models. Attributes: provider: Embedding provider name (Mem0 SDK identifier). @@ -106,13 +107,14 @@ def build_mem0_config_dict(config: Mem0BackendConfig) -> dict[str, Any]: Returns: Configuration dict suitable for ``Memory.from_config()``. """ + base_path = PurePosixPath(config.data_dir) return { "vector_store": { "provider": "qdrant", "config": { "collection_name": config.collection_name, "embedding_model_dims": config.embedder.dims, - "path": f"{config.data_dir}/qdrant", + "path": str(base_path / "qdrant"), }, }, "embedder": { @@ -121,7 +123,7 @@ def build_mem0_config_dict(config: Mem0BackendConfig) -> dict[str, Any]: "model": config.embedder.model, }, }, - "history_db_path": f"{config.data_dir}/history.db", + "history_db_path": str(base_path / "history.db"), # Mem0 config schema version — required by Memory.from_config(). "version": "v1.1", } diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 5fda6621d9..572e3ac61c 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -288,12 +288,23 @@ def mem0_result_to_entry( created_at = parse_mem0_datetime(raw.get("created_at")) updated_at = parse_mem0_datetime(raw.get("updated_at")) + + raw_metadata = raw.get("metadata") + category, metadata, expires_at = parse_mem0_metadata(raw_metadata) + if created_at is None: - # Prefer updated_at as a fallback — it is a closer - # approximation than now() and avoids violating the - # MemoryEntry invariant (updated_at >= created_at). - fallback = updated_at or datetime.now(UTC) - fallback_source = "updated_at" if updated_at else "now()" + # Pick the best available fallback to avoid violating the + # MemoryEntry invariants (updated_at >= created_at, + # expires_at >= created_at). + if updated_at is not None: + fallback = updated_at + fallback_source = "updated_at" + elif expires_at is not None: + fallback = expires_at + fallback_source = "expires_at" + else: + fallback = datetime.now(UTC) + fallback_source = "now()" logger.warning( MEMORY_MODEL_INVALID, field="created_at", @@ -303,9 +314,6 @@ def mem0_result_to_entry( ) created_at = fallback - raw_metadata = raw.get("metadata") - category, metadata, expires_at = parse_mem0_metadata(raw_metadata) - raw_score = raw.get("score") relevance_score = normalize_relevance_score(raw_score) diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index 6c1d181f1a..ac1b6340e0 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -79,7 +79,7 @@ def _create_mem0_backend( config, embedder=embedder, ) - except ValueError as exc: + except (ValueError, ValidationError) as exc: msg = f"Invalid Mem0 configuration: {exc}" logger.warning( MEMORY_BACKEND_CONFIG_INVALID, diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index 0783421b0f..f0eeae58d0 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -302,7 +302,10 @@ async def test_get_existing( backend: Mem0MemoryBackend, mock_client: MagicMock, ) -> None: - mock_client.get.return_value = mem0_get_result("mem-001") + mock_client.get.return_value = mem0_get_result( + "mem-001", + user_id="test-agent-001", + ) entry = await backend.get("test-agent-001", "mem-001") @@ -373,6 +376,17 @@ async def test_get_ownership_mismatch_returns_none( entry = await backend.get("test-agent-001", "mem-001") assert entry is None + async def test_get_orphan_returns_none( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """get() returns None when memory has no user_id (orphan).""" + mock_client.get.return_value = mem0_get_result("mem-001") + + entry = await backend.get("test-agent-001", "mem-001") + assert entry is None + # ── Delete ──────────────────────────────────────────────────────── @@ -384,7 +398,10 @@ async def test_delete_existing( backend: Mem0MemoryBackend, mock_client: MagicMock, ) -> None: - mock_client.get.return_value = mem0_get_result("mem-001") + mock_client.get.return_value = mem0_get_result( + "mem-001", + user_id="test-agent-001", + ) mock_client.delete.return_value = None result = await backend.delete("test-agent-001", "mem-001") @@ -419,7 +436,10 @@ async def test_delete_get_ok_but_delete_fails( backend: Mem0MemoryBackend, mock_client: MagicMock, ) -> None: - mock_client.get.return_value = mem0_get_result("mem-001") + mock_client.get.return_value = mem0_get_result( + "mem-001", + user_id="test-agent-001", + ) mock_client.delete.side_effect = RuntimeError("delete failed") with pytest.raises(MemoryStoreError, match="Failed to delete"): @@ -473,6 +493,18 @@ async def test_delete_ownership_mismatch_raises( with pytest.raises(MemoryStoreError, match="cannot delete"): await backend.delete("test-agent-001", "mem-001") + async def test_delete_orphan_raises( + self, + backend: Mem0MemoryBackend, + mock_client: MagicMock, + ) -> None: + """delete() raises when memory has no user_id (orphan).""" + mock_client.get.return_value = mem0_get_result("mem-001") + + with pytest.raises(MemoryStoreError, match="unverifiable"): + await backend.delete("test-agent-001", "mem-001") + mock_client.delete.assert_not_called() + # ── Count ───────────────────────────────────────────────────────── diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py index c8aa457a45..22aa719d67 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_shared.py +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -254,6 +254,20 @@ async def test_search_shared_with_category_post_filter( assert len(entries) == 1 assert entries[0].category == MemoryCategory.SEMANTIC + async def test_search_shared_rejects_shared_namespace_exclude( + self, + backend: Mem0MemoryBackend, + ) -> None: + """search_shared() rejects exclude_agent == shared namespace.""" + with pytest.raises( + MemoryRetrievalError, + match="reserved shared namespace", + ): + await backend.search_shared( + MemoryQuery(text="test"), + exclude_agent=_SHARED_NAMESPACE, + ) + async def test_search_shared_reraises_recursion_error( self, backend: Mem0MemoryBackend, diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index 5b774fe960..52ddbf2d69 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -2,6 +2,7 @@ from datetime import UTC, datetime, timedelta from typing import Any +from unittest.mock import patch import pytest @@ -425,11 +426,16 @@ def test_expired_entries_excluded(self) -> None: def test_exactly_expired_entry_excluded(self) -> None: """Entry with expires_at == now is excluded (<=).""" - now = datetime.now(UTC) - past = now - timedelta(days=7) - entries = (_make_entry(memory_id="m1", created_at=past, expires_at=now),) + fixed_now = datetime(2026, 3, 13, 12, 0, 0, tzinfo=UTC) + past = fixed_now - timedelta(days=7) + entries = (_make_entry(memory_id="m1", created_at=past, expires_at=fixed_now),) query = MemoryQuery() - result = apply_post_filters(entries, query) + with patch( + "ai_company.memory.backends.mem0.mappers.datetime", + ) as mock_dt: + mock_dt.now.return_value = fixed_now + mock_dt.side_effect = lambda *a, **kw: datetime(*a, **kw) # noqa: DTZ001, PLW0108 + result = apply_post_filters(entries, query) assert len(result) == 0 def test_combined_filters(self) -> None: @@ -465,39 +471,26 @@ def test_valid_result(self) -> None: memory_id = validate_add_result(result, context="test") assert memory_id == "mem-001" - def test_empty_results_raises(self) -> None: - result: dict[str, Any] = {"results": []} - with pytest.raises(MemoryStoreError, match="no results"): - validate_add_result(result, context="test") - - def test_missing_results_key_raises(self) -> None: - result = {"data": "something"} - with pytest.raises(MemoryStoreError, match="no results"): - validate_add_result(result, context="test") - - def test_non_list_results_raises(self) -> None: - result = {"results": "not-a-list"} - with pytest.raises(MemoryStoreError, match="no results"): - validate_add_result(result, context="test") - - def test_missing_id_raises(self) -> None: - result = {"results": [{"memory": "no id"}]} - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - validate_add_result(result, context="test") - - def test_none_id_raises(self) -> None: - result = {"results": [{"id": None, "event": "ADD"}]} - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - validate_add_result(result, context="test") - - def test_blank_id_raises(self) -> None: - result = {"results": [{"id": "", "event": "ADD"}]} - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): - validate_add_result(result, context="test") - - def test_whitespace_only_id_raises(self) -> None: - result = {"results": [{"id": " ", "event": "ADD"}]} - with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): + @pytest.mark.parametrize( + ("result", "expected_match"), + [ + ({"results": []}, "no results"), + ({"data": "something"}, "no results"), + ({"results": "not-a-list"}, "no results"), + ({"results": [{"memory": "no id"}]}, "missing or blank 'id'"), + ({"results": [{"id": None, "event": "ADD"}]}, "missing or blank 'id'"), + ({"results": [{"id": "", "event": "ADD"}]}, "missing or blank 'id'"), + ({"results": [{"id": " ", "event": "ADD"}]}, "missing or blank 'id'"), + ("not-a-dict", "unexpected type"), + ({"results": ["not-a-dict"]}, "not a dict"), + ], + ) + def test_malformed_result_raises( + self, + result: Any, + expected_match: str, + ) -> None: + with pytest.raises(MemoryStoreError, match=expected_match): validate_add_result(result, context="test") def test_numeric_id_coerced_to_string(self) -> None: @@ -505,15 +498,6 @@ def test_numeric_id_coerced_to_string(self) -> None: memory_id = validate_add_result(result, context="test") assert memory_id == "42" - def test_non_dict_result_raises(self) -> None: - with pytest.raises(MemoryStoreError, match="unexpected type"): - validate_add_result("not-a-dict", context="test") - - def test_non_dict_first_item_raises(self) -> None: - result: dict[str, Any] = {"results": ["not-a-dict"]} - with pytest.raises(MemoryStoreError, match="not a dict"): - validate_add_result(result, context="test") - @pytest.mark.unit class TestExtractCategory: From fc9fcb79324853a98961757272f6e9c3fd765b22 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 11:21:55 +0100 Subject: [PATCH 15/17] fix: type-safe client, idiomatic patterns, and expanded test coverage Pre-reviewed by 8 agents, 11 findings addressed: - _require_connected() now returns Mem0Client for type-safe call sites - Replaced truthiness checks with explicit `is not None` in post-filters - Replaced dynamic log_fn dispatch with explicit branches in mappers - Added DEBUG logging for dropped tags in _normalize_tags - Extracted _resolve_publisher helper, replaced intermediate list with generator - Added logging to search_shared namespace guard - Removed dead owner-is-not-None conditions in delete() - Widened factory catch from (ValueError, ValidationError) to Exception - Added tests: orphan rejection, shared namespace guards, fallback dates --- .../memory/backends/mem0/adapter.py | 98 +++++++++++-------- .../memory/backends/mem0/mappers.py | 33 +++++-- src/ai_company/memory/factory.py | 8 +- .../memory/backends/mem0/test_adapter_crud.py | 8 ++ .../backends/mem0/test_adapter_shared.py | 16 +++ .../unit/memory/backends/mem0/test_mappers.py | 28 ++++++ 6 files changed, 138 insertions(+), 53 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 0e3c3a2424..d8c549aba3 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -175,6 +175,22 @@ def _validate_mem0_result( return raw_list +def _resolve_publisher(item: dict[str, Any]) -> str: + """Extract publisher from a shared memory, defaulting to namespace. + + Logs at DEBUG when publisher metadata is missing. + """ + publisher = extract_publisher(item) + if publisher is None: + logger.debug( + MEMORY_SHARED_SEARCHED, + memory_id=item.get("id", "?"), + reason="no publisher metadata — attributing to shared namespace", + ) + return _SHARED_NAMESPACE + return publisher + + class Mem0MemoryBackend: """Mem0-backed agent memory backend. @@ -194,7 +210,7 @@ def __init__( ) -> None: self._mem0_config = mem0_config self._max_memories_per_agent = max_memories_per_agent - self._client: Any = None + self._client: Mem0Client | None = None self._connected = False self._connect_lock = asyncio.Lock() @@ -348,8 +364,12 @@ def max_memories_per_agent(self) -> int | None: # ── Guards ──────────────────────────────────────────────────── - def _require_connected(self) -> None: - """Raise ``MemoryConnectionError`` if not connected.""" + def _require_connected(self) -> Mem0Client: + """Return the client or raise ``MemoryConnectionError``. + + Returns: + The connected Mem0 client (enables mypy type narrowing). + """ if not self._connected or self._client is None: logger.warning( MEMORY_BACKEND_NOT_CONNECTED, @@ -357,6 +377,7 @@ def _require_connected(self) -> None: ) msg = "Not connected — call connect() first" raise MemoryConnectionError(msg) + return self._client def _validate_agent_id( self, @@ -411,7 +432,7 @@ async def store( MemoryConnectionError: If the backend is not connected. MemoryStoreError: If the store operation fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id) try: kwargs = { @@ -422,7 +443,7 @@ async def store( "metadata": build_mem0_metadata(request), "infer": False, } - result = await asyncio.to_thread(self._client.add, **kwargs) + result = await asyncio.to_thread(client.add, **kwargs) memory_id = validate_add_result(result, context="store") except MemoryStoreError as exc: logger.warning( @@ -474,15 +495,15 @@ async def retrieve( MemoryConnectionError: If the backend is not connected. MemoryRetrievalError: If the retrieval fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id, error_cls=MemoryRetrievalError) try: if query.text is not None: kwargs = query_to_mem0_search_args(str(agent_id), query) - raw_result = await asyncio.to_thread(self._client.search, **kwargs) + raw_result = await asyncio.to_thread(client.search, **kwargs) else: kwargs = query_to_mem0_getall_args(str(agent_id), query) - raw_result = await asyncio.to_thread(self._client.get_all, **kwargs) + raw_result = await asyncio.to_thread(client.get_all, **kwargs) raw_list = _validate_mem0_result(raw_result, context="retrieve") entries = tuple(mem0_result_to_entry(item, agent_id) for item in raw_list) entries = apply_post_filters(entries, query) @@ -534,10 +555,10 @@ async def get( MemoryConnectionError: If the backend is not connected. MemoryRetrievalError: If the backend query fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id, error_cls=MemoryRetrievalError) try: - raw = await asyncio.to_thread(self._client.get, str(memory_id)) + raw = await asyncio.to_thread(client.get, str(memory_id)) if raw is None: logger.debug( MEMORY_ENTRY_FETCHED, @@ -620,12 +641,12 @@ async def delete( MemoryStoreError: If the delete operation fails or ownership verification fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id) try: # Check existence first — Mem0 delete doesn't indicate # whether the entry existed. - existing = await asyncio.to_thread(self._client.get, str(memory_id)) + existing = await asyncio.to_thread(client.get, str(memory_id)) if existing is None: logger.debug( MEMORY_ENTRY_DELETED, @@ -648,7 +669,7 @@ async def delete( reason="unverifiable_ownership", ) raise MemoryStoreError(msg) # noqa: TRY301 - if owner is not None and str(owner) == _SHARED_NAMESPACE: + if str(owner) == _SHARED_NAMESPACE: msg = ( f"Memory {memory_id} belongs to the shared namespace — " f"use retract() to remove shared entries" @@ -661,7 +682,7 @@ async def delete( ) raise MemoryStoreError(msg) # noqa: TRY301 # Verify ownership — reject cross-agent deletion. - if owner is not None and str(owner) != str(agent_id): + if str(owner) != str(agent_id): msg = ( f"Agent {agent_id} cannot delete memory " f"{memory_id} owned by {owner}" @@ -674,7 +695,7 @@ async def delete( actual_owner=str(owner), ) raise MemoryStoreError(msg) # noqa: TRY301 - await asyncio.to_thread(self._client.delete, str(memory_id)) + await asyncio.to_thread(client.delete, str(memory_id)) except MemoryStoreError: raise except builtins.MemoryError, RecursionError: @@ -731,11 +752,11 @@ async def count( MemoryConnectionError: If the backend is not connected. MemoryRetrievalError: If the count query fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id, error_cls=MemoryRetrievalError) try: raw_result = await asyncio.to_thread( - self._client.get_all, + client.get_all, user_id=str(agent_id), limit=self._max_memories_per_agent, ) @@ -810,7 +831,7 @@ async def publish( MemoryConnectionError: If the backend is not connected. MemoryStoreError: If the publish operation fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id) try: metadata = { @@ -825,7 +846,7 @@ async def publish( "metadata": metadata, "infer": False, } - result = await asyncio.to_thread(self._client.add, **kwargs) + result = await asyncio.to_thread(client.add, **kwargs) memory_id = validate_add_result(result, context="shared publish") except MemoryStoreError as exc: logger.warning( @@ -873,24 +894,29 @@ async def search_shared( MemoryConnectionError: If the backend is not connected. MemoryRetrievalError: If the search fails. """ - self._require_connected() + client = self._require_connected() if exclude_agent is not None and str(exclude_agent) == _SHARED_NAMESPACE: msg = ( "exclude_agent must not be the reserved shared namespace: " f"{_SHARED_NAMESPACE!r}" ) + logger.warning( + MEMORY_BACKEND_AGENT_ID_REJECTED, + agent_id=exclude_agent, + reason="reserved shared namespace used as exclude_agent", + ) raise MemoryRetrievalError(msg) try: if query.text is not None: raw_result = await asyncio.to_thread( - self._client.search, + client.search, query=str(query.text), user_id=_SHARED_NAMESPACE, limit=query.limit, ) else: raw_result = await asyncio.to_thread( - self._client.get_all, + client.get_all, user_id=_SHARED_NAMESPACE, limit=query.limit, ) @@ -899,21 +925,15 @@ async def search_shared( context="search_shared", ) - entries_list: list[MemoryEntry] = [] - for item in raw_list: - publisher = extract_publisher(item) - if publisher is None: - logger.debug( - MEMORY_SHARED_SEARCHED, - memory_id=item.get("id", "?"), - reason="no publisher metadata — " - "attributing to shared namespace", - ) - publisher = _SHARED_NAMESPACE - entries_list.append( - mem0_result_to_entry(item, NotBlankStr(publisher)), + raw_entries = tuple( + mem0_result_to_entry( + item, + NotBlankStr( + _resolve_publisher(item), + ), ) - raw_entries = tuple(entries_list) + for item in raw_list + ) filtered = apply_post_filters(raw_entries, query) if exclude_agent is not None: @@ -968,10 +988,10 @@ async def retract( MemoryStoreError: If the retraction operation fails or ownership verification fails. """ - self._require_connected() + client = self._require_connected() self._validate_agent_id(agent_id) try: - raw = await asyncio.to_thread(self._client.get, str(memory_id)) + raw = await asyncio.to_thread(client.get, str(memory_id)) if raw is None: logger.debug( MEMORY_SHARED_RETRACTED, @@ -1025,7 +1045,7 @@ async def retract( ) raise MemoryStoreError(msg) # noqa: TRY301 - await asyncio.to_thread(self._client.delete, str(memory_id)) + await asyncio.to_thread(client.delete, str(memory_id)) except MemoryStoreError: # Ownership-check MemoryStoreErrors are already logged # with context (reason, publisher) above — re-raise diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 572e3ac61c..0fca5852fc 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -190,7 +190,18 @@ def _normalize_tags( reason="unexpected tags type, ignoring", ) raw_tags = () - return tuple(NotBlankStr(str(t)) for t in raw_tags if t and str(t).strip()) + valid: list[NotBlankStr] = [] + for t in raw_tags: + if t and str(t).strip(): + valid.append(NotBlankStr(str(t))) + else: + logger.debug( + MEMORY_MODEL_INVALID, + field="tags", + raw_value=t, + reason="blank or falsy tag dropped", + ) + return tuple(valid) def parse_mem0_metadata( @@ -205,13 +216,15 @@ def parse_mem0_metadata( Tuple of (category, metadata, expires_at). """ if not raw_metadata or not isinstance(raw_metadata, dict): - log_fn = logger.warning if raw_metadata is not None else logger.debug - log_fn( - MEMORY_MODEL_INVALID, - field="metadata", - raw_value=type(raw_metadata).__name__ if raw_metadata else None, - reason="missing or non-dict metadata, using defaults", - ) + log_kwargs = { + "field": "metadata", + "raw_value": type(raw_metadata).__name__ if raw_metadata else None, + "reason": "missing or non-dict metadata, using defaults", + } + if raw_metadata is not None: + logger.warning(MEMORY_MODEL_INVALID, **log_kwargs) + else: + logger.debug(MEMORY_MODEL_INVALID, **log_kwargs) return ( MemoryCategory.WORKING, MemoryMetadata(), @@ -412,9 +425,9 @@ def apply_post_filters( continue if query.tags and not all(tag in entry.metadata.tags for tag in query.tags): continue - if query.since and entry.created_at < query.since: + if query.since is not None and entry.created_at < query.since: continue - if query.until and entry.created_at >= query.until: + if query.until is not None and entry.created_at >= query.until: continue if ( query.min_relevance > 0.0 diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index ac1b6340e0..f58d666f1f 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -7,8 +7,6 @@ from typing import TYPE_CHECKING -from pydantic import ValidationError - from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 if TYPE_CHECKING: @@ -79,13 +77,14 @@ def _create_mem0_backend( config, embedder=embedder, ) - except (ValueError, ValidationError) as exc: + except Exception as exc: msg = f"Invalid Mem0 configuration: {exc}" logger.warning( MEMORY_BACKEND_CONFIG_INVALID, backend="mem0", reason="config_build_failed", error=msg, + error_type=type(exc).__name__, ) raise MemoryConfigError(msg) from exc try: @@ -93,13 +92,14 @@ def _create_mem0_backend( mem0_config=mem0_config, max_memories_per_agent=config.options.max_memories_per_agent, ) - except (ValueError, ValidationError) as exc: + except Exception as exc: msg = f"Failed to create Mem0 backend: {exc}" logger.warning( MEMORY_BACKEND_CONFIG_INVALID, backend="mem0", reason="backend_init_failed", error=msg, + error_type=type(exc).__name__, ) raise MemoryConfigError(msg) from exc logger.info( diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index f0eeae58d0..025bf8836e 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -505,6 +505,14 @@ async def test_delete_orphan_raises( await backend.delete("test-agent-001", "mem-001") mock_client.delete.assert_not_called() + async def test_delete_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """delete() rejects the shared namespace as agent_id.""" + with pytest.raises(MemoryStoreError, match="reserved shared namespace"): + await backend.delete(_SHARED_NAMESPACE, "mem-001") + # ── Count ───────────────────────────────────────────────────────── diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py index 22aa719d67..d0aabc35ca 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_shared.py +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -91,6 +91,14 @@ async def test_publish_reraises_memory_error( with pytest.raises(builtins.MemoryError): await backend.publish("test-agent-001", make_store_request()) + async def test_publish_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """publish() rejects the shared namespace as agent_id.""" + with pytest.raises(MemoryStoreError, match="reserved shared namespace"): + await backend.publish(_SHARED_NAMESPACE, make_store_request()) + async def test_publish_reraises_recursion_error( self, backend: Mem0MemoryBackend, @@ -388,6 +396,14 @@ async def test_retract_reraises_memory_error( with pytest.raises(builtins.MemoryError): await backend.retract("test-agent-001", "shared-001") + async def test_retract_rejects_shared_namespace_agent_id( + self, + backend: Mem0MemoryBackend, + ) -> None: + """retract() rejects the shared namespace as agent_id.""" + with pytest.raises(MemoryStoreError, match="reserved shared namespace"): + await backend.retract(_SHARED_NAMESPACE, "shared-001") + async def test_retract_reraises_recursion_error( self, backend: Mem0MemoryBackend, diff --git a/tests/unit/memory/backends/mem0/test_mappers.py b/tests/unit/memory/backends/mem0/test_mappers.py index 52ddbf2d69..edcb2369ef 100644 --- a/tests/unit/memory/backends/mem0/test_mappers.py +++ b/tests/unit/memory/backends/mem0/test_mappers.py @@ -180,6 +180,13 @@ def test_empty_metadata(self) -> None: assert category == MemoryCategory.WORKING assert metadata.confidence == 1.0 + def test_non_dict_truthy_metadata(self) -> None: + """Non-dict truthy metadata (e.g. string) uses defaults.""" + category, metadata, expires_at = parse_mem0_metadata("not-a-dict") # type: ignore[arg-type] + assert category == MemoryCategory.WORKING + assert metadata.confidence == 1.0 + assert expires_at is None + def test_full_metadata(self) -> None: raw = { f"{_PREFIX}category": "semantic", @@ -269,6 +276,27 @@ def test_missing_created_at_uses_now(self) -> None: assert before <= entry.created_at <= after + def test_missing_created_at_falls_back_to_updated_at(self) -> None: + """When created_at is missing, falls back to updated_at.""" + raw = { + "id": "no-created", + "memory": "content", + "updated_at": "2026-03-10T08:00:00+00:00", + "metadata": {}, + } + entry = mem0_result_to_entry(raw, "test-agent-001") + assert entry.created_at == datetime(2026, 3, 10, 8, 0, 0, tzinfo=UTC) + + def test_missing_created_at_falls_back_to_expires_at(self) -> None: + """When created_at and updated_at are missing, falls back to expires_at.""" + raw = { + "id": "no-created", + "memory": "content", + "metadata": {"_synthorg_expires_at": "2026-04-01T00:00:00+00:00"}, + } + entry = mem0_result_to_entry(raw, "test-agent-001") + assert entry.created_at == datetime(2026, 4, 1, 0, 0, 0, tzinfo=UTC) + def test_no_metadata(self) -> None: raw = { "id": "no-meta", From 155d2636336d3c81777311f345b5a3cc8aa99527 Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 12:55:46 +0100 Subject: [PATCH 16/17] =?UTF-8?q?fix:=20address=20round-9=20PR=20review=20?= =?UTF-8?q?findings=20=E2=80=94=20module=20extraction,=20system=20error=20?= =?UTF-8?q?guards,=20docs?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Extract SharedKnowledgeStore methods from adapter.py to shared.py (797 < 800 line limit) - Move validation helpers (validate_mem0_result, resolve_publisher, check_delete_ownership) to mappers.py - Rename _SHARED_NAMESPACE to SHARED_NAMESPACE (public constant in mappers) - Add MemoryError/RecursionError guards with logger.exception() in adapter, shared, factory - Add MEMORY_BACKEND_SYSTEM_ERROR event constant - Add Windows path rejection (backslashes, drive letters) to config validator - Update CLAUDE.md event exemplar list - Update docs/design/memory.md config comments (hardcoded stores) - Update all test imports for renamed/moved symbols --- CLAUDE.md | 2 +- docs/design/memory.md | 8 +- .../memory/backends/mem0/adapter.py | 471 ++++-------------- src/ai_company/memory/backends/mem0/config.py | 21 +- .../memory/backends/mem0/mappers.py | 257 ++++++++-- src/ai_company/memory/backends/mem0/shared.py | 343 +++++++++++++ src/ai_company/memory/factory.py | 18 + src/ai_company/observability/events/memory.py | 1 + tests/integration/memory/test_mem0_backend.py | 9 +- .../memory/backends/mem0/test_adapter_crud.py | 34 +- .../backends/mem0/test_adapter_shared.py | 27 +- 11 files changed, 729 insertions(+), 462 deletions(-) create mode 100644 src/ai_company/memory/backends/mem0/shared.py diff --git a/CLAUDE.md b/CLAUDE.md index 201443436a..b0bb0bed72 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -127,7 +127,7 @@ src/ai_company/ - **Every module** with business logic MUST have: `from ai_company.observability import get_logger` then `logger = get_logger(__name__)` - **Never** use `import logging` / `logging.getLogger()` / `print()` in application code - **Variable name**: always `logger` (not `_logger`, not `log`) -- **Event names**: always use constants from the domain-specific module under `ai_company.observability.events` (e.g. `PROVIDER_CALL_START` from `events.provider`, `BUDGET_RECORD_ADDED` from `events.budget`, `CFO_ANOMALY_DETECTED` from `events.cfo`, `CONFLICT_DETECTED` from `events.conflict`, `MEETING_STARTED` from `events.meeting`, `CLASSIFICATION_START` from `events.classification`, `CONSOLIDATION_START` from `events.consolidation`, `ORG_MEMORY_QUERY_START` from `events.org_memory`, `API_REQUEST_STARTED` from `events.api`, `CODE_RUNNER_EXECUTE_START` from `events.code_runner`, `DOCKER_EXECUTE_START` from `events.docker`, `MCP_INVOKE_START` from `events.mcp`, `SECURITY_EVALUATE_START` from `events.security`, `HR_HIRING_REQUEST_CREATED` from `events.hr`, `PERF_METRIC_RECORDED` from `events.performance`, `TRUST_EVALUATE_START` from `events.trust`, `PROMOTION_EVALUATE_START` from `events.promotion`, `PROMPT_BUILD_START` from `events.prompt`, `MEMORY_RETRIEVAL_START` from `events.memory`, `AUTONOMY_ACTION_AUTO_APPROVED` from `events.autonomy`, `TIMEOUT_POLICY_EVALUATED` from `events.timeout`, `PERSISTENCE_AUDIT_ENTRY_SAVED` from `events.persistence`, `TASK_ENGINE_STARTED` from `events.task_engine`, `COORDINATION_STARTED` from `events.coordination`, `COMMUNICATION_DISPATCH_START` from `events.communication`, `COMPANY_STARTED` from `events.company`, `CONFIG_LOADED` from `events.config`, `CORRELATION_ID_CREATED` from `events.correlation`, `DECOMPOSITION_STARTED` from `events.decomposition`, `DELEGATION_STARTED` from `events.delegation`, `EXECUTION_LOOP_STARTED` from `events.execution`, `GIT_OPERATION_START` from `events.git`, `PARALLEL_EXECUTION_STARTED` from `events.parallel`, `PERSONALITY_LOADED` from `events.personality`, `QUOTA_CHECKED` from `events.quota`, `ROLE_ASSIGNED` from `events.role`, `ROUTING_STARTED` from `events.routing`, `SANDBOX_EXECUTE_START` from `events.sandbox`, `TASK_CREATED` from `events.task`, `TASK_ASSIGNMENT_STARTED` from `events.task_assignment`, `TASK_ROUTING_STARTED` from `events.task_routing`, `TEMPLATE_LOADED` from `events.template`, `TOOL_INVOKE_START` from `events.tool`, `WORKSPACE_CREATED` from `events.workspace`). Import directly: `from ai_company.observability.events. import EVENT_CONSTANT` +- **Event names**: always use constants from the domain-specific module under `ai_company.observability.events` (e.g. `PROVIDER_CALL_START` from `events.provider`, `BUDGET_RECORD_ADDED` from `events.budget`, `CFO_ANOMALY_DETECTED` from `events.cfo`, `CONFLICT_DETECTED` from `events.conflict`, `MEETING_STARTED` from `events.meeting`, `CLASSIFICATION_START` from `events.classification`, `CONSOLIDATION_START` from `events.consolidation`, `ORG_MEMORY_QUERY_START` from `events.org_memory`, `API_REQUEST_STARTED` from `events.api`, `CODE_RUNNER_EXECUTE_START` from `events.code_runner`, `DOCKER_EXECUTE_START` from `events.docker`, `MCP_INVOKE_START` from `events.mcp`, `SECURITY_EVALUATE_START` from `events.security`, `HR_HIRING_REQUEST_CREATED` from `events.hr`, `PERF_METRIC_RECORDED` from `events.performance`, `TRUST_EVALUATE_START` from `events.trust`, `PROMOTION_EVALUATE_START` from `events.promotion`, `PROMPT_BUILD_START` from `events.prompt`, `MEMORY_RETRIEVAL_START` from `events.memory`, `MEMORY_BACKEND_CONNECTED` from `events.memory`, `MEMORY_ENTRY_STORED` from `events.memory`, `MEMORY_BACKEND_SYSTEM_ERROR` from `events.memory`, `AUTONOMY_ACTION_AUTO_APPROVED` from `events.autonomy`, `TIMEOUT_POLICY_EVALUATED` from `events.timeout`, `PERSISTENCE_AUDIT_ENTRY_SAVED` from `events.persistence`, `TASK_ENGINE_STARTED` from `events.task_engine`, `COORDINATION_STARTED` from `events.coordination`, `COMMUNICATION_DISPATCH_START` from `events.communication`, `COMPANY_STARTED` from `events.company`, `CONFIG_LOADED` from `events.config`, `CORRELATION_ID_CREATED` from `events.correlation`, `DECOMPOSITION_STARTED` from `events.decomposition`, `DELEGATION_STARTED` from `events.delegation`, `EXECUTION_LOOP_STARTED` from `events.execution`, `GIT_OPERATION_START` from `events.git`, `PARALLEL_EXECUTION_STARTED` from `events.parallel`, `PERSONALITY_LOADED` from `events.personality`, `QUOTA_CHECKED` from `events.quota`, `ROLE_ASSIGNED` from `events.role`, `ROUTING_STARTED` from `events.routing`, `SANDBOX_EXECUTE_START` from `events.sandbox`, `TASK_CREATED` from `events.task`, `TASK_ASSIGNMENT_STARTED` from `events.task_assignment`, `TASK_ROUTING_STARTED` from `events.task_routing`, `TEMPLATE_LOADED` from `events.template`, `TOOL_INVOKE_START` from `events.tool`, `WORKSPACE_CREATED` from `events.workspace`). Import directly: `from ai_company.observability.events. import EVENT_CONSTANT` - **Structured kwargs**: always `logger.info(EVENT, key=value)` — never `logger.info("msg %s", val)` - **All error paths** must log at WARNING or ERROR with context before raising - **All state transitions** must log at INFO diff --git a/docs/design/memory.md b/docs/design/memory.md index 7972a78396..39d116be3b 100644 --- a/docs/design/memory.md +++ b/docs/design/memory.md @@ -64,8 +64,8 @@ Memory persistence is configurable per agent, from no persistence to fully persi backend: "mem0" # mem0 | custom | cognee | graphiti (future) -- see Decision Log storage: data_dir: "/data/memory" # mounted Docker volume path - vector_store: "qdrant" # qdrant (embedded), qdrant-external, etc. - history_store: "sqlite" # sqlite, postgresql + vector_store: "qdrant" # hardcoded to embedded qdrant in Mem0 backend + history_store: "sqlite" # hardcoded to sqlite in Mem0 backend options: retention_days: null # null = forever max_memories_per_agent: 10000 @@ -292,8 +292,8 @@ memory: level: "persistent" # none, session, project, persistent (default: session) storage: data_dir: "/data/memory" - vector_store: "qdrant" - history_store: "sqlite" + vector_store: "qdrant" # hardcoded to embedded qdrant in Mem0 backend + history_store: "sqlite" # hardcoded to sqlite in Mem0 backend options: retention_days: null # null = forever max_memories_per_agent: 10000 diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index d8c549aba3..95bca88a52 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -1,20 +1,10 @@ """Mem0 memory backend adapter. -Implements ``MemoryBackend``, ``MemoryCapabilities``, and -``SharedKnowledgeStore`` protocols using Mem0 as the storage layer -(default: Qdrant (embedded by default) + SQLite). +Implements ``MemoryBackend`` and ``MemoryCapabilities`` protocols. +``SharedKnowledgeStore`` methods delegate to ``shared.py``. -All Mem0 calls are synchronous — they run in ``asyncio.to_thread()`` -to avoid blocking the event loop. - -All methods re-raise ``builtins.MemoryError`` and ``RecursionError`` -immediately without wrapping, to avoid masking system-level failures. - -Note: This file exceeds the 800-line guideline because the single -``Mem0MemoryBackend`` class implements three protocols cohesively -(``MemoryBackend``, ``MemoryCapabilities``, ``SharedKnowledgeStore``). -Splitting would fragment the unified client lifecycle and connection -guard logic. +All Mem0 SDK calls run in ``asyncio.to_thread()``. +``builtins.MemoryError`` / ``RecursionError`` re-raise immediately. """ import asyncio @@ -28,15 +18,21 @@ build_mem0_config_dict, ) from ai_company.memory.backends.mem0.mappers import ( - PUBLISHER_KEY, + SHARED_NAMESPACE, apply_post_filters, build_mem0_metadata, + check_delete_ownership, extract_category, - extract_publisher, mem0_result_to_entry, query_to_mem0_getall_args, query_to_mem0_search_args, validate_add_result, + validate_mem0_result, +) +from ai_company.memory.backends.mem0.shared import ( + publish_shared, + retract_shared, + search_shared_memories, ) from ai_company.memory.errors import ( MemoryConnectionError, @@ -56,6 +52,7 @@ MEMORY_BACKEND_DISCONNECTING, MEMORY_BACKEND_HEALTH_CHECK, MEMORY_BACKEND_NOT_CONNECTED, + MEMORY_BACKEND_SYSTEM_ERROR, MEMORY_ENTRY_COUNT_FAILED, MEMORY_ENTRY_COUNTED, MEMORY_ENTRY_DELETE_FAILED, @@ -66,12 +63,6 @@ MEMORY_ENTRY_RETRIEVED, MEMORY_ENTRY_STORE_FAILED, MEMORY_ENTRY_STORED, - MEMORY_SHARED_PUBLISH_FAILED, - MEMORY_SHARED_PUBLISHED, - MEMORY_SHARED_RETRACT_FAILED, - MEMORY_SHARED_RETRACTED, - MEMORY_SHARED_SEARCH_FAILED, - MEMORY_SHARED_SEARCHED, ) if TYPE_CHECKING: @@ -84,112 +75,17 @@ ) class Mem0Client(Protocol): - """Structural type for the Mem0 ``Memory`` client. - - Defines the subset of ``Memory`` methods that the adapter - uses, so the rest of the codebase does not depend on - ``Any``. - """ - - def add(self, **kwargs: Any) -> dict[str, Any]: - """Add a memory entry.""" - ... + """Subset of ``Memory`` methods used by the adapter.""" - def search(self, **kwargs: Any) -> dict[str, Any]: - """Search memories.""" - ... - - def get_all(self, **kwargs: Any) -> dict[str, Any]: - """Get all memories for a user.""" - ... - - def get(self, memory_id: str) -> dict[str, Any] | None: - """Get a single memory by ID.""" - ... - - def delete(self, memory_id: str) -> None: - """Delete a memory by ID.""" - ... + def add(self, **kwargs: Any) -> dict[str, Any]: ... # noqa: D102 + def search(self, **kwargs: Any) -> dict[str, Any]: ... # noqa: D102 + def get_all(self, **kwargs: Any) -> dict[str, Any]: ... # noqa: D102 + def get(self, memory_id: str) -> dict[str, Any] | None: ... # noqa: D102 + def delete(self, memory_id: str) -> None: ... # noqa: D102 logger = get_logger(__name__) -# Reserved user_id for the shared knowledge namespace. -# All shared memories are stored under this Mem0 ``user_id`` so they -# are isolated from per-agent memories and can be queried centrally. -_SHARED_NAMESPACE: str = "__synthorg_shared__" - - -def _validate_mem0_result( - raw_result: Any, - *, - context: str, -) -> list[dict[str, Any]]: - """Validate and extract the results list from a Mem0 response. - - Args: - raw_result: Raw return value from a Mem0 SDK call. - context: Human-readable context for error messages. - - Returns: - The ``"results"`` list from the response. - - Raises: - MemoryRetrievalError: If the response is not a dict or - ``"results"`` is not a list. - """ - if not isinstance(raw_result, dict): - msg = ( - f"Unexpected Mem0 response type for {context}: " - f"{type(raw_result).__name__}, expected dict" - ) - logger.warning( - MEMORY_ENTRY_RETRIEVAL_FAILED, - context=context, - error=msg, - ) - raise MemoryRetrievalError(msg) - if "results" not in raw_result: - msg = ( - f"Mem0 response missing 'results' key for {context}: " - f"keys={list(raw_result.keys())}" - ) - logger.warning( - MEMORY_ENTRY_RETRIEVAL_FAILED, - context=context, - error=msg, - ) - raise MemoryRetrievalError(msg) - raw_list = raw_result["results"] - if not isinstance(raw_list, list): - msg = ( - f"Unexpected Mem0 results type for {context}: " - f"{type(raw_list).__name__}, expected list" - ) - logger.warning( - MEMORY_ENTRY_RETRIEVAL_FAILED, - context=context, - error=msg, - ) - raise MemoryRetrievalError(msg) - return raw_list - - -def _resolve_publisher(item: dict[str, Any]) -> str: - """Extract publisher from a shared memory, defaulting to namespace. - - Logs at DEBUG when publisher metadata is missing. - """ - publisher = extract_publisher(item) - if publisher is None: - logger.debug( - MEMORY_SHARED_SEARCHED, - memory_id=item.get("id", "?"), - reason="no publisher metadata — attributing to shared namespace", - ) - return _SHARED_NAMESPACE - return publisher - class Mem0MemoryBackend: """Mem0-backed agent memory backend. @@ -208,6 +104,9 @@ def __init__( mem0_config: Mem0BackendConfig, max_memories_per_agent: int = 10_000, ) -> None: + if max_memories_per_agent < 1: + msg = f"max_memories_per_agent must be >= 1, got {max_memories_per_agent}" + raise ValueError(msg) self._mem0_config = mem0_config self._max_memories_per_agent = max_memories_per_agent self._client: Mem0Client | None = None @@ -250,7 +149,13 @@ async def connect(self) -> None: try: config_dict = build_mem0_config_dict(self._mem0_config) client = await asyncio.to_thread(Memory.from_config, config_dict) - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="connect", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -274,6 +179,13 @@ async def disconnect(self) -> None: in-progress ``connect()`` call. """ async with self._connect_lock: + if not self._connected and self._client is None: + logger.debug( + MEMORY_BACKEND_DISCONNECTED, + backend="mem0", + reason="already disconnected — no-op", + ) + return logger.info(MEMORY_BACKEND_DISCONNECTING, backend="mem0") self._client = None self._connected = False @@ -299,10 +211,16 @@ async def health_check(self) -> bool: try: await asyncio.to_thread( self._client.get_all, - user_id=_SHARED_NAMESPACE, + user_id=SHARED_NAMESPACE, limit=1, ) - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="health_check", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -396,11 +314,11 @@ def _validate_agent_id( Raises: MemoryStoreError: If ``agent_id`` collides with - ``_SHARED_NAMESPACE`` (default). + ``SHARED_NAMESPACE`` (default). MemoryRetrievalError: If ``error_cls`` was set to ``MemoryRetrievalError``. """ - if str(agent_id) == _SHARED_NAMESPACE: + if str(agent_id) == SHARED_NAMESPACE: logger.warning( MEMORY_BACKEND_AGENT_ID_REJECTED, agent_id=agent_id, @@ -408,7 +326,7 @@ def _validate_agent_id( ) msg = ( f"agent_id must not be the reserved shared namespace: " - f"{_SHARED_NAMESPACE!r}" + f"{SHARED_NAMESPACE!r}" ) raise error_cls(msg) @@ -453,7 +371,13 @@ async def store( error_type="MemoryStoreError", ) raise - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="store", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -504,7 +428,7 @@ async def retrieve( else: kwargs = query_to_mem0_getall_args(str(agent_id), query) raw_result = await asyncio.to_thread(client.get_all, **kwargs) - raw_list = _validate_mem0_result(raw_result, context="retrieve") + raw_list = validate_mem0_result(raw_result, context="retrieve") entries = tuple(mem0_result_to_entry(item, agent_id) for item in raw_list) entries = apply_post_filters(entries, query) except MemoryRetrievalError as exc: @@ -515,7 +439,13 @@ async def retrieve( error_type="MemoryRetrievalError", ) raise - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="retrieve", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -579,7 +509,7 @@ async def get( ) return None if str(owner) != str(agent_id): - logger.debug( + logger.info( MEMORY_ENTRY_FETCHED, agent_id=agent_id, memory_id=memory_id, @@ -598,7 +528,13 @@ async def get( error_type="MemoryRetrievalError", ) raise - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="get", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -644,8 +580,6 @@ async def delete( client = self._require_connected() self._validate_agent_id(agent_id) try: - # Check existence first — Mem0 delete doesn't indicate - # whether the entry existed. existing = await asyncio.to_thread(client.get, str(memory_id)) if existing is None: logger.debug( @@ -655,50 +589,17 @@ async def delete( found=False, ) return False - # Block deletion of shared-namespace entries — use retract(). - owner = existing.get("user_id") - if owner is None: - msg = ( - f"Memory {memory_id} has no user_id — ownership " - f"unverifiable, refusing deletion" - ) - logger.warning( - MEMORY_ENTRY_DELETE_FAILED, - agent_id=agent_id, - memory_id=memory_id, - reason="unverifiable_ownership", - ) - raise MemoryStoreError(msg) # noqa: TRY301 - if str(owner) == _SHARED_NAMESPACE: - msg = ( - f"Memory {memory_id} belongs to the shared namespace — " - f"use retract() to remove shared entries" - ) - logger.warning( - MEMORY_ENTRY_DELETE_FAILED, - agent_id=agent_id, - memory_id=memory_id, - reason="shared namespace entry", - ) - raise MemoryStoreError(msg) # noqa: TRY301 - # Verify ownership — reject cross-agent deletion. - if str(owner) != str(agent_id): - msg = ( - f"Agent {agent_id} cannot delete memory " - f"{memory_id} owned by {owner}" - ) - logger.warning( - MEMORY_ENTRY_DELETE_FAILED, - agent_id=agent_id, - memory_id=memory_id, - reason="ownership mismatch", - actual_owner=str(owner), - ) - raise MemoryStoreError(msg) # noqa: TRY301 + check_delete_ownership(existing, agent_id, memory_id) await asyncio.to_thread(client.delete, str(memory_id)) except MemoryStoreError: raise - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="delete", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -760,7 +661,7 @@ async def count( user_id=str(agent_id), limit=self._max_memories_per_agent, ) - raw_list = _validate_mem0_result(raw_result, context="count") + raw_list = validate_mem0_result(raw_result, context="count") if category is None: total = len(raw_list) else: @@ -775,7 +676,13 @@ async def count( error_type="MemoryRetrievalError", ) raise - except builtins.MemoryError, RecursionError: + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="count", + error=str(exc), + error_type=type(exc).__name__, + ) raise except Exception as exc: logger.warning( @@ -809,6 +716,9 @@ async def count( return total # ── SharedKnowledgeStore ────────────────────────────────────── + # Implementations live in shared.py to keep this file under + # the 800-line guideline. These methods validate preconditions + # (connection, agent ID) and delegate to the standalone functions. async def publish( self, @@ -817,9 +727,6 @@ async def publish( ) -> NotBlankStr: """Publish a memory to the shared knowledge store. - Uses a reserved namespace (``__synthorg_shared__``) and - records the publisher in metadata for ownership tracking. - Args: agent_id: Publishing agent identifier. request: Memory content and metadata. @@ -833,47 +740,7 @@ async def publish( """ client = self._require_connected() self._validate_agent_id(agent_id) - try: - metadata = { - **build_mem0_metadata(request), - PUBLISHER_KEY: str(agent_id), - } - kwargs = { - "messages": [ - {"role": "user", "content": request.content}, - ], - "user_id": _SHARED_NAMESPACE, - "metadata": metadata, - "infer": False, - } - result = await asyncio.to_thread(client.add, **kwargs) - memory_id = validate_add_result(result, context="shared publish") - except MemoryStoreError as exc: - logger.warning( - MEMORY_SHARED_PUBLISH_FAILED, - agent_id=agent_id, - error=str(exc), - error_type="MemoryStoreError", - ) - raise - except builtins.MemoryError, RecursionError: - raise - except Exception as exc: - logger.warning( - MEMORY_SHARED_PUBLISH_FAILED, - agent_id=agent_id, - error=str(exc), - error_type=type(exc).__name__, - ) - msg = f"Failed to publish shared memory: {exc}" - raise MemoryStoreError(msg) from exc - else: - logger.info( - MEMORY_SHARED_PUBLISHED, - agent_id=agent_id, - memory_id=memory_id, - ) - return memory_id + return await publish_shared(client, agent_id, request) async def search_shared( self, @@ -895,77 +762,11 @@ async def search_shared( MemoryRetrievalError: If the search fails. """ client = self._require_connected() - if exclude_agent is not None and str(exclude_agent) == _SHARED_NAMESPACE: - msg = ( - "exclude_agent must not be the reserved shared namespace: " - f"{_SHARED_NAMESPACE!r}" - ) - logger.warning( - MEMORY_BACKEND_AGENT_ID_REJECTED, - agent_id=exclude_agent, - reason="reserved shared namespace used as exclude_agent", - ) - raise MemoryRetrievalError(msg) - try: - if query.text is not None: - raw_result = await asyncio.to_thread( - client.search, - query=str(query.text), - user_id=_SHARED_NAMESPACE, - limit=query.limit, - ) - else: - raw_result = await asyncio.to_thread( - client.get_all, - user_id=_SHARED_NAMESPACE, - limit=query.limit, - ) - raw_list = _validate_mem0_result( - raw_result, - context="search_shared", - ) - - raw_entries = tuple( - mem0_result_to_entry( - item, - NotBlankStr( - _resolve_publisher(item), - ), - ) - for item in raw_list - ) - filtered = apply_post_filters(raw_entries, query) - - if exclude_agent is not None: - filtered = tuple(e for e in filtered if e.agent_id != exclude_agent) - except MemoryRetrievalError as exc: - logger.warning( - MEMORY_SHARED_SEARCH_FAILED, - error=str(exc), - error_type="MemoryRetrievalError", - query_text=query.text, - exclude_agent=exclude_agent, - ) - raise - except builtins.MemoryError, RecursionError: - raise - except Exception as exc: - logger.warning( - MEMORY_SHARED_SEARCH_FAILED, - error=str(exc), - error_type=type(exc).__name__, - query_text=query.text, - exclude_agent=exclude_agent, - ) - msg = f"Failed to search shared knowledge: {exc}" - raise MemoryRetrievalError(msg) from exc - else: - logger.info( - MEMORY_SHARED_SEARCHED, - count=len(filtered), - exclude_agent=exclude_agent, - ) - return filtered + return await search_shared_memories( + client, + query, + exclude_agent=exclude_agent, + ) async def retract( self, @@ -990,84 +791,4 @@ async def retract( """ client = self._require_connected() self._validate_agent_id(agent_id) - try: - raw = await asyncio.to_thread(client.get, str(memory_id)) - if raw is None: - logger.debug( - MEMORY_SHARED_RETRACTED, - agent_id=agent_id, - memory_id=memory_id, - found=False, - ) - return False - - # Verify this memory belongs to the shared namespace. - owner_ns = raw.get("user_id") - if owner_ns != _SHARED_NAMESPACE: - logger.warning( - MEMORY_SHARED_RETRACT_FAILED, - agent_id=agent_id, - memory_id=memory_id, - reason="not in shared namespace", - actual_namespace=str(owner_ns), - ) - msg = ( - f"Memory {memory_id} is not in the shared namespace — " - f"use delete() to remove private entries" - ) - raise MemoryStoreError(msg) # noqa: TRY301 - - publisher = extract_publisher(raw) - if publisher is None: - logger.warning( - MEMORY_SHARED_RETRACT_FAILED, - agent_id=agent_id, - memory_id=memory_id, - reason="not a shared memory entry (no publisher)", - ) - msg = ( - f"Memory {memory_id} is not a shared memory entry " - f"(no publisher metadata)" - ) - raise MemoryStoreError(msg) # noqa: TRY301 - - if publisher != str(agent_id): - logger.warning( - MEMORY_SHARED_RETRACT_FAILED, - agent_id=agent_id, - memory_id=memory_id, - reason="ownership mismatch", - publisher=publisher, - ) - msg = ( - f"Agent {agent_id} cannot retract memory " - f"{memory_id} published by {publisher}" - ) - raise MemoryStoreError(msg) # noqa: TRY301 - - await asyncio.to_thread(client.delete, str(memory_id)) - except MemoryStoreError: - # Ownership-check MemoryStoreErrors are already logged - # with context (reason, publisher) above — re-raise - # without duplicate logging. - raise - except builtins.MemoryError, RecursionError: - raise - except Exception as exc: - logger.warning( - MEMORY_SHARED_RETRACT_FAILED, - agent_id=agent_id, - memory_id=memory_id, - error=str(exc), - error_type=type(exc).__name__, - ) - msg = f"Failed to retract shared memory {memory_id}: {exc}" - raise MemoryStoreError(msg) from exc - else: - logger.info( - MEMORY_SHARED_RETRACTED, - agent_id=agent_id, - memory_id=memory_id, - found=True, - ) - return True + return await retract_shared(client, agent_id, memory_id) diff --git a/src/ai_company/memory/backends/mem0/config.py b/src/ai_company/memory/backends/mem0/config.py index 696a4119db..1b20b939e6 100644 --- a/src/ai_company/memory/backends/mem0/config.py +++ b/src/ai_company/memory/backends/mem0/config.py @@ -76,7 +76,11 @@ class Mem0BackendConfig(BaseModel): @model_validator(mode="after") def _reject_traversal(self) -> Self: - """Reject parent-directory traversal to prevent path escapes. + """Reject parent-directory traversal and Windows paths. + + The Mem0 backend targets Linux/Docker containers where paths + must be POSIX. Windows-style paths (drive letters, backslashes) + are rejected to prevent accidental host-path leaks. Note: ``build_config_from_company_config`` passes ``data_dir`` from ``CompanyMemoryConfig``, so this check also protects @@ -95,6 +99,21 @@ def _reject_traversal(self) -> Self: reason=msg, ) raise ValueError(msg) + if "\\" in self.data_dir or ( + len(self.data_dir) >= 2 and self.data_dir[1] == ":" # noqa: PLR2004 # drive-letter check + ): + msg = ( + "data_dir must be a POSIX path (no backslashes or " + "drive letters) — the Mem0 backend targets Linux containers" + ) + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + field="data_dir", + value=self.data_dir, + reason=msg, + ) + raise ValueError(msg) return self diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 0fca5852fc..63574cc93f 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -11,7 +11,10 @@ from ai_company.core.enums import MemoryCategory from ai_company.core.types import NotBlankStr -from ai_company.memory.errors import MemoryRetrievalError, MemoryStoreError +from ai_company.memory.errors import ( + MemoryRetrievalError, + MemoryStoreError, +) from ai_company.memory.models import ( MemoryEntry, MemoryMetadata, @@ -20,6 +23,8 @@ ) from ai_company.observability import get_logger from ai_company.observability.events.memory import ( + MEMORY_ENTRY_DELETE_FAILED, + MEMORY_ENTRY_RETRIEVAL_FAILED, MEMORY_ENTRY_STORE_FAILED, MEMORY_MODEL_INVALID, ) @@ -34,7 +39,12 @@ # Metadata key to track who published a shared memory. # Public because the adapter module needs it for ownership tracking. -PUBLISHER_KEY: str = "_synthorg_publisher" +PUBLISHER_KEY: str = f"{_PREFIX}publisher" + +# Reserved user_id for the shared knowledge namespace. +# All shared memories are stored under this Mem0 ``user_id`` so they +# are isolated from per-agent memories and can be queried centrally. +SHARED_NAMESPACE: str = "__synthorg_shared__" def build_mem0_metadata(request: MemoryStoreRequest) -> dict[str, Any]: @@ -192,8 +202,9 @@ def _normalize_tags( raw_tags = () valid: list[NotBlankStr] = [] for t in raw_tags: - if t and str(t).strip(): - valid.append(NotBlankStr(str(t))) + stripped = str(t).strip() if t else "" + if stripped: + valid.append(NotBlankStr(stripped)) else: logger.debug( MEMORY_MODEL_INVALID, @@ -231,20 +242,8 @@ def parse_mem0_metadata( None, ) - category_str = raw_metadata.get(f"{_PREFIX}category") - if category_str: - try: - category = MemoryCategory(category_str) - except ValueError: - logger.warning( - MEMORY_MODEL_INVALID, - field="category", - raw_value=category_str, - reason="unrecognized category, defaulting to WORKING", - ) - category = MemoryCategory.WORKING - else: - category = MemoryCategory.WORKING + # Delegate to extract_category for consistent fallback logic. + category = extract_category({"metadata": raw_metadata}) confidence = _coerce_confidence(raw_metadata) source = _coerce_source(raw_metadata) @@ -261,6 +260,45 @@ def parse_mem0_metadata( return category, metadata, expires_at +def _resolve_created_at( + raw: dict[str, Any], + *, + updated_at: AwareDatetime | None, + expires_at: AwareDatetime | None, +) -> AwareDatetime: + """Pick the best fallback when ``created_at`` is missing. + + Uses the earliest available candidate to avoid violating + ``MemoryEntry`` invariants (``updated_at >= created_at``, + ``expires_at >= created_at``). + """ + candidates: list[datetime] = [] + if updated_at is not None: + candidates.append(updated_at) + if expires_at is not None: + candidates.append(expires_at) + if candidates: + fallback = min(candidates) + sources = [] + if updated_at is not None: + sources.append("updated_at") + if expires_at is not None: + sources.append("expires_at") + fallback_source = ( + f"min({', '.join(sources)})" if len(sources) > 1 else sources[0] + ) + else: + fallback = datetime.now(UTC) + fallback_source = "now()" + logger.warning( + MEMORY_MODEL_INVALID, + field="created_at", + memory_id=str(raw.get("id", "?")), + reason=f"missing or unparseable created_at, defaulting to {fallback_source}", + ) + return fallback + + def mem0_result_to_entry( raw: dict[str, Any], agent_id: NotBlankStr, @@ -306,26 +344,11 @@ def mem0_result_to_entry( category, metadata, expires_at = parse_mem0_metadata(raw_metadata) if created_at is None: - # Pick the best available fallback to avoid violating the - # MemoryEntry invariants (updated_at >= created_at, - # expires_at >= created_at). - if updated_at is not None: - fallback = updated_at - fallback_source = "updated_at" - elif expires_at is not None: - fallback = expires_at - fallback_source = "expires_at" - else: - fallback = datetime.now(UTC) - fallback_source = "now()" - logger.warning( - MEMORY_MODEL_INVALID, - field="created_at", - memory_id=str(raw.get("id", "?")), - reason=f"missing or unparseable created_at, " - f"defaulting to {fallback_source}", + created_at = _resolve_created_at( + raw, + updated_at=updated_at, + expires_at=expires_at, ) - created_at = fallback raw_score = raw.get("score") relevance_score = normalize_relevance_score(raw_score) @@ -344,7 +367,7 @@ def mem0_result_to_entry( def query_to_mem0_search_args( - agent_id: str, + agent_id: NotBlankStr, query: MemoryQuery, ) -> dict[str, Any]: """Convert a ``MemoryQuery`` to ``Memory.search()`` kwargs. @@ -370,13 +393,13 @@ def query_to_mem0_search_args( raise ValueError(msg) return { "query": query.text, - "user_id": agent_id, + "user_id": str(agent_id), "limit": query.limit, } def query_to_mem0_getall_args( - agent_id: str, + agent_id: NotBlankStr, query: MemoryQuery, ) -> dict[str, Any]: """Convert a ``MemoryQuery`` to ``Memory.get_all()`` kwargs. @@ -389,7 +412,7 @@ def query_to_mem0_getall_args( Dict of kwargs for ``Memory.get_all()``. """ return { - "user_id": agent_id, + "user_id": str(agent_id), "limit": query.limit, } @@ -417,6 +440,7 @@ def apply_post_filters( Filtered entries (order preserved). """ now = datetime.now(UTC) + pre_count = len(entries) result: list[MemoryEntry] = [] for entry in entries: if entry.expires_at is not None and entry.expires_at <= now: @@ -436,6 +460,22 @@ def apply_post_filters( ): continue result.append(entry) + post_count = len(result) + if pre_count > 0 and post_count == 0: + logger.warning( + MEMORY_MODEL_INVALID, + field="post_filter", + reason="all entries filtered out by post-filters", + pre_filter_count=pre_count, + ) + elif pre_count != post_count: + logger.debug( + MEMORY_MODEL_INVALID, + field="post_filter", + pre_filter_count=pre_count, + post_filter_count=post_count, + reason="entries filtered by post-filters", + ) return tuple(result) @@ -493,6 +533,12 @@ def extract_category(raw: dict[str, Any]) -> MemoryCategory: """ metadata = raw.get("metadata", {}) if not metadata or not isinstance(metadata, dict): + logger.debug( + MEMORY_MODEL_INVALID, + field="category", + raw_value=type(metadata).__name__ if metadata else None, + reason="missing or non-dict metadata, defaulting to WORKING", + ) return MemoryCategory.WORKING cat_str = metadata.get(f"{_PREFIX}category") if cat_str: @@ -507,10 +553,86 @@ def extract_category(raw: dict[str, Any]) -> MemoryCategory: "defaulting to WORKING", ) return MemoryCategory.WORKING + logger.debug( + MEMORY_MODEL_INVALID, + field="category", + reason="category key absent from metadata, defaulting to WORKING", + ) return MemoryCategory.WORKING -def extract_publisher(raw: dict[str, Any]) -> str | None: +def validate_mem0_result( + raw_result: Any, + *, + context: str, +) -> list[dict[str, Any]]: + """Validate and extract the results list from a Mem0 response. + + Args: + raw_result: Raw return value from a Mem0 SDK call. + context: Human-readable context for error messages. + + Returns: + The ``"results"`` list from the response. + + Raises: + MemoryRetrievalError: If the response is not a dict or + ``"results"`` is not a list. + """ + if not isinstance(raw_result, dict): + msg = ( + f"Unexpected Mem0 response type for {context}: " + f"{type(raw_result).__name__}, expected dict" + ) + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + context=context, + error=msg, + ) + raise MemoryRetrievalError(msg) + if "results" not in raw_result: + msg = ( + f"Mem0 response missing 'results' key for {context}: " + f"keys={list(raw_result.keys())}" + ) + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + context=context, + error=msg, + ) + raise MemoryRetrievalError(msg) + raw_list = raw_result["results"] + if not isinstance(raw_list, list): + msg = ( + f"Unexpected Mem0 results type for {context}: " + f"{type(raw_list).__name__}, expected list" + ) + logger.warning( + MEMORY_ENTRY_RETRIEVAL_FAILED, + context=context, + error=msg, + ) + raise MemoryRetrievalError(msg) + return raw_list + + +def resolve_publisher(item: dict[str, Any]) -> str: + """Extract publisher from a shared memory, defaulting to namespace. + + Logs at DEBUG when publisher metadata is missing. + """ + publisher = extract_publisher(item) + if publisher is None: + logger.debug( + MEMORY_MODEL_INVALID, + memory_id=item.get("id", "?"), + reason="no publisher metadata — attributing to shared namespace", + ) + return SHARED_NAMESPACE + return publisher + + +def extract_publisher(raw: dict[str, Any]) -> NotBlankStr | None: """Extract the publisher agent ID from a shared memory dict. Returns ``None`` if the publisher key is missing, non-dict @@ -523,4 +645,53 @@ def extract_publisher(raw: dict[str, Any]) -> str | None: if value is None: return None coerced = str(value).strip() - return coerced or None + return NotBlankStr(coerced) if coerced else None + + +def check_delete_ownership( + existing: dict[str, Any], + agent_id: NotBlankStr, + memory_id: NotBlankStr, +) -> None: + """Verify the caller owns this private memory entry. + + Raises: + MemoryStoreError: If ownership cannot be verified + (missing user_id, shared namespace entry, or + ownership mismatch). + """ + owner = existing.get("user_id") + if owner is None: + msg = ( + f"Memory {memory_id} has no user_id — ownership " + f"unverifiable, refusing deletion" + ) + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="unverifiable_ownership", + ) + raise MemoryStoreError(msg) + if str(owner) == SHARED_NAMESPACE: + msg = ( + f"Memory {memory_id} belongs to the shared namespace — " + f"use retract() to remove shared entries" + ) + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="shared namespace entry", + ) + raise MemoryStoreError(msg) + if str(owner) != str(agent_id): + msg = f"Agent {agent_id} cannot delete memory {memory_id} owned by {owner}" + logger.warning( + MEMORY_ENTRY_DELETE_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="ownership mismatch", + actual_owner=str(owner), + ) + raise MemoryStoreError(msg) diff --git a/src/ai_company/memory/backends/mem0/shared.py b/src/ai_company/memory/backends/mem0/shared.py new file mode 100644 index 0000000000..511a9de129 --- /dev/null +++ b/src/ai_company/memory/backends/mem0/shared.py @@ -0,0 +1,343 @@ +"""SharedKnowledgeStore operations for the Mem0 backend. + +Standalone async functions that implement the ``SharedKnowledgeStore`` +protocol methods. The ``Mem0MemoryBackend`` class delegates to these +functions after performing connection and agent-ID validation. + +Separated from ``adapter.py`` to keep individual modules under the +800-line guideline while maintaining a single cohesive backend package. +""" + +import asyncio +import builtins +from typing import TYPE_CHECKING, Any + +from ai_company.core.types import NotBlankStr +from ai_company.memory.backends.mem0.mappers import ( + PUBLISHER_KEY, + SHARED_NAMESPACE, + apply_post_filters, + build_mem0_metadata, + extract_publisher, + mem0_result_to_entry, + resolve_publisher, + validate_add_result, + validate_mem0_result, +) +from ai_company.memory.errors import ( + MemoryRetrievalError, + MemoryStoreError, +) +from ai_company.observability import get_logger +from ai_company.observability.events.memory import ( + MEMORY_BACKEND_AGENT_ID_REJECTED, + MEMORY_BACKEND_SYSTEM_ERROR, + MEMORY_SHARED_PUBLISH_FAILED, + MEMORY_SHARED_PUBLISHED, + MEMORY_SHARED_RETRACT_FAILED, + MEMORY_SHARED_RETRACTED, + MEMORY_SHARED_SEARCH_FAILED, + MEMORY_SHARED_SEARCHED, +) + +if TYPE_CHECKING: + from ai_company.memory.backends.mem0.adapter import Mem0Client + from ai_company.memory.models import MemoryEntry, MemoryQuery, MemoryStoreRequest + + +logger = get_logger(__name__) + + +def _check_retract_ownership( + raw: dict[str, Any], + agent_id: NotBlankStr, + memory_id: NotBlankStr, +) -> None: + """Verify the caller published this shared memory entry. + + Raises: + MemoryStoreError: If ownership cannot be verified. + """ + owner_ns = raw.get("user_id") + if owner_ns is None: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="unverifiable_ownership", + ) + msg = ( + f"Memory {memory_id} has no user_id — " + f"ownership unverifiable, refusing retraction" + ) + raise MemoryStoreError(msg) + if owner_ns != SHARED_NAMESPACE: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="not in shared namespace", + actual_namespace=str(owner_ns), + ) + msg = ( + f"Memory {memory_id} is not in the shared namespace — " + f"use delete() to remove private entries" + ) + raise MemoryStoreError(msg) + + publisher = extract_publisher(raw) + if publisher is None: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="not a shared memory entry (no publisher)", + ) + msg = f"Memory {memory_id} is not a shared memory entry (no publisher metadata)" + raise MemoryStoreError(msg) + + if publisher != str(agent_id): + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + reason="ownership mismatch", + publisher=publisher, + ) + msg = ( + f"Agent {agent_id} cannot retract memory " + f"{memory_id} published by {publisher}" + ) + raise MemoryStoreError(msg) + + +async def publish_shared( + client: Mem0Client, + agent_id: NotBlankStr, + request: MemoryStoreRequest, +) -> NotBlankStr: + """Publish a memory to the shared knowledge store. + + Args: + client: Connected Mem0 client. + agent_id: Publishing agent identifier. + request: Memory content and metadata. + + Returns: + The backend-assigned shared memory ID. + + Raises: + MemoryStoreError: If the publish operation fails. + """ + try: + metadata = { + **build_mem0_metadata(request), + PUBLISHER_KEY: str(agent_id), + } + kwargs = { + "messages": [ + {"role": "user", "content": request.content}, + ], + "user_id": SHARED_NAMESPACE, + "metadata": metadata, + "infer": False, + } + result = await asyncio.to_thread(client.add, **kwargs) + memory_id = validate_add_result(result, context="shared publish") + except MemoryStoreError as exc: + logger.warning( + MEMORY_SHARED_PUBLISH_FAILED, + agent_id=agent_id, + error=str(exc), + error_type="MemoryStoreError", + ) + raise + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="publish", + error=str(exc), + error_type=type(exc).__name__, + ) + raise + except Exception as exc: + logger.warning( + MEMORY_SHARED_PUBLISH_FAILED, + agent_id=agent_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to publish shared memory: {exc}" + raise MemoryStoreError(msg) from exc + else: + logger.info( + MEMORY_SHARED_PUBLISHED, + agent_id=agent_id, + memory_id=memory_id, + ) + return memory_id + + +async def search_shared_memories( + client: Mem0Client, + query: MemoryQuery, + *, + exclude_agent: NotBlankStr | None = None, +) -> tuple[MemoryEntry, ...]: + """Search the shared knowledge store across agents. + + Args: + client: Connected Mem0 client. + query: Search parameters. + exclude_agent: Optional agent ID to exclude from results. + + Returns: + Matching shared memory entries ordered by relevance. + + Raises: + MemoryRetrievalError: If the search fails. + """ + if exclude_agent is not None and str(exclude_agent) == SHARED_NAMESPACE: + msg = ( + "exclude_agent must not be the reserved shared namespace: " + f"{SHARED_NAMESPACE!r}" + ) + logger.warning( + MEMORY_BACKEND_AGENT_ID_REJECTED, + agent_id=exclude_agent, + reason="reserved shared namespace used as exclude_agent", + ) + raise MemoryRetrievalError(msg) + try: + if query.text is not None: + raw_result = await asyncio.to_thread( + client.search, + query=str(query.text), + user_id=SHARED_NAMESPACE, + limit=query.limit, + ) + else: + raw_result = await asyncio.to_thread( + client.get_all, + user_id=SHARED_NAMESPACE, + limit=query.limit, + ) + raw_list = validate_mem0_result( + raw_result, + context="search_shared", + ) + + raw_entries = tuple( + mem0_result_to_entry( + item, + NotBlankStr( + resolve_publisher(item), + ), + ) + for item in raw_list + ) + filtered = apply_post_filters(raw_entries, query) + + if exclude_agent is not None: + filtered = tuple(e for e in filtered if e.agent_id != exclude_agent) + except MemoryRetrievalError as exc: + logger.warning( + MEMORY_SHARED_SEARCH_FAILED, + error=str(exc), + error_type="MemoryRetrievalError", + query_text=query.text, + exclude_agent=exclude_agent, + ) + raise + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="search_shared", + error=str(exc), + error_type=type(exc).__name__, + ) + raise + except Exception as exc: + logger.warning( + MEMORY_SHARED_SEARCH_FAILED, + error=str(exc), + error_type=type(exc).__name__, + query_text=query.text, + exclude_agent=exclude_agent, + ) + msg = f"Failed to search shared knowledge: {exc}" + raise MemoryRetrievalError(msg) from exc + else: + logger.info( + MEMORY_SHARED_SEARCHED, + count=len(filtered), + exclude_agent=exclude_agent, + ) + return filtered + + +async def retract_shared( + client: Mem0Client, + agent_id: NotBlankStr, + memory_id: NotBlankStr, +) -> bool: + """Remove a memory from the shared knowledge store. + + Verifies publisher ownership before deletion. + + Args: + client: Connected Mem0 client. + agent_id: Retracting agent identifier. + memory_id: Shared memory identifier. + + Returns: + ``True`` if retracted, ``False`` if not found. + + Raises: + MemoryStoreError: If the retraction operation fails or + ownership verification fails. + """ + try: + raw = await asyncio.to_thread(client.get, str(memory_id)) + if raw is None: + logger.debug( + MEMORY_SHARED_RETRACTED, + agent_id=agent_id, + memory_id=memory_id, + found=False, + ) + return False + + _check_retract_ownership(raw, agent_id, memory_id) + await asyncio.to_thread(client.delete, str(memory_id)) + except MemoryStoreError: + # Ownership-check MemoryStoreErrors are already logged + # with context (reason, publisher) above — re-raise + # without duplicate logging. + raise + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="retract", + error=str(exc), + error_type=type(exc).__name__, + ) + raise + except Exception as exc: + logger.warning( + MEMORY_SHARED_RETRACT_FAILED, + agent_id=agent_id, + memory_id=memory_id, + error=str(exc), + error_type=type(exc).__name__, + ) + msg = f"Failed to retract shared memory {memory_id}: {exc}" + raise MemoryStoreError(msg) from exc + else: + logger.info( + MEMORY_SHARED_RETRACTED, + agent_id=agent_id, + memory_id=memory_id, + found=True, + ) + return True diff --git a/src/ai_company/memory/factory.py b/src/ai_company/memory/factory.py index f58d666f1f..ef03d4a9c9 100644 --- a/src/ai_company/memory/factory.py +++ b/src/ai_company/memory/factory.py @@ -5,6 +5,7 @@ ``config.backend``. """ +import builtins from typing import TYPE_CHECKING from ai_company.memory.config import CompanyMemoryConfig # noqa: TC001 @@ -17,6 +18,7 @@ from ai_company.observability.events.memory import ( MEMORY_BACKEND_CONFIG_INVALID, MEMORY_BACKEND_CREATED, + MEMORY_BACKEND_SYSTEM_ERROR, MEMORY_BACKEND_UNKNOWN, ) @@ -77,6 +79,14 @@ def _create_mem0_backend( config, embedder=embedder, ) + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="create_mem0_backend", + error=str(exc), + error_type=type(exc).__name__, + ) + raise except Exception as exc: msg = f"Invalid Mem0 configuration: {exc}" logger.warning( @@ -92,6 +102,14 @@ def _create_mem0_backend( mem0_config=mem0_config, max_memories_per_agent=config.options.max_memories_per_agent, ) + except (builtins.MemoryError, RecursionError) as exc: + logger.exception( + MEMORY_BACKEND_SYSTEM_ERROR, + operation="create_mem0_backend", + error=str(exc), + error_type=type(exc).__name__, + ) + raise except Exception as exc: msg = f"Failed to create Mem0 backend: {exc}" logger.warning( diff --git a/src/ai_company/observability/events/memory.py b/src/ai_company/observability/events/memory.py index 1d7570a3ff..19c56a093c 100644 --- a/src/ai_company/observability/events/memory.py +++ b/src/ai_company/observability/events/memory.py @@ -21,6 +21,7 @@ MEMORY_BACKEND_CONFIG_INVALID: Final[str] = "memory.backend.config_invalid" MEMORY_BACKEND_NOT_CONNECTED: Final[str] = "memory.backend.not_connected" MEMORY_BACKEND_AGENT_ID_REJECTED: Final[str] = "memory.backend.agent_id_rejected" +MEMORY_BACKEND_SYSTEM_ERROR: Final[str] = "memory.backend.system_error" # ── Entry operations ────────────────────────────────────────────── diff --git a/tests/integration/memory/test_mem0_backend.py b/tests/integration/memory/test_mem0_backend.py index c6a6317080..b2076825bb 100644 --- a/tests/integration/memory/test_mem0_backend.py +++ b/tests/integration/memory/test_mem0_backend.py @@ -11,15 +11,12 @@ import pytest from ai_company.core.enums import MemoryCategory -from ai_company.memory.backends.mem0.adapter import ( - _SHARED_NAMESPACE, - Mem0MemoryBackend, -) +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend from ai_company.memory.backends.mem0.config import ( Mem0BackendConfig, Mem0EmbedderConfig, ) -from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY +from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY, SHARED_NAMESPACE from ai_company.memory.models import MemoryQuery, MemoryStoreRequest from ai_company.memory.retrieval_config import MemoryRetrievalConfig from ai_company.memory.retriever import ContextInjectionStrategy @@ -216,7 +213,7 @@ async def test_shared_knowledge_flow( "id": "shared-001", "memory": "company policy: always test code", "created_at": datetime.now(UTC).isoformat(), - "user_id": _SHARED_NAMESPACE, + "user_id": SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.return_value = None diff --git a/tests/unit/memory/backends/mem0/test_adapter_crud.py b/tests/unit/memory/backends/mem0/test_adapter_crud.py index 025bf8836e..28fa185680 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_crud.py +++ b/tests/unit/memory/backends/mem0/test_adapter_crud.py @@ -6,10 +6,10 @@ import pytest from ai_company.core.enums import MemoryCategory -from ai_company.memory.backends.mem0.adapter import ( - _SHARED_NAMESPACE, - Mem0MemoryBackend, - _validate_mem0_result, +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend +from ai_company.memory.backends.mem0.mappers import ( + SHARED_NAMESPACE, + validate_mem0_result, ) from ai_company.memory.errors import ( MemoryRetrievalError, @@ -135,7 +135,7 @@ async def test_store_rejects_shared_namespace_agent_id( ) -> None: """Storing with the shared namespace agent ID is rejected.""" with pytest.raises(MemoryStoreError, match="reserved shared namespace"): - await backend.store(_SHARED_NAMESPACE, make_store_request()) + await backend.store(SHARED_NAMESPACE, make_store_request()) mock_client.add.assert_not_called() async def test_store_reraises_recursion_error( @@ -270,7 +270,7 @@ async def test_retrieve_rejects_shared_namespace_agent_id( """retrieve() rejects the shared namespace with MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="reserved shared namespace"): await backend.retrieve( - _SHARED_NAMESPACE, + SHARED_NAMESPACE, MemoryQuery(text="test"), ) @@ -360,7 +360,7 @@ async def test_get_rejects_shared_namespace_agent_id( ) -> None: """get() rejects the shared namespace with MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="reserved shared namespace"): - await backend.get(_SHARED_NAMESPACE, "mem-001") + await backend.get(SHARED_NAMESPACE, "mem-001") async def test_get_ownership_mismatch_returns_none( self, @@ -473,7 +473,7 @@ async def test_delete_shared_namespace_entry_raises( """delete() rejects entries belonging to the shared namespace.""" mock_client.get.return_value = mem0_get_result( "mem-001", - user_id=_SHARED_NAMESPACE, + user_id=SHARED_NAMESPACE, ) with pytest.raises(MemoryStoreError, match="shared namespace"): @@ -511,7 +511,7 @@ async def test_delete_rejects_shared_namespace_agent_id( ) -> None: """delete() rejects the shared namespace as agent_id.""" with pytest.raises(MemoryStoreError, match="reserved shared namespace"): - await backend.delete(_SHARED_NAMESPACE, "mem-001") + await backend.delete(SHARED_NAMESPACE, "mem-001") # ── Count ───────────────────────────────────────────────────────── @@ -599,7 +599,7 @@ async def test_count_rejects_shared_namespace_agent_id( ) -> None: """count() rejects the shared namespace with MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="reserved shared namespace"): - await backend.count(_SHARED_NAMESPACE) + await backend.count(SHARED_NAMESPACE) async def test_count_exception_wraps( self, @@ -659,7 +659,7 @@ async def test_count_reraises_recursion_error( await backend.count("test-agent-001") -# ── _validate_mem0_result ──────────────────────────────────────── +# ── validate_mem0_result ──────────────────────────────────────── @pytest.mark.unit @@ -667,30 +667,30 @@ class TestValidateMem0Result: def test_non_dict_raises(self) -> None: """Non-dict response raises MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="Unexpected Mem0 response type"): - _validate_mem0_result("not-a-dict", context="test") + validate_mem0_result("not-a-dict", context="test") def test_missing_results_key_raises(self) -> None: """Dict without 'results' key raises MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="missing 'results' key"): - _validate_mem0_result({"data": []}, context="test") + validate_mem0_result({"data": []}, context="test") def test_non_list_results_raises(self) -> None: """Non-list 'results' value raises MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="Unexpected Mem0 results type"): - _validate_mem0_result({"results": "not-a-list"}, context="test") + validate_mem0_result({"results": "not-a-list"}, context="test") def test_valid_response(self) -> None: """Valid response returns the results list.""" items = [{"id": "m1"}] - result = _validate_mem0_result({"results": items}, context="test") + result = validate_mem0_result({"results": items}, context="test") assert result == items def test_empty_results(self) -> None: """Empty results list is valid.""" - result = _validate_mem0_result({"results": []}, context="test") + result = validate_mem0_result({"results": []}, context="test") assert result == [] def test_none_raises(self) -> None: """None response raises MemoryRetrievalError.""" with pytest.raises(MemoryRetrievalError, match="Unexpected Mem0 response type"): - _validate_mem0_result(None, context="test") + validate_mem0_result(None, context="test") diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py index d0aabc35ca..ab1a7cc9ac 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_shared.py +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -5,11 +5,8 @@ import pytest -from ai_company.memory.backends.mem0.adapter import ( - _SHARED_NAMESPACE, - Mem0MemoryBackend, -) -from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY +from ai_company.memory.backends.mem0.adapter import Mem0MemoryBackend +from ai_company.memory.backends.mem0.mappers import PUBLISHER_KEY, SHARED_NAMESPACE from ai_company.memory.errors import ( MemoryRetrievalError, MemoryStoreError, @@ -44,7 +41,7 @@ async def test_publish_success( assert memory_id == "shared-mem-001" call_kwargs = mock_client.add.call_args[1] - assert call_kwargs["user_id"] == _SHARED_NAMESPACE + assert call_kwargs["user_id"] == SHARED_NAMESPACE assert PUBLISHER_KEY in call_kwargs["metadata"] assert call_kwargs["metadata"][PUBLISHER_KEY] == "test-agent-001" @@ -97,7 +94,7 @@ async def test_publish_rejects_shared_namespace_agent_id( ) -> None: """publish() rejects the shared namespace as agent_id.""" with pytest.raises(MemoryStoreError, match="reserved shared namespace"): - await backend.publish(_SHARED_NAMESPACE, make_store_request()) + await backend.publish(SHARED_NAMESPACE, make_store_request()) async def test_publish_reraises_recursion_error( self, @@ -142,7 +139,7 @@ async def test_search_shared_with_text( assert entries[0].agent_id == "test-agent-002" mock_client.search.assert_called_once() call_kwargs = mock_client.search.call_args[1] - assert call_kwargs["user_id"] == _SHARED_NAMESPACE + assert call_kwargs["user_id"] == SHARED_NAMESPACE async def test_search_shared_without_text( self, @@ -273,7 +270,7 @@ async def test_search_shared_rejects_shared_namespace_exclude( ): await backend.search_shared( MemoryQuery(text="test"), - exclude_agent=_SHARED_NAMESPACE, + exclude_agent=SHARED_NAMESPACE, ) async def test_search_shared_reraises_recursion_error( @@ -306,7 +303,7 @@ async def test_search_shared_no_publisher_uses_namespace( entries = await backend.search_shared(MemoryQuery(text="test")) assert len(entries) == 1 - assert entries[0].agent_id == _SHARED_NAMESPACE + assert entries[0].agent_id == SHARED_NAMESPACE # ── Retract ────────────────────────────────────────────────────── @@ -323,7 +320,7 @@ async def test_retract_success( "id": "shared-001", "memory": "shared content", "created_at": "2026-03-12T10:00:00+00:00", - "user_id": _SHARED_NAMESPACE, + "user_id": SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.return_value = None @@ -353,7 +350,7 @@ async def test_retract_ownership_mismatch( "id": "shared-001", "memory": "content", "created_at": "2026-03-12T10:00:00+00:00", - "user_id": _SHARED_NAMESPACE, + "user_id": SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-002"}, } @@ -369,7 +366,7 @@ async def test_retract_no_publisher_raises( "id": "not-shared-001", "memory": "private content", "created_at": "2026-03-12T10:00:00+00:00", - "user_id": _SHARED_NAMESPACE, + "user_id": SHARED_NAMESPACE, "metadata": {}, } @@ -402,7 +399,7 @@ async def test_retract_rejects_shared_namespace_agent_id( ) -> None: """retract() rejects the shared namespace as agent_id.""" with pytest.raises(MemoryStoreError, match="reserved shared namespace"): - await backend.retract(_SHARED_NAMESPACE, "shared-001") + await backend.retract(SHARED_NAMESPACE, "shared-001") async def test_retract_reraises_recursion_error( self, @@ -441,7 +438,7 @@ async def test_retract_delete_failure_wraps( "id": "shared-001", "memory": "content", "created_at": "2026-03-12T10:00:00+00:00", - "user_id": _SHARED_NAMESPACE, + "user_id": SHARED_NAMESPACE, "metadata": {PUBLISHER_KEY: "test-agent-001"}, } mock_client.delete.side_effect = RuntimeError("delete failed") From 1c20792714af3d5fbbe1ae3437c4bb25c29cc89a Mon Sep 17 00:00:00 2001 From: Aurelio <19254254+Aureliolo@users.noreply.github.com> Date: Fri, 13 Mar 2026 13:15:46 +0100 Subject: [PATCH 17/17] =?UTF-8?q?fix:=20address=20round-10=20review=20find?= =?UTF-8?q?ings=20=E2=80=94=20event=20misuse,=20constructor=20log,=20test?= =?UTF-8?q?=20parametrize?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../memory/backends/mem0/adapter.py | 8 +++ .../memory/backends/mem0/mappers.py | 5 +- .../backends/mem0/test_adapter_shared.py | 72 ++++++++----------- 3 files changed, 41 insertions(+), 44 deletions(-) diff --git a/src/ai_company/memory/backends/mem0/adapter.py b/src/ai_company/memory/backends/mem0/adapter.py index 95bca88a52..cfcf20b9ab 100644 --- a/src/ai_company/memory/backends/mem0/adapter.py +++ b/src/ai_company/memory/backends/mem0/adapter.py @@ -45,6 +45,7 @@ from ai_company.observability import get_logger from ai_company.observability.events.memory import ( MEMORY_BACKEND_AGENT_ID_REJECTED, + MEMORY_BACKEND_CONFIG_INVALID, MEMORY_BACKEND_CONNECTED, MEMORY_BACKEND_CONNECTING, MEMORY_BACKEND_CONNECTION_FAILED, @@ -106,6 +107,13 @@ def __init__( ) -> None: if max_memories_per_agent < 1: msg = f"max_memories_per_agent must be >= 1, got {max_memories_per_agent}" + logger.warning( + MEMORY_BACKEND_CONFIG_INVALID, + backend="mem0", + reason="invalid_max_memories_per_agent", + value=max_memories_per_agent, + error=msg, + ) raise ValueError(msg) self._mem0_config = mem0_config self._max_memories_per_agent = max_memories_per_agent diff --git a/src/ai_company/memory/backends/mem0/mappers.py b/src/ai_company/memory/backends/mem0/mappers.py index 63574cc93f..4d9b70dec8 100644 --- a/src/ai_company/memory/backends/mem0/mappers.py +++ b/src/ai_company/memory/backends/mem0/mappers.py @@ -26,6 +26,7 @@ MEMORY_ENTRY_DELETE_FAILED, MEMORY_ENTRY_RETRIEVAL_FAILED, MEMORY_ENTRY_STORE_FAILED, + MEMORY_FILTER_APPLIED, MEMORY_MODEL_INVALID, ) @@ -463,14 +464,14 @@ def apply_post_filters( post_count = len(result) if pre_count > 0 and post_count == 0: logger.warning( - MEMORY_MODEL_INVALID, + MEMORY_FILTER_APPLIED, field="post_filter", reason="all entries filtered out by post-filters", pre_filter_count=pre_count, ) elif pre_count != post_count: logger.debug( - MEMORY_MODEL_INVALID, + MEMORY_FILTER_APPLIED, field="post_filter", pre_filter_count=pre_count, post_filter_count=post_count, diff --git a/tests/unit/memory/backends/mem0/test_adapter_shared.py b/tests/unit/memory/backends/mem0/test_adapter_shared.py index ab1a7cc9ac..7b6a2d72c6 100644 --- a/tests/unit/memory/backends/mem0/test_adapter_shared.py +++ b/tests/unit/memory/backends/mem0/test_adapter_shared.py @@ -78,14 +78,20 @@ async def test_publish_missing_id_raises( with pytest.raises(MemoryStoreError, match="missing or blank 'id'"): await backend.publish("test-agent-001", make_store_request()) - async def test_publish_reraises_memory_error( + @pytest.mark.parametrize( + "exc_type", + [builtins.MemoryError, RecursionError], + ids=["MemoryError", "RecursionError"], + ) + async def test_publish_reraises_system_error( self, backend: Mem0MemoryBackend, mock_client: MagicMock, + exc_type: type[BaseException], ) -> None: - """builtins.MemoryError is re-raised without wrapping.""" - mock_client.add.side_effect = builtins.MemoryError("out of memory") - with pytest.raises(builtins.MemoryError): + """System errors are re-raised without wrapping.""" + mock_client.add.side_effect = exc_type("system failure") + with pytest.raises(exc_type): await backend.publish("test-agent-001", make_store_request()) async def test_publish_rejects_shared_namespace_agent_id( @@ -96,16 +102,6 @@ async def test_publish_rejects_shared_namespace_agent_id( with pytest.raises(MemoryStoreError, match="reserved shared namespace"): await backend.publish(SHARED_NAMESPACE, make_store_request()) - async def test_publish_reraises_recursion_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """RecursionError is re-raised without wrapping.""" - mock_client.add.side_effect = RecursionError("infinite loop") - with pytest.raises(RecursionError): - await backend.publish("test-agent-001", make_store_request()) - # ── SearchShared ───────────────────────────────────────────────── @@ -208,14 +204,20 @@ async def test_search_shared_exception_wraps( with pytest.raises(MemoryRetrievalError, match="Failed to search"): await backend.search_shared(MemoryQuery(text="test")) - async def test_search_shared_reraises_memory_error( + @pytest.mark.parametrize( + "exc_type", + [builtins.MemoryError, RecursionError], + ids=["MemoryError", "RecursionError"], + ) + async def test_search_shared_reraises_system_error( self, backend: Mem0MemoryBackend, mock_client: MagicMock, + exc_type: type[BaseException], ) -> None: - """builtins.MemoryError is re-raised without wrapping.""" - mock_client.search.side_effect = builtins.MemoryError("out of memory") - with pytest.raises(builtins.MemoryError): + """System errors are re-raised without wrapping.""" + mock_client.search.side_effect = exc_type("system failure") + with pytest.raises(exc_type): await backend.search_shared(MemoryQuery(text="test")) async def test_search_shared_with_category_post_filter( @@ -273,16 +275,6 @@ async def test_search_shared_rejects_shared_namespace_exclude( exclude_agent=SHARED_NAMESPACE, ) - async def test_search_shared_reraises_recursion_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """RecursionError is re-raised without wrapping.""" - mock_client.search.side_effect = RecursionError("infinite loop") - with pytest.raises(RecursionError): - await backend.search_shared(MemoryQuery(text="test")) - async def test_search_shared_no_publisher_uses_namespace( self, backend: Mem0MemoryBackend, @@ -383,14 +375,20 @@ async def test_retract_exception_wraps( with pytest.raises(MemoryStoreError, match="Failed to retract"): await backend.retract("test-agent-001", "shared-001") - async def test_retract_reraises_memory_error( + @pytest.mark.parametrize( + "exc_type", + [builtins.MemoryError, RecursionError], + ids=["MemoryError", "RecursionError"], + ) + async def test_retract_reraises_system_error( self, backend: Mem0MemoryBackend, mock_client: MagicMock, + exc_type: type[BaseException], ) -> None: - """builtins.MemoryError is re-raised without wrapping.""" - mock_client.get.side_effect = builtins.MemoryError("out of memory") - with pytest.raises(builtins.MemoryError): + """System errors are re-raised without wrapping.""" + mock_client.get.side_effect = exc_type("system failure") + with pytest.raises(exc_type): await backend.retract("test-agent-001", "shared-001") async def test_retract_rejects_shared_namespace_agent_id( @@ -401,16 +399,6 @@ async def test_retract_rejects_shared_namespace_agent_id( with pytest.raises(MemoryStoreError, match="reserved shared namespace"): await backend.retract(SHARED_NAMESPACE, "shared-001") - async def test_retract_reraises_recursion_error( - self, - backend: Mem0MemoryBackend, - mock_client: MagicMock, - ) -> None: - """RecursionError is re-raised without wrapping.""" - mock_client.get.side_effect = RecursionError("infinite loop") - with pytest.raises(RecursionError): - await backend.retract("test-agent-001", "shared-001") - async def test_retract_not_shared_namespace_raises( self, backend: Mem0MemoryBackend,