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

326 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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}