"""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()