"""SQLAlchemy-backed task-button 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.task_buttons.base import TaskButtonStore from deerflow.persistence.task_buttons.model import TaskButtonRow def _to_dict(row: TaskButtonRow) -> dict[str, Any]: data = row.to_dict() value = data.get("updated_at") if isinstance(value, datetime): data["updated_at"] = value.isoformat() # 旧行 / NULL 归一化为默认值。 data["id_param"] = data.get("id_param") or "id" data["id_kind"] = data.get("id_kind") or "task" data["open_mode"] = data.get("open_mode") or "blank" # JSON 列改名 login_params_json → login_params;旧行 / NULL 归一化为 []。 data["login_params"] = data.pop("login_params_json", None) or [] data["append_auth"] = bool(data.get("append_auth", False)) data["auth_token_param"] = data.get("auth_token_param") or "token" data["auth_name_param"] = data.get("auth_name_param") or "username" return data def _row_from(data: dict[str, Any], updated_by: str | None) -> TaskButtonRow: return TaskButtonRow( id=str(data["id"]), business=str(data.get("business") or ""), label=str(data.get("label") or ""), link_type=str(data.get("link_type") or "url"), target=str(data.get("target") or ""), append_task_id=bool(data.get("append_task_id", True)), id_param=str(data.get("id_param") or "id"), id_kind=str(data.get("id_kind") or "task"), open_mode=str(data.get("open_mode") or "blank"), login_params_json=list(data.get("login_params") or []), append_auth=bool(data.get("append_auth", False)), auth_token_param=str(data.get("auth_token_param") or "token"), auth_name_param=str(data.get("auth_name_param") or "username"), enabled=bool(data.get("enabled", True)), sort_order=int(data.get("sort_order", 0)), updated_by=updated_by, ) class TaskButtonRepository(TaskButtonStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def list_buttons(self) -> list[dict[str, Any]]: stmt = select(TaskButtonRow).order_by(TaskButtonRow.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 id (last wins) so a malformed payload can't violate the PK. by_id: dict[str, dict[str, Any]] = {} for r in rows: row_id = r.get("id") if row_id: by_id[str(row_id)] = r async with self._sf() as session: await session.execute(delete(TaskButtonRow)) 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]