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

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