"""In-memory custom agent store for memory persistence mode.""" from __future__ import annotations from datetime import UTC, datetime from typing import Any from deerflow.config.agent_name_desensitize import name_matches from deerflow.persistence.agents.base import AgentStore def _now() -> str: return datetime.now(UTC).isoformat() def _copy(data: dict[str, Any]) -> dict[str, Any]: copied = dict(data) copied["builtin"] = copied.get("user_id") is None return copied class MemoryAgentStore(AgentStore): def __init__(self) -> None: self._agents: dict[str, dict[str, Any]] = {} self._favorites: set[tuple[str, str]] = set() self._pins: set[tuple[str, str]] = set() # Listing cache mirrors (see AgentExtrasRow / AgentSkillRow): self._skills: dict[str, list[str]] = {} self._extras: dict[str, dict[str, Any]] = {} @staticmethod def _is_visible(agent: dict[str, Any], user_id: str) -> bool: return agent.get("user_id") is None or agent.get("user_id") == user_id or bool(agent.get("published")) def _is_fetchable(self, agent: dict[str, Any], user_id: str) -> bool: """Single-record visibility (``get_visible``) — the union of every listing surface, so a list-visible agent always opens. Mirrors ``AgentRepository._visible_filter``: baseline visibility OR private-square (password is the discovery gate, reads are not re-gated) OR favorited (the "我的" listing keeps favorites even after the owner unpublishes).""" return self._is_visible(agent, user_id) or bool(agent.get("is_private")) or (user_id, agent["id"]) in self._favorites async def list_visible(self, user_id: str) -> list[dict[str, Any]]: rows = [_copy(agent) for agent in self._agents.values() if self._is_visible(agent, user_id)] rows.sort(key=lambda agent: agent.get("created_at") or "") return rows async def get_visible(self, agent_id: str, user_id: str) -> dict[str, Any] | None: agent = self._agents.get(agent_id) if agent is None or not self._is_fetchable(agent, user_id): return None return _copy(agent) async def get_any(self, agent_id: str) -> dict[str, Any] | None: agent = self._agents.get(agent_id) return _copy(agent) if agent is not None else None async def get_owned(self, agent_id: str, user_id: str) -> dict[str, Any] | None: agent = self._agents.get(agent_id) if agent is None or agent.get("user_id") != user_id: return None return _copy(agent) async def create(self, data: dict[str, Any]) -> dict[str, Any]: now = _now() row = { "id": data["id"], "name": data["name"], "description": data.get("description") or "", "user_id": data.get("user_id"), "published": bool(data.get("published", False)), "is_private": bool(data.get("is_private", False)), "square_id": str(data.get("square_id") or ""), "featured_order": None, "created_at": now, "updated_at": now, } self._agents[row["id"]] = row return _copy(row) async def ensure_builtin(self, data: dict[str, Any]) -> dict[str, Any]: row = self._agents.get(data["id"]) if row is None: row = { "id": data["id"], "name": data["name"], "description": data.get("description") or "", "user_id": None, "published": bool(data.get("published", True)), "featured_order": None, "created_at": _now(), "updated_at": _now(), } self._agents[row["id"]] = row elif row.get("user_id") is None: row["name"] = data["name"] row["description"] = data.get("description") or "" row["published"] = bool(data.get("published", True)) row["updated_at"] = _now() return _copy(row) async def update(self, agent_id: str, user_id: str, data: dict[str, Any]) -> dict[str, Any] | None: row = self._agents.get(agent_id) if row is None or row.get("user_id") != user_id: return None for key in ("name", "description", "published", "is_private", "square_id"): if key in data: row[key] = data[key] row["updated_at"] = _now() return _copy(row) async def update_builtin(self, agent_id: str, data: dict[str, Any]) -> dict[str, Any] | None: row = self._agents.get(agent_id) if row is None or row.get("user_id") is not None: return None for key in ("name", "description", "published", "is_private", "square_id"): if key in data: row[key] = data[key] row["updated_at"] = _now() return _copy(row) async def delete(self, agent_id: str, user_id: str) -> bool: row = self._agents.get(agent_id) if row is None or row.get("user_id") != user_id: return False self._agents.pop(agent_id, None) return True async def delete_builtin(self, agent_id: str) -> bool: row = self._agents.get(agent_id) if row is None or row.get("user_id") is not None: return False self._agents.pop(agent_id, None) return True async def count_owned(self, user_id: str) -> int: return sum(1 for a in self._agents.values() if a.get("user_id") == user_id) async def set_featured(self, agent_id: str, order: int | None) -> dict[str, Any] | None: row = self._agents.get(agent_id) if row is None: return None row["featured_order"] = order row["updated_at"] = _now() return _copy(row) async def list_featured(self, user_id: str) -> list[dict[str, Any]]: candidates = [a for a in self._agents.values() if a.get("featured_order") is not None and self._is_visible(a, user_id)] candidates.sort(key=lambda a: (a.get("featured_order"), a.get("created_at") or "")) return [_copy(a) for a in candidates] async def list_by_ids(self, agent_ids: list[str], user_id: str) -> list[dict[str, Any]]: order_map = {aid: idx for idx, aid in enumerate(agent_ids)} rows = [self._agents[aid] for aid in agent_ids if aid in self._agents and self._is_visible(self._agents[aid], user_id)] rows.sort(key=lambda a: order_map.get(a["id"], 0)) return [_copy(a) for a in rows] async def list_private(self) -> list[dict[str, Any]]: rows = [_copy(a) for a in self._agents.values() if bool(a.get("is_private"))] rows.sort(key=lambda a: a.get("created_at") or "") return rows async def list_all(self) -> list[dict[str, Any]]: rows = [_copy(a) for a in self._agents.values()] rows.sort(key=lambda a: a.get("created_at") or "") return rows async def update_any(self, agent_id: str, data: dict[str, Any]) -> dict[str, Any] | None: row = self._agents.get(agent_id) if row is None: return None for key in ("name", "description", "published", "is_private", "square_id"): if key in data: row[key] = data[key] row["updated_at"] = _now() return _copy(row) async def delete_any(self, agent_id: str) -> bool: if agent_id not in self._agents: return False self._agents.pop(agent_id, None) return True # --- Paginated listing --- def _matches_scope(self, agent: dict[str, Any], *, user_id: str, scope: str | None, admin: bool) -> bool: if admin: return True if scope == "square": return (agent.get("user_id") is None or bool(agent.get("published"))) and not bool(agent.get("is_private")) if scope == "mine": owned = agent.get("user_id") is not None and agent.get("user_id") == user_id favorited = (user_id, agent["id"]) in self._favorites return owned or favorited if scope == "private": return bool(agent.get("is_private")) if scope == "builtin": # 内置智能体 = user_id IS NULL(仅管理员;路由层 admin 闸门)。 return agent.get("user_id") is None return self._is_visible(agent, user_id) async def list_paginated( self, *, user_id: str, scope: str | None, search: str | None, tag_ids: list[str] | None, page: int, page_size: int, admin: bool = False, sort: str = "time_desc", owner: str | None = None, square_id: str | None = None, square_is_default: bool = False, ) -> tuple[list[dict[str, Any]], int]: # NOTE: the in-memory store has no tag-assignment data, so ``tag_ids`` # is ignored here. Tag-filter behavior is covered by the SQL repo tests. owner_key = owner.strip() if owner and owner.strip() else None square_key = square_id.strip() if square_id and square_id.strip() else None def _square_ok(a: dict[str, Any]) -> bool: if square_key is None: return True sid = str(a.get("square_id") or "") return sid == square_key or (square_is_default and sid == "") matched = [ a for a in self._agents.values() if self._matches_scope(a, user_id=user_id, scope=scope, admin=admin) and (not search or not search.strip() or name_matches(a.get("name") or "", search)) and (owner_key is None or a.get("user_id") == owner_key) and _square_ok(a) ] if sort == "time_asc": matched.sort(key=lambda a: (a.get("created_at") or "", a.get("id") or "")) elif sort == "updated_desc": matched.sort( key=lambda a: (a.get("updated_at") or "", a.get("id") or ""), reverse=True, ) elif sort == "updated_asc": matched.sort(key=lambda a: (a.get("updated_at") or "", a.get("id") or "")) elif sort == "name_asc": matched.sort(key=lambda a: ((a.get("name") or "").lower(), a.get("id") or "")) elif sort == "name_desc": matched.sort( key=lambda a: ((a.get("name") or "").lower(), a.get("id") or ""), reverse=True, ) else: matched.sort( key=lambda a: (a.get("created_at") or "", a.get("id") or ""), reverse=True, ) pin_ids = {agent_id for uid, agent_id in self._pins if uid == user_id} matched.sort(key=lambda a: (0 if a.get("id") in pin_ids else 1,)) total = len(matched) start = max(page - 1, 0) * page_size page_rows = matched[start : start + page_size] return [_copy(a) for a in page_rows], total # --- Listing cache (skills / tool_groups / model) --- async def replace_skills(self, agent_id: str, skills: list[str] | None) -> None: self._skills[agent_id] = [s for s in (skills or []) if isinstance(s, str)] async def set_extras(self, agent_id: str, *, tool_groups: list[str] | None, model: str | None) -> None: self._extras[agent_id] = {"tool_groups": tool_groups, "model": model} async def get_extras_for(self, agent_ids: list[str]) -> dict[str, dict[str, Any]]: out: dict[str, dict[str, Any]] = {} for agent_id in agent_ids: has_extras = agent_id in self._extras has_skills = agent_id in self._skills if not has_extras and not has_skills: continue extras = self._extras.get(agent_id, {}) out[agent_id] = { "skills": list(self._skills.get(agent_id, [])), "tool_groups": extras.get("tool_groups"), "model": extras.get("model"), } return out async def delete_extras_for_agent(self, agent_id: str) -> None: self._skills.pop(agent_id, None) self._extras.pop(agent_id, None) # --- Favorites ("收藏") --- async def add_favorite(self, user_id: str, agent_id: str) -> None: self._favorites.add((user_id, agent_id)) async def remove_favorite(self, user_id: str, agent_id: str) -> bool: if (user_id, agent_id) in self._favorites: self._favorites.discard((user_id, agent_id)) return True return False async def list_favorite_ids(self, user_id: str) -> set[str]: return {agent_id for (uid, agent_id) in self._favorites if uid == user_id} async def delete_favorites_for_agent(self, agent_id: str) -> None: self._favorites = {(uid, aid) for (uid, aid) in self._favorites if aid != agent_id} # --- Pins ("置顶") --- async def add_pin(self, user_id: str, agent_id: str) -> None: self._pins.add((user_id, agent_id)) async def remove_pin(self, user_id: str, agent_id: str) -> bool: key = (user_id, agent_id) if key not in self._pins: return False self._pins.remove(key) return True async def list_pin_ids(self, user_id: str) -> set[str]: return {agent_id for uid, agent_id in self._pins if uid == user_id} async def delete_pins_for_agent(self, agent_id: str) -> None: self._pins = {(uid, aid) for (uid, aid) in self._pins if aid != agent_id}