78 lines
3.3 KiB
Python
78 lines
3.3 KiB
Python
"""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]
|