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

78 lines
3.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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]