Skip to content
← How we build

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.

No spam. Unsubscribe anytime.