Coverage for src/lilbee/providers/routing_provider.py: 100%

195 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-09-28 17:20 +0000

1"""Routing provider: prefix-based dispatch between the SDK backend and the local engine.""" 

2 

3from __future__ import annotations 

4 

5import contextlib 

6import threading 

7from collections.abc import Callable 

8from pathlib import Path 

9from typing import TYPE_CHECKING, Any, Literal, overload 

10 

11from lilbee.app.services import get_services 

12from lilbee.catalog.refs import is_bare_hf_repo 

13from lilbee.core.config import cfg 

14from lilbee.core.health_warnings import HealthWarning, WarningCode 

15from lilbee.core.vectors import Vector 

16from lilbee.providers.base import ( 

17 ChatMessage, 

18 ChatResult, 

19 ChatStreamItem, 

20 ChatToolResult, 

21 ClosableIterator, 

22 LLMProvider, 

23 ProviderError, 

24 require_role_ref, 

25) 

26from lilbee.providers.litellm_sdk import LitellmSdkBackend 

27from lilbee.providers.model_ref import ProviderModelRef, parse_model_ref, routes_to_native_gguf 

28from lilbee.providers.roles import ROLE_REGISTRY, WorkerRole 

29from lilbee.providers.sdk_llm_provider import SdkLLMProvider 

30 

31if TYPE_CHECKING: 

32 from lilbee.providers.warm_progress import WarmProgress 

33 

34 

35class RoutingProvider(LLMProvider): 

36 """Dispatches calls based on the model ref prefix. 

37 

38 ``ollama/``, ``openai/``, ``anthropic/``, ``gemini/`` go to the SDK 

39 provider. Other refs (the HuggingFace ``<org>/<repo>/<file>.gguf`` 

40 shape) go to the local llama-server engine, which resolves them against the native 

41 registry. A registry miss surfaces the native ProviderError 

42 unchanged, rather than silently falling through to a remote backend. 

43 """ 

44 

45 def __init__(self, *, hold_warm: bool = False) -> None: 

46 self._local: LLMProvider | None = None 

47 self._sdk_provider: SdkLLMProvider | None = None 

48 # Carried until the local fleet is lazily built, so a provider created for 

49 # an interactive session hands that down to the FleetProvider it composes. 

50 self._hold_warm = hold_warm 

51 # Guards both lazy inits so two concurrent first-callers on the shared 

52 # daemon don't each build a backend and leak the loser. (Construction is 

53 # cheap field init; FleetProvider defers the role-server spawn to first 

54 # use and single-flights it internally, so this lock is the singleton 

55 # guard, not the spawn guard.) 

56 self._init_lock = threading.Lock() 

57 

58 def _get_local(self) -> LLMProvider: 

59 if self._local is None: 

60 # FleetProvider composes the llama-server stack; its role servers 

61 # spawn lazily on first use, not here. Double-checked under 

62 # _init_lock so only the first concurrent caller builds it. 

63 with self._init_lock: 

64 if self._local is None: 

65 from lilbee.providers.fleet.provider import FleetProvider 

66 

67 self._local = FleetProvider(hold_warm=self._hold_warm) 

68 return self._local 

69 

70 def _get_sdk_provider(self) -> SdkLLMProvider: 

71 if self._sdk_provider is None: 

72 with self._init_lock: 

73 if self._sdk_provider is None: 

74 self._sdk_provider = SdkLLMProvider(LitellmSdkBackend()) 

75 return self._sdk_provider 

76 

77 def _pick_backend(self, ref: ProviderModelRef) -> LLMProvider: 

78 """Pick the backend for *ref* purely by prefix.""" 

79 if ref.is_remote: 

80 return self._get_sdk_provider() 

81 return self._get_local() 

82 

83 def embed(self, texts: list[str]) -> list[Vector]: 

84 ref = parse_model_ref(require_role_ref(cfg.embedding_model, WorkerRole.EMBED)) 

85 return self._pick_backend(ref).embed(texts) 

86 

87 def count_tokens(self, text: str) -> int: 

88 ref = parse_model_ref(require_role_ref(cfg.embedding_model, WorkerRole.EMBED)) 

89 return self._pick_backend(ref).count_tokens(text) 

90 

91 def count_chat_prompt_tokens( 

92 self, 

93 messages: list[ChatMessage], 

94 *, 

95 options: dict[str, Any] | None = None, 

96 model: str | None = None, 

97 tools: list[dict[str, Any]] | None = None, 

98 tool_choice: str | dict[str, Any] | None = None, 

99 ) -> int: 

100 """Count on the backend the chat ref routes to, same rules as :meth:`chat`.""" 

101 ref = parse_model_ref(require_role_ref(model or cfg.chat_model, WorkerRole.CHAT)) 

102 return self._pick_backend(ref).count_chat_prompt_tokens( 

103 messages, options=options, model=model, tools=tools, tool_choice=tool_choice 

104 ) 

105 

106 @overload 

107 def chat( 

108 self, 

109 messages: list[dict[str, str]], 

110 *, 

111 stream: Literal[False] = False, 

112 options: dict[str, Any] | None = None, 

113 model: str | None = None, 

114 tools: list[dict[str, Any]] | None = None, 

115 tool_choice: str | dict[str, Any] | None = None, 

116 ) -> ChatResult: ... 

117 

118 @overload 

119 def chat( 

120 self, 

121 messages: list[dict[str, str]], 

122 *, 

123 stream: Literal[True], 

124 options: dict[str, Any] | None = None, 

125 model: str | None = None, 

126 tools: list[dict[str, Any]] | None = None, 

127 tool_choice: str | dict[str, Any] | None = None, 

128 ) -> ClosableIterator[ChatStreamItem]: ... 

129 

130 def chat( 

131 self, 

132 messages: list[dict[str, str]], 

133 *, 

134 stream: bool = False, 

135 options: dict[str, Any] | None = None, 

136 model: str | None = None, 

137 tools: list[dict[str, Any]] | None = None, 

138 tool_choice: str | dict[str, Any] | None = None, 

139 ) -> ChatResult | ClosableIterator[ChatStreamItem]: 

140 ref = parse_model_ref(require_role_ref(model or cfg.chat_model, WorkerRole.CHAT)) 

141 backend = self._pick_backend(ref) 

142 # Split on stream so each call resolves to a specific overload; the 

143 # base impl signature accepts bool but the @overloads on the LLMProvider 

144 # Protocol require Literal narrowing at the boundary. 

145 if stream: 

146 return backend.chat( 

147 messages, 

148 stream=True, 

149 options=options, 

150 model=model, 

151 tools=tools, 

152 tool_choice=tool_choice, 

153 ) 

154 return backend.chat( 

155 messages, 

156 stream=False, 

157 options=options, 

158 model=model, 

159 tools=tools, 

160 tool_choice=tool_choice, 

161 ) 

162 

163 def supports_tools(self, model_ref: str) -> bool: 

164 """Delegate the tool-capability probe to the backend the ref routes to.""" 

165 resolved = model_ref or cfg.chat_model 

166 if not resolved: 

167 return False # no model configured: nothing to advertise tools 

168 return self._pick_backend(parse_model_ref(resolved)).supports_tools(resolved) 

169 

170 def chat_with_tools( 

171 self, 

172 messages: list[dict[str, str]], 

173 *, 

174 tools: list[dict[str, Any]], 

175 tool_choice: str | dict[str, Any] | None = None, 

176 options: dict[str, Any] | None = None, 

177 model: str | None = None, 

178 ) -> ChatToolResult: 

179 """Dispatch a tool-enabled chat turn to the backend the ref routes to.""" 

180 ref = parse_model_ref(require_role_ref(model or cfg.chat_model, WorkerRole.CHAT)) 

181 backend = self._pick_backend(ref) 

182 return backend.chat_with_tools( 

183 messages, tools=tools, tool_choice=tool_choice, options=options, model=model 

184 ) 

185 

186 def vision_ocr( 

187 self, 

188 png_bytes: bytes, 

189 model: str, 

190 prompt: str = "", 

191 *, 

192 timeout: float | None = None, 

193 ) -> str: 

194 """Dispatch by ``model``'s ref prefix, same rules as :meth:`chat`.""" 

195 ref = parse_model_ref(model) 

196 return self._pick_backend(ref).vision_ocr(png_bytes, model, prompt, timeout=timeout) 

197 

198 def vision_slot_capacity(self) -> int | None: 

199 """Delegate to the local fleet, but never build it just to size the fan-out.""" 

200 return self._local.vision_slot_capacity() if self._local is not None else None 

201 

202 def list_models(self) -> list[str]: 

203 """Return the union of native and SDK-visible models. 

204 

205 Both halves are wrapped so an unreachable remote backend or a 

206 missing native registry does not mask the other. 

207 """ 

208 native: set[str] = set() 

209 with contextlib.suppress(Exception): 

210 native = set(self._get_local().list_models()) 

211 sdk = self._get_sdk_provider() 

212 if not sdk.available(): 

213 return sorted(native) 

214 try: 

215 remote = set(sdk.list_models()) 

216 except Exception: 

217 return sorted(native) 

218 return sorted(native | remote) 

219 

220 def list_chat_models(self, provider: str) -> list[str]: 

221 """Delegate to the SDK backend; the native engine has no catalog.""" 

222 sdk = self._get_sdk_provider() 

223 if not sdk.available(): 

224 return [] 

225 return sdk.list_chat_models(provider) 

226 

227 def pull_model(self, model: str, *, on_progress: Callable[..., Any] | None = None) -> None: 

228 """Pull via the SDK backend if installed, otherwise raise.""" 

229 sdk = self._get_sdk_provider() 

230 if not sdk.available(): 

231 raise ProviderError(f"Cannot pull model {model!r}: no pull-capable backend available") 

232 sdk.pull_model(model, on_progress=on_progress) 

233 

234 def show_model(self, model: str) -> dict[str, Any] | None: 

235 """Show model info from the backend selected by the ref prefix.""" 

236 ref = parse_model_ref(model) 

237 return self._pick_backend(ref).show_model(model) 

238 

239 def get_capabilities(self, model: str) -> list[str]: 

240 """Return capability tags from the backend selected by the ref prefix.""" 

241 ref = parse_model_ref(model) 

242 return self._pick_backend(ref).get_capabilities(model) 

243 

244 def rerank(self, query: str, candidates: list[str]) -> list[float]: 

245 """Dispatch rerank to the backend that owns ``cfg.reranker_model``. 

246 

247 Native GGUF refs go to the local engine; hosted refs go through the SDK 

248 provider. Raises ``ProviderError`` when ``cfg.reranker_model`` is 

249 empty or the selected backend does not support reranking. 

250 """ 

251 if not cfg.reranker_model: 

252 raise ProviderError("No reranker configured. Set cfg.reranker_model first.") 

253 if _is_native_rerank_ref(cfg.reranker_model): 

254 return self._get_local().rerank(query, candidates) 

255 sdk = self._get_sdk_provider() 

256 if not sdk.supports_rerank(): 

257 raise ProviderError( 

258 f"Cannot rerank with {cfg.reranker_model!r}: " 

259 "hosted rerank backend not available. " 

260 "Install the 'litellm' extra to enable hosted reranking." 

261 ) 

262 return sdk.rerank(query, candidates) 

263 

264 def supports_rerank(self) -> bool: 

265 """Capability probe: can the routed backend rerank if configured? 

266 

267 Pure capability check, NOT "a reranker is currently active". An 

268 empty ``cfg.reranker_model`` returns ``True`` so the settings UI 

269 keeps the picker visible; callers that need to know whether 

270 reranking is actually configured must check ``bool(cfg.reranker_model)`` 

271 separately. Delegates to the backend that would handle the 

272 configured model when one is set. 

273 """ 

274 model = cfg.reranker_model 

275 if not model: 

276 return True 

277 if _is_native_rerank_ref(model): 

278 return self._get_local().supports_rerank() 

279 return self._get_sdk_provider().supports_rerank() 

280 

281 def shutdown(self) -> None: 

282 """Shut down sub-providers to release resources.""" 

283 if self._local is not None: 

284 self._local.shutdown() 

285 if self._sdk_provider is not None: 

286 self._sdk_provider.shutdown() 

287 

288 def invalidate_load_cache(self, model_path: Path | None = None) -> None: 

289 """Forward to the native side only; the SDK side has no local cache.""" 

290 if self._local is not None: 

291 self._local.invalidate_load_cache(model_path) 

292 

293 def drop_loaded_models_async(self) -> None: 

294 """Forward the off-thread fleet drop to the native side; SDK has no cache.""" 

295 if self._local is not None: 

296 self._local.drop_loaded_models_async() 

297 

298 def warm_up_pool(self) -> None: 

299 """Forward to the native side; the SDK side has no servers to warm. 

300 

301 Lazily constructs the local engine if it isn't already up so 

302 eager-start during ``Services`` boot still warms the configured 

303 native roles, even when the user hasn't issued a chat call yet. 

304 """ 

305 self._get_local().warm_up_pool() 

306 

307 def cancel_inference(self) -> None: 

308 """Forward to the native engine; the SDK side has nothing to interrupt.""" 

309 if self._local is not None: 

310 self._local.cancel_inference() 

311 

312 def reload_role(self, role: WorkerRole, *, wait: bool = False) -> None: 

313 """Forward to the native engine; the SDK side has no per-role servers.""" 

314 if self._local is not None: 

315 self._local.reload_role(role, wait=wait) 

316 

317 def reload_placement(self, *, wait: bool = False) -> None: 

318 """Forward to the native engine; the SDK side has no GPU placement.""" 

319 if self._local is not None: 

320 self._local.reload_placement(wait=wait) 

321 

322 def role_ready(self, role: WorkerRole) -> bool: 

323 """Whether *role* can serve a request right now. 

324 

325 A role whose configured ref routes to the SDK backend needs no local 

326 server, so it is always ready; the local fleet's readiness is 

327 irrelevant to it. A native ref with no local engine built yet cannot 

328 serve a token, so it reports not-ready (without building the engine); 

329 health's ``chat_ready`` and the cold-start waits all rely on this 

330 being positive readiness, not reachability. 

331 """ 

332 if self._role_routes_remote(role): 

333 return True 

334 if self._local is None: 

335 return False 

336 return self._local.role_ready(role) 

337 

338 @staticmethod 

339 def _role_routes_remote(role: WorkerRole) -> bool: 

340 """Whether *role*'s configured model ref dispatches to the SDK backend.""" 

341 ref = str(getattr(cfg, ROLE_REGISTRY[role].config_field)) 

342 return bool(ref) and parse_model_ref(ref).is_remote 

343 

344 def max_concurrent_chats(self) -> int: 

345 """Chat concurrency of the local engine; 1 until one exists.""" 

346 if self._local is None: 

347 return 1 

348 return self._local.max_concurrent_chats() 

349 

350 def served_chat_ctx(self) -> int | None: 

351 """Per-slot chat context of the local engine, or None when none exists.""" 

352 if self._local is None: 

353 return None 

354 return self._local.served_chat_ctx() 

355 

356 def served_chat_slots(self) -> int | None: 

357 """Chat batching slots of the local engine, or None when none exists.""" 

358 if self._local is None: 

359 return None 

360 return self._local.served_chat_slots() 

361 

362 def _local_embedder(self) -> LLMProvider | None: 

363 """The local engine when it serves the configured embedder, else None.""" 

364 ref = cfg.embedding_model 

365 if not ref or parse_model_ref(ref).is_remote: 

366 return None 

367 return self._get_local() 

368 

369 def embed_token_cap(self) -> int | None: 

370 """Embed token cap of the local engine; None for a remote or unset embedder.""" 

371 local = self._local_embedder() 

372 return None if local is None else local.embed_token_cap() 

373 

374 def health_warnings(self) -> list[HealthWarning]: 

375 """Serving degradations of the local engine; the embed warning needs a local embedder.""" 

376 local = self._local_embedder() 

377 if local is not None: 

378 return local.health_warnings() 

379 if self._local is None: 

380 return [] 

381 return [ 

382 w 

383 for w in self._local.health_warnings() 

384 if w.code != WarningCode.EMBED_WINDOW_BELOW_CHUNK 

385 ] 

386 

387 def chat_prefill_progress(self) -> tuple[int, int] | None: 

388 """Chat prefill progress of the local engine, or None when none exists.""" 

389 if self._local is None: 

390 return None 

391 return self._local.chat_prefill_progress() 

392 

393 def warm_pending(self) -> bool: 

394 """Forward to the native side; the SDK side never warms.""" 

395 return self._get_local().warm_pending() 

396 

397 def warm_progress(self) -> WarmProgress | None: 

398 """Cold-load progress of the local engine, or None when none exists yet.""" 

399 if self._local is None: 

400 return None 

401 return self._local.warm_progress() 

402 

403 def add_spawn_listener( 

404 self, 

405 *, 

406 on_spawning: Callable[[WorkerRole], None] | None = None, 

407 on_spawned: Callable[[WorkerRole], None] | None = None, 

408 ) -> None: 

409 """Register on the native engine so its server spawns reach the TUI. 

410 

411 Builds the local engine if it isn't up yet so the listener is attached 

412 before the first spawn, matching ``warm_up_pool``'s eager construction. 

413 """ 

414 self._get_local().add_spawn_listener(on_spawning=on_spawning, on_spawned=on_spawned) 

415 

416 

417def _is_native_rerank_ref(model: str) -> bool: 

418 """Return True iff *model* should route to the native llama-server rerank path. 

419 

420 Two acceptance paths: 

421 

422 1. :func:`routes_to_native_gguf` accepts the ref: a native GGUF shape that 

423 no local-server prefix (``ollama/``, ``lm_studio/``) claims, matching 

424 :func:`parse_model_ref`'s exemption. 

425 2. The bare ``<org>/<repo>`` names a repo with an installed quant. 

426 

427 The model's name is deliberately not consulted. Hosted rerankers are 

428 usually called rerankers too (``cohere/rerank-english-v3.0``), so matching 

429 on the name captures them and starves the SDK backend. The registry answers 

430 "is this one of ours" without guessing. Non-GGUF refs without a known SDK 

431 prefix still raise downstream through :func:`parse_model_ref`. 

432 

433 An empty registry reports nothing installed; a registry that cannot be read 

434 raises, and the fault surfaces to the caller. 

435 """ 

436 if not model: 

437 return False 

438 if routes_to_native_gguf(model): 

439 return True 

440 if not is_bare_hf_repo(model): 

441 return False 

442 return get_services().registry.installed_ref_for_repo(model) is not None