89 lines
3.1 KiB
Python
89 lines
3.1 KiB
Python
"""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]
|