80 lines
2.2 KiB
Python
80 lines
2.2 KiB
Python
from uuid import uuid4
|
|
|
|
from fastapi import Depends, FastAPI, Header, HTTPException, Request, Response, status
|
|
|
|
from aiprovider.config import settings
|
|
from aiprovider.provider_service import ProviderService
|
|
from aiprovider.schemas import (
|
|
AIProviderStatusResponse,
|
|
SituationalAnalysisRequest,
|
|
SituationalAnalysisResponse,
|
|
)
|
|
|
|
app = FastAPI(
|
|
title=settings.SERVICE_NAME,
|
|
version=settings.SERVICE_VERSION,
|
|
description="AI provider adapter service for Planet",
|
|
)
|
|
|
|
|
|
@app.middleware("http")
|
|
async def request_id_middleware(request: Request, call_next):
|
|
request_id = request.headers.get("X-Request-ID") or str(uuid4())
|
|
request.state.request_id = request_id
|
|
response = await call_next(request)
|
|
response.headers["X-Request-ID"] = request_id
|
|
return response
|
|
|
|
|
|
def verify_service_token(x_provider_token: str | None = Header(default=None)) -> None:
|
|
expected = settings.AI_PROVIDER_SERVICE_TOKEN
|
|
if not expected:
|
|
return
|
|
if x_provider_token != expected:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid provider service token",
|
|
)
|
|
|
|
|
|
def get_provider_service() -> ProviderService:
|
|
return ProviderService()
|
|
|
|
|
|
@app.get("/health")
|
|
async def health_check():
|
|
return {
|
|
"status": "healthy",
|
|
"service": settings.SERVICE_NAME,
|
|
"version": settings.SERVICE_VERSION,
|
|
}
|
|
|
|
|
|
@app.get(
|
|
"/v1/provider/status",
|
|
response_model=AIProviderStatusResponse,
|
|
dependencies=[Depends(verify_service_token)],
|
|
)
|
|
async def get_provider_status(
|
|
response: Response,
|
|
request: Request,
|
|
provider_service: ProviderService = Depends(get_provider_service),
|
|
):
|
|
response.headers["X-Request-ID"] = request.state.request_id
|
|
return provider_service.get_status()
|
|
|
|
|
|
@app.post(
|
|
"/v1/analyze",
|
|
response_model=SituationalAnalysisResponse,
|
|
dependencies=[Depends(verify_service_token)],
|
|
)
|
|
async def analyze(
|
|
payload: SituationalAnalysisRequest,
|
|
response: Response,
|
|
request: Request,
|
|
provider_service: ProviderService = Depends(get_provider_service),
|
|
):
|
|
response.headers["X-Request-ID"] = request.state.request_id
|
|
return await provider_service.analyze(payload)
|