mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
entity repository tests
This commit is contained in:
@@ -1,21 +1,25 @@
|
||||
"""Knowledge graph models."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional, List
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import Integer, String, Text, ForeignKey, UniqueConstraint, text, DateTime, Index
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from basic_memory.models.base import Base
|
||||
|
||||
|
||||
class Entity(Base):
|
||||
"""
|
||||
Core entity in the knowledge graph.
|
||||
|
||||
|
||||
Entities represent semantic nodes maintained by the AI layer. Each entity:
|
||||
- Has a unique numeric ID (database-generated)
|
||||
- Maps to a document file on disk (optional)
|
||||
- Maps to a document file on disk (optional)
|
||||
- Maintains a checksum for change detection
|
||||
- Tracks both source document and semantic properties
|
||||
"""
|
||||
|
||||
__tablename__ = "entity"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("entity_type", "name", name="uix_entity_type_name"),
|
||||
@@ -27,32 +31,27 @@ class Entity(Base):
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String)
|
||||
entity_type: Mapped[str] = mapped_column(String)
|
||||
|
||||
|
||||
# Content and validation
|
||||
description: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
|
||||
path: Mapped[Optional[str]] = mapped_column(String, nullable=True)
|
||||
checksum: Mapped[Optional[str]] = mapped_column(String, nullable=True)
|
||||
|
||||
|
||||
# Metadata and tracking
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime,
|
||||
server_default=text("CURRENT_TIMESTAMP")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=text("CURRENT_TIMESTAMP"))
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime,
|
||||
server_default=text("CURRENT_TIMESTAMP"),
|
||||
onupdate=text("CURRENT_TIMESTAMP")
|
||||
DateTime, server_default=text("CURRENT_TIMESTAMP"), onupdate=text("CURRENT_TIMESTAMP")
|
||||
)
|
||||
|
||||
|
||||
# Relations
|
||||
doc_id: Mapped[Optional[int]] = mapped_column(
|
||||
Integer,
|
||||
ForeignKey("documents.id", ondelete="SET NULL"),
|
||||
nullable=True
|
||||
Integer, ForeignKey("documents.id", ondelete="SET NULL"), nullable=True
|
||||
)
|
||||
|
||||
# Relationships
|
||||
observations = relationship("Observation", back_populates="entity", cascade="all, delete-orphan")
|
||||
observations = relationship(
|
||||
"Observation", back_populates="entity", cascade="all, delete-orphan"
|
||||
)
|
||||
from_relations = relationship(
|
||||
"Relation",
|
||||
back_populates="from_entity",
|
||||
@@ -66,6 +65,10 @@ class Entity(Base):
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
@property
|
||||
def relations(self):
|
||||
return self.to_relations + self.from_relations
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"Entity(id={self.id}, name='{self.name}', type='{self.entity_type}')"
|
||||
|
||||
@@ -73,18 +76,16 @@ class Entity(Base):
|
||||
class Observation(Base):
|
||||
"""
|
||||
An observation about an entity.
|
||||
|
||||
|
||||
Observations are atomic facts or notes about an entity.
|
||||
"""
|
||||
|
||||
__tablename__ = "observations"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
entity_id: Mapped[int] = mapped_column(Integer, ForeignKey("entity.id"))
|
||||
content: Mapped[str] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime,
|
||||
server_default=text("CURRENT_TIMESTAMP")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=text("CURRENT_TIMESTAMP"))
|
||||
|
||||
# Relationships
|
||||
entity = relationship("Entity", back_populates="observations")
|
||||
@@ -97,6 +98,7 @@ class Relation(Base):
|
||||
"""
|
||||
A directed relation between two entities.
|
||||
"""
|
||||
|
||||
__tablename__ = "relations"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("from_id", "to_id", "relation_type", name="uix_relation"),
|
||||
@@ -107,14 +109,11 @@ class Relation(Base):
|
||||
from_id: Mapped[int] = mapped_column(Integer, ForeignKey("entity.id"))
|
||||
to_id: Mapped[int] = mapped_column(Integer, ForeignKey("entity.id"))
|
||||
relation_type: Mapped[str] = mapped_column(String)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime,
|
||||
server_default=text("CURRENT_TIMESTAMP")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=text("CURRENT_TIMESTAMP"))
|
||||
|
||||
# Relationships
|
||||
from_entity = relationship("Entity", foreign_keys=[from_id], back_populates="from_relations")
|
||||
to_entity = relationship("Entity", foreign_keys=[to_id], back_populates="to_relations")
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"Relation(from_id={self.from_id}, to_id={self.to_id}, type='{self.relation_type}')"
|
||||
return f"Relation(from_id={self.from_id}, to_id={self.to_id}, type='{self.relation_type}')"
|
||||
|
||||
@@ -15,6 +15,10 @@ from basic_memory.repository.repository import Repository
|
||||
class EntityRepository(Repository[Entity]):
|
||||
"""Repository for Entity model."""
|
||||
|
||||
def __init__(self, session_maker: async_sessionmaker[AsyncSession]):
|
||||
"""Initialize with session maker."""
|
||||
super().__init__(session_maker, Entity)
|
||||
|
||||
async def create_entity(
|
||||
self,
|
||||
name: str,
|
||||
@@ -93,8 +97,8 @@ class EntityRepository(Repository[Entity]):
|
||||
updates: Dict[str, Any]
|
||||
) -> Optional[Entity]:
|
||||
"""Update an entity with the given fields."""
|
||||
return await self.update(str(entity_id), updates)
|
||||
return await self.update(entity_id, updates)
|
||||
|
||||
async def delete_entities_by_doc_id(self, doc_id: int) -> bool:
|
||||
"""Delete all entities associated with a document."""
|
||||
return await self.delete_by_fields(doc_id=doc_id)
|
||||
return await self.delete_by_fields(doc_id=doc_id)
|
||||
@@ -91,7 +91,7 @@ class Repository[T: Base]:
|
||||
logger.debug(f"Found {len(items)} {self.Model.__name__} records")
|
||||
return items
|
||||
|
||||
async def find_by_id(self, entity_id: str) -> Optional[T]:
|
||||
async def find_by_id(self, entity_id: int) -> Optional[T]:
|
||||
"""Fetch an entity by its unique identifier."""
|
||||
logger.debug(f"Finding {self.Model.__name__} by ID: {entity_id}")
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
@@ -117,7 +117,7 @@ class Repository[T: Base]:
|
||||
logger.debug(f"No {self.Model.__name__} found")
|
||||
return entity
|
||||
|
||||
async def find_by_ids(self, ids: List[str]) -> Sequence[T]:
|
||||
async def find_by_ids(self, ids: List[int]) -> Sequence[T]:
|
||||
"""Fetch multiple entities by their identifiers in a single query."""
|
||||
logger.debug(f"Finding {self.Model.__name__} by IDs: {ids}")
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
@@ -146,7 +146,7 @@ class Repository[T: Base]:
|
||||
session.add_all(model_list)
|
||||
return model_list
|
||||
|
||||
async def update(self, entity_id: str, entity_data: dict) -> Optional[T]:
|
||||
async def update(self, entity_id: int, entity_data: dict) -> Optional[T]:
|
||||
"""Update an entity with the given data."""
|
||||
logger.debug(f"Updating {self.Model.__name__} {entity_id} with data: {entity_data}")
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
@@ -169,7 +169,7 @@ class Repository[T: Base]:
|
||||
logger.debug(f"No {self.Model.__name__} found to update: {entity_id}")
|
||||
return None
|
||||
|
||||
async def delete(self, entity_id: str) -> bool:
|
||||
async def delete(self, entity_id: int) -> bool:
|
||||
"""Delete an entity from the database."""
|
||||
logger.debug(f"Deleting {self.Model.__name__}: {entity_id}")
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
@@ -186,8 +186,8 @@ class Repository[T: Base]:
|
||||
logger.debug(f"No {self.Model.__name__} found to delete: {entity_id}")
|
||||
return False
|
||||
|
||||
async def delete_by_ids(self, ids: List[str]) -> int:
|
||||
"""Delete records matching given field values."""
|
||||
async def delete_by_ids(self, ids: List[int]) -> int:
|
||||
"""Delete records matching given IDs."""
|
||||
logger.debug(f"Deleting {self.Model.__name__} by ids: {ids}")
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
query = delete(self.Model).where(self.primary_key.in_(ids))
|
||||
@@ -223,4 +223,4 @@ class Repository[T: Base]:
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
result = await session.execute(query)
|
||||
logger.debug("Query executed successfully")
|
||||
return result
|
||||
return result
|
||||
Reference in New Issue
Block a user