- Added `auth.py` for mock login and JWT issuance. - Introduced `chat.py` for handling chat requests with role-based access control. - Enhanced `main.py` to include new routers and middleware for tracing. - Implemented input validation in `input_guard.py` to prevent SQL injection. - Created repositories for managing agent sessions and audit logs. - Added exception handling for authorization errors. - Updated settings to include JWT configuration. - Introduced tests for authentication and input validation.
91 lines
2.9 KiB
Python
91 lines
2.9 KiB
Python
"""Agent 会话与消息持久化。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import uuid
|
|
|
|
from sqlalchemy import text
|
|
from sqlalchemy.engine import Engine
|
|
|
|
from app.config.database import get_agent_engine
|
|
from app.model.schemas import AgentType, AuthContext
|
|
|
|
|
|
class AgentSessionRepository:
|
|
def __init__(self, engine: Engine | None = None) -> None:
|
|
self._engine = engine or get_agent_engine()
|
|
|
|
def ensure_session(
|
|
self,
|
|
ctx: AuthContext,
|
|
session_id: str | None,
|
|
customer_id: str | None,
|
|
) -> str:
|
|
sid = session_id or str(uuid.uuid4())
|
|
sql = text(
|
|
"""
|
|
INSERT INTO agent_session
|
|
(session_id, trace_id, agent_type, actor_id, actor_role,
|
|
customer_id, advisor_id, status)
|
|
VALUES
|
|
(:session_id, :trace_id, :agent_type, :actor_id, :actor_role,
|
|
:customer_id, :advisor_id, 'active')
|
|
ON DUPLICATE KEY UPDATE
|
|
trace_id = VALUES(trace_id),
|
|
updated_at = CURRENT_TIMESTAMP(3)
|
|
"""
|
|
)
|
|
actor_role = ctx.roles[0] if ctx.roles else "unknown"
|
|
advisor_id = ctx.sub if "advisor" in ctx.roles else None
|
|
with self._engine.begin() as conn:
|
|
conn.execute(
|
|
sql,
|
|
{
|
|
"session_id": sid,
|
|
"trace_id": ctx.trace_id,
|
|
"agent_type": ctx.agent_type,
|
|
"actor_id": ctx.sub,
|
|
"actor_role": actor_role,
|
|
"customer_id": customer_id,
|
|
"advisor_id": advisor_id,
|
|
},
|
|
)
|
|
return sid
|
|
|
|
def next_seq(self, session_id: str) -> int:
|
|
sql = text("SELECT COALESCE(MAX(seq_no), 0) + 1 AS n FROM agent_message WHERE session_id = :sid")
|
|
with self._engine.connect() as conn:
|
|
row = conn.execute(sql, {"sid": session_id}).mappings().first()
|
|
return int(row["n"]) if row else 1
|
|
|
|
def insert_message(
|
|
self,
|
|
*,
|
|
session_id: str,
|
|
trace_id: str,
|
|
seq_no: int,
|
|
role: str,
|
|
content: str,
|
|
has_disclaimer: bool = False,
|
|
) -> None:
|
|
sql = text(
|
|
"""
|
|
INSERT INTO agent_message
|
|
(session_id, trace_id, seq_no, role, content, has_disclaimer)
|
|
VALUES
|
|
(:session_id, :trace_id, :seq_no, :role, :content, :has_disclaimer)
|
|
"""
|
|
)
|
|
with self._engine.begin() as conn:
|
|
conn.execute(
|
|
sql,
|
|
{
|
|
"session_id": session_id,
|
|
"trace_id": trace_id,
|
|
"seq_no": seq_no,
|
|
"role": role,
|
|
"content": content,
|
|
"has_disclaimer": 1 if has_disclaimer else 0,
|
|
},
|
|
)
|