574 lines
23 KiB
Python
574 lines
23 KiB
Python
"""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()
|