diff --git a/src/basic_memory/api/routers/knowledge_router.py b/src/basic_memory/api/routers/knowledge_router.py index 87ed322f..08f260a5 100644 --- a/src/basic_memory/api/routers/knowledge_router.py +++ b/src/basic_memory/api/routers/knowledge_router.py @@ -5,8 +5,7 @@ from loguru import logger from basic_memory.deps import ( EntityServiceDep, - KnowledgeServiceDep, - get_search_service, + get_search_service, RelationServiceDep, ObservationServiceDep, FileServiceDep, ) from basic_memory.schemas import ( CreateEntityRequest, @@ -33,11 +32,11 @@ router = APIRouter(prefix="/knowledge", tags=["knowledge"]) async def create_entities( data: CreateEntityRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + entity_service: EntityServiceDep, search_service=Depends(get_search_service), ) -> EntityListResponse: """Create new entities in the knowledge graph and index them.""" - entities = await knowledge_service.create_entities(data.entities) + entities = await entity_service.create_entities(data.entities) # Index each entity for entity in entities: @@ -53,7 +52,7 @@ async def update_entity( path_id: PathId, data: UpdateEntityRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + entity_service: EntityServiceDep, search_service=Depends(get_search_service), ) -> EntityResponse: """Update an existing entity and reindex it.""" @@ -62,7 +61,7 @@ async def update_entity( update_data = data.model_dump(exclude_none=True) # Update the entity - updated_entity = await knowledge_service.update_entity(path_id, **update_data) + updated_entity = await entity_service.update_entity(path_id, **update_data) # Reindex since content changed await search_service.index_entity(updated_entity, background_tasks=background_tasks) @@ -77,11 +76,11 @@ async def update_entity( async def create_relations( data: CreateRelationsRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + relation_service: RelationServiceDep, search_service=Depends(get_search_service), ) -> EntityListResponse: """Create relations between entities and update search index.""" - updated_entities = await knowledge_service.create_relations(data.relations) + updated_entities = await relation_service.create_relations(data.relations) # Reindex updated entities since relations have changed for entity in updated_entities: @@ -96,12 +95,12 @@ async def create_relations( async def add_observations( data: AddObservationsRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + observation_service: ObservationServiceDep, search_service=Depends(get_search_service), ) -> EntityResponse: """Add observations to an entity and update search index.""" logger.debug(f"Adding observations to entity: {data.path_id}") - updated_entity = await knowledge_service.add_observations( + updated_entity = await observation_service.add_observations( data.path_id, data.observations, data.context ) @@ -116,7 +115,8 @@ async def add_observations( @router.get("/entities/{path_id:path}", response_model=EntityResponse) async def get_entity( - knowledge_service: KnowledgeServiceDep, + entity_service: EntityServiceDep, + file_service: FileServiceDep, path_id: PathId, content: bool = False, # New parameter ) -> EntityResponse: @@ -125,13 +125,14 @@ async def get_entity( Args: path_id: Entity path ID content: If True, include full file content + :param entity_service: EntityService """ try: - entity = await knowledge_service.get_entity_by_path_id(path_id) + entity = await entity_service.get_by_path_id(path_id) entity_response = EntityResponse.model_validate(entity) if content: # Load content if requested - content = await knowledge_service.read_entity_content(entity) + content = await file_service.read_entity_content(entity) entity_response.content = content return entity_response @@ -156,11 +157,11 @@ async def open_nodes( async def delete_entities( data: DeleteEntitiesRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + entity_service: EntityServiceDep, search_service=Depends(get_search_service), ) -> DeleteEntitiesResponse: """Delete entities and remove from search index.""" - deleted = await knowledge_service.delete_entities(data.path_ids) + deleted = await entity_service.delete_entities(data.path_ids) # Remove each deleted entity from search index for path_id in data.path_ids: @@ -173,12 +174,12 @@ async def delete_entities( async def delete_observations( data: DeleteObservationsRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + observation_service: ObservationServiceDep, search_service=Depends(get_search_service), ) -> EntityResponse: """Delete observations and update search index.""" path_id = data.path_id - updated_entity = await knowledge_service.delete_observations(path_id, data.observations) + updated_entity = await observation_service.delete_observations(path_id, data.observations) # Reindex the entity since observations changed await search_service.index_entity(updated_entity, background_tasks=background_tasks) @@ -190,11 +191,11 @@ async def delete_observations( async def delete_relations( data: DeleteRelationsRequest, background_tasks: BackgroundTasks, - knowledge_service: KnowledgeServiceDep, + relation_service: RelationServiceDep, search_service=Depends(get_search_service), ) -> EntityListResponse: """Delete relations and update search index.""" - updated_entities = await knowledge_service.delete_relations(data.relations) + updated_entities = await relation_service.delete_relations(data.relations) # Reindex entities since relations changed for entity in updated_entities: diff --git a/src/basic_memory/deps.py b/src/basic_memory/deps.py index 18147371..b96842e8 100644 --- a/src/basic_memory/deps.py +++ b/src/basic_memory/deps.py @@ -23,7 +23,6 @@ from basic_memory.services import ( ) from basic_memory.services.activity_service import ActivityService from basic_memory.services.file_service import FileService -from basic_memory.services.knowledge import KnowledgeService from basic_memory.services.search_service import SearchService @@ -107,16 +106,18 @@ SearchRepositoryDep = Annotated[SearchRepository, Depends(get_search_repository) ## services -async def get_file_service() -> FileService: - return FileService() +async def get_file_service(project_config: ProjectConfigDep) -> FileService: + return FileService(project_config.home, KnowledgeWriter()) FileServiceDep = Annotated[FileService, Depends(get_file_service)] -async def get_entity_service(entity_repository: EntityRepositoryDep) -> EntityService: +async def get_entity_service( + entity_repository: EntityRepositoryDep, file_service: FileServiceDep +) -> EntityService: """Create EntityService with repository.""" - return EntityService(entity_repository) + return EntityService(entity_repository=entity_repository, file_service=file_service) EntityServiceDep = Annotated[EntityService, Depends(get_entity_service)] @@ -124,17 +125,25 @@ EntityServiceDep = Annotated[EntityService, Depends(get_entity_service)] async def get_observation_service( observation_repository: ObservationRepositoryDep, + entity_repository: EntityRepositoryDep, + file_service: FileServiceDep, ) -> ObservationService: """Create ObservationService with repository.""" - return ObservationService(observation_repository) + return ObservationService( + observation_repository=observation_repository, + entity_repository=entity_repository, + file_service=file_service, + ) ObservationServiceDep = Annotated[ObservationService, Depends(get_observation_service)] -async def get_relation_service(relation_repository: RelationRepositoryDep) -> RelationService: +async def get_relation_service( + relation_repository: RelationRepositoryDep, entity_repository: EntityRepositoryDep, file_service: FileServiceDep +) -> RelationService: """Create RelationService with repository.""" - return RelationService(relation_repository) + return RelationService(relation_repository=relation_repository, entity_repository=entity_repository, file_service=file_service) RelationServiceDep = Annotated[RelationService, Depends(get_relation_service)] @@ -169,27 +178,3 @@ async def get_knowledge_writer() -> KnowledgeWriter: KnowledgeWriterDep = Annotated[KnowledgeWriter, Depends(get_knowledge_writer)] - - - - -async def get_knowledge_service( - entity_service: EntityServiceDep, - observation_service: ObservationServiceDep, - relation_service: RelationServiceDep, - file_service: FileServiceDep, - knowledge_writer: KnowledgeWriterDep, - project_config: ProjectConfigDep, -) -> KnowledgeService: - """Create KnowledgeService with dependencies.""" - return KnowledgeService( - entity_service=entity_service, - observation_service=observation_service, - relation_service=relation_service, - file_service=file_service, - knowledge_writer=knowledge_writer, - base_path=project_config.knowledge_dir, - ) - - -KnowledgeServiceDep = Annotated[KnowledgeService, Depends(get_knowledge_service)] diff --git a/src/basic_memory/schemas/base.py b/src/basic_memory/schemas/base.py index 0069ef67..223512c9 100644 --- a/src/basic_memory/schemas/base.py +++ b/src/basic_memory/schemas/base.py @@ -123,8 +123,8 @@ class Relation(BaseModel): or recipient entity. """ - from_path_id: PathId - to_path_id: PathId + from_id: PathId + to_id: PathId relation_type: RelationType context: Optional[str] = None diff --git a/src/basic_memory/schemas/response.py b/src/basic_memory/schemas/response.py index ad8165c6..aadfb344 100644 --- a/src/basic_memory/schemas/response.py +++ b/src/basic_memory/schemas/response.py @@ -61,7 +61,7 @@ class RelationResponse(Relation, SQLAlchemyModel): "context": "Comprehensive test suite" } """ - from_path_id: PathId = Field( + from_id: PathId = Field( # use the path_id from the associated Entity # or the from_id value validation_alias=AliasChoices( @@ -69,7 +69,7 @@ class RelationResponse(Relation, SQLAlchemyModel): 'from_id', ) ) - to_path_id: PathId = Field( + to_id: PathId = Field( # use the path_id from the associated Entity # or the to_id value validation_alias=AliasChoices( diff --git a/src/basic_memory/services/activity_service.py b/src/basic_memory/services/activity_service.py index a3727043..32464e0e 100644 --- a/src/basic_memory/services/activity_service.py +++ b/src/basic_memory/services/activity_service.py @@ -122,10 +122,10 @@ class ActivityService: activity_type=ActivityType.RELATION, change_type=change_type, timestamp=updated_at, - path_id=f"{relation.from_path_id}->{relation.to_path_id}", + path_id=f"{relation.from_id}->{relation.to_id}", summary=( f"{change_type.value.title()} relation: " - f"{relation.from_path_id} {relation.relation_type} {relation.to_path_id}" + f"{relation.from_id} {relation.relation_type} {relation.to_id}" ), content=relation.context ) diff --git a/src/basic_memory/services/observation_service.py b/src/basic_memory/services/observation_service.py index 6da499f6..95dcf2ad 100644 --- a/src/basic_memory/services/observation_service.py +++ b/src/basic_memory/services/observation_service.py @@ -1,9 +1,8 @@ """Service for managing observations in the database.""" -from typing import List, Sequence, Optional +from typing import List, Sequence from loguru import logger -from sqlalchemy import select from basic_memory.models import Observation as ObservationModel from basic_memory.models import Entity as EntityModel @@ -22,16 +21,18 @@ class ObservationService(BaseService[ObservationRepository]): File operations are handled by MemoryService. """ - def __init__(self, observation_repository: ObservationRepository, entity_repository: EntityRepository, file_service: FileService): + def __init__( + self, + observation_repository: ObservationRepository, + entity_repository: EntityRepository, + file_service: FileService, + ): super().__init__(observation_repository) self.entity_repository = entity_repository self.file_operations = file_service async def add_observations( - self, - path_id: str, - observations: List[ObservationCreate], - context: str | None = None + self, path_id: str, observations: List[ObservationCreate], context: str | None = None ) -> EntityModel: """Add observations to entity and update its file. @@ -44,7 +45,7 @@ class ObservationService(BaseService[ObservationRepository]): observations: List of observations with categories context: Optional shared context for all observations """ - logger.debug(f"Adding observations to entity {path_id}") + logger.debug(f"Adding observations to entity: {path_id}") try: # Get entity to update @@ -80,11 +81,7 @@ class ObservationService(BaseService[ObservationRepository]): logger.error(f"Failed to add observations: {e}") raise - async def delete_observations( - self, - path_id: str, - observations: List[str] - ) -> EntityModel: + async def delete_observations(self, path_id: str, observations: List[str]) -> EntityModel: """Delete observations from entity and update its file. Args: @@ -101,8 +98,10 @@ class ObservationService(BaseService[ObservationRepository]): # Delete observations from DB by comparing the string value to the Observation content for observation in observations: - result = await self.repository.delete_by_fields(entity_id=entity.id, content=observation) - + result = await self.repository.delete_by_fields( + entity_id=entity.id, content=observation + ) + # Write updated file _, checksum = await self.file_operations.write_entity_file(entity) await self.entity_repository.update(entity.id, {"checksum": checksum}) @@ -114,19 +113,19 @@ class ObservationService(BaseService[ObservationRepository]): logger.error(f"Failed to delete observations: {e}") raise - async def delete_by_entity(self, entity_id: int) -> bool: """Delete all observations for an entity.""" logger.debug(f"Deleting all observations for entity: {entity_id}") return await self.repository.delete_by_fields(entity_id=entity_id) - async def get_observations_by_context(self, context: str) -> Sequence[ObservationModel]: """Get all observations with a specific context.""" logger.debug(f"Getting observations for context: {context}") return await self.repository.find_by_context(context) - async def get_observations_by_category(self, category: ObservationCategory) -> Sequence[ObservationModel]: + async def get_observations_by_category( + self, category: ObservationCategory + ) -> Sequence[ObservationModel]: """Get all observations with a specific context.""" logger.debug(f"Getting observations for context: {category}") return await self.repository.find_by_category(category) diff --git a/src/basic_memory/services/relation_service.py b/src/basic_memory/services/relation_service.py index d7a55192..3b704bca 100644 --- a/src/basic_memory/services/relation_service.py +++ b/src/basic_memory/services/relation_service.py @@ -37,8 +37,8 @@ class RelationService(BaseService[RelationRepository]): for rs in relations: try: - from_entity = await self.entity_repository.get_by_path_id(rs.from_path_id) - to_entity = await self.entity_repository.get_by_path_id(rs.to_path_id) + from_entity = await self.entity_repository.get_by_path_id(rs.from_id) + to_entity = await self.entity_repository.get_by_path_id(rs.to_id) relation = RelationModel( from_id=from_entity.id, @@ -50,8 +50,8 @@ class RelationService(BaseService[RelationRepository]): await self.repository.add(relation) # Keep track of entities we need to update - entities_to_update.add(rs.from_path_id) - entities_to_update.add(rs.to_path_id) + entities_to_update.add(rs.from_id) + entities_to_update.add(rs.to_id) except Exception as e: logger.error(f"Failed to create relation: {e}") @@ -86,11 +86,11 @@ class RelationService(BaseService[RelationRepository]): try: # Delete relations from DB for relation in to_delete: - entities_to_update.add(relation.from_path_id) - entities_to_update.add(relation.to_path_id) + entities_to_update.add(relation.from_id) + entities_to_update.add(relation.to_id) relation = await self.find_relation( - relation.from_path_id, relation.to_path_id, relation.relation_type + relation.from_id, relation.to_id, relation.relation_type ) if relation: relations.append(relation) diff --git a/src/basic_memory/services/search_service.py b/src/basic_memory/services/search_service.py index ba4f1525..d9afc696 100644 --- a/src/basic_memory/services/search_service.py +++ b/src/basic_memory/services/search_service.py @@ -67,7 +67,7 @@ class SearchService: *[f"{obs.category}: {obs.content}" for obs in entity.observations], # Add relations *[ - f"{rel.relation_type} {rel.to_path_id}: {rel.context or ''}" + f"{rel.relation_type} {rel.to_entity.path_id}: {rel.context or ''}" for rel in entity.relations ], ] diff --git a/tests/api/test_knowledge_router.py b/tests/api/test_knowledge_router.py index 36af0d8b..91837b64 100644 --- a/tests/api/test_knowledge_router.py +++ b/tests/api/test_knowledge_router.py @@ -67,6 +67,7 @@ async def create_related_entities(client) -> List[RelationResponse]: # pyright: {"name": "TargetEntity", "entity_type": "test"}, ] create_response = await client.post("/knowledge/entities", json={"entities": entities}) + assert create_response.status_code == 200 created = create_response.json()["entities"] source_path_id = "source_entity" target_path_id = "target_entity" @@ -396,7 +397,7 @@ async def test_full_knowledge_flow(client: AsyncClient): # 6. Search should find all related entities search = await client.post("/search/", json={"text": "Related"}) matches = search.json()["results"] - assert len(matches) == 1 + assert len(matches) == 3 # 7. Delete main entity response = await client.post( diff --git a/tests/api/test_search_router.py b/tests/api/test_search_router.py index dec9a818..f78ea668 100644 --- a/tests/api/test_search_router.py +++ b/tests/api/test_search_router.py @@ -7,6 +7,7 @@ import pytest_asyncio from sqlalchemy import text from basic_memory import db +from basic_memory.schemas import Entity as EntitySchema from basic_memory.schemas.search import SearchItemType, SearchResponse @@ -140,7 +141,13 @@ async def test_search_empty(search_service, client): async def test_reindex(client, search_service, entity_service, test_entity, session_maker): """Test reindex endpoint.""" # Create test entity and document - await entity_service.create_entity(test_entity) + await entity_service.create_entity( EntitySchema( + name="TestEntity1", + entity_type="test", + summary="A test entity description", + observations=["this is a test observation"], + ), +) # Clear search index async with db.scoped_session(session_maker) as session: diff --git a/tests/schemas/test_schemas.py b/tests/schemas/test_schemas.py index 4a2b45d5..c3d7f8d7 100644 --- a/tests/schemas/test_schemas.py +++ b/tests/schemas/test_schemas.py @@ -56,8 +56,8 @@ def test_relation_in_validation(): """Test RelationIn validation.""" data = {"from_id": "test/123", "to_id": "test/456", "relation_type": "test"} relation = Relation.model_validate(data) - assert relation.from_path_id == "test/123" - assert relation.to_path_id == "test/456" + assert relation.from_id == "test/123" + assert relation.to_id == "test/456" assert relation.relation_type == "test" assert relation.context is None diff --git a/tests/sync/test_knowledge_sync_service.py b/tests/sync/test_knowledge_sync_service.py index 85a57f41..8159c4a8 100644 --- a/tests/sync/test_knowledge_sync_service.py +++ b/tests/sync/test_knowledge_sync_service.py @@ -144,12 +144,12 @@ async def test_update_entity_relations( relations = sorted(updated.relations, key=lambda r: r.relation_type) assert relations[0].relation_type == "depends_on" - assert relations[0].from_path_id == entity.id - assert relations[0].to_path_id == other_entity.id + assert relations[0].from_id == entity.id + assert relations[0].to_id == other_entity.id assert relations[1].relation_type == "related_to" - assert relations[1].from_path_id == entity.id - assert relations[1].to_path_id == another_entity.id + assert relations[1].from_id == entity.id + assert relations[1].to_id == another_entity.id # Check checksum set assert updated.checksum == test_checksum @@ -192,5 +192,5 @@ async def test_two_pass_sync_flow( assert updated.checksum == checksum relations = sorted(updated.relations, key=lambda r: r.relation_type) - assert relations[0].to_path_id == other_entity.id - assert relations[1].to_path_id == another_entity.id + assert relations[0].to_id == other_entity.id + assert relations[1].to_id == another_entity.id diff --git a/tests/sync/test_sync_knowledge.py b/tests/sync/test_sync_knowledge.py index 1c494b91..0461e54d 100644 --- a/tests/sync/test_sync_knowledge.py +++ b/tests/sync/test_sync_knowledge.py @@ -178,10 +178,10 @@ modified: 2024-01-01 # Verify circular reference works a_relation = entity_a.outgoing_relations[0] - assert a_relation.to_path_id == entity_b.id + assert a_relation.to_id == entity_b.id b_relation = entity_b.outgoing_relations[0] - assert b_relation.to_path_id == entity_a.id + assert b_relation.to_id == entity_a.id @pytest.mark.asyncio