144 lines
5.8 KiB
Python
144 lines
5.8 KiB
Python
"""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])
|
|
]
|