- 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.
198 lines
6.5 KiB
Python
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()
|