deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/menu_overrides/sql.py
2026-09-07 18:24:55 +08:00

59 lines
2.2 KiB
Python

"""SQLAlchemy-backed sidebar-menu override store (global / shared)."""
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.menu_overrides.base import MenuOverrideStore
from deerflow.persistence.menu_overrides.model import MenuOverrideRow
def _to_dict(row: MenuOverrideRow) -> dict[str, Any]:
data = row.to_dict()
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) -> MenuOverrideRow:
return MenuOverrideRow(
node_id=str(data["node_id"]),
parent_id=(str(data["parent_id"]) if data.get("parent_id") not in (None, "") else None),
sort_order=int(data.get("sort_order", 0)),
label=(data.get("label") or None),
disabled=bool(data.get("disabled", False)),
updated_by=updated_by,
)
class MenuOverrideRepository(MenuOverrideStore):
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
async def list_overrides(self) -> list[dict[str, Any]]:
stmt = select(MenuOverrideRow).order_by(MenuOverrideRow.sort_order.asc())
async with self._sf() as session:
result = await session.execute(stmt)
return [_to_dict(row) for row in result.scalars()]
async def replace_all(self, rows: list[dict[str, Any]], *, updated_by: str | None = None) -> list[dict[str, Any]]:
# De-dupe by node_id (last wins) so a malformed payload can't violate the PK.
by_id: dict[str, dict[str, Any]] = {}
for r in rows:
node_id = r.get("node_id")
if node_id:
by_id[str(node_id)] = r
async with self._sf() as session:
await session.execute(delete(MenuOverrideRow))
new_rows = [_row_from(r, updated_by) for r in by_id.values()]
session.add_all(new_rows)
await session.commit()
for row in new_rows:
await session.refresh(row)
return [_to_dict(row) for row in new_rows]