Files

144 lines
5.2 KiB
Python
Raw Permalink Normal View History

2026-09-09 21:55:37 +08:00
from datetime import UTC, datetime
from math import exp
from uuid import uuid4
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.model.memory import MemoryConflict, MemoryUnit
class MemoryService:
"""MySQL authoritative memory operations; vector/graph stores are projections."""
def __init__(self, session: AsyncSession) -> None:
self.session = session
@staticmethod
def should_extract_memory(
*,
conversation_content: str,
role: str,
tool_result: bool = False,
event_type: str | None = None,
) -> bool:
"""Only extract durable user facts/preferences or explicit business events."""
if not conversation_content.strip():
return False
if role == "user":
return len(conversation_content.strip()) >= 4
return tool_result or event_type in {"risk.assessment_completed", "trade.completed"}
async def recall(self, customer_id: int, *, limit: int = 10) -> list[MemoryUnit]:
now = datetime.now(UTC).replace(tzinfo=None)
result = await self.session.scalars(
select(MemoryUnit)
.where(
MemoryUnit.customer_id == customer_id,
MemoryUnit.status == "active",
(MemoryUnit.valid_until.is_(None) | (MemoryUnit.valid_until > now)),
)
.order_by(MemoryUnit.confidence.desc(), MemoryUnit.updated_at.desc())
.limit(max(1, min(limit, 100)))
)
return list(result)
async def recall_with_decay(
self, customer_id: int, query: str | None = None, *, limit: int = 10
) -> list[MemoryUnit]:
memories = await self.recall(customer_id, limit=100)
now = datetime.now(UTC).replace(tzinfo=None)
if query:
terms = {term.lower() for term in query.split() if term}
memories = [
memory
for memory in memories
if not terms
or any(term in memory.content.lower() for term in terms)
or any(term in memory.memory_key.lower() for term in terms)
]
memories.sort(
key=lambda memory: memory.confidence
* exp(-max(0, (now - memory.updated_at).days) / 365),
reverse=True,
)
return memories[: max(1, min(limit, 100))]
async def upsert(
self,
customer_id: int,
memory_key: str,
content: str,
*,
memory_type: str = "fact",
confidence: float = 0.5,
source_type: str = "conversation",
) -> MemoryUnit:
now = datetime.now(UTC).replace(tzinfo=None)
memory = await self.session.scalar(
select(MemoryUnit).where(
MemoryUnit.customer_id == customer_id,
MemoryUnit.memory_key == memory_key,
MemoryUnit.status == "active",
)
)
if memory is None:
memory = MemoryUnit(
id=0, memory_uuid=str(uuid4()), customer_id=customer_id,
memory_key=memory_key, content=content, memory_type=memory_type,
source_type=source_type, source_confidence=confidence,
confidence=confidence, status="active", valid_from=now,
version=1, created_at=now, updated_at=now,
)
self.session.add(memory)
else:
if memory.content != content:
conflict = MemoryConflict(
id=0,
left_memory_id=memory.id,
right_memory_id=memory.id,
conflict_type="content_changed",
resolution_status="resolved",
winner_memory_id=memory.id,
created_at=now,
resolved_at=now,
)
self.session.add(conflict)
memory.content = content
memory.confidence = confidence
memory.version += 1
memory.updated_at = now
await self.session.flush()
return memory
async def expire_stale(self, *, customer_id: int | None = None) -> int:
now = datetime.now(UTC).replace(tzinfo=None)
statement = select(MemoryUnit).where(
MemoryUnit.status == "active",
MemoryUnit.valid_until.is_not(None),
MemoryUnit.valid_until <= now,
)
if customer_id is not None:
statement = statement.where(MemoryUnit.customer_id == customer_id)
memories = list(await self.session.scalars(statement))
for memory in memories:
memory.status = "expired"
memory.updated_at = now
await self.session.flush()
return len(memories)
async def invalidate(self, memory_uuid: str, customer_id: int) -> bool:
memory = await self.session.scalar(
select(MemoryUnit).where(
MemoryUnit.memory_uuid == memory_uuid,
MemoryUnit.customer_id == customer_id,
MemoryUnit.status == "active",
)
)
if memory is None:
return False
memory.status = "invalidated"
memory.updated_at = datetime.now(UTC).replace(tzinfo=None)
await self.session.flush()
return True