142 lines
5.7 KiB
Python
142 lines
5.7 KiB
Python
"""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
|