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