"""SQLAlchemy repository for generated HTML pages and favorites.""" from __future__ import annotations from datetime import datetime from typing import Any from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.html_pages.base import HtmlPageStore from deerflow.persistence.html_pages.model import HtmlPageFavoriteRow, ScheduledHtmlPageRow def _to_dict(row) -> dict[str, Any]: data = row.to_dict() for key in ("created_at", "updated_at"): value = data.get(key) if isinstance(value, datetime): data[key] = value.isoformat() # Normalize the JSON column names to API-friendly keys. if "data_snapshot_json" in data: data["data_snapshot"] = data.pop("data_snapshot_json") or {} if "prompt_snapshot_json" in data: data["prompt_snapshot"] = data.pop("prompt_snapshot_json") or {} if "tags_json" in data: data["tags"] = data.pop("tags_json") or [] return data class HtmlPageRepository(HtmlPageStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory # --- Generated pages --- async def create_page(self, data: dict[str, Any]) -> dict[str, Any]: row = ScheduledHtmlPageRow( id=data["id"], task_id=data["task_id"], run_id=data["run_id"], user_id=data["user_id"], title=data.get("title") or "", html_content=data.get("html_content") or "", html_file_path=data.get("html_file_path"), data_snapshot_json=data.get("data_snapshot") or {}, prompt_snapshot_json=data.get("prompt_snapshot") or {}, reference_favorite_id=data.get("reference_favorite_id"), style_summary=data.get("style_summary"), layout_summary=data.get("layout_summary"), ) async with self._sf() as session: session.add(row) await session.commit() await session.refresh(row) return _to_dict(row) async def get_page(self, page_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(ScheduledHtmlPageRow, page_id) return _to_dict(row) if row is not None else None async def get_page_by_run(self, run_id: str) -> dict[str, Any] | None: async with self._sf() as session: result = await session.execute( select(ScheduledHtmlPageRow) .where(ScheduledHtmlPageRow.run_id == run_id) .order_by(ScheduledHtmlPageRow.created_at.desc()) .limit(1) ) row = result.scalar_one_or_none() return _to_dict(row) if row is not None else None async def latest_page_for_task(self, task_id: str) -> dict[str, Any] | None: async with self._sf() as session: result = await session.execute( select(ScheduledHtmlPageRow) .where(ScheduledHtmlPageRow.task_id == task_id) .order_by(ScheduledHtmlPageRow.created_at.desc()) .limit(1) ) row = result.scalar_one_or_none() return _to_dict(row) if row is not None else None # --- Favorites --- async def create_favorite(self, data: dict[str, Any]) -> dict[str, Any]: row = HtmlPageFavoriteRow( id=data["id"], user_id=data["user_id"], page_id=data["page_id"], title=data.get("title") or "", description=data.get("description"), tags_json=data.get("tags") or [], html_snapshot=data.get("html_snapshot") or "", style_summary=data.get("style_summary"), layout_summary=data.get("layout_summary"), content_structure_summary=data.get("content_structure_summary"), reuse_prompt=data.get("reuse_prompt"), ) async with self._sf() as session: session.add(row) await session.commit() await session.refresh(row) return _to_dict(row) async def list_favorites(self, user_id: str, *, limit: int = 100) -> list[dict[str, Any]]: async with self._sf() as session: result = await session.execute( select(HtmlPageFavoriteRow) .where(HtmlPageFavoriteRow.user_id == user_id) .order_by(HtmlPageFavoriteRow.created_at.desc()) .limit(max(1, limit)) ) return [_to_dict(row) for row in result.scalars()] async def get_favorite(self, favorite_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(HtmlPageFavoriteRow, favorite_id) return _to_dict(row) if row is not None else None async def get_favorite_by_page(self, user_id: str, page_id: str) -> dict[str, Any] | None: async with self._sf() as session: result = await session.execute( select(HtmlPageFavoriteRow) .where(HtmlPageFavoriteRow.user_id == user_id, HtmlPageFavoriteRow.page_id == page_id) .limit(1) ) row = result.scalar_one_or_none() return _to_dict(row) if row is not None else None async def delete_favorite(self, favorite_id: str, user_id: str) -> bool: async with self._sf() as session: result = await session.execute( delete(HtmlPageFavoriteRow).where( HtmlPageFavoriteRow.id == favorite_id, HtmlPageFavoriteRow.user_id == user_id, ) ) await session.commit() return (result.rowcount or 0) > 0