from __future__ import annotations import asyncio import json from time import perf_counter import httpx from fastapi import Depends, HTTPException, status from sqlalchemy.ext.asyncio import AsyncSession from app.core.config import settings from app.core.logging import get_logger from app.db.session import get_db from app.schemas.ai import ( AIProviderStatusResponse, SituationalAnalysisRequest, SituationalAnalysisResponse, ) from app.services.business_logs import emit_business_log, exception_context logger = get_logger(__name__, service="ai") class AIProviderClient: def __init__( self, *, service_url: str | None = None, service_token: str | None = None, timeout: int | None = None, retry_attempts: int | None = None, llm_config: dict | None = None, ) -> None: self.service_url = ( service_url if service_url is not None else settings.AI_PROVIDER_SERVICE_URL ).rstrip("/") self.service_token = ( service_token if service_token is not None else settings.AI_PROVIDER_SERVICE_TOKEN ) self.timeout = timeout if timeout is not None else settings.AI_PROVIDER_TIMEOUT_SECONDS self.retry_attempts = max( retry_attempts if retry_attempts is not None else settings.AI_PROVIDER_RETRY_ATTEMPTS, 1, ) self.llm_config = llm_config or {} 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 llm_header_map = { "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", "max_tokens": "X-AI-Max-Tokens", "anthropic_version": "X-AI-Anthropic-Version", } for key, header_name in llm_header_map.items(): value = self.llm_config.get(key) if value not in (None, ""): headers[header_name] = str(value) model_provider_apis = self.llm_config.get("model_provider_apis") if isinstance(model_provider_apis, dict) and model_provider_apis: headers["X-AI-Model-Provider-APIs"] = json.dumps(model_provider_apis) return headers async def get_status(self, request_id: str | None = None) -> AIProviderStatusResponse: context = self._base_log_context(operation="status") if not self.service_url: await emit_business_log( logger, event="ai.provider.status.failed", message="AI provider status skipped because service URL is not configured", category="ai", level="warning", service="ai", module=__name__, request_id=request_id, context={**context, "status": "unconfigured"}, ) return AIProviderStatusResponse( provider="unconfigured", enabled=False, configured=False, model=None, base_url=None, ) started_at = perf_counter() await emit_business_log( logger, event="ai.provider.status.start", message="AI provider status request started", category="ai", service="ai", module=__name__, request_id=request_id, context=context, ) try: data = await self._request("GET", "/v1/provider/status", request_id=request_id, operation="status") result = AIProviderStatusResponse.model_validate(data) await emit_business_log( logger, event="ai.provider.status.success", message="AI provider status request completed", category="ai", service="ai", module=__name__, request_id=request_id, context={ **context, "status": "success", "duration_ms": self._duration_ms(started_at), "result_provider": result.provider, "result_model": result.model, "configured": result.configured, "enabled": result.enabled, }, ) return result except Exception as exc: await emit_business_log( logger, event="ai.provider.status.failed", message="AI provider status request failed", category="ai", level="error", service="ai", module=__name__, request_id=request_id, context=exception_context(exc, {**context, "status": "failed", "duration_ms": self._duration_ms(started_at)}), ) raise async def analyze( self, payload: SituationalAnalysisRequest, request_id: str | None = None, ) -> SituationalAnalysisResponse: context = self._base_log_context( operation="analyze", preferred_model=payload.preferred_model, input_summary=self._summarize_analysis_payload(payload), ) if not self.service_url: await emit_business_log( logger, event="ai.provider.analyze.failed", message="AI provider analyze skipped because service URL is not configured", category="ai", level="warning", service="ai", module=__name__, request_id=request_id, context={**context, "status": "unconfigured"}, ) raise HTTPException( status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="AI provider service URL is not configured.", ) started_at = perf_counter() await emit_business_log( logger, event="ai.provider.analyze.start", message="AI provider analyze request started", category="ai", service="ai", module=__name__, request_id=request_id, context=context, ) try: data = await self._request( "POST", "/v1/analyze", json=payload.model_dump(), request_id=request_id, operation="analyze", payload_summary=context["input_summary"], ) result = SituationalAnalysisResponse.model_validate(data) await emit_business_log( logger, event="ai.provider.analyze.success", message="AI provider analyze request completed", category="ai", service="ai", module=__name__, request_id=request_id, context={ **context, "status": "success", "duration_ms": self._duration_ms(started_at), "result_provider": result.provider, "result_model": result.model, "content_block_count": len(result.content_blocks or []), "thinking_block_count": len(result.thinking_blocks or []), }, ) return result except Exception as exc: await emit_business_log( logger, event="ai.provider.analyze.failed", message="AI provider analyze request failed", category="ai", level="error", service="ai", module=__name__, request_id=request_id, context=exception_context(exc, {**context, "status": "failed", "duration_ms": self._duration_ms(started_at)}), ) raise async def _request( self, method: str, path: str, json: dict | None = None, request_id: str | None = None, operation: str = "request", payload_summary: dict | None = None, ) -> dict: last_error: Exception | None = None for attempt in range(1, self.retry_attempts + 1): attempt_started_at = perf_counter() 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 self._log_retry( operation=operation, request_id=request_id, attempt=attempt, status_code=exc.response.status_code, duration_ms=self._duration_ms(attempt_started_at), error=exc, payload_summary=payload_summary, ) 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 self._log_retry( operation=operation, request_id=request_id, attempt=attempt, duration_ms=self._duration_ms(attempt_started_at), error=exc, payload_summary=payload_summary, ) 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 _base_log_context(self, **extra: object) -> dict: llm_provider_apis = self.llm_config.get("model_provider_apis") return { "provider": self.llm_config.get("provider") or "", "provider_api": self.llm_config.get("provider_api") or "", "model": self.llm_config.get("model") or "", "base_url_configured": bool(self.llm_config.get("base_url")), "service_url_configured": bool(self.service_url), "timeout_seconds": self.timeout, "retry_attempts": self.retry_attempts, "model_provider_api_count": len(llm_provider_apis or {}) if isinstance(llm_provider_apis, dict) else 0, **extra, } @staticmethod def _duration_ms(started_at: float) -> int: return int((perf_counter() - started_at) * 1000) @staticmethod def _summarize_analysis_payload(payload: SituationalAnalysisRequest) -> dict: context = payload.context if isinstance(payload.context, dict) else {} thinking = payload.thinking if isinstance(payload.thinking, dict) else payload.thinking return { "title_length": len(payload.title or ""), "objective_length": len(payload.objective or ""), "observation_count": len(payload.observations or []), "constraint_count": len(payload.constraints or []), "has_system_prompt": bool(payload.system_prompt), "thinking_enabled": bool(thinking), "context_keys": sorted(str(key) for key in context.keys()), } async def _log_retry( self, *, operation: str, request_id: str | None, attempt: int, duration_ms: int, error: BaseException, status_code: int | None = None, payload_summary: dict | None = None, ) -> None: await emit_business_log( logger, event=f"ai.provider.{operation}.retry", message="AI provider request will retry", category="ai", level="warning", service="ai", module=__name__, request_id=request_id, context=exception_context( error, { **self._base_log_context(operation=operation), "attempt": attempt, "next_attempt": attempt + 1, "status_code": status_code, "duration_ms": duration_ms, "input_summary": payload_summary, }, ), ) async def get_ai_provider_client(db: AsyncSession = Depends(get_db)) -> AIProviderClient: from app.api.v1.settings import get_runtime_ai_provider_config runtime_config = await get_runtime_ai_provider_config(db) return AIProviderClient( service_url=runtime_config["service_url"], service_token=runtime_config["service_token"], timeout=runtime_config["timeout_seconds"], retry_attempts=runtime_config["retry_attempts"], llm_config=runtime_config.get("llm_config") or {}, )