326 lines
13 KiB
Python
326 lines
13 KiB
Python
"""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}
|