Files
planet/aiprovider/main.py
2026-04-28 16:10:17 +08:00

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)