from __future__ import annotations import asyncio import uuid from collections.abc import AsyncIterator from contextlib import asynccontextmanager from datetime import UTC, datetime from typing import Any from sqlalchemy import JSON, DateTime, Integer, String, Text, UniqueConstraint, delete, select from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column class Base(DeclarativeBase): pass class ConversationMode(Base): __tablename__ = "agent_conversation_modes" __table_args__ = (UniqueConstraint("user_id", "chat_id"),) id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) user_id: Mapped[str] = mapped_column(String(128), index=True) chat_id: Mapped[str] = mapped_column(String(256), index=True) mode: Mapped[str] = mapped_column(String(16)) updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) class RunEvent(Base): __tablename__ = "agent_run_events" id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) run_id: Mapped[str] = mapped_column(String(64), index=True) user_id: Mapped[str] = mapped_column(String(128), index=True) chat_id: Mapped[str] = mapped_column(String(256), index=True) sequence: Mapped[int] = mapped_column(Integer) event_type: Mapped[str] = mapped_column(String(80), index=True) payload: Mapped[dict[str, Any]] = mapped_column(JSON, default=dict) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) class Memory(Base): __tablename__ = "agent_memories" id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: str(uuid.uuid4())) user_id: Mapped[str] = mapped_column(String(128), index=True) content: Mapped[str] = mapped_column(Text) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) class Plan(Base): __tablename__ = "agent_plans" __table_args__ = (UniqueConstraint("user_id", "chat_id"),) id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) user_id: Mapped[str] = mapped_column(String(128), index=True) chat_id: Mapped[str] = mapped_column(String(256), index=True) items: Mapped[list[dict[str, Any]]] = mapped_column(JSON, default=list) updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) class RuntimeStore: def __init__(self, url: str) -> None: self.engine: AsyncEngine = create_async_engine(url, pool_pre_ping=True) self.sessions = async_sessionmaker(self.engine, expire_on_commit=False) self._mode_locks: dict[tuple[str, str], asyncio.Lock] = {} self._lock_guard = asyncio.Lock() async def initialize(self) -> None: async with self.engine.begin() as connection: await connection.run_sync(Base.metadata.create_all) async def close(self) -> None: await self.engine.dispose() @asynccontextmanager async def session(self) -> AsyncIterator[AsyncSession]: async with self.sessions() as session: yield session async def _mode_lock(self, user_id: str, chat_id: str) -> asyncio.Lock: key = (user_id, chat_id) async with self._lock_guard: return self._mode_locks.setdefault(key, asyncio.Lock()) async def select_mode(self, user_id: str, chat_id: str, requested: str) -> str: if requested not in {"chat", "work"}: raise ValueError("Invalid conversation mode") if not chat_id: return requested lock = await self._mode_lock(user_id, chat_id) async with lock: async with self.session() as session: row = await session.scalar( select(ConversationMode).where( ConversationMode.user_id == user_id, ConversationMode.chat_id == chat_id, ) ) if row is None: session.add(ConversationMode(user_id=user_id, chat_id=chat_id, mode=requested)) await session.commit() return requested if row.mode == "work" and requested == "chat": return "work" if row.mode != requested: row.mode = requested row.updated_at = datetime.now(UTC) await session.commit() return row.mode async def append_event( self, run_id: str, user_id: str, chat_id: str, sequence: int, event_type: str, payload: dict[str, Any] | None = None, ) -> None: async with self.session() as session: session.add( RunEvent( run_id=run_id, user_id=user_id, chat_id=chat_id, sequence=sequence, event_type=event_type, payload=payload or {}, ) ) await session.commit() async def events_for_chat(self, user_id: str, chat_id: str, limit: int = 500) -> list[dict[str, Any]]: async with self.session() as session: rows = ( await session.scalars( select(RunEvent) .where(RunEvent.user_id == user_id, RunEvent.chat_id == chat_id) .order_by(RunEvent.id.desc()) .limit(max(1, min(limit, 2000))) ) ).all() return [ { "run_id": row.run_id, "sequence": row.sequence, "type": row.event_type, "payload": row.payload, "created_at": row.created_at.isoformat(), } for row in reversed(rows) ] async def remember(self, user_id: str, content: str) -> str: item = Memory(user_id=user_id, content=content) async with self.session() as session: session.add(item) await session.commit() return item.id async def recall(self, user_id: str, query: str = "", limit: int = 8) -> list[dict[str, str]]: statement = select(Memory).where(Memory.user_id == user_id) if query.strip(): statement = statement.where(Memory.content.ilike(f"%{query.strip()}%")) statement = statement.order_by(Memory.created_at.desc()).limit(max(1, min(limit, 50))) async with self.session() as session: rows = (await session.scalars(statement)).all() return [{"id": row.id, "content": row.content, "created_at": row.created_at.isoformat()} for row in rows] async def forget(self, user_id: str, memory_id: str) -> bool: async with self.session() as session: result = await session.execute(delete(Memory).where(Memory.id == memory_id, Memory.user_id == user_id)) await session.commit() return bool(result.rowcount) async def update_plan(self, user_id: str, chat_id: str, items: list[dict[str, Any]]) -> None: async with self.session() as session: row = await session.scalar(select(Plan).where(Plan.user_id == user_id, Plan.chat_id == chat_id)) if row is None: session.add(Plan(user_id=user_id, chat_id=chat_id, items=items)) else: row.items = items row.updated_at = datetime.now(UTC) await session.commit()