From 370ab791c88df0d57d2b5d625f23449d9396a183 Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Tue, 27 Jan 2026 16:30:03 -0800 Subject: [PATCH 1/9] Add milvus backend unit test --- AGENTS.md | 5 +- kaizen/backend/milvus.py | 176 ++++++++++---------- tests/unit/test_milvus_backend.py | 258 ++++++++++++++++++++++++++++++ 3 files changed, 345 insertions(+), 94 deletions(-) create mode 100644 tests/unit/test_milvus_backend.py diff --git a/AGENTS.md b/AGENTS.md index 79b4562c..ebfe12e6 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -47,13 +47,16 @@ cp .env.example .env # Configure any environment variables, defined in `./kaize pre-commit install ``` +## Development Tips +- This project is managed by `uv`, not `python` or `pip`, so any python commands need to go through `uv`. All dependencies are defined in `pyproject.toml`. + ## Testing Instructions - Run pytest verbosely with the `-v` flag by default so that you have more context when tests fail. - Use `uv run pytest tests/.../` to run tests individually. - We use the pytest markers `e2e` for end-to-end tests, and `unit` for unit tests, and `phoenix` to test integration with Phoenix. - When running `uv run pytest` it will skip the tests marked with `phoenix`. - To run specific markers: `uv run pytest -m e2e` or `uv run pytest -m unit` -- To override and run all: `uv run pytest -m "e2e or unit or phoenix"` +- To run all tests: `uv run pytest -m "e2e or unit or phoenix"` ## Available Interfaces - MCP Server: `get_guidelines()`, `save_trajectory()` diff --git a/kaizen/backend/milvus.py b/kaizen/backend/milvus.py index d5e2d49a..4ce91362 100644 --- a/kaizen/backend/milvus.py +++ b/kaizen/backend/milvus.py @@ -33,12 +33,8 @@ def deserialize_content(content: str): class MilvusEntityBackend(BaseEntityBackend): - # Removed class attributes - - def __init__(self, config=None): - super().__init__(config) - self.milvus = MilvusClient(**milvus_client_settings.model_dump()) - self.embedding_model = SentenceTransformer(milvus_other_settings.embedding_model) + milvus = MilvusClient(**milvus_client_settings.model_dump()) + embedding_model = SentenceTransformer(milvus_other_settings.embedding_model) def ready(self): _ = self.milvus.list_collections() @@ -48,15 +44,12 @@ def validate_namespace(self, namespace_id: str): if not self.milvus.has_collection(namespace_id): raise NamespaceNotFoundException(f"Namespace `{namespace_id}` not found") - def create_namespace( - self, - namespace_id: str | None = None - ) -> Namespace: + def create_namespace(self, namespace_id: str | None = None) -> Namespace: """Create a new namespace for entities to exist in.""" - namespace_id = namespace_id or 'ns_' + str(uuid.uuid4()).replace('-', '_') + namespace_id = namespace_id or "ns_" + str(uuid.uuid4()).replace("-", "_") if not self.milvus.has_collection(namespace_id): - self.milvus.create_collection(collection_name=namespace_id, dimension=768, auto_id=False, schema=entity_schema) + self.milvus.create_collection(collection_name=namespace_id, dimension=384, auto_id=False, schema=entity_schema) with SQLiteManager() as db_manager: return db_manager.create_namespace(namespace_id) @@ -66,17 +59,14 @@ def get_namespace_details(self, namespace_id: str) -> Namespace: with SQLiteManager() as db_manager: namespace = db_manager.get_namespace(namespace_id) - namespace.num_entities = self.milvus.get_collection_stats(namespace_id)['row_count'] + namespace.num_entities = self.milvus.get_collection_stats(namespace_id)["row_count"] return namespace - def search_namespaces( - self, - limit: int = 10 - ) -> list[Namespace]: + def search_namespaces(self, limit: int = 10) -> list[Namespace]: with SQLiteManager() as db_manager: namespaces = [] for namespace in db_manager.search_namespaces(limit): - namespace.num_entities = self.milvus.get_collection_stats(namespace.id)['row_count'] + namespace.num_entities = self.milvus.get_collection_stats(namespace.id)["row_count"] namespaces.append(namespace) return namespaces @@ -87,12 +77,7 @@ def delete_namespace(self, namespace_id: str): with SQLiteManager() as db_manager: db_manager.delete_namespace(namespace_id) - def update_entities( - self, - namespace_id: str, - entities: list[Entity], - enable_conflict_resolution: bool = True - ) -> list[EntityUpdate]: + def update_entities(self, namespace_id: str, entities: list[Entity], enable_conflict_resolution: bool = True) -> list[EntityUpdate]: self.validate_namespace(namespace_id) entity_type = entities[0].type if not all(entity.type == entity_type for entity in entities): @@ -103,13 +88,11 @@ def update_entities( entities_with_temporary_ids = [] for i, entity in enumerate(entities): entity_data = entity.model_dump() - if entity_data.get('metadata') is None: - entity_data['metadata'] = {} - entities_with_temporary_ids.append(RecordedEntity( - **entity_data, - created_at=datetime.datetime.now(datetime.UTC), - id=f'Unprocessed_Entity_{i}' - )) + if entity_data.get("metadata") is None: + entity_data["metadata"] = {} + entities_with_temporary_ids.append( + RecordedEntity(**entity_data, created_at=datetime.datetime.now(datetime.UTC), id=f"Unprocessed_Entity_{i}") + ) if enable_conflict_resolution: old_entities = [] @@ -121,56 +104,58 @@ def update_entities( for update in updates: content_str = serialize_content(update.content) match update.event: - case 'ADD': - entity_id = str(self.milvus.insert(collection_name=namespace_id, data={ - 'type': entity_type, - 'content': content_str, - 'created_at': int(now.timestamp()), - 'embedding': self.embedding_model.encode(content_str), - 'metadata': update.metadata, - })['ids'][0]) + case "ADD": + entity_id = str( + self.milvus.insert( + collection_name=namespace_id, + data={ + "type": entity_type, + "content": content_str, + "created_at": int(now.timestamp()), + "embedding": self.embedding_model.encode(content_str), + "metadata": update.metadata, + }, + )["ids"][0] + ) update.id = entity_id - case 'UPDATE': - self.milvus.upsert(collection_name=namespace_id, data={ - 'type': entity_type, - 'id': int(update.id), - 'content': content_str, - 'created_at': int(now.timestamp()), - 'embedding': self.embedding_model.encode(content_str), - 'metadata': update.metadata - }, partial_update=True) - case 'DELETE': + case "UPDATE": + self.milvus.upsert( + collection_name=namespace_id, + data={ + "type": entity_type, + "id": int(update.id), + "content": content_str, + "created_at": int(now.timestamp()), + "embedding": self.embedding_model.encode(content_str), + "metadata": update.metadata, + }, + partial_update=True, + ) + case "DELETE": self.delete_entity_by_id(namespace_id=namespace_id, entity_id=update.id) - case 'NONE': + case "NONE": pass else: updates = [] for entity in entities: content_str = serialize_content(entity.content) - # Convert None metadata to empty dict for Milvus compatibility - metadata = entity.metadata if entity.metadata is not None else {} - entity_id = str(self.milvus.insert(collection_name=namespace_id, data={ - 'type': entity_type, - 'content': content_str, - 'created_at': int(now.timestamp()), - 'embedding': self.embedding_model.encode(content_str), - 'metadata': metadata - })['ids'][0]) - updates.append(EntityUpdate( - id=entity_id, - type=entity_type, - content=entity.content, - event='ADD', - metadata=metadata - )) + entity_id = str( + self.milvus.insert( + collection_name=namespace_id, + data={ + "type": entity_type, + "content": content_str, + "created_at": int(now.timestamp()), + "embedding": self.embedding_model.encode(content_str), + "metadata": entity.metadata, + }, + )["ids"][0] + ) + updates.append(EntityUpdate(id=entity_id, type=entity_type, content=entity.content, event="ADD", metadata=entity.metadata)) return updates def search_entities( - self, - namespace_id: str, - query: str | None = None, - filters: dict | None = None, - limit: int = 10 + self, namespace_id: str, query: str | None = None, filters: dict | None = None, limit: int = 10 ) -> list[RecordedEntity]: self.validate_namespace(namespace_id) filters = filters or {} @@ -178,17 +163,16 @@ def search_entities( if query is None: results = self.milvus.query( collection_name=namespace_id, - filter=' AND '.join([f"{k} == '{v}'" for k, v in filters.items()]) if len(filters) > 0 else 'id > 0' + filter=" AND ".join([f"{k} == '{v}'" for k, v in filters.items()]) if len(filters) > 0 else "id > 0", ) else: - results = self.milvus.query( collection_name=namespace_id, - anns_field='embedding', + anns_field="embedding", data=[self.embedding_model.encode(query)], - filter=' AND '.join([f"{k} == '{v}'" for k, v in filters.items()]), + filter=" AND ".join([f"{k} == '{v}'" for k, v in filters.items()]), limit=limit, - search_params={"metric_type": "IP"} + search_params={"metric_type": "IP"}, ) return [parse_milvus_entity(i) for i in results] @@ -198,7 +182,7 @@ def delete_entity_by_id(self, namespace_id: str, entity_id: str): except ValueError: raise KaizenException(f"Invalid entity ID: {entity_id}. Entity IDs must be numeric.") self.validate_namespace(namespace_id) - + # Check if entity exists before deleting existing = self.milvus.query( collection_name=namespace_id, @@ -207,7 +191,7 @@ def delete_entity_by_id(self, namespace_id: str, entity_id: str): ) if not existing: raise KaizenException(f"Entity with ID {entity_id} not found in namespace {namespace_id}.") - + self.milvus.delete(collection_name=namespace_id, ids=[entity_id_int]) def close(self): @@ -218,20 +202,26 @@ def close(self): except Exception as e: logger.warning(f"Error closing Milvus client: {e}") -entity_schema = CollectionSchema(fields=[ - # Keep it as an INT64 or else you won't be able to list all entities. - FieldSchema(name='id', is_primary=True, auto_id=True, dtype=DataType.INT64, max_length=128), - FieldSchema(name='type', dtype=DataType.VARCHAR, max_length=128), - FieldSchema(name='content', dtype=DataType.VARCHAR, max_length=65535), - FieldSchema(name='created_at', dtype=DataType.INT64), - FieldSchema(name='embedding', dtype=DataType.FLOAT_VECTOR, dim=384), - FieldSchema(name='metadata', dtype=DataType.JSON), -]) + +entity_schema = CollectionSchema( + fields=[ + # Keep it as an INT64 or else you won't be able to list all entities. + FieldSchema(name="id", is_primary=True, auto_id=True, dtype=DataType.INT64, max_length=128), + FieldSchema(name="type", dtype=DataType.VARCHAR, max_length=128), + FieldSchema(name="content", dtype=DataType.VARCHAR, max_length=65535), + FieldSchema(name="created_at", dtype=DataType.INT64), + FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=384), + FieldSchema(name="metadata", dtype=DataType.JSON), + ] +) + def parse_milvus_entity(entity: dict) -> RecordedEntity: - return RecordedEntity.model_validate({ - **entity, - 'id': str(entity['id']), - 'content': deserialize_content(entity['content']), - 'created_at': datetime.datetime.fromtimestamp(entity['created_at'], datetime.UTC), - }) \ No newline at end of file + return RecordedEntity.model_validate( + { + **entity, + "id": str(entity["id"]), + "content": deserialize_content(entity["content"]), + "created_at": datetime.datetime.fromtimestamp(entity["created_at"], datetime.UTC), + } + ) diff --git a/tests/unit/test_milvus_backend.py b/tests/unit/test_milvus_backend.py new file mode 100644 index 00000000..c0caa41c --- /dev/null +++ b/tests/unit/test_milvus_backend.py @@ -0,0 +1,258 @@ +""" +Unit tests for MilvusEntityBackend. +Tests all methods with mocked Milvus client, SQLiteManager, and embedding model. +""" + +import datetime +import pytest +from unittest.mock import Mock, MagicMock, patch + +from kaizen.backend.milvus import MilvusEntityBackend +from kaizen.schema.core import Entity, Namespace, RecordedEntity +from kaizen.schema.conflict_resolution import EntityUpdate +from kaizen.schema.exceptions import NamespaceNotFoundException, KaizenException + + +@pytest.fixture(scope="module") +def milvus_backend() -> MilvusEntityBackend: + """Create a MilvusEntityBackend instance for testing.""" + with patch("kaizen.backend.milvus.MilvusClient"), patch("kaizen.backend.milvus.SentenceTransformer"): + backend = MilvusEntityBackend() + return backend + + +@pytest.fixture +def db_manager(): + """Create a mock SQLiteManager for testing.""" + + def create_namespace(namespace_id: str) -> Namespace: + return Namespace(id=namespace_id, created_at=datetime.datetime.now(datetime.UTC)) + + manager = MagicMock() + manager.__enter__ = Mock(return_value=manager) + manager.__exit__ = Mock(return_value=False) + manager.create_namespace = create_namespace + return manager + + +def always_has_collection(collection_name: str): + return True + + +def never_has_collection(collection_name: str): + return False + + +def noop_create_collection(collection_name: str, dimension, auto_id, schema): + pass + + +def arbitrary_namespace(namespace_id: str) -> Namespace: + return Namespace(id=namespace_id, created_at=datetime.datetime.now(datetime.UTC)) + + +def arbitrary_collection_stats(collection_name: str): + return {"row_count": 42} + + +def arbitrary_embedding(text: str): + return [0.1] * 384 + + +@pytest.mark.unit +def test_ready(milvus_backend: MilvusEntityBackend, monkeypatch): + """Test the ready() health check method.""" + + def list_collections(): + return ["collection1", "collection2"] + + monkeypatch.setattr(milvus_backend.milvus, "list_collections", list_collections) + result = milvus_backend.ready() + assert result == {"status": "ok"} + + +@pytest.mark.unit +def test_create_namespace(milvus_backend: MilvusEntityBackend, db_manager, monkeypatch): + """Test creating a new namespace.""" + namespace_id = "test_namespace" + monkeypatch.setattr(milvus_backend.milvus, "has_collection", never_has_collection) + monkeypatch.setattr(milvus_backend.milvus, "create_collection", noop_create_collection) + + with patch("kaizen.backend.milvus.SQLiteManager", return_value=db_manager): + result = milvus_backend.create_namespace(namespace_id=namespace_id) + + assert result.id == namespace_id + assert isinstance(result.created_at, datetime.datetime) + + # create a namespace with auto-generated id + with patch("kaizen.backend.milvus.SQLiteManager", return_value=db_manager): + result = milvus_backend.create_namespace() + + assert result.id.startswith("ns_") + assert isinstance(result.created_at, datetime.datetime) + + +@pytest.mark.unit +def test_get_namespace_details(milvus_backend: MilvusEntityBackend, db_manager, monkeypatch): + """Test retrieving namespace details.""" + monkeypatch.setattr(milvus_backend.milvus, "has_collection", never_has_collection) + + # Test nonexistent namespace + with pytest.raises(NamespaceNotFoundException): + milvus_backend.get_namespace_details(namespace_id="nonexistent_namespace") + + monkeypatch.setattr(milvus_backend.milvus, "has_collection", always_has_collection) + monkeypatch.setattr(milvus_backend.milvus, "get_collection_stats", arbitrary_collection_stats) + db_manager.get_namespace = arbitrary_namespace + + # Test existing namespace + with patch("kaizen.backend.milvus.SQLiteManager", return_value=db_manager): + result = milvus_backend.get_namespace_details(namespace_id="test_namespace") + + assert result.id == "test_namespace" + assert isinstance(result.created_at, datetime.datetime) + assert result.num_entities == 42 + + +@pytest.mark.unit +def test_search_namespaces(milvus_backend: MilvusEntityBackend, db_manager, monkeypatch): + """Test searching for namespaces.""" + created_at = datetime.datetime.now(datetime.UTC) + + db_manager.search_namespaces = Mock( + return_value=[Namespace(id="namespace1", created_at=created_at), Namespace(id="namespace2", created_at=created_at)] + ) + + monkeypatch.setattr(milvus_backend.milvus, "get_collection_stats", arbitrary_collection_stats) + + with patch("kaizen.backend.milvus.SQLiteManager", return_value=db_manager): + result = milvus_backend.search_namespaces(limit=10) + + assert len(result) == 2 + assert result[0].id == "namespace1" + assert result[0].num_entities == 42 + assert result[1].id == "namespace2" + assert result[1].num_entities == 42 + + +@pytest.mark.unit +def test_delete_namespace(milvus_backend: MilvusEntityBackend, db_manager, monkeypatch): + """Test deleting a namespace.""" + namespace_id = "test_namespace" + drop_collection = Mock() + db_manager.delete_namespace = Mock() + monkeypatch.setattr(milvus_backend.milvus, "drop_collection", drop_collection) + + with patch("kaizen.backend.milvus.SQLiteManager", return_value=db_manager): + milvus_backend.delete_namespace(namespace_id=namespace_id) + + drop_collection.assert_called_once_with(collection_name=namespace_id) + db_manager.delete_namespace.assert_called_once_with(namespace_id) + + +@pytest.mark.unit +def test_update_entities(milvus_backend: MilvusEntityBackend, monkeypatch): + """Test updating entities.""" + entity_update = EntityUpdate(id="12345", type="Test entity content", content="fact", event="ADD") + + # No potential conflicts to resolve + def search_entities(self, namespace_id, query, filters=None, limit=10): + return [] + + def insert(collection_name, data): + return {"ids": [12345]} + + def resolve_conflicts(old_entities, new_entities): + return [entity_update] + + monkeypatch.setattr(milvus_backend.milvus, "has_collection", always_has_collection) + monkeypatch.setattr(milvus_backend.milvus, "insert", insert) + monkeypatch.setattr(milvus_backend.embedding_model, "encode", arbitrary_embedding) + monkeypatch.setattr(milvus_backend, "search_entities", search_entities.__get__(milvus_backend, MilvusEntityBackend)) + + with patch("kaizen.backend.milvus.resolve_conflicts", resolve_conflicts): + entities = [Entity(type=entity_update.type, content=entity_update.content, metadata={"key": "value"})] + result = milvus_backend.update_entities(namespace_id="test_namespace", entities=entities, enable_conflict_resolution=True) + + assert len(result) == 1 + assert result[0] == entity_update + + +@pytest.mark.unit +def test_update_entities_mixed_types_raises_exception(milvus_backend: MilvusEntityBackend, monkeypatch): + """Test that updating entities with mixed types raises an exception.""" + monkeypatch.setattr(milvus_backend.milvus, "has_collection", always_has_collection) + + with pytest.raises(KaizenException, match="All entities must have the same type"): + milvus_backend.update_entities( + namespace_id="test_namespace", + entities=[Entity(type="fact", content="Content 1"), Entity(type="guideline", content="Content 2")], + enable_conflict_resolution=False, + ) + + +@pytest.mark.unit +def test_search_entities(milvus_backend: MilvusEntityBackend, monkeypatch): + """Test searching entities with a query string.""" + + def query(collection_name, filter="", output_fields=None, timeout=None, ids=None, partition_names=None, **kwargs): + return [ + { + "id": 123, + "type": "fact", + "content": "Test content", + "created_at": int(datetime.datetime.now(datetime.UTC).timestamp()), + "metadata": {}, + } + ] + + monkeypatch.setattr(milvus_backend.milvus, "has_collection", always_has_collection) + monkeypatch.setattr(milvus_backend.milvus, "query", query) + monkeypatch.setattr(milvus_backend.embedding_model, "encode", arbitrary_embedding) + + # Test searching entities with a query (list all). + result = milvus_backend.search_entities(namespace_id="test_namespace", query="test query", limit=10) + + assert len(result) == 1 + assert result[0].id == "123" + assert result[0].type == "fact" + assert result[0].content == "Test content" + + # Test searching entities without a query (list all). + result: list[RecordedEntity] = milvus_backend.search_entities(namespace_id="test_namespace", query=None) + + assert len(result) == 1 + assert result[0].id == "123" + assert result[0].type == "fact" + assert result[0].content == "Test content" + + # Test searching entities with filters. + result = milvus_backend.search_entities(namespace_id="test_namespace", query="test_query", filters={"type": "fact"}, limit=10) + + assert len(result) == 1 + assert result[0].id == "123" + assert result[0].type == "fact" + assert result[0].content == "Test content" + + +@pytest.mark.unit +def test_delete_entity_by_id(milvus_backend: MilvusEntityBackend, monkeypatch): + """Test deleting an entity by ID.""" + delete = Mock() + + monkeypatch.setattr(milvus_backend.milvus, "has_collection", always_has_collection) + monkeypatch.setattr(milvus_backend.milvus, "delete", delete) + + milvus_backend.delete_entity_by_id(namespace_id="test_namespace", entity_id="12345") + + # Milvus uses integers for its IDs, so the backend converted it. + delete.assert_called_once_with(collection_name="test_namespace", ids=[12345]) + + +@pytest.mark.unit +def test_delete_entity_nonexistent_namespace(milvus_backend: MilvusEntityBackend, monkeypatch): + """Test deleting an entity from a non-existent namespace.""" + monkeypatch.setattr(milvus_backend.milvus, "has_collection", never_has_collection) + + with pytest.raises(NamespaceNotFoundException): + milvus_backend.delete_entity_by_id(namespace_id="nonexistent_namespace", entity_id="12345") From 2a83921805b0f6a13e3d800efa89027a8166c317 Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Wed, 28 Jan 2026 10:35:57 -0800 Subject: [PATCH 2/9] Update milvus.py --- kaizen/backend/milvus.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/kaizen/backend/milvus.py b/kaizen/backend/milvus.py index 4ce91362..2458669c 100644 --- a/kaizen/backend/milvus.py +++ b/kaizen/backend/milvus.py @@ -3,7 +3,7 @@ import logging import uuid -from kaizen.backend.base import BaseEntityBackend +from kaizen.backend.base import BaseEntityBackend, BaseSettings from kaizen.config.milvus import milvus_client_settings, milvus_other_settings from kaizen.db.sqlite_manager import SQLiteManager from kaizen.llm.conflict_resolution.conflict_resolution import resolve_conflicts @@ -33,8 +33,13 @@ def deserialize_content(content: str): class MilvusEntityBackend(BaseEntityBackend): - milvus = MilvusClient(**milvus_client_settings.model_dump()) - embedding_model = SentenceTransformer(milvus_other_settings.embedding_model) + milvus: MilvusClient + embedding_model: SentenceTransformer + + def __init__(self, config: BaseSettings | None = None): + super().__init__(config) + self.milvus = MilvusClient(**milvus_client_settings.model_dump()) + self.embedding_model = SentenceTransformer(milvus_other_settings.embedding_model) def ready(self): _ = self.milvus.list_collections() From 59abd72f42f227682ecb3217a540c43cbf0b486c Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Wed, 28 Jan 2026 10:52:02 -0800 Subject: [PATCH 3/9] Update milvus.py --- kaizen/backend/milvus.py | 1 + 1 file changed, 1 insertion(+) diff --git a/kaizen/backend/milvus.py b/kaizen/backend/milvus.py index 2458669c..f01e60a2 100644 --- a/kaizen/backend/milvus.py +++ b/kaizen/backend/milvus.py @@ -166,6 +166,7 @@ def search_entities( filters = filters or {} if query is None: + # Default query: Get all entities results = self.milvus.query( collection_name=namespace_id, filter=" AND ".join([f"{k} == '{v}'" for k, v in filters.items()]) if len(filters) > 0 else "id > 0", From 81085205796bfff4d247e380f102d64d2b9a9fea Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Thu, 29 Jan 2026 10:29:15 -0800 Subject: [PATCH 4/9] metadata default --- kaizen/schema/conflict_resolution.py | 23 +++++++++++++---------- kaizen/schema/core.py | 23 ++++++++++++++--------- 2 files changed, 27 insertions(+), 19 deletions(-) diff --git a/kaizen/schema/conflict_resolution.py b/kaizen/schema/conflict_resolution.py index 5dc36dee..464dac50 100644 --- a/kaizen/schema/conflict_resolution.py +++ b/kaizen/schema/conflict_resolution.py @@ -2,22 +2,25 @@ from pydantic import BaseModel, Field from typing import Literal + class SimpleEntity(BaseModel): """Derived from either a `Entity` or `RecordedEntity`. Optimized for LLM-based conflict resolution.""" - id: str = Field(description='The unique ID of an entity.') - type: str = Field(description='The type of the entity.') - content: str | list | dict = Field(description='The content of the entity.') + + id: str = Field(description="The unique ID of an entity.") + type: str = Field(description="The type of the entity.") + content: str | list | dict = Field(description="The content of the entity.") @staticmethod - def from_recorded_entities(entities: list[RecordedEntity]) -> list['SimpleEntity']: + def from_recorded_entities(entities: list[RecordedEntity]) -> list["SimpleEntity"]: return [SimpleEntity(id=entity.id, type=entity.type, content=entity.content) for entity in entities] class EntityUpdate(BaseModel): """Produced by the LLM, to be processed by a entity backend.""" - id: str = Field(description='The unique ID of an entity.') - type: str = Field(description='The type of the entity.') - content: str | list | dict = Field(description='The content of the entity.') - event: Literal['ADD', 'UPDATE', 'DELETE', 'NONE'] = Field(description='The type of update operation to perform.') - old_entity: str | None = Field(default=None, description='The entity before it was updated.') - metadata: dict | None = Field(default=None, description='Arbitrary metadata which is related to the entity.') \ No newline at end of file + + id: str = Field(description="The unique ID of an entity.") + type: str = Field(description="The type of the entity.") + content: str | list | dict = Field(description="The content of the entity.") + event: Literal["ADD", "UPDATE", "DELETE", "NONE"] = Field(description="The type of update operation to perform.") + old_entity: str | None = Field(default=None, description="The entity before it was updated.") + metadata: dict = Field(default_factory=dict, description="Arbitrary metadata which is related to the entity.") diff --git a/kaizen/schema/core.py b/kaizen/schema/core.py index 0f42c728..cfc98851 100644 --- a/kaizen/schema/core.py +++ b/kaizen/schema/core.py @@ -5,22 +5,27 @@ class Namespace(BaseModel): """Details of a namespace containing memories.""" - id: str = Field(description='The unique ID of a namespace.') - created_at: datetime = Field(description='The time the namespace was created.') - num_entities: int | None = Field(default= None, description='The number of entities in the namespace. May not be accurate.') + + id: str = Field(description="The unique ID of a namespace.") + created_at: datetime = Field(description="The time the namespace was created.") + num_entities: int | None = Field(default=None, description="The number of entities in the namespace. May not be accurate.") @staticmethod - def row_factory(cursor: Cursor, row: Row) -> 'Namespace': + def row_factory(cursor: Cursor, row: Row) -> "Namespace": fields = [column[0] for column in cursor.description] return Namespace(**{k: v for k, v in zip(fields, row)}) + class Entity(BaseModel): """Basic data stored in the DB""" - content: str | list | dict = Field(description='Searchable text or structured data.') - metadata: dict | None = Field(default=None, description='Arbitrary metadata which is related to the entity.') - type: str = Field(description='The type of the entity.') + + content: str | list | dict = Field(description="Searchable text or structured data.") + metadata: dict = Field(default_factory=dict, description="Arbitrary metadata which is related to the entity.") + type: str = Field(description="The type of the entity.") + class RecordedEntity(Entity): """A statement about a person, place, or thing.""" - id: str = Field(description='The unique ID of an entity.') - created_at: datetime = Field(description='The date and time the entity was created.') + + id: str = Field(description="The unique ID of an entity.") + created_at: datetime = Field(description="The date and time the entity was created.") From 01e36b8856e1d23d849f1c0fe8a22d419981d1c7 Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Thu, 29 Jan 2026 10:52:47 -0800 Subject: [PATCH 5/9] fix filter --- kaizen/backend/base.py | 20 +++++++++----------- kaizen/backend/filesystem.py | 28 +++++++++++----------------- kaizen/backend/milvus.py | 18 ++++++++++++++---- 3 files changed, 34 insertions(+), 32 deletions(-) diff --git a/kaizen/backend/base.py b/kaizen/backend/base.py index 29c0e985..7cc8cce7 100644 --- a/kaizen/backend/base.py +++ b/kaizen/backend/base.py @@ -13,14 +13,15 @@ def __init__(self, config: BaseSettings | None = None): pass @abstractmethod - def ready(self): + def ready(self) -> bool: pass @abstractmethod - def create_namespace( - self, - namespace_id: str | None = None - ) -> Namespace: + def details(self) -> dict: + pass + + @abstractmethod + def create_namespace(self, namespace_id: str | None = None) -> Namespace: pass @abstractmethod @@ -39,15 +40,12 @@ def update_entities( enable_conflict_resolution: bool = True, ) -> list[EntityUpdate]: pass + def search_entities( - self, - namespace_id: str, - query: str | None = None, - filters: dict | None = None, - limit: int = 10 + self, namespace_id: str, query: str | None = None, filters: dict | None = None, limit: int = 10 ) -> list[RecordedEntity]: pass @abstractmethod def delete_entity_by_id(self, namespace_id: str, entity_id: str): - pass \ No newline at end of file + pass diff --git a/kaizen/backend/filesystem.py b/kaizen/backend/filesystem.py index 4b380f99..31157b7b 100644 --- a/kaizen/backend/filesystem.py +++ b/kaizen/backend/filesystem.py @@ -50,9 +50,13 @@ def _save_namespace_data(self, namespace_id: str, data: dict): with open(file_path, "w") as f: json.dump(data, f, indent=2, default=str) - def ready(self): + def ready(self) -> bool: """Check if the backend is healthy.""" - return {"status": "ok", "data_dir": str(self.data_dir)} + return True + + def details(self) -> dict: + """Return details about the backend.""" + return {"data_dir": str(self.data_dir)} def create_namespace(self, namespace_id: str | None = None) -> Namespace: """Create a new namespace for entities to exist in.""" @@ -61,9 +65,7 @@ def create_namespace(self, namespace_id: str | None = None) -> Namespace: with self._lock: if file_path.exists(): - raise NamespaceAlreadyExistsException( - f'Namespace "{namespace_id}" already exists.' - ) + raise NamespaceAlreadyExistsException(f'Namespace "{namespace_id}" already exists.') now = datetime.datetime.now(datetime.UTC) data = { @@ -97,9 +99,7 @@ def search_namespaces(self, limit: int = 10) -> list[Namespace]: namespaces.append( Namespace( id=data["id"], - created_at=datetime.datetime.fromisoformat( - data["created_at"] - ), + created_at=datetime.datetime.fromisoformat(data["created_at"]), num_entities=len(data["entities"]), ) ) @@ -151,9 +151,7 @@ def update_entities( # Find similar existing entities for conflict resolution old_entities = [] for entity in entities: - similar = self._search_entities_internal( - data, query=entity.content, filters=None, limit=10 - ) + similar = self._search_entities_internal(data, query=entity.content, filters=None, limit=10) old_entities.extend(similar) updates = resolve_conflicts(old_entities, entities_with_temporary_ids) @@ -181,9 +179,7 @@ def update_entities( ent["metadata"] = update.metadata break case "DELETE": - data["entities"] = [ - e for e in data["entities"] if e["id"] != update.id - ] + data["entities"] = [e for e in data["entities"] if e["id"] != update.id] case "NONE": pass else: @@ -286,9 +282,7 @@ def delete_entity_by_id(self, namespace_id: str, entity_id: str): with self._lock: data = self._load_namespace_data(namespace_id) original_count = len(data["entities"]) - data["entities"] = [ - e for e in data["entities"] if str(e["id"]) != entity_id - ] + data["entities"] = [e for e in data["entities"] if str(e["id"]) != entity_id] if len(data["entities"]) == original_count: raise KaizenException(f"Entity `{entity_id}` not found") self._save_namespace_data(namespace_id, data) diff --git a/kaizen/backend/milvus.py b/kaizen/backend/milvus.py index f01e60a2..01e6e8a9 100644 --- a/kaizen/backend/milvus.py +++ b/kaizen/backend/milvus.py @@ -32,6 +32,10 @@ def deserialize_content(content: str): return content +def _escape_filter(value: str) -> str: + return value.replace("\\", "\\\\").replace("'", "\\'") + + class MilvusEntityBackend(BaseEntityBackend): milvus: MilvusClient embedding_model: SentenceTransformer @@ -41,9 +45,13 @@ def __init__(self, config: BaseSettings | None = None): self.milvus = MilvusClient(**milvus_client_settings.model_dump()) self.embedding_model = SentenceTransformer(milvus_other_settings.embedding_model) - def ready(self): + def ready(self) -> bool: _ = self.milvus.list_collections() - return {"status": "ok"} + return True + + def details(self) -> dict: + """Return details about the backend.""" + return {} def validate_namespace(self, namespace_id: str): if not self.milvus.has_collection(namespace_id): @@ -169,14 +177,16 @@ def search_entities( # Default query: Get all entities results = self.milvus.query( collection_name=namespace_id, - filter=" AND ".join([f"{k} == '{v}'" for k, v in filters.items()]) if len(filters) > 0 else "id > 0", + filter=" AND ".join([f"{_escape_filter(k)} == '{_escape_filter(v)}'" for k, v in filters.items()]) + if len(filters) > 0 + else "id > 0", ) else: results = self.milvus.query( collection_name=namespace_id, anns_field="embedding", data=[self.embedding_model.encode(query)], - filter=" AND ".join([f"{k} == '{v}'" for k, v in filters.items()]), + filter=" AND ".join([f"{_escape_filter(k)} == '{_escape_filter(v)}'" for k, v in filters.items()]), limit=limit, search_params={"metric_type": "IP"}, ) From 4b04683bcb2c9ecf1a7535b1f1d52e6c0492073a Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Mon, 2 Feb 2026 10:50:16 -0800 Subject: [PATCH 6/9] fix errors --- kaizen/cli/cli.py | 2 + kaizen/frontend/mcp/mcp_server.py | 80 +++------ kaizen/schema/core.py | 2 +- tests/e2e/test_e2e_pipeline.py | 94 +++++------ tests/e2e/test_mcp.py | 272 +++++++++++++++--------------- tests/unit/test_cli.py | 8 +- tests/unit/test_milvus_backend.py | 3 +- 7 files changed, 203 insertions(+), 258 deletions(-) diff --git a/kaizen/cli/cli.py b/kaizen/cli/cli.py index bfa93266..10b3b199 100644 --- a/kaizen/cli/cli.py +++ b/kaizen/cli/cli.py @@ -188,6 +188,8 @@ def add_entity( except json.JSONDecodeError: console.print("[red]Invalid JSON metadata.[/red]") raise typer.Exit(1) + else: + parsed_metadata = {} # Ensure namespace exists if not client.namespace_exists(namespace): diff --git a/kaizen/frontend/mcp/mcp_server.py b/kaizen/frontend/mcp/mcp_server.py index c2564f04..62fcfdd4 100644 --- a/kaizen/frontend/mcp/mcp_server.py +++ b/kaizen/frontend/mcp/mcp_server.py @@ -24,7 +24,7 @@ def get_client() -> KaizenClient: """Get or create the KaizenClient singleton. - + This lazy initialization allows tests to configure settings before the client is created. """ @@ -69,9 +69,7 @@ def get_guidelines(task: str) -> str: @mcp.tool() -def save_trajectory( - trajectory_data: str, task_id: str | None = None -) -> list[RecordedEntity]: +def save_trajectory(trajectory_data: str, task_id: str | None = None) -> list[RecordedEntity]: """ Save the full agent trajectory to the Entity DB and generate tips @@ -87,9 +85,7 @@ def save_trajectory( entities.append( Entity( type="trajectory", - content=message["content"] - if isinstance(message["content"], str) - else str(message["content"]), + content=message["content"] if isinstance(message["content"], str) else str(message["content"]), metadata={ "task_id": task_id, "message": message, # store the original message for reference @@ -129,27 +125,22 @@ def save_trajectory( @mcp.tool() -def create_entity( - content: str, - entity_type: str, - metadata: str | None = None, - enable_conflict_resolution: bool = False -) -> str: +def create_entity(content: str, entity_type: str, metadata: str | None = None, enable_conflict_resolution: bool = False) -> str: """ Create a single entity in the namespace. - + Args: content: The searchable text or structured data for the entity entity_type: The type/category of the entity (e.g., 'guideline', 'note', 'fact') metadata: Optional JSON string containing arbitrary metadata related to the entity enable_conflict_resolution: If True, uses LLM to check for conflicts with existing entities - + Returns: JSON string with the entity update details (ADD/UPDATE/DELETE/NONE) and entity ID """ logger.info(f"Creating entity of type: {entity_type}") ensure_namespace() - + # Parse metadata if provided metadata_dict = None if metadata: @@ -157,36 +148,26 @@ def create_entity( metadata_dict = json.loads(metadata) except json.JSONDecodeError as e: logger.exception(f"Invalid JSON in metadata parameter: {str(e)}") - return json.dumps({ - "error": "Invalid metadata JSON", - "message": f"Failed to parse metadata: {str(e)}", - "invalid_metadata": metadata - }) - + return json.dumps( + {"error": "Invalid metadata JSON", "message": f"Failed to parse metadata: {str(e)}", "invalid_metadata": metadata} + ) + else: + metadata_dict = {} + # Create the entity using the Entity schema - entity = Entity( - type=entity_type, - content=content, - metadata=metadata_dict - ) - + entity = Entity(type=entity_type, content=content, metadata=metadata_dict) + # Use KaizenClient.update_entities() to create the entity updates = get_client().update_entities( - namespace_id=kaizen_config.namespace_id, - entities=[entity], - enable_conflict_resolution=enable_conflict_resolution + namespace_id=kaizen_config.namespace_id, entities=[entity], enable_conflict_resolution=enable_conflict_resolution ) - + # Return the first (and only) update result if updates: update = updates[0] - return json.dumps({ - "event": update.event, - "id": update.id, - "type": update.type, - "content": update.content, - "metadata": update.metadata - }) + return json.dumps( + {"event": update.event, "id": update.id, "type": update.type, "content": update.content, "metadata": update.metadata} + ) else: return json.dumps({"error": "Entity creation failed"}) @@ -195,29 +176,20 @@ def create_entity( def delete_entity(entity_id: str) -> str: """ Delete a specific entity by its ID. - + Args: entity_id: The unique identifier of the entity to delete - + Returns: JSON string confirming deletion or error message """ logger.info(f"Deleting entity: {entity_id}") ensure_namespace() - + try: # Use KaizenClient.delete_entity_by_id() to delete the entity - get_client().delete_entity_by_id( - namespace_id=kaizen_config.namespace_id, - entity_id=entity_id - ) - return json.dumps({ - "success": True, - "message": f"Entity {entity_id} deleted successfully" - }) + get_client().delete_entity_by_id(namespace_id=kaizen_config.namespace_id, entity_id=entity_id) + return json.dumps({"success": True, "message": f"Entity {entity_id} deleted successfully"}) except KaizenException as e: logger.exception(f"Error deleting entity {entity_id}: {str(e)}") - return json.dumps({ - "success": False, - "error": str(e) - }) + return json.dumps({"success": False, "error": str(e)}) diff --git a/kaizen/schema/core.py b/kaizen/schema/core.py index cfc98851..ffd10944 100644 --- a/kaizen/schema/core.py +++ b/kaizen/schema/core.py @@ -20,8 +20,8 @@ class Entity(BaseModel): """Basic data stored in the DB""" content: str | list | dict = Field(description="Searchable text or structured data.") - metadata: dict = Field(default_factory=dict, description="Arbitrary metadata which is related to the entity.") type: str = Field(description="The type of the entity.") + metadata: dict = Field(default_factory=dict, description="Arbitrary metadata which is related to the entity.") class RecordedEntity(Entity): diff --git a/tests/e2e/test_e2e_pipeline.py b/tests/e2e/test_e2e_pipeline.py index 89a1bee0..ea42daa5 100644 --- a/tests/e2e/test_e2e_pipeline.py +++ b/tests/e2e/test_e2e_pipeline.py @@ -8,34 +8,20 @@ # Configuration PHOENIX_URL = phoenix_settings.url -# Use a session-scope timestamp or generate per test? +# Use a session-scope timestamp or generate per test? # Per-test ensures no collisions even if run in parallel (though these should satisfy sequential) TIMESTAMP = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") AGENTS_TO_TEST = [ - { - "name": "smolagents", - "script": "examples/low_code/smolagents_demo.py", - "project_prefix": "verify-smolagents" - }, - { - "name": "openai_agents", - "script": "examples/low_code/openai_agents_demo.py", - "project_prefix": "verify-openai" - }, - { - "name": "manual_phoenix", - "script": "examples/low_code/manual_phoenix_demo.py", - "project_prefix": "verify-manual" - }, - { - "name": "simple_openai", - "script": "examples/low_code/simple_openai.py", - "project_prefix": "verify-simple-openai" - } + {"name": "smolagents", "script": "examples/low_code/smolagents_demo.py", "project_prefix": "verify-smolagents"}, + {"name": "openai_agents", "script": "examples/low_code/openai_agents_demo.py", "project_prefix": "verify-openai"}, + {"name": "manual_phoenix", "script": "examples/low_code/manual_phoenix_demo.py", "project_prefix": "verify-manual"}, + {"name": "simple_openai", "script": "examples/low_code/simple_openai.py", "project_prefix": "verify-simple-openai"}, ] -@pytest.mark.skipif(os.getenv("KAIZEN_E2E") != "true", reason="E2E tests disabled unless KAIZEN_E2E=true") + +@pytest.mark.e2e +@pytest.mark.phoenix @pytest.mark.parametrize("agent_config", AGENTS_TO_TEST, ids=[a["name"] for a in AGENTS_TO_TEST]) def test_e2e_pipeline_agent(agent_config): """ @@ -46,12 +32,12 @@ def test_e2e_pipeline_agent(agent_config): """ agent_name = agent_config["name"] script_path = agent_config["script"] - + # Generate unique project name for this run # Using a fresh timestamp per run to avoid collisions if tests run slowly current_timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") project_name = f"{agent_config['project_prefix']}-{current_timestamp}" - + print("\n==================================================") print(f" TESTING AGENT: {agent_name}") print(f" Script: {script_path}") @@ -67,32 +53,27 @@ def test_e2e_pipeline_agent(agent_config): # kaizen.auto prioritizes KAIZEN_TRACING_PROJECT over PHOENIX_PROJECT_NAME env["KAIZEN_TRACING_PROJECT"] = project_name env["PHOENIX_PROJECT_NAME"] = project_name - + # Ensure script exists if not os.path.exists(script_path): pytest.fail(f"Script not found: {script_path}") - result = subprocess.run( - ["uv", "run", "python", script_path], - env=env, - capture_output=True, - text=True - ) - + result = subprocess.run(["uv", "run", "python", script_path], env=env, capture_output=True, text=True) + if result.returncode != 0: print(f"❌ Agent failed with exit code {result.returncode}") print("STDERR:", result.stderr) print("STDOUT:", result.stdout) pytest.fail(f"Agent execution failed: {result.stderr}") - + print(f"✅ Agent finished in {time.time() - start_time:.2f}s") # --- Step 2: Verify Traces --- print(f"\n--- Step 2: Verifying Phoenix Traces ({project_name}) ---") - + # Wait briefly for traces to be flushed/indexed - time.sleep(2) - + time.sleep(2) + check_script = f""" import phoenix as px import sys @@ -106,12 +87,8 @@ def test_e2e_pipeline_agent(agent_config): except Exception as e: print(f"ERROR:{{e}}") """ - result = subprocess.run( - ["uv", "run", "python", "-c", check_script], - capture_output=True, - text=True - ) - + result = subprocess.run(["uv", "run", "python", "-c", check_script], capture_output=True, text=True) + output = result.stdout + result.stderr if "FOUND_TRACES" in output: count = output.split("FOUND_TRACES:")[1].split()[0] @@ -124,41 +101,48 @@ def test_e2e_pipeline_agent(agent_config): # --- Step 3: Sync & Generate Tips --- print("\n--- Step 3: Running Kaizen Sync & Monitoring ---") sync_command = [ - "uv", "run", "python", "-m", "kaizen.frontend.cli.cli", - "sync", "phoenix", - "--project", project_name, + "uv", + "run", + "python", + "-m", + "kaizen.frontend.cli.cli", + "sync", + "phoenix", + "--project", + project_name, "--include-errors", - "--limit", "500" + "--limit", + "500", ] print(f"Command: {' '.join(sync_command)}") - + process = subprocess.Popen( sync_command, stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, # Merge stderr to monitor everything + stderr=subprocess.STDOUT, # Merge stderr to monitor everything text=True, bufsize=1, - universal_newlines=True + universal_newlines=True, ) - + tips_found = False sync_start = time.time() - timeout = 120 # 2 minute timeout for sync - + timeout = 120 # 2 minute timeout for sync + try: while True: if time.time() - sync_start > timeout: print("❌ Timeout waiting for tips generation") break - + line = process.stdout.readline() if not line and process.poll() is not None: break - + if line: line_stripped = line.strip() # print(f"[Sync] {line_stripped}") # Optional: verbose logging - + # Check target log pattern match = re.search(r"generated (\d+) tips", line_stripped) if match: diff --git a/tests/e2e/test_mcp.py b/tests/e2e/test_mcp.py index f54324de..06fb2077 100644 --- a/tests/e2e/test_mcp.py +++ b/tests/e2e/test_mcp.py @@ -8,45 +8,49 @@ import uuid from kaizen.config.milvus import milvus_client_settings -__data__ = Path(__file__).parent.parent / 'data' +__data__ = Path(__file__).parent.parent / "data" load_dotenv() + @pytest.fixture def mcp(): - os.environ['KAIZEN_NAMESPACE_ID'] = 'test' + os.environ["KAIZEN_NAMESPACE_ID"] = "test" from kaizen.frontend.client.kaizen_client import KaizenClient from kaizen.config.kaizen import kaizen_config + # we change the namespace ID for these tests so we have to reset the loaded settings kaizen_config.__init__() - + # Use a unique DB file for each test to avoid socket/locking issues # Milvus Lite has a 36 character limit on DB filenames db_file = f"test_{uuid.uuid4().hex[:8]}.db" original_uri = milvus_client_settings.uri milvus_client_settings.uri = db_file - + # Reset the MCP server client to ensure it uses the new DB file import kaizen.frontend.mcp.mcp_server as mcp_server_module + mcp_server_module._client = None - + kaizen_client = KaizenClient() # Create the test namespace try: - kaizen_client.create_namespace('test') + kaizen_client.create_namespace("test") except Exception: pass - + yield mcp_server_module.mcp - + # Cleanup - close the backend connection properly try: kaizen_client.backend.close() except Exception: pass - + # Disconnect all pymilvus connections to ensure clean state for next test try: from pymilvus import connections + for alias, _ in connections.list_connections(): try: connections.disconnect(alias) @@ -54,20 +58,21 @@ def mcp(): pass except Exception: pass - + # Release all Milvus Lite servers to fully clean up between tests try: from milvus_lite.server_manager import server_manager_instance + server_manager_instance.release_all() except Exception: pass - + # Reset the MCP server client mcp_server_module._client = None - + # Restore original URI milvus_client_settings.uri = original_uri - + # Clean up temp DB files if os.path.exists(db_file): try: @@ -81,49 +86,44 @@ def mcp(): pass - - - @pytest.mark.e2e async def test_save_trajectory_and_retrieve_guidelines(mcp): async with Client(transport=mcp) as kaizen_mcp: - trajectory = (__data__ / 'trajectory.json').read_text() - response = await kaizen_mcp.call_tool_mcp('save_trajectory', { - 'trajectory_data': trajectory, - 'task_id': '123' - }) + trajectory = (__data__ / "trajectory.json").read_text() + response = await kaizen_mcp.call_tool_mcp("save_trajectory", {"trajectory_data": trajectory, "task_id": "123"}) saved_trajectory = json.loads(response.content[0].text) # MCP server should return entity versions of the trajectory assert len(saved_trajectory) == 19 - response = await kaizen_mcp.call_tool_mcp('get_guidelines', { - 'task': 'What states do I have teammates in? Read the list from the states.txt file. use the filesystem mcp tool' - }) + response = await kaizen_mcp.call_tool_mcp( + "get_guidelines", + {"task": "What states do I have teammates in? Read the list from the states.txt file. use the filesystem mcp tool"}, + ) guidelines = response.content[0].text - assert '# Guidelines for: ' in guidelines + assert "# Guidelines for: " in guidelines @pytest.mark.e2e async def test_create_entity_without_conflict_resolution(mcp): """Test creating a single entity without conflict resolution.""" async with Client(transport=mcp) as kaizen_mcp: - response = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Always use type hints in Python functions', - 'entity_type': 'guideline', - 'metadata': json.dumps({ - 'category': 'code_quality', - 'language': 'python' - }), - 'enable_conflict_resolution': False - }) - + response = await kaizen_mcp.call_tool_mcp( + "create_entity", + { + "content": "Always use type hints in Python functions", + "entity_type": "guideline", + "metadata": json.dumps({"category": "code_quality", "language": "python"}), + "enable_conflict_resolution": False, + }, + ) + result = json.loads(response.content[0].text) - + # Verify entity was created (ADD event) - assert result['event'] == 'ADD' - assert 'id' in result - assert result['type'] == 'guideline' - assert result['content'] == 'Always use type hints in Python functions' - assert result['metadata']['category'] == 'code_quality' + assert result["event"] == "ADD" + assert "id" in result + assert result["type"] == "guideline" + assert result["content"] == "Always use type hints in Python functions" + assert result["metadata"]["category"] == "code_quality" @pytest.mark.e2e @@ -131,61 +131,56 @@ async def test_create_entity_with_conflict_resolution(mcp): """Test creating an entity with conflict resolution enabled.""" from unittest.mock import patch from kaizen.schema.conflict_resolution import EntityUpdate - + async with Client(transport=mcp) as kaizen_mcp: # Create first entity - response1 = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Use descriptive variable names', - 'entity_type': 'guideline', - 'enable_conflict_resolution': False - }) - result1 = json.loads(response1.content[0].text) - assert result1['event'] == 'ADD' - first_entity_id = result1['id'] - + response = await kaizen_mcp.call_tool_mcp( + "create_entity", {"content": "Use descriptive variable names", "entity_type": "guideline", "enable_conflict_resolution": False} + ) + + assert not response.isError + result = json.loads(response.content[0].text) + assert result["event"] == "ADD" + first_entity_id = result["id"] + # Mock resolve_conflicts to avoid LLM call timeout - with patch('kaizen.backend.milvus.resolve_conflicts') as mock_resolve: + with patch("kaizen.backend.milvus.resolve_conflicts") as mock_resolve: # Configure mock to return an UPDATE event mock_resolve.return_value = [ EntityUpdate( - id=str(first_entity_id), - type='guideline', - content='Always use descriptive variable names', - event='UPDATE', - metadata={} + id=str(first_entity_id), type="guideline", content="Always use descriptive variable names", event="UPDATE", metadata={} ) ] - + # Create similar entity with conflict resolution - response2 = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Always use descriptive variable names', - 'entity_type': 'guideline', - 'enable_conflict_resolution': True - }) - result2 = json.loads(response2.content[0].text) - + response = await kaizen_mcp.call_tool_mcp( + "create_entity", + {"content": "Always use descriptive variable names", "entity_type": "guideline", "enable_conflict_resolution": True}, + ) + + assert not response.isError + result = json.loads(response.content[0].text) + # Should return what our mock returned - assert result2['event'] == 'UPDATE' - assert result2['id'] == str(first_entity_id) + assert result["event"] == "UPDATE" + assert result["id"] == str(first_entity_id) @pytest.mark.e2e async def test_create_entity_without_metadata(mcp): """Test creating an entity without optional metadata.""" async with Client(transport=mcp) as kaizen_mcp: - response = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Simple entity without metadata', - 'entity_type': 'note', - 'enable_conflict_resolution': False - }) - + response = await kaizen_mcp.call_tool_mcp( + "create_entity", {"content": "Simple entity without metadata", "entity_type": "note", "enable_conflict_resolution": False} + ) + result = json.loads(response.content[0].text) - + # Verify entity was created - assert result['event'] == 'ADD' - assert 'id' in result - assert result['type'] == 'note' - assert result['content'] == 'Simple entity without metadata' + assert result["event"] == "ADD" + assert "id" in result + assert result["type"] == "note" + assert result["content"] == "Simple entity without metadata" @pytest.mark.e2e @@ -193,26 +188,27 @@ async def test_delete_entity(mcp): """Test deleting an entity via MCP.""" async with Client(transport=mcp) as kaizen_mcp: # Create an entity - create_response = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Temporary test entity', - 'entity_type': 'test', - 'metadata': json.dumps({'temp': True}), - 'enable_conflict_resolution': False - }) - + create_response = await kaizen_mcp.call_tool_mcp( + "create_entity", + { + "content": "Temporary test entity", + "entity_type": "test", + "metadata": json.dumps({"temp": True}), + "enable_conflict_resolution": False, + }, + ) + created_entity = json.loads(create_response.content[0].text) - entity_id = created_entity['id'] - + entity_id = created_entity["id"] + # Delete the entity - delete_response = await kaizen_mcp.call_tool_mcp('delete_entity', { - 'entity_id': entity_id - }) - + delete_response = await kaizen_mcp.call_tool_mcp("delete_entity", {"entity_id": entity_id}) + result = json.loads(delete_response.content[0].text) - + # Verify deletion was successful - assert result['success'] is True - assert entity_id in result['message'] + assert result["success"] is True + assert entity_id in result["message"] @pytest.mark.e2e @@ -220,15 +216,13 @@ async def test_delete_nonexistent_entity(mcp): """Test deleting an entity that doesn't exist.""" async with Client(transport=mcp) as kaizen_mcp: # Use a numeric ID that doesn't exist (entity IDs are integers) - delete_response = await kaizen_mcp.call_tool_mcp('delete_entity', { - 'entity_id': '99999999' - }) - + delete_response = await kaizen_mcp.call_tool_mcp("delete_entity", {"entity_id": "99999999"}) + result = json.loads(delete_response.content[0].text) - + # Should return an error since entity doesn't exist - assert result['success'] is False - assert 'error' in result + assert result["success"] is False + assert "error" in result @pytest.mark.e2e @@ -236,22 +230,23 @@ async def test_create_and_delete_workflow(mcp): """Test complete workflow: create then delete.""" async with Client(transport=mcp) as kaizen_mcp: # Create entity - create_response = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Workflow test entity', - 'entity_type': 'test', - 'metadata': json.dumps({'workflow': 'test'}), - 'enable_conflict_resolution': False - }) + create_response = await kaizen_mcp.call_tool_mcp( + "create_entity", + { + "content": "Workflow test entity", + "entity_type": "test", + "metadata": json.dumps({"workflow": "test"}), + "enable_conflict_resolution": False, + }, + ) created = json.loads(create_response.content[0].text) - entity_id = created['id'] - assert created['event'] == 'ADD' - + entity_id = created["id"] + assert created["event"] == "ADD" + # Delete entity - delete_response = await kaizen_mcp.call_tool_mcp('delete_entity', { - 'entity_id': entity_id - }) + delete_response = await kaizen_mcp.call_tool_mcp("delete_entity", {"entity_id": entity_id}) delete_result = json.loads(delete_response.content[0].text) - assert delete_result['success'] is True + assert delete_result["success"] is True @pytest.mark.e2e @@ -259,43 +254,42 @@ async def test_create_multiple_entities_same_type(mcp): """Test creating multiple entities of the same type.""" async with Client(transport=mcp) as kaizen_mcp: entity_ids = [] - + # Create 3 entities for i in range(3): - response = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': f'Test guideline number {i}', - 'entity_type': 'guideline', - 'enable_conflict_resolution': False - }) + response = await kaizen_mcp.call_tool_mcp( + "create_entity", {"content": f"Test guideline number {i}", "entity_type": "guideline", "enable_conflict_resolution": False} + ) result = json.loads(response.content[0].text) - assert result['event'] == 'ADD' - entity_ids.append(result['id']) - + assert result["event"] == "ADD" + entity_ids.append(result["id"]) + # Verify all have unique IDs assert len(set(entity_ids)) == 3 - + # Clean up for entity_id in entity_ids: - await kaizen_mcp.call_tool_mcp('delete_entity', { - 'entity_id': entity_id - }) + await kaizen_mcp.call_tool_mcp("delete_entity", {"entity_id": entity_id}) @pytest.mark.e2e async def test_create_entity_with_invalid_json_metadata(mcp): """Test creating an entity with invalid JSON metadata.""" async with Client(transport=mcp) as kaizen_mcp: - response = await kaizen_mcp.call_tool_mcp('create_entity', { - 'content': 'Test entity with bad metadata', - 'entity_type': 'test', - 'metadata': '{invalid json here}', - 'enable_conflict_resolution': False - }) - + response = await kaizen_mcp.call_tool_mcp( + "create_entity", + { + "content": "Test entity with bad metadata", + "entity_type": "test", + "metadata": "{invalid json here}", + "enable_conflict_resolution": False, + }, + ) + result = json.loads(response.content[0].text) - + # Should return an error - assert 'error' in result - assert result['error'] == 'Invalid metadata JSON' - assert 'message' in result - assert 'invalid_metadata' in result \ No newline at end of file + assert "error" in result + assert result["error"] == "Invalid metadata JSON" + assert "message" in result + assert "invalid_metadata" in result diff --git a/tests/unit/test_cli.py b/tests/unit/test_cli.py index addfe214..2ea16162 100644 --- a/tests/unit/test_cli.py +++ b/tests/unit/test_cli.py @@ -495,13 +495,7 @@ def test_show_entity_without_metadata(self, mock_client): """Test showing entity without metadata.""" created_at = datetime.datetime(2024, 1, 15, 10, 30, 0, tzinfo=datetime.UTC) mock_client.get_all_entities.return_value = [ - RecordedEntity( - id="123", - type="guideline", - content="Content without metadata", - created_at=created_at, - metadata=None, - ), + RecordedEntity(id="123", type="guideline", content="Content without metadata", created_at=created_at), ] result = runner.invoke(app, ["entities", "show", "my_namespace", "123"]) diff --git a/tests/unit/test_milvus_backend.py b/tests/unit/test_milvus_backend.py index c0caa41c..976e17ef 100644 --- a/tests/unit/test_milvus_backend.py +++ b/tests/unit/test_milvus_backend.py @@ -67,8 +67,7 @@ def list_collections(): return ["collection1", "collection2"] monkeypatch.setattr(milvus_backend.milvus, "list_collections", list_collections) - result = milvus_backend.ready() - assert result == {"status": "ok"} + assert milvus_backend.ready() @pytest.mark.unit From ebb01cc79f551f27b9c6033209d6bc0c9d0d33a8 Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Mon, 2 Feb 2026 12:48:01 -0800 Subject: [PATCH 7/9] Update milvus.py --- kaizen/backend/milvus.py | 13 ++----------- 1 file changed, 2 insertions(+), 11 deletions(-) diff --git a/kaizen/backend/milvus.py b/kaizen/backend/milvus.py index 01e6e8a9..711e21d6 100644 --- a/kaizen/backend/milvus.py +++ b/kaizen/backend/milvus.py @@ -198,22 +198,13 @@ def delete_entity_by_id(self, namespace_id: str, entity_id: str): except ValueError: raise KaizenException(f"Invalid entity ID: {entity_id}. Entity IDs must be numeric.") self.validate_namespace(namespace_id) - - # Check if entity exists before deleting - existing = self.milvus.query( - collection_name=namespace_id, - filter=f"id == {entity_id_int}", - output_fields=["id"] - ) - if not existing: - raise KaizenException(f"Entity with ID {entity_id} not found in namespace {namespace_id}.") - + # Entity deletion is idempotent and does not require validation. self.milvus.delete(collection_name=namespace_id, ids=[entity_id_int]) def close(self): """Close Milvus connection.""" try: - if hasattr(self, 'milvus'): + if hasattr(self, "milvus"): self.milvus.close() except Exception as e: logger.warning(f"Error closing Milvus client: {e}") From 87367a660dc6b0d1b40ddde01870e838810eea6a Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Mon, 2 Feb 2026 12:56:26 -0800 Subject: [PATCH 8/9] Update milvus.py --- kaizen/backend/milvus.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/kaizen/backend/milvus.py b/kaizen/backend/milvus.py index 711e21d6..a902f9f8 100644 --- a/kaizen/backend/milvus.py +++ b/kaizen/backend/milvus.py @@ -92,6 +92,10 @@ def delete_namespace(self, namespace_id: str): def update_entities(self, namespace_id: str, entities: list[Entity], enable_conflict_resolution: bool = True) -> list[EntityUpdate]: self.validate_namespace(namespace_id) + if len(entities) == 0: + logger.warning("No entities to update.") + return [] + entity_type = entities[0].type if not all(entity.type == entity_type for entity in entities): raise KaizenException("All entities must have the same type.") From c04ea153d30fdbe9a3708efb50e94e013fc353cb Mon Sep 17 00:00:00 2001 From: Punleuk Oum Date: Mon, 2 Feb 2026 12:58:49 -0800 Subject: [PATCH 9/9] Update test_mcp.py --- tests/e2e/test_mcp.py | 14 -------------- 1 file changed, 14 deletions(-) diff --git a/tests/e2e/test_mcp.py b/tests/e2e/test_mcp.py index 06fb2077..2a2f4925 100644 --- a/tests/e2e/test_mcp.py +++ b/tests/e2e/test_mcp.py @@ -211,20 +211,6 @@ async def test_delete_entity(mcp): assert entity_id in result["message"] -@pytest.mark.e2e -async def test_delete_nonexistent_entity(mcp): - """Test deleting an entity that doesn't exist.""" - async with Client(transport=mcp) as kaizen_mcp: - # Use a numeric ID that doesn't exist (entity IDs are integers) - delete_response = await kaizen_mcp.call_tool_mcp("delete_entity", {"entity_id": "99999999"}) - - result = json.loads(delete_response.content[0].text) - - # Should return an error since entity doesn't exist - assert result["success"] is False - assert "error" in result - - @pytest.mark.e2e async def test_create_and_delete_workflow(mcp): """Test complete workflow: create then delete."""