59 lines
2.2 KiB
Python
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]
|