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

574 lines
23 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.

"""SQLAlchemy-backed custom agent repository."""
from __future__ import annotations
from datetime import UTC, datetime
from typing import Any
from uuid import uuid4
from sqlalchemy import and_, case, delete, func, literal, not_, or_, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from deerflow.config.agent_name_desensitize import source_terms_for_needle
from deerflow.persistence.agents.base import AgentStore
from deerflow.persistence.agents.model import (
AgentExtrasRow,
AgentFavoriteRow,
AgentPinRow,
AgentRow,
AgentSkillRow,
)
from deerflow.persistence.roundtable_chains.model import RoundtableChainRow
from deerflow.persistence.tags.model import TagAssignmentRow
def _row_to_dict(row: AgentRow) -> dict[str, Any]:
data = row.to_dict()
data["builtin"] = data.get("user_id") is None
for key in ("created_at", "updated_at"):
value = data.get(key)
if isinstance(value, datetime):
data[key] = value.isoformat()
return data
class AgentRepository(AgentStore):
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
@staticmethod
def _business_chain_agent_exists(agent_id_col):
return (
select(RoundtableChainRow.id)
.where(
or_(
RoundtableChainRow.seats.like(literal('%"agent_id": "') + agent_id_col + literal('"%')),
RoundtableChainRow.seats.like(literal('%"agent_id":"') + agent_id_col + literal('"%')),
)
)
.exists()
)
@classmethod
def _valid_favorite_filter(cls):
return or_(
AgentFavoriteRow.origin == "user",
and_(
AgentFavoriteRow.origin == "position",
AgentRow.user_id.is_not(None),
not_(cls._business_chain_agent_exists(AgentFavoriteRow.agent_id)),
),
)
@staticmethod
def _visible_filter(user_id: str):
"""Visibility for single-record fetches (``get_visible``).
Lenient — must cover the union of every listing surface, so any agent
a user can see in a list always opens (detail + chat) instead of
404ing:
- built-in / owned / published — the baseline gallery visibility;
- ``is_private`` — the private-square password is the *discovery*
gate (listings); individual reads are not re-gated, so a private
agent opened from the square (or a bookmark/link) stays openable.
Note moving an agent into the private square forces
``published=false``, so this clause cannot lean on ``published``;
- favorited — the "我的" listing keeps favorited agents even after
the owner unpublishes them, so reads keep honoring the favorite.
"""
valid_favorite_exists = (
select(AgentFavoriteRow.id)
.where(
AgentFavoriteRow.user_id == user_id,
AgentFavoriteRow.agent_id == AgentRow.id,
AgentRepository._valid_favorite_filter(),
)
.exists()
)
return or_(
AgentRow.user_id.is_(None),
AgentRow.user_id == user_id,
AgentRow.published.is_(True),
AgentRow.is_private.is_(True),
valid_favorite_exists,
)
@staticmethod
def _listing_filter(user_id: str):
"""Visibility for *listings* — hides private-square agents from
non-owners so they only surface inside the password-gated endpoint.
Owners keep seeing their own private agents in '我的'.
"""
return or_(
AgentRow.user_id == user_id,
and_(
or_(AgentRow.user_id.is_(None), AgentRow.published.is_(True)),
AgentRow.is_private.is_(False),
),
)
async def list_visible(self, user_id: str) -> list[dict[str, Any]]:
stmt = select(AgentRow).where(self._listing_filter(user_id)).order_by(AgentRow.created_at.asc())
async with self._sf() as session:
result = await session.execute(stmt)
return [_row_to_dict(row) for row in result.scalars()]
async def get_visible(self, agent_id: str, user_id: str) -> dict[str, Any] | None:
stmt = select(AgentRow).where(AgentRow.id == agent_id, self._visible_filter(user_id))
async with self._sf() as session:
row = (await session.execute(stmt)).scalar_one_or_none()
return _row_to_dict(row) if row is not None else None
async def get_any(self, agent_id: str) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
return _row_to_dict(row) if row is not None else None
async def get_owned(self, agent_id: str, user_id: str) -> dict[str, Any] | None:
stmt = select(AgentRow).where(AgentRow.id == agent_id, AgentRow.user_id == user_id)
async with self._sf() as session:
row = (await session.execute(stmt)).scalar_one_or_none()
return _row_to_dict(row) if row is not None else None
async def create(self, data: dict[str, Any]) -> dict[str, Any]:
now = datetime.now(UTC)
row = AgentRow(
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 ""),
created_at=now,
updated_at=now,
)
async with self._sf() as session:
session.add(row)
await session.commit()
await session.refresh(row)
return _row_to_dict(row)
async def ensure_builtin(self, data: dict[str, Any]) -> dict[str, Any]:
async with self._sf() as session:
row = await session.get(AgentRow, data["id"])
if row is None:
now = datetime.now(UTC)
row = AgentRow(
id=data["id"],
name=data["name"],
description=data.get("description") or "",
user_id=None,
published=bool(data.get("published", True)),
created_at=now,
updated_at=now,
)
session.add(row)
elif row.user_id is None:
row.name = data["name"]
row.description = data.get("description") or ""
row.published = bool(data.get("published", True))
row.updated_at = datetime.now(UTC)
await session.commit()
await session.refresh(row)
return _row_to_dict(row)
async def update(self, agent_id: str, user_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None or row.user_id != user_id:
return None
for key in ("name", "description", "published", "is_private", "square_id"):
if key in data:
setattr(row, key, data[key])
row.updated_at = datetime.now(UTC)
await session.commit()
await session.refresh(row)
return _row_to_dict(row)
async def update_builtin(self, agent_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None or row.user_id is not None:
return None
for key in ("name", "description", "published", "is_private", "square_id"):
if key in data:
setattr(row, key, data[key])
row.updated_at = datetime.now(UTC)
await session.commit()
await session.refresh(row)
return _row_to_dict(row)
async def delete(self, agent_id: str, user_id: str) -> bool:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None or row.user_id != user_id:
return False
await session.delete(row)
await session.commit()
return True
async def delete_builtin(self, agent_id: str) -> bool:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None or row.user_id is not None:
return False
await session.delete(row)
await session.commit()
return True
async def count_owned(self, user_id: str) -> int:
stmt = select(func.count()).select_from(AgentRow).where(AgentRow.user_id == user_id)
async with self._sf() as session:
return await session.scalar(stmt) or 0
async def set_featured(self, agent_id: str, order: int | None) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None:
return None
row.featured_order = order
row.updated_at = datetime.now(UTC)
await session.commit()
await session.refresh(row)
return _row_to_dict(row)
async def list_featured(self, user_id: str) -> list[dict[str, Any]]:
stmt = (
select(AgentRow)
.where(AgentRow.featured_order.is_not(None), self._listing_filter(user_id))
.order_by(AgentRow.featured_order.asc(), AgentRow.created_at.asc())
)
async with self._sf() as session:
result = await session.execute(stmt)
return [_row_to_dict(row) for row in result.scalars()]
async def list_by_ids(self, agent_ids: list[str], user_id: str) -> list[dict[str, Any]]:
if not agent_ids:
return []
order_map = {aid: idx for idx, aid in enumerate(agent_ids)}
stmt = select(AgentRow).where(AgentRow.id.in_(agent_ids), self._listing_filter(user_id))
async with self._sf() as session:
result = await session.execute(stmt)
rows = list(result.scalars())
rows.sort(key=lambda r: order_map.get(r.id, len(order_map)))
return [_row_to_dict(r) for r in rows]
async def list_private(self) -> list[dict[str, Any]]:
stmt = select(AgentRow).where(AgentRow.is_private.is_(True)).order_by(AgentRow.created_at.asc())
async with self._sf() as session:
result = await session.execute(stmt)
return [_row_to_dict(row) for row in result.scalars()]
async def list_all(self) -> list[dict[str, Any]]:
stmt = select(AgentRow).order_by(AgentRow.created_at.asc())
async with self._sf() as session:
result = await session.execute(stmt)
return [_row_to_dict(row) for row in result.scalars()]
async def update_any(self, agent_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None:
return None
for key in ("name", "description", "published", "is_private", "square_id"):
if key in data:
setattr(row, key, data[key])
row.updated_at = datetime.now(UTC)
await session.commit()
await session.refresh(row)
return _row_to_dict(row)
async def delete_any(self, agent_id: str) -> bool:
async with self._sf() as session:
row = await session.get(AgentRow, agent_id)
if row is None:
return False
await session.delete(row)
await session.commit()
return True
# --- Paginated listing ---
def _scope_filter(self, *, user_id: str, scope: str | None, admin: bool):
"""Build the visibility predicate for a paginated listing."""
if admin:
return None # no visibility filter — every agent
if scope == "square":
return and_(
or_(AgentRow.user_id.is_(None), AgentRow.published.is_(True)),
AgentRow.is_private.is_(False),
)
if scope == "mine":
# owned OR valid favorites. Built-in agents and business-chain
# seats only enter "mine" when explicitly favorited by the user;
# the original position-granted behavior stays for ordinary agents.
valid_favorite_exists = (
select(AgentFavoriteRow.id)
.where(
AgentFavoriteRow.user_id == user_id,
AgentFavoriteRow.agent_id == AgentRow.id,
self._valid_favorite_filter(),
)
.exists()
)
return or_(AgentRow.user_id == user_id, valid_favorite_exists)
if scope == "private":
return AgentRow.is_private.is_(True)
if scope == "builtin":
# 内置智能体 = user_id IS NULL。仅管理员可用此 scope(路由层已 admin 闸门)。
return AgentRow.user_id.is_(None)
return self._listing_filter(user_id)
@staticmethod
def _search_clause(search: str):
"""Match the raw name OR (via the desensitize mirror) its display form."""
needle = search.strip().lower()
clauses = [func.lower(AgentRow.name).like(f"%{needle}%")]
for term in source_terms_for_needle(needle):
clauses.append(AgentRow.name.like(f"%{term}%"))
return or_(*clauses)
@staticmethod
def _tag_clause(tag_ids: list[str]):
subq = select(TagAssignmentRow.target_id).where(
TagAssignmentRow.target_type == "agent",
TagAssignmentRow.tag_id.in_(tag_ids),
)
return AgentRow.id.in_(subq)
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]:
filters = []
scope_filter = self._scope_filter(user_id=user_id, scope=scope, admin=admin)
if scope_filter is not None:
filters.append(scope_filter)
if search and search.strip():
filters.append(self._search_clause(search))
if tag_ids:
filters.append(self._tag_clause(tag_ids))
if owner and owner.strip():
# Filter to one publisher. Built-in agents (user_id IS NULL) never match.
filters.append(AgentRow.user_id == owner.strip())
if square_id and square_id.strip():
sid = square_id.strip()
if square_is_default:
# Legacy / unassigned rows (square_id == "") belong to the
# default square, so include them alongside explicit matches.
filters.append(or_(AgentRow.square_id == sid, AgentRow.square_id == ""))
else:
filters.append(AgentRow.square_id == sid)
count_stmt = select(func.count()).select_from(AgentRow)
pin_subq = select(AgentPinRow.agent_id).where(AgentPinRow.user_id == user_id)
pin_rank = case((AgentRow.id.in_(pin_subq), 0), else_=1)
if sort == "time_asc":
order_by = (AgentRow.created_at.asc(), AgentRow.id.asc())
elif sort == "updated_desc":
order_by = (AgentRow.updated_at.desc(), AgentRow.id.asc())
elif sort == "updated_asc":
order_by = (AgentRow.updated_at.asc(), AgentRow.id.asc())
elif sort == "name_asc":
order_by = (func.lower(AgentRow.name).asc(), AgentRow.id.asc())
elif sort == "name_desc":
order_by = (func.lower(AgentRow.name).desc(), AgentRow.id.asc())
else:
order_by = (AgentRow.created_at.desc(), AgentRow.id.asc())
page_stmt = (
select(AgentRow)
.order_by(pin_rank.asc(), *order_by)
.limit(page_size)
.offset(max(page - 1, 0) * page_size)
)
if filters:
where = and_(*filters)
count_stmt = count_stmt.where(where)
page_stmt = page_stmt.where(where)
async with self._sf() as session:
total = await session.scalar(count_stmt) or 0
result = await session.execute(page_stmt)
rows = [_row_to_dict(row) for row in result.scalars()]
return rows, total
# --- Listing cache (skills / tool_groups / model) ---
async def replace_skills(self, agent_id: str, skills: list[str] | None) -> None:
async with self._sf() as session:
await session.execute(delete(AgentSkillRow).where(AgentSkillRow.agent_id == agent_id))
for position, name in enumerate(skills or []):
if not isinstance(name, str):
continue
session.add(
AgentSkillRow(
id=uuid4().hex,
agent_id=agent_id,
skill_name=name,
position=position,
)
)
await session.commit()
async def set_extras(self, agent_id: str, *, tool_groups: list[str] | None, model: str | None) -> None:
async with self._sf() as session:
row = await session.get(AgentExtrasRow, agent_id)
now = datetime.now(UTC)
if row is None:
session.add(
AgentExtrasRow(
agent_id=agent_id,
tool_groups=tool_groups,
model=model,
skills_synced_at=now,
)
)
else:
row.tool_groups = tool_groups
row.model = model
row.skills_synced_at = now
await session.commit()
async def get_extras_for(self, agent_ids: list[str]) -> dict[str, dict[str, Any]]:
if not agent_ids:
return {}
out: dict[str, dict[str, Any]] = {}
async with self._sf() as session:
extras = await session.execute(
select(AgentExtrasRow).where(AgentExtrasRow.agent_id.in_(agent_ids))
)
for row in extras.scalars():
out[row.agent_id] = {
"skills": [], # default; filled below if skill rows exist
"tool_groups": row.tool_groups,
"model": row.model,
}
skill_rows = await session.execute(
select(AgentSkillRow.agent_id, AgentSkillRow.skill_name)
.where(AgentSkillRow.agent_id.in_(agent_ids))
.order_by(AgentSkillRow.agent_id.asc(), AgentSkillRow.position.asc())
)
for agent_id, skill_name in skill_rows.all():
entry = out.setdefault(
agent_id, {"skills": [], "tool_groups": None, "model": None}
)
entry["skills"].append(skill_name)
return out
async def delete_extras_for_agent(self, agent_id: str) -> None:
async with self._sf() as session:
await session.execute(delete(AgentSkillRow).where(AgentSkillRow.agent_id == agent_id))
await session.execute(delete(AgentExtrasRow).where(AgentExtrasRow.agent_id == agent_id))
await session.commit()
# --- Favorites ("收藏") ---
async def add_favorite(self, user_id: str, agent_id: str) -> None:
async with self._sf() as session:
existing = (
await session.execute(
select(AgentFavoriteRow)
.where(AgentFavoriteRow.user_id == user_id, AgentFavoriteRow.agent_id == agent_id)
.limit(1)
)
).scalar_one_or_none()
if existing is not None:
if existing.origin != "user" or existing.position_id is not None:
existing.origin = "user"
existing.position_id = None
await session.commit()
return
session.add(AgentFavoriteRow(id=uuid4().hex, user_id=user_id, agent_id=agent_id))
try:
await session.commit()
except IntegrityError:
# Lost a race against a concurrent favorite of the same agent —
# the unique constraint already guarantees the desired state.
await session.rollback()
async def remove_favorite(self, user_id: str, agent_id: str) -> bool:
async with self._sf() as session:
result = await session.execute(
delete(AgentFavoriteRow).where(
AgentFavoriteRow.user_id == user_id,
AgentFavoriteRow.agent_id == agent_id,
)
)
await session.commit()
return (result.rowcount or 0) > 0
async def list_favorite_ids(self, user_id: str) -> set[str]:
async with self._sf() as session:
result = await session.execute(
select(AgentFavoriteRow.agent_id)
.join(AgentRow, AgentRow.id == AgentFavoriteRow.agent_id)
.where(
AgentFavoriteRow.user_id == user_id,
self._valid_favorite_filter(),
)
)
return {row[0] for row in result.all()}
async def delete_favorites_for_agent(self, agent_id: str) -> None:
async with self._sf() as session:
await session.execute(delete(AgentFavoriteRow).where(AgentFavoriteRow.agent_id == agent_id))
await session.commit()
# --- Pins ("置顶") ---
async def add_pin(self, user_id: str, agent_id: str) -> None:
async with self._sf() as session:
existing = (
await session.execute(
select(AgentPinRow.id)
.where(AgentPinRow.user_id == user_id, AgentPinRow.agent_id == agent_id)
.limit(1)
)
).scalar_one_or_none()
if existing is not None:
return
session.add(AgentPinRow(id=uuid4().hex, user_id=user_id, agent_id=agent_id))
try:
await session.commit()
except IntegrityError:
await session.rollback()
async def remove_pin(self, user_id: str, agent_id: str) -> bool:
async with self._sf() as session:
result = await session.execute(
delete(AgentPinRow).where(
AgentPinRow.user_id == user_id,
AgentPinRow.agent_id == agent_id,
)
)
await session.commit()
return (result.rowcount or 0) > 0
async def list_pin_ids(self, user_id: str) -> set[str]:
async with self._sf() as session:
result = await session.execute(
select(AgentPinRow.agent_id).where(AgentPinRow.user_id == user_id)
)
return {row[0] for row in result.all()}
async def delete_pins_for_agent(self, agent_id: str) -> None:
async with self._sf() as session:
await session.execute(delete(AgentPinRow).where(AgentPinRow.agent_id == agent_id))
await session.commit()