173 lines
5.6 KiB
Python
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)
|