Files
AItrackwalker/ai-services/vision-inference/app/main.py
T
2026-07-22 15:09:42 +08:00

80 lines
2.4 KiB
Python

from __future__ import annotations
from fastapi import FastAPI
from .engines import VisionInferenceEngine
from .schemas import (
HealthResponse,
InferenceRequest,
InferenceResponse,
ModelTestInferenceRequest,
ModelTestInferenceResponse,
)
app = FastAPI(title="Rail UAV Vision Inference Service", version="0.1.0")
engine = VisionInferenceEngine()
@app.get("/health", response_model=HealthResponse)
def health() -> HealthResponse:
return HealthResponse(status="ok")
@app.get("/api/v1/runtime")
def runtime_status() -> dict:
return engine.runtime_status()
@app.post("/api/v1/runtime/unload")
def unload_models(model_group: str | None = None) -> dict:
return {"unloaded": engine.unload(model_group), "runtime": engine.runtime_status()}
@app.post("/api/v1/runtime/load")
def load_model(model_group: str, model_version: str | None = None) -> dict:
return {"model": engine.load(model_group, model_version), "runtime": engine.runtime_status()}
@app.post("/api/v1/inference/detect", response_model=InferenceResponse)
def detect(request: InferenceRequest) -> InferenceResponse:
selected_group = request.model_group
selected_version = request.model_version or "v1.0.0"
results = (
engine.test_infer(
request.resource,
request.scenes,
selected_group,
selected_version,
request.parameters,
)
if selected_group
else engine.infer(request.resource, request.scenes)
)
if request.deployment_id:
for result in results:
result.attributes["deployment_id"] = request.deployment_id
return InferenceResponse(
job_id=request.job_id,
model={"name": selected_group or "rail-multi-scene-inference", "version": selected_version},
results=results,
)
@app.post("/api/v1/inference/test", response_model=ModelTestInferenceResponse)
def test_inference(request: ModelTestInferenceRequest) -> ModelTestInferenceResponse:
results = engine.test_infer(
request.resource,
request.scenes,
request.model_group,
request.model_version,
request.parameters,
)
if request.deployment_id:
for result in results:
result.attributes["deployment_id"] = request.deployment_id
return ModelTestInferenceResponse(
model={"name": request.model_group, "version": request.model_version},
parameter_snapshot=request.parameters,
results=results,
)