98 lines
3.0 KiB
Python
98 lines
3.0 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(
|
|
x_ai_provider: str | None = Header(default=None),
|
|
x_ai_provider_api: str | None = Header(default=None),
|
|
x_ai_base_url: str | None = Header(default=None),
|
|
x_ai_api_key: str | None = Header(default=None),
|
|
x_ai_model: str | None = Header(default=None),
|
|
x_ai_max_tokens: str | None = Header(default=None),
|
|
x_ai_anthropic_version: str | None = Header(default=None),
|
|
) -> ProviderService:
|
|
overrides = {
|
|
"provider": x_ai_provider,
|
|
"provider_api": x_ai_provider_api,
|
|
"base_url": x_ai_base_url,
|
|
"api_key": x_ai_api_key,
|
|
"model": x_ai_model,
|
|
"anthropic_version": x_ai_anthropic_version,
|
|
}
|
|
if x_ai_max_tokens:
|
|
overrides["max_tokens"] = x_ai_max_tokens
|
|
return ProviderService({key: value for key, value in overrides.items() if value not in (None, "")})
|
|
|
|
|
|
@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)
|