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

202 lines
9.4 KiB
Python

"""SQLAlchemy-backed position + tag-binding repository."""
from __future__ import annotations
import uuid
from datetime import UTC, datetime
from typing import Any
from sqlalchemy import delete, func, select, update
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from deerflow.persistence.positions.base import PositionStore
from deerflow.persistence.positions.model import POSITION_RESOURCE_TYPES, PositionRow, PositionTagBindingRow
from deerflow.persistence.user.model import UserRow
def _position_to_dict(row: PositionRow, member_count: int = 0) -> dict[str, Any]:
data = row.to_dict()
for key in ("created_at", "updated_at"):
value = data.get(key)
if isinstance(value, datetime):
data[key] = value.isoformat()
data["member_count"] = member_count
return data
class PositionRepository(PositionStore):
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
@staticmethod
async def _member_counts(session: AsyncSession, position_ids: list[str]) -> dict[str, int]:
if not position_ids:
return {}
stmt = select(UserRow.position_id, func.count()).where(UserRow.position_id.in_(position_ids)).group_by(UserRow.position_id)
result = await session.execute(stmt)
return {pid: count for pid, count in result.all() if pid is not None}
@staticmethod
async def _clear_default(session: AsyncSession, keep_id: str | None) -> None:
stmt = update(PositionRow).where(PositionRow.is_default.is_(True))
if keep_id is not None:
stmt = stmt.where(PositionRow.id != keep_id)
await session.execute(stmt.values(is_default=False))
# -------------------------------------------------------------- positions
async def list_positions(self) -> list[dict[str, Any]]:
async with self._sf() as session:
rows = list((await session.execute(select(PositionRow).order_by(PositionRow.created_at.asc()))).scalars())
counts = await self._member_counts(session, [r.id for r in rows])
return [_position_to_dict(r, counts.get(r.id, 0)) for r in rows]
async def get_position(self, position_id: str) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(PositionRow, position_id)
if row is None:
return None
counts = await self._member_counts(session, [position_id])
return _position_to_dict(row, counts.get(position_id, 0))
async def create_position(self, data: dict[str, Any]) -> dict[str, Any]:
now = datetime.now(UTC)
is_default = bool(data.get("is_default"))
row = PositionRow(
id=data["id"],
name=data["name"],
description=data.get("description") or "",
is_default=is_default,
created_at=now,
updated_at=now,
)
async with self._sf() as session:
if is_default:
await self._clear_default(session, keep_id=row.id)
session.add(row)
try:
await session.commit()
except IntegrityError as exc:
await session.rollback()
raise ValueError(f"Position '{data['name']}' already exists") from exc
await session.refresh(row)
return _position_to_dict(row, 0)
async def update_position(self, position_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(PositionRow, position_id)
if row is None:
return None
if "name" in data and data["name"] is not None:
row.name = data["name"]
if "description" in data and data["description"] is not None:
row.description = data["description"]
if "is_default" in data and data["is_default"] is not None:
row.is_default = bool(data["is_default"])
if row.is_default:
await self._clear_default(session, keep_id=position_id)
row.updated_at = datetime.now(UTC)
try:
await session.commit()
except IntegrityError as exc:
await session.rollback()
raise ValueError(f"Position '{data.get('name')}' already exists") from exc
await session.refresh(row)
counts = await self._member_counts(session, [position_id])
return _position_to_dict(row, counts.get(position_id, 0))
async def delete_position(self, position_id: str) -> bool:
async with self._sf() as session:
row = await session.get(PositionRow, position_id)
if row is None:
return False
await session.execute(delete(PositionTagBindingRow).where(PositionTagBindingRow.position_id == position_id))
await session.execute(update(UserRow).where(UserRow.position_id == position_id).values(position_id=None))
await session.delete(row)
await session.commit()
return True
async def get_default_position_id(self) -> str | None:
async with self._sf() as session:
return (await session.execute(select(PositionRow.id).where(PositionRow.is_default.is_(True)).limit(1))).scalar_one_or_none()
# ------------------------------------------------------------ tag bindings
async def get_tag_bindings(self, position_id: str) -> dict[str, list[str]]:
async with self._sf() as session:
rows = list(
(
await session.execute(
select(PositionTagBindingRow).where(PositionTagBindingRow.position_id == position_id).order_by(PositionTagBindingRow.created_at.asc())
)
).scalars()
)
out: dict[str, list[str]] = {}
for r in rows:
# Drop bindings for resource types no longer supported (e.g. the
# removed ``user_prompt`` type) so stale rows can't poison the
# editor or be echoed back into a 422 on save.
if r.resource_type not in POSITION_RESOURCE_TYPES:
continue
out.setdefault(r.resource_type, []).append(r.tag_id)
return out
async def set_tag_bindings(self, position_id: str, bindings: dict[str, list[str]]) -> dict[str, list[str]]:
now = datetime.now(UTC)
async with self._sf() as session:
await session.execute(delete(PositionTagBindingRow).where(PositionTagBindingRow.position_id == position_id))
for resource_type, tag_ids in bindings.items():
for tag_id in dict.fromkeys(tag_ids):
session.add(
PositionTagBindingRow(
id=uuid.uuid4().hex,
position_id=position_id,
resource_type=resource_type,
tag_id=tag_id,
created_at=now,
)
)
await session.commit()
return await self.get_tag_bindings(position_id)
# --------------------------------------------------------------- members
async def get_user_position_id(self, user_id: str) -> str | None:
async with self._sf() as session:
return (await session.execute(select(UserRow.position_id).where(UserRow.id == user_id))).scalar_one_or_none()
async def list_members(self, position_id: str) -> list[dict[str, Any]]:
async with self._sf() as session:
rows = (await session.execute(select(UserRow.id, UserRow.email).where(UserRow.position_id == position_id).order_by(UserRow.email.asc()))).all()
return [{"id": uid, "email": email} for uid, email in rows]
async def assign_users(self, position_id: str, user_ids: list[str]) -> int:
if not user_ids:
return 0
async with self._sf() as session:
result = await session.execute(update(UserRow).where(UserRow.id.in_(user_ids)).values(position_id=position_id))
await session.commit()
return result.rowcount or 0
async def unassign_users(self, user_ids: list[str]) -> int:
if not user_ids:
return 0
async with self._sf() as session:
result = await session.execute(update(UserRow).where(UserRow.id.in_(user_ids)).values(position_id=None))
await session.commit()
return result.rowcount or 0
async def list_member_ids(self, position_id: str) -> list[str]:
async with self._sf() as session:
return [uid for (uid,) in (await session.execute(select(UserRow.id).where(UserRow.position_id == position_id))).all()]
async def list_assigned_user_ids(self) -> list[str]:
async with self._sf() as session:
return [uid for (uid,) in (await session.execute(select(UserRow.id).where(UserRow.position_id.isnot(None)))).all()]
async def list_all_users(self) -> list[dict[str, Any]]:
async with self._sf() as session:
rows = (await session.execute(select(UserRow.id, UserRow.email, UserRow.system_role, UserRow.position_id).order_by(UserRow.email.asc()))).all()
return [{"id": uid, "email": email, "system_role": role, "position_id": pid} for uid, email, role, pid in rows]