Files
inquiry_robot/inquiry-agent/agent/runtime_db/session_store.py
T

173 lines
5.6 KiB
Python

"""
会话仓储壳:conversation_session 读写。
本文件职责:按 sender_id / thread_id 存取活动工单与打断栈。
禁止:写入主账六态;本壳不做询价业务分支。
线程:短连接;Memory 后端仅自检。
"""
from __future__ import annotations
import logging
import threading
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Optional
from agent.runtime_db import assert_safe_database, connect_psycopg
logger = logging.getLogger(__name__)
def _utcnow() -> datetime:
return datetime.now(timezone.utc)
@dataclass
class SessionRecord:
"""一条会话快照(非六态真相)。"""
thread_id: str
sender_id: str
active_inquiry_no: str = ""
interrupt_stack: list[Any] = field(default_factory=list)
meta: dict[str, Any] = field(default_factory=dict)
class MemorySessionStore:
"""内存会话壳(自检 / 无 PG)。"""
def __init__(self) -> None:
self._by_thread: dict[str, SessionRecord] = {}
self._lock = threading.Lock()
def get(self, thread_id: str) -> Optional[SessionRecord]:
with self._lock:
rec = self._by_thread.get(thread_id)
return None if rec is None else SessionRecord(
thread_id=rec.thread_id,
sender_id=rec.sender_id,
active_inquiry_no=rec.active_inquiry_no,
interrupt_stack=list(rec.interrupt_stack),
meta=dict(rec.meta),
)
def upsert(self, record: SessionRecord) -> None:
with self._lock:
self._by_thread[record.thread_id] = SessionRecord(
thread_id=record.thread_id,
sender_id=record.sender_id,
active_inquiry_no=record.active_inquiry_no,
interrupt_stack=list(record.interrupt_stack),
meta=dict(record.meta),
)
logger.info(
"session.memory.upsert thread=%s sender=%s",
record.thread_id,
record.sender_id,
)
def find_by_sender(self, sender_id: str) -> list[SessionRecord]:
with self._lock:
return [
SessionRecord(
thread_id=r.thread_id,
sender_id=r.sender_id,
active_inquiry_no=r.active_inquiry_no,
interrupt_stack=list(r.interrupt_stack),
meta=dict(r.meta),
)
for r in self._by_thread.values()
if r.sender_id == sender_id
]
class PostgresSessionStore:
"""
PG 会话壳。
须已执行 001 迁移;读写 conversation_session,无业务推导。
"""
def __init__(self, dsn: str) -> None:
assert_safe_database(dsn)
self._dsn = dsn
def get(self, thread_id: str) -> Optional[SessionRecord]:
with connect_psycopg(self._dsn) as conn:
row = conn.execute(
"""
SELECT thread_id, sender_id, active_inquiry_no, interrupt_stack, meta
FROM conversation_session WHERE thread_id = %s
""",
(thread_id,),
).fetchone()
if not row:
return None
return SessionRecord(
thread_id=str(row[0]),
sender_id=str(row[1]),
active_inquiry_no=str(row[2] or ""),
interrupt_stack=list(row[3] or []),
meta=dict(row[4] or {}),
)
def upsert(self, record: SessionRecord) -> None:
import json
with connect_psycopg(self._dsn) as conn:
conn.execute(
"""
INSERT INTO conversation_session (
thread_id, sender_id, active_inquiry_no, interrupt_stack, meta, updated_at
) VALUES (%s, %s, %s, %s::jsonb, %s::jsonb, %s)
ON CONFLICT (thread_id) DO UPDATE SET
sender_id = EXCLUDED.sender_id,
active_inquiry_no = EXCLUDED.active_inquiry_no,
interrupt_stack = EXCLUDED.interrupt_stack,
meta = EXCLUDED.meta,
updated_at = EXCLUDED.updated_at
""",
(
record.thread_id,
record.sender_id,
record.active_inquiry_no,
json.dumps(record.interrupt_stack, ensure_ascii=False),
json.dumps(record.meta, ensure_ascii=False),
_utcnow(),
),
)
conn.commit()
logger.info("session.pg.upsert thread=%s", record.thread_id)
def find_by_sender(self, sender_id: str) -> list[SessionRecord]:
with connect_psycopg(self._dsn) as conn:
rows = conn.execute(
"""
SELECT thread_id, sender_id, active_inquiry_no, interrupt_stack, meta
FROM conversation_session WHERE sender_id = %s
""",
(sender_id,),
).fetchall()
return [
SessionRecord(
thread_id=str(r[0]),
sender_id=str(r[1]),
active_inquiry_no=str(r[2] or ""),
interrupt_stack=list(r[3] or []),
meta=dict(r[4] or {}),
)
for r in rows
]
def create_session_store(*, dsn: str = "", backend: str = "auto"):
"""
工厂:有 DSN 且 backend!=memory → PG;否则 Memory。
框架壳,不自动建业务会话。
"""
if backend == "memory" or not dsn:
return MemorySessionStore()
return PostgresSessionStore(dsn)