1"""Base definitions and shared helpers for routing strategies."""
2
3from __future__ import annotations
4
5from datetime import UTC, datetime, timedelta
6from typing import TYPE_CHECKING, Any
7
8from lexigram.ai.llm.exceptions import LLMQuotaExceededError, LLMRateLimitError
9from lexigram.ai.llm.routing.types import InferenceResult
10from lexigram.ai.llm.types import ChatMessage, Role
11from lexigram.contracts.ai.multimodal import ImageUrlPart, TextPart
12from lexigram.contracts.web.http_models import HttpStatusError
13from lexigram.logging import (
14 get_logger,
15)
16from lexigram.primitives import clock as ambient_clock
17
18if TYPE_CHECKING:
19 from lexigram.ai.llm.routing.config import (
20 LLMConfig,
21 )
22 from lexigram.ai.llm.types import AIError
23 from lexigram.contracts.ai import LLMClientProtocol
24 from lexigram.contracts.ai.routing import QuotaBackendProtocol
25
26logger = get_logger(__name__)
27
28
29def _extract_status_code(exc: AIError) -> int | None:
30 """Extract the HTTP status code from an ``AIError``, if available.
31
32 Typed errors map directly (rate limit → 429, quota → 402) since clients
33 that return ``Err(...)`` carry no ``__cause__`` chain to inspect.
34 """
35 if isinstance(exc, LLMRateLimitError):
36 return 429
37 if isinstance(exc, LLMQuotaExceededError):
38 return 402
39 status_code: int | None = getattr(exc, "status_code", None)
40 if status_code is None and exc.__cause__ is not None:
41 status_code = getattr(exc.__cause__, "status_code", None)
42 if status_code is None and isinstance(exc.__cause__, HttpStatusError):
43 status_code = exc.__cause__.status
44 return status_code
45
46
47async def _attempt_provider(
48 *,
49 client: LLMClientProtocol,
50 provider_name: str,
51 model: str,
52 messages: list[Any],
53 temperature: float,
54 max_tokens: int | None,
55) -> InferenceResult:
56 """Make a single completion request against *client*.
57
58 Accepts both ``ChatMessage`` objects and OpenAI-compatible dicts;
59 dicts are normalised to ``ChatMessage`` before dispatch.
60 Multimodal content (list payloads with ``image_url``/``text`` parts) is
61 preserved as typed ``ContentPart`` objects so vision models receive the
62 full payload.
63 """
64 t0 = ambient_clock.monotonic()
65
66 chat_messages: list[ChatMessage] = []
67 for msg in messages:
68 if isinstance(msg, ChatMessage):
69 chat_messages.append(msg)
70 elif isinstance(msg, dict):
71 raw_content = msg.get("content", "")
72 if isinstance(raw_content, list):
73 content_parts: list[TextPart | ImageUrlPart] = []
74 for part in raw_content:
75 if not isinstance(part, dict):
76 continue
77 part_type = part.get("type")
78 if part_type == "text":
79 content_parts.append(TextPart(text=part.get("text", "")))
80 elif part_type == "image_url":
81 img = part.get("image_url", {})
82 if isinstance(img, dict):
83 url = img.get("url", "")
84 detail = img.get("detail", "auto")
85 else:
86 url = str(img)
87 detail = "auto"
88 content_parts.append(ImageUrlPart(url=url, detail=detail))
89 content: str | list = content_parts if content_parts else ""
90 else:
91 content = str(raw_content)
92 role_str = msg.get("role", "user")
93 try:
94 role = Role(role_str)
95 except ValueError:
96 role = Role.USER
97 chat_messages.append(ChatMessage(role=role, content=content))
98
99 result = await client.complete(
100 messages=chat_messages,
101 model=model,
102 temperature=temperature,
103 max_tokens=max_tokens,
104 )
105 if result.is_err():
106 raise result.unwrap_err()
107 completion = result.unwrap()
108 latency_ms = (ambient_clock.monotonic() - t0) * 1000.0
109 usage = completion.usage
110 if usage is None:
111 prompt_tokens = 0
112 completion_tokens = 0
113 else:
114 prompt_tokens = usage.prompt_tokens # type: ignore[attr-defined]
115 completion_tokens = usage.completion_tokens # type: ignore[attr-defined]
116 return InferenceResult(
117 provider=provider_name,
118 model=model,
119 content=completion.content,
120 latency_ms=latency_ms,
121 prompt_tokens=prompt_tokens,
122 completion_tokens=completion_tokens,
123 )
124
125
126async def _handle_free_failure(
127 *,
128 exc: AIError,
129 provider_key: str,
130 model: str,
131 quota: QuotaBackendProtocol,
132 cooldown_seconds: int,
133) -> None:
134 """Update quota state after a cascade-entry failure.
135
136 429 (throttle) cools the entry down for *cooldown_seconds*; 402
137 (payment) exhausts it for the rest of the UTC day; anything else is
138 recorded as a plain error. State is keyed per cascade entry
139 (``ProviderConfig.key``), so one entry's failure never poisons its
140 siblings.
141 """
142 status_code = _extract_status_code(exc)
143 if status_code == 429:
144 until = datetime.now(UTC) + timedelta(seconds=cooldown_seconds)
145 logger.warning(
146 "strategy: entry %s throttled (model=%s), cooling down %ds",
147 provider_key,
148 model,
149 cooldown_seconds,
150 )
151 await quota.mark_exhausted(provider_key, until=until)
152 elif status_code == 402:
153 logger.warning(
154 "strategy: entry %s payment-exhausted (model=%s) for today",
155 provider_key,
156 model,
157 )
158 await quota.mark_exhausted(provider_key)
159 else:
160 logger.warning(
161 "strategy: entry %s error (model=%s status=%s): %s",
162 provider_key,
163 model,
164 status_code,
165 exc,
166 )
167 await quota.record_error(provider_key)
168
169
170def _gen_defaults(
171 config: LLMConfig,
172 kwargs: dict[str, Any],
173) -> tuple[float, int | None]:
174 """Resolve temperature and max_tokens from config + overrides."""
175 defaults = config.defaults
176 return (
177 kwargs.get("temperature", defaults.temperature),
178 kwargs.get("max_tokens", defaults.max_tokens),
179 )
180
181
182def _estimate_prompt_tokens(messages: list[Any]) -> int:
183 """Rough token estimate for prompt messages (~4 chars per token)."""
184 total_chars = 0
185 for msg in messages:
186 if isinstance(msg, dict):
187 total_chars += len(str(msg.get("content", "")))
188 elif hasattr(msg, "content"):
189 total_chars += len(str(msg.content))
190 else:
191 total_chars += len(str(msg))
192 return max(1, total_chars // 4)