Files
ModelForge/backend/src/modelforge_api/api/routes/agent.py
T

372 lines
11 KiB
Python

from __future__ import annotations
import time
import uuid
from collections import defaultdict, deque
from typing import Annotated
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
from pydantic import BaseModel
from sqlalchemy.orm import Session
from modelforge_api.api.authorization import Admin
from modelforge_api.db import get_session
from modelforge_api.domain.acquisition import (
AgentArtifactJobLease,
AgentJobComplete,
AgentJobControl,
AgentJobFailure,
AgentJobProgress,
ArtifactJobResponse,
)
from modelforge_api.domain.agent_protocol import (
EnrollmentRequest,
EnrollmentResponse,
EnrollmentTokenCreate,
EnrollmentTokenCreated,
EnrollmentTokenSummary,
HeartbeatRequest,
InventoryReport,
NodeCredentialCreated,
NodeManagementUpdate,
ObservationAck,
TelemetryReport,
)
from modelforge_api.domain.node_decommission import (
NodeDecommissionExecute,
NodeDecommissionPreview,
NodeDecommissionResult,
)
from modelforge_api.domain.runtime import (
AgentRuntimeProbeComplete,
AgentRuntimeProbeControl,
AgentRuntimeProbeFailure,
AgentRuntimeProbeLease,
AgentRuntimeProbeProgress,
RuntimeProbeResponse,
)
from modelforge_api.persistence.models import ComputeNode, NodeCredential
from modelforge_api.providers.huggingface import OfficialHuggingFaceProvider
from modelforge_api.services.acquisition import AcquisitionService
from modelforge_api.services.node_agent import NodeAgentService, NodeAuthenticationEvidence
from modelforge_api.services.node_decommission import NodeDecommissionService
from modelforge_api.services.runtime import RuntimeService
from modelforge_api.settings import Settings, get_settings
router = APIRouter(tags=["node-agent"])
class ActionResponse(BaseModel):
status: str = "ok"
class AttemptLimiter:
def __init__(self, limit: int = 10, window_seconds: int = 60) -> None:
self.limit = limit
self.window_seconds = window_seconds
self.attempts: dict[str, deque[float]] = defaultdict(deque)
def check(self, key: str) -> None:
now = time.monotonic()
bucket = self.attempts[key]
while bucket and bucket[0] <= now - self.window_seconds:
bucket.popleft()
if len(bucket) >= self.limit:
raise HTTPException(status_code=429, detail="enrollment rate limit exceeded")
bucket.append(now)
enrollment_limiter = AttemptLimiter()
def get_agent_service(
session: Annotated[Session, Depends(get_session)],
settings: Annotated[Settings, Depends(get_settings)],
) -> NodeAgentService:
return NodeAgentService(session, settings)
Service = Annotated[NodeAgentService, Depends(get_agent_service)]
def get_decommission_service(
session: Annotated[Session, Depends(get_session)],
) -> NodeDecommissionService:
return NodeDecommissionService(session)
DecommissionService = Annotated[NodeDecommissionService, Depends(get_decommission_service)]
@router.post(
"/api/v1/admin/node-enrollments",
response_model=EnrollmentTokenCreated,
status_code=status.HTTP_201_CREATED,
)
def create_enrollment(
request: EnrollmentTokenCreate, service: Service, _admin: Admin, response: Response
) -> EnrollmentTokenCreated:
response.headers["Cache-Control"] = "no-store"
return service.create_enrollment(request)
@router.get("/api/v1/admin/node-enrollments", response_model=list[EnrollmentTokenSummary])
def list_enrollments(service: Service, _admin: Admin) -> list[EnrollmentTokenSummary]:
return [
EnrollmentTokenSummary(
id=str(item.id),
created_at=item.created_at,
expires_at=item.expires_at,
used_at=item.used_at,
revoked_at=item.revoked_at,
)
for item in service.repository.enrollments()
]
@router.delete("/api/v1/admin/node-enrollments/{enrollment_id}", response_model=ActionResponse)
def revoke_enrollment(enrollment_id: uuid.UUID, service: Service, _admin: Admin) -> ActionResponse:
service.revoke_enrollment(enrollment_id)
return ActionResponse()
@router.patch("/api/v1/admin/hardware/nodes/{node_id}", response_model=ActionResponse)
def update_node(
node_id: uuid.UUID,
request: NodeManagementUpdate,
service: Service,
_admin: Admin,
) -> ActionResponse:
service.update_node(node_id, request)
return ActionResponse()
@router.post(
"/api/v1/admin/hardware/nodes/{node_id}/decommission/preview",
response_model=NodeDecommissionPreview,
)
def preview_node_decommission(
node_id: uuid.UUID, service: DecommissionService, _admin: Admin
) -> NodeDecommissionPreview:
return service.preview(node_id)
@router.post(
"/api/v1/admin/hardware/nodes/{node_id}/decommission",
response_model=NodeDecommissionResult,
)
def execute_node_decommission(
node_id: uuid.UUID,
request: NodeDecommissionExecute,
service: DecommissionService,
_admin: Admin,
) -> NodeDecommissionResult:
return service.execute(node_id, request)
@router.delete("/api/v1/admin/hardware/nodes/{node_id}/credential", response_model=ActionResponse)
def revoke_credential(node_id: uuid.UUID, service: Service, _admin: Admin) -> ActionResponse:
service.revoke_credential(node_id)
return ActionResponse()
@router.post(
"/api/v1/admin/hardware/nodes/{node_id}/credential/rotate",
response_model=NodeCredentialCreated,
)
def rotate_credential(
node_id: uuid.UUID, service: Service, _admin: Admin, response: Response
) -> NodeCredentialCreated:
response.headers["Cache-Control"] = "no-store"
return service.rotate_credential(node_id)
@router.post(
"/api/v1/agent/enroll",
response_model=EnrollmentResponse,
status_code=status.HTTP_201_CREATED,
)
def enroll(
request: EnrollmentRequest, http_request: Request, service: Service, response: Response
) -> EnrollmentResponse:
enrollment_limiter.check(http_request.client.host if http_request.client else "unknown")
response.headers["Cache-Control"] = "no-store"
return service.enroll(request)
def authenticated_node(
request: Request,
service: Service,
authorization: Annotated[str | None, Header(alias="Authorization")] = None,
) -> tuple[NodeCredential, ComputeNode]:
evidence = getattr(request.state, "node_authentication", None)
if isinstance(evidence, NodeAuthenticationEvidence):
return service.reuse_authentication(evidence, authorization)
return service.authenticate(authorization)
NodeIdentity = tuple[NodeCredential, ComputeNode]
AgentIdentity = Annotated[NodeIdentity, Depends(authenticated_node)]
def get_agent_acquisition_service(
session: Annotated[Session, Depends(get_session)],
settings: Annotated[Settings, Depends(get_settings)],
) -> AcquisitionService:
token = settings.hf_token.get_secret_value() if settings.hf_token else None
return AcquisitionService(
session,
settings,
OfficialHuggingFaceProvider(token=token, timeout=settings.hf_timeout_seconds),
actor_type="node_agent",
actor_id="authenticated-node",
)
AgentAcquisition = Annotated[AcquisitionService, Depends(get_agent_acquisition_service)]
def get_agent_runtime_service(
session: Annotated[Session, Depends(get_session)],
settings: Annotated[Settings, Depends(get_settings)],
) -> RuntimeService:
return RuntimeService(
session,
settings,
actor_type="runtime_worker",
actor_id="authenticated-node",
)
AgentRuntime = Annotated[RuntimeService, Depends(get_agent_runtime_service)]
@router.post("/api/v1/agent/heartbeat", response_model=ObservationAck)
def heartbeat(
request: HeartbeatRequest, service: Service, identity: AgentIdentity
) -> ObservationAck:
_credential, node = identity
return service.heartbeat(node, request)
@router.put("/api/v1/agent/inventory", response_model=ObservationAck)
def publish_inventory(
request: InventoryReport, service: Service, identity: AgentIdentity
) -> ObservationAck:
_credential, node = identity
return service.publish_inventory(node, request)
@router.put("/api/v1/agent/telemetry", response_model=ObservationAck)
def publish_telemetry(
request: TelemetryReport, service: Service, identity: AgentIdentity
) -> ObservationAck:
_credential, node = identity
return service.publish_telemetry(node, request)
@router.get(
"/api/v1/agent/artifact-jobs/next",
response_model=AgentArtifactJobLease | None,
)
def claim_artifact_job(
acquisition: AgentAcquisition, identity: AgentIdentity
) -> AgentArtifactJobLease | None:
_credential, node = identity
return acquisition.claim_next(node)
@router.post(
"/api/v1/agent/artifact-jobs/{job_id}/progress",
response_model=AgentJobControl,
)
def artifact_job_progress(
job_id: uuid.UUID,
request: AgentJobProgress,
acquisition: AgentAcquisition,
identity: AgentIdentity,
) -> AgentJobControl:
_credential, node = identity
return acquisition.progress(job_id, node, request)
@router.post(
"/api/v1/agent/artifact-jobs/{job_id}/complete",
response_model=ArtifactJobResponse,
)
def artifact_job_complete(
job_id: uuid.UUID,
request: AgentJobComplete,
acquisition: AgentAcquisition,
identity: AgentIdentity,
) -> ArtifactJobResponse:
_credential, node = identity
return acquisition.complete(job_id, node, request)
@router.post(
"/api/v1/agent/artifact-jobs/{job_id}/fail",
response_model=ArtifactJobResponse,
)
def artifact_job_fail(
job_id: uuid.UUID,
request: AgentJobFailure,
acquisition: AgentAcquisition,
identity: AgentIdentity,
) -> ArtifactJobResponse:
_credential, node = identity
return acquisition.fail(job_id, node, request)
@router.get(
"/api/v1/agent/runtime-probes/next",
response_model=AgentRuntimeProbeLease | None,
)
def claim_runtime_probe(
runtime: AgentRuntime, identity: AgentIdentity
) -> AgentRuntimeProbeLease | None:
_credential, node = identity
return runtime.claim_next(node)
@router.post(
"/api/v1/agent/runtime-probes/{probe_id}/progress",
response_model=AgentRuntimeProbeControl,
)
def runtime_probe_progress(
probe_id: uuid.UUID,
request: AgentRuntimeProbeProgress,
runtime: AgentRuntime,
identity: AgentIdentity,
) -> AgentRuntimeProbeControl:
_credential, node = identity
return runtime.progress(probe_id, node, request)
@router.post(
"/api/v1/agent/runtime-probes/{probe_id}/complete",
response_model=RuntimeProbeResponse,
)
def runtime_probe_complete(
probe_id: uuid.UUID,
request: AgentRuntimeProbeComplete,
runtime: AgentRuntime,
identity: AgentIdentity,
) -> RuntimeProbeResponse:
_credential, node = identity
return runtime.complete(probe_id, node, request)
@router.post(
"/api/v1/agent/runtime-probes/{probe_id}/fail",
response_model=RuntimeProbeResponse,
)
def runtime_probe_fail(
probe_id: uuid.UUID,
request: AgentRuntimeProbeFailure,
runtime: AgentRuntime,
identity: AgentIdentity,
) -> RuntimeProbeResponse:
_credential, node = identity
return runtime.fail(probe_id, node, request)