backend/src/services/llm_service.py
"""LLM Service — policy layer over llm/factory.py (retry, failover, cost).""" from __future__ import annotations import logging import time from typing import Callable from api.schemas.chat import Message from config.llm_pricing import estimate_cost_usd from config.settings import Settings, get_settings from llm.factory import LLMFactory from llm.providers.base_provider import BaseProvider from models.llm import LLMRequest, LLMResponse from observability import log_event, record_metric logger = logging.getLogger("agentic.request") trace_logger = logging.getLogger("agentic.trace") # Static placeholder until a real prompt-template registry exists. PROMPT_VERSION_PLACEHOLDER = "v1" # Approximate tokens when providers do not return usage (Phase 1). _CHARS_PER_TOKEN = 4 FactoryCreate = Callable[[str, str], BaseProvider] class LLMServiceError(RuntimeError): """Raised when all providers/retries are exhausted.""" class LLMService: """ Platform inference entry point for Skills. Wraps ``llm/factory.py`` without modifying provider internals. Owns model selection, fixed-count retries, ordered failover, and token/cost accounting. """ def __init__( self, settings: Settings | None = None, *, factory_create: FactoryCreate | None = None, ) -> None: self._settings = settings or get_settings() self._factory_create = factory_create or LLMFactory.create async def generate(self, request: LLMRequest) -> LLMResponse: candidates = self._build_candidate_chain(request.model_hint) total_retries = 0 last_error: Exception | None = None started = time.perf_counter() trace_id = request.context.trace_id if request.context is not None else None log_event( "generate_start", component="LLMService", trace_id=trace_id, fields={ "candidates": [f"{p}:{m}" for p, m in candidates], "max_retries": self._settings.llm_max_retries, "prompt_chars": len(request.prompt), }, ) if request.context is not None: trace_logger.debug( "llm_generate_context trace_id=%s session_id=%s", request.context.trace_id, request.context.session_id, ) for index, (provider_name, model) in enumerate(candidates): attempts = self._settings.llm_max_retries + 1 for attempt in range(attempts): try: content = await self._call_provider(provider_name, model, request.prompt) latency_ms = (time.perf_counter() - started) * 1000 input_tokens = self._estimate_tokens(request.prompt) output_tokens = self._estimate_tokens(content) cost = estimate_cost_usd( provider_name, model, input_tokens, output_tokens ) # Retries = failed attempts before this success (not counting the success). response = LLMResponse( content=content, model_used=f"{provider_name}:{model}", input_tokens=input_tokens, output_tokens=output_tokens, latency_ms=round(latency_ms, 2), estimated_cost_usd=cost, retries_used=total_retries, metadata={"prompt_version": PROMPT_VERSION_PLACEHOLDER}, ) log_event( "generate_success", component="LLMService", trace_id=trace_id, fields={ "provider": provider_name, "model": model, "retries": total_retries, "input_tokens": input_tokens, "output_tokens": output_tokens, "cost_usd": cost, "latency_ms": response.latency_ms, }, ) record_metric( "llm.latency_ms", latency_ms, tags={"provider": provider_name, "model": model}, ) record_metric( "llm.estimated_cost_usd", cost, tags={"provider": provider_name, "model": model}, ) trace_logger.debug( "llm_generate_detail candidate_index=%s attempt=%s prompt_version=%s", index, attempt + 1, PROMPT_VERSION_PLACEHOLDER, ) return response except Exception as exc: total_retries += 1 last_error = exc log_event( "generate_attempt_failed", component="LLMService", trace_id=trace_id, fields={ "provider": provider_name, "model": model, "attempt": f"{attempt + 1}/{attempts}", "error": str(exc), }, level=logging.WARNING, ) trace_logger.debug( "llm_generate_retry provider=%s model=%s attempt=%s error=%s", provider_name, model, attempt + 1, exc, exc_info=True, ) raise LLMServiceError( f"All LLM providers failed after {total_retries} attempt(s). " f"Last error: {last_error}" ) from last_error def _build_candidate_chain(self, model_hint: str | None) -> list[tuple[str, str]]: primary = self._resolve_selection(model_hint) chain: list[tuple[str, str]] = [primary] seen = {primary} for entry in self._settings.llm_failover_order: entry = entry.strip() if not entry: continue candidate = self._parse_provider_model(entry, allow_provider_only=True) if candidate not in seen: chain.append(candidate) seen.add(candidate) return chain def _resolve_selection(self, model_hint: str | None) -> tuple[str, str]: if model_hint and model_hint.strip(): return self._parse_provider_model(model_hint.strip(), allow_provider_only=False) return self._settings.default_provider, self._settings.default_model def _parse_provider_model( self, value: str, *, allow_provider_only: bool, ) -> tuple[str, str]: """ Parse ``provider:model`` or a bare provider/model string. ``model_hint`` may be ``provider:model`` or a bare model name (keeps default provider). Failover entries may be ``provider`` alone (uses that provider's default model from settings) or ``provider:model``. OpenRouter-style ids that contain ``:`` (e.g. ``...:free``) are treated as model names unless the prefix before the first ``:`` is a known provider. """ known = {name.lower() for name in LLMFactory.available_providers()} lowered = value.lower() if allow_provider_only and lowered in known: return lowered, self._default_model_for(lowered) if ":" in value: provider, model = value.split(":", 1) provider = provider.strip().lower() model = model.strip() if provider in known and model: return provider, model return self._settings.default_provider, value return self._settings.default_provider, value def _default_model_for(self, provider: str) -> str: if provider == "openrouter": return self._settings.openrouter_default_model return self._settings.default_model async def _call_provider(self, provider_name: str, model: str, prompt: str) -> str: provider = self._factory_create(provider_name, model) messages = [Message(role="user", content=prompt)] reply = await provider.complete(messages=messages, system_prompt=None) return reply.content @staticmethod def _estimate_tokens(text: str) -> int: if not text: return 0 return max(1, len(text) // _CHARS_PER_TOKEN)
Follow the work.
Occasional updates on SignalFoundry, MarketCompass, and what we are building at CompassFoundry Labs.