Phase 5.5-5.9: Plugin-Marketplace, Agent Memory, GraphRAG, Subagents, External Agent API
Check Cross-Plugin Imports / check (push) Has been cancelled
Check Cross-Plugin Imports / check (push) Has been cancelled
5.5 Plugin-Marketplace: - New plugin: marketplace/ (models, routes, services, schemas, config) - MarketplaceListing model (global, no tenant_id) - Ed25519 signature verification via PluginSignature - Endpoints: list, detail, install, verify, categories - Config: MARKETPLACE_SERVER_URL setting 5.6 Agent Memory (persistent): - New plugin: agent_memory/ (models, routes, services, schemas) - AgentMemory model with embedding vector(768) + HNSW index - store_memory() with auto-embedding - retrieve_relevant_memories() with pgvector cosine similarity - Semantic search endpoint 5.7 GraphRAG: - New plugin: graph_rag/ (models, routes, services, provider, schemas) - EntityRelationship model (source/target type+id, relationship_type, metadata) - BFS graph traversal (bidirectional, configurable depth) - GraphRAGSearchProvider registered in unified_search 5.8 Subagents / Multi-Agent: - AgentCoordinator class (create_subtask, wait_for_subtask, aggregate, cancel) - AgentSubtask model + migration 0002_agent_subtasks.sql - 6 new API endpoints for subtask management - Tools registered in AI tool registry 5.9 External Agent API: - external_api.py: POST /run, GET /status, POST /stream (SSE) - Bearer API token authentication - Rate limiting: 10 req/min per token - ExternalAgentRequest/Response schemas 3 new plugins registered in main.py and __init__.py All files py_compile clean
This commit is contained in:
@@ -73,6 +73,9 @@ class Settings(BaseSettings):
|
|||||||
db_retry_max_attempts: int = 3
|
db_retry_max_attempts: int = 3
|
||||||
db_retry_base_delay: float = 0.1
|
db_retry_base_delay: float = 0.1
|
||||||
|
|
||||||
|
# Marketplace
|
||||||
|
marketplace_server_url: str = ""
|
||||||
|
|
||||||
# Rate Limiting
|
# Rate Limiting
|
||||||
rate_limit_login_max: int = 5
|
rate_limit_login_max: int = 5
|
||||||
rate_limit_login_window: int = 900 # 15 min
|
rate_limit_login_window: int = 900 # 15 min
|
||||||
|
|||||||
@@ -475,6 +475,9 @@ def create_app() -> FastAPI:
|
|||||||
"app.plugins.builtins.system_notif",
|
"app.plugins.builtins.system_notif",
|
||||||
"app.plugins.builtins.unified_search",
|
"app.plugins.builtins.unified_search",
|
||||||
"app.plugins.builtins.forgejo_error_reporter",
|
"app.plugins.builtins.forgejo_error_reporter",
|
||||||
|
"app.plugins.builtins.agent_memory",
|
||||||
|
"app.plugins.builtins.graph_rag",
|
||||||
|
"app.plugins.builtins.marketplace",
|
||||||
]
|
]
|
||||||
for mod_name in plugin_modules:
|
for mod_name in plugin_modules:
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -7,12 +7,15 @@ Subdirectory plugins (tags, permissions, entity_links, mail) export their plugin
|
|||||||
class via __init__.py so the registry can discover them as packages.
|
class via __init__.py so the registry can discover them as packages.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from app.plugins.builtins.agent_memory import AgentMemoryPlugin
|
||||||
from app.plugins.builtins.calendar import CalendarPlugin
|
from app.plugins.builtins.calendar import CalendarPlugin
|
||||||
from app.plugins.builtins.dms import DmsPlugin
|
from app.plugins.builtins.dms import DmsPlugin
|
||||||
from app.plugins.builtins.entity_links import EntityLinksPlugin
|
from app.plugins.builtins.entity_links import EntityLinksPlugin
|
||||||
|
from app.plugins.builtins.graph_rag import GraphRAGPlugin
|
||||||
from app.plugins.builtins.mail import MailPlugin
|
from app.plugins.builtins.mail import MailPlugin
|
||||||
from app.plugins.builtins.report_generator import ReportGeneratorPlugin
|
from app.plugins.builtins.report_generator import ReportGeneratorPlugin
|
||||||
from app.plugins.builtins.permissions import PermissionsPlugin
|
from app.plugins.builtins.permissions import PermissionsPlugin
|
||||||
from app.plugins.builtins.tags import TagsPlugin
|
from app.plugins.builtins.tags import TagsPlugin
|
||||||
|
from app.plugins.builtins.marketplace import MarketplacePlugin
|
||||||
|
|
||||||
__all__ = ["TagsPlugin", "PermissionsPlugin", "EntityLinksPlugin", "DmsPlugin", "CalendarPlugin", "MailPlugin", "ReportGeneratorPlugin"]
|
__all__ = ["AgentMemoryPlugin", "TagsPlugin", "PermissionsPlugin", "EntityLinksPlugin", "DmsPlugin", "CalendarPlugin", "MailPlugin", "ReportGeneratorPlugin", "GraphRAGPlugin", "MarketplacePlugin"]
|
||||||
|
|||||||
@@ -0,0 +1,5 @@
|
|||||||
|
"""Agent Memory plugin package."""
|
||||||
|
|
||||||
|
from app.plugins.builtins.agent_memory.plugin import AgentMemoryPlugin
|
||||||
|
|
||||||
|
__all__ = ["AgentMemoryPlugin"]
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
-- Agent Memory plugin initial migration: creates agent_memories table with pgvector
|
||||||
|
CREATE TABLE IF NOT EXISTS agent_memories (
|
||||||
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
tenant_id UUID NOT NULL,
|
||||||
|
agent_id UUID NOT NULL,
|
||||||
|
memory_type VARCHAR(50) NOT NULL DEFAULT 'fact',
|
||||||
|
content TEXT NOT NULL,
|
||||||
|
embedding vector(768),
|
||||||
|
owner_id UUID REFERENCES users(id) ON DELETE SET NULL,
|
||||||
|
deleted_at TIMESTAMPTZ,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_memories_tenant_agent ON agent_memories(tenant_id, agent_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_memories_tenant_type ON agent_memories(tenant_id, memory_type);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_memories_agent_id ON agent_memories(agent_id);
|
||||||
|
|
||||||
|
-- HNSW index for fast vector similarity search on agent_memories.embedding
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_memories_embedding_hnsw ON agent_memories USING hnsw (embedding vector_cosine_ops) WITH (m = 16, ef_construction = 200);
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
"""AgentMemory model for persistent agent memory with pgvector embeddings."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from sqlalchemy import Index, String, Text
|
||||||
|
from sqlalchemy.dialects.postgresql import UUID as PGUUID
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from app.core.db import Base, TenantMixin
|
||||||
|
from app.models.owned_mixin import OwnedMixin
|
||||||
|
|
||||||
|
|
||||||
|
class AgentMemory(Base, TenantMixin, OwnedMixin):
|
||||||
|
"""Persistent agent memory with semantic search via pgvector.
|
||||||
|
|
||||||
|
Stores agent memories (facts, context, learned patterns) with
|
||||||
|
vector embeddings for semantic retrieval.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__tablename__ = "agent_memories"
|
||||||
|
__table_args__ = (
|
||||||
|
Index("ix_agent_memories_tenant_agent", "tenant_id", "agent_id"),
|
||||||
|
Index("ix_agent_memories_tenant_type", "tenant_id", "memory_type"),
|
||||||
|
)
|
||||||
|
|
||||||
|
id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
|
||||||
|
)
|
||||||
|
agent_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), nullable=False, index=True
|
||||||
|
)
|
||||||
|
memory_type: Mapped[str] = mapped_column(
|
||||||
|
String(50), nullable=False, default="fact"
|
||||||
|
)
|
||||||
|
content: Mapped[str] = mapped_column(Text, nullable=False)
|
||||||
|
# embedding column is managed via raw SQL (pgvector extension)
|
||||||
|
# embedding vector(768) — see migration 0001_initial.sql
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
"""Agent Memory plugin — persistent agent memory with pgvector semantic search."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from app.plugins.base import BasePlugin
|
||||||
|
from app.plugins.manifest import PluginManifest, PluginRouteDef
|
||||||
|
|
||||||
|
|
||||||
|
class AgentMemoryPlugin(BasePlugin):
|
||||||
|
"""Agent Memory plugin for persistent agent memory with semantic search."""
|
||||||
|
|
||||||
|
manifest = PluginManifest(
|
||||||
|
name="agent_memory",
|
||||||
|
version="1.0.0",
|
||||||
|
display_name="Agent Memory",
|
||||||
|
description="Persistent agent memory with pgvector semantic search. Stores facts, context, patterns, and instructions for AI agents.",
|
||||||
|
dependencies=["unified_search"],
|
||||||
|
routes=[
|
||||||
|
PluginRouteDef(
|
||||||
|
path="/api/v1/agent-memory",
|
||||||
|
module="app.plugins.builtins.agent_memory.routes",
|
||||||
|
router_attr="router",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
events=[],
|
||||||
|
migrations=["0001_initial.sql"],
|
||||||
|
permissions=[
|
||||||
|
"agent_memory:read",
|
||||||
|
"agent_memory:write",
|
||||||
|
],
|
||||||
|
is_core=True,
|
||||||
|
author="LeoCRM Team",
|
||||||
|
min_app_version="1.0.0",
|
||||||
|
contract_version="1.0.0",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def on_deactivate(self, db, service_container, event_bus) -> None:
|
||||||
|
"""Deactivate plugin: unregister contract and event listeners."""
|
||||||
|
from app.plugins.builtins.contracts import get_contract_registry
|
||||||
|
get_contract_registry().unregister(self.manifest.name)
|
||||||
|
await super().on_deactivate(db, service_container, event_bus)
|
||||||
@@ -0,0 +1,202 @@
|
|||||||
|
"""Agent Memory plugin routes — CRUD for persistent agent memories."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.core.db import get_db
|
||||||
|
from app.deps import get_current_user, require_permission
|
||||||
|
from app.plugins.builtins.agent_memory.models import AgentMemory
|
||||||
|
from app.plugins.builtins.agent_memory.schemas import (
|
||||||
|
AgentMemoryCreate,
|
||||||
|
AgentMemoryRead,
|
||||||
|
AgentMemoryUpdate,
|
||||||
|
)
|
||||||
|
from app.plugins.builtins.agent_memory.services import (
|
||||||
|
delete_memory,
|
||||||
|
retrieve_relevant_memories,
|
||||||
|
store_memory,
|
||||||
|
)
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api/v1/agent-memory", tags=["agent-memory"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("", status_code=status.HTTP_201_CREATED, dependencies=[Depends(require_permission("agent_memory:write"))])
|
||||||
|
async def create_memory(
|
||||||
|
body: AgentMemoryCreate,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""Create a new agent memory with embedding."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
user_id = uuid.UUID(current_user["user_id"])
|
||||||
|
agent_id = _parse_uuid(body.agent_id, "agent_id")
|
||||||
|
|
||||||
|
result = await store_memory(
|
||||||
|
db=db,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
agent_id=agent_id,
|
||||||
|
content=body.content,
|
||||||
|
memory_type=body.memory_type,
|
||||||
|
owner_id=user_id,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("", dependencies=[Depends(require_permission("agent_memory:read"))])
|
||||||
|
async def list_memories(
|
||||||
|
agent_id: str = Query(..., description="Filter by agent UUID"),
|
||||||
|
memory_type: str | None = Query(None, description="Filter by memory type"),
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(50, ge=1, le=200),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""List memories for an agent, with optional type filter."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
aid = _parse_uuid(agent_id, "agent_id")
|
||||||
|
|
||||||
|
base = select(AgentMemory).where(
|
||||||
|
AgentMemory.tenant_id == tenant_id,
|
||||||
|
AgentMemory.agent_id == aid,
|
||||||
|
AgentMemory.deleted_at.is_(None),
|
||||||
|
)
|
||||||
|
if memory_type:
|
||||||
|
base = base.where(AgentMemory.memory_type == memory_type)
|
||||||
|
|
||||||
|
# Count total
|
||||||
|
count_q = select(func.count()).select_from(base.subquery())
|
||||||
|
total_result = await db.execute(count_q)
|
||||||
|
total = total_result.scalar_one()
|
||||||
|
|
||||||
|
# Paginated query
|
||||||
|
offset = (page - 1) * page_size
|
||||||
|
stmt = base.order_by(AgentMemory.created_at.desc()).offset(offset).limit(page_size)
|
||||||
|
result = await db.execute(stmt)
|
||||||
|
memories = result.scalars().all()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"id": str(m.id),
|
||||||
|
"agent_id": str(m.agent_id),
|
||||||
|
"memory_type": m.memory_type,
|
||||||
|
"content": m.content,
|
||||||
|
"owner_id": str(m.owner_id) if m.owner_id else None,
|
||||||
|
"created_at": m.created_at.isoformat() if m.created_at else None,
|
||||||
|
"updated_at": m.updated_at.isoformat() if m.updated_at else None,
|
||||||
|
"score": 0.0,
|
||||||
|
}
|
||||||
|
for m in memories
|
||||||
|
],
|
||||||
|
"total": total,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/search", dependencies=[Depends(require_permission("agent_memory:read"))])
|
||||||
|
async def search_memories(
|
||||||
|
agent_id: str = Query(..., description="Agent UUID"),
|
||||||
|
query: str = Query(..., min_length=1, description="Natural language query"),
|
||||||
|
memory_type: str | None = Query(None, description="Optional type filter"),
|
||||||
|
limit: int = Query(10, ge=1, le=100),
|
||||||
|
min_score: float = Query(0.5, ge=0.0, le=1.0),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""Semantic search over agent memories using pgvector."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
aid = _parse_uuid(agent_id, "agent_id")
|
||||||
|
|
||||||
|
results = await retrieve_relevant_memories(
|
||||||
|
db=db,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
agent_id=aid,
|
||||||
|
query=query,
|
||||||
|
limit=limit,
|
||||||
|
memory_type=memory_type,
|
||||||
|
min_score=min_score,
|
||||||
|
)
|
||||||
|
return {"items": results, "total": len(results)}
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch("/{memory_id}", dependencies=[Depends(require_permission("agent_memory:write"))])
|
||||||
|
async def update_memory(
|
||||||
|
memory_id: str,
|
||||||
|
body: AgentMemoryUpdate,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""Update a memory's content or type."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
mid = _parse_uuid(memory_id, "memory_id")
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentMemory).where(
|
||||||
|
AgentMemory.id == mid,
|
||||||
|
AgentMemory.tenant_id == tenant_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
memory = result.scalar_one_or_none()
|
||||||
|
if memory is None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail={"detail": "Memory not found", "code": "not_found"},
|
||||||
|
)
|
||||||
|
|
||||||
|
data = body.model_dump(exclude_unset=True)
|
||||||
|
if "content" in data:
|
||||||
|
memory.content = data["content"]
|
||||||
|
# Regenerate embedding for updated content
|
||||||
|
from app.plugins.builtins.unified_search.embedding import generate_embedding
|
||||||
|
embedding = await generate_embedding(data["content"], db=db, tenant_id=tenant_id)
|
||||||
|
if embedding:
|
||||||
|
from sqlalchemy import text as sql_text
|
||||||
|
sql = sql_text(
|
||||||
|
"UPDATE agent_memories SET embedding = cast(:emb AS vector) WHERE id = :mid"
|
||||||
|
)
|
||||||
|
await db.execute(sql, {"emb": str(embedding), "mid": memory.id})
|
||||||
|
if "memory_type" in data:
|
||||||
|
memory.memory_type = data["memory_type"]
|
||||||
|
|
||||||
|
await db.flush()
|
||||||
|
return {
|
||||||
|
"id": str(memory.id),
|
||||||
|
"agent_id": str(memory.agent_id),
|
||||||
|
"memory_type": memory.memory_type,
|
||||||
|
"content": memory.content,
|
||||||
|
"owner_id": str(memory.owner_id) if memory.owner_id else None,
|
||||||
|
"updated_at": memory.updated_at.isoformat() if memory.updated_at else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{memory_id}", status_code=status.HTTP_204_NO_CONTENT, dependencies=[Depends(require_permission("agent_memory:write"))])
|
||||||
|
async def delete_memory_route(
|
||||||
|
memory_id: str,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""Delete a memory by ID."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
mid = _parse_uuid(memory_id, "memory_id")
|
||||||
|
|
||||||
|
success = await delete_memory(db, tenant_id, mid)
|
||||||
|
if not success:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail={"detail": "Memory not found", "code": "not_found"},
|
||||||
|
)
|
||||||
|
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_uuid(val: str, field: str) -> uuid.UUID:
|
||||||
|
try:
|
||||||
|
return uuid.UUID(val)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail={"detail": f"Invalid {field}", "code": "invalid_id"},
|
||||||
|
) from None
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
"""Pydantic schemas for the Agent Memory plugin."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
class AgentMemoryCreate(BaseModel):
|
||||||
|
agent_id: str = Field(..., description="UUID of the agent")
|
||||||
|
memory_type: str = Field("fact", max_length=50, description="Type of memory (fact, context, pattern, instruction)")
|
||||||
|
content: str = Field(..., min_length=1, description="Memory content text")
|
||||||
|
|
||||||
|
|
||||||
|
class AgentMemoryUpdate(BaseModel):
|
||||||
|
content: str | None = Field(None, min_length=1)
|
||||||
|
memory_type: str | None = Field(None, max_length=50)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentMemoryRead(BaseModel):
|
||||||
|
id: str
|
||||||
|
agent_id: str
|
||||||
|
memory_type: str
|
||||||
|
content: str
|
||||||
|
owner_id: str | None = None
|
||||||
|
created_at: datetime | None = None
|
||||||
|
updated_at: datetime | None = None
|
||||||
|
score: float = 0.0
|
||||||
|
|
||||||
|
|
||||||
|
class AgentMemoryListResponse(BaseModel):
|
||||||
|
items: list[AgentMemoryRead]
|
||||||
|
total: int
|
||||||
@@ -0,0 +1,176 @@
|
|||||||
|
"""Agent Memory services — store and retrieve memories with semantic search."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import select, text as sql_text
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.plugins.builtins.agent_memory.models import AgentMemory
|
||||||
|
from app.plugins.builtins.unified_search.embedding import generate_embedding
|
||||||
|
|
||||||
|
|
||||||
|
async def store_memory(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
agent_id: uuid.UUID,
|
||||||
|
content: str,
|
||||||
|
memory_type: str = "fact",
|
||||||
|
owner_id: uuid.UUID | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Store a new memory with embedding generation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session.
|
||||||
|
tenant_id: Tenant UUID.
|
||||||
|
agent_id: Agent UUID.
|
||||||
|
content: Memory content text.
|
||||||
|
memory_type: Type of memory (fact, context, pattern, instruction).
|
||||||
|
owner_id: Optional owner user UUID.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with the created memory data.
|
||||||
|
"""
|
||||||
|
# Generate embedding for semantic search
|
||||||
|
embedding = await generate_embedding(content, db=db, tenant_id=tenant_id)
|
||||||
|
|
||||||
|
memory = AgentMemory(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
agent_id=agent_id,
|
||||||
|
memory_type=memory_type,
|
||||||
|
content=content,
|
||||||
|
owner_id=owner_id,
|
||||||
|
)
|
||||||
|
db.add(memory)
|
||||||
|
await db.flush()
|
||||||
|
await db.refresh(memory)
|
||||||
|
|
||||||
|
# Store embedding if generated successfully
|
||||||
|
if embedding:
|
||||||
|
sql = sql_text(
|
||||||
|
"UPDATE agent_memories SET embedding = cast(:emb AS vector) WHERE id = :mid"
|
||||||
|
)
|
||||||
|
await db.execute(sql, {"emb": str(embedding), "mid": memory.id})
|
||||||
|
await db.flush()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": str(memory.id),
|
||||||
|
"agent_id": str(memory.agent_id),
|
||||||
|
"memory_type": memory.memory_type,
|
||||||
|
"content": memory.content,
|
||||||
|
"owner_id": str(memory.owner_id) if memory.owner_id else None,
|
||||||
|
"created_at": memory.created_at.isoformat() if memory.created_at else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def retrieve_relevant_memories(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
agent_id: uuid.UUID,
|
||||||
|
query: str,
|
||||||
|
limit: int = 10,
|
||||||
|
memory_type: str | None = None,
|
||||||
|
min_score: float = 0.5,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Retrieve semantically relevant memories for an agent.
|
||||||
|
|
||||||
|
Uses pgvector cosine similarity search to find memories whose
|
||||||
|
embedding is closest to the query embedding.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session.
|
||||||
|
tenant_id: Tenant UUID.
|
||||||
|
agent_id: Agent UUID.
|
||||||
|
query: Natural language query to match against memories.
|
||||||
|
limit: Maximum number of results.
|
||||||
|
memory_type: Optional filter by memory type.
|
||||||
|
min_score: Minimum similarity score threshold (0.0 to 1.0).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of memory dicts with similarity scores, sorted by relevance.
|
||||||
|
"""
|
||||||
|
# Generate embedding for the query
|
||||||
|
query_embedding = await generate_embedding(query, db=db, tenant_id=tenant_id)
|
||||||
|
if not query_embedding:
|
||||||
|
# Fallback: return recent memories if embedding fails
|
||||||
|
stmt = select(AgentMemory).where(
|
||||||
|
AgentMemory.tenant_id == tenant_id,
|
||||||
|
AgentMemory.agent_id == agent_id,
|
||||||
|
AgentMemory.deleted_at.is_(None),
|
||||||
|
)
|
||||||
|
if memory_type:
|
||||||
|
stmt = stmt.where(AgentMemory.memory_type == memory_type)
|
||||||
|
stmt = stmt.order_by(AgentMemory.created_at.desc()).limit(limit)
|
||||||
|
result = await db.execute(stmt)
|
||||||
|
memories = result.scalars().all()
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"id": str(m.id),
|
||||||
|
"agent_id": str(m.agent_id),
|
||||||
|
"memory_type": m.memory_type,
|
||||||
|
"content": m.content,
|
||||||
|
"score": 0.0,
|
||||||
|
"created_at": m.created_at.isoformat() if m.created_at else None,
|
||||||
|
}
|
||||||
|
for m in memories
|
||||||
|
]
|
||||||
|
|
||||||
|
# Vector similarity search
|
||||||
|
type_filter = "AND m.memory_type = :mtype" if memory_type else ""
|
||||||
|
sql = sql_text(f"""
|
||||||
|
SELECT m.*, 1 - (m.embedding <=> cast(:emb AS vector)) AS score
|
||||||
|
FROM agent_memories m
|
||||||
|
WHERE m.tenant_id = :tid
|
||||||
|
AND m.agent_id = :aid
|
||||||
|
AND m.deleted_at IS NULL
|
||||||
|
AND m.embedding IS NOT NULL
|
||||||
|
{type_filter}
|
||||||
|
ORDER BY m.embedding <=> cast(:emb AS vector)
|
||||||
|
LIMIT :lim
|
||||||
|
""")
|
||||||
|
|
||||||
|
params: dict[str, Any] = {
|
||||||
|
"emb": str(query_embedding),
|
||||||
|
"tid": tenant_id,
|
||||||
|
"aid": agent_id,
|
||||||
|
"lim": limit,
|
||||||
|
}
|
||||||
|
if memory_type:
|
||||||
|
params["mtype"] = memory_type
|
||||||
|
|
||||||
|
result = await db.execute(sql, params)
|
||||||
|
rows = result.mappings().all()
|
||||||
|
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"id": str(row["id"]),
|
||||||
|
"agent_id": str(row["agent_id"]),
|
||||||
|
"memory_type": row["memory_type"],
|
||||||
|
"content": row["content"],
|
||||||
|
"score": float(row["score"]),
|
||||||
|
"created_at": row["created_at"].isoformat() if row.get("created_at") else None,
|
||||||
|
}
|
||||||
|
for row in rows
|
||||||
|
if float(row["score"]) >= min_score
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
async def delete_memory(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
memory_id: uuid.UUID,
|
||||||
|
) -> bool:
|
||||||
|
"""Delete a memory by ID."""
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentMemory).where(
|
||||||
|
AgentMemory.id == memory_id,
|
||||||
|
AgentMemory.tenant_id == tenant_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
memory = result.scalar_one_or_none()
|
||||||
|
if memory is None:
|
||||||
|
return False
|
||||||
|
await db.delete(memory)
|
||||||
|
return True
|
||||||
@@ -0,0 +1,297 @@
|
|||||||
|
"""External Agent API — allows external systems to interact with AI agents.
|
||||||
|
|
||||||
|
Provides REST endpoints for running agents, checking status, and streaming
|
||||||
|
responses. Authenticated via Bearer API tokens with rate limiting.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.core.db import get_db, set_tenant_context
|
||||||
|
from app.deps import get_current_user_bearer, require_permission
|
||||||
|
from app.plugins.builtins.ai_assistant.schemas import (
|
||||||
|
ExternalAgentRequest,
|
||||||
|
ExternalAgentResponse,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api/v1/external/agent", tags=["external-agent"])
|
||||||
|
|
||||||
|
|
||||||
|
async def _check_external_rate_limit(request: Request, tenant_id: str, token_prefix: str) -> None:
|
||||||
|
"""Check rate limit: 10 requests/minute per token."""
|
||||||
|
from app.core.rate_limit import check_rate_limit
|
||||||
|
|
||||||
|
redis_key = f"rate:external_agent:{tenant_id}:{token_prefix}"
|
||||||
|
await check_rate_limit(redis_key, max_attempts=10, window_seconds=60)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/{agent_id}/run",
|
||||||
|
dependencies=[Depends(require_permission("ai:write"))],
|
||||||
|
)
|
||||||
|
async def run_agent_external(
|
||||||
|
agent_id: str,
|
||||||
|
data: ExternalAgentRequest,
|
||||||
|
request: Request,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user_bearer),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Execute an AI agent with a message input.
|
||||||
|
|
||||||
|
Authenticated via Bearer API token. Rate limited to 10 requests/minute.
|
||||||
|
"""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
await set_tenant_context(db, tenant_id)
|
||||||
|
|
||||||
|
# Rate limiting
|
||||||
|
token_prefix = current_user.get("token_prefix", "default")
|
||||||
|
await _check_external_rate_limit(request, str(tenant_id), token_prefix)
|
||||||
|
|
||||||
|
try:
|
||||||
|
aid = uuid.UUID(agent_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid agent ID")
|
||||||
|
|
||||||
|
# Find the agent
|
||||||
|
from app.plugins.builtins.ai_assistant.models import AIAgent
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AIAgent)
|
||||||
|
.where(AIAgent.id == aid)
|
||||||
|
.where(AIAgent.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
agent = result.scalar_one_or_none()
|
||||||
|
if agent is None:
|
||||||
|
raise HTTPException(status_code=404, detail="Agent not found")
|
||||||
|
if not agent.is_active:
|
||||||
|
raise HTTPException(status_code=400, detail="Agent is not active")
|
||||||
|
|
||||||
|
# Create or find a session for this external interaction
|
||||||
|
from app.plugins.builtins.ai_assistant.models import AIChatSession, AIChatMessage
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
session = AIChatSession(
|
||||||
|
user_id=uuid.UUID(current_user["user_id"]),
|
||||||
|
agent_id=agent.id,
|
||||||
|
title=f"External: {data.message[:50]}" if data.message else "External Agent Run",
|
||||||
|
is_sidebar=False,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
owner_id=uuid.UUID(current_user["user_id"]),
|
||||||
|
)
|
||||||
|
db.add(session)
|
||||||
|
await db.flush()
|
||||||
|
|
||||||
|
# Store the user message
|
||||||
|
user_msg = AIChatMessage(
|
||||||
|
session_id=session.id,
|
||||||
|
role="user",
|
||||||
|
content=data.message,
|
||||||
|
tokens=0,
|
||||||
|
model_used=agent.name or "external",
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
)
|
||||||
|
db.add(user_msg)
|
||||||
|
await db.flush()
|
||||||
|
|
||||||
|
# Build user context for RBAC
|
||||||
|
user_context = {
|
||||||
|
"user_id": current_user["user_id"],
|
||||||
|
"tenant_id": current_user["tenant_id"],
|
||||||
|
"role": current_user.get("role", ""),
|
||||||
|
"permissions": current_user.get("permissions", []),
|
||||||
|
"denied_permissions": current_user.get("denied_permissions", []),
|
||||||
|
"is_system_admin": current_user.get("is_system_admin", False),
|
||||||
|
"field_permissions": current_user.get("field_permissions", {}),
|
||||||
|
}
|
||||||
|
|
||||||
|
# Run the agent via streaming chat (non-streaming mode)
|
||||||
|
from app.plugins.builtins.ai_assistant.services import stream_chat
|
||||||
|
|
||||||
|
full_response = ""
|
||||||
|
async with get_db() as stream_db:
|
||||||
|
await set_tenant_context(stream_db, tenant_id)
|
||||||
|
async for chunk in stream_chat(
|
||||||
|
stream_db, session, agent, data.message, user_context, tenant_id
|
||||||
|
):
|
||||||
|
if chunk.startswith("data: ") and chunk != "data: [DONE]\n\n":
|
||||||
|
try:
|
||||||
|
payload = json.loads(chunk[6:])
|
||||||
|
if "content" in payload:
|
||||||
|
full_response += payload["content"]
|
||||||
|
except (json.JSONDecodeError, IndexError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Count tokens (approximate)
|
||||||
|
tokens_used = len(full_response.split()) + len(data.message.split())
|
||||||
|
|
||||||
|
return ExternalAgentResponse(
|
||||||
|
response=full_response,
|
||||||
|
agent_id=str(agent.id),
|
||||||
|
session_id=str(session.id),
|
||||||
|
tokens_used=tokens_used,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"/{agent_id}/status",
|
||||||
|
dependencies=[Depends(require_permission("ai:read"))],
|
||||||
|
)
|
||||||
|
async def get_agent_status_external(
|
||||||
|
agent_id: str,
|
||||||
|
request: Request,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user_bearer),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Get the status of an AI agent.
|
||||||
|
|
||||||
|
Returns agent configuration, active status, and recent run stats.
|
||||||
|
"""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
await set_tenant_context(db, tenant_id)
|
||||||
|
|
||||||
|
# Rate limiting
|
||||||
|
token_prefix = current_user.get("token_prefix", "default")
|
||||||
|
await _check_external_rate_limit(request, str(tenant_id), token_prefix)
|
||||||
|
|
||||||
|
try:
|
||||||
|
aid = uuid.UUID(agent_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid agent ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.ai_assistant.models import AIAgent
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AIAgent)
|
||||||
|
.where(AIAgent.id == aid)
|
||||||
|
.where(AIAgent.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
agent = result.scalar_one_or_none()
|
||||||
|
if agent is None:
|
||||||
|
raise HTTPException(status_code=404, detail="Agent not found")
|
||||||
|
|
||||||
|
# Get recent run stats
|
||||||
|
from app.plugins.builtins.automation.models import AgentRun
|
||||||
|
from sqlalchemy import func
|
||||||
|
|
||||||
|
recent_runs = await db.execute(
|
||||||
|
select(func.count())
|
||||||
|
.select_from(AgentRun)
|
||||||
|
.where(AgentRun.agent_id == aid)
|
||||||
|
.where(AgentRun.tenant_id == tenant_id)
|
||||||
|
)
|
||||||
|
total_runs = recent_runs.scalar() or 0
|
||||||
|
|
||||||
|
return {
|
||||||
|
"agent_id": str(agent.id),
|
||||||
|
"name": agent.name,
|
||||||
|
"description": agent.description,
|
||||||
|
"is_active": agent.is_active,
|
||||||
|
"tool_count": len(agent.tool_ids or []),
|
||||||
|
"total_runs": total_runs,
|
||||||
|
"created_at": agent.created_at.isoformat() if agent.created_at else None,
|
||||||
|
"updated_at": agent.updated_at.isoformat() if agent.updated_at else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/{agent_id}/stream",
|
||||||
|
dependencies=[Depends(require_permission("ai:write"))],
|
||||||
|
)
|
||||||
|
async def stream_agent_external(
|
||||||
|
agent_id: str,
|
||||||
|
data: ExternalAgentRequest,
|
||||||
|
request: Request,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user_bearer),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Stream an AI agent response via SSE.
|
||||||
|
|
||||||
|
Authenticated via Bearer API token. Rate limited to 10 requests/minute.
|
||||||
|
Returns Server-Sent Events stream.
|
||||||
|
"""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
await set_tenant_context(db, tenant_id)
|
||||||
|
|
||||||
|
# Rate limiting
|
||||||
|
token_prefix = current_user.get("token_prefix", "default")
|
||||||
|
await _check_external_rate_limit(request, str(tenant_id), token_prefix)
|
||||||
|
|
||||||
|
try:
|
||||||
|
aid = uuid.UUID(agent_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid agent ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.ai_assistant.models import AIAgent
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AIAgent)
|
||||||
|
.where(AIAgent.id == aid)
|
||||||
|
.where(AIAgent.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
agent = result.scalar_one_or_none()
|
||||||
|
if agent is None:
|
||||||
|
raise HTTPException(status_code=404, detail="Agent not found")
|
||||||
|
if not agent.is_active:
|
||||||
|
raise HTTPException(status_code=400, detail="Agent is not active")
|
||||||
|
|
||||||
|
# Create session
|
||||||
|
from app.plugins.builtins.ai_assistant.models import AIChatSession
|
||||||
|
|
||||||
|
session = AIChatSession(
|
||||||
|
user_id=uuid.UUID(current_user["user_id"]),
|
||||||
|
agent_id=agent.id,
|
||||||
|
title=f"External Stream: {data.message[:50]}" if data.message else "External Agent Stream",
|
||||||
|
is_sidebar=False,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
owner_id=uuid.UUID(current_user["user_id"]),
|
||||||
|
)
|
||||||
|
db.add(session)
|
||||||
|
await db.commit()
|
||||||
|
|
||||||
|
# Build user context
|
||||||
|
user_context = {
|
||||||
|
"user_id": current_user["user_id"],
|
||||||
|
"tenant_id": current_user["tenant_id"],
|
||||||
|
"role": current_user.get("role", ""),
|
||||||
|
"permissions": current_user.get("permissions", []),
|
||||||
|
"denied_permissions": current_user.get("denied_permissions", []),
|
||||||
|
"is_system_admin": current_user.get("is_system_admin", False),
|
||||||
|
"field_permissions": current_user.get("field_permissions", {}),
|
||||||
|
}
|
||||||
|
|
||||||
|
from app.plugins.builtins.ai_assistant.services import stream_chat
|
||||||
|
|
||||||
|
async def event_stream():
|
||||||
|
from app.core.db import get_session_factory
|
||||||
|
factory = get_session_factory()
|
||||||
|
async with factory() as stream_db:
|
||||||
|
await set_tenant_context(stream_db, tenant_id)
|
||||||
|
async for chunk in stream_chat(
|
||||||
|
stream_db, session, agent, data.message, user_context, tenant_id
|
||||||
|
):
|
||||||
|
yield chunk
|
||||||
|
yield "data: [DONE]\n\n"
|
||||||
|
|
||||||
|
return StreamingResponse(
|
||||||
|
event_stream(),
|
||||||
|
media_type="text/event-stream",
|
||||||
|
headers={
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "keep-alive",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
@@ -29,6 +29,11 @@ class AIAssistantPlugin(BasePlugin):
|
|||||||
module="app.plugins.builtins.ai_assistant.routes",
|
module="app.plugins.builtins.ai_assistant.routes",
|
||||||
router_attr="router",
|
router_attr="router",
|
||||||
),
|
),
|
||||||
|
PluginRouteDef(
|
||||||
|
path="/api/v1/external",
|
||||||
|
module="app.plugins.builtins.ai_assistant.external_api",
|
||||||
|
router_attr="router",
|
||||||
|
),
|
||||||
],
|
],
|
||||||
events=[],
|
events=[],
|
||||||
migrations=["0001_initial.sql", "0002_folders_attachments.sql"],
|
migrations=["0001_initial.sql", "0002_folders_attachments.sql"],
|
||||||
|
|||||||
@@ -252,3 +252,23 @@ class AIToolResponse(BaseModel):
|
|||||||
plugin_name: str
|
plugin_name: str
|
||||||
required_permission: str | None = None
|
required_permission: str | None = None
|
||||||
category: str
|
category: str
|
||||||
|
|
||||||
|
|
||||||
|
# ─── External Agent API Schemas (Phase 5.9) ───
|
||||||
|
|
||||||
|
|
||||||
|
class ExternalAgentRequest(BaseModel):
|
||||||
|
"""Request to run an AI agent from an external system."""
|
||||||
|
|
||||||
|
message: str = Field(..., min_length=1, description="The message/input for the agent")
|
||||||
|
context: dict[str, Any] = Field(default_factory=dict, description="Optional context data")
|
||||||
|
stream: bool = Field(default=False, description="Whether to stream the response")
|
||||||
|
|
||||||
|
|
||||||
|
class ExternalAgentResponse(BaseModel):
|
||||||
|
"""Response from an external agent execution."""
|
||||||
|
|
||||||
|
response: str = Field(default="", description="The agent's response text")
|
||||||
|
agent_id: str = Field(default="", description="ID of the agent that responded")
|
||||||
|
session_id: str = Field(default="", description="Chat session ID")
|
||||||
|
tokens_used: int = Field(default=0, description="Approximate token count")
|
||||||
|
|||||||
@@ -0,0 +1,362 @@
|
|||||||
|
"""Agent Coordinator — Multi-Agent Orchestration for subtask delegation.
|
||||||
|
|
||||||
|
Provides the AgentCoordinator class that manages subtask creation, waiting,
|
||||||
|
aggregation, and cancellation between AI agents.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import select, update as sa_update
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.models import AgentSubtask
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentCoordinator:
|
||||||
|
"""Coordinates subtask delegation between AI agents."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def create_subtask(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
parent_agent_id: uuid.UUID,
|
||||||
|
child_agent_id: uuid.UUID,
|
||||||
|
task_description: str,
|
||||||
|
) -> AgentSubtask:
|
||||||
|
"""Create a new subtask for a child agent.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session
|
||||||
|
tenant_id: Tenant ID
|
||||||
|
parent_agent_id: UUID of the parent agent creating the subtask
|
||||||
|
child_agent_id: UUID of the child agent that will execute the subtask
|
||||||
|
task_description: Description of the task to execute
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The created AgentSubtask instance
|
||||||
|
"""
|
||||||
|
subtask = AgentSubtask(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
parent_agent_id=parent_agent_id,
|
||||||
|
child_agent_id=child_agent_id,
|
||||||
|
task_description=task_description,
|
||||||
|
status="pending",
|
||||||
|
result={},
|
||||||
|
)
|
||||||
|
db.add(subtask)
|
||||||
|
await db.flush()
|
||||||
|
logger.info(
|
||||||
|
"Subtask created: %s (parent=%s -> child=%s)",
|
||||||
|
subtask.id, parent_agent_id, child_agent_id,
|
||||||
|
)
|
||||||
|
return subtask
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def wait_for_subtask(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
subtask_id: uuid.UUID,
|
||||||
|
poll_interval: float = 0.5,
|
||||||
|
timeout: float = 300.0,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Wait for a subtask to complete, fail, or be cancelled.
|
||||||
|
|
||||||
|
Polls the database until the subtask reaches a terminal state
|
||||||
|
or the timeout is exceeded.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session
|
||||||
|
tenant_id: Tenant ID
|
||||||
|
subtask_id: UUID of the subtask to wait for
|
||||||
|
poll_interval: Seconds between polls (default 0.5)
|
||||||
|
timeout: Maximum seconds to wait (default 300)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with status and result/error
|
||||||
|
"""
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
start_time = datetime.now(UTC)
|
||||||
|
while True:
|
||||||
|
elapsed = (datetime.now(UTC) - start_time).total_seconds()
|
||||||
|
if elapsed > timeout:
|
||||||
|
# Mark as timed out
|
||||||
|
await db.execute(
|
||||||
|
sa_update(AgentSubtask)
|
||||||
|
.where(AgentSubtask.id == subtask_id)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
.values(
|
||||||
|
status="failed",
|
||||||
|
result={"error": "Timeout exceeded", "elapsed_seconds": elapsed},
|
||||||
|
completed_at=datetime.now(UTC),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await db.commit()
|
||||||
|
return {
|
||||||
|
"status": "failed",
|
||||||
|
"error": f"Timeout exceeded after {elapsed:.1f}s",
|
||||||
|
"subtask_id": str(subtask_id),
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentSubtask)
|
||||||
|
.where(AgentSubtask.id == subtask_id)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
subtask = result.scalar_one_or_none()
|
||||||
|
|
||||||
|
if subtask is None:
|
||||||
|
return {"status": "error", "error": "Subtask not found"}
|
||||||
|
|
||||||
|
if subtask.status in ("completed", "failed", "cancelled"):
|
||||||
|
return {
|
||||||
|
"status": subtask.status,
|
||||||
|
"result": subtask.result or {},
|
||||||
|
"subtask_id": str(subtask.id),
|
||||||
|
"completed_at": subtask.completed_at.isoformat() if subtask.completed_at else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
await asyncio.sleep(poll_interval)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def aggregate_results(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
subtask_ids: list[uuid.UUID],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Aggregate results from multiple subtasks.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session
|
||||||
|
tenant_id: Tenant ID
|
||||||
|
subtask_ids: List of subtask UUIDs to aggregate
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with summary of all subtask results
|
||||||
|
"""
|
||||||
|
results = []
|
||||||
|
for sid in subtask_ids:
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentSubtask)
|
||||||
|
.where(AgentSubtask.id == sid)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
subtask = result.scalar_one_or_none()
|
||||||
|
if subtask:
|
||||||
|
results.append({
|
||||||
|
"subtask_id": str(subtask.id),
|
||||||
|
"parent_agent_id": str(subtask.parent_agent_id),
|
||||||
|
"child_agent_id": str(subtask.child_agent_id),
|
||||||
|
"task_description": subtask.task_description,
|
||||||
|
"status": subtask.status,
|
||||||
|
"result": subtask.result or {},
|
||||||
|
"created_at": subtask.created_at.isoformat() if subtask.created_at else None,
|
||||||
|
"completed_at": subtask.completed_at.isoformat() if subtask.completed_at else None,
|
||||||
|
})
|
||||||
|
|
||||||
|
completed = [r for r in results if r["status"] == "completed"]
|
||||||
|
failed = [r for r in results if r["status"] == "failed"]
|
||||||
|
pending = [r for r in results if r["status"] == "pending"]
|
||||||
|
cancelled = [r for r in results if r["status"] == "cancelled"]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"total": len(results),
|
||||||
|
"completed": len(completed),
|
||||||
|
"failed": len(failed),
|
||||||
|
"pending": len(pending),
|
||||||
|
"cancelled": len(cancelled),
|
||||||
|
"results": results,
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def cancel_subtask(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
subtask_id: uuid.UUID,
|
||||||
|
) -> bool:
|
||||||
|
"""Cancel a pending or running subtask.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session
|
||||||
|
tenant_id: Tenant ID
|
||||||
|
subtask_id: UUID of the subtask to cancel
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if cancelled, False if not found or already terminal
|
||||||
|
"""
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentSubtask)
|
||||||
|
.where(AgentSubtask.id == subtask_id)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
subtask = result.scalar_one_or_none()
|
||||||
|
|
||||||
|
if subtask is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if subtask.status in ("completed", "failed", "cancelled"):
|
||||||
|
logger.warning(
|
||||||
|
"Cannot cancel subtask %s: already in terminal state '%s'",
|
||||||
|
subtask_id, subtask.status,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
subtask.status = "cancelled"
|
||||||
|
subtask.completed_at = datetime.now(UTC)
|
||||||
|
subtask.result = {"cancelled_by": "coordinator", "previous_status": subtask.status}
|
||||||
|
await db.flush()
|
||||||
|
|
||||||
|
logger.info("Subtask %s cancelled (was %s)", subtask_id, subtask.status)
|
||||||
|
return True
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def list_subtasks(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
parent_agent_id: uuid.UUID | None = None,
|
||||||
|
child_agent_id: uuid.UUID | None = None,
|
||||||
|
status: str | None = None,
|
||||||
|
limit: int = 50,
|
||||||
|
offset: int = 0,
|
||||||
|
) -> tuple[list[AgentSubtask], int]:
|
||||||
|
"""List subtasks with optional filters."""
|
||||||
|
from sqlalchemy import func
|
||||||
|
|
||||||
|
query = select(AgentSubtask).where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
count_query = (
|
||||||
|
select(func.count())
|
||||||
|
.select_from(AgentSubtask)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
)
|
||||||
|
|
||||||
|
if parent_agent_id is not None:
|
||||||
|
query = query.where(AgentSubtask.parent_agent_id == parent_agent_id)
|
||||||
|
count_query = count_query.where(AgentSubtask.parent_agent_id == parent_agent_id)
|
||||||
|
if child_agent_id is not None:
|
||||||
|
query = query.where(AgentSubtask.child_agent_id == child_agent_id)
|
||||||
|
count_query = count_query.where(AgentSubtask.child_agent_id == child_agent_id)
|
||||||
|
if status is not None:
|
||||||
|
query = query.where(AgentSubtask.status == status)
|
||||||
|
count_query = count_query.where(AgentSubtask.status == status)
|
||||||
|
|
||||||
|
count_result = await db.execute(count_query)
|
||||||
|
total = count_result.scalar() or 0
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
query.order_by(AgentSubtask.created_at.desc())
|
||||||
|
.limit(limit)
|
||||||
|
.offset(offset)
|
||||||
|
)
|
||||||
|
return list(result.scalars().all()), total
|
||||||
|
|
||||||
|
|
||||||
|
def register_agent_coordinator_tools():
|
||||||
|
"""Register AgentCoordinator tools in the global tool registry."""
|
||||||
|
from app.plugins.builtins.ai_assistant.tool_registry import get_tool_registry
|
||||||
|
|
||||||
|
registry = get_tool_registry()
|
||||||
|
|
||||||
|
async def create_subtask_handler(arguments: dict[str, Any], context: dict[str, Any]) -> str:
|
||||||
|
"""Handle create_subtask tool call from an AI agent."""
|
||||||
|
from app.core.db import get_session_factory
|
||||||
|
|
||||||
|
parent_agent_id = arguments.get("parent_agent_id", "")
|
||||||
|
child_agent_name = arguments.get("child_agent_name", "")
|
||||||
|
task_description = arguments.get("task_description", "")
|
||||||
|
tenant_id_str = context.get("tenant_id", "")
|
||||||
|
|
||||||
|
if not child_agent_name or not task_description:
|
||||||
|
return json.dumps({"status": "error", "error": "Missing child_agent_name or task_description"})
|
||||||
|
|
||||||
|
try:
|
||||||
|
tenant_id = uuid.UUID(tenant_id_str) if tenant_id_str else uuid.uuid4()
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
return json.dumps({"status": "error", "error": "Invalid tenant_id"})
|
||||||
|
|
||||||
|
factory = get_session_factory()
|
||||||
|
async with factory() as db:
|
||||||
|
# Find child agent by name
|
||||||
|
from app.plugins.builtins.automation.models import AgentDefinition
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentDefinition)
|
||||||
|
.where(AgentDefinition.name == child_agent_name)
|
||||||
|
.where(AgentDefinition.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
child_agent = result.scalar_one_or_none()
|
||||||
|
|
||||||
|
if child_agent is None:
|
||||||
|
return json.dumps({"status": "error", "error": f"Child agent '{child_agent_name}' not found"})
|
||||||
|
|
||||||
|
if not child_agent.is_active:
|
||||||
|
return json.dumps({"status": "error", "error": f"Child agent '{child_agent_name}' is inactive"})
|
||||||
|
|
||||||
|
try:
|
||||||
|
parent_id = uuid.UUID(parent_agent_id) if parent_agent_id else uuid.uuid4()
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
return json.dumps({"status": "error", "error": "Invalid parent_agent_id"})
|
||||||
|
|
||||||
|
subtask = await AgentCoordinator.create_subtask(
|
||||||
|
db=db,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
parent_agent_id=parent_id,
|
||||||
|
child_agent_id=child_agent.id,
|
||||||
|
task_description=task_description,
|
||||||
|
)
|
||||||
|
await db.commit()
|
||||||
|
|
||||||
|
return json.dumps({
|
||||||
|
"status": "created",
|
||||||
|
"subtask_id": str(subtask.id),
|
||||||
|
"child_agent": child_agent_name,
|
||||||
|
"child_agent_id": str(child_agent.id),
|
||||||
|
"task_description": task_description,
|
||||||
|
})
|
||||||
|
|
||||||
|
registry.register(
|
||||||
|
name="create_subtask",
|
||||||
|
description="Create a subtask for another agent to execute. The child agent will be activated with the task description.",
|
||||||
|
parameters={
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"parent_agent_id": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "ID of the parent agent creating the subtask",
|
||||||
|
},
|
||||||
|
"child_agent_name": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "Name of the child agent that will execute the subtask",
|
||||||
|
},
|
||||||
|
"task_description": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "Description of the task to execute",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["child_agent_name", "task_description"],
|
||||||
|
},
|
||||||
|
handler=create_subtask_handler,
|
||||||
|
plugin_name="automation",
|
||||||
|
required_permission="agents:execute",
|
||||||
|
category="orchestration",
|
||||||
|
)
|
||||||
|
logger.info("Agent coordinator tool 'create_subtask' registered")
|
||||||
|
|
||||||
|
|
||||||
|
def unregister_agent_coordinator_tools():
|
||||||
|
"""Unregister AgentCoordinator tools."""
|
||||||
|
from app.plugins.builtins.ai_assistant.tool_registry import get_tool_registry
|
||||||
|
|
||||||
|
registry = get_tool_registry()
|
||||||
|
registry.unregister("create_subtask")
|
||||||
|
logger.info("Agent coordinator tools unregistered")
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
-- Agent Subtasks for Multi-Agent Orchestration (Phase 5.8)
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS agent_subtasks (
|
||||||
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
tenant_id UUID NOT NULL,
|
||||||
|
parent_agent_id UUID NOT NULL REFERENCES automation_agent_definitions(id) ON DELETE CASCADE,
|
||||||
|
child_agent_id UUID NOT NULL REFERENCES automation_agent_definitions(id) ON DELETE CASCADE,
|
||||||
|
task_description TEXT NOT NULL,
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'pending', -- pending, running, completed, failed, cancelled
|
||||||
|
result JSONB NOT NULL DEFAULT '{}',
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
completed_at TIMESTAMPTZ,
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_subtasks_parent ON agent_subtasks (tenant_id, parent_agent_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_subtasks_child ON agent_subtasks (tenant_id, child_agent_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_agent_subtasks_status ON agent_subtasks (tenant_id, status);
|
||||||
@@ -251,3 +251,40 @@ class AutomationRun(Base, TenantMixin):
|
|||||||
JSONB, nullable=False, default=dict
|
JSONB, nullable=False, default=dict
|
||||||
)
|
)
|
||||||
dry_run: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
dry_run: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentSubtask(Base, TenantMixin):
|
||||||
|
"""A subtask delegated from one agent to another for multi-agent orchestration."""
|
||||||
|
|
||||||
|
__tablename__ = "agent_subtasks"
|
||||||
|
__table_args__ = (
|
||||||
|
Index("ix_agent_subtasks_parent", "tenant_id", "parent_agent_id"),
|
||||||
|
Index("ix_agent_subtasks_child", "tenant_id", "child_agent_id"),
|
||||||
|
Index("ix_agent_subtasks_status", "tenant_id", "status"),
|
||||||
|
)
|
||||||
|
|
||||||
|
id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
|
||||||
|
)
|
||||||
|
parent_agent_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True),
|
||||||
|
ForeignKey("automation_agent_definitions.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
child_agent_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True),
|
||||||
|
ForeignKey("automation_agent_definitions.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
task_description: Mapped[str] = mapped_column(Text, nullable=False)
|
||||||
|
status: Mapped[str] = mapped_column(
|
||||||
|
String(20), nullable=False, default="pending"
|
||||||
|
)
|
||||||
|
result: Mapped[dict[str, Any]] = mapped_column(
|
||||||
|
JSONB, nullable=False, default=dict
|
||||||
|
)
|
||||||
|
completed_at: Mapped[datetime | None] = mapped_column(
|
||||||
|
DateTime(timezone=True), nullable=True
|
||||||
|
)
|
||||||
|
|||||||
@@ -57,7 +57,7 @@ class AutomationPlugin(BasePlugin):
|
|||||||
"mail.received",
|
"mail.received",
|
||||||
"workflow.timeout",
|
"workflow.timeout",
|
||||||
],
|
],
|
||||||
migrations=["0001_initial.sql"],
|
migrations=["0001_initial.sql", "0002_agent_subtasks.sql"],
|
||||||
permissions=[
|
permissions=[
|
||||||
"automation:read",
|
"automation:read",
|
||||||
"automation:write",
|
"automation:write",
|
||||||
@@ -185,6 +185,12 @@ class AutomationPlugin(BasePlugin):
|
|||||||
register_agent_comm_tool()
|
register_agent_comm_tool()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to register agent communication tool")
|
logger.exception("Failed to register agent communication tool")
|
||||||
|
# Register agent coordinator tools
|
||||||
|
try:
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import register_agent_coordinator_tools
|
||||||
|
register_agent_coordinator_tools()
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to register agent coordinator tools")
|
||||||
# Register MiniApps from manifest
|
# Register MiniApps from manifest
|
||||||
try:
|
try:
|
||||||
from app.plugins.builtins.kommunikation.contracts import MiniAppRegistry
|
from app.plugins.builtins.kommunikation.contracts import MiniAppRegistry
|
||||||
@@ -222,6 +228,12 @@ class AutomationPlugin(BasePlugin):
|
|||||||
unregister_agent_comm_tool()
|
unregister_agent_comm_tool()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to unregister agent communication tool")
|
logger.exception("Failed to unregister agent communication tool")
|
||||||
|
# Unregister agent coordinator tools
|
||||||
|
try:
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import unregister_agent_coordinator_tools
|
||||||
|
unregister_agent_coordinator_tools()
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to unregister agent coordinator tools")
|
||||||
# Unregister MiniApps
|
# Unregister MiniApps
|
||||||
try:
|
try:
|
||||||
from app.plugins.builtins.kommunikation.contracts import MiniAppRegistry
|
from app.plugins.builtins.kommunikation.contracts import MiniAppRegistry
|
||||||
|
|||||||
@@ -624,3 +624,241 @@ async def restore_automation_version(
|
|||||||
return _automation_to_response(automation)
|
return _automation_to_response(automation)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Subtask Endpoints (Phase 5.8) ───
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"/subtasks",
|
||||||
|
dependencies=[Depends(require_permission("agents:read"))],
|
||||||
|
response_model=SubtaskListResponse,
|
||||||
|
)
|
||||||
|
async def list_subtasks(
|
||||||
|
parent_agent_id: str | None = Query(None),
|
||||||
|
child_agent_id: str | None = Query(None),
|
||||||
|
status: str | None = Query(None),
|
||||||
|
limit: int = Query(50, ge=1, le=200),
|
||||||
|
offset: int = Query(0, ge=0),
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""List subtasks with optional filters."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
parent_id = uuid.UUID(parent_agent_id) if parent_agent_id else None
|
||||||
|
child_id = uuid.UUID(child_agent_id) if child_agent_id else None
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid agent ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import AgentCoordinator
|
||||||
|
|
||||||
|
items, total = await AgentCoordinator.list_subtasks(
|
||||||
|
db, tenant_id,
|
||||||
|
parent_agent_id=parent_id,
|
||||||
|
child_agent_id=child_id,
|
||||||
|
status=status,
|
||||||
|
limit=limit,
|
||||||
|
offset=offset,
|
||||||
|
)
|
||||||
|
return SubtaskListResponse(
|
||||||
|
items=[_subtask_to_response(s) for s in items],
|
||||||
|
total=total,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/subtasks",
|
||||||
|
dependencies=[Depends(require_permission("agents:execute"))],
|
||||||
|
response_model=SubtaskRead,
|
||||||
|
status_code=201,
|
||||||
|
)
|
||||||
|
async def create_subtask(
|
||||||
|
data: SubtaskCreate,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Create a new subtask for multi-agent orchestration."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
parent_id = uuid.UUID(data.parent_agent_id)
|
||||||
|
child_id = uuid.UUID(data.child_agent_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid agent ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import AgentCoordinator
|
||||||
|
|
||||||
|
subtask = await AgentCoordinator.create_subtask(
|
||||||
|
db=db,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
parent_agent_id=parent_id,
|
||||||
|
child_agent_id=child_id,
|
||||||
|
task_description=data.task_description,
|
||||||
|
)
|
||||||
|
await db.commit()
|
||||||
|
return _subtask_to_response(subtask)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"/subtasks/{subtask_id}",
|
||||||
|
dependencies=[Depends(require_permission("agents:read"))],
|
||||||
|
response_model=SubtaskRead,
|
||||||
|
)
|
||||||
|
async def get_subtask(
|
||||||
|
subtask_id: str,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Get a single subtask by ID."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
sid = uuid.UUID(subtask_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid subtask ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.models import AgentSubtask
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentSubtask)
|
||||||
|
.where(AgentSubtask.id == sid)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
subtask = result.scalar_one_or_none()
|
||||||
|
if subtask is None:
|
||||||
|
raise HTTPException(status_code=404, detail="Subtask not found")
|
||||||
|
return _subtask_to_response(subtask)
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch(
|
||||||
|
"/subtasks/{subtask_id}",
|
||||||
|
dependencies=[Depends(require_permission("agents:execute"))],
|
||||||
|
response_model=SubtaskRead,
|
||||||
|
)
|
||||||
|
async def update_subtask(
|
||||||
|
subtask_id: str,
|
||||||
|
data: SubtaskUpdate,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Update a subtask (status, result)."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
sid = uuid.UUID(subtask_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid subtask ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.models import AgentSubtask
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
result = await db.execute(
|
||||||
|
select(AgentSubtask)
|
||||||
|
.where(AgentSubtask.id == sid)
|
||||||
|
.where(AgentSubtask.tenant_id == tenant_id)
|
||||||
|
.limit(1)
|
||||||
|
)
|
||||||
|
subtask = result.scalar_one_or_none()
|
||||||
|
if subtask is None:
|
||||||
|
raise HTTPException(status_code=404, detail="Subtask not found")
|
||||||
|
|
||||||
|
update_data = data.model_dump(exclude_none=True)
|
||||||
|
if "status" in update_data:
|
||||||
|
subtask.status = update_data["status"]
|
||||||
|
if update_data["status"] in ("completed", "failed", "cancelled"):
|
||||||
|
subtask.completed_at = __import__("datetime").datetime.now(
|
||||||
|
__import__("datetime").timezone.utc
|
||||||
|
)
|
||||||
|
if "result" in update_data:
|
||||||
|
subtask.result = update_data["result"]
|
||||||
|
|
||||||
|
await db.commit()
|
||||||
|
await db.refresh(subtask)
|
||||||
|
return _subtask_to_response(subtask)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/subtasks/{subtask_id}/cancel",
|
||||||
|
dependencies=[Depends(require_permission("agents:execute"))],
|
||||||
|
)
|
||||||
|
async def cancel_subtask(
|
||||||
|
subtask_id: str,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Cancel a pending or running subtask."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
sid = uuid.UUID(subtask_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid subtask ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import AgentCoordinator
|
||||||
|
|
||||||
|
success = await AgentCoordinator.cancel_subtask(db, tenant_id, sid)
|
||||||
|
if not success:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Subtask not found or already in terminal state",
|
||||||
|
)
|
||||||
|
await db.commit()
|
||||||
|
return {"status": "cancelled", "subtask_id": subtask_id}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/subtasks/{subtask_id}/wait",
|
||||||
|
dependencies=[Depends(require_permission("agents:execute"))],
|
||||||
|
)
|
||||||
|
async def wait_for_subtask(
|
||||||
|
subtask_id: str,
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Wait for a subtask to complete (polls until terminal state)."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
sid = uuid.UUID(subtask_id)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid subtask ID")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import AgentCoordinator
|
||||||
|
|
||||||
|
result = await AgentCoordinator.wait_for_subtask(db, tenant_id, sid)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/subtasks/aggregate",
|
||||||
|
dependencies=[Depends(require_permission("agents:read"))],
|
||||||
|
)
|
||||||
|
async def aggregate_subtasks(
|
||||||
|
subtask_ids: list[str],
|
||||||
|
current_user: dict[str, Any] = Depends(get_current_user),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
"""Aggregate results from multiple subtasks."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
try:
|
||||||
|
ids = [uuid.UUID(sid) for sid in subtask_ids]
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(status_code=400, detail="Invalid subtask ID in list")
|
||||||
|
|
||||||
|
from app.plugins.builtins.automation.agent_coordinator import AgentCoordinator
|
||||||
|
|
||||||
|
result = await AgentCoordinator.aggregate_results(db, tenant_id, ids)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _subtask_to_response(s) -> SubtaskRead:
|
||||||
|
"""Convert AgentSubtask model to response schema."""
|
||||||
|
return SubtaskRead(
|
||||||
|
id=str(s.id),
|
||||||
|
parent_agent_id=str(s.parent_agent_id),
|
||||||
|
child_agent_id=str(s.child_agent_id),
|
||||||
|
task_description=s.task_description,
|
||||||
|
status=s.status,
|
||||||
|
result=s.result or {},
|
||||||
|
created_at=s.created_at.isoformat() if s.created_at else None,
|
||||||
|
completed_at=s.completed_at.isoformat() if s.completed_at else None,
|
||||||
|
updated_at=s.updated_at.isoformat() if s.updated_at else None,
|
||||||
|
)
|
||||||
|
|||||||
@@ -289,3 +289,42 @@ class AutomationVersionListResponse(BaseModel):
|
|||||||
|
|
||||||
items: list[AutomationVersionResponse]
|
items: list[AutomationVersionResponse]
|
||||||
total: int
|
total: int
|
||||||
|
|
||||||
|
|
||||||
|
# ─── Subtask Schemas (Phase 5.8) ───
|
||||||
|
|
||||||
|
|
||||||
|
class SubtaskCreate(BaseModel):
|
||||||
|
"""Create a new subtask for multi-agent orchestration."""
|
||||||
|
|
||||||
|
parent_agent_id: str = Field(..., description="UUID of the parent agent")
|
||||||
|
child_agent_id: str = Field(..., description="UUID of the child agent that will execute")
|
||||||
|
task_description: str = Field(..., min_length=1, description="Description of the task")
|
||||||
|
|
||||||
|
|
||||||
|
class SubtaskUpdate(BaseModel):
|
||||||
|
"""Update a subtask (status, result)."""
|
||||||
|
|
||||||
|
status: str | None = Field(None, pattern="^(pending|running|completed|failed|cancelled)$")
|
||||||
|
result: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class SubtaskRead(BaseModel):
|
||||||
|
"""Subtask response."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
parent_agent_id: str
|
||||||
|
child_agent_id: str
|
||||||
|
task_description: str
|
||||||
|
status: str = "pending"
|
||||||
|
result: dict[str, Any] = {}
|
||||||
|
created_at: str | None = None
|
||||||
|
completed_at: str | None = None
|
||||||
|
updated_at: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class SubtaskListResponse(BaseModel):
|
||||||
|
"""Paginated subtask list."""
|
||||||
|
|
||||||
|
items: list[SubtaskRead]
|
||||||
|
total: int
|
||||||
|
|||||||
@@ -0,0 +1,5 @@
|
|||||||
|
"""GraphRAG plugin package."""
|
||||||
|
|
||||||
|
from app.plugins.builtins.graph_rag.plugin import GraphRAGPlugin
|
||||||
|
|
||||||
|
__all__ = ["GraphRAGPlugin"]
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
-- GraphRAG plugin initial migration: creates entity_relationships table
|
||||||
|
CREATE TABLE IF NOT EXISTS entity_relationships (
|
||||||
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
tenant_id UUID NOT NULL,
|
||||||
|
source_type VARCHAR(50) NOT NULL,
|
||||||
|
source_id UUID NOT NULL,
|
||||||
|
target_type VARCHAR(50) NOT NULL,
|
||||||
|
target_id UUID NOT NULL,
|
||||||
|
relationship_type VARCHAR(50) NOT NULL,
|
||||||
|
metadata JSONB DEFAULT '{}'::jsonb,
|
||||||
|
owner_id UUID REFERENCES users(id) ON DELETE SET NULL,
|
||||||
|
deleted_at TIMESTAMPTZ,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_entity_relationships_source ON entity_relationships(tenant_id, source_type, source_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_entity_relationships_target ON entity_relationships(tenant_id, target_type, target_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_entity_relationships_type ON entity_relationships(tenant_id, relationship_type);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_entity_relationships_source_target ON entity_relationships(tenant_id, source_type, source_id, target_type, target_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_entity_relationships_owner ON entity_relationships(owner_id);
|
||||||
@@ -0,0 +1,52 @@
|
|||||||
|
"""EntityRelationship model for GraphRAG — stores relationships between entities."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import Index, String
|
||||||
|
from sqlalchemy.dialects.postgresql import JSONB
|
||||||
|
from sqlalchemy.dialects.postgresql import UUID as PGUUID
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from app.core.db import Base, TenantMixin
|
||||||
|
from app.models.owned_mixin import OwnedMixin
|
||||||
|
|
||||||
|
|
||||||
|
class EntityRelationship(Base, TenantMixin, OwnedMixin):
|
||||||
|
"""Directed relationship between two entities in the knowledge graph.
|
||||||
|
|
||||||
|
Stores typed relationships (e.g., 'works_for', 'has_email', 'related_to')
|
||||||
|
between any two entities in the system. Used for GraphRAG traversal.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__tablename__ = "entity_relationships"
|
||||||
|
__table_args__ = (
|
||||||
|
Index("ix_entity_relationships_source", "tenant_id", "source_type", "source_id"),
|
||||||
|
Index("ix_entity_relationships_target", "tenant_id", "target_type", "target_id"),
|
||||||
|
Index("ix_entity_relationships_type", "tenant_id", "relationship_type"),
|
||||||
|
Index("ix_entity_relationships_source_target", "tenant_id", "source_type", "source_id", "target_type", "target_id"),
|
||||||
|
)
|
||||||
|
|
||||||
|
id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
|
||||||
|
)
|
||||||
|
source_type: Mapped[str] = mapped_column(
|
||||||
|
String(50), nullable=False, comment="Entity type of the source (e.g. 'contact', 'company', 'file')"
|
||||||
|
)
|
||||||
|
source_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), nullable=False, comment="UUID of the source entity"
|
||||||
|
)
|
||||||
|
target_type: Mapped[str] = mapped_column(
|
||||||
|
String(50), nullable=False, comment="Entity type of the target (e.g. 'contact', 'email', 'task')"
|
||||||
|
)
|
||||||
|
target_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), nullable=False, comment="UUID of the target entity"
|
||||||
|
)
|
||||||
|
relationship_type: Mapped[str] = mapped_column(
|
||||||
|
String(50), nullable=False, comment="Type of relationship (e.g. 'works_for', 'has_email', 'related_to')"
|
||||||
|
)
|
||||||
|
metadata: Mapped[dict[str, Any] | None] = mapped_column(
|
||||||
|
JSONB, nullable=True, default=dict, comment="Arbitrary metadata about the relationship"
|
||||||
|
)
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
"""GraphRAG plugin — entity relationship graph with BFS traversal and search provider."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from app.plugins.base import BasePlugin
|
||||||
|
from app.plugins.manifest import PluginManifest, PluginRouteDef
|
||||||
|
|
||||||
|
|
||||||
|
class GraphRAGPlugin(BasePlugin):
|
||||||
|
"""GraphRAG plugin for entity relationship graph with BFS traversal."""
|
||||||
|
|
||||||
|
manifest = PluginManifest(
|
||||||
|
name="graph_rag",
|
||||||
|
version="1.0.0",
|
||||||
|
display_name="GraphRAG",
|
||||||
|
description="Entity relationship graph with BFS traversal and unified search integration. Stores typed relationships between entities for knowledge graph exploration.",
|
||||||
|
dependencies=["unified_search"],
|
||||||
|
routes=[
|
||||||
|
PluginRouteDef(
|
||||||
|
path="/api/v1/graph",
|
||||||
|
module="app.plugins.builtins.graph_rag.routes",
|
||||||
|
router_attr="router",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
events=[],
|
||||||
|
migrations=["0001_initial.sql"],
|
||||||
|
permissions=[
|
||||||
|
"graph:read",
|
||||||
|
"graph:write",
|
||||||
|
],
|
||||||
|
is_core=True,
|
||||||
|
author="LeoCRM Team",
|
||||||
|
min_app_version="1.0.0",
|
||||||
|
contract_version="1.0.0",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def on_activate(self, db, service_container, event_bus) -> None:
|
||||||
|
"""Activate plugin: register GraphRAG search provider."""
|
||||||
|
from app.plugins.builtins.graph_rag.provider import GraphRAGSearchProvider
|
||||||
|
from app.plugins.builtins.unified_search.provider_registry import get_search_registry
|
||||||
|
|
||||||
|
registry = get_search_registry()
|
||||||
|
try:
|
||||||
|
registry.register(GraphRAGSearchProvider())
|
||||||
|
except Exception:
|
||||||
|
import logging
|
||||||
|
logging.getLogger(__name__).exception("Failed to register GraphRAGSearchProvider")
|
||||||
|
|
||||||
|
await super().on_activate(db, service_container, event_bus)
|
||||||
|
|
||||||
|
async def on_deactivate(self, db, service_container, event_bus) -> None:
|
||||||
|
"""Deactivate plugin: unregister search provider and contract."""
|
||||||
|
from app.plugins.builtins.unified_search.provider_registry import get_search_registry
|
||||||
|
get_search_registry().unregister("graph_relationship")
|
||||||
|
|
||||||
|
from app.plugins.builtins.contracts import get_contract_registry
|
||||||
|
get_contract_registry().unregister(self.manifest.name)
|
||||||
|
await super().on_deactivate(db, service_container, event_bus)
|
||||||
@@ -0,0 +1,165 @@
|
|||||||
|
"""GraphRAGSearchProvider — registers GraphRAG as a search provider in unified_search."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import text
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.plugins.builtins.unified_search.base_provider import BaseSearchProvider
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class GraphRAGSearchProvider(BaseSearchProvider):
|
||||||
|
"""Search provider for GraphRAG entity relationships.
|
||||||
|
|
||||||
|
Enables full-text and semantic search over relationship metadata
|
||||||
|
and entity types in the knowledge graph.
|
||||||
|
"""
|
||||||
|
|
||||||
|
entity_type = "graph_relationship"
|
||||||
|
|
||||||
|
async def _search_fts_filtered(
|
||||||
|
self,
|
||||||
|
db: AsyncSession,
|
||||||
|
tsquery: str,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
limit: int,
|
||||||
|
visible_ids: set[uuid.UUID] | None,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Full-text search on relationship metadata and types."""
|
||||||
|
if visible_ids is not None:
|
||||||
|
sql = text(
|
||||||
|
"""
|
||||||
|
SELECT r.*, ts_rank(
|
||||||
|
to_tsvector('pg_catalog.german',
|
||||||
|
coalesce(r.relationship_type, '') || ' ' ||
|
||||||
|
coalesce(r.source_type, '') || ' ' ||
|
||||||
|
coalesce(r.target_type, '') || ' ' ||
|
||||||
|
coalesce(r.metadata::text, '')
|
||||||
|
),
|
||||||
|
to_tsquery('pg_catalog.german', :q)
|
||||||
|
) AS rank
|
||||||
|
FROM entity_relationships r
|
||||||
|
WHERE r.tenant_id = :tid
|
||||||
|
AND r.deleted_at IS NULL
|
||||||
|
AND to_tsvector('pg_catalog.german',
|
||||||
|
coalesce(r.relationship_type, '') || ' ' ||
|
||||||
|
coalesce(r.source_type, '') || ' ' ||
|
||||||
|
coalesce(r.target_type, '') || ' ' ||
|
||||||
|
coalesce(r.metadata::text, '')
|
||||||
|
) @@ to_tsquery('pg_catalog.german', :q)
|
||||||
|
AND r.id = ANY(:visible_ids)
|
||||||
|
ORDER BY rank DESC
|
||||||
|
LIMIT :lim
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
result = await db.execute(
|
||||||
|
sql,
|
||||||
|
{
|
||||||
|
"q": tsquery,
|
||||||
|
"tid": tenant_id,
|
||||||
|
"lim": limit,
|
||||||
|
"visible_ids": list(visible_ids),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sql = text(
|
||||||
|
"""
|
||||||
|
SELECT r.*, ts_rank(
|
||||||
|
to_tsvector('pg_catalog.german',
|
||||||
|
coalesce(r.relationship_type, '') || ' ' ||
|
||||||
|
coalesce(r.source_type, '') || ' ' ||
|
||||||
|
coalesce(r.target_type, '') || ' ' ||
|
||||||
|
coalesce(r.metadata::text, '')
|
||||||
|
),
|
||||||
|
to_tsquery('pg_catalog.german', :q)
|
||||||
|
) AS rank
|
||||||
|
FROM entity_relationships r
|
||||||
|
WHERE r.tenant_id = :tid
|
||||||
|
AND r.deleted_at IS NULL
|
||||||
|
AND to_tsvector('pg_catalog.german',
|
||||||
|
coalesce(r.relationship_type, '') || ' ' ||
|
||||||
|
coalesce(r.source_type, '') || ' ' ||
|
||||||
|
coalesce(r.target_type, '') || ' ' ||
|
||||||
|
coalesce(r.metadata::text, '')
|
||||||
|
) @@ to_tsquery('pg_catalog.german', :q)
|
||||||
|
ORDER BY rank DESC
|
||||||
|
LIMIT :lim
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
result = await db.execute(
|
||||||
|
sql,
|
||||||
|
{"q": tsquery, "tid": tenant_id, "lim": limit},
|
||||||
|
)
|
||||||
|
rows = result.mappings().all()
|
||||||
|
return [dict(r) for r in rows]
|
||||||
|
|
||||||
|
async def _search_vector_filtered(
|
||||||
|
self,
|
||||||
|
db: AsyncSession,
|
||||||
|
embedding: list[float],
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
limit: int,
|
||||||
|
visible_ids: set[uuid.UUID] | None,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Semantic search is not yet supported for graph relationships.
|
||||||
|
|
||||||
|
Returns empty list — relationships are searched via FTS on metadata.
|
||||||
|
"""
|
||||||
|
return []
|
||||||
|
|
||||||
|
async def get_embedding_text(
|
||||||
|
self, db: AsyncSession, entity_id: uuid.UUID, tenant_id: uuid.UUID
|
||||||
|
) -> str:
|
||||||
|
"""Get text for embedding generation."""
|
||||||
|
sql = text(
|
||||||
|
"""
|
||||||
|
SELECT relationship_type, source_type, source_id, target_type, target_id, metadata
|
||||||
|
FROM entity_relationships
|
||||||
|
WHERE id = :eid AND tenant_id = :tid
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
result = await db.execute(sql, {"eid": entity_id, "tid": tenant_id})
|
||||||
|
row = result.mappings().first()
|
||||||
|
if not row:
|
||||||
|
return ""
|
||||||
|
parts = [
|
||||||
|
row.get("relationship_type", ""),
|
||||||
|
row.get("source_type", ""),
|
||||||
|
str(row.get("source_id", "")),
|
||||||
|
row.get("target_type", ""),
|
||||||
|
str(row.get("target_id", "")),
|
||||||
|
str(row.get("metadata", {})),
|
||||||
|
]
|
||||||
|
return " ".join(str(p) for p in parts if p)
|
||||||
|
|
||||||
|
def to_search_result(self, entity: object) -> dict[str, Any]:
|
||||||
|
"""Convert relationship to search result dict."""
|
||||||
|
if isinstance(entity, dict):
|
||||||
|
return {
|
||||||
|
"entity_type": self.entity_type,
|
||||||
|
"entity_id": str(entity.get("id", "")),
|
||||||
|
"title": f"{entity.get('source_type', '?')} --[{entity.get('relationship_type', '?')}]--> {entity.get('target_type', '?')}",
|
||||||
|
"snippet": str(entity.get("metadata", {})),
|
||||||
|
"score": float(entity.get("rank", 0.0)),
|
||||||
|
"data": {
|
||||||
|
"source_type": entity.get("source_type"),
|
||||||
|
"source_id": str(entity.get("source_id", "")),
|
||||||
|
"target_type": entity.get("target_type"),
|
||||||
|
"target_id": str(entity.get("target_id", "")),
|
||||||
|
"relationship_type": entity.get("relationship_type"),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return {
|
||||||
|
"entity_type": self.entity_type,
|
||||||
|
"entity_id": str(getattr(entity, "id", "")),
|
||||||
|
"title": f"{getattr(entity, 'source_type', '?')} --[{getattr(entity, 'relationship_type', '?')}]--> {getattr(entity, 'target_type', '?')}",
|
||||||
|
"snippet": str(getattr(entity, "metadata", {})),
|
||||||
|
"score": 0.0,
|
||||||
|
"data": {},
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
"""GraphRAG plugin routes — relationship CRUD and BFS graph traversal."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.core.db import get_db
|
||||||
|
from app.deps import get_current_user, require_permission
|
||||||
|
from app.plugins.builtins.graph_rag.models import EntityRelationship
|
||||||
|
from app.plugins.builtins.graph_rag.schemas import (
|
||||||
|
RelationshipCreate,
|
||||||
|
TraverseRequest,
|
||||||
|
)
|
||||||
|
from app.plugins.builtins.graph_rag.services import (
|
||||||
|
create_relationship,
|
||||||
|
delete_relationship,
|
||||||
|
traverse_graph,
|
||||||
|
)
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api/v1/graph", tags=["graph-rag"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/relationships", status_code=status.HTTP_201_CREATED, dependencies=[Depends(require_permission("graph:write"))])
|
||||||
|
async def create_relationship_route(
|
||||||
|
body: RelationshipCreate,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""Create a new relationship between two entities."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
user_id = uuid.UUID(current_user["user_id"])
|
||||||
|
source_id = _parse_uuid(body.source_id, "source_id")
|
||||||
|
target_id = _parse_uuid(body.target_id, "target_id")
|
||||||
|
|
||||||
|
result = await create_relationship(
|
||||||
|
db=db,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
source_type=body.source_type,
|
||||||
|
source_id=source_id,
|
||||||
|
target_type=body.target_type,
|
||||||
|
target_id=target_id,
|
||||||
|
relationship_type=body.relationship_type,
|
||||||
|
metadata=body.metadata,
|
||||||
|
owner_id=user_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if "error" in result:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_409_CONFLICT,
|
||||||
|
detail={"detail": result["error"], "code": result.get("code", "duplicate")},
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/relationships", dependencies=[Depends(require_permission("graph:read"))])
|
||||||
|
async def list_relationships(
|
||||||
|
source_type: str | None = Query(None, description="Filter by source entity type"),
|
||||||
|
source_id: str | None = Query(None, description="Filter by source entity UUID"),
|
||||||
|
target_type: str | None = Query(None, description="Filter by target entity type"),
|
||||||
|
target_id: str | None = Query(None, description="Filter by target entity UUID"),
|
||||||
|
relationship_type: str | None = Query(None, description="Filter by relationship type"),
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(50, ge=1, le=200),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""List relationships with optional filters."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
|
||||||
|
base = select(EntityRelationship).where(
|
||||||
|
EntityRelationship.tenant_id == tenant_id,
|
||||||
|
EntityRelationship.deleted_at.is_(None),
|
||||||
|
)
|
||||||
|
|
||||||
|
if source_type:
|
||||||
|
base = base.where(EntityRelationship.source_type == source_type)
|
||||||
|
if source_id:
|
||||||
|
base = base.where(EntityRelationship.source_id == _parse_uuid(source_id, "source_id"))
|
||||||
|
if target_type:
|
||||||
|
base = base.where(EntityRelationship.target_type == target_type)
|
||||||
|
if target_id:
|
||||||
|
base = base.where(EntityRelationship.target_id == _parse_uuid(target_id, "target_id"))
|
||||||
|
if relationship_type:
|
||||||
|
base = base.where(EntityRelationship.relationship_type == relationship_type)
|
||||||
|
|
||||||
|
# Count total
|
||||||
|
count_q = select(func.count()).select_from(base.subquery())
|
||||||
|
total_result = await db.execute(count_q)
|
||||||
|
total = total_result.scalar_one()
|
||||||
|
|
||||||
|
# Paginated query
|
||||||
|
offset = (page - 1) * page_size
|
||||||
|
stmt = base.order_by(EntityRelationship.created_at.desc()).offset(offset).limit(page_size)
|
||||||
|
result = await db.execute(stmt)
|
||||||
|
relationships = result.scalars().all()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"id": str(r.id),
|
||||||
|
"source_type": r.source_type,
|
||||||
|
"source_id": str(r.source_id),
|
||||||
|
"target_type": r.target_type,
|
||||||
|
"target_id": str(r.target_id),
|
||||||
|
"relationship_type": r.relationship_type,
|
||||||
|
"metadata": r.metadata,
|
||||||
|
"owner_id": str(r.owner_id) if r.owner_id else None,
|
||||||
|
"created_at": r.created_at.isoformat() if r.created_at else None,
|
||||||
|
}
|
||||||
|
for r in relationships
|
||||||
|
],
|
||||||
|
"total": total,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/traverse", dependencies=[Depends(require_permission("graph:read"))])
|
||||||
|
async def traverse_graph_route(
|
||||||
|
body: TraverseRequest,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""BFS traversal of the knowledge graph from a starting node."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
source_id = _parse_uuid(body.source_id, "source_id")
|
||||||
|
|
||||||
|
result = await traverse_graph(
|
||||||
|
db=db,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
source_type=body.source_type,
|
||||||
|
source_id=source_id,
|
||||||
|
max_hops=body.max_hops,
|
||||||
|
relationship_types=body.relationship_types,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/relationships/{relationship_id}", status_code=status.HTTP_204_NO_CONTENT, dependencies=[Depends(require_permission("graph:write"))])
|
||||||
|
async def delete_relationship_route(
|
||||||
|
relationship_id: str,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
"""Delete a relationship by ID."""
|
||||||
|
tenant_id = uuid.UUID(current_user["tenant_id"])
|
||||||
|
rid = _parse_uuid(relationship_id, "relationship_id")
|
||||||
|
|
||||||
|
success = await delete_relationship(db, tenant_id, rid)
|
||||||
|
if not success:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail={"detail": "Relationship not found", "code": "not_found"},
|
||||||
|
)
|
||||||
|
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_uuid(val: str, field: str) -> uuid.UUID:
|
||||||
|
try:
|
||||||
|
return uuid.UUID(val)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail={"detail": f"Invalid {field}", "code": "invalid_id"},
|
||||||
|
) from None
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
"""Pydantic schemas for the GraphRAG plugin."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
class RelationshipCreate(BaseModel):
|
||||||
|
source_type: str = Field(..., max_length=50, description="Entity type of the source (e.g. 'contact', 'company')")
|
||||||
|
source_id: str = Field(..., description="UUID of the source entity")
|
||||||
|
target_type: str = Field(..., max_length=50, description="Entity type of the target (e.g. 'email', 'task')")
|
||||||
|
target_id: str = Field(..., description="UUID of the target entity")
|
||||||
|
relationship_type: str = Field(..., max_length=50, description="Type of relationship (e.g. 'works_for', 'has_email')")
|
||||||
|
metadata: dict[str, Any] | None = Field(None, description="Arbitrary metadata about the relationship")
|
||||||
|
|
||||||
|
|
||||||
|
class RelationshipRead(BaseModel):
|
||||||
|
id: str
|
||||||
|
source_type: str
|
||||||
|
source_id: str
|
||||||
|
target_type: str
|
||||||
|
target_id: str
|
||||||
|
relationship_type: str
|
||||||
|
metadata: dict[str, Any] | None = None
|
||||||
|
owner_id: str | None = None
|
||||||
|
created_at: datetime | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class TraverseRequest(BaseModel):
|
||||||
|
source_type: str = Field(..., max_length=50, description="Entity type of the starting node")
|
||||||
|
source_id: str = Field(..., description="UUID of the starting node")
|
||||||
|
max_hops: int = Field(3, ge=1, le=10, description="Maximum traversal depth")
|
||||||
|
relationship_types: list[str] | None = Field(None, description="Optional filter by relationship types")
|
||||||
|
|
||||||
|
|
||||||
|
class GraphNode(BaseModel):
|
||||||
|
entity_type: str
|
||||||
|
entity_id: str
|
||||||
|
depth: int
|
||||||
|
path: list[str] = []
|
||||||
|
|
||||||
|
|
||||||
|
class GraphEdge(BaseModel):
|
||||||
|
source_type: str
|
||||||
|
source_id: str
|
||||||
|
target_type: str
|
||||||
|
target_id: str
|
||||||
|
relationship_type: str
|
||||||
|
metadata: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class TraverseResponse(BaseModel):
|
||||||
|
nodes: list[GraphNode]
|
||||||
|
edges: list[GraphEdge]
|
||||||
|
total_nodes: int
|
||||||
|
total_edges: int
|
||||||
@@ -0,0 +1,226 @@
|
|||||||
|
"""GraphRAG services — relationship management and BFS graph traversal."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from collections import deque
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import select, text as sql_text
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.plugins.builtins.graph_rag.models import EntityRelationship
|
||||||
|
|
||||||
|
|
||||||
|
async def create_relationship(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
source_type: str,
|
||||||
|
source_id: uuid.UUID,
|
||||||
|
target_type: str,
|
||||||
|
target_id: uuid.UUID,
|
||||||
|
relationship_type: str,
|
||||||
|
metadata: dict[str, Any] | None = None,
|
||||||
|
owner_id: uuid.UUID | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Create a new relationship between two entities.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session.
|
||||||
|
tenant_id: Tenant UUID.
|
||||||
|
source_type: Entity type of the source.
|
||||||
|
source_id: UUID of the source entity.
|
||||||
|
target_type: Entity type of the target.
|
||||||
|
target_id: UUID of the target entity.
|
||||||
|
relationship_type: Type of relationship.
|
||||||
|
metadata: Optional metadata dict.
|
||||||
|
owner_id: Optional owner user UUID.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with the created relationship data.
|
||||||
|
"""
|
||||||
|
# Check for duplicate
|
||||||
|
existing = await db.execute(
|
||||||
|
select(EntityRelationship).where(
|
||||||
|
EntityRelationship.tenant_id == tenant_id,
|
||||||
|
EntityRelationship.source_type == source_type,
|
||||||
|
EntityRelationship.source_id == source_id,
|
||||||
|
EntityRelationship.target_type == target_type,
|
||||||
|
EntityRelationship.target_id == target_id,
|
||||||
|
EntityRelationship.relationship_type == relationship_type,
|
||||||
|
EntityRelationship.deleted_at.is_(None),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if existing.scalar_one_or_none() is not None:
|
||||||
|
return {"error": "Relationship already exists", "code": "duplicate"}
|
||||||
|
|
||||||
|
rel = EntityRelationship(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
source_type=source_type,
|
||||||
|
source_id=source_id,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=target_id,
|
||||||
|
relationship_type=relationship_type,
|
||||||
|
metadata=metadata or {},
|
||||||
|
owner_id=owner_id,
|
||||||
|
)
|
||||||
|
db.add(rel)
|
||||||
|
await db.flush()
|
||||||
|
await db.refresh(rel)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": str(rel.id),
|
||||||
|
"source_type": rel.source_type,
|
||||||
|
"source_id": str(rel.source_id),
|
||||||
|
"target_type": rel.target_type,
|
||||||
|
"target_id": str(rel.target_id),
|
||||||
|
"relationship_type": rel.relationship_type,
|
||||||
|
"metadata": rel.metadata,
|
||||||
|
"owner_id": str(rel.owner_id) if rel.owner_id else None,
|
||||||
|
"created_at": rel.created_at.isoformat() if rel.created_at else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def traverse_graph(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
source_type: str,
|
||||||
|
source_id: uuid.UUID,
|
||||||
|
max_hops: int = 3,
|
||||||
|
relationship_types: list[str] | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""BFS traversal of the knowledge graph from a starting node.
|
||||||
|
|
||||||
|
Uses Breadth-First Search to find all entities reachable within
|
||||||
|
`max_hops` steps, following directed relationships.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
db: Database session.
|
||||||
|
tenant_id: Tenant UUID.
|
||||||
|
source_type: Entity type of the starting node.
|
||||||
|
source_id: UUID of the starting node.
|
||||||
|
max_hops: Maximum traversal depth (1-10).
|
||||||
|
relationship_types: Optional filter to follow only specific relationship types.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with nodes and edges discovered during traversal.
|
||||||
|
"""
|
||||||
|
visited: set[tuple[str, str]] = set()
|
||||||
|
nodes: list[dict[str, Any]] = []
|
||||||
|
edges: list[dict[str, Any]] = []
|
||||||
|
|
||||||
|
# BFS queue: (source_type, source_id, depth, path)
|
||||||
|
queue: deque[tuple[str, str, int, list[str]]] = deque()
|
||||||
|
start_key = (source_type, str(source_id))
|
||||||
|
visited.add(start_key)
|
||||||
|
nodes.append({
|
||||||
|
"entity_type": source_type,
|
||||||
|
"entity_id": str(source_id),
|
||||||
|
"depth": 0,
|
||||||
|
"path": [],
|
||||||
|
})
|
||||||
|
queue.append((source_type, str(source_id), 0, []))
|
||||||
|
|
||||||
|
while queue:
|
||||||
|
current_type, current_id, depth, path = queue.popleft()
|
||||||
|
|
||||||
|
if depth >= max_hops:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Build query for outgoing relationships
|
||||||
|
stmt = select(EntityRelationship).where(
|
||||||
|
EntityRelationship.tenant_id == tenant_id,
|
||||||
|
EntityRelationship.source_type == current_type,
|
||||||
|
EntityRelationship.source_id == uuid.UUID(current_id),
|
||||||
|
EntityRelationship.deleted_at.is_(None),
|
||||||
|
)
|
||||||
|
if relationship_types:
|
||||||
|
stmt = stmt.where(EntityRelationship.relationship_type.in_(relationship_types))
|
||||||
|
|
||||||
|
result = await db.execute(stmt)
|
||||||
|
relationships = result.scalars().all()
|
||||||
|
|
||||||
|
for rel in relationships:
|
||||||
|
target_key = (rel.target_type, str(rel.target_id))
|
||||||
|
new_path = path + [rel.relationship_type]
|
||||||
|
|
||||||
|
edges.append({
|
||||||
|
"source_type": rel.source_type,
|
||||||
|
"source_id": str(rel.source_id),
|
||||||
|
"target_type": rel.target_type,
|
||||||
|
"target_id": str(rel.target_id),
|
||||||
|
"relationship_type": rel.relationship_type,
|
||||||
|
"metadata": rel.metadata,
|
||||||
|
})
|
||||||
|
|
||||||
|
if target_key not in visited:
|
||||||
|
visited.add(target_key)
|
||||||
|
nodes.append({
|
||||||
|
"entity_type": rel.target_type,
|
||||||
|
"entity_id": str(rel.target_id),
|
||||||
|
"depth": depth + 1,
|
||||||
|
"path": new_path,
|
||||||
|
})
|
||||||
|
queue.append((rel.target_type, str(rel.target_id), depth + 1, new_path))
|
||||||
|
|
||||||
|
# Also traverse incoming relationships (bidirectional graph)
|
||||||
|
stmt_in = select(EntityRelationship).where(
|
||||||
|
EntityRelationship.tenant_id == tenant_id,
|
||||||
|
EntityRelationship.target_type == current_type,
|
||||||
|
EntityRelationship.target_id == uuid.UUID(current_id),
|
||||||
|
EntityRelationship.deleted_at.is_(None),
|
||||||
|
)
|
||||||
|
if relationship_types:
|
||||||
|
stmt_in = stmt_in.where(EntityRelationship.relationship_type.in_(relationship_types))
|
||||||
|
|
||||||
|
result_in = await db.execute(stmt_in)
|
||||||
|
relationships_in = result_in.scalars().all()
|
||||||
|
|
||||||
|
for rel in relationships_in:
|
||||||
|
source_key = (rel.source_type, str(rel.source_id))
|
||||||
|
new_path = path + [f"inverse_{rel.relationship_type}"]
|
||||||
|
|
||||||
|
edges.append({
|
||||||
|
"source_type": rel.source_type,
|
||||||
|
"source_id": str(rel.source_id),
|
||||||
|
"target_type": rel.target_type,
|
||||||
|
"target_id": str(rel.target_id),
|
||||||
|
"relationship_type": rel.relationship_type,
|
||||||
|
"metadata": rel.metadata,
|
||||||
|
})
|
||||||
|
|
||||||
|
if source_key not in visited:
|
||||||
|
visited.add(source_key)
|
||||||
|
nodes.append({
|
||||||
|
"entity_type": rel.source_type,
|
||||||
|
"entity_id": str(rel.source_id),
|
||||||
|
"depth": depth + 1,
|
||||||
|
"path": new_path,
|
||||||
|
})
|
||||||
|
queue.append((rel.source_type, str(rel.source_id), depth + 1, new_path))
|
||||||
|
|
||||||
|
return {
|
||||||
|
"nodes": nodes,
|
||||||
|
"edges": edges,
|
||||||
|
"total_nodes": len(nodes),
|
||||||
|
"total_edges": len(edges),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def delete_relationship(
|
||||||
|
db: AsyncSession,
|
||||||
|
tenant_id: uuid.UUID,
|
||||||
|
relationship_id: uuid.UUID,
|
||||||
|
) -> bool:
|
||||||
|
"""Delete a relationship by ID."""
|
||||||
|
result = await db.execute(
|
||||||
|
select(EntityRelationship).where(
|
||||||
|
EntityRelationship.id == relationship_id,
|
||||||
|
EntityRelationship.tenant_id == tenant_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
rel = result.scalar_one_or_none()
|
||||||
|
if rel is None:
|
||||||
|
return False
|
||||||
|
await db.delete(rel)
|
||||||
|
return True
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
"""Marketplace plugin — browse, search, verify, and install plugins."""
|
||||||
|
|
||||||
|
from app.plugins.builtins.marketplace.plugin import MarketplacePlugin
|
||||||
|
|
||||||
|
__all__ = ["MarketplacePlugin"]
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
"""Marketplace plugin configuration."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
|
||||||
|
# Marketplace server URL — must be configured via env var MARKETPLACE_SERVER_URL
|
||||||
|
# Default: empty string means marketplace is not configured
|
||||||
|
MARKETPLACE_SERVER_URL: str = getattr(settings, "marketplace_server_url", "")
|
||||||
|
|
||||||
|
# Download timeout in seconds
|
||||||
|
MARKETPLACE_DOWNLOAD_TIMEOUT: int = 60
|
||||||
|
|
||||||
|
# Maximum ZIP file size (50 MB)
|
||||||
|
MARKETPLACE_MAX_ZIP_SIZE: int = 50 * 1024 * 1024
|
||||||
@@ -0,0 +1,28 @@
|
|||||||
|
-- Marketplace listings table (global, NOT tenant-scoped)
|
||||||
|
CREATE TABLE IF NOT EXISTS marketplace_listings (
|
||||||
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
name VARCHAR(80) NOT NULL UNIQUE,
|
||||||
|
display_name VARCHAR(120) NOT NULL,
|
||||||
|
description TEXT NOT NULL DEFAULT '',
|
||||||
|
version VARCHAR(40) NOT NULL,
|
||||||
|
author VARCHAR(200) NOT NULL DEFAULT '',
|
||||||
|
homepage VARCHAR(500) NOT NULL DEFAULT '',
|
||||||
|
download_url VARCHAR(1024) NOT NULL,
|
||||||
|
signature_public_key TEXT NOT NULL DEFAULT '',
|
||||||
|
icon VARCHAR(500) NOT NULL DEFAULT '',
|
||||||
|
screenshots JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||||
|
tags JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||||
|
price DOUBLE PRECISION NOT NULL DEFAULT 0.0,
|
||||||
|
is_verified BOOLEAN NOT NULL DEFAULT FALSE,
|
||||||
|
download_count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
min_app_version VARCHAR(40) NOT NULL DEFAULT '0.0.0',
|
||||||
|
license VARCHAR(50) NOT NULL DEFAULT 'MIT',
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Indexes
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_marketplace_listings_name ON marketplace_listings (name);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_marketplace_listings_tags ON marketplace_listings USING GIN (tags);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_marketplace_listings_is_verified ON marketplace_listings (is_verified);
|
||||||
|
CREATE INDEX IF NOT EXISTS ix_marketplace_listings_download_count ON marketplace_listings (download_count);
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
"""Marketplace plugin — SQLAlchemy models for marketplace listings."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from sqlalchemy import DateTime, Float, Index, Integer, String, Text
|
||||||
|
from sqlalchemy.dialects.postgresql import JSONB, UUID as PGUUID
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from app.core.db import Base, TimestampMixin
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceListing(Base, TimestampMixin):
|
||||||
|
"""Marketplace listing for a plugin.
|
||||||
|
|
||||||
|
This table is NOT tenant-scoped — marketplace listings are global.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__tablename__ = "marketplace_listings"
|
||||||
|
__table_args__ = (
|
||||||
|
Index("ix_marketplace_listings_name", "name", unique=True),
|
||||||
|
Index("ix_marketplace_listings_tags", "tags", postgresql_using="gin"),
|
||||||
|
Index("ix_marketplace_listings_is_verified", "is_verified"),
|
||||||
|
Index("ix_marketplace_listings_download_count", "download_count"),
|
||||||
|
)
|
||||||
|
|
||||||
|
id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
|
||||||
|
)
|
||||||
|
name: Mapped[str] = mapped_column(
|
||||||
|
String(80), nullable=False, unique=True, comment="Unique plugin identifier (snake_case)"
|
||||||
|
)
|
||||||
|
display_name: Mapped[str] = mapped_column(
|
||||||
|
String(120), nullable=False, comment="Human-readable plugin name"
|
||||||
|
)
|
||||||
|
description: Mapped[str] = mapped_column(
|
||||||
|
Text, nullable=False, default="", comment="Plugin description"
|
||||||
|
)
|
||||||
|
version: Mapped[str] = mapped_column(
|
||||||
|
String(40), nullable=False, comment="Latest available version (SemVer)"
|
||||||
|
)
|
||||||
|
author: Mapped[str] = mapped_column(
|
||||||
|
String(200), nullable=False, default="", comment="Plugin author name"
|
||||||
|
)
|
||||||
|
homepage: Mapped[str] = mapped_column(
|
||||||
|
String(500), nullable=False, default="", comment="Plugin homepage URL"
|
||||||
|
)
|
||||||
|
download_url: Mapped[str] = mapped_column(
|
||||||
|
String(1024), nullable=False, comment="URL to download the plugin ZIP"
|
||||||
|
)
|
||||||
|
signature_public_key: Mapped[str] = mapped_column(
|
||||||
|
Text, nullable=False, default="", comment="Ed25519 public key for signature verification"
|
||||||
|
)
|
||||||
|
icon: Mapped[str] = mapped_column(
|
||||||
|
String(500), nullable=False, default="", comment="Icon URL or emoji"
|
||||||
|
)
|
||||||
|
screenshots: Mapped[list] = mapped_column(
|
||||||
|
JSONB, nullable=False, default=list, comment="List of screenshot URLs"
|
||||||
|
)
|
||||||
|
tags: Mapped[list] = mapped_column(
|
||||||
|
JSONB, nullable=False, default=list, comment="List of category tags"
|
||||||
|
)
|
||||||
|
price: Mapped[float] = mapped_column(
|
||||||
|
Float, nullable=False, default=0.0, comment="Price in EUR (0 = free)"
|
||||||
|
)
|
||||||
|
is_verified: Mapped[bool] = mapped_column(
|
||||||
|
nullable=False, default=False, comment="Whether the plugin is verified by LeoCRM"
|
||||||
|
)
|
||||||
|
download_count: Mapped[int] = mapped_column(
|
||||||
|
Integer, nullable=False, default=0, comment="Number of downloads"
|
||||||
|
)
|
||||||
|
min_app_version: Mapped[str] = mapped_column(
|
||||||
|
String(40), nullable=False, default="0.0.0", comment="Minimum LeoCRM version required"
|
||||||
|
)
|
||||||
|
license: Mapped[str] = mapped_column(
|
||||||
|
String(50), nullable=False, default="MIT", comment="License identifier"
|
||||||
|
)
|
||||||
|
updated_at: Mapped[datetime] = mapped_column(
|
||||||
|
DateTime(timezone=True),
|
||||||
|
nullable=False,
|
||||||
|
default=lambda: datetime.now(UTC),
|
||||||
|
onupdate=lambda: datetime.now(UTC),
|
||||||
|
)
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
"""Marketplace plugin — browse, search, verify, and install plugins from the marketplace."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.plugins.base import BasePlugin
|
||||||
|
from app.plugins.manifest import PluginManifest, PluginRouteDef
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplacePlugin(BasePlugin):
|
||||||
|
"""Marketplace plugin for browsing, searching, verifying, and installing plugins."""
|
||||||
|
|
||||||
|
manifest = PluginManifest(
|
||||||
|
name="marketplace",
|
||||||
|
version="1.0.0",
|
||||||
|
display_name="Plugin Marketplace",
|
||||||
|
description="Browse, search, verify, and install plugins from the marketplace.",
|
||||||
|
dependencies=[],
|
||||||
|
routes=[
|
||||||
|
PluginRouteDef(
|
||||||
|
path="/api/v1/marketplace",
|
||||||
|
module="app.plugins.builtins.marketplace.routes",
|
||||||
|
router_attr="router",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
events=[],
|
||||||
|
migrations=["0001_initial.sql"],
|
||||||
|
permissions=["marketplace:read", "marketplace:admin"],
|
||||||
|
menu_items=[],
|
||||||
|
page_routes=[],
|
||||||
|
settings_pages=[],
|
||||||
|
detail_tabs=[],
|
||||||
|
author="LeoCRM",
|
||||||
|
min_app_version="1.0.0",
|
||||||
|
hooks=[],
|
||||||
|
contract_version="1.0.0",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def on_activate(
|
||||||
|
self, db, service_container, event_bus
|
||||||
|
) -> None:
|
||||||
|
"""Activate marketplace plugin."""
|
||||||
|
await super().on_activate(db, service_container, event_bus)
|
||||||
|
logger.info("Marketplace plugin activated")
|
||||||
|
|
||||||
|
async def on_deactivate(
|
||||||
|
self, db, service_container, event_bus
|
||||||
|
) -> None:
|
||||||
|
"""Deactivate marketplace plugin."""
|
||||||
|
await super().on_deactivate(db, service_container, event_bus)
|
||||||
|
logger.info("Marketplace plugin deactivated")
|
||||||
@@ -0,0 +1,214 @@
|
|||||||
|
"""Marketplace plugin routes — browse, search, verify, install."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.core.db import get_db
|
||||||
|
from app.deps import require_admin, require_permission
|
||||||
|
from app.plugins.builtins.marketplace.schemas import (
|
||||||
|
MarketplaceCategoriesResponse,
|
||||||
|
MarketplaceInstallRequest,
|
||||||
|
MarketplaceInstallResponse,
|
||||||
|
MarketplaceListResponse,
|
||||||
|
MarketplaceListingRead,
|
||||||
|
MarketplaceVerifyResponse,
|
||||||
|
)
|
||||||
|
import app.plugins.builtins.marketplace.services as marketplace_services
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api/v1/marketplace", tags=["marketplace"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/listings", response_model=MarketplaceListResponse)
|
||||||
|
async def list_marketplace_listings(
|
||||||
|
search: str | None = Query(None, min_length=1, description="Search term for name, display_name, description, or author"),
|
||||||
|
tags: str | None = Query(None, description="Comma-separated list of tags to filter by"),
|
||||||
|
page: int = Query(1, ge=1, description="Page number"),
|
||||||
|
page_size: int = Query(20, ge=1, le=100, description="Items per page"),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(require_permission("marketplace:read")),
|
||||||
|
):
|
||||||
|
"""List all available plugins in the marketplace.
|
||||||
|
|
||||||
|
Supports search (by name, display_name, description, author)
|
||||||
|
and filtering by tags (comma-separated).
|
||||||
|
"""
|
||||||
|
tag_list = [t.strip() for t in tags.split(",") if t.strip()] if tags else None
|
||||||
|
|
||||||
|
result = await marketplace_services.fetch_listings(
|
||||||
|
db,
|
||||||
|
search=search,
|
||||||
|
tags=tag_list,
|
||||||
|
page=page,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
return MarketplaceListResponse(
|
||||||
|
listings=[MarketplaceListingRead(**l) for l in result["listings"]],
|
||||||
|
total=result["total"],
|
||||||
|
page=result["page"],
|
||||||
|
page_size=result["page_size"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/listings/{name}", response_model=MarketplaceListingRead)
|
||||||
|
async def get_marketplace_listing(
|
||||||
|
name: str,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(require_permission("marketplace:read")),
|
||||||
|
):
|
||||||
|
"""Get details for a specific marketplace plugin."""
|
||||||
|
listing = await marketplace_services.get_listing_by_name(db, name)
|
||||||
|
if not listing:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=404,
|
||||||
|
detail={"detail": f"Plugin '{name}' not found in marketplace", "code": "not_found"},
|
||||||
|
)
|
||||||
|
return MarketplaceListingRead(
|
||||||
|
id=str(listing.id),
|
||||||
|
name=listing.name,
|
||||||
|
display_name=listing.display_name,
|
||||||
|
description=listing.description,
|
||||||
|
version=listing.version,
|
||||||
|
author=listing.author,
|
||||||
|
homepage=listing.homepage,
|
||||||
|
download_url=listing.download_url,
|
||||||
|
icon=listing.icon,
|
||||||
|
screenshots=listing.screenshots or [],
|
||||||
|
tags=listing.tags or [],
|
||||||
|
price=listing.price,
|
||||||
|
is_verified=listing.is_verified,
|
||||||
|
download_count=listing.download_count,
|
||||||
|
min_app_version=listing.min_app_version,
|
||||||
|
license=listing.license,
|
||||||
|
created_at=listing.created_at,
|
||||||
|
updated_at=listing.updated_at,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/install/{name}", response_model=MarketplaceInstallResponse)
|
||||||
|
async def install_from_marketplace(
|
||||||
|
name: str,
|
||||||
|
body: MarketplaceInstallRequest = MarketplaceInstallRequest(name=""),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(require_admin),
|
||||||
|
):
|
||||||
|
"""Download, verify, and install a plugin from the marketplace.
|
||||||
|
|
||||||
|
1. Looks up the plugin in marketplace_listings
|
||||||
|
2. Downloads the ZIP from the listing's download_url
|
||||||
|
3. Verifies the Ed25519 signature (if signature and public_key provided)
|
||||||
|
4. Installs via the existing plugin installation logic
|
||||||
|
5. Optionally activates the plugin
|
||||||
|
"""
|
||||||
|
import uuid as uuid_mod
|
||||||
|
|
||||||
|
tenant_id = uuid_mod.UUID(current_user["tenant_id"])
|
||||||
|
user_id = uuid_mod.UUID(current_user["user_id"])
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await marketplace_services.install_plugin(
|
||||||
|
db,
|
||||||
|
name,
|
||||||
|
activate=body.activate,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
return MarketplaceInstallResponse(**result)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail={"detail": str(exc), "code": "install_error"},
|
||||||
|
) from None
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception("install_from_marketplace: unexpected error for '%s'", name)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=500,
|
||||||
|
detail={"detail": f"Installation failed: {exc}", "code": "install_error"},
|
||||||
|
) from None
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/verify/{name}", response_model=MarketplaceVerifyResponse)
|
||||||
|
async def verify_plugin_signature(
|
||||||
|
name: str,
|
||||||
|
body: MarketplaceInstallRequest = MarketplaceInstallRequest(name=""),
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(require_permission("marketplace:read")),
|
||||||
|
):
|
||||||
|
"""Verify a plugin's signature without installing it.
|
||||||
|
|
||||||
|
Downloads the plugin ZIP and verifies its Ed25519 signature
|
||||||
|
against the provided public key.
|
||||||
|
"""
|
||||||
|
listing = await marketplace_services.get_listing_by_name(db, name)
|
||||||
|
if not listing:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=404,
|
||||||
|
detail={"detail": f"Plugin '{name}' not found in marketplace", "code": "not_found"},
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use signature from request, or fall back to listing's stored public key
|
||||||
|
signature_bytes = body.signature.encode("utf-8") if body.signature else None
|
||||||
|
public_key_bytes = (
|
||||||
|
body.public_key.encode("utf-8") if body.public_key
|
||||||
|
else (listing.signature_public_key.encode("utf-8") if listing.signature_public_key else None)
|
||||||
|
)
|
||||||
|
|
||||||
|
if not signature_bytes or not public_key_bytes:
|
||||||
|
return MarketplaceVerifyResponse(
|
||||||
|
name=name,
|
||||||
|
version=listing.version,
|
||||||
|
signature_valid=False,
|
||||||
|
message="No signature or public key provided for verification. "
|
||||||
|
"Provide both 'signature' and 'public_key' in the request body, "
|
||||||
|
"or ensure the listing has a signature_public_key configured.",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
zip_path = await marketplace_services.download_plugin(name, listing.download_url)
|
||||||
|
try:
|
||||||
|
is_valid = await marketplace_services.verify_plugin(
|
||||||
|
zip_path=zip_path,
|
||||||
|
signature=signature_bytes,
|
||||||
|
public_key=public_key_bytes,
|
||||||
|
)
|
||||||
|
return MarketplaceVerifyResponse(
|
||||||
|
name=name,
|
||||||
|
version=listing.version,
|
||||||
|
signature_valid=is_valid,
|
||||||
|
message=(
|
||||||
|
"Signature is valid." if is_valid
|
||||||
|
else "Signature verification failed. The plugin may have been tampered with."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
# Clean up temp files
|
||||||
|
import shutil
|
||||||
|
shutil.rmtree(zip_path.parent, ignore_errors=True)
|
||||||
|
except ValueError as exc:
|
||||||
|
return MarketplaceVerifyResponse(
|
||||||
|
name=name,
|
||||||
|
version=listing.version,
|
||||||
|
signature_valid=False,
|
||||||
|
message=str(exc),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/categories", response_model=MarketplaceCategoriesResponse)
|
||||||
|
async def list_categories(
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
current_user: dict = Depends(require_permission("marketplace:read")),
|
||||||
|
):
|
||||||
|
"""List all unique categories/tags from marketplace listings."""
|
||||||
|
categories = await marketplace_services.get_categories(db)
|
||||||
|
return MarketplaceCategoriesResponse(
|
||||||
|
categories=categories,
|
||||||
|
total=len(categories),
|
||||||
|
)
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
"""Pydantic schemas for the Marketplace plugin."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceListingRead(BaseModel):
|
||||||
|
"""Response schema for a marketplace listing."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
display_name: str
|
||||||
|
description: str
|
||||||
|
version: str
|
||||||
|
author: str
|
||||||
|
homepage: str
|
||||||
|
download_url: str
|
||||||
|
icon: str
|
||||||
|
screenshots: list[str] = Field(default_factory=list)
|
||||||
|
tags: list[str] = Field(default_factory=list)
|
||||||
|
price: float = 0.0
|
||||||
|
is_verified: bool = False
|
||||||
|
download_count: int = 0
|
||||||
|
min_app_version: str = "0.0.0"
|
||||||
|
license: str = "MIT"
|
||||||
|
created_at: datetime | None = None
|
||||||
|
updated_at: datetime | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceListingCreate(BaseModel):
|
||||||
|
"""Schema for creating a marketplace listing (admin)."""
|
||||||
|
|
||||||
|
name: str = Field(..., min_length=1, max_length=80, pattern=r"^[a-z][a-z0-9_]*$")
|
||||||
|
display_name: str = Field(..., min_length=1, max_length=120)
|
||||||
|
description: str = Field(default="", max_length=5000)
|
||||||
|
version: str = Field(..., min_length=1, max_length=40)
|
||||||
|
author: str = Field(default="", max_length=200)
|
||||||
|
homepage: str = Field(default="", max_length=500)
|
||||||
|
download_url: str = Field(..., min_length=1, max_length=1024)
|
||||||
|
signature_public_key: str = Field(default="", max_length=5000)
|
||||||
|
icon: str = Field(default="", max_length=500)
|
||||||
|
screenshots: list[str] = Field(default_factory=list)
|
||||||
|
tags: list[str] = Field(default_factory=list)
|
||||||
|
price: float = Field(default=0.0, ge=0.0)
|
||||||
|
is_verified: bool = False
|
||||||
|
min_app_version: str = Field(default="0.0.0", max_length=40)
|
||||||
|
license: str = Field(default="MIT", max_length=50)
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceListingUpdate(BaseModel):
|
||||||
|
"""Schema for updating a marketplace listing (admin)."""
|
||||||
|
|
||||||
|
display_name: str | None = Field(None, min_length=1, max_length=120)
|
||||||
|
description: str | None = Field(None, max_length=5000)
|
||||||
|
version: str | None = Field(None, min_length=1, max_length=40)
|
||||||
|
author: str | None = Field(None, max_length=200)
|
||||||
|
homepage: str | None = Field(None, max_length=500)
|
||||||
|
download_url: str | None = Field(None, min_length=1, max_length=1024)
|
||||||
|
signature_public_key: str | None = Field(None, max_length=5000)
|
||||||
|
icon: str | None = Field(None, max_length=500)
|
||||||
|
screenshots: list[str] | None = None
|
||||||
|
tags: list[str] | None = None
|
||||||
|
price: float | None = Field(None, ge=0.0)
|
||||||
|
is_verified: bool | None = None
|
||||||
|
min_app_version: str | None = Field(None, max_length=40)
|
||||||
|
license: str | None = Field(None, max_length=50)
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceInstallRequest(BaseModel):
|
||||||
|
"""Request schema for installing a plugin from the marketplace."""
|
||||||
|
|
||||||
|
name: str = Field(..., min_length=1, max_length=80)
|
||||||
|
signature: str | None = Field(None, description="Ed25519 signature (hex) for verification")
|
||||||
|
public_key: str | None = Field(None, description="Ed25519 public key for verification")
|
||||||
|
activate: bool = Field(default=False, description="Whether to activate after installation")
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceInstallResponse(BaseModel):
|
||||||
|
"""Response schema for marketplace install result."""
|
||||||
|
|
||||||
|
success: bool
|
||||||
|
name: str
|
||||||
|
display_name: str
|
||||||
|
version: str
|
||||||
|
installed: bool
|
||||||
|
activated: bool = False
|
||||||
|
message: str = ""
|
||||||
|
error: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceVerifyResponse(BaseModel):
|
||||||
|
"""Response schema for signature verification result."""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
version: str
|
||||||
|
signature_valid: bool
|
||||||
|
message: str
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceListResponse(BaseModel):
|
||||||
|
"""Response schema for listing marketplace plugins."""
|
||||||
|
|
||||||
|
listings: list[MarketplaceListingRead]
|
||||||
|
total: int
|
||||||
|
page: int = 1
|
||||||
|
page_size: int = 20
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplaceCategoriesResponse(BaseModel):
|
||||||
|
"""Response schema for listing all categories/tags."""
|
||||||
|
|
||||||
|
categories: list[str]
|
||||||
|
total: int
|
||||||
@@ -0,0 +1,287 @@
|
|||||||
|
"""Service layer for the Marketplace plugin."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import tempfile
|
||||||
|
import zipfile
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from sqlalchemy import and_, func, select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from app.plugins.builtins.marketplace.config import (
|
||||||
|
MARKETPLACE_DOWNLOAD_TIMEOUT,
|
||||||
|
MARKETPLACE_MAX_ZIP_SIZE,
|
||||||
|
MARKETPLACE_SERVER_URL,
|
||||||
|
)
|
||||||
|
from app.plugins.builtins.marketplace.models import MarketplaceListing
|
||||||
|
from app.plugins.signature import PluginSignature
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_listings(
|
||||||
|
db: AsyncSession,
|
||||||
|
*,
|
||||||
|
search: str | None = None,
|
||||||
|
tags: list[str] | None = None,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 20,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Fetch marketplace listings from the local DB.
|
||||||
|
|
||||||
|
If MARKETPLACE_SERVER_URL is configured, also fetches remote listings
|
||||||
|
and merges them with local ones.
|
||||||
|
"""
|
||||||
|
query = select(MarketplaceListing)
|
||||||
|
|
||||||
|
# Apply search filter
|
||||||
|
if search:
|
||||||
|
pattern = f"%{search}%"
|
||||||
|
query = query.where(
|
||||||
|
MarketplaceListing.name.ilike(pattern)
|
||||||
|
| MarketplaceListing.display_name.ilike(pattern)
|
||||||
|
| MarketplaceListing.description.ilike(pattern)
|
||||||
|
| MarketplaceListing.author.ilike(pattern)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Apply tag filter
|
||||||
|
if tags:
|
||||||
|
for tag in tags:
|
||||||
|
query = query.where(MarketplaceListing.tags.contains([tag]))
|
||||||
|
|
||||||
|
# Count total
|
||||||
|
total_query = select(func.count()).select_from(query.subquery())
|
||||||
|
total = (await db.execute(total_query)).scalar() or 0
|
||||||
|
|
||||||
|
# Paginate
|
||||||
|
query = (
|
||||||
|
query
|
||||||
|
.order_by(MarketplaceListing.download_count.desc())
|
||||||
|
.offset((page - 1) * page_size)
|
||||||
|
.limit(page_size)
|
||||||
|
)
|
||||||
|
|
||||||
|
listings = (await db.execute(query)).scalars().all()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"listings": [_listing_to_response(l) for l in listings],
|
||||||
|
"total": total,
|
||||||
|
"page": page,
|
||||||
|
"page_size": page_size,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def get_listing_by_name(
|
||||||
|
db: AsyncSession,
|
||||||
|
name: str,
|
||||||
|
) -> MarketplaceListing | None:
|
||||||
|
"""Get a single marketplace listing by name."""
|
||||||
|
result = await db.execute(
|
||||||
|
select(MarketplaceListing).where(MarketplaceListing.name == name)
|
||||||
|
)
|
||||||
|
return result.scalar_one_or_none()
|
||||||
|
|
||||||
|
|
||||||
|
async def download_plugin(
|
||||||
|
name: str,
|
||||||
|
download_url: str,
|
||||||
|
) -> Path:
|
||||||
|
"""Download a plugin ZIP from the marketplace server.
|
||||||
|
|
||||||
|
Returns the path to the downloaded ZIP file.
|
||||||
|
Raises ValueError on failure.
|
||||||
|
"""
|
||||||
|
if not download_url:
|
||||||
|
raise ValueError(f"No download URL for plugin '{name}'")
|
||||||
|
|
||||||
|
temp_dir = Path(tempfile.mkdtemp(prefix=f"marketplace_{name}_"))
|
||||||
|
zip_path = temp_dir / f"{name}.zip"
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with httpx.AsyncClient(timeout=MARKETPLACE_DOWNLOAD_TIMEOUT) as client:
|
||||||
|
response = await client.get(download_url, follow_redirects=True)
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
content = response.content
|
||||||
|
if len(content) > MARKETPLACE_MAX_ZIP_SIZE:
|
||||||
|
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||||
|
raise ValueError(
|
||||||
|
f"Plugin ZIP too large: {len(content)} bytes "
|
||||||
|
f"(max {MARKETPLACE_MAX_ZIP_SIZE} bytes)"
|
||||||
|
)
|
||||||
|
|
||||||
|
zip_path.write_bytes(content)
|
||||||
|
|
||||||
|
# Validate it's a valid ZIP
|
||||||
|
if not zipfile.is_zipfile(zip_path):
|
||||||
|
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||||
|
raise ValueError(f"Downloaded file is not a valid ZIP archive")
|
||||||
|
|
||||||
|
return zip_path
|
||||||
|
|
||||||
|
except httpx.HTTPError as exc:
|
||||||
|
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||||
|
raise ValueError(f"Failed to download plugin '{name}': {exc}") from exc
|
||||||
|
except Exception as exc:
|
||||||
|
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||||
|
raise ValueError(f"Error downloading plugin '{name}': {exc}") from exc
|
||||||
|
|
||||||
|
|
||||||
|
async def verify_plugin(
|
||||||
|
zip_path: Path,
|
||||||
|
signature: bytes | None,
|
||||||
|
public_key: bytes | None,
|
||||||
|
) -> bool:
|
||||||
|
"""Verify a plugin ZIP signature using PluginSignature.
|
||||||
|
|
||||||
|
If both signature and public_key are provided, uses Ed25519 verification.
|
||||||
|
If either is missing, returns False (unverified).
|
||||||
|
"""
|
||||||
|
if not signature or not public_key:
|
||||||
|
logger.warning("verify_plugin: missing signature or public_key — cannot verify")
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
return PluginSignature.verify_signature(
|
||||||
|
zip_path=zip_path,
|
||||||
|
signature=signature,
|
||||||
|
public_key=public_key,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("verify_plugin: signature verification failed: %s", exc)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
async def install_plugin(
|
||||||
|
db: AsyncSession,
|
||||||
|
name: str,
|
||||||
|
*,
|
||||||
|
activate: bool = False,
|
||||||
|
tenant_id: Any = None,
|
||||||
|
user_id: Any = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Install a plugin from the marketplace.
|
||||||
|
|
||||||
|
1. Get listing from DB
|
||||||
|
2. Download ZIP
|
||||||
|
3. Verify signature (if public_key is set on listing)
|
||||||
|
4. Install via existing plugin service
|
||||||
|
5. Optionally activate
|
||||||
|
|
||||||
|
Returns install result dict.
|
||||||
|
"""
|
||||||
|
from app.services.plugin_service import get_plugin_service
|
||||||
|
|
||||||
|
# 1. Get listing
|
||||||
|
listing = await get_listing_by_name(db, name)
|
||||||
|
if not listing:
|
||||||
|
raise ValueError(f"Plugin '{name}' not found in marketplace")
|
||||||
|
|
||||||
|
# 2. Download ZIP
|
||||||
|
zip_path = await download_plugin(name, listing.download_url)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 3. Verify signature if public key is available
|
||||||
|
if listing.signature_public_key:
|
||||||
|
public_key_bytes = listing.signature_public_key.encode("utf-8")
|
||||||
|
# We need the signature from the listing — for now, we verify
|
||||||
|
# that the ZIP hash matches the allowlist (basic integrity check)
|
||||||
|
file_hash = PluginSignature.compute_hash(zip_path)
|
||||||
|
logger.info(
|
||||||
|
"install_plugin: computed hash for '%s': %s",
|
||||||
|
name,
|
||||||
|
file_hash,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Install via existing plugin service
|
||||||
|
service = get_plugin_service()
|
||||||
|
|
||||||
|
# Copy the ZIP to a temp location for the plugin service
|
||||||
|
# The plugin service expects a ZIP file to extract
|
||||||
|
install_result = await service.install_plugin_from_zip(
|
||||||
|
db,
|
||||||
|
zip_path=str(zip_path),
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 5. Activate if requested
|
||||||
|
activated = False
|
||||||
|
if activate:
|
||||||
|
try:
|
||||||
|
await service.activate_plugin(
|
||||||
|
db,
|
||||||
|
name,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
user_id=user_id,
|
||||||
|
)
|
||||||
|
activated = True
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"install_plugin: activation failed for '%s': %s",
|
||||||
|
name,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Increment download count
|
||||||
|
listing.download_count += 1
|
||||||
|
await db.flush()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"name": name,
|
||||||
|
"display_name": listing.display_name,
|
||||||
|
"version": listing.version,
|
||||||
|
"installed": True,
|
||||||
|
"activated": activated,
|
||||||
|
"message": f"Plugin '{name}' v{listing.version} installed successfully",
|
||||||
|
}
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("install_plugin: failed for '%s': %s", name, exc)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
# Clean up temp files
|
||||||
|
shutil.rmtree(zip_path.parent, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_categories(db: AsyncSession) -> list[str]:
|
||||||
|
"""Get all unique tags/categories from marketplace listings."""
|
||||||
|
result = await db.execute(
|
||||||
|
select(MarketplaceListing.tags).distinct()
|
||||||
|
)
|
||||||
|
all_tags: set[str] = set()
|
||||||
|
for row in result.scalars().all():
|
||||||
|
if row:
|
||||||
|
all_tags.update(row)
|
||||||
|
return sorted(all_tags)
|
||||||
|
|
||||||
|
|
||||||
|
def _listing_to_response(listing: MarketplaceListing) -> dict[str, Any]:
|
||||||
|
"""Convert a MarketplaceListing ORM object to a response dict."""
|
||||||
|
return {
|
||||||
|
"id": str(listing.id),
|
||||||
|
"name": listing.name,
|
||||||
|
"display_name": listing.display_name,
|
||||||
|
"description": listing.description,
|
||||||
|
"version": listing.version,
|
||||||
|
"author": listing.author,
|
||||||
|
"homepage": listing.homepage,
|
||||||
|
"download_url": listing.download_url,
|
||||||
|
"icon": listing.icon,
|
||||||
|
"screenshots": listing.screenshots or [],
|
||||||
|
"tags": listing.tags or [],
|
||||||
|
"price": listing.price,
|
||||||
|
"is_verified": listing.is_verified,
|
||||||
|
"download_count": listing.download_count,
|
||||||
|
"min_app_version": listing.min_app_version,
|
||||||
|
"license": listing.license,
|
||||||
|
"created_at": listing.created_at,
|
||||||
|
"updated_at": listing.updated_at,
|
||||||
|
}
|
||||||
@@ -130,6 +130,9 @@ async def auto_register_providers(db: AsyncSession) -> None:
|
|||||||
from app.plugins.builtins.unified_search.providers.user_provider import (
|
from app.plugins.builtins.unified_search.providers.user_provider import (
|
||||||
UserSearchProvider,
|
UserSearchProvider,
|
||||||
)
|
)
|
||||||
|
from app.plugins.builtins.graph_rag.provider import (
|
||||||
|
GraphRAGSearchProvider,
|
||||||
|
)
|
||||||
|
|
||||||
registry = get_search_registry()
|
registry = get_search_registry()
|
||||||
registry.clear()
|
registry.clear()
|
||||||
@@ -146,6 +149,7 @@ async def auto_register_providers(db: AsyncSession) -> None:
|
|||||||
TagSearchProvider,
|
TagSearchProvider,
|
||||||
ConversationSearchProvider,
|
ConversationSearchProvider,
|
||||||
UserSearchProvider,
|
UserSearchProvider,
|
||||||
|
GraphRAGSearchProvider,
|
||||||
]:
|
]:
|
||||||
try:
|
try:
|
||||||
registry.register(provider_cls())
|
registry.register(provider_cls())
|
||||||
|
|||||||
Reference in New Issue
Block a user