200 lines
5.3 KiB
Python
200 lines
5.3 KiB
Python
"""
|
||
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
|