Initial public ModelForge release
This commit is contained in:
@@ -0,0 +1,371 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user