372 lines
11 KiB
Python
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)
|