"""SQLAlchemy-backed fixed-question store.""" from __future__ import annotations import hashlib from datetime import datetime from typing import Any from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.fixed_questions.base import FixedQuestionStore from deerflow.persistence.fixed_questions.model import FixedQuestionRow def question_hash(question: str) -> str: return hashlib.sha256(question.encode("utf-8")).hexdigest() def _to_dict(row: FixedQuestionRow) -> dict[str, Any]: data = row.to_dict() data.pop("question_hash", None) value = data.get("updated_at") if isinstance(value, datetime): data["updated_at"] = value.isoformat() return data def _row_from(data: dict[str, Any], updated_by: str | None) -> FixedQuestionRow: question = str(data.get("question") or "") return FixedQuestionRow( id=str(data["id"]), question=question, question_hash=question_hash(question), answer=str(data.get("answer") or ""), enabled=bool(data.get("enabled", True)), tokens_per_second=int(data.get("tokens_per_second", 100)), sort_order=int(data.get("sort_order", 0)), updated_by=updated_by, ) class FixedQuestionRepository(FixedQuestionStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def list_questions(self) -> list[dict[str, Any]]: stmt = select(FixedQuestionRow).order_by( FixedQuestionRow.sort_order.asc(), FixedQuestionRow.question.asc(), ) async with self._sf() as session: result = await session.execute(stmt) return [_to_dict(row) for row in result.scalars()] async def find_enabled(self, question: str) -> dict[str, Any] | None: stmt = select(FixedQuestionRow).where( FixedQuestionRow.question_hash == question_hash(question), FixedQuestionRow.enabled.is_(True), ) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() if row is None or row.question != question: return None return _to_dict(row) async def replace_all( self, rows: list[dict[str, Any]], *, updated_by: str | None = None, ) -> list[dict[str, Any]]: # Last duplicate question wins, preventing a unique-hash violation even # if a non-UI caller submits malformed data. by_id: dict[str, dict[str, Any]] = {} for row in rows: question = row.get("question") row_id = row.get("id") if row_id and question not in (None, ""): by_id[str(row_id)] = row by_question = {str(row["question"]): row for row in by_id.values()} async with self._sf() as session: await session.execute(delete(FixedQuestionRow)) new_rows = [_row_from(row, updated_by) for row in by_question.values()] session.add_all(new_rows) await session.commit() return [_to_dict(row) for row in new_rows]