Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-llm/src/lexigram/ai/llm/routing/strategies/base.py: 18%

89 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 07:19 +0800

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)