Files
zk-data-agent/agent_platform/store.py
T
2026-07-26 06:52:12 +08:00

187 lines
7.5 KiB
Python

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()