Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions src/basic_memory/index/local_dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,8 @@
)
from basic_memory.indexing.orphan_cleanup import OrphanEntityRepository, OrphanSearchIndex
from basic_memory.indexing.relation_resolution import (
BatchRelationResolutionEntityIndexer,
BatchRelationResolutionEntityRepository,
RelationResolutionEntityIndexer,
RelationResolutionEntityRepository,
RelationResolutionLinkResolver,
Expand Down Expand Up @@ -110,6 +112,7 @@ class LocalIndexEntityRepository(
IndexedFileChecksumRepository,
CurrentMaterializedNoteEntityRepository,
OrphanEntityRepository[Entity],
BatchRelationResolutionEntityRepository,
RelationResolutionEntityRepository,
Protocol,
):
Expand Down Expand Up @@ -163,6 +166,7 @@ async def delete_by_fields(

class LocalIndexSearchService(
OrphanSearchIndex[Entity],
BatchRelationResolutionEntityIndexer,
RelationResolutionEntityIndexer,
Protocol,
):
Expand Down
155 changes: 85 additions & 70 deletions src/basic_memory/indexing/relation_resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,15 @@

import logfire
from loguru import logger
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker

from basic_memory import db
from basic_memory.indexing.models import IndexFileJobStatus
from basic_memory.models import Entity, Relation
from basic_memory.models import Entity
from basic_memory.repository.relation_repository import (
ResolvedRelationWrite,
ResolvedRelationWriteResult,
)

type EntityId = int
type AffectedEntityIds = set[EntityId]
Expand Down Expand Up @@ -63,21 +66,12 @@ async def find_unresolved_relations_for_entity(
) -> Sequence[UnresolvedRelation]:
"""Return unresolved relations for one source entity."""

# Positional-only parameters: the concrete implementation is the generic
# model repository, whose parameters are named for entities. `/` lets this
# contract name the relation id honestly without renaming the shared
# repository method.
async def update(
async def apply_resolved_targets(
self,
session: AsyncSession,
relation_id: int,
resolved_target_fields: dict[str, int | str],
/,
) -> Relation | None:
"""Apply resolved target fields (to_id, to_name) to one relation row."""

async def delete(self, session: AsyncSession, relation_id: int, /) -> bool:
"""Delete one redundant unresolved relation row."""
writes: Sequence[ResolvedRelationWrite],
) -> ResolvedRelationWriteResult:
"""Apply canonical targets and remove duplicate edges as one batch."""


class RelationResolutionEntityRepository(Protocol):
Expand All @@ -87,6 +81,17 @@ async def find_by_id(self, session: AsyncSession, entity_id: EntityId) -> Entity
"""Return one source entity by database id."""


class BatchRelationResolutionEntityRepository(Protocol):
"""Repository capability for loading an affected source-entity batch."""

async def find_by_ids(
self,
session: AsyncSession,
ids: list[EntityId],
) -> Sequence[Entity]:
"""Return source entities by database id."""


class RelationResolutionLinkResolver(Protocol):
"""Capability for resolving a relation target by link text."""

Expand All @@ -107,15 +112,22 @@ async def index_entity(self, entity: Entity) -> None:
"""Refresh derived index rows for one entity."""


class BatchRelationResolutionEntityIndexer(Protocol):
"""Capability for refreshing a batch of derived entity search rows."""

async def index_entities(self, entities: Sequence[Entity]) -> None:
"""Refresh derived index rows for a group of entities."""


@dataclass(frozen=True, slots=True)
class RepositoryRelationResolutionRuntime:
"""Resolve forward references with project-scoped repositories and services."""

session_maker: async_sessionmaker[AsyncSession]
relation_repository: RelationResolutionRelationRepository
entity_repository: RelationResolutionEntityRepository
entity_repository: BatchRelationResolutionEntityRepository
link_resolver: RelationResolutionLinkResolver
entity_indexer: RelationResolutionEntityIndexer
entity_indexer: BatchRelationResolutionEntityIndexer

async def count_unresolved_relations(self) -> int:
"""Return the current unresolved relation count for this project."""
Expand All @@ -127,6 +139,7 @@ async def resolve_relations(
entity_id: EntityId | None = None,
) -> AffectedEntityIds:
"""Resolve visible forward references and refresh affected entities."""
resolved_targets_by_link_text: dict[str, ResolvedRelationTarget | None] = {}
async with db.scoped_session(self.session_maker) as session:
if entity_id is None:
unresolved_relations = await self.relation_repository.find_unresolved_relations(
Expand All @@ -145,64 +158,66 @@ async def resolve_relations(
count=len(unresolved_relations),
)

affected_entity_ids: AffectedEntityIds = set()

for relation in unresolved_relations:
logger.trace(
"Attempting to resolve relation "
f"relation_id={relation.id} "
f"from_id={relation.from_id} "
f"to_name={relation.to_name}"
)
async with db.scoped_session(self.session_maker) as session:
resolved_entity = await self.link_resolver.resolve_link(
relation.to_name,
strict=True,
session=session,
writes: list[ResolvedRelationWrite] = []
for relation in unresolved_relations:
logger.trace(
"Attempting to resolve relation "
f"relation_id={relation.id} "
f"from_id={relation.from_id} "
f"to_name={relation.to_name}"
)

if resolved_entity is None or resolved_entity.id == relation.from_id:
continue

logger.debug(
"Resolved forward reference "
f"relation_id={relation.id} "
f"from_id={relation.from_id} "
f"to_name={relation.to_name} "
f"resolved_id={resolved_entity.id} "
f"resolved_title={resolved_entity.title}",
)
try:
async with db.scoped_session(self.session_maker) as session:
await self.relation_repository.update(
session,
relation.id,
{
"to_id": resolved_entity.id,
"to_name": resolved_entity.title,
},
if relation.to_name not in resolved_targets_by_link_text:
resolved_targets_by_link_text[
relation.to_name
] = await self.link_resolver.resolve_link(
relation.to_name,
strict=True,
session=session,
)
except IntegrityError:
with logfire.span(
"indexing.relation.resolve_conflict",
relation_id=relation.id,
relation_type=relation.relation_type,
):
# Another resolved row already represents this edge. Remove
# the redundant unresolved row so future passes do not keep
# retrying the same conflict.
async with db.scoped_session(self.session_maker) as session:
await self.relation_repository.delete(session, relation.id)
affected_entity_ids.add(relation.from_id)

for affected_entity_id in sorted(affected_entity_ids):
resolved_entity = resolved_targets_by_link_text[relation.to_name]
if resolved_entity is None or resolved_entity.id == relation.from_id:
continue

logger.debug(
"Resolved forward reference "
f"relation_id={relation.id} "
f"from_id={relation.from_id} "
f"to_name={relation.to_name} "
f"resolved_id={resolved_entity.id} "
f"resolved_title={resolved_entity.title}",
)
writes.append(
ResolvedRelationWrite(
relation_id=relation.id,
from_id=relation.from_id,
target_id=resolved_entity.id,
target_name=resolved_entity.title,
relation_type=relation.relation_type,
)
)
write_result = await self.relation_repository.apply_resolved_targets(session, writes)

affected_entity_ids: AffectedEntityIds = set(write_result.affected_entity_ids)
if write_result.duplicate_relation_ids:
with logfire.span(
"indexing.relation.resolve_conflicts",
relation_ids=write_result.duplicate_relation_ids,
conflict_count=len(write_result.duplicate_relation_ids),
):
logger.debug(
"Removed redundant unresolved relations",
relation_ids=write_result.duplicate_relation_ids,
)

if affected_entity_ids:
async with db.scoped_session(self.session_maker) as session:
source_entity = await self.entity_repository.find_by_id(
source_entities = await self.entity_repository.find_by_ids(
session,
affected_entity_id,
sorted(affected_entity_ids),
)
if source_entity is not None:
await self.entity_indexer.index_entity(source_entity)
await self.entity_indexer.index_entities(
sorted(source_entities, key=lambda entity: entity.id)
)

return affected_entity_ids

Expand Down
135 changes: 134 additions & 1 deletion src/basic_memory/repository/relation_repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from dataclasses import dataclass
from typing import Sequence, List, Optional, Any, cast

from sqlalchemy import and_, delete, select
from sqlalchemy import and_, case, delete, select, update
from sqlalchemy.engine import CursorResult
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
Expand Down Expand Up @@ -31,6 +31,25 @@ class AcceptedRelationWrite:
target_id: int | None = None


@dataclass(frozen=True, slots=True)
class ResolvedRelationWrite:
"""One unresolved relation with its canonical resolved target."""

relation_id: int
from_id: int
target_id: int
target_name: str
relation_type: str


@dataclass(frozen=True, slots=True)
class ResolvedRelationWriteResult:
"""Outcome of applying one set of resolved relation targets."""

affected_entity_ids: frozenset[int]
duplicate_relation_ids: tuple[int, ...]


class RelationRepository(Repository[Relation]):
"""Repository for Relation model with memory-specific operations."""

Expand Down Expand Up @@ -115,6 +134,120 @@ async def find_unresolved_relations_for_entity(
result = await self.execute_query(session, query)
return result.scalars().all()

async def apply_resolved_targets(
self,
session: AsyncSession,
writes: Sequence[ResolvedRelationWrite],
) -> ResolvedRelationWriteResult:
"""Resolve relation targets in one transaction without duplicate edges.

Both relation uniqueness constraints can collide when aliases resolve to
the same canonical entity. The old row-at-a-time path handled that by
deleting the later unresolved row after an ``IntegrityError``. Planning
the accepted and redundant rows up front preserves that behavior while
allowing the mutations to run as set-based statements.
"""
if not writes:
return ResolvedRelationWriteResult(frozenset(), ())

ordered_writes = sorted(writes, key=lambda write: write.relation_id)
relation_ids = {write.relation_id for write in ordered_writes}
source_entity_ids = {write.from_id for write in ordered_writes}
result = await session.execute(
select(
Relation.id,
Relation.from_id,
Relation.to_id,
Relation.to_name,
Relation.relation_type,
).where(
Relation.project_id == self.project_id,
Relation.from_id.in_(source_entity_ids),
)
)
existing_relations = result.tuples().all()

occupied_target_keys: set[tuple[int, int, str]] = set()
occupied_name_keys: set[tuple[int, str, str]] = set()
all_name_keys: set[tuple[int, str, str]] = set()
for relation_id, from_id, to_id, to_name, relation_type in existing_relations:
name_key = (from_id, to_name, relation_type)
all_name_keys.add(name_key)
if relation_id in relation_ids:
continue
occupied_name_keys.add(name_key)
if to_id is not None:
occupied_target_keys.add((from_id, to_id, relation_type))

accepted_writes: list[ResolvedRelationWrite] = []
duplicate_relation_ids: list[int] = []
for write in ordered_writes:
target_key = (write.from_id, write.target_id, write.relation_type)
name_key = (write.from_id, write.target_name, write.relation_type)
if target_key in occupied_target_keys or name_key in occupied_name_keys:
duplicate_relation_ids.append(write.relation_id)
continue
accepted_writes.append(write)
occupied_target_keys.add(target_key)
occupied_name_keys.add(name_key)

if duplicate_relation_ids:
await session.execute(
delete(Relation).where(
Relation.project_id == self.project_id,
Relation.id.in_(duplicate_relation_ids),
)
)

if accepted_writes:
accepted_relation_ids = [write.relation_id for write in accepted_writes]
temporary_names_by_relation_id: dict[int, str] = {}
for write in accepted_writes:
temporary_name = f"__basic_memory_resolving_relation_{write.relation_id}__"
while (write.from_id, temporary_name, write.relation_type) in all_name_keys:
temporary_name += "_"
temporary_names_by_relation_id[write.relation_id] = temporary_name
all_name_keys.add((write.from_id, temporary_name, write.relation_type))

# Clear both unique keys before assigning canonical targets. This
# makes alias swaps safe on databases that check uniqueness row by
# row inside a multi-row UPDATE.
await session.execute(
update(Relation)
.where(
Relation.project_id == self.project_id,
Relation.id.in_(accepted_relation_ids),
)
.values(
to_id=None,
to_name=case(temporary_names_by_relation_id, value=Relation.id),
)
.execution_options(synchronize_session=False)
)
await session.execute(
update(Relation)
.where(
Relation.project_id == self.project_id,
Relation.id.in_(accepted_relation_ids),
)
.values(
to_id=case(
{write.relation_id: write.target_id for write in accepted_writes},
value=Relation.id,
),
to_name=case(
{write.relation_id: write.target_name for write in accepted_writes},
value=Relation.id,
),
)
.execution_options(synchronize_session=False)
)

return ResolvedRelationWriteResult(
affected_entity_ids=frozenset(source_entity_ids),
duplicate_relation_ids=tuple(duplicate_relation_ids),
)

async def add_all_ignore_duplicates(
self, session: AsyncSession, relations: List[Relation]
) -> int:
Expand Down
Loading
Loading