144 lines
5.8 KiB
Python
144 lines
5.8 KiB
Python
"""SQLAlchemy-backed global business → agent / chain mapping store.
|
||
|
||
Global / shared (no ``user_id`` scoping): every account reads and writes the same
|
||
rows. Fixed business rows are seeded once on startup (``ensure_seed_defaults``)
|
||
with empty agent/chain; users fill those in from the management page via
|
||
``update_mapping``.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from datetime import UTC, datetime
|
||
from typing import Any
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
||
|
||
from deerflow.persistence.business_mapping.model import BusinessMappingRow
|
||
|
||
_UNSET = object()
|
||
|
||
# Fixed business rows (业务码 → 展示名). Seeded on first boot; agent/chain left empty.
|
||
# 8BF 页面尚未建好,先占位保留行。顺序即管理表的展示顺序。
|
||
# AFX = goPath=agentfx 深链(新系统 iframe 嵌入的智能体问答,任务详情取本系统 cop-task-detail)。
|
||
# 6-1…6-4 / 7-1…7-4 = 独立配置位(与 3Q/6BF/7BF 等无关),固定排在表尾。
|
||
DEFAULT_BUSINESS_ROWS: list[dict[str, Any]] = [
|
||
{"business_code": "3Q", "label": "3QFX"},
|
||
{"business_code": "6BF", "label": "任务FX"},
|
||
{"business_code": "7BF", "label": "XDFX"},
|
||
{"business_code": "8BF", "label": "8BF"},
|
||
{"business_code": "AFX", "label": "AGENTFX"},
|
||
# 独立配置位,与上方业务码无关;固定排在表尾。
|
||
{"business_code": "6-1", "label": "6-1"},
|
||
{"business_code": "6-2", "label": "6-2"},
|
||
{"business_code": "6-3", "label": "6-3"},
|
||
{"business_code": "6-4", "label": "6-4"},
|
||
{"business_code": "7-1", "label": "7-1"},
|
||
{"business_code": "7-2", "label": "7-2"},
|
||
{"business_code": "7-3", "label": "7-3"},
|
||
{"business_code": "7-4", "label": "7-4"},
|
||
]
|
||
|
||
|
||
class BusinessMappingRepository:
|
||
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
||
self._sf = session_factory
|
||
|
||
@staticmethod
|
||
def _to_dict(row: BusinessMappingRow) -> dict[str, Any]:
|
||
return {
|
||
"business_code": row.business_code,
|
||
"label": row.label,
|
||
"agent_id": row.agent_id,
|
||
"agent_name": row.agent_name,
|
||
"chain_id": row.chain_id,
|
||
"chain_title": row.chain_title,
|
||
"sort_order": row.sort_order,
|
||
"created_at": row.created_at.isoformat() if isinstance(row.created_at, datetime) else row.created_at,
|
||
"updated_at": row.updated_at.isoformat() if isinstance(row.updated_at, datetime) else row.updated_at,
|
||
}
|
||
|
||
async def list_mappings(self) -> list[dict[str, Any]]:
|
||
stmt = select(BusinessMappingRow).order_by(
|
||
BusinessMappingRow.sort_order.asc(), BusinessMappingRow.business_code.asc()
|
||
)
|
||
async with self._sf() as session:
|
||
result = await session.execute(stmt)
|
||
return [self._to_dict(row) for row in result.scalars()]
|
||
|
||
async def get_mapping(self, business_code: str) -> dict[str, Any] | None:
|
||
async with self._sf() as session:
|
||
row = await session.get(BusinessMappingRow, business_code)
|
||
return self._to_dict(row) if row is not None else None
|
||
|
||
async def update_mapping(
|
||
self,
|
||
business_code: str,
|
||
*,
|
||
agent_id=_UNSET,
|
||
agent_name=_UNSET,
|
||
chain_id=_UNSET,
|
||
chain_title=_UNSET,
|
||
) -> dict[str, Any] | None:
|
||
"""Update an existing (seeded) business row. Returns None if the code is unknown.
|
||
|
||
Only forwarded fields are touched; explicit ``None`` clears that field.
|
||
"""
|
||
async with self._sf() as session:
|
||
row = await session.get(BusinessMappingRow, business_code)
|
||
if row is None:
|
||
return None
|
||
if agent_id is not _UNSET:
|
||
row.agent_id = agent_id or None
|
||
if agent_name is not _UNSET:
|
||
row.agent_name = agent_name or None
|
||
if chain_id is not _UNSET:
|
||
row.chain_id = chain_id or None
|
||
if chain_title is not _UNSET:
|
||
row.chain_title = chain_title or None
|
||
row.updated_at = datetime.now(UTC)
|
||
await session.commit()
|
||
await session.refresh(row)
|
||
return self._to_dict(row)
|
||
|
||
async def ensure_seed_defaults(self) -> int:
|
||
"""Idempotently insert the fixed business rows if missing. Returns inserted count.
|
||
|
||
Never overwrites an existing row's agent/chain (respects user config).
|
||
Re-syncs ``sort_order`` (and empty ``label``) so newly inserted codes keep
|
||
the intended table order next to older seeded rows.
|
||
"""
|
||
inserted = 0
|
||
dirty = False
|
||
now = datetime.now(UTC)
|
||
async with self._sf() as session:
|
||
for idx, seed in enumerate(DEFAULT_BUSINESS_ROWS):
|
||
existing = await session.get(BusinessMappingRow, seed["business_code"])
|
||
if existing is not None:
|
||
if existing.sort_order != idx:
|
||
existing.sort_order = idx
|
||
dirty = True
|
||
seed_label = str(seed.get("label") or "")
|
||
if seed_label and not (existing.label or "").strip():
|
||
existing.label = seed_label
|
||
dirty = True
|
||
continue
|
||
session.add(
|
||
BusinessMappingRow(
|
||
business_code=seed["business_code"],
|
||
label=seed.get("label", ""),
|
||
agent_id=None,
|
||
agent_name=None,
|
||
chain_id=None,
|
||
chain_title=None,
|
||
sort_order=idx,
|
||
created_at=now,
|
||
updated_at=now,
|
||
)
|
||
)
|
||
inserted += 1
|
||
dirty = True
|
||
if dirty:
|
||
await session.commit()
|
||
return inserted
|