Files
group_xinghuo_jinrong/app/repository/kyc_session_repository.py
T
zhanghongyu_0626 70aa861983 feat(advisor-agent): Introduce advisor agent functionalities with compliance, KYC, and script templates
- Added new modules for advisor compliance, KYC sessions, and script templates, enhancing the advisor agent's capabilities.
- Implemented a comprehensive API structure under the `/api/advisor-agent` prefix, ensuring clear organization and access to new features.
- Established database models and repositories for compliance rules and KYC sessions, facilitating robust data management.
- Integrated exception handling and response models to improve error management and user feedback.
- Updated settings to include new configurations for compliance and KYC features, ensuring flexibility and adaptability.

This update significantly expands the advisor agent's functionality, providing essential tools for compliance and customer interaction while maintaining a structured API design.
2026-09-12 16:33:07 +08:00

198 lines
6.5 KiB
Python

"""Persistence for KYC sessions and their shared agent-session records."""
from __future__ import annotations
from collections.abc import Callable
from datetime import datetime, timezone
from hashlib import sha256
from sqlalchemy import func
from sqlalchemy.orm import Session
from app.advisor_db import AgentSessionLocal
from app.model.entities import AgentMessage, AgentSession
from app.model.entities_advisor import KycSession
class KycSessionRepository:
def __init__(self, session_factory: Callable[[], Session] | None = None) -> None:
self._override = session_factory
def _session(self) -> Session:
return (self._override or AgentSessionLocal)()
def create(self, kyc_session: KycSession, agent_session: AgentSession) -> KycSession:
session = self._session()
try:
session.add(agent_session)
session.add(kyc_session)
session.commit()
session.refresh(kyc_session)
session.expunge(kyc_session)
return kyc_session
except Exception:
session.rollback()
raise
finally:
session.close()
def get_by_session_id(self, session_id: str) -> KycSession | None:
session = self._session()
try:
row = (
session.query(KycSession)
.filter(KycSession.session_id == session_id)
.one_or_none()
)
if row is not None:
session.expunge(row)
return row
finally:
session.close()
def update_state(
self,
*,
session_id: str,
status: str | None = None,
current_node: str | None = None,
completed_at: datetime | None = None,
duration_seconds: int | None = None,
) -> KycSession | None:
session = self._session()
try:
row = (
session.query(KycSession)
.filter(KycSession.session_id == session_id)
.with_for_update()
.one_or_none()
)
if row is None:
return None
if status is not None:
row.status = status
if current_node is not None:
row.current_node = current_node
if completed_at is not None:
row.completed_at = completed_at
if duration_seconds is not None:
row.duration_seconds = duration_seconds
if status in {"completed", "abandoned"}:
agent_row = (
session.query(AgentSession)
.filter(AgentSession.session_id == session_id)
.one_or_none()
)
if agent_row is not None:
agent_row.status = "closed"
agent_row.closed_at = completed_at or datetime.now(timezone.utc).replace(tzinfo=None)
session.commit()
session.refresh(row)
session.expunge(row)
return row
except Exception:
session.rollback()
raise
finally:
session.close()
def append_chat(
self,
*,
session_id: str,
trace_id: str,
customer_input: str,
parsed_fields: dict,
state_builder: Callable[[KycSession, dict], tuple[dict, list[str], int, str, str]],
) -> KycSession | None:
session = self._session()
try:
row = (
session.query(KycSession)
.filter(KycSession.session_id == session_id)
.with_for_update()
.one_or_none()
)
if row is None:
return None
if row.status != "in_progress":
session.expunge(row)
return row
last_seq = (
session.query(func.max(AgentMessage.seq_no))
.filter(AgentMessage.session_id == session_id)
.scalar()
or 0
)
collected_fields, missing_fields, progress_pct, current_node, assistant_content = (
state_builder(row, parsed_fields)
)
row.collected_fields = collected_fields
row.missing_fields = missing_fields
row.progress_pct = progress_pct
row.current_node = current_node
row.dialog_turns += 1
session.add_all(
[
AgentMessage(
session_id=session_id,
trace_id=trace_id,
seq_no=last_seq + 1,
role="user",
content=customer_input,
content_hash=sha256(customer_input.encode("utf-8")).hexdigest(),
has_disclaimer=0,
),
AgentMessage(
session_id=session_id,
trace_id=trace_id,
seq_no=last_seq + 2,
role="assistant",
content=assistant_content,
content_hash=sha256(assistant_content.encode("utf-8")).hexdigest(),
has_disclaimer=0,
),
]
)
session.commit()
session.refresh(row)
session.expunge(row)
return row
except Exception:
session.rollback()
raise
finally:
session.close()
def expire_stale(self, cutoff: datetime) -> list[KycSession]:
session = self._session()
try:
rows = (
session.query(KycSession)
.filter(KycSession.status == "in_progress", KycSession.updated_at < cutoff)
.all()
)
for row in rows:
row.status = "abandoned"
agent_row = (
session.query(AgentSession)
.filter(AgentSession.session_id == row.session_id)
.one_or_none()
)
if agent_row is not None:
agent_row.status = "closed"
agent_row.closed_at = cutoff
session.commit()
for row in rows:
session.refresh(row)
session.expunge(row)
return rows
except Exception:
session.rollback()
raise
finally:
session.close()