Files
leocrm/app/plugins/builtins/graph_rag/services.py
T

227 lines
7.4 KiB
Python
Raw Normal View History

"""GraphRAG services — relationship management and BFS graph traversal."""
from __future__ import annotations
import uuid
from collections import deque
from typing import Any
from sqlalchemy import select, text as sql_text
from sqlalchemy.ext.asyncio import AsyncSession
from app.plugins.builtins.graph_rag.models import EntityRelationship
async def create_relationship(
db: AsyncSession,
tenant_id: uuid.UUID,
source_type: str,
source_id: uuid.UUID,
target_type: str,
target_id: uuid.UUID,
relationship_type: str,
metadata: dict[str, Any] | None = None,
owner_id: uuid.UUID | None = None,
) -> dict[str, Any]:
"""Create a new relationship between two entities.
Args:
db: Database session.
tenant_id: Tenant UUID.
source_type: Entity type of the source.
source_id: UUID of the source entity.
target_type: Entity type of the target.
target_id: UUID of the target entity.
relationship_type: Type of relationship.
metadata: Optional metadata dict.
owner_id: Optional owner user UUID.
Returns:
Dict with the created relationship data.
"""
# Check for duplicate
existing = await db.execute(
select(EntityRelationship).where(
EntityRelationship.tenant_id == tenant_id,
EntityRelationship.source_type == source_type,
EntityRelationship.source_id == source_id,
EntityRelationship.target_type == target_type,
EntityRelationship.target_id == target_id,
EntityRelationship.relationship_type == relationship_type,
EntityRelationship.deleted_at.is_(None),
)
)
if existing.scalar_one_or_none() is not None:
return {"error": "Relationship already exists", "code": "duplicate"}
rel = EntityRelationship(
tenant_id=tenant_id,
source_type=source_type,
source_id=source_id,
target_type=target_type,
target_id=target_id,
relationship_type=relationship_type,
meta=metadata or {},
owner_id=owner_id,
)
db.add(rel)
await db.flush()
await db.refresh(rel)
return {
"id": str(rel.id),
"source_type": rel.source_type,
"source_id": str(rel.source_id),
"target_type": rel.target_type,
"target_id": str(rel.target_id),
"relationship_type": rel.relationship_type,
"metadata": rel.meta,
"owner_id": str(rel.owner_id) if rel.owner_id else None,
"created_at": rel.created_at.isoformat() if rel.created_at else None,
}
async def traverse_graph(
db: AsyncSession,
tenant_id: uuid.UUID,
source_type: str,
source_id: uuid.UUID,
max_hops: int = 3,
relationship_types: list[str] | None = None,
) -> dict[str, Any]:
"""BFS traversal of the knowledge graph from a starting node.
Uses Breadth-First Search to find all entities reachable within
`max_hops` steps, following directed relationships.
Args:
db: Database session.
tenant_id: Tenant UUID.
source_type: Entity type of the starting node.
source_id: UUID of the starting node.
max_hops: Maximum traversal depth (1-10).
relationship_types: Optional filter to follow only specific relationship types.
Returns:
Dict with nodes and edges discovered during traversal.
"""
visited: set[tuple[str, str]] = set()
nodes: list[dict[str, Any]] = []
edges: list[dict[str, Any]] = []
# BFS queue: (source_type, source_id, depth, path)
queue: deque[tuple[str, str, int, list[str]]] = deque()
start_key = (source_type, str(source_id))
visited.add(start_key)
nodes.append({
"entity_type": source_type,
"entity_id": str(source_id),
"depth": 0,
"path": [],
})
queue.append((source_type, str(source_id), 0, []))
while queue:
current_type, current_id, depth, path = queue.popleft()
if depth >= max_hops:
continue
# Build query for outgoing relationships
stmt = select(EntityRelationship).where(
EntityRelationship.tenant_id == tenant_id,
EntityRelationship.source_type == current_type,
EntityRelationship.source_id == uuid.UUID(current_id),
EntityRelationship.deleted_at.is_(None),
)
if relationship_types:
stmt = stmt.where(EntityRelationship.relationship_type.in_(relationship_types))
result = await db.execute(stmt)
relationships = result.scalars().all()
for rel in relationships:
target_key = (rel.target_type, str(rel.target_id))
new_path = path + [rel.relationship_type]
edges.append({
"source_type": rel.source_type,
"source_id": str(rel.source_id),
"target_type": rel.target_type,
"target_id": str(rel.target_id),
"relationship_type": rel.relationship_type,
"metadata": rel.meta,
})
if target_key not in visited:
visited.add(target_key)
nodes.append({
"entity_type": rel.target_type,
"entity_id": str(rel.target_id),
"depth": depth + 1,
"path": new_path,
})
queue.append((rel.target_type, str(rel.target_id), depth + 1, new_path))
# Also traverse incoming relationships (bidirectional graph)
stmt_in = select(EntityRelationship).where(
EntityRelationship.tenant_id == tenant_id,
EntityRelationship.target_type == current_type,
EntityRelationship.target_id == uuid.UUID(current_id),
EntityRelationship.deleted_at.is_(None),
)
if relationship_types:
stmt_in = stmt_in.where(EntityRelationship.relationship_type.in_(relationship_types))
result_in = await db.execute(stmt_in)
relationships_in = result_in.scalars().all()
for rel in relationships_in:
source_key = (rel.source_type, str(rel.source_id))
new_path = path + [f"inverse_{rel.relationship_type}"]
edges.append({
"source_type": rel.source_type,
"source_id": str(rel.source_id),
"target_type": rel.target_type,
"target_id": str(rel.target_id),
"relationship_type": rel.relationship_type,
"metadata": rel.meta,
})
if source_key not in visited:
visited.add(source_key)
nodes.append({
"entity_type": rel.source_type,
"entity_id": str(rel.source_id),
"depth": depth + 1,
"path": new_path,
})
queue.append((rel.source_type, str(rel.source_id), depth + 1, new_path))
return {
"nodes": nodes,
"edges": edges,
"total_nodes": len(nodes),
"total_edges": len(edges),
}
async def delete_relationship(
db: AsyncSession,
tenant_id: uuid.UUID,
relationship_id: uuid.UUID,
) -> bool:
"""Delete a relationship by ID."""
result = await db.execute(
select(EntityRelationship).where(
EntityRelationship.id == relationship_id,
EntityRelationship.tenant_id == tenant_id,
)
)
rel = result.scalar_one_or_none()
if rel is None:
return False
await db.delete(rel)
return True