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)