"""In-memory ``PositionStore`` implementation for tests.""" from __future__ import annotations import copy from datetime import UTC, datetime from typing import Any def _now() -> str: return datetime.now(UTC).isoformat() class MemoryPositionStore: def __init__(self) -> None: self._positions: dict[str, dict[str, Any]] = {} # position_id -> {resource_type -> [tag_id]} self._bindings: dict[str, dict[str, list[str]]] = {} # user_id -> position_id self._user_position: dict[str, str] = {} # user_id -> email (optional; populated via register_user for nicer member lists) self._user_email: dict[str, str] = {} # test helper — not part of the abstract API def register_user(self, user_id: str, email: str = "") -> None: self._user_email[user_id] = email def _member_count(self, position_id: str) -> int: return sum(1 for pid in self._user_position.values() if pid == position_id) def _to_dict(self, record: dict[str, Any]) -> dict[str, Any]: out = copy.deepcopy(record) out["member_count"] = self._member_count(record["id"]) return out # -------------------------------------------------------------- positions async def list_positions(self) -> list[dict[str, Any]]: rows = sorted(self._positions.values(), key=lambda r: r["created_at"]) return [self._to_dict(r) for r in rows] async def get_position(self, position_id: str) -> dict[str, Any] | None: record = self._positions.get(position_id) return self._to_dict(record) if record is not None else None def _clear_default(self, keep_id: str | None) -> None: for r in self._positions.values(): if r["id"] != keep_id: r["is_default"] = False async def create_position(self, data: dict[str, Any]) -> dict[str, Any]: name = data["name"] if any(r["name"] == name for r in self._positions.values()): raise ValueError(f"Position '{name}' already exists") is_default = bool(data.get("is_default")) if is_default: self._clear_default(keep_id=data["id"]) record = { "id": data["id"], "name": name, "description": data.get("description") or "", "is_default": is_default, "created_at": _now(), "updated_at": _now(), } self._positions[data["id"]] = record return self._to_dict(record) async def update_position(self, position_id: str, data: dict[str, Any]) -> dict[str, Any] | None: record = self._positions.get(position_id) if record is None: return None new_name = data.get("name", record["name"]) if new_name is not None and any(r["id"] != position_id and r["name"] == new_name for r in self._positions.values()): raise ValueError(f"Position '{new_name}' already exists") if "name" in data and data["name"] is not None: record["name"] = data["name"] if "description" in data and data["description"] is not None: record["description"] = data["description"] if "is_default" in data and data["is_default"] is not None: record["is_default"] = bool(data["is_default"]) if record["is_default"]: self._clear_default(keep_id=position_id) record["updated_at"] = _now() return self._to_dict(record) async def delete_position(self, position_id: str) -> bool: if position_id not in self._positions: return False del self._positions[position_id] self._bindings.pop(position_id, None) for uid in [uid for uid, pid in self._user_position.items() if pid == position_id]: del self._user_position[uid] return True async def get_default_position_id(self) -> str | None: for r in self._positions.values(): if r["is_default"]: return r["id"] return None # ------------------------------------------------------------ tag bindings async def get_tag_bindings(self, position_id: str) -> dict[str, list[str]]: return copy.deepcopy(self._bindings.get(position_id, {})) async def set_tag_bindings(self, position_id: str, bindings: dict[str, list[str]]) -> dict[str, list[str]]: cleaned = {rt: list(dict.fromkeys(tag_ids)) for rt, tag_ids in bindings.items() if tag_ids} self._bindings[position_id] = cleaned return copy.deepcopy(cleaned) # --------------------------------------------------------------- members async def get_user_position_id(self, user_id: str) -> str | None: return self._user_position.get(user_id) async def list_members(self, position_id: str) -> list[dict[str, Any]]: return [{"id": uid, "email": self._user_email.get(uid, "")} for uid, pid in self._user_position.items() if pid == position_id] async def assign_users(self, position_id: str, user_ids: list[str]) -> int: for uid in user_ids: self._user_position[uid] = position_id return len(user_ids) async def unassign_users(self, user_ids: list[str]) -> int: count = 0 for uid in user_ids: if uid in self._user_position: del self._user_position[uid] count += 1 return count async def list_member_ids(self, position_id: str) -> list[str]: return [uid for uid, pid in self._user_position.items() if pid == position_id] async def list_assigned_user_ids(self) -> list[str]: return list(self._user_position.keys()) async def list_all_users(self) -> list[dict[str, Any]]: return [ {"id": uid, "email": email, "system_role": "user", "position_id": self._user_position.get(uid)} for uid, email in sorted(self._user_email.items(), key=lambda kv: kv[1]) ]