deletes for repositories

This commit is contained in:
phernandez
2024-12-15 11:31:51 -06:00
parent b1ec3bca26
commit fa4e55f75f
10 changed files with 789 additions and 489 deletions
+1 -1
View File
@@ -68,10 +68,10 @@ async def create_relations(
@router.delete("/relations/{from_id:path}/{to_id:path}", response_model=DeleteEntityResponse)
async def delete_relation(
memory_service: MemoryServiceDep,
from_id: str,
to_id: str,
relation_type: str | None = None,
memory_service: MemoryServiceDep
) -> DeleteEntityResponse:
"""Delete relations between entities, optionally filtered by type."""
request = DeleteRelationRequest(from_id=from_id, to_id=to_id, relation_type=relation_type)
+36 -2
View File
@@ -1,6 +1,6 @@
"""Base repository implementation."""
from typing import Type, Optional, Any, Sequence, TypeVar, List
from sqlalchemy import select, func, Select, Executable, inspect, Result, Column, insert
from typing import Type, Optional, Any, Sequence, TypeVar, List, Dict
from sqlalchemy import select, func, Select, Executable, inspect, Result, Column, insert, and_, delete
from sqlalchemy.exc import NoResultFound
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import Mapped
@@ -22,6 +22,17 @@ class Repository[T: Base]:
logger.debug(f"Initialized {self.__class__.__name__} for {Model.__name__}")
logger.debug(f"Valid columns: {self.valid_columns}")
def select(self, *entities: Any) -> Select:
"""Create a new SELECT statement.
Returns:
A SQLAlchemy Select object configured with the provided entities
or this repository's model if no entities provided.
"""
if not entities:
entities = (self.Model,)
return select(*entities)
async def refresh(self, instance: T, relationships: list[str] | None = None) -> None:
"""Refresh instance and optionally specified relationships."""
@@ -162,6 +173,29 @@ class Repository[T: Base]:
logger.exception(f"Failed to delete {self.Model.__name__}: {entity_id}")
raise
async def delete_by_fields(self, **filters: Dict[str, Any]) -> bool:
"""
Delete records matching given field values.
Args:
**filters: Field names and values to filter by
Returns:
bool: True if any records were deleted
"""
logger.debug(f"Deleting {self.Model.__name__} by fields: {filters}")
try:
conditions = [getattr(self.Model, field) == value for field, value in filters.items()]
query = delete(self.Model).where(and_(*conditions))
result = await self.execute_query(query)
await self.session.flush()
deleted = result.rowcount > 0
logger.debug(f"Deleted {result.rowcount} records")
return deleted # pyright: ignore [reportAttributeAccessIssue]
except Exception as e:
logger.exception(f"Failed to delete {self.Model.__name__} by fields")
raise
async def count(self, query: Executable | None = None) -> int:
"""Count entities in the database table."""
try:
@@ -1,6 +1,6 @@
"""Repository for managing Observation objects."""
from typing import Sequence, Dict, Any
from sqlalchemy import select, and_, delete
from typing import Sequence
from sqlalchemy import select
from basic_memory.models import Observation
from basic_memory.repository import Repository
@@ -22,12 +22,4 @@ 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()
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]
return result.scalars().all()
@@ -1,6 +1,6 @@
"""Repository for managing Relation objects."""
from typing import Sequence, Any, Dict
from sqlalchemy import select, and_, delete
from typing import Sequence
from sqlalchemy import select, and_
from basic_memory.models import Relation
from basic_memory.repository import Repository
@@ -27,12 +27,4 @@ class RelationRepository(Repository[Relation]):
"""Find all relations of a specific type."""
query = select(Relation).filter(Relation.relation_type == relation_type)
result = await self.execute_query(query)
return result.scalars().all()
async def delete_by_fields(self, **filters: Dict[str, Any]) -> bool:
"""Delete relations matching the given field values."""
conditions = [getattr(Relation, field) == value for field, value in filters.items()]
query = delete(Relation).where(and_(*conditions))
result = await self.execute_query(query)
await self.session.flush()
return result.rowcount > 0
return result.scalars().all()
+30 -1
View File
@@ -1,5 +1,6 @@
"""Service for managing relations in the database."""
from pathlib import Path
from typing import List, Dict, Any
from basic_memory.repository.relation_repository import RelationRepository
from basic_memory.schemas import Entity, Relation
@@ -46,4 +47,32 @@ class RelationService:
return result
except Exception as e:
raise DatabaseSyncError(f"Failed to delete relation: {str(e)}") from e
raise DatabaseSyncError(f"Failed to delete relation: {str(e)}") from e
async def delete_relations(self, relations: List[Dict[str, Any]]) -> bool:
"""
Delete relations matching specified criteria.
Args:
relations: List of dicts with from_id, to_id, and optional relation_type
Returns:
True if any relations were deleted
"""
try:
deleted = False
for relation in relations:
filters = {
'from_id': relation['from_id'],
'to_id': relation['to_id']
}
if 'relation_type' in relation:
filters['relation_type'] = relation['relation_type']
result = await self.relation_repo.delete_by_fields(**filters)
if result:
deleted = True
return deleted
except Exception as e:
raise DatabaseSyncError(f"Failed to delete relations: {str(e)}") from e