Files
inquiry_robot/inquiry-agent/agent/channel/clarification.py
T

200 lines
5.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
clarification / H5 token 绑定(骨架)。
本文件职责:签发与校验一次性 token;优先写 PG clarification_binding,无库则内存。
禁止:URL 携带 userid;身份以绑定时的 sender_id 为准。
"""
from __future__ import annotations
import logging
import secrets
import threading
import time
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Optional
from agent.config import get_settings
from agent.runtime_db import assert_safe_database, connect_psycopg
logger = logging.getLogger(__name__)
# H5 TTL 30 分钟(prompt/14)
DEFAULT_TTL_SECONDS = 30 * 60
@dataclass
class ClarificationRecord:
token: str
thread_id: str
wait_version: int
sender_id: str
allowed_actions: list[str]
expires_at: float
consumed: bool = False
_MEM: dict[str, ClarificationRecord] = {}
_MEM_LOCK = threading.Lock()
def issue_token(
*,
thread_id: str,
wait_version: int,
sender_id: str,
allowed_actions: list[str],
ttl_seconds: int = DEFAULT_TTL_SECONDS,
) -> str:
"""
签发 clarification token。
副作用:写 PG 或内存;返回随机 token(不进 URL 的 userid)。
"""
token = secrets.token_urlsafe(24)
expires = time.time() + ttl_seconds
rec = ClarificationRecord(
token=token,
thread_id=thread_id,
wait_version=wait_version,
sender_id=sender_id,
allowed_actions=list(allowed_actions),
expires_at=expires,
)
if _try_pg_insert(rec, ttl_seconds):
return token
with _MEM_LOCK:
_MEM[token] = rec
logger.info("clarification 内存签发 thread=%s wait=%s", thread_id, wait_version)
return token
def validate_token(
token: str,
*,
action: str = "",
consume: bool = False,
) -> Optional[ClarificationRecord]:
"""
校验 token:存在、未过期、未消费;若指定 action 须在允许集合。
consume=True:一次性消费。
"""
if not (token or "").strip():
return None
rec = _try_pg_load(token) or _mem_load(token)
if rec is None:
return None
if rec.consumed or rec.expires_at < time.time():
return None
if action and action not in rec.allowed_actions:
return None
if consume:
_consume(token)
rec.consumed = True
return rec
def _mem_load(token: str) -> Optional[ClarificationRecord]:
with _MEM_LOCK:
return _MEM.get(token)
def _consume(token: str) -> None:
if _try_pg_consume(token):
return
with _MEM_LOCK:
rec = _MEM.get(token)
if rec:
rec.consumed = True
def _try_pg_insert(rec: ClarificationRecord, ttl_seconds: int) -> bool:
settings = get_settings()
dsn = settings.resolve_pg_dsn()
if not dsn:
return False
try:
assert_safe_database(dsn)
expires = datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds)
import json
with connect_psycopg(dsn) as conn:
conn.execute(
"""
INSERT INTO clarification_binding (
token, thread_id, wait_version, sender_id, allowed_actions, expires_at
) VALUES (%s, %s, %s, %s, %s::jsonb, %s)
""",
(
rec.token,
rec.thread_id,
rec.wait_version,
rec.sender_id,
json.dumps(rec.allowed_actions, ensure_ascii=False),
expires,
),
)
conn.commit()
return True
except Exception as exc: # noqa: BLE001
logger.warning("clarification PG 写入失败,改用内存:%s", exc)
return False
def _try_pg_load(token: str) -> Optional[ClarificationRecord]:
settings = get_settings()
dsn = settings.resolve_pg_dsn()
if not dsn:
return None
try:
with connect_psycopg(dsn) as conn:
row = conn.execute(
"""
SELECT thread_id, wait_version, sender_id, allowed_actions,
EXTRACT(EPOCH FROM expires_at), consumed_at
FROM clarification_binding WHERE token = %s
""",
(token,),
).fetchone()
if not row:
return None
actions = row[3]
if isinstance(actions, str):
import json
actions = json.loads(actions)
return ClarificationRecord(
token=token,
thread_id=str(row[0]),
wait_version=int(row[1]),
sender_id=str(row[2]),
allowed_actions=list(actions or []),
expires_at=float(row[4] or 0),
consumed=row[5] is not None,
)
except Exception: # noqa: BLE001
return None
def _try_pg_consume(token: str) -> bool:
settings = get_settings()
dsn = settings.resolve_pg_dsn()
if not dsn:
return False
try:
with connect_psycopg(dsn) as conn:
conn.execute(
"""
UPDATE clarification_binding
SET consumed_at = NOW()
WHERE token = %s AND consumed_at IS NULL
""",
(token,),
)
conn.commit()
return True
except Exception: # noqa: BLE001
return False