delete observation repository

This commit is contained in:
phernandez
2024-12-15 10:50:25 -06:00
parent b03e7229a9
commit b1ec3bca26
6 changed files with 342 additions and 75 deletions
+18 -11
View File
@@ -1,7 +1,6 @@
"""Router for knowledge graph operations."""
from fastapi import APIRouter
from basic_memory.deps import MemoryServiceDep
from basic_memory.schemas import (
CreateEntityRequest, CreateEntityResponse,
@@ -10,7 +9,8 @@ from basic_memory.schemas import (
EntityResponse, AddObservationsRequest, ObservationResponse,
OpenNodesRequest, OpenNodesResponse,
DeleteEntityResponse,
DeleteObservationsRequest, DeleteObservationsResponse, AddObservationsResponse, RelationResponse
DeleteObservationsRequest, DeleteObservationsResponse, AddObservationsResponse, RelationResponse,
DeleteRelationRequest
)
router = APIRouter(prefix="/knowledge", tags=["knowledge"])
@@ -66,14 +66,17 @@ async def create_relations(
return CreateRelationsResponse(relations=[RelationResponse.model_validate(relation) for relation in relations])
@router.delete("/relations/{relation_id}", response_model=DeleteEntityResponse)
@router.delete("/relations/{from_id:path}/{to_id:path}", response_model=DeleteEntityResponse)
async def delete_relation(
relation_id: int,
from_id: str,
to_id: str,
relation_type: str | None = None,
memory_service: MemoryServiceDep
) -> DeleteEntityResponse:
"""Delete a specific relation by ID."""
# TODO: Implement delete_relation in memory service
raise NotImplementedError("Delete relation not implemented yet")
"""Delete relations between entities, optionally filtered by type."""
request = DeleteRelationRequest(from_id=from_id, to_id=to_id, relation_type=relation_type)
deleted = await memory_service.delete_relations([request])
return DeleteEntityResponse(deleted=deleted)
@router.post("/observations", response_model=AddObservationsResponse)
@@ -86,14 +89,18 @@ async def add_observations(
return AddObservationsResponse(entity_id=data.entity_id, observations=[ObservationResponse.model_validate(observation) for observation in observations])
@router.delete("/observations", response_model=DeleteObservationsResponse)
@router.delete("/entities/{entity_id:path}/observations", response_model=DeleteObservationsResponse)
async def delete_observations(
entity_id: str,
data: DeleteObservationsRequest,
memory_service: MemoryServiceDep
) -> DeleteObservationsResponse:
"""Delete observations from an entity."""
# TODO: Implement delete_observations in memory service
raise NotImplementedError("Delete observations not implemented yet")
# Ensure entity_id matches the request
if data.entity_id != entity_id:
data.entity_id = entity_id
deleted = await memory_service.delete_observations(data)
return DeleteObservationsResponse(deleted=deleted)
@router.post("/search", response_model=SearchNodesResponse)
@@ -103,4 +110,4 @@ async def search_nodes(
) -> SearchNodesResponse:
"""Search for entities in the knowledge graph."""
matches = await memory_service.search_nodes(data.query)
return SearchNodesResponse(matches=[EntityResponse.model_validate(entity) for entity in matches], query=data.query)
return SearchNodesResponse(matches=[EntityResponse.model_validate(entity) for entity in matches], query=data.query)
@@ -1,6 +1,6 @@
"""Repository for managing Observation objects."""
from typing import Sequence
from sqlalchemy import select
from typing import Sequence, Dict, Any
from sqlalchemy import select, and_, delete
from basic_memory.models import Observation
from basic_memory.repository import Repository
@@ -22,4 +22,12 @@ class ObservationRepository(Repository[Observation]):
"""Find observations with a specific context."""
query = select(Observation).filter(Observation.context == context)
result = await self.execute_query(query)
return result.scalars().all()
return result.scalars().all()
async def delete_by_fields(self, **filters: Dict[str, Any]) -> bool:
"""Delete observations matching the given field values."""
conditions = [getattr(Observation, field) == value for field, value in filters.items()]
query = delete(Observation).where(and_(*conditions))
result = await self.execute_query(query)
await self.session.flush()
return result.rowcount > 0 # pyright: ignore [reportAttributeAccessIssue]
@@ -1,6 +1,6 @@
"""Service for managing observations in both filesystem and database."""
from pathlib import Path
from typing import List, Sequence
from typing import List, Sequence, Dict, Any
from sqlalchemy import select
from basic_memory.models import Observation as ObservationModel
@@ -9,8 +9,6 @@ from . import DatabaseSyncError
from basic_memory.schemas import Observation
#from basic_memory.schemas import Observation
class ObservationService:
"""
Service for managing observations in the database.
@@ -37,6 +35,45 @@ class ObservationService:
except Exception as e:
raise DatabaseSyncError(f"Failed to add observations to database: {str(e)}") from e
async def delete_observations(self, entity_id: str, contents: List[str]) -> int:
"""
Delete specific observations from an entity.
Args:
entity_id: ID of the entity
contents: List of observation contents to delete
Returns:
Number of observations deleted
"""
try:
deleted = False
for content in contents:
result = await self.observation_repo.delete_by_fields(
entity_id=entity_id,
content=content
)
if result:
deleted = True
return deleted
except Exception as e:
raise DatabaseSyncError(f"Failed to delete observations from database: {str(e)}") from e
async def delete_by_entity(self, entity_id: str) -> bool:
"""
Delete all observations for an entity.
Args:
entity_id: ID of the entity
Returns:
True if any observations were deleted
"""
try:
return await self.observation_repo.delete_by_fields(entity_id=entity_id)
except Exception as e:
raise DatabaseSyncError(f"Failed to delete observations from database: {str(e)}") from e
async def search_observations(self, query: str) -> List[ObservationModel]:
"""
Search for observations across all entities.