110 lines
3.8 KiB
Python
110 lines
3.8 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
|
|
import httpx
|
|
from fastapi import HTTPException, status
|
|
|
|
from app.core.config import settings
|
|
from app.schemas.ai import (
|
|
AIProviderStatusResponse,
|
|
SituationalAnalysisRequest,
|
|
SituationalAnalysisResponse,
|
|
)
|
|
|
|
|
|
class AIProviderClient:
|
|
def __init__(self) -> None:
|
|
self.service_url = settings.AI_PROVIDER_SERVICE_URL.rstrip("/")
|
|
self.service_token = settings.AI_PROVIDER_SERVICE_TOKEN
|
|
self.timeout = settings.AI_PROVIDER_TIMEOUT_SECONDS
|
|
self.retry_attempts = max(settings.AI_PROVIDER_RETRY_ATTEMPTS, 1)
|
|
|
|
def _headers(self, request_id: str | None = None) -> dict[str, str]:
|
|
headers = {"Content-Type": "application/json"}
|
|
if self.service_token:
|
|
headers["X-Provider-Token"] = self.service_token
|
|
if request_id:
|
|
headers["X-Request-ID"] = request_id
|
|
return headers
|
|
|
|
async def get_status(self, request_id: str | None = None) -> AIProviderStatusResponse:
|
|
if not self.service_url:
|
|
return AIProviderStatusResponse(
|
|
provider="unconfigured",
|
|
enabled=False,
|
|
configured=False,
|
|
model=None,
|
|
base_url=None,
|
|
)
|
|
|
|
data = await self._request("GET", "/v1/provider/status", request_id=request_id)
|
|
return AIProviderStatusResponse.model_validate(data)
|
|
|
|
async def analyze(
|
|
self,
|
|
payload: SituationalAnalysisRequest,
|
|
request_id: str | None = None,
|
|
) -> SituationalAnalysisResponse:
|
|
if not self.service_url:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
|
detail="AI provider service URL is not configured.",
|
|
)
|
|
|
|
data = await self._request(
|
|
"POST",
|
|
"/v1/analyze",
|
|
json=payload.model_dump(),
|
|
request_id=request_id,
|
|
)
|
|
return SituationalAnalysisResponse.model_validate(data)
|
|
|
|
async def _request(
|
|
self,
|
|
method: str,
|
|
path: str,
|
|
json: dict | None = None,
|
|
request_id: str | None = None,
|
|
) -> dict:
|
|
last_error: Exception | None = None
|
|
for attempt in range(1, self.retry_attempts + 1):
|
|
try:
|
|
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
|
response = await client.request(
|
|
method,
|
|
f"{self.service_url}{path}",
|
|
headers=self._headers(request_id),
|
|
json=json,
|
|
)
|
|
response.raise_for_status()
|
|
return response.json()
|
|
except httpx.HTTPStatusError as exc:
|
|
last_error = exc
|
|
if attempt < self.retry_attempts and exc.response.status_code >= 500:
|
|
await asyncio.sleep(0.3 * attempt)
|
|
continue
|
|
detail = exc.response.text or "AI provider service returned an error"
|
|
raise HTTPException(
|
|
status_code=status.HTTP_502_BAD_GATEWAY,
|
|
detail=f"AI provider service request failed: {detail}",
|
|
) from exc
|
|
except httpx.HTTPError as exc:
|
|
last_error = exc
|
|
if attempt < self.retry_attempts:
|
|
await asyncio.sleep(0.3 * attempt)
|
|
continue
|
|
raise HTTPException(
|
|
status_code=status.HTTP_502_BAD_GATEWAY,
|
|
detail=f"Failed to reach AI provider service: {exc}",
|
|
) from exc
|
|
|
|
raise HTTPException(
|
|
status_code=status.HTTP_502_BAD_GATEWAY,
|
|
detail=f"AI provider service request failed: {last_error}",
|
|
)
|
|
|
|
|
|
def get_ai_provider_client() -> AIProviderClient:
|
|
return AIProviderClient()
|