Files
microfish/backend/app/services/memory_tools.py
Kunthawat Greethong 8b84378fe1 feat: SaaS foundation for CrowdSight
Elevate MiroFish/CrowdSight from single-container dev to a SaaS foundation:

- Local memory backend (Zep-compatible): memory services/models, local graph
  builder + updater, AgentActivity seam, import-boundary isolation; Zep stays
  default, local is opt-in behind MEMORY_BACKEND. Semantic parity not yet proven.
- Durable product persistence: projects/simulations/reports schema (migration
  0007) + tenant/owner-scoped ProductRepository + dual-write + scoped_project
  read-first + ArtifactStore abstraction; durable JobQueue + worker.py.
- SaaS hardening: durable RateLimiter (wired to login), UsageService (LLM
  accounting), redacted AuditService, idempotency, CORS allowlist, safe API
  errors, single-use PasswordResetService + endpoints (covers invite-pending).
- Exactly 3 roles (super_admin/admin/user) with tenant authz policy.
- Admin UI: GET/POST/PATCH /api/admin/users + GET/PUT /api/admin/settings
  (super-admin only, encrypted/masked); AdminView.vue + SettingsView.vue with
  admin/super-admin route guards, th/en i18n.
- Production deploy topology: multi-stage Dockerfile (frontend build + gunicorn
  wsgi + nginx SPA-proxy + supervisord worker), backend/wsgi.py, gunicorn dep.

Backend 197 passed; frontend 10 tests + build green. ruff unavailable (gap).
No commit of credentials; secrets handled via env/.env.example.
Deferred: Zep semantic A/B parity, object storage cutover, mobile QA, EasyPanel
container build of deploy topology.
2026-08-31 13:05:21 +07:00

507 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Local compatibility tools for the legacy ZepToolsService result contract."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Optional
from .memory_repository import SqlAlchemyMemoryRepository
@dataclass
class LocalSearchResult:
facts: list[str]
edges: list[dict[str, Any]]
nodes: list[dict[str, Any]]
query: str
total_count: int
def to_dict(self) -> dict[str, Any]:
return {
"facts": self.facts,
"edges": self.edges,
"nodes": self.nodes,
"query": self.query,
"total_count": self.total_count,
}
def to_text(self) -> str:
parts = [f"Search query: {self.query}", f"Found {self.total_count} relevant items"]
if self.facts:
parts.append("\n### Relevant facts:")
parts.extend(f"{index}. {fact}" for index, fact in enumerate(self.facts, 1))
return "\n".join(parts)
@dataclass
class LocalNodeInfo:
uuid: str
name: str
labels: list[str]
summary: str
attributes: dict[str, Any]
def to_dict(self) -> dict[str, Any]:
return {
"uuid": self.uuid,
"name": self.name,
"labels": self.labels,
"summary": self.summary,
"attributes": self.attributes,
}
@dataclass
class LocalInsightForgeResult:
query: str
simulation_requirement: str
sub_queries: list[str]
semantic_facts: list[str] = field(default_factory=list)
entity_insights: list[dict[str, Any]] = field(default_factory=list)
relationship_chains: list[str] = field(default_factory=list)
total_facts: int = 0
total_entities: int = 0
total_relationships: int = 0
def to_text(self) -> str:
parts = [
"## Local Memory Deep Analysis",
f"Analysis question: {self.query}",
f"Prediction scenario: {self.simulation_requirement}",
f"\n### Statistics\n- Facts: {self.total_facts}\n- Entities: {self.total_entities}\n- Relationships: {self.total_relationships}",
]
if self.semantic_facts:
parts.append("\n### Key facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.semantic_facts, 1)))
if self.entity_insights:
parts.append(
"\n### Entities\n"
+ "\n".join(
f"- {item.get('name', 'Unknown')} ({item.get('type', 'Entity')}): {item.get('summary', '')}"
for item in self.entity_insights
)
)
if self.relationship_chains:
parts.append("\n### Relationships\n" + "\n".join(f"- {chain}" for chain in self.relationship_chains))
return "\n".join(parts)
@dataclass
class LocalPanoramaResult:
query: str
all_nodes: list[LocalNodeInfo] = field(default_factory=list)
all_edges: list[LocalEdgeInfo] = field(default_factory=list)
active_facts: list[str] = field(default_factory=list)
historical_facts: list[str] = field(default_factory=list)
def to_text(self) -> str:
parts = [
"## Local Memory Panorama",
f"Query: {self.query}",
f"\n### Statistics\n- Total nodes: {len(self.all_nodes)}\n- Total edges: {len(self.all_edges)}\n- Active facts: {len(self.active_facts)}\n- Historical facts: {len(self.historical_facts)}",
]
if self.active_facts:
parts.append("\n### Active facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.active_facts, 1)))
if self.historical_facts:
parts.append("\n### Historical facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.historical_facts, 1)))
if self.all_nodes:
parts.append("\n### Entities\n" + "\n".join(f"- {node.name}: {node.summary}" for node in self.all_nodes))
return "\n".join(parts)
@dataclass
class LocalInterviewResult:
interview_topic: str
summary: str
def to_text(self) -> str:
return (
"## Local Memory Interview\n"
f"Topic: {self.interview_topic}\n\n"
f"{self.summary}"
)
@dataclass
class LocalEdgeInfo:
uuid: str
name: str
fact: str
source_node_uuid: str
target_node_uuid: str
source_node_name: Optional[str] = None
target_node_name: Optional[str] = None
created_at: Optional[str] = None
valid_at: Optional[str] = None
invalid_at: Optional[str] = None
expired_at: Optional[str] = None
def to_dict(self) -> dict[str, Any]:
return {
"uuid": self.uuid,
"name": self.name,
"fact": self.fact,
"source_node_uuid": self.source_node_uuid,
"target_node_uuid": self.target_node_uuid,
"source_node_name": self.source_node_name,
"target_node_name": self.target_node_name,
"created_at": self.created_at,
"valid_at": self.valid_at,
"invalid_at": self.invalid_at,
"expired_at": self.expired_at,
}
@property
def is_expired(self) -> bool:
return self.expired_at is not None
@property
def is_invalid(self) -> bool:
return self.invalid_at is not None
class LocalMemoryTools:
def __init__(self, session_or_repository, *, organization_id: str | None = None, graph_id: str | None = None):
if isinstance(session_or_repository, SqlAlchemyMemoryRepository):
self.repository = session_or_repository
else:
if organization_id is None or graph_id is None:
raise ValueError("memory_tools_scope_required")
self.repository = SqlAlchemyMemoryRepository(
session_or_repository,
organization_id=organization_id,
graph_id=graph_id,
)
def _validate_graph_id(self, graph_id: str | None = None) -> None:
if graph_id is not None and graph_id != self.repository.graph_id:
raise ValueError("memory_graph_scope_conflict")
@staticmethod
def _node_info(node) -> LocalNodeInfo:
return LocalNodeInfo(
uuid=node.id,
name=node.canonical_name,
labels=list(node.labels or []),
summary=node.summary or "",
attributes=dict(node.attributes or {}),
)
def _edge_info(self, edge, nodes_by_id: dict[str, Any]) -> LocalEdgeInfo:
source = nodes_by_id.get(edge.source_node_id)
target = nodes_by_id.get(edge.target_node_id)
return LocalEdgeInfo(
uuid=edge.id,
name=edge.relation,
fact=edge.fact,
source_node_uuid=edge.source_node_id,
target_node_uuid=edge.target_node_id,
source_node_name=source.canonical_name if source else None,
target_node_name=target.canonical_name if target else None,
created_at=edge.created_at.isoformat() if edge.created_at else None,
valid_at=edge.valid_at.isoformat() if edge.valid_at else None,
invalid_at=edge.invalid_at.isoformat() if edge.invalid_at else None,
expired_at=edge.expired_at.isoformat() if edge.expired_at else None,
)
def get_all_nodes(self, *, limit: int = 10_000) -> list[LocalNodeInfo]:
return [self._node_info(node) for node in self.repository.list_nodes(limit=limit)]
def get_all_edges(self, *, limit: int = 20_000) -> list[LocalEdgeInfo]:
nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)}
return [self._edge_info(edge, nodes) for edge in self.repository.list_edges(limit=limit)]
@staticmethod
def _score(query: str, *values: str) -> int:
normalized_query = query.casefold().strip()
haystack = " ".join(value or "" for value in values).casefold()
if not normalized_query or not haystack:
return 0
if normalized_query in haystack:
return 100
return sum(10 for token in normalized_query.split() if token in haystack)
def search_graph(
self,
query: str,
*,
limit: int = 10,
scope: str = "edges",
graph_id: str | None = None,
) -> LocalSearchResult:
self._validate_graph_id(graph_id)
if scope not in {"edges", "nodes", "both"}:
raise ValueError("invalid_memory_search_scope")
safe_limit = min(max(int(limit), 1), 100)
facts: list[str] = []
edges: list[dict[str, Any]] = []
nodes: list[dict[str, Any]] = []
all_nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)}
if scope in {"edges", "both"}:
scored_edges = [
(self._score(query, edge.fact, edge.relation), edge)
for edge in self.repository.list_edges(limit=20_000)
]
for score, edge in sorted(scored_edges, key=lambda item: (-item[0], item[1].id))[:safe_limit]:
if score <= 0:
continue
edge_info = self._edge_info(edge, all_nodes).to_dict()
edges.append(edge_info)
if edge.fact:
facts.append(edge.fact)
if scope in {"nodes", "both"}:
scored_nodes = [
(self._score(query, node.canonical_name, node.summary), node)
for node in all_nodes.values()
]
for score, node in sorted(scored_nodes, key=lambda item: (-item[0], item[1].id))[:safe_limit]:
if score <= 0:
continue
node_info = self._node_info(node).to_dict()
node_info.pop("attributes", None)
nodes.append(node_info)
if node.summary:
facts.append(f"[{node.canonical_name}]: {node.summary}")
return LocalSearchResult(
facts=facts,
edges=edges,
nodes=nodes,
query=query,
total_count=len(facts),
)
def quick_search(
self,
query: str,
*,
limit: int = 10,
graph_id: str | None = None,
) -> LocalSearchResult:
return self.search_graph(query, limit=limit, scope="edges", graph_id=graph_id)
def get_node_detail(self, node_uuid: str) -> LocalNodeInfo | None:
node = self.repository.get_node(node_uuid)
return self._node_info(node) if node is not None else None
def get_node_edges(self, node_uuid: str, *, limit: int = 500) -> list[LocalEdgeInfo]:
nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)}
return [self._edge_info(edge, nodes) for edge in self.repository.get_node_edges(node_uuid, limit=limit)]
def insight_forge(
self,
*,
graph_id: str,
query: str,
simulation_requirement: str = "",
report_context: str = "",
) -> LocalInsightForgeResult:
self._validate_graph_id(graph_id)
result = self.search_graph(query, limit=100, scope="edges", graph_id=graph_id)
facts = list(dict.fromkeys(result.facts))
nodes_by_id = {
node.id: node for node in self.repository.list_nodes(limit=10_000)
}
related_node_ids = list(
dict.fromkeys(
node_id
for edge in result.edges
for node_id in (edge.get("source_node_uuid"), edge.get("target_node_uuid"))
if node_id
)
)
entity_insights = []
for node_id in related_node_ids:
node = nodes_by_id.get(node_id)
if node is None:
continue
entity_insights.append(
{
"uuid": node.id,
"name": node.canonical_name,
"type": next(
(label for label in (node.labels or []) if label not in {"Entity", "Node"}),
"Entity",
),
"summary": node.summary or "",
"related_facts": [
fact for fact in facts if node.canonical_name.casefold() in fact.casefold()
],
}
)
relationship_chains = []
seen_chains = set()
for edge in result.edges:
chain = (
f"{edge.get('source_node_name') or edge.get('source_node_uuid')} "
f"--[{edge.get('name', '')}]--> "
f"{edge.get('target_node_name') or edge.get('target_node_uuid')}"
)
if chain not in seen_chains:
seen_chains.add(chain)
relationship_chains.append(chain)
return LocalInsightForgeResult(
query=query,
simulation_requirement=simulation_requirement,
sub_queries=[query] if query else [],
semantic_facts=facts,
entity_insights=entity_insights,
relationship_chains=relationship_chains,
total_facts=len(facts),
total_entities=len(entity_insights),
total_relationships=len(relationship_chains),
)
def panorama_search(
self,
*,
graph_id: str,
query: str,
include_expired: bool = True,
limit: int = 50,
) -> LocalPanoramaResult:
self._validate_graph_id(graph_id)
safe_limit = min(max(int(limit), 1), 100)
nodes = self.get_all_nodes(limit=10_000)
edges = self.get_all_edges(limit=20_000)
active_facts = []
historical_facts = []
for edge in edges:
if not edge.fact:
continue
if edge.is_invalid or edge.is_expired:
valid_at = edge.valid_at or "Unknown"
invalid_at = edge.invalid_at or edge.expired_at or "Unknown"
historical_facts.append(f"[{valid_at} - {invalid_at}] {edge.fact}")
else:
active_facts.append(edge.fact)
query_lower = query.casefold()
keywords = [
word.strip()
for word in query_lower.replace(",", " ").replace(",", " ").split()
if len(word.strip()) > 1
]
def relevance_score(fact: str) -> int:
fact_lower = fact.casefold()
score = 100 if query_lower in fact_lower else 0
return score + sum(10 for keyword in keywords if keyword in fact_lower)
active_facts.sort(key=relevance_score, reverse=True)
historical_facts.sort(key=relevance_score, reverse=True)
return LocalPanoramaResult(
query=query,
all_nodes=nodes,
all_edges=edges,
active_facts=active_facts[:safe_limit],
historical_facts=historical_facts[:safe_limit] if include_expired else [],
)
def get_entity_summary(self, *, graph_id: str, entity_name: str) -> dict[str, Any]:
self._validate_graph_id(graph_id)
normalized = " ".join((entity_name or "").casefold().split())
node = next(
(
item
for item in self.repository.list_nodes(limit=10_000)
if item.normalized_name == normalized
),
None,
)
if node is None:
return {"entity_name": entity_name, "found": False, "summary": "", "related_facts": []}
edges = self.get_node_edges(node.id, limit=500)
return {
"entity_name": node.canonical_name,
"found": True,
"uuid": node.id,
"labels": list(node.labels or []),
"summary": node.summary or "",
"attributes": dict(node.attributes or {}),
"related_facts": [edge.fact for edge in edges],
}
def get_entities_by_type(self, *, graph_id: str, entity_type: str) -> list[LocalNodeInfo]:
self._validate_graph_id(graph_id)
target = (entity_type or "").casefold().strip()
return [
self._node_info(node)
for node in self.repository.list_nodes(limit=10_000)
if target in {str(label).casefold() for label in (node.labels or [])}
]
def interview_agents(
self,
*,
simulation_id: str,
interview_requirement: str,
simulation_requirement: str,
max_agents: int = 5,
) -> LocalInterviewResult:
del simulation_id, max_agents
return LocalInterviewResult(
interview_topic=interview_requirement,
summary=(
"Local memory stores graph facts, not simulation transcripts. "
"Use the retrieved entities and relationships as evidence; no synthetic interview response was generated."
),
)
def get_simulation_context(
self,
*,
graph_id: str,
simulation_requirement: str,
limit: int = 30,
) -> dict[str, Any]:
self._validate_graph_id(graph_id)
result = self.search_graph(
simulation_requirement,
limit=limit,
scope="both",
graph_id=graph_id,
)
nodes = self.repository.list_nodes(limit=10_000)
entities = [
{
"name": node.canonical_name,
"type": next((label for label in (node.labels or []) if label not in {"Entity", "Node"}), "Entity"),
"summary": node.summary or "",
}
for node in nodes
if any(label not in {"Entity", "Node"} for label in (node.labels or []))
]
return {
"simulation_requirement": simulation_requirement,
"related_facts": result.facts,
"graph_statistics": self.get_graph_statistics(graph_id),
"entities": entities[: max(int(limit), 1)],
"total_entities": len(entities),
}
def get_graph_statistics(self, graph_id: str | None = None) -> dict[str, Any]:
self._validate_graph_id(graph_id)
nodes = self.repository.list_nodes(limit=10_000)
edges = self.repository.list_edges(limit=20_000)
entity_types: dict[str, int] = {}
for node in nodes:
for label in node.labels or []:
if label not in {"Entity", "Node"}:
entity_types[label] = entity_types.get(label, 0) + 1
relation_types: dict[str, int] = {}
for edge in edges:
relation_types[edge.relation] = relation_types.get(edge.relation, 0) + 1
return {
"graph_id": self.repository.graph_id,
"node_count": len(nodes),
"edge_count": len(edges),
"total_nodes": len(nodes),
"total_edges": len(edges),
"entity_types": entity_types,
"relation_types": relation_types,
}