ML_AOTS / app.py
Raja27
Add missing app entrypoint
2f8c53f
Raw
History Blame Contribute Delete
3.78 kB
from __future__ import annotations
import os
from pathlib import Path
from typing import Any
import joblib
from fastapi import FastAPI, Header, HTTPException
from pydantic import BaseModel, Field
from aots_pipeline import DEFAULT_CONFIDENCE_THRESHOLD, predict_records
MODEL_PATH = Path(os.getenv("MODEL_PATH", "artifacts/aots_technology_model.joblib"))
SPACE_API_KEY = os.getenv("SPACE_API_KEY")
class AlertRequest(BaseModel):
aots_number: str | None = Field(default=None, description="Ticket or AOTS number.")
asset_id: str = Field(default="", description="Asset ID from the alert email.")
problem_abstract: str = Field(
default="",
description="Problem abstract from the alert email.",
)
confidence_threshold: float | None = Field(
default=None,
ge=0.0,
le=1.0,
description="Optional override for manual review routing.",
)
top_k: int = Field(
default=3,
ge=1,
le=5,
description="How many ranked predictions to return.",
)
class BatchRequest(BaseModel):
alerts: list[AlertRequest]
confidence_threshold: float | None = Field(default=None, ge=0.0, le=1.0)
top_k: int = Field(default=3, ge=1, le=5)
app = FastAPI(
title="AOTS Technology Predictor",
version="1.0.0",
description="Predicts the storage/network technology for an AOTS alert.",
)
_model_bundle: dict[str, Any] | None = None
def get_model_bundle() -> dict[str, Any]:
global _model_bundle
if _model_bundle is None:
if not MODEL_PATH.exists():
raise HTTPException(
status_code=503,
detail=(
f"Model bundle not found at '{MODEL_PATH}'. "
"Train the model with train_aots_model.py before calling /predict."
),
)
_model_bundle = joblib.load(MODEL_PATH)
return _model_bundle
def validate_api_key(x_api_key: str | None) -> None:
if not SPACE_API_KEY:
return
if x_api_key != SPACE_API_KEY:
raise HTTPException(status_code=401, detail="Invalid or missing API key.")
@app.get("/health")
def health(x_api_key: str | None = Header(default=None)) -> dict[str, Any]:
validate_api_key(x_api_key)
model_ready = MODEL_PATH.exists()
threshold = DEFAULT_CONFIDENCE_THRESHOLD
if model_ready:
bundle = get_model_bundle()
threshold = bundle.get("confidence_threshold", DEFAULT_CONFIDENCE_THRESHOLD)
return {
"status": "ok" if model_ready else "degraded",
"model_path": str(MODEL_PATH),
"model_ready": model_ready,
"confidence_threshold": threshold,
}
@app.post("/predict")
def predict(
request: AlertRequest,
x_api_key: str | None = Header(default=None),
) -> dict[str, Any]:
validate_api_key(x_api_key)
bundle = get_model_bundle()
result = predict_records(
bundle=bundle,
records=[request.model_dump(exclude_none=True)],
confidence_threshold=request.confidence_threshold,
top_k=request.top_k,
)[0]
return result
@app.post("/predict-batch")
def predict_batch(
request: BatchRequest,
x_api_key: str | None = Header(default=None),
) -> dict[str, Any]:
validate_api_key(x_api_key)
bundle = get_model_bundle()
results = predict_records(
bundle=bundle,
records=[alert.model_dump(exclude_none=True) for alert in request.alerts],
confidence_threshold=request.confidence_threshold,
top_k=request.top_k,
)
return {"count": len(results), "results": results}
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"app:app",
host="0.0.0.0",
port=int(os.getenv("PORT", "7860")),
reload=False,
)