ai-memory-backend / api /capability_catalog.py
Baida07's picture
sync: 171 files from Baida98/AI [deploy-all]
24480a0 verified
Raw
History Blame Contribute Delete
16.2 kB
"""
backend/api/capability_catalog.py β€” Capability Marketplace Catalog (ARCH-E3.1)
Catalogo dinamico dove ogni Worker registra le proprie capabilities con metadati
completi: versione, SLA, latenza target, GPU, costo, tag, disponibilitΓ , regione.
Il Brain NON conosce i Worker β€” chiede una capability, il Kernel + Fabric scelgono.
Il Catalog Γ¨ il registro centrale di discovery; il Fabric usa il Catalog per lo scoring.
Flusso:
Worker β†’ POST /api/catalog/register β†’ entry creata/aggiornata con TTL
Worker β†’ POST /api/catalog/heartbeat β†’ TTL rinnovato
Fabric β†’ (auto) register all'init β†’ fleet registrata automaticamente
Client β†’ GET /api/catalog/capabilities β†’ lista capabilities vive
Client β†’ GET /api/catalog/capabilities/{name} β†’ providers per una capability
Client β†’ GET /api/catalog/status β†’ diagnostica + contatori
Invarianti ADR:
S9: ogni servizio ignora l'impl interna degli altri
S19: nessun vendor lock-in β€” qualsiasi Worker puΓ² registrarsi
S20: routing intent-based, non hardcoded
S27: ogni capability tracciabile via provider_id + correlation_id
"""
from __future__ import annotations
import asyncio
import logging
import time
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
from .auth_guard import AuthRole, require_role
_logger = logging.getLogger("api.capability_catalog")
# ── TTL / cleanup config ────────────────────────────────────────────────────────
_ENTRY_TTL_S: float = float(__import__("os").getenv("CATALOG_TTL_S", "300")) # 5 min
_CLEANUP_INTERVAL_S: int = int(__import__("os").getenv("CATALOG_CLEANUP_S", "60")) # 1 min
# ── Models ──────────────────────────────────────────────────────────────────────
class CapabilityDescriptor(BaseModel):
"""Descrittore completo di una singola capability offerta da un provider."""
name: str = Field(..., description="Nome capability, es. 'browser', 'python_sandbox'")
version: str = Field("1.0.0", description="Versione semantica (semver)")
provider_id: str = Field(..., description="ID univoco del provider")
provider_name: str = Field("", description="Nome human-readable del provider")
description: str = Field("", description="Descrizione funzionale breve")
tags: list[str] = Field(default_factory=list, description="Tag per discovery intent-based")
requires_gpu: bool = Field(False, description="Richiede GPU")
sla_ms: float = Field(5000.0, description="Target latency SLA in ms (p95)")
max_payload_kb: int = Field(1024, description="Payload massimo accettato in KB")
cost_unit: float = Field(0.0, description="Costo per invocazione (0 = free)")
region: str = Field("us", description="Regione di deployment")
always_on: bool = Field(True, description="Provider sempre attivo (no cold start)")
registered_at: float = Field(default_factory=time.time)
last_heartbeat: float = Field(default_factory=time.time)
metadata: dict[str, Any] = Field(default_factory=dict, description="Metadati extra provider-specifici")
class RegisterRequest(BaseModel):
descriptors: list[CapabilityDescriptor] = Field(
..., description="Lista capability da registrare per questo provider"
)
class HeartbeatRequest(BaseModel):
provider_id: str
capability_names: list[str] | None = None # None = tutte le capability del provider
# ── CapabilityCatalog singleton ─────────────────────────────────────────────────
class CapabilityCatalog:
"""
Registro dinamico di tutte le capabilities disponibili nel sistema.
Struttura interna:
_entries: { (provider_id, capability_name) β†’ CapabilityDescriptor }
Thread/task safety: lock asyncio su tutte le mutazioni.
"""
def __init__(self) -> None:
self._entries: dict[tuple[str, str], CapabilityDescriptor] = {}
self._lock = asyncio.Lock()
self._cleanup_task: asyncio.Task | None = None
# ── Registration ──────────────────────────────────────────────────────────
async def register(self, descriptors: list[CapabilityDescriptor]) -> int:
"""Registra/aggiorna N capabilities. Ritorna il numero di entry salvate."""
async with self._lock:
now = time.time()
for d in descriptors:
d.registered_at = now
d.last_heartbeat = now
self._entries[(d.provider_id, d.name)] = d
count = len(descriptors)
_logger.info("[catalog] registered %d capabilities from provider=%s",
count, descriptors[0].provider_id if descriptors else "?")
return count
async def deregister(self, provider_id: str, capability_names: list[str] | None = None) -> int:
"""Rimuove capability di un provider (o subset se specificato)."""
async with self._lock:
to_del = [
k for k in self._entries
if k[0] == provider_id and (capability_names is None or k[1] in capability_names)
]
for k in to_del:
del self._entries[k]
if to_del:
_logger.info("[catalog] deregistered %d capabilities from provider=%s",
len(to_del), provider_id)
return len(to_del)
async def heartbeat(self, provider_id: str, capability_names: list[str] | None = None) -> int:
"""Aggiorna last_heartbeat. Ritorna il numero di entry aggiornate."""
async with self._lock:
now = time.time()
count = 0
for (pid, cname), d in self._entries.items():
if pid == provider_id and (capability_names is None or cname in capability_names):
d.last_heartbeat = now
count += 1
return count
# ── Query ─────────────────────────────────────────────────────────────────
def query(
self,
name: str | None = None,
tags: list[str] | None = None,
requires_gpu: bool | None = None,
max_sla_ms: float | None = None,
region: str | None = None,
include_stale: bool = False,
) -> list[CapabilityDescriptor]:
"""
Ricerca nel catalogo con filtri combinabili.
Di default ritorna solo entry vive (last_heartbeat entro TTL).
"""
now = time.time()
results = []
for d in self._entries.values():
if not include_stale and (now - d.last_heartbeat) > _ENTRY_TTL_S:
continue
if name and d.name != name:
continue
if tags and not any(t in d.tags for t in tags):
continue
if requires_gpu is not None and d.requires_gpu != requires_gpu:
continue
if max_sla_ms is not None and d.sla_ms > max_sla_ms:
continue
if region and d.region != region:
continue
results.append(d)
return results
def get_sla(self, capability_name: str, provider_id: str) -> float:
"""
Ritorna sla_ms per una capability specifica, o 9999.0 se non trovata.
Non-blocking: lookup puro dict β€” safe da chiamare in _select().
"""
d = self._entries.get((provider_id, capability_name))
return d.sla_ms if d else 9999.0
def all_entries(self, include_stale: bool = False) -> list[CapabilityDescriptor]:
"""Lista completa (per diagnostica)."""
if include_stale:
return list(self._entries.values())
now = time.time()
return [d for d in self._entries.values() if (now - d.last_heartbeat) <= _ENTRY_TTL_S]
# ── Cleanup ───────────────────────────────────────────────────────────────
async def cleanup_stale(self) -> int:
"""Rimuove entry con TTL scaduto. Chiamato dal loop interno."""
async with self._lock:
now = time.time()
stale = [k for k, d in self._entries.items()
if (now - d.last_heartbeat) > _ENTRY_TTL_S]
for k in stale:
del self._entries[k]
if stale:
_logger.warning("[catalog] cleanup: removed %d stale entries", len(stale))
return len(stale)
async def _cleanup_loop(self) -> None:
while True:
await asyncio.sleep(_CLEANUP_INTERVAL_S)
try:
await self.cleanup_stale()
except Exception as exc:
_logger.warning("[catalog] cleanup error: %s", exc)
def start_cleanup_loop(self) -> None:
"""Avvia background cleanup. Chiamare in on_startup."""
if self._cleanup_task is None or self._cleanup_task.done():
self._cleanup_task = asyncio.create_task(self._cleanup_loop())
_logger.info("[catalog] cleanup loop started (TTL=%ds, interval=%ds)",
int(_ENTRY_TTL_S), _CLEANUP_INTERVAL_S)
# ── Singleton ───────────────────────────────────────────────────────────────────
catalog = CapabilityCatalog()
# ── HTTP Router ─────────────────────────────────────────────────────────────────
router = APIRouter(
prefix="/api/catalog",
tags=["capability-catalog"],
dependencies=[Depends(require_role(AuthRole.MACHINE))],
)
@router.post("/register", summary="Registra capabilities di un Worker nel catalogo")
async def route_register(req: RegisterRequest) -> dict:
if not req.descriptors:
raise HTTPException(400, "descriptors lista vuota")
count = await catalog.register(req.descriptors)
return {
"registered": count,
"provider_id": req.descriptors[0].provider_id,
}
@router.post("/heartbeat", summary="Rinnova TTL capabilities (keep-alive)")
async def route_heartbeat(req: HeartbeatRequest) -> dict:
count = await catalog.heartbeat(req.provider_id, req.capability_names)
return {"updated": count, "provider_id": req.provider_id}
@router.delete("/providers/{provider_id}", summary="Deregistra capabilities di un provider")
async def route_deregister(provider_id: str) -> dict:
count = await catalog.deregister(provider_id)
return {"removed": count, "provider_id": provider_id}
@router.get("/capabilities", summary="Lista capabilities disponibili con filtri")
async def route_list_capabilities(
name: str | None = None,
tag: str | None = None,
requires_gpu: bool | None = None,
max_sla_ms: float | None = None,
region: str | None = None,
) -> dict:
tags = [tag] if tag else None
entries = catalog.query(name=name, tags=tags, requires_gpu=requires_gpu,
max_sla_ms=max_sla_ms, region=region)
return {
"count": len(entries),
"capabilities": [e.model_dump() for e in entries],
}
@router.get("/capabilities/{capability_name}", summary="Dettaglio capability per nome")
async def route_get_capability(capability_name: str) -> dict:
entries = catalog.query(name=capability_name)
if not entries:
raise HTTPException(404, f"Capability '{capability_name}' non trovata nel catalogo")
best = min(entries, key=lambda e: e.sla_ms)
return {
"capability": capability_name,
"providers": len(entries),
"best_sla_ms": best.sla_ms,
"best_provider": best.provider_id,
"descriptors": [e.model_dump() for e in sorted(entries, key=lambda e: e.sla_ms)],
}
@router.get("/status", summary="Stato del catalogo e contatori")
async def route_status() -> dict:
all_e = catalog.all_entries(include_stale=True)
live = catalog.all_entries()
stale = len(all_e) - len(live)
by_prov: dict[str, int] = {}
for e in live:
by_prov[e.provider_id] = by_prov.get(e.provider_id, 0) + 1
unique_caps = sorted({e.name for e in live})
return {
"total_entries": len(all_e),
"live_entries": len(live),
"stale_entries": stale,
"unique_capabilities": unique_caps,
"providers": by_prov,
"ttl_s": _ENTRY_TTL_S,
"cleanup_interval_s": _CLEANUP_INTERVAL_S,
}
# ── ARCH-E3.4: Worker self-announcement ────────────────────────────────────────
class WorkerAnnouncement(BaseModel):
"""
Payload che ogni Worker invia al boot per auto-registrare le proprie capabilities.
Sostituisce la registrazione manuale β€” il Worker conosce se stesso.
"""
worker_id: str = Field(..., description="ID univoco del Worker, es. 'hf-space-browser'")
worker_name: str = Field(...)
worker_kind: str = Field("http", description="http | grpc | ws")
base_url: str = Field(..., description="URL base del Worker")
capabilities: list[str] = Field(..., description="Lista capability esposte")
region: str = Field("global")
requires_gpu: bool = Field(False)
sla_ms: dict[str, float] = Field(default_factory=dict,
description="SLA per capability, es. {'browser': 8000}")
cost_unit: float = Field(0.0)
version: str = Field("1.0.0")
always_on: bool = Field(True)
tags: list[str] = Field(default_factory=list)
metadata: dict = Field(default_factory=dict)
@router.post("/worker-announce", summary="Worker auto-registra le proprie capabilities al boot (ARCH-E3.4)")
async def route_worker_announce(ann: WorkerAnnouncement) -> dict:
"""
Endpoint chiamato dai Worker all'avvio per registrare capabilities nel Catalog.
Ogni capability riceve un CapabilityDescriptor auto-costruito dall'annuncio.
Esegue anche heartbeat se il worker Γ¨ giΓ  registrato (idempotente).
"""
descs = [
CapabilityDescriptor(
name = cap,
version = ann.version,
provider_id = ann.worker_id,
provider_name = ann.worker_name,
sla_ms = ann.sla_ms.get(cap, 5000.0),
requires_gpu = ann.requires_gpu,
cost_unit = ann.cost_unit,
region = ann.region,
always_on = ann.always_on,
tags = ann.tags + ["worker", ann.worker_kind],
metadata = {"base_url": ann.base_url, **ann.metadata},
)
for cap in ann.capabilities
]
if descs:
await catalog.register(descs)
return {
"announced": True,
"worker_id": ann.worker_id,
"capabilities": ann.capabilities,
"registered": len(descs),
}