202 lines
9.4 KiB
Python
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]
|