Refactor provider runtime ownership (#925)
Browse files## Problem
Provider construction, model discovery, validation, and cleanup lived in
one registry module. API and admin routes depended on registry-shaped
app state and legacy process-level provider helpers.
## Changes
| Before | After |
| --- | --- |
| `providers.registry` mixed provider factories, config, cache,
discovery, validation, and cleanup. | `providers.runtime` splits
factories, config, cache, model cache, discovery, validation, and
runtime orchestration. |
| API and admin routes read `app.state.provider_registry` and sometimes
created registries ad hoc. | API and admin routes use app-scoped
`ProviderRuntime` through `app.state.provider_runtime`. |
| `api.dependencies` kept process-global provider cache helpers. |
`api.dependencies` resolves providers only through the app-scoped
runtime. |
| Registry-shaped tests preserved old internal boundaries. |
Runtime-shaped tests assert provider config, construction, cache,
discovery, validation, and import boundaries. |
<!-- greptile_comment -->
<details open><summary><h3>Greptile Summary</h3></summary>
This PR moves provider lifecycle ownership from the old registry module
into an app-scoped runtime package. The main changes are:
- Split provider config, factory wiring, instance cache, model cache,
discovery, validation, and cleanup into `providers.runtime` modules.
- Updated API and admin routes to resolve providers and model metadata
through `app.state.provider_runtime`.
- Removed legacy process-global provider helpers and the deleted
`providers.registry` module.
- Updated docs, smoke metadata, import-boundary checks, and tests for
the new runtime ownership model.
- Bumped the package version and lockfile metadata for the production
refactor.
</details>
<h3>Confidence Score: 5/5</h3>
The provider runtime refactor appears merge-safe with no identified
blocking issues.
The changes consistently move provider ownership to app-scoped runtime
modules and update API, admin, docs, smoke metadata, import-boundary
checks, and tests around that architecture.
<details><summary><h3><a href="https://www.greptile.com/trex"><img
alt="T-Rex"
src="https://greptile-static-assets.s3.amazonaws.com/trex/trex_green.svg"
height="20" align="absmiddle"></a> T-Rex Logs</h3></summary>
**What T-Rex did**
- Ran a baseline and head comparison of provider registry and runtime
states, verifying the after-state shows head
state\_has\_provider\_registry=False and
state\_has\_provider\_runtime=True, that GET /v1/models and admin
endpoints respond with 200, and that provider\_resolver\_called via
runtime, with assertions passing.
- Verified that the four focused provider-runtime contract tests passed
in both the before and after refactor runs, including runtime split
checks, with exit code 0.
- Identified environmental blockers that prevented the smoke-runtime
workflow from running, including uv unavailability, missing pytest for
/usr/local/bin/python, and Python 3.11 being used despite pyproject.toml
requiring \>=3.14.
<a
href="https://app.greptile.com/trex/runs/12528505/artifacts"><picture><source
media="(prefers-color-scheme: dark)"
srcset="https://greptile-static-assets.s3.amazonaws.com/badges/ViewAllArtifactsDark.svg?v=1"><source
media="(prefers-color-scheme: light)"
srcset="https://greptile-static-assets.s3.amazonaws.com/badges/ViewAllArtifacts.svg?v=1"><img
alt="View all artifacts"
src="https://greptile-static-assets.s3.amazonaws.com/badges/ViewAllArtifacts.svg?v=1"
height="32"></picture></a>
<sub><a href="https://www.greptile.com/trex"><img alt="T-Rex"
src="https://greptile-static-assets.s3.amazonaws.com/trex/trex_green.svg"
height="14" align="absmiddle"></a> Ran code and verified through
T-Rex</sub>
</details>
<sub>Reviews (1): Last reviewed commit: ["Refactor provider runtime
ownership"](https://github.com/alishahryar1/free-claude-code/commit/01d589488185c1f85112f1a49c47f04512846161)
| [Re-trigger
Greptile](https://app.greptile.com/api/retrigger?id=40312173)</sub>
<!-- /greptile_comment -->
- ARCHITECTURE.md +15 -13
- README.md +2 -2
- api/admin_routes.py +25 -20
- api/dependencies.py +25 -63
- api/model_catalog.py +6 -6
- api/routes.py +3 -5
- api/runtime.py +12 -12
- config/provider_catalog.py +1 -1
- providers/registry.py +0 -527
- providers/runtime/__init__.py +12 -0
- providers/runtime/cache.py +54 -0
- providers/runtime/config.py +65 -0
- providers/runtime/discovery.py +140 -0
- providers/runtime/factory.py +166 -0
- providers/runtime/model_cache.py +72 -0
- providers/runtime/runtime.py +93 -0
- providers/runtime/validation.py +123 -0
- pyproject.toml +1 -1
- smoke/README.md +1 -1
- smoke/capabilities.py +8 -8
- smoke/features.py +3 -3
- smoke/product/test_config_extensibility_product_live.py +8 -6
- tests/api/test_api.py +2 -2
- tests/api/test_app_lifespan_and_errors.py +28 -28
- tests/api/test_dependencies.py +93 -582
- tests/api/test_model_listing.py +16 -16
- tests/contracts/test_import_boundaries.py +17 -3
- tests/providers/test_model_validation.py +62 -59
- tests/providers/{test_registry.py → test_provider_runtime.py} +30 -32
- uv.lock +1 -1
|
@@ -35,7 +35,7 @@ flowchart LR
|
|
| 35 |
ProxyAPI --> Handlers[API Product Handlers]
|
| 36 |
Handlers --> Router[ModelRouter]
|
| 37 |
Handlers --> Executor[ProviderExecutionService]
|
| 38 |
-
Executor --> Providers[
|
| 39 |
Providers --> OpenAIChat[OpenAI Chat Providers]
|
| 40 |
Providers --> NativeAnthropic[Anthropic Messages Providers]
|
| 41 |
```
|
|
@@ -156,7 +156,7 @@ reported without noisy Starlette tracebacks.
|
|
| 156 |
[api/runtime.py](api/runtime.py) owns process-lifetime resources through
|
| 157 |
`AppRuntime`:
|
| 158 |
|
| 159 |
-
- creates and publishes an app-scoped `
|
| 160 |
- validates configured models best-effort without blocking first-run admin access;
|
| 161 |
- starts provider model-list refresh;
|
| 162 |
- starts optional Discord or Telegram messaging when configured;
|
|
@@ -198,7 +198,7 @@ Model routing configuration is tiered:
|
|
| 198 |
and writes managed env updates. [api/admin_routes.py](api/admin_routes.py)
|
| 199 |
exposes local-only admin endpoints that load, validate, apply, and test config.
|
| 200 |
After an apply, settings are cache-cleared. Depending on the changed fields, the
|
| 201 |
-
server either replaces the app provider
|
| 202 |
restart.
|
| 203 |
|
| 204 |
Admin routes call `require_loopback_admin()`, which rejects non-loopback clients
|
|
@@ -241,7 +241,7 @@ sequenceDiagram
|
|
| 241 |
participant Handler as ProductHandler
|
| 242 |
participant Router as ModelRouter
|
| 243 |
participant Exec as ProviderExecution
|
| 244 |
-
participant
|
| 245 |
participant Provider
|
| 246 |
|
| 247 |
Client->>Route: POST /v1/messages
|
|
@@ -250,8 +250,8 @@ sequenceDiagram
|
|
| 250 |
Handler->>Router: resolve model and thinking
|
| 251 |
Handler->>Handler: server tools or optimizations
|
| 252 |
Handler->>Exec: stream routed request
|
| 253 |
-
Exec->>
|
| 254 |
-
|
| 255 |
Exec->>Provider: preflight_stream
|
| 256 |
Exec->>Provider: stream_response
|
| 257 |
Provider-->>Client: Anthropic SSE events
|
|
@@ -287,7 +287,7 @@ overrides or the global setting.
|
|
| 287 |
- no-thinking variants when appropriate;
|
| 288 |
- built-in Claude model IDs for compatibility with Claude clients.
|
| 289 |
|
| 290 |
-
Provider model discovery is app-scoped through `
|
| 291 |
model IDs and optional thinking capability metadata for the model-list route and
|
| 292 |
admin status.
|
| 293 |
|
|
@@ -304,10 +304,12 @@ Provider metadata is neutral and centralized in
|
|
| 304 |
`ProviderDescriptor` declares provider ID, transport type, capabilities,
|
| 305 |
credential env var, default base URL, settings attribute names, and proxy support.
|
| 306 |
|
| 307 |
-
[providers/
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
|
|
|
|
|
|
|
| 311 |
|
| 312 |
[providers/base.py](providers/base.py) defines:
|
| 313 |
|
|
@@ -342,7 +344,7 @@ where supported, and returning Anthropic SSE strings to the service layer.
|
|
| 342 |
the setting should be editable in the Admin UI.
|
| 343 |
4. Implement the provider under [providers/](providers/) using the appropriate
|
| 344 |
shared transport family.
|
| 345 |
-
5. Add a factory in [providers/
|
| 346 |
6. Add deterministic tests under [tests/providers/](tests/providers/) and any
|
| 347 |
relevant contract tests.
|
| 348 |
7. Add smoke coverage or smoke config in [smoke/](smoke/) when the provider can
|
|
@@ -697,7 +699,7 @@ Update this file when a change adds or meaningfully changes:
|
|
| 697 |
- a public route or wire protocol;
|
| 698 |
- startup, shutdown, or resource ownership;
|
| 699 |
- configuration precedence or managed config behavior;
|
| 700 |
-
- provider
|
| 701 |
- model routing or thinking behavior;
|
| 702 |
- CLI adapter behavior;
|
| 703 |
- messaging platform behavior;
|
|
|
|
| 35 |
ProxyAPI --> Handlers[API Product Handlers]
|
| 36 |
Handlers --> Router[ModelRouter]
|
| 37 |
Handlers --> Executor[ProviderExecutionService]
|
| 38 |
+
Executor --> Providers[ProviderRuntime]
|
| 39 |
Providers --> OpenAIChat[OpenAI Chat Providers]
|
| 40 |
Providers --> NativeAnthropic[Anthropic Messages Providers]
|
| 41 |
```
|
|
|
|
| 156 |
[api/runtime.py](api/runtime.py) owns process-lifetime resources through
|
| 157 |
`AppRuntime`:
|
| 158 |
|
| 159 |
+
- creates and publishes an app-scoped `ProviderRuntime`;
|
| 160 |
- validates configured models best-effort without blocking first-run admin access;
|
| 161 |
- starts provider model-list refresh;
|
| 162 |
- starts optional Discord or Telegram messaging when configured;
|
|
|
|
| 198 |
and writes managed env updates. [api/admin_routes.py](api/admin_routes.py)
|
| 199 |
exposes local-only admin endpoints that load, validate, apply, and test config.
|
| 200 |
After an apply, settings are cache-cleared. Depending on the changed fields, the
|
| 201 |
+
server either replaces the app provider runtime or asks the supervised server to
|
| 202 |
restart.
|
| 203 |
|
| 204 |
Admin routes call `require_loopback_admin()`, which rejects non-loopback clients
|
|
|
|
| 241 |
participant Handler as ProductHandler
|
| 242 |
participant Router as ModelRouter
|
| 243 |
participant Exec as ProviderExecution
|
| 244 |
+
participant Runtime as ProviderRuntime
|
| 245 |
participant Provider
|
| 246 |
|
| 247 |
Client->>Route: POST /v1/messages
|
|
|
|
| 250 |
Handler->>Router: resolve model and thinking
|
| 251 |
Handler->>Handler: server tools or optimizations
|
| 252 |
Handler->>Exec: stream routed request
|
| 253 |
+
Exec->>Runtime: resolve provider
|
| 254 |
+
Runtime->>Provider: cached or new provider
|
| 255 |
Exec->>Provider: preflight_stream
|
| 256 |
Exec->>Provider: stream_response
|
| 257 |
Provider-->>Client: Anthropic SSE events
|
|
|
|
| 287 |
- no-thinking variants when appropriate;
|
| 288 |
- built-in Claude model IDs for compatibility with Claude clients.
|
| 289 |
|
| 290 |
+
Provider model discovery is app-scoped through `ProviderRuntime`, which caches
|
| 291 |
model IDs and optional thinking capability metadata for the model-list route and
|
| 292 |
admin status.
|
| 293 |
|
|
|
|
| 304 |
`ProviderDescriptor` declares provider ID, transport type, capabilities,
|
| 305 |
credential env var, default base URL, settings attribute names, and proxy support.
|
| 306 |
|
| 307 |
+
[providers/runtime/](providers/runtime/) owns the app-scoped provider runtime.
|
| 308 |
+
It validates that descriptors, factories, and supported IDs are in sync, builds
|
| 309 |
+
shared `ProviderConfig`, checks required credentials, creates providers lazily,
|
| 310 |
+
caches them, refreshes model lists, validates configured models, and cleans up
|
| 311 |
+
transports. The package splits factory wiring, config building, provider instance
|
| 312 |
+
cache, model metadata cache, discovery, and validation into separate modules.
|
| 313 |
|
| 314 |
[providers/base.py](providers/base.py) defines:
|
| 315 |
|
|
|
|
| 344 |
the setting should be editable in the Admin UI.
|
| 345 |
4. Implement the provider under [providers/](providers/) using the appropriate
|
| 346 |
shared transport family.
|
| 347 |
+
5. Add a factory in [providers/runtime/factory.py](providers/runtime/factory.py).
|
| 348 |
6. Add deterministic tests under [tests/providers/](tests/providers/) and any
|
| 349 |
relevant contract tests.
|
| 350 |
7. Add smoke coverage or smoke config in [smoke/](smoke/) when the provider can
|
|
|
|
| 699 |
- a public route or wire protocol;
|
| 700 |
- startup, shutdown, or resource ownership;
|
| 701 |
- configuration precedence or managed config behavior;
|
| 702 |
+
- provider runtime, catalog, or transport architecture;
|
| 703 |
- model routing or thinking behavior;
|
| 704 |
- CLI adapter behavior;
|
| 705 |
- messaging platform behavior;
|
|
@@ -577,7 +577,7 @@ free-claude-code/
|
|
| 577 |
├── api/ # FastAPI routes, service layer, routing, optimizations
|
| 578 |
├── core/ # Shared Anthropic protocol helpers, SSE, OpenAI Responses
|
| 579 |
│ └── openai_responses/ # Responses ↔ Anthropic conversion and SSE mapping
|
| 580 |
-
├── providers/ # Provider
|
| 581 |
├── messaging/ # Discord/Telegram runtimes, outbound ports, voice
|
| 582 |
├── cli/ # Package entry points and client CLI process management
|
| 583 |
├── config/ # Settings, provider catalog, logging
|
|
@@ -639,7 +639,7 @@ CI also enforces a ban on `# type: ignore` / `# ty: ignore` suppressions; `scrip
|
|
| 639 |
- Add OpenAI-compatible providers by extending `OpenAIChatTransport`.
|
| 640 |
- Add Anthropic Messages providers by extending `AnthropicMessagesTransport`.
|
| 641 |
- Extend OpenAI Responses conversion in `core/openai_responses/` when Codex adds new request or stream shapes.
|
| 642 |
-
- Register provider metadata in `config.provider_catalog` and factory wiring in `providers.
|
| 643 |
- Add messaging platforms by wiring runtime, outbound, and inbound-normalizer ports in `messaging/platforms/`.
|
| 644 |
|
| 645 |
## Contributing
|
|
|
|
| 577 |
├── api/ # FastAPI routes, service layer, routing, optimizations
|
| 578 |
├── core/ # Shared Anthropic protocol helpers, SSE, OpenAI Responses
|
| 579 |
│ └── openai_responses/ # Responses ↔ Anthropic conversion and SSE mapping
|
| 580 |
+
├── providers/ # Provider runtime, transports, rate limiting
|
| 581 |
├── messaging/ # Discord/Telegram runtimes, outbound ports, voice
|
| 582 |
├── cli/ # Package entry points and client CLI process management
|
| 583 |
├── config/ # Settings, provider catalog, logging
|
|
|
|
| 639 |
- Add OpenAI-compatible providers by extending `OpenAIChatTransport`.
|
| 640 |
- Add Anthropic Messages providers by extending `AnthropicMessagesTransport`.
|
| 641 |
- Extend OpenAI Responses conversion in `core/openai_responses/` when Codex adds new request or stream shapes.
|
| 642 |
+
- Register provider metadata in `config.provider_catalog` and factory wiring in `providers.runtime`.
|
| 643 |
- Add messaging platforms by wiring runtime, outbound, and inbound-normalizer ports in `messaging/platforms/`.
|
| 644 |
|
| 645 |
## Contributing
|
|
@@ -15,7 +15,7 @@ from pydantic import BaseModel, Field
|
|
| 15 |
|
| 16 |
from config.settings import Settings
|
| 17 |
from config.settings import get_settings as get_cached_settings
|
| 18 |
-
from providers.
|
| 19 |
|
| 20 |
from .admin_config import (
|
| 21 |
FIELD_BY_KEY,
|
|
@@ -126,10 +126,10 @@ async def apply_admin_config(
|
|
| 126 |
request.app.state.admin_pending_fields = []
|
| 127 |
return result
|
| 128 |
|
| 129 |
-
|
| 130 |
-
if isinstance(
|
| 131 |
-
await
|
| 132 |
-
request.app.state.
|
| 133 |
request.app.state.admin_pending_fields = result["pending_fields"]
|
| 134 |
return result
|
| 135 |
|
|
@@ -138,12 +138,12 @@ async def apply_admin_config(
|
|
| 138 |
async def admin_status(request: Request):
|
| 139 |
require_loopback_admin(request)
|
| 140 |
settings = get_cached_settings()
|
| 141 |
-
|
| 142 |
cached_models: dict[str, list[str]] = {}
|
| 143 |
-
if isinstance(
|
| 144 |
cached_models = {
|
| 145 |
provider_id: sorted(model_ids)
|
| 146 |
-
for provider_id, model_ids in
|
| 147 |
}
|
| 148 |
return {
|
| 149 |
"status": "running",
|
|
@@ -173,12 +173,9 @@ async def local_provider_status(request: Request):
|
|
| 173 |
async def test_provider(provider_id: str, request: Request):
|
| 174 |
require_loopback_admin(request)
|
| 175 |
settings = get_cached_settings()
|
| 176 |
-
|
| 177 |
-
if not isinstance(registry, ProviderRegistry):
|
| 178 |
-
registry = ProviderRegistry()
|
| 179 |
-
request.app.state.provider_registry = registry
|
| 180 |
try:
|
| 181 |
-
provider =
|
| 182 |
infos = await provider.list_model_infos()
|
| 183 |
except Exception as exc:
|
| 184 |
return {
|
|
@@ -186,7 +183,7 @@ async def test_provider(provider_id: str, request: Request):
|
|
| 186 |
"ok": False,
|
| 187 |
"error_type": type(exc).__name__,
|
| 188 |
}
|
| 189 |
-
|
| 190 |
return {
|
| 191 |
"provider_id": provider_id,
|
| 192 |
"ok": True,
|
|
@@ -198,19 +195,27 @@ async def test_provider(provider_id: str, request: Request):
|
|
| 198 |
async def refresh_models(request: Request):
|
| 199 |
require_loopback_admin(request)
|
| 200 |
settings = get_cached_settings()
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
registry = ProviderRegistry()
|
| 204 |
-
request.app.state.provider_registry = registry
|
| 205 |
-
await registry.refresh_model_list_cache(settings)
|
| 206 |
return {
|
| 207 |
"cached_models": {
|
| 208 |
provider_id: sorted(model_ids)
|
| 209 |
-
for provider_id, model_ids in
|
| 210 |
}
|
| 211 |
}
|
| 212 |
|
| 213 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 214 |
def _filtered_values(values: dict[str, Any]) -> dict[str, Any]:
|
| 215 |
return {key: value for key, value in values.items() if key in FIELD_BY_KEY}
|
| 216 |
|
|
|
|
| 15 |
|
| 16 |
from config.settings import Settings
|
| 17 |
from config.settings import get_settings as get_cached_settings
|
| 18 |
+
from providers.runtime import ProviderRuntime
|
| 19 |
|
| 20 |
from .admin_config import (
|
| 21 |
FIELD_BY_KEY,
|
|
|
|
| 126 |
request.app.state.admin_pending_fields = []
|
| 127 |
return result
|
| 128 |
|
| 129 |
+
old_runtime = getattr(request.app.state, "provider_runtime", None)
|
| 130 |
+
if isinstance(old_runtime, ProviderRuntime):
|
| 131 |
+
await old_runtime.cleanup()
|
| 132 |
+
request.app.state.provider_runtime = ProviderRuntime(get_cached_settings())
|
| 133 |
request.app.state.admin_pending_fields = result["pending_fields"]
|
| 134 |
return result
|
| 135 |
|
|
|
|
| 138 |
async def admin_status(request: Request):
|
| 139 |
require_loopback_admin(request)
|
| 140 |
settings = get_cached_settings()
|
| 141 |
+
runtime = getattr(request.app.state, "provider_runtime", None)
|
| 142 |
cached_models: dict[str, list[str]] = {}
|
| 143 |
+
if isinstance(runtime, ProviderRuntime):
|
| 144 |
cached_models = {
|
| 145 |
provider_id: sorted(model_ids)
|
| 146 |
+
for provider_id, model_ids in runtime.cached_model_ids().items()
|
| 147 |
}
|
| 148 |
return {
|
| 149 |
"status": "running",
|
|
|
|
| 173 |
async def test_provider(provider_id: str, request: Request):
|
| 174 |
require_loopback_admin(request)
|
| 175 |
settings = get_cached_settings()
|
| 176 |
+
runtime = _provider_runtime_for_admin(request, settings)
|
|
|
|
|
|
|
|
|
|
| 177 |
try:
|
| 178 |
+
provider = runtime.resolve_provider(provider_id)
|
| 179 |
infos = await provider.list_model_infos()
|
| 180 |
except Exception as exc:
|
| 181 |
return {
|
|
|
|
| 183 |
"ok": False,
|
| 184 |
"error_type": type(exc).__name__,
|
| 185 |
}
|
| 186 |
+
runtime.cache_model_infos(provider_id, infos)
|
| 187 |
return {
|
| 188 |
"provider_id": provider_id,
|
| 189 |
"ok": True,
|
|
|
|
| 195 |
async def refresh_models(request: Request):
|
| 196 |
require_loopback_admin(request)
|
| 197 |
settings = get_cached_settings()
|
| 198 |
+
runtime = _provider_runtime_for_admin(request, settings)
|
| 199 |
+
await runtime.refresh_model_list_cache()
|
|
|
|
|
|
|
|
|
|
| 200 |
return {
|
| 201 |
"cached_models": {
|
| 202 |
provider_id: sorted(model_ids)
|
| 203 |
+
for provider_id, model_ids in runtime.cached_model_ids().items()
|
| 204 |
}
|
| 205 |
}
|
| 206 |
|
| 207 |
|
| 208 |
+
def _provider_runtime_for_admin(
|
| 209 |
+
request: Request, settings: Settings
|
| 210 |
+
) -> ProviderRuntime:
|
| 211 |
+
runtime = getattr(request.app.state, "provider_runtime", None)
|
| 212 |
+
if isinstance(runtime, ProviderRuntime):
|
| 213 |
+
return runtime
|
| 214 |
+
runtime = ProviderRuntime(settings)
|
| 215 |
+
request.app.state.provider_runtime = runtime
|
| 216 |
+
return runtime
|
| 217 |
+
|
| 218 |
+
|
| 219 |
def _filtered_values(values: dict[str, Any]) -> dict[str, Any]:
|
| 220 |
return {key: value for key, value in values.items() if key in FIELD_BY_KEY}
|
| 221 |
|
|
@@ -6,6 +6,7 @@ from fastapi import Depends, HTTPException, Request
|
|
| 6 |
from loguru import logger
|
| 7 |
from starlette.applications import Starlette
|
| 8 |
|
|
|
|
| 9 |
from config.settings import Settings
|
| 10 |
from config.settings import get_settings as _get_settings
|
| 11 |
from core.anthropic import get_user_facing_error_message
|
|
@@ -15,12 +16,7 @@ from providers.exceptions import (
|
|
| 15 |
ServiceUnavailableError,
|
| 16 |
UnknownProviderTypeError,
|
| 17 |
)
|
| 18 |
-
from providers.
|
| 19 |
-
|
| 20 |
-
# Process-level cache: only for :func:`get_provider_for_type` / :func:`get_provider`
|
| 21 |
-
# when there is no ``Request``/``app`` (unit tests, scripts). HTTP handlers must pass
|
| 22 |
-
# ``app`` to :func:`resolve_provider` so the app-scoped registry is used.
|
| 23 |
-
_providers: dict[str, BaseProvider] = {}
|
| 24 |
|
| 25 |
|
| 26 |
def get_settings() -> Settings:
|
|
@@ -28,39 +24,33 @@ def get_settings() -> Settings:
|
|
| 28 |
return _get_settings()
|
| 29 |
|
| 30 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
def resolve_provider(
|
| 32 |
provider_type: str,
|
| 33 |
*,
|
| 34 |
-
app: Starlette
|
| 35 |
-
settings: Settings,
|
| 36 |
-
) -> BaseProvider:
|
| 37 |
-
"""Resolve a provider using the app-scoped registry when ``app`` is set.
|
| 38 |
-
|
| 39 |
-
When ``app`` is not ``None``, the app-owned :attr:`app.state.provider_registry`
|
| 40 |
-
must exist (installed by :class:`~api.runtime.AppRuntime` during startup).
|
| 41 |
-
Callers that construct a bare ``FastAPI`` without lifespan must set
|
| 42 |
-
``app.state.provider_registry`` explicitly.
|
| 43 |
-
|
| 44 |
-
When ``app`` is ``None`` (no HTTP context), uses the process-level
|
| 45 |
-
:data:`_providers` cache only.
|
| 46 |
-
"""
|
| 47 |
-
if app is not None:
|
| 48 |
-
reg = getattr(app.state, "provider_registry", None)
|
| 49 |
-
if reg is None:
|
| 50 |
-
raise ServiceUnavailableError(
|
| 51 |
-
"Provider registry is not configured. Ensure AppRuntime startup ran "
|
| 52 |
-
"or assign app.state.provider_registry for test apps."
|
| 53 |
-
)
|
| 54 |
-
return _resolve_with_registry(reg, provider_type, settings)
|
| 55 |
-
return _resolve_with_registry(ProviderRegistry(_providers), provider_type, settings)
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
def _resolve_with_registry(
|
| 59 |
-
registry: ProviderRegistry, provider_type: str, settings: Settings
|
| 60 |
) -> BaseProvider:
|
| 61 |
-
|
|
|
|
|
|
|
| 62 |
try:
|
| 63 |
-
provider =
|
| 64 |
except AuthenticationError as e:
|
| 65 |
# Provider :class:`~providers.exceptions.AuthenticationError` messages are
|
| 66 |
# curated configuration hints (env var names, docs links), not upstream noise.
|
|
@@ -70,7 +60,7 @@ def _resolve_with_registry(
|
|
| 70 |
logger.error(
|
| 71 |
"Unknown provider_type: '{}'. Supported: {}",
|
| 72 |
provider_type,
|
| 73 |
-
", ".join(f"'{key}'" for key in
|
| 74 |
)
|
| 75 |
raise
|
| 76 |
if should_log_init:
|
|
@@ -78,16 +68,6 @@ def _resolve_with_registry(
|
|
| 78 |
return provider
|
| 79 |
|
| 80 |
|
| 81 |
-
def get_provider_for_type(provider_type: str) -> BaseProvider:
|
| 82 |
-
"""Get or create a provider in the process-level cache (no ``app``/Request).
|
| 83 |
-
|
| 84 |
-
HTTP route handlers should call :func:`resolve_provider` with the active
|
| 85 |
-
:attr:`request.app` (via :class:`~api.runtime.AppRuntime`) instead of this
|
| 86 |
-
process-wide cache.
|
| 87 |
-
"""
|
| 88 |
-
return resolve_provider(provider_type, app=None, settings=get_settings())
|
| 89 |
-
|
| 90 |
-
|
| 91 |
def require_api_key(
|
| 92 |
request: Request, settings: Settings = Depends(get_settings)
|
| 93 |
) -> None:
|
|
@@ -124,21 +104,3 @@ def require_api_key(
|
|
| 124 |
token.encode("utf-8"), anthropic_auth_token.encode("utf-8")
|
| 125 |
):
|
| 126 |
raise HTTPException(status_code=401, detail="Invalid API key")
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
def get_provider() -> BaseProvider:
|
| 130 |
-
"""Get or create the default provider (``MODEL`` / ``provider_type``).
|
| 131 |
-
|
| 132 |
-
Process-cache helper for scripts, unit tests, and non-FastAPI callers. HTTP
|
| 133 |
-
handlers must use :func:`resolve_provider` with :attr:`request.app` so the
|
| 134 |
-
app-scoped :class:`~providers.registry.ProviderRegistry` is used.
|
| 135 |
-
"""
|
| 136 |
-
return get_provider_for_type(get_settings().provider_type)
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
async def cleanup_provider():
|
| 140 |
-
"""Cleanup all provider resources."""
|
| 141 |
-
global _providers
|
| 142 |
-
await ProviderRegistry(_providers).cleanup()
|
| 143 |
-
_providers = {}
|
| 144 |
-
logger.debug("Provider cleanup completed")
|
|
|
|
| 6 |
from loguru import logger
|
| 7 |
from starlette.applications import Starlette
|
| 8 |
|
| 9 |
+
from config.provider_catalog import PROVIDER_CATALOG
|
| 10 |
from config.settings import Settings
|
| 11 |
from config.settings import get_settings as _get_settings
|
| 12 |
from core.anthropic import get_user_facing_error_message
|
|
|
|
| 16 |
ServiceUnavailableError,
|
| 17 |
UnknownProviderTypeError,
|
| 18 |
)
|
| 19 |
+
from providers.runtime import ProviderRuntime
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
|
| 21 |
|
| 22 |
def get_settings() -> Settings:
|
|
|
|
| 24 |
return _get_settings()
|
| 25 |
|
| 26 |
|
| 27 |
+
def get_provider_runtime(app: Starlette) -> ProviderRuntime:
|
| 28 |
+
"""Return the app-scoped provider runtime installed by ``AppRuntime``."""
|
| 29 |
+
runtime = getattr(app.state, "provider_runtime", None)
|
| 30 |
+
if isinstance(runtime, ProviderRuntime):
|
| 31 |
+
return runtime
|
| 32 |
+
raise ServiceUnavailableError(
|
| 33 |
+
"Provider runtime is not configured. Ensure AppRuntime startup ran "
|
| 34 |
+
"or assign app.state.provider_runtime for test apps."
|
| 35 |
+
)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def maybe_provider_runtime(app: Starlette) -> ProviderRuntime | None:
|
| 39 |
+
"""Return the app-scoped provider runtime when it is installed."""
|
| 40 |
+
runtime = getattr(app.state, "provider_runtime", None)
|
| 41 |
+
return runtime if isinstance(runtime, ProviderRuntime) else None
|
| 42 |
+
|
| 43 |
+
|
| 44 |
def resolve_provider(
|
| 45 |
provider_type: str,
|
| 46 |
*,
|
| 47 |
+
app: Starlette,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 48 |
) -> BaseProvider:
|
| 49 |
+
"""Resolve a provider through the app-scoped provider runtime."""
|
| 50 |
+
runtime = get_provider_runtime(app)
|
| 51 |
+
should_log_init = not runtime.is_cached(provider_type)
|
| 52 |
try:
|
| 53 |
+
provider = runtime.resolve_provider(provider_type)
|
| 54 |
except AuthenticationError as e:
|
| 55 |
# Provider :class:`~providers.exceptions.AuthenticationError` messages are
|
| 56 |
# curated configuration hints (env var names, docs links), not upstream noise.
|
|
|
|
| 60 |
logger.error(
|
| 61 |
"Unknown provider_type: '{}'. Supported: {}",
|
| 62 |
provider_type,
|
| 63 |
+
", ".join(f"'{key}'" for key in PROVIDER_CATALOG),
|
| 64 |
)
|
| 65 |
raise
|
| 66 |
if should_log_init:
|
|
|
|
| 68 |
return provider
|
| 69 |
|
| 70 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
def require_api_key(
|
| 72 |
request: Request, settings: Settings = Depends(get_settings)
|
| 73 |
) -> None:
|
|
|
|
| 104 |
token.encode("utf-8"), anthropic_auth_token.encode("utf-8")
|
| 105 |
):
|
| 106 |
raise HTTPException(status_code=401, detail="Invalid API key")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -3,7 +3,7 @@
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
from config.settings import Settings
|
| 6 |
-
from providers.
|
| 7 |
|
| 8 |
from .gateway_model_ids import gateway_model_id, no_thinking_gateway_model_id
|
| 9 |
from .models.responses import ModelResponse, ModelsListResponse
|
|
@@ -51,7 +51,7 @@ SUPPORTED_CLAUDE_MODELS = [
|
|
| 51 |
|
| 52 |
|
| 53 |
def build_models_list_response(
|
| 54 |
-
settings: Settings,
|
| 55 |
) -> ModelsListResponse:
|
| 56 |
"""Return configured, cached, and compatibility model ids."""
|
| 57 |
models: list[ModelResponse] = []
|
|
@@ -59,8 +59,8 @@ def build_models_list_response(
|
|
| 59 |
|
| 60 |
for ref in settings.configured_chat_model_refs():
|
| 61 |
supports_thinking = None
|
| 62 |
-
if
|
| 63 |
-
supports_thinking =
|
| 64 |
ref.provider_id, ref.model_id
|
| 65 |
)
|
| 66 |
_append_provider_model_variants(
|
|
@@ -70,8 +70,8 @@ def build_models_list_response(
|
|
| 70 |
supports_thinking=supports_thinking,
|
| 71 |
)
|
| 72 |
|
| 73 |
-
if
|
| 74 |
-
for model_info in
|
| 75 |
_append_provider_model_variants(
|
| 76 |
models,
|
| 77 |
seen,
|
|
|
|
| 3 |
from __future__ import annotations
|
| 4 |
|
| 5 |
from config.settings import Settings
|
| 6 |
+
from providers.runtime import ProviderRuntime
|
| 7 |
|
| 8 |
from .gateway_model_ids import gateway_model_id, no_thinking_gateway_model_id
|
| 9 |
from .models.responses import ModelResponse, ModelsListResponse
|
|
|
|
| 51 |
|
| 52 |
|
| 53 |
def build_models_list_response(
|
| 54 |
+
settings: Settings, provider_runtime: ProviderRuntime | None
|
| 55 |
) -> ModelsListResponse:
|
| 56 |
"""Return configured, cached, and compatibility model ids."""
|
| 57 |
models: list[ModelResponse] = []
|
|
|
|
| 59 |
|
| 60 |
for ref in settings.configured_chat_model_refs():
|
| 61 |
supports_thinking = None
|
| 62 |
+
if provider_runtime is not None:
|
| 63 |
+
supports_thinking = provider_runtime.cached_model_supports_thinking(
|
| 64 |
ref.provider_id, ref.model_id
|
| 65 |
)
|
| 66 |
_append_provider_model_variants(
|
|
|
|
| 70 |
supports_thinking=supports_thinking,
|
| 71 |
)
|
| 72 |
|
| 73 |
+
if provider_runtime is not None:
|
| 74 |
+
for model_info in provider_runtime.cached_prefixed_model_infos():
|
| 75 |
_append_provider_model_variants(
|
| 76 |
models,
|
| 77 |
seen,
|
|
@@ -6,7 +6,6 @@ from loguru import logger
|
|
| 6 |
from config.settings import Settings
|
| 7 |
from core.anthropic import get_token_count
|
| 8 |
from core.trace import trace_event
|
| 9 |
-
from providers.registry import ProviderRegistry
|
| 10 |
|
| 11 |
from . import dependencies
|
| 12 |
from .dependencies import get_settings, require_api_key
|
|
@@ -21,7 +20,7 @@ router = APIRouter()
|
|
| 21 |
|
| 22 |
def _provider_getter(request: Request, settings: Settings):
|
| 23 |
return lambda provider_type: dependencies.resolve_provider(
|
| 24 |
-
provider_type, app=request.app
|
| 25 |
)
|
| 26 |
|
| 27 |
|
|
@@ -149,9 +148,8 @@ async def list_models(
|
|
| 149 |
):
|
| 150 |
"""List the model ids this proxy advertises to Claude-compatible clients."""
|
| 151 |
trace_event(stage="ingress", event="api.models.list", source="api")
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
return build_models_list_response(settings, provider_registry)
|
| 155 |
|
| 156 |
|
| 157 |
@router.post("/stop")
|
|
|
|
| 6 |
from config.settings import Settings
|
| 7 |
from core.anthropic import get_token_count
|
| 8 |
from core.trace import trace_event
|
|
|
|
| 9 |
|
| 10 |
from . import dependencies
|
| 11 |
from .dependencies import get_settings, require_api_key
|
|
|
|
| 20 |
|
| 21 |
def _provider_getter(request: Request, settings: Settings):
|
| 22 |
return lambda provider_type: dependencies.resolve_provider(
|
| 23 |
+
provider_type, app=request.app
|
| 24 |
)
|
| 25 |
|
| 26 |
|
|
|
|
| 148 |
):
|
| 149 |
"""List the model ids this proxy advertises to Claude-compatible clients."""
|
| 150 |
trace_event(stage="ingress", event="api.models.list", source="api")
|
| 151 |
+
provider_runtime = dependencies.maybe_provider_runtime(request.app)
|
| 152 |
+
return build_models_list_response(settings, provider_runtime)
|
|
|
|
| 153 |
|
| 154 |
|
| 155 |
@router.post("/stop")
|
|
@@ -14,7 +14,7 @@ from loguru import logger
|
|
| 14 |
from api.admin_urls import local_admin_url
|
| 15 |
from config.settings import Settings, get_settings
|
| 16 |
from providers.exceptions import ServiceUnavailableError
|
| 17 |
-
from providers.
|
| 18 |
|
| 19 |
if TYPE_CHECKING:
|
| 20 |
from cli.managed import ManagedClaudeSessionManager
|
|
@@ -87,7 +87,7 @@ class AppRuntime:
|
|
| 87 |
|
| 88 |
app: FastAPI
|
| 89 |
settings: Settings
|
| 90 |
-
|
| 91 |
messaging_runtime: MessagingRuntime | None = None
|
| 92 |
messaging_workflow: MessagingWorkflow | None = None
|
| 93 |
cli_manager: ManagedClaudeSessionManager | None = None
|
|
@@ -103,12 +103,12 @@ class AppRuntime:
|
|
| 103 |
async def startup(self) -> None:
|
| 104 |
logger.info("Starting Claude Code Proxy...")
|
| 105 |
admin_url = local_admin_url(self.settings)
|
| 106 |
-
self.
|
| 107 |
-
self.app.state.
|
| 108 |
try:
|
| 109 |
warn_if_process_auth_token(self.settings)
|
| 110 |
await self._validate_configured_models_best_effort()
|
| 111 |
-
self.
|
| 112 |
await self._start_messaging_if_configured()
|
| 113 |
self._publish_state()
|
| 114 |
logging.getLogger("uvicorn.error").info(
|
|
@@ -117,18 +117,18 @@ class AppRuntime:
|
|
| 117 |
except Exception as exc:
|
| 118 |
log_startup_failure(self.settings, exc)
|
| 119 |
await best_effort(
|
| 120 |
-
"
|
| 121 |
-
self.
|
| 122 |
log_verbose_errors=self.settings.log_api_error_tracebacks,
|
| 123 |
)
|
| 124 |
raise
|
| 125 |
|
| 126 |
async def _validate_configured_models_best_effort(self) -> None:
|
| 127 |
"""Warm validation status without blocking first-run/admin access."""
|
| 128 |
-
if self.
|
| 129 |
return
|
| 130 |
try:
|
| 131 |
-
await self.
|
| 132 |
except ServiceUnavailableError as exc:
|
| 133 |
self.app.state.startup_validation_error = exc.message
|
| 134 |
logger.warning(
|
|
@@ -165,10 +165,10 @@ class AppRuntime:
|
|
| 165 |
self.cli_manager.stop_all(),
|
| 166 |
log_verbose_errors=verbose,
|
| 167 |
)
|
| 168 |
-
if self.
|
| 169 |
await best_effort(
|
| 170 |
-
"
|
| 171 |
-
self.
|
| 172 |
log_verbose_errors=verbose,
|
| 173 |
)
|
| 174 |
await self._shutdown_limiter()
|
|
|
|
| 14 |
from api.admin_urls import local_admin_url
|
| 15 |
from config.settings import Settings, get_settings
|
| 16 |
from providers.exceptions import ServiceUnavailableError
|
| 17 |
+
from providers.runtime import ProviderRuntime
|
| 18 |
|
| 19 |
if TYPE_CHECKING:
|
| 20 |
from cli.managed import ManagedClaudeSessionManager
|
|
|
|
| 87 |
|
| 88 |
app: FastAPI
|
| 89 |
settings: Settings
|
| 90 |
+
_provider_runtime: ProviderRuntime | None = field(default=None, init=False)
|
| 91 |
messaging_runtime: MessagingRuntime | None = None
|
| 92 |
messaging_workflow: MessagingWorkflow | None = None
|
| 93 |
cli_manager: ManagedClaudeSessionManager | None = None
|
|
|
|
| 103 |
async def startup(self) -> None:
|
| 104 |
logger.info("Starting Claude Code Proxy...")
|
| 105 |
admin_url = local_admin_url(self.settings)
|
| 106 |
+
self._provider_runtime = ProviderRuntime(self.settings)
|
| 107 |
+
self.app.state.provider_runtime = self._provider_runtime
|
| 108 |
try:
|
| 109 |
warn_if_process_auth_token(self.settings)
|
| 110 |
await self._validate_configured_models_best_effort()
|
| 111 |
+
self._provider_runtime.start_model_list_refresh()
|
| 112 |
await self._start_messaging_if_configured()
|
| 113 |
self._publish_state()
|
| 114 |
logging.getLogger("uvicorn.error").info(
|
|
|
|
| 117 |
except Exception as exc:
|
| 118 |
log_startup_failure(self.settings, exc)
|
| 119 |
await best_effort(
|
| 120 |
+
"provider_runtime.cleanup",
|
| 121 |
+
self._provider_runtime.cleanup(),
|
| 122 |
log_verbose_errors=self.settings.log_api_error_tracebacks,
|
| 123 |
)
|
| 124 |
raise
|
| 125 |
|
| 126 |
async def _validate_configured_models_best_effort(self) -> None:
|
| 127 |
"""Warm validation status without blocking first-run/admin access."""
|
| 128 |
+
if self._provider_runtime is None:
|
| 129 |
return
|
| 130 |
try:
|
| 131 |
+
await self._provider_runtime.validate_configured_models()
|
| 132 |
except ServiceUnavailableError as exc:
|
| 133 |
self.app.state.startup_validation_error = exc.message
|
| 134 |
logger.warning(
|
|
|
|
| 165 |
self.cli_manager.stop_all(),
|
| 166 |
log_verbose_errors=verbose,
|
| 167 |
)
|
| 168 |
+
if self._provider_runtime is not None:
|
| 169 |
await best_effort(
|
| 170 |
+
"provider_runtime.cleanup",
|
| 171 |
+
self._provider_runtime.cleanup(),
|
| 172 |
log_verbose_errors=verbose,
|
| 173 |
)
|
| 174 |
await self._shutdown_limiter()
|
|
@@ -1,6 +1,6 @@
|
|
| 1 |
"""Neutral provider catalog: IDs, credentials, defaults, proxy and capability metadata.
|
| 2 |
|
| 3 |
-
Adapter factories live in :mod:`providers.
|
| 4 |
provider implementation imports (see contract tests).
|
| 5 |
"""
|
| 6 |
|
|
|
|
| 1 |
"""Neutral provider catalog: IDs, credentials, defaults, proxy and capability metadata.
|
| 2 |
|
| 3 |
+
Adapter factories live in :mod:`providers.runtime.factory`; this module stays free of
|
| 4 |
provider implementation imports (see contract tests).
|
| 5 |
"""
|
| 6 |
|
|
@@ -1,527 +0,0 @@
|
|
| 1 |
-
"""Provider descriptors, factory, and runtime registry."""
|
| 2 |
-
|
| 3 |
-
from __future__ import annotations
|
| 4 |
-
|
| 5 |
-
import asyncio
|
| 6 |
-
from collections import defaultdict
|
| 7 |
-
from collections.abc import Callable, Iterable, MutableMapping
|
| 8 |
-
from contextlib import suppress
|
| 9 |
-
|
| 10 |
-
import httpx
|
| 11 |
-
from loguru import logger
|
| 12 |
-
|
| 13 |
-
from config.provider_catalog import (
|
| 14 |
-
PROVIDER_CATALOG,
|
| 15 |
-
SUPPORTED_PROVIDER_IDS,
|
| 16 |
-
ProviderDescriptor,
|
| 17 |
-
)
|
| 18 |
-
from config.settings import ConfiguredChatModelRef, Settings
|
| 19 |
-
from providers.base import BaseProvider, ProviderConfig
|
| 20 |
-
from providers.exceptions import (
|
| 21 |
-
AuthenticationError,
|
| 22 |
-
ModelListResponseError,
|
| 23 |
-
ProviderError,
|
| 24 |
-
ServiceUnavailableError,
|
| 25 |
-
UnknownProviderTypeError,
|
| 26 |
-
)
|
| 27 |
-
from providers.model_listing import ProviderModelInfo, model_infos_from_ids
|
| 28 |
-
|
| 29 |
-
ProviderFactory = Callable[[ProviderConfig, Settings], BaseProvider]
|
| 30 |
-
|
| 31 |
-
# Backwards-compatible name for the catalog (single source: ``config.provider_catalog``).
|
| 32 |
-
PROVIDER_DESCRIPTORS: dict[str, ProviderDescriptor] = PROVIDER_CATALOG
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
def _create_nvidia_nim(config: ProviderConfig, settings: Settings) -> BaseProvider:
|
| 36 |
-
from providers.nvidia_nim import NvidiaNimProvider
|
| 37 |
-
|
| 38 |
-
return NvidiaNimProvider(config, nim_settings=settings.nim)
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
def _create_open_router(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 42 |
-
from providers.open_router import OpenRouterProvider
|
| 43 |
-
|
| 44 |
-
return OpenRouterProvider(config)
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
def _create_mistral(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 48 |
-
from providers.mistral import MistralProvider
|
| 49 |
-
|
| 50 |
-
return MistralProvider(config)
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
def _create_mistral_codestral(
|
| 54 |
-
config: ProviderConfig, _settings: Settings
|
| 55 |
-
) -> BaseProvider:
|
| 56 |
-
from providers.codestral import CodestralProvider
|
| 57 |
-
|
| 58 |
-
return CodestralProvider(config)
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
def _create_deepseek(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 62 |
-
from providers.deepseek import DeepSeekProvider
|
| 63 |
-
|
| 64 |
-
return DeepSeekProvider(config)
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
def _create_lmstudio(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 68 |
-
from providers.lmstudio import LMStudioProvider
|
| 69 |
-
|
| 70 |
-
return LMStudioProvider(config)
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
def _create_llamacpp(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 74 |
-
from providers.llamacpp import LlamaCppProvider
|
| 75 |
-
|
| 76 |
-
return LlamaCppProvider(config)
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
def _create_ollama(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 80 |
-
from providers.ollama import OllamaProvider
|
| 81 |
-
|
| 82 |
-
return OllamaProvider(config)
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
def _create_kimi(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 86 |
-
from providers.kimi import KimiProvider
|
| 87 |
-
|
| 88 |
-
return KimiProvider(config)
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
def _create_wafer(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 92 |
-
from providers.wafer import WaferProvider
|
| 93 |
-
|
| 94 |
-
return WaferProvider(config)
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
def _create_opencode(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 98 |
-
from providers.opencode import OpenCodeProvider
|
| 99 |
-
|
| 100 |
-
return OpenCodeProvider(config)
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
def _create_opencode_go(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 104 |
-
from providers.opencode import OpenCodeProvider
|
| 105 |
-
|
| 106 |
-
return OpenCodeProvider(config, provider_name="OPENCODE_GO")
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
def _create_zai(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 110 |
-
from providers.zai import ZaiProvider
|
| 111 |
-
|
| 112 |
-
return ZaiProvider(config)
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
def _create_fireworks(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 116 |
-
from providers.fireworks import FireworksProvider
|
| 117 |
-
|
| 118 |
-
return FireworksProvider(config)
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
def _create_gemini(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 122 |
-
from providers.gemini import GeminiProvider
|
| 123 |
-
|
| 124 |
-
return GeminiProvider(config)
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
def _create_groq(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 128 |
-
from providers.groq import GroqProvider
|
| 129 |
-
|
| 130 |
-
return GroqProvider(config)
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
def _create_cerebras(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 134 |
-
from providers.cerebras import CerebrasProvider
|
| 135 |
-
|
| 136 |
-
return CerebrasProvider(config)
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
PROVIDER_FACTORIES: dict[str, ProviderFactory] = {
|
| 140 |
-
"nvidia_nim": _create_nvidia_nim,
|
| 141 |
-
"open_router": _create_open_router,
|
| 142 |
-
"gemini": _create_gemini,
|
| 143 |
-
"deepseek": _create_deepseek,
|
| 144 |
-
"mistral": _create_mistral,
|
| 145 |
-
"mistral_codestral": _create_mistral_codestral,
|
| 146 |
-
"opencode": _create_opencode,
|
| 147 |
-
"opencode_go": _create_opencode_go,
|
| 148 |
-
"wafer": _create_wafer,
|
| 149 |
-
"kimi": _create_kimi,
|
| 150 |
-
"cerebras": _create_cerebras,
|
| 151 |
-
"groq": _create_groq,
|
| 152 |
-
"fireworks": _create_fireworks,
|
| 153 |
-
"zai": _create_zai,
|
| 154 |
-
"lmstudio": _create_lmstudio,
|
| 155 |
-
"llamacpp": _create_llamacpp,
|
| 156 |
-
"ollama": _create_ollama,
|
| 157 |
-
}
|
| 158 |
-
|
| 159 |
-
if set(PROVIDER_DESCRIPTORS) != set(SUPPORTED_PROVIDER_IDS) or set(
|
| 160 |
-
PROVIDER_FACTORIES
|
| 161 |
-
) != set(SUPPORTED_PROVIDER_IDS):
|
| 162 |
-
raise AssertionError(
|
| 163 |
-
"PROVIDER_DESCRIPTORS, PROVIDER_FACTORIES, and SUPPORTED_PROVIDER_IDS are out of sync: "
|
| 164 |
-
f"descriptors={set(PROVIDER_DESCRIPTORS)!r} factories={set(PROVIDER_FACTORIES)!r} "
|
| 165 |
-
f"ids={set(SUPPORTED_PROVIDER_IDS)!r}"
|
| 166 |
-
)
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
def _string_attr(settings: Settings, attr_name: str | None, default: str = "") -> str:
|
| 170 |
-
if attr_name is None:
|
| 171 |
-
return default
|
| 172 |
-
value = getattr(settings, attr_name, default)
|
| 173 |
-
return value if isinstance(value, str) else default
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
def _credential_for(descriptor: ProviderDescriptor, settings: Settings) -> str:
|
| 177 |
-
if descriptor.static_credential is not None:
|
| 178 |
-
return descriptor.static_credential
|
| 179 |
-
if descriptor.credential_attr:
|
| 180 |
-
return _string_attr(settings, descriptor.credential_attr)
|
| 181 |
-
return ""
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
def _require_credential(descriptor: ProviderDescriptor, credential: str) -> None:
|
| 185 |
-
if descriptor.credential_env is None:
|
| 186 |
-
return
|
| 187 |
-
if credential and credential.strip():
|
| 188 |
-
return
|
| 189 |
-
message = f"{descriptor.credential_env} is not set. Add it to your .env file."
|
| 190 |
-
if descriptor.credential_url:
|
| 191 |
-
message = f"{message} Get a key at {descriptor.credential_url}"
|
| 192 |
-
raise AuthenticationError(message)
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
def build_provider_config(
|
| 196 |
-
descriptor: ProviderDescriptor, settings: Settings
|
| 197 |
-
) -> ProviderConfig:
|
| 198 |
-
credential = _credential_for(descriptor, settings)
|
| 199 |
-
_require_credential(descriptor, credential)
|
| 200 |
-
base_url = _string_attr(
|
| 201 |
-
settings, descriptor.base_url_attr, descriptor.default_base_url or ""
|
| 202 |
-
)
|
| 203 |
-
proxy = _string_attr(settings, descriptor.proxy_attr)
|
| 204 |
-
return ProviderConfig(
|
| 205 |
-
api_key=credential,
|
| 206 |
-
base_url=base_url or descriptor.default_base_url,
|
| 207 |
-
rate_limit=settings.provider_rate_limit,
|
| 208 |
-
rate_window=settings.provider_rate_window,
|
| 209 |
-
max_concurrency=settings.provider_max_concurrency,
|
| 210 |
-
http_read_timeout=settings.http_read_timeout,
|
| 211 |
-
http_write_timeout=settings.http_write_timeout,
|
| 212 |
-
http_connect_timeout=settings.http_connect_timeout,
|
| 213 |
-
enable_thinking=settings.enable_model_thinking,
|
| 214 |
-
proxy=proxy,
|
| 215 |
-
log_raw_sse_events=settings.log_raw_sse_events,
|
| 216 |
-
log_api_error_tracebacks=settings.log_api_error_tracebacks,
|
| 217 |
-
)
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
def create_provider(provider_id: str, settings: Settings) -> BaseProvider:
|
| 221 |
-
descriptor = PROVIDER_DESCRIPTORS.get(provider_id)
|
| 222 |
-
if descriptor is None:
|
| 223 |
-
supported = "', '".join(PROVIDER_DESCRIPTORS)
|
| 224 |
-
raise UnknownProviderTypeError(
|
| 225 |
-
f"Unknown provider_type: '{provider_id}'. Supported: '{supported}'"
|
| 226 |
-
)
|
| 227 |
-
|
| 228 |
-
config = build_provider_config(descriptor, settings)
|
| 229 |
-
factory = PROVIDER_FACTORIES.get(provider_id)
|
| 230 |
-
if factory is None:
|
| 231 |
-
raise AssertionError(f"Unhandled provider descriptor: {provider_id}")
|
| 232 |
-
return factory(config, settings)
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
def _format_provider_query_failures(
|
| 236 |
-
refs: list[ConfiguredChatModelRef],
|
| 237 |
-
exc: BaseException,
|
| 238 |
-
settings: Settings,
|
| 239 |
-
) -> list[str]:
|
| 240 |
-
reason = _provider_query_failure_reason(exc, settings)
|
| 241 |
-
return [_format_model_validation_failure(ref, reason) for ref in refs]
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
def _format_missing_model_failure(ref: ConfiguredChatModelRef) -> str:
|
| 245 |
-
return _format_model_validation_failure(ref, "missing model")
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
def _format_model_validation_failure(ref: ConfiguredChatModelRef, problem: str) -> str:
|
| 249 |
-
return (
|
| 250 |
-
f"sources={','.join(ref.sources)} provider={ref.provider_id} "
|
| 251 |
-
f"model={ref.model_id} problem={problem}"
|
| 252 |
-
)
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
def _provider_query_failure_reason(
|
| 256 |
-
exc: BaseException,
|
| 257 |
-
settings: Settings,
|
| 258 |
-
) -> str:
|
| 259 |
-
if isinstance(exc, ModelListResponseError):
|
| 260 |
-
return f"malformed model-list response: {exc.message}"
|
| 261 |
-
if isinstance(exc, httpx.HTTPStatusError):
|
| 262 |
-
return f"query failure: HTTP {exc.response.status_code}"
|
| 263 |
-
if isinstance(exc, AuthenticationError):
|
| 264 |
-
return f"query failure: {exc.message}"
|
| 265 |
-
if isinstance(exc, ProviderError) and settings.log_api_error_tracebacks:
|
| 266 |
-
return f"query failure: {exc.message}"
|
| 267 |
-
return f"query failure: {type(exc).__name__}"
|
| 268 |
-
|
| 269 |
-
|
| 270 |
-
def _referenced_provider_ids(settings: Settings) -> frozenset[str]:
|
| 271 |
-
return frozenset(ref.provider_id for ref in settings.configured_chat_model_refs())
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
def _model_list_provider_ids_for_settings(settings: Settings) -> tuple[str, ...]:
|
| 275 |
-
"""Return providers worth discovering for this process configuration."""
|
| 276 |
-
referenced_provider_ids = _referenced_provider_ids(settings)
|
| 277 |
-
provider_ids: list[str] = []
|
| 278 |
-
for provider_id, descriptor in PROVIDER_DESCRIPTORS.items():
|
| 279 |
-
if descriptor.static_credential is not None:
|
| 280 |
-
if provider_id in referenced_provider_ids:
|
| 281 |
-
provider_ids.append(provider_id)
|
| 282 |
-
continue
|
| 283 |
-
if (
|
| 284 |
-
descriptor.credential_env is not None
|
| 285 |
-
and _credential_for(descriptor, settings).strip()
|
| 286 |
-
):
|
| 287 |
-
provider_ids.append(provider_id)
|
| 288 |
-
return tuple(provider_ids)
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
def _log_model_discovery_failure(
|
| 292 |
-
provider_id: str, exc: BaseException, settings: Settings
|
| 293 |
-
) -> None:
|
| 294 |
-
logger.warning(
|
| 295 |
-
"Provider model discovery skipped: provider={} reason={}",
|
| 296 |
-
provider_id,
|
| 297 |
-
_provider_query_failure_reason(exc, settings),
|
| 298 |
-
)
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
class ProviderRegistry:
|
| 302 |
-
"""Cache and clean up provider instances by provider id."""
|
| 303 |
-
|
| 304 |
-
def __init__(self, providers: MutableMapping[str, BaseProvider] | None = None):
|
| 305 |
-
self._providers = providers if providers is not None else {}
|
| 306 |
-
self._model_ids_by_provider: dict[str, frozenset[str]] = {}
|
| 307 |
-
self._model_infos_by_provider: dict[str, dict[str, ProviderModelInfo]] = {}
|
| 308 |
-
self._model_list_refresh_task: asyncio.Task[None] | None = None
|
| 309 |
-
|
| 310 |
-
def is_cached(self, provider_id: str) -> bool:
|
| 311 |
-
"""Return whether a provider for this id is already in the cache."""
|
| 312 |
-
return provider_id in self._providers
|
| 313 |
-
|
| 314 |
-
def get(self, provider_id: str, settings: Settings) -> BaseProvider:
|
| 315 |
-
if provider_id not in self._providers:
|
| 316 |
-
self._providers[provider_id] = create_provider(provider_id, settings)
|
| 317 |
-
return self._providers[provider_id]
|
| 318 |
-
|
| 319 |
-
def cache_model_ids(self, provider_id: str, model_ids: Iterable[str]) -> None:
|
| 320 |
-
"""Store a provider model-list result for later instant API responses."""
|
| 321 |
-
self.cache_model_infos(provider_id, model_infos_from_ids(model_ids))
|
| 322 |
-
|
| 323 |
-
def cache_model_infos(
|
| 324 |
-
self, provider_id: str, model_infos: Iterable[ProviderModelInfo]
|
| 325 |
-
) -> None:
|
| 326 |
-
"""Store provider model metadata for later instant API responses."""
|
| 327 |
-
clean_infos = {
|
| 328 |
-
info.model_id: info for info in model_infos if info.model_id.strip()
|
| 329 |
-
}
|
| 330 |
-
self._model_infos_by_provider[provider_id] = clean_infos
|
| 331 |
-
self._model_ids_by_provider[provider_id] = frozenset(clean_infos)
|
| 332 |
-
|
| 333 |
-
def cached_model_ids(self) -> dict[str, frozenset[str]]:
|
| 334 |
-
"""Return a copy of cached raw provider model ids."""
|
| 335 |
-
return dict(self._model_ids_by_provider)
|
| 336 |
-
|
| 337 |
-
def cached_model_supports_thinking(
|
| 338 |
-
self, provider_id: str, model_id: str
|
| 339 |
-
) -> bool | None:
|
| 340 |
-
"""Return cached thinking support when a provider exposes it."""
|
| 341 |
-
info = self._model_infos_by_provider.get(provider_id, {}).get(model_id)
|
| 342 |
-
if info is None:
|
| 343 |
-
return None
|
| 344 |
-
return info.supports_thinking
|
| 345 |
-
|
| 346 |
-
def cached_prefixed_model_refs(self) -> tuple[str, ...]:
|
| 347 |
-
"""Return cached provider models in user-selectable ``provider/model`` form."""
|
| 348 |
-
return tuple(info.model_id for info in self.cached_prefixed_model_infos())
|
| 349 |
-
|
| 350 |
-
def cached_prefixed_model_infos(self) -> tuple[ProviderModelInfo, ...]:
|
| 351 |
-
"""Return cached provider models with user-selectable prefixed ids."""
|
| 352 |
-
infos: list[ProviderModelInfo] = []
|
| 353 |
-
for provider_id in SUPPORTED_PROVIDER_IDS:
|
| 354 |
-
provider_infos = self._model_infos_by_provider.get(provider_id, {})
|
| 355 |
-
infos.extend(
|
| 356 |
-
ProviderModelInfo(
|
| 357 |
-
model_id=f"{provider_id}/{info.model_id}",
|
| 358 |
-
supports_thinking=info.supports_thinking,
|
| 359 |
-
)
|
| 360 |
-
for info in sorted(
|
| 361 |
-
provider_infos.values(), key=lambda item: item.model_id
|
| 362 |
-
)
|
| 363 |
-
)
|
| 364 |
-
return tuple(infos)
|
| 365 |
-
|
| 366 |
-
async def refresh_model_list_cache(
|
| 367 |
-
self, settings: Settings, *, only_missing: bool = False
|
| 368 |
-
) -> None:
|
| 369 |
-
"""Best-effort refresh of model lists for providers usable in this process."""
|
| 370 |
-
provider_ids = _model_list_provider_ids_for_settings(settings)
|
| 371 |
-
if only_missing:
|
| 372 |
-
provider_ids = tuple(
|
| 373 |
-
provider_id
|
| 374 |
-
for provider_id in provider_ids
|
| 375 |
-
if provider_id not in self._model_ids_by_provider
|
| 376 |
-
)
|
| 377 |
-
await self._refresh_model_ids(settings, provider_ids)
|
| 378 |
-
|
| 379 |
-
def start_model_list_refresh(self, settings: Settings) -> None:
|
| 380 |
-
"""Start a non-blocking cache warmup for missing eligible provider lists."""
|
| 381 |
-
if (
|
| 382 |
-
self._model_list_refresh_task is not None
|
| 383 |
-
and not self._model_list_refresh_task.done()
|
| 384 |
-
):
|
| 385 |
-
return
|
| 386 |
-
|
| 387 |
-
provider_ids = tuple(
|
| 388 |
-
provider_id
|
| 389 |
-
for provider_id in _model_list_provider_ids_for_settings(settings)
|
| 390 |
-
if provider_id not in self._model_ids_by_provider
|
| 391 |
-
)
|
| 392 |
-
if not provider_ids:
|
| 393 |
-
logger.info(
|
| 394 |
-
"Provider model discovery cache already warm: providers={}",
|
| 395 |
-
len(self._model_ids_by_provider),
|
| 396 |
-
)
|
| 397 |
-
return
|
| 398 |
-
|
| 399 |
-
self._model_list_refresh_task = asyncio.create_task(
|
| 400 |
-
self._run_model_list_refresh(settings, provider_ids)
|
| 401 |
-
)
|
| 402 |
-
|
| 403 |
-
async def _run_model_list_refresh(
|
| 404 |
-
self, settings: Settings, provider_ids: tuple[str, ...]
|
| 405 |
-
) -> None:
|
| 406 |
-
try:
|
| 407 |
-
await self._refresh_model_ids(settings, provider_ids)
|
| 408 |
-
except asyncio.CancelledError:
|
| 409 |
-
raise
|
| 410 |
-
except Exception as exc:
|
| 411 |
-
logger.warning(
|
| 412 |
-
"Provider model discovery task failed: exc_type={}",
|
| 413 |
-
type(exc).__name__,
|
| 414 |
-
)
|
| 415 |
-
|
| 416 |
-
async def _refresh_model_ids(
|
| 417 |
-
self, settings: Settings, provider_ids: tuple[str, ...]
|
| 418 |
-
) -> None:
|
| 419 |
-
tasks: dict[str, asyncio.Task[frozenset[ProviderModelInfo]]] = {}
|
| 420 |
-
for provider_id in provider_ids:
|
| 421 |
-
try:
|
| 422 |
-
provider = self.get(provider_id, settings)
|
| 423 |
-
except Exception as exc:
|
| 424 |
-
_log_model_discovery_failure(provider_id, exc, settings)
|
| 425 |
-
continue
|
| 426 |
-
tasks[provider_id] = asyncio.create_task(provider.list_model_infos())
|
| 427 |
-
|
| 428 |
-
if not tasks:
|
| 429 |
-
return
|
| 430 |
-
|
| 431 |
-
results = await asyncio.gather(*tasks.values(), return_exceptions=True)
|
| 432 |
-
for (provider_id, _task), result in zip(tasks.items(), results, strict=True):
|
| 433 |
-
if isinstance(result, BaseException):
|
| 434 |
-
if isinstance(result, asyncio.CancelledError):
|
| 435 |
-
raise result
|
| 436 |
-
_log_model_discovery_failure(provider_id, result, settings)
|
| 437 |
-
continue
|
| 438 |
-
self.cache_model_infos(provider_id, result)
|
| 439 |
-
logger.info(
|
| 440 |
-
"Provider model discovery cached: provider={} models={}",
|
| 441 |
-
provider_id,
|
| 442 |
-
len(result),
|
| 443 |
-
)
|
| 444 |
-
|
| 445 |
-
async def validate_configured_models(self, settings: Settings) -> None:
|
| 446 |
-
"""Fail fast unless every configured chat model exists upstream."""
|
| 447 |
-
refs = settings.configured_chat_model_refs()
|
| 448 |
-
refs_by_provider: dict[str, list[ConfiguredChatModelRef]] = defaultdict(list)
|
| 449 |
-
for ref in refs:
|
| 450 |
-
refs_by_provider[ref.provider_id].append(ref)
|
| 451 |
-
|
| 452 |
-
failures: list[str] = []
|
| 453 |
-
tasks: dict[str, asyncio.Task[frozenset[ProviderModelInfo]]] = {}
|
| 454 |
-
for provider_id, provider_refs in refs_by_provider.items():
|
| 455 |
-
try:
|
| 456 |
-
provider = self.get(provider_id, settings)
|
| 457 |
-
except Exception as exc:
|
| 458 |
-
failures.extend(
|
| 459 |
-
_format_provider_query_failures(provider_refs, exc, settings)
|
| 460 |
-
)
|
| 461 |
-
continue
|
| 462 |
-
tasks[provider_id] = asyncio.create_task(provider.list_model_infos())
|
| 463 |
-
|
| 464 |
-
if tasks:
|
| 465 |
-
results = await asyncio.gather(*tasks.values(), return_exceptions=True)
|
| 466 |
-
for (provider_id, _task), result in zip(
|
| 467 |
-
tasks.items(), results, strict=True
|
| 468 |
-
):
|
| 469 |
-
provider_refs = refs_by_provider[provider_id]
|
| 470 |
-
if isinstance(result, BaseException):
|
| 471 |
-
if isinstance(result, asyncio.CancelledError):
|
| 472 |
-
raise result
|
| 473 |
-
failures.extend(
|
| 474 |
-
_format_provider_query_failures(provider_refs, result, settings)
|
| 475 |
-
)
|
| 476 |
-
continue
|
| 477 |
-
self.cache_model_infos(provider_id, result)
|
| 478 |
-
model_ids = self._model_ids_by_provider[provider_id]
|
| 479 |
-
failures.extend(
|
| 480 |
-
_format_missing_model_failure(ref)
|
| 481 |
-
for ref in provider_refs
|
| 482 |
-
if ref.model_id not in model_ids
|
| 483 |
-
)
|
| 484 |
-
|
| 485 |
-
if failures:
|
| 486 |
-
message = "Configured model validation failed:\n" + "\n".join(
|
| 487 |
-
f"- {failure}" for failure in failures
|
| 488 |
-
)
|
| 489 |
-
raise ServiceUnavailableError(message)
|
| 490 |
-
|
| 491 |
-
logger.info(
|
| 492 |
-
"Configured provider models validated: models={} providers={}",
|
| 493 |
-
len(refs),
|
| 494 |
-
len(refs_by_provider),
|
| 495 |
-
)
|
| 496 |
-
|
| 497 |
-
async def cleanup(self) -> None:
|
| 498 |
-
"""Call ``cleanup`` on every cached provider, then clear the cache.
|
| 499 |
-
|
| 500 |
-
Attempts all providers even if one fails. A single failure is re-raised
|
| 501 |
-
as-is; multiple failures are wrapped in :exc:`ExceptionGroup`.
|
| 502 |
-
"""
|
| 503 |
-
if (
|
| 504 |
-
self._model_list_refresh_task is not None
|
| 505 |
-
and not self._model_list_refresh_task.done()
|
| 506 |
-
):
|
| 507 |
-
self._model_list_refresh_task.cancel()
|
| 508 |
-
with suppress(asyncio.CancelledError):
|
| 509 |
-
await self._model_list_refresh_task
|
| 510 |
-
|
| 511 |
-
items = list(self._providers.items())
|
| 512 |
-
errors: list[Exception] = []
|
| 513 |
-
try:
|
| 514 |
-
for _pid, provider in items:
|
| 515 |
-
try:
|
| 516 |
-
await provider.cleanup()
|
| 517 |
-
except Exception as e:
|
| 518 |
-
errors.append(e)
|
| 519 |
-
finally:
|
| 520 |
-
self._providers.clear()
|
| 521 |
-
self._model_ids_by_provider.clear()
|
| 522 |
-
self._model_infos_by_provider.clear()
|
| 523 |
-
if len(errors) == 1:
|
| 524 |
-
raise errors[0]
|
| 525 |
-
if len(errors) > 1:
|
| 526 |
-
msg = "One or more provider cleanups failed"
|
| 527 |
-
raise ExceptionGroup(msg, errors)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""App-scoped provider runtime facade."""
|
| 2 |
+
|
| 3 |
+
from .config import build_provider_config
|
| 4 |
+
from .factory import PROVIDER_FACTORIES, create_provider
|
| 5 |
+
from .runtime import ProviderRuntime
|
| 6 |
+
|
| 7 |
+
__all__ = [
|
| 8 |
+
"PROVIDER_FACTORIES",
|
| 9 |
+
"ProviderRuntime",
|
| 10 |
+
"build_provider_config",
|
| 11 |
+
"create_provider",
|
| 12 |
+
]
|
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Provider instance cache and cleanup."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from collections.abc import Callable, MutableMapping
|
| 6 |
+
|
| 7 |
+
from config.settings import Settings
|
| 8 |
+
from providers.base import BaseProvider
|
| 9 |
+
|
| 10 |
+
from .factory import create_provider
|
| 11 |
+
|
| 12 |
+
ProviderCreator = Callable[[str, Settings], BaseProvider]
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class ProviderCache:
|
| 16 |
+
"""Cache provider instances for one settings snapshot."""
|
| 17 |
+
|
| 18 |
+
def __init__(
|
| 19 |
+
self,
|
| 20 |
+
settings: Settings,
|
| 21 |
+
providers: MutableMapping[str, BaseProvider] | None = None,
|
| 22 |
+
*,
|
| 23 |
+
creator: ProviderCreator = create_provider,
|
| 24 |
+
) -> None:
|
| 25 |
+
self._settings = settings
|
| 26 |
+
self._providers = providers if providers is not None else {}
|
| 27 |
+
self._creator = creator
|
| 28 |
+
|
| 29 |
+
def is_cached(self, provider_id: str) -> bool:
|
| 30 |
+
"""Return whether a provider for this id is already cached."""
|
| 31 |
+
return provider_id in self._providers
|
| 32 |
+
|
| 33 |
+
def get(self, provider_id: str) -> BaseProvider:
|
| 34 |
+
"""Return an existing provider or create it lazily."""
|
| 35 |
+
if provider_id not in self._providers:
|
| 36 |
+
self._providers[provider_id] = self._creator(provider_id, self._settings)
|
| 37 |
+
return self._providers[provider_id]
|
| 38 |
+
|
| 39 |
+
async def cleanup(self) -> None:
|
| 40 |
+
"""Clean up every cached provider, then clear the cache."""
|
| 41 |
+
items = list(self._providers.items())
|
| 42 |
+
errors: list[Exception] = []
|
| 43 |
+
try:
|
| 44 |
+
for _provider_id, provider in items:
|
| 45 |
+
try:
|
| 46 |
+
await provider.cleanup()
|
| 47 |
+
except Exception as exc:
|
| 48 |
+
errors.append(exc)
|
| 49 |
+
finally:
|
| 50 |
+
self._providers.clear()
|
| 51 |
+
if len(errors) == 1:
|
| 52 |
+
raise errors[0]
|
| 53 |
+
if len(errors) > 1:
|
| 54 |
+
raise ExceptionGroup("One or more provider cleanups failed", errors)
|
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Provider configuration construction from neutral catalog metadata."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from config.provider_catalog import ProviderDescriptor
|
| 6 |
+
from config.settings import Settings
|
| 7 |
+
from providers.base import ProviderConfig
|
| 8 |
+
from providers.exceptions import AuthenticationError
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def string_setting(settings: Settings, attr_name: str | None, default: str = "") -> str:
|
| 12 |
+
"""Return a string-valued settings attribute, ignoring non-string mocks."""
|
| 13 |
+
if attr_name is None:
|
| 14 |
+
return default
|
| 15 |
+
value = getattr(settings, attr_name, default)
|
| 16 |
+
return value if isinstance(value, str) else default
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def provider_credential(descriptor: ProviderDescriptor, settings: Settings) -> str:
|
| 20 |
+
"""Return the configured credential for a provider descriptor."""
|
| 21 |
+
if descriptor.static_credential is not None:
|
| 22 |
+
return descriptor.static_credential
|
| 23 |
+
if descriptor.credential_attr:
|
| 24 |
+
return string_setting(settings, descriptor.credential_attr)
|
| 25 |
+
return ""
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def require_provider_credential(
|
| 29 |
+
descriptor: ProviderDescriptor, credential: str
|
| 30 |
+
) -> None:
|
| 31 |
+
"""Raise a user-facing configuration error when a required key is missing."""
|
| 32 |
+
if descriptor.credential_env is None:
|
| 33 |
+
return
|
| 34 |
+
if credential and credential.strip():
|
| 35 |
+
return
|
| 36 |
+
message = f"{descriptor.credential_env} is not set. Add it to your .env file."
|
| 37 |
+
if descriptor.credential_url:
|
| 38 |
+
message = f"{message} Get a key at {descriptor.credential_url}"
|
| 39 |
+
raise AuthenticationError(message)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def build_provider_config(
|
| 43 |
+
descriptor: ProviderDescriptor, settings: Settings
|
| 44 |
+
) -> ProviderConfig:
|
| 45 |
+
"""Build shared provider configuration for one provider descriptor."""
|
| 46 |
+
credential = provider_credential(descriptor, settings)
|
| 47 |
+
require_provider_credential(descriptor, credential)
|
| 48 |
+
base_url = string_setting(
|
| 49 |
+
settings, descriptor.base_url_attr, descriptor.default_base_url or ""
|
| 50 |
+
)
|
| 51 |
+
proxy = string_setting(settings, descriptor.proxy_attr)
|
| 52 |
+
return ProviderConfig(
|
| 53 |
+
api_key=credential,
|
| 54 |
+
base_url=base_url or descriptor.default_base_url,
|
| 55 |
+
rate_limit=settings.provider_rate_limit,
|
| 56 |
+
rate_window=settings.provider_rate_window,
|
| 57 |
+
max_concurrency=settings.provider_max_concurrency,
|
| 58 |
+
http_read_timeout=settings.http_read_timeout,
|
| 59 |
+
http_write_timeout=settings.http_write_timeout,
|
| 60 |
+
http_connect_timeout=settings.http_connect_timeout,
|
| 61 |
+
enable_thinking=settings.enable_model_thinking,
|
| 62 |
+
proxy=proxy,
|
| 63 |
+
log_raw_sse_events=settings.log_raw_sse_events,
|
| 64 |
+
log_api_error_tracebacks=settings.log_api_error_tracebacks,
|
| 65 |
+
)
|
|
@@ -0,0 +1,140 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Provider model-list discovery and background refresh."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import asyncio
|
| 6 |
+
from collections.abc import Callable
|
| 7 |
+
from contextlib import suppress
|
| 8 |
+
|
| 9 |
+
from loguru import logger
|
| 10 |
+
|
| 11 |
+
from config.provider_catalog import PROVIDER_CATALOG
|
| 12 |
+
from config.settings import Settings
|
| 13 |
+
from providers.base import BaseProvider
|
| 14 |
+
from providers.model_listing import ProviderModelInfo
|
| 15 |
+
|
| 16 |
+
from .config import provider_credential
|
| 17 |
+
from .model_cache import ProviderModelCache
|
| 18 |
+
from .validation import provider_query_failure_reason
|
| 19 |
+
|
| 20 |
+
ProviderResolver = Callable[[str], BaseProvider]
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def referenced_provider_ids(settings: Settings) -> frozenset[str]:
|
| 24 |
+
"""Return provider ids referenced by configured chat model refs."""
|
| 25 |
+
return frozenset(ref.provider_id for ref in settings.configured_chat_model_refs())
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def model_list_provider_ids_for_settings(settings: Settings) -> tuple[str, ...]:
|
| 29 |
+
"""Return providers worth discovering for this process configuration."""
|
| 30 |
+
referenced_ids = referenced_provider_ids(settings)
|
| 31 |
+
provider_ids: list[str] = []
|
| 32 |
+
for provider_id, descriptor in PROVIDER_CATALOG.items():
|
| 33 |
+
if descriptor.static_credential is not None:
|
| 34 |
+
if provider_id in referenced_ids:
|
| 35 |
+
provider_ids.append(provider_id)
|
| 36 |
+
continue
|
| 37 |
+
if (
|
| 38 |
+
descriptor.credential_env is not None
|
| 39 |
+
and provider_credential(descriptor, settings).strip()
|
| 40 |
+
):
|
| 41 |
+
provider_ids.append(provider_id)
|
| 42 |
+
return tuple(provider_ids)
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
class ProviderModelDiscovery:
|
| 46 |
+
"""Refresh provider model-list metadata for one provider runtime."""
|
| 47 |
+
|
| 48 |
+
def __init__(
|
| 49 |
+
self,
|
| 50 |
+
settings: Settings,
|
| 51 |
+
provider_resolver: ProviderResolver,
|
| 52 |
+
model_cache: ProviderModelCache,
|
| 53 |
+
) -> None:
|
| 54 |
+
self._settings = settings
|
| 55 |
+
self._provider_resolver = provider_resolver
|
| 56 |
+
self._model_cache = model_cache
|
| 57 |
+
self._refresh_task: asyncio.Task[None] | None = None
|
| 58 |
+
|
| 59 |
+
async def refresh_model_list_cache(self, *, only_missing: bool = False) -> None:
|
| 60 |
+
"""Best-effort refresh of model lists for usable providers."""
|
| 61 |
+
provider_ids = model_list_provider_ids_for_settings(self._settings)
|
| 62 |
+
if only_missing:
|
| 63 |
+
provider_ids = tuple(
|
| 64 |
+
provider_id
|
| 65 |
+
for provider_id in provider_ids
|
| 66 |
+
if not self._model_cache.has_provider(provider_id)
|
| 67 |
+
)
|
| 68 |
+
await self._refresh_model_infos(provider_ids)
|
| 69 |
+
|
| 70 |
+
def start_model_list_refresh(self) -> None:
|
| 71 |
+
"""Start a non-blocking cache warmup for missing eligible provider lists."""
|
| 72 |
+
if self._refresh_task is not None and not self._refresh_task.done():
|
| 73 |
+
return
|
| 74 |
+
|
| 75 |
+
provider_ids = tuple(
|
| 76 |
+
provider_id
|
| 77 |
+
for provider_id in model_list_provider_ids_for_settings(self._settings)
|
| 78 |
+
if not self._model_cache.has_provider(provider_id)
|
| 79 |
+
)
|
| 80 |
+
if not provider_ids:
|
| 81 |
+
logger.info(
|
| 82 |
+
"Provider model discovery cache already warm: providers={}",
|
| 83 |
+
len(self._model_cache.cached_model_ids()),
|
| 84 |
+
)
|
| 85 |
+
return
|
| 86 |
+
|
| 87 |
+
self._refresh_task = asyncio.create_task(self._run_refresh(provider_ids))
|
| 88 |
+
|
| 89 |
+
async def cleanup(self) -> None:
|
| 90 |
+
"""Cancel any background model-list refresh."""
|
| 91 |
+
if self._refresh_task is None or self._refresh_task.done():
|
| 92 |
+
return
|
| 93 |
+
self._refresh_task.cancel()
|
| 94 |
+
with suppress(asyncio.CancelledError):
|
| 95 |
+
await self._refresh_task
|
| 96 |
+
|
| 97 |
+
async def _run_refresh(self, provider_ids: tuple[str, ...]) -> None:
|
| 98 |
+
try:
|
| 99 |
+
await self._refresh_model_infos(provider_ids)
|
| 100 |
+
except asyncio.CancelledError:
|
| 101 |
+
raise
|
| 102 |
+
except Exception as exc:
|
| 103 |
+
logger.warning(
|
| 104 |
+
"Provider model discovery task failed: exc_type={}",
|
| 105 |
+
type(exc).__name__,
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
async def _refresh_model_infos(self, provider_ids: tuple[str, ...]) -> None:
|
| 109 |
+
tasks: dict[str, asyncio.Task[frozenset[ProviderModelInfo]]] = {}
|
| 110 |
+
for provider_id in provider_ids:
|
| 111 |
+
try:
|
| 112 |
+
provider = self._provider_resolver(provider_id)
|
| 113 |
+
except Exception as exc:
|
| 114 |
+
self._log_discovery_failure(provider_id, exc)
|
| 115 |
+
continue
|
| 116 |
+
tasks[provider_id] = asyncio.create_task(provider.list_model_infos())
|
| 117 |
+
|
| 118 |
+
if not tasks:
|
| 119 |
+
return
|
| 120 |
+
|
| 121 |
+
results = await asyncio.gather(*tasks.values(), return_exceptions=True)
|
| 122 |
+
for (provider_id, _task), result in zip(tasks.items(), results, strict=True):
|
| 123 |
+
if isinstance(result, BaseException):
|
| 124 |
+
if isinstance(result, asyncio.CancelledError):
|
| 125 |
+
raise result
|
| 126 |
+
self._log_discovery_failure(provider_id, result)
|
| 127 |
+
continue
|
| 128 |
+
self._model_cache.cache_model_infos(provider_id, result)
|
| 129 |
+
logger.info(
|
| 130 |
+
"Provider model discovery cached: provider={} models={}",
|
| 131 |
+
provider_id,
|
| 132 |
+
len(result),
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
def _log_discovery_failure(self, provider_id: str, exc: BaseException) -> None:
|
| 136 |
+
logger.warning(
|
| 137 |
+
"Provider model discovery skipped: provider={} reason={}",
|
| 138 |
+
provider_id,
|
| 139 |
+
provider_query_failure_reason(exc, self._settings),
|
| 140 |
+
)
|
|
@@ -0,0 +1,166 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Provider factory wiring and lazy adapter construction."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from collections.abc import Callable
|
| 6 |
+
|
| 7 |
+
from config.provider_catalog import (
|
| 8 |
+
PROVIDER_CATALOG,
|
| 9 |
+
SUPPORTED_PROVIDER_IDS,
|
| 10 |
+
)
|
| 11 |
+
from config.settings import Settings
|
| 12 |
+
from providers.base import BaseProvider, ProviderConfig
|
| 13 |
+
from providers.exceptions import UnknownProviderTypeError
|
| 14 |
+
|
| 15 |
+
from .config import build_provider_config
|
| 16 |
+
|
| 17 |
+
ProviderFactory = Callable[[ProviderConfig, Settings], BaseProvider]
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def _create_nvidia_nim(config: ProviderConfig, settings: Settings) -> BaseProvider:
|
| 21 |
+
from providers.nvidia_nim import NvidiaNimProvider
|
| 22 |
+
|
| 23 |
+
return NvidiaNimProvider(config, nim_settings=settings.nim)
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def _create_open_router(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 27 |
+
from providers.open_router import OpenRouterProvider
|
| 28 |
+
|
| 29 |
+
return OpenRouterProvider(config)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _create_mistral(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 33 |
+
from providers.mistral import MistralProvider
|
| 34 |
+
|
| 35 |
+
return MistralProvider(config)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def _create_mistral_codestral(
|
| 39 |
+
config: ProviderConfig, _settings: Settings
|
| 40 |
+
) -> BaseProvider:
|
| 41 |
+
from providers.codestral import CodestralProvider
|
| 42 |
+
|
| 43 |
+
return CodestralProvider(config)
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def _create_deepseek(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 47 |
+
from providers.deepseek import DeepSeekProvider
|
| 48 |
+
|
| 49 |
+
return DeepSeekProvider(config)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def _create_lmstudio(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 53 |
+
from providers.lmstudio import LMStudioProvider
|
| 54 |
+
|
| 55 |
+
return LMStudioProvider(config)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def _create_llamacpp(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 59 |
+
from providers.llamacpp import LlamaCppProvider
|
| 60 |
+
|
| 61 |
+
return LlamaCppProvider(config)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def _create_ollama(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 65 |
+
from providers.ollama import OllamaProvider
|
| 66 |
+
|
| 67 |
+
return OllamaProvider(config)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def _create_kimi(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 71 |
+
from providers.kimi import KimiProvider
|
| 72 |
+
|
| 73 |
+
return KimiProvider(config)
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def _create_wafer(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 77 |
+
from providers.wafer import WaferProvider
|
| 78 |
+
|
| 79 |
+
return WaferProvider(config)
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def _create_opencode(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 83 |
+
from providers.opencode import OpenCodeProvider
|
| 84 |
+
|
| 85 |
+
return OpenCodeProvider(config)
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def _create_opencode_go(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 89 |
+
from providers.opencode import OpenCodeProvider
|
| 90 |
+
|
| 91 |
+
return OpenCodeProvider(config, provider_name="OPENCODE_GO")
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
def _create_zai(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 95 |
+
from providers.zai import ZaiProvider
|
| 96 |
+
|
| 97 |
+
return ZaiProvider(config)
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def _create_fireworks(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 101 |
+
from providers.fireworks import FireworksProvider
|
| 102 |
+
|
| 103 |
+
return FireworksProvider(config)
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def _create_gemini(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 107 |
+
from providers.gemini import GeminiProvider
|
| 108 |
+
|
| 109 |
+
return GeminiProvider(config)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
def _create_groq(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 113 |
+
from providers.groq import GroqProvider
|
| 114 |
+
|
| 115 |
+
return GroqProvider(config)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def _create_cerebras(config: ProviderConfig, _settings: Settings) -> BaseProvider:
|
| 119 |
+
from providers.cerebras import CerebrasProvider
|
| 120 |
+
|
| 121 |
+
return CerebrasProvider(config)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
PROVIDER_FACTORIES: dict[str, ProviderFactory] = {
|
| 125 |
+
"nvidia_nim": _create_nvidia_nim,
|
| 126 |
+
"open_router": _create_open_router,
|
| 127 |
+
"gemini": _create_gemini,
|
| 128 |
+
"deepseek": _create_deepseek,
|
| 129 |
+
"mistral": _create_mistral,
|
| 130 |
+
"mistral_codestral": _create_mistral_codestral,
|
| 131 |
+
"opencode": _create_opencode,
|
| 132 |
+
"opencode_go": _create_opencode_go,
|
| 133 |
+
"wafer": _create_wafer,
|
| 134 |
+
"kimi": _create_kimi,
|
| 135 |
+
"cerebras": _create_cerebras,
|
| 136 |
+
"groq": _create_groq,
|
| 137 |
+
"fireworks": _create_fireworks,
|
| 138 |
+
"zai": _create_zai,
|
| 139 |
+
"lmstudio": _create_lmstudio,
|
| 140 |
+
"llamacpp": _create_llamacpp,
|
| 141 |
+
"ollama": _create_ollama,
|
| 142 |
+
}
|
| 143 |
+
|
| 144 |
+
if set(PROVIDER_CATALOG) != set(SUPPORTED_PROVIDER_IDS) or set(
|
| 145 |
+
PROVIDER_FACTORIES
|
| 146 |
+
) != set(SUPPORTED_PROVIDER_IDS):
|
| 147 |
+
raise AssertionError(
|
| 148 |
+
"PROVIDER_CATALOG, PROVIDER_FACTORIES, and SUPPORTED_PROVIDER_IDS are out of sync: "
|
| 149 |
+
f"catalog={set(PROVIDER_CATALOG)!r} factories={set(PROVIDER_FACTORIES)!r} "
|
| 150 |
+
f"ids={set(SUPPORTED_PROVIDER_IDS)!r}"
|
| 151 |
+
)
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
def create_provider(provider_id: str, settings: Settings) -> BaseProvider:
|
| 155 |
+
"""Create a provider instance for a supported provider id."""
|
| 156 |
+
descriptor = PROVIDER_CATALOG.get(provider_id)
|
| 157 |
+
if descriptor is None:
|
| 158 |
+
supported = "', '".join(PROVIDER_CATALOG)
|
| 159 |
+
raise UnknownProviderTypeError(
|
| 160 |
+
f"Unknown provider_type: '{provider_id}'. Supported: '{supported}'"
|
| 161 |
+
)
|
| 162 |
+
|
| 163 |
+
factory = PROVIDER_FACTORIES.get(provider_id)
|
| 164 |
+
if factory is None:
|
| 165 |
+
raise AssertionError(f"Unhandled provider descriptor: {provider_id}")
|
| 166 |
+
return factory(build_provider_config(descriptor, settings), settings)
|
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Provider model-list metadata cache."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from collections.abc import Iterable
|
| 6 |
+
|
| 7 |
+
from config.provider_catalog import SUPPORTED_PROVIDER_IDS
|
| 8 |
+
from providers.model_listing import ProviderModelInfo, model_infos_from_ids
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class ProviderModelCache:
|
| 12 |
+
"""Store provider model metadata for instant model-list responses."""
|
| 13 |
+
|
| 14 |
+
def __init__(self) -> None:
|
| 15 |
+
self._model_infos_by_provider: dict[str, dict[str, ProviderModelInfo]] = {}
|
| 16 |
+
|
| 17 |
+
def cache_model_ids(self, provider_id: str, model_ids: Iterable[str]) -> None:
|
| 18 |
+
"""Store raw provider model ids with unknown capability metadata."""
|
| 19 |
+
self.cache_model_infos(provider_id, model_infos_from_ids(model_ids))
|
| 20 |
+
|
| 21 |
+
def cache_model_infos(
|
| 22 |
+
self, provider_id: str, model_infos: Iterable[ProviderModelInfo]
|
| 23 |
+
) -> None:
|
| 24 |
+
"""Store provider model metadata by raw provider model id."""
|
| 25 |
+
clean_infos = {
|
| 26 |
+
info.model_id: info for info in model_infos if info.model_id.strip()
|
| 27 |
+
}
|
| 28 |
+
self._model_infos_by_provider[provider_id] = clean_infos
|
| 29 |
+
|
| 30 |
+
def cached_model_ids(self) -> dict[str, frozenset[str]]:
|
| 31 |
+
"""Return cached raw provider model ids by provider."""
|
| 32 |
+
return {
|
| 33 |
+
provider_id: frozenset(infos)
|
| 34 |
+
for provider_id, infos in self._model_infos_by_provider.items()
|
| 35 |
+
}
|
| 36 |
+
|
| 37 |
+
def has_provider(self, provider_id: str) -> bool:
|
| 38 |
+
"""Return whether this provider has any cached model-list result."""
|
| 39 |
+
return provider_id in self._model_infos_by_provider
|
| 40 |
+
|
| 41 |
+
def cached_model_supports_thinking(
|
| 42 |
+
self, provider_id: str, model_id: str
|
| 43 |
+
) -> bool | None:
|
| 44 |
+
"""Return cached thinking support when a provider exposes it."""
|
| 45 |
+
info = self._model_infos_by_provider.get(provider_id, {}).get(model_id)
|
| 46 |
+
if info is None:
|
| 47 |
+
return None
|
| 48 |
+
return info.supports_thinking
|
| 49 |
+
|
| 50 |
+
def cached_prefixed_model_refs(self) -> tuple[str, ...]:
|
| 51 |
+
"""Return cached provider models in user-selectable ``provider/model`` form."""
|
| 52 |
+
return tuple(info.model_id for info in self.cached_prefixed_model_infos())
|
| 53 |
+
|
| 54 |
+
def cached_prefixed_model_infos(self) -> tuple[ProviderModelInfo, ...]:
|
| 55 |
+
"""Return cached provider models with user-selectable prefixed ids."""
|
| 56 |
+
infos: list[ProviderModelInfo] = []
|
| 57 |
+
for provider_id in SUPPORTED_PROVIDER_IDS:
|
| 58 |
+
provider_infos = self._model_infos_by_provider.get(provider_id, {})
|
| 59 |
+
infos.extend(
|
| 60 |
+
ProviderModelInfo(
|
| 61 |
+
model_id=f"{provider_id}/{info.model_id}",
|
| 62 |
+
supports_thinking=info.supports_thinking,
|
| 63 |
+
)
|
| 64 |
+
for info in sorted(
|
| 65 |
+
provider_infos.values(), key=lambda item: item.model_id
|
| 66 |
+
)
|
| 67 |
+
)
|
| 68 |
+
return tuple(infos)
|
| 69 |
+
|
| 70 |
+
def clear(self) -> None:
|
| 71 |
+
"""Clear all cached model metadata."""
|
| 72 |
+
self._model_infos_by_provider.clear()
|
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""App-scoped provider runtime orchestration."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from collections.abc import Iterable, MutableMapping
|
| 6 |
+
|
| 7 |
+
from config.settings import Settings
|
| 8 |
+
from providers.base import BaseProvider
|
| 9 |
+
from providers.model_listing import ProviderModelInfo
|
| 10 |
+
|
| 11 |
+
from .cache import ProviderCache
|
| 12 |
+
from .discovery import ProviderModelDiscovery
|
| 13 |
+
from .model_cache import ProviderModelCache
|
| 14 |
+
from .validation import ConfiguredModelValidator
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class ProviderRuntime:
|
| 18 |
+
"""Own provider instances, model discovery, validation, and cleanup."""
|
| 19 |
+
|
| 20 |
+
def __init__(
|
| 21 |
+
self,
|
| 22 |
+
settings: Settings,
|
| 23 |
+
providers: MutableMapping[str, BaseProvider] | None = None,
|
| 24 |
+
) -> None:
|
| 25 |
+
self.settings = settings
|
| 26 |
+
self._provider_cache = ProviderCache(settings, providers)
|
| 27 |
+
self._model_cache = ProviderModelCache()
|
| 28 |
+
self._discovery = ProviderModelDiscovery(
|
| 29 |
+
settings,
|
| 30 |
+
self.resolve_provider,
|
| 31 |
+
self._model_cache,
|
| 32 |
+
)
|
| 33 |
+
self._validator = ConfiguredModelValidator(
|
| 34 |
+
settings,
|
| 35 |
+
self.resolve_provider,
|
| 36 |
+
self._model_cache,
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
def is_cached(self, provider_id: str) -> bool:
|
| 40 |
+
"""Return whether a provider for this id is already cached."""
|
| 41 |
+
return self._provider_cache.is_cached(provider_id)
|
| 42 |
+
|
| 43 |
+
def resolve_provider(self, provider_id: str) -> BaseProvider:
|
| 44 |
+
"""Return an existing provider or create it lazily."""
|
| 45 |
+
return self._provider_cache.get(provider_id)
|
| 46 |
+
|
| 47 |
+
def cache_model_ids(self, provider_id: str, model_ids: Iterable[str]) -> None:
|
| 48 |
+
"""Store raw provider model ids for later instant API responses."""
|
| 49 |
+
self._model_cache.cache_model_ids(provider_id, model_ids)
|
| 50 |
+
|
| 51 |
+
def cache_model_infos(
|
| 52 |
+
self, provider_id: str, model_infos: Iterable[ProviderModelInfo]
|
| 53 |
+
) -> None:
|
| 54 |
+
"""Store provider model metadata for later instant API responses."""
|
| 55 |
+
self._model_cache.cache_model_infos(provider_id, model_infos)
|
| 56 |
+
|
| 57 |
+
def cached_model_ids(self) -> dict[str, frozenset[str]]:
|
| 58 |
+
"""Return cached raw provider model ids by provider."""
|
| 59 |
+
return self._model_cache.cached_model_ids()
|
| 60 |
+
|
| 61 |
+
def cached_model_supports_thinking(
|
| 62 |
+
self, provider_id: str, model_id: str
|
| 63 |
+
) -> bool | None:
|
| 64 |
+
"""Return cached thinking support when a provider exposes it."""
|
| 65 |
+
return self._model_cache.cached_model_supports_thinking(provider_id, model_id)
|
| 66 |
+
|
| 67 |
+
def cached_prefixed_model_refs(self) -> tuple[str, ...]:
|
| 68 |
+
"""Return cached provider models in user-selectable ``provider/model`` form."""
|
| 69 |
+
return self._model_cache.cached_prefixed_model_refs()
|
| 70 |
+
|
| 71 |
+
def cached_prefixed_model_infos(self) -> tuple[ProviderModelInfo, ...]:
|
| 72 |
+
"""Return cached provider models with user-selectable prefixed ids."""
|
| 73 |
+
return self._model_cache.cached_prefixed_model_infos()
|
| 74 |
+
|
| 75 |
+
async def refresh_model_list_cache(self, *, only_missing: bool = False) -> None:
|
| 76 |
+
"""Best-effort refresh of model lists for usable providers."""
|
| 77 |
+
await self._discovery.refresh_model_list_cache(only_missing=only_missing)
|
| 78 |
+
|
| 79 |
+
def start_model_list_refresh(self) -> None:
|
| 80 |
+
"""Start a non-blocking cache warmup for missing eligible provider lists."""
|
| 81 |
+
self._discovery.start_model_list_refresh()
|
| 82 |
+
|
| 83 |
+
async def validate_configured_models(self) -> None:
|
| 84 |
+
"""Fail unless every configured chat model exists upstream."""
|
| 85 |
+
await self._validator.validate_configured_models()
|
| 86 |
+
|
| 87 |
+
async def cleanup(self) -> None:
|
| 88 |
+
"""Cancel discovery, clean provider instances, and clear model metadata."""
|
| 89 |
+
try:
|
| 90 |
+
await self._discovery.cleanup()
|
| 91 |
+
await self._provider_cache.cleanup()
|
| 92 |
+
finally:
|
| 93 |
+
self._model_cache.clear()
|
|
@@ -0,0 +1,123 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configured provider model validation."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import asyncio
|
| 6 |
+
from collections import defaultdict
|
| 7 |
+
from collections.abc import Callable
|
| 8 |
+
|
| 9 |
+
import httpx
|
| 10 |
+
from loguru import logger
|
| 11 |
+
|
| 12 |
+
from config.settings import ConfiguredChatModelRef, Settings
|
| 13 |
+
from providers.base import BaseProvider
|
| 14 |
+
from providers.exceptions import (
|
| 15 |
+
AuthenticationError,
|
| 16 |
+
ModelListResponseError,
|
| 17 |
+
ProviderError,
|
| 18 |
+
ServiceUnavailableError,
|
| 19 |
+
)
|
| 20 |
+
from providers.model_listing import ProviderModelInfo
|
| 21 |
+
|
| 22 |
+
from .model_cache import ProviderModelCache
|
| 23 |
+
|
| 24 |
+
ProviderResolver = Callable[[str], BaseProvider]
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def provider_query_failure_reason(exc: BaseException, settings: Settings) -> str:
|
| 28 |
+
"""Return a concise model-list query failure reason for user-facing logs."""
|
| 29 |
+
if isinstance(exc, ModelListResponseError):
|
| 30 |
+
return f"malformed model-list response: {exc.message}"
|
| 31 |
+
if isinstance(exc, httpx.HTTPStatusError):
|
| 32 |
+
return f"query failure: HTTP {exc.response.status_code}"
|
| 33 |
+
if isinstance(exc, AuthenticationError):
|
| 34 |
+
return f"query failure: {exc.message}"
|
| 35 |
+
if isinstance(exc, ProviderError) and settings.log_api_error_tracebacks:
|
| 36 |
+
return f"query failure: {exc.message}"
|
| 37 |
+
return f"query failure: {type(exc).__name__}"
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class ConfiguredModelValidator:
|
| 41 |
+
"""Validate configured provider/model refs against upstream model lists."""
|
| 42 |
+
|
| 43 |
+
def __init__(
|
| 44 |
+
self,
|
| 45 |
+
settings: Settings,
|
| 46 |
+
provider_resolver: ProviderResolver,
|
| 47 |
+
model_cache: ProviderModelCache,
|
| 48 |
+
) -> None:
|
| 49 |
+
self._settings = settings
|
| 50 |
+
self._provider_resolver = provider_resolver
|
| 51 |
+
self._model_cache = model_cache
|
| 52 |
+
|
| 53 |
+
async def validate_configured_models(self) -> None:
|
| 54 |
+
"""Fail unless every configured chat model exists upstream."""
|
| 55 |
+
refs = self._settings.configured_chat_model_refs()
|
| 56 |
+
refs_by_provider: dict[str, list[ConfiguredChatModelRef]] = defaultdict(list)
|
| 57 |
+
for ref in refs:
|
| 58 |
+
refs_by_provider[ref.provider_id].append(ref)
|
| 59 |
+
|
| 60 |
+
failures: list[str] = []
|
| 61 |
+
tasks: dict[str, asyncio.Task[frozenset[ProviderModelInfo]]] = {}
|
| 62 |
+
for provider_id, provider_refs in refs_by_provider.items():
|
| 63 |
+
try:
|
| 64 |
+
provider = self._provider_resolver(provider_id)
|
| 65 |
+
except Exception as exc:
|
| 66 |
+
failures.extend(
|
| 67 |
+
self._format_provider_query_failures(provider_refs, exc)
|
| 68 |
+
)
|
| 69 |
+
continue
|
| 70 |
+
tasks[provider_id] = asyncio.create_task(provider.list_model_infos())
|
| 71 |
+
|
| 72 |
+
if tasks:
|
| 73 |
+
results = await asyncio.gather(*tasks.values(), return_exceptions=True)
|
| 74 |
+
for (provider_id, _task), result in zip(
|
| 75 |
+
tasks.items(), results, strict=True
|
| 76 |
+
):
|
| 77 |
+
provider_refs = refs_by_provider[provider_id]
|
| 78 |
+
if isinstance(result, BaseException):
|
| 79 |
+
if isinstance(result, asyncio.CancelledError):
|
| 80 |
+
raise result
|
| 81 |
+
failures.extend(
|
| 82 |
+
self._format_provider_query_failures(provider_refs, result)
|
| 83 |
+
)
|
| 84 |
+
continue
|
| 85 |
+
self._model_cache.cache_model_infos(provider_id, result)
|
| 86 |
+
model_ids = self._model_cache.cached_model_ids()[provider_id]
|
| 87 |
+
failures.extend(
|
| 88 |
+
self._format_missing_model_failure(ref)
|
| 89 |
+
for ref in provider_refs
|
| 90 |
+
if ref.model_id not in model_ids
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
if failures:
|
| 94 |
+
message = "Configured model validation failed:\n" + "\n".join(
|
| 95 |
+
f"- {failure}" for failure in failures
|
| 96 |
+
)
|
| 97 |
+
raise ServiceUnavailableError(message)
|
| 98 |
+
|
| 99 |
+
logger.info(
|
| 100 |
+
"Configured provider models validated: models={} providers={}",
|
| 101 |
+
len(refs),
|
| 102 |
+
len(refs_by_provider),
|
| 103 |
+
)
|
| 104 |
+
|
| 105 |
+
def _format_provider_query_failures(
|
| 106 |
+
self,
|
| 107 |
+
refs: list[ConfiguredChatModelRef],
|
| 108 |
+
exc: BaseException,
|
| 109 |
+
) -> list[str]:
|
| 110 |
+
reason = provider_query_failure_reason(exc, self._settings)
|
| 111 |
+
return [self._format_model_validation_failure(ref, reason) for ref in refs]
|
| 112 |
+
|
| 113 |
+
def _format_missing_model_failure(self, ref: ConfiguredChatModelRef) -> str:
|
| 114 |
+
return self._format_model_validation_failure(ref, "missing model")
|
| 115 |
+
|
| 116 |
+
@staticmethod
|
| 117 |
+
def _format_model_validation_failure(
|
| 118 |
+
ref: ConfiguredChatModelRef, problem: str
|
| 119 |
+
) -> str:
|
| 120 |
+
return (
|
| 121 |
+
f"sources={','.join(ref.sources)} provider={ref.provider_id} "
|
| 122 |
+
f"model={ref.model_id} problem={problem}"
|
| 123 |
+
)
|
|
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
|
| 4 |
|
| 5 |
[project]
|
| 6 |
name = "free-claude-code"
|
| 7 |
-
version = "2.3.
|
| 8 |
description = "Middleware between Claude Code CLI (Anthropic API) and NVIDIA NIM"
|
| 9 |
readme = "README.md"
|
| 10 |
requires-python = ">=3.14.0"
|
|
|
|
| 4 |
|
| 5 |
[project]
|
| 6 |
name = "free-claude-code"
|
| 7 |
+
version = "2.3.17"
|
| 8 |
description = "Middleware between Claude Code CLI (Anthropic API) and NVIDIA NIM"
|
| 9 |
readme = "README.md"
|
| 10 |
requires-python = ">=3.14.0"
|
|
@@ -58,7 +58,7 @@ Default targets do not send real bot messages or load voice backends:
|
|
| 58 |
| `cli` | `fcc-init`, server entrypoint, Claude CLI adaptive thinking, session cleanup | Claude CLI binary and provider only for real CLI |
|
| 59 |
| `clients` | VS Code and JetBrains protocol payloads | configured provider |
|
| 60 |
| `config` | env precedence, removed-env migration, proxy/timeouts | none |
|
| 61 |
-
| `extensibility` | provider
|
| 62 |
| `messaging` | fake Discord/Telegram full flow, commands, trees, persistence, voice cancel | none |
|
| 63 |
| `providers` | multi-turn text, adaptive thinking history, tools, disconnect, errors | configured providers, optional `FCC_SMOKE_MODEL_*` |
|
| 64 |
| `tools` | forced tool_use and tool_result continuation | tool-capable configured provider |
|
|
|
|
| 58 |
| `cli` | `fcc-init`, server entrypoint, Claude CLI adaptive thinking, session cleanup | Claude CLI binary and provider only for real CLI |
|
| 59 |
| `clients` | VS Code and JetBrains protocol payloads | configured provider |
|
| 60 |
| `config` | env precedence, removed-env migration, proxy/timeouts | none |
|
| 61 |
+
| `extensibility` | provider runtime and platform factory construction | none |
|
| 62 |
| `messaging` | fake Discord/Telegram full flow, commands, trees, persistence, voice cancel | none |
|
| 63 |
| `providers` | multi-turn text, adaptive thinking history, tools, disconnect, errors | configured providers, optional `FCC_SMOKE_MODEL_*` |
|
| 64 |
| `tools` | forced tool_use and tool_result continuation | tool-capable configured provider |
|
|
@@ -112,13 +112,13 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = (
|
|
| 112 |
),
|
| 113 |
CapabilityContract(
|
| 114 |
"provider_routing",
|
| 115 |
-
"
|
| 116 |
"provider_matrix",
|
| 117 |
-
"providers.
|
| 118 |
"provider id and Settings",
|
| 119 |
"configured BaseProvider instance",
|
| 120 |
"503 for missing credentials; invalid_request_error for unknown provider",
|
| 121 |
-
("tests/api/test_dependencies.py", "tests/providers/
|
| 122 |
(
|
| 123 |
"test_configured_provider_models_stream_successfully",
|
| 124 |
"test_provider_matrix_presence_e2e",
|
|
@@ -150,17 +150,17 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = (
|
|
| 150 |
"provider_routing",
|
| 151 |
"provider_runtime_config",
|
| 152 |
"provider_proxy_timeout_config",
|
| 153 |
-
"providers.
|
| 154 |
"provider proxy, timeout, and rate-limit settings",
|
| 155 |
"provider client and scoped limiter config",
|
| 156 |
"provider construction failure",
|
| 157 |
-
("tests/api/test_dependencies.py", "tests/providers/
|
| 158 |
),
|
| 159 |
CapabilityContract(
|
| 160 |
"provider_routing",
|
| 161 |
"zero_cost_backends",
|
| 162 |
"zero_cost_provider_access",
|
| 163 |
-
"providers.
|
| 164 |
"configured free/local provider",
|
| 165 |
"streaming response from selected backend",
|
| 166 |
"missing env or upstream unavailable skip in smoke",
|
|
@@ -454,13 +454,13 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = (
|
|
| 454 |
"extensibility",
|
| 455 |
"provider_platform_abcs",
|
| 456 |
"extensible_provider_platform_abcs",
|
| 457 |
-
"providers.
|
| 458 |
"new provider/platform implementations",
|
| 459 |
"registered BaseProvider or messaging component bundle",
|
| 460 |
"unknown platform returns None; unknown provider errors",
|
| 461 |
(
|
| 462 |
"tests/contracts/test_feature_manifest.py",
|
| 463 |
-
"tests/providers/
|
| 464 |
),
|
| 465 |
),
|
| 466 |
)
|
|
|
|
| 112 |
),
|
| 113 |
CapabilityContract(
|
| 114 |
"provider_routing",
|
| 115 |
+
"provider_runtime",
|
| 116 |
"provider_matrix",
|
| 117 |
+
"providers.runtime.ProviderRuntime",
|
| 118 |
"provider id and Settings",
|
| 119 |
"configured BaseProvider instance",
|
| 120 |
"503 for missing credentials; invalid_request_error for unknown provider",
|
| 121 |
+
("tests/api/test_dependencies.py", "tests/providers/test_provider_runtime.py"),
|
| 122 |
(
|
| 123 |
"test_configured_provider_models_stream_successfully",
|
| 124 |
"test_provider_matrix_presence_e2e",
|
|
|
|
| 150 |
"provider_routing",
|
| 151 |
"provider_runtime_config",
|
| 152 |
"provider_proxy_timeout_config",
|
| 153 |
+
"providers.runtime.ProviderRuntime",
|
| 154 |
"provider proxy, timeout, and rate-limit settings",
|
| 155 |
"provider client and scoped limiter config",
|
| 156 |
"provider construction failure",
|
| 157 |
+
("tests/api/test_dependencies.py", "tests/providers/test_provider_runtime.py"),
|
| 158 |
),
|
| 159 |
CapabilityContract(
|
| 160 |
"provider_routing",
|
| 161 |
"zero_cost_backends",
|
| 162 |
"zero_cost_provider_access",
|
| 163 |
+
"providers.runtime.ProviderRuntime",
|
| 164 |
"configured free/local provider",
|
| 165 |
"streaming response from selected backend",
|
| 166 |
"missing env or upstream unavailable skip in smoke",
|
|
|
|
| 454 |
"extensibility",
|
| 455 |
"provider_platform_abcs",
|
| 456 |
"extensible_provider_platform_abcs",
|
| 457 |
+
"providers.runtime and messaging.platforms.factory",
|
| 458 |
"new provider/platform implementations",
|
| 459 |
"registered BaseProvider or messaging component bundle",
|
| 460 |
"unknown platform returns None; unknown provider errors",
|
| 461 |
(
|
| 462 |
"tests/contracts/test_feature_manifest.py",
|
| 463 |
+
"tests/providers/test_provider_runtime.py",
|
| 464 |
),
|
| 465 |
),
|
| 466 |
)
|
|
@@ -235,10 +235,10 @@ FEATURE_INVENTORY: tuple[FeatureCoverage, ...] = (
|
|
| 235 |
"readme",
|
| 236 |
(
|
| 237 |
"tests/contracts/test_feature_manifest.py",
|
| 238 |
-
"tests/providers/
|
| 239 |
),
|
| 240 |
(),
|
| 241 |
-
("
|
| 242 |
("extensibility",),
|
| 243 |
(),
|
| 244 |
"always runnable with isolated settings",
|
|
@@ -341,7 +341,7 @@ FEATURE_INVENTORY: tuple[FeatureCoverage, ...] = (
|
|
| 341 |
"provider_proxy_timeout_config",
|
| 342 |
"Provider proxies and HTTP timeout settings reach provider config",
|
| 343 |
"public_surface",
|
| 344 |
-
("tests/api/test_dependencies.py", "tests/providers/
|
| 345 |
(),
|
| 346 |
("test_proxy_timeout_config_e2e",),
|
| 347 |
("config",),
|
|
|
|
| 235 |
"readme",
|
| 236 |
(
|
| 237 |
"tests/contracts/test_feature_manifest.py",
|
| 238 |
+
"tests/providers/test_provider_runtime.py",
|
| 239 |
),
|
| 240 |
(),
|
| 241 |
+
("test_provider_runtime_config_e2e", "test_platform_factory_e2e"),
|
| 242 |
("extensibility",),
|
| 243 |
(),
|
| 244 |
"always runnable with isolated settings",
|
|
|
|
| 341 |
"provider_proxy_timeout_config",
|
| 342 |
"Provider proxies and HTTP timeout settings reach provider config",
|
| 343 |
"public_surface",
|
| 344 |
+
("tests/api/test_dependencies.py", "tests/providers/test_provider_runtime.py"),
|
| 345 |
(),
|
| 346 |
("test_proxy_timeout_config_e2e",),
|
| 347 |
("config",),
|
|
@@ -5,9 +5,10 @@ import subprocess
|
|
| 5 |
|
| 6 |
import pytest
|
| 7 |
|
|
|
|
| 8 |
from config.settings import Settings
|
| 9 |
from messaging.platforms.factory import create_messaging_components
|
| 10 |
-
from providers.
|
| 11 |
from smoke.lib.child_process import cmd_free_claude_code_serve, cmd_python_c
|
| 12 |
from smoke.lib.config import SmokeConfig
|
| 13 |
from smoke.lib.e2e import SmokeServerDriver
|
|
@@ -112,8 +113,9 @@ def test_proxy_timeout_config_e2e(smoke_config: SmokeConfig, tmp_path) -> None:
|
|
| 112 |
env["FCC_ENV_FILE"] = str(env_file)
|
| 113 |
script = (
|
| 114 |
"from config.settings import Settings; "
|
| 115 |
-
"from
|
| 116 |
-
"
|
|
|
|
| 117 |
"print(c.proxy); print(c.http_read_timeout); "
|
| 118 |
"print(c.http_connect_timeout); print(c.http_write_timeout)"
|
| 119 |
)
|
|
@@ -136,9 +138,9 @@ def test_proxy_timeout_config_e2e(smoke_config: SmokeConfig, tmp_path) -> None:
|
|
| 136 |
|
| 137 |
|
| 138 |
@pytest.mark.smoke_target("extensibility")
|
| 139 |
-
def
|
| 140 |
settings_kwargs: dict[str, str] = {}
|
| 141 |
-
for descriptor in
|
| 142 |
if descriptor.credential_attr is not None:
|
| 143 |
settings_kwargs[_settings_init_key(descriptor.credential_attr)] = (
|
| 144 |
f"{descriptor.provider_id}-key"
|
|
@@ -148,7 +150,7 @@ def test_provider_registry_e2e() -> None:
|
|
| 148 |
descriptor.default_base_url
|
| 149 |
)
|
| 150 |
settings = Settings.model_validate(settings_kwargs)
|
| 151 |
-
for descriptor in
|
| 152 |
config = build_provider_config(descriptor, settings)
|
| 153 |
assert config.base_url
|
| 154 |
assert config.api_key
|
|
|
|
| 5 |
|
| 6 |
import pytest
|
| 7 |
|
| 8 |
+
from config.provider_catalog import PROVIDER_CATALOG
|
| 9 |
from config.settings import Settings
|
| 10 |
from messaging.platforms.factory import create_messaging_components
|
| 11 |
+
from providers.runtime import build_provider_config
|
| 12 |
from smoke.lib.child_process import cmd_free_claude_code_serve, cmd_python_c
|
| 13 |
from smoke.lib.config import SmokeConfig
|
| 14 |
from smoke.lib.e2e import SmokeServerDriver
|
|
|
|
| 113 |
env["FCC_ENV_FILE"] = str(env_file)
|
| 114 |
script = (
|
| 115 |
"from config.settings import Settings; "
|
| 116 |
+
"from config.provider_catalog import PROVIDER_CATALOG; "
|
| 117 |
+
"from providers.runtime import build_provider_config; "
|
| 118 |
+
"s=Settings(); c=build_provider_config(PROVIDER_CATALOG['open_router'], s); "
|
| 119 |
"print(c.proxy); print(c.http_read_timeout); "
|
| 120 |
"print(c.http_connect_timeout); print(c.http_write_timeout)"
|
| 121 |
)
|
|
|
|
| 138 |
|
| 139 |
|
| 140 |
@pytest.mark.smoke_target("extensibility")
|
| 141 |
+
def test_provider_runtime_config_e2e() -> None:
|
| 142 |
settings_kwargs: dict[str, str] = {}
|
| 143 |
+
for descriptor in PROVIDER_CATALOG.values():
|
| 144 |
if descriptor.credential_attr is not None:
|
| 145 |
settings_kwargs[_settings_init_key(descriptor.credential_attr)] = (
|
| 146 |
f"{descriptor.provider_id}-key"
|
|
|
|
| 150 |
descriptor.default_base_url
|
| 151 |
)
|
| 152 |
settings = Settings.model_validate(settings_kwargs)
|
| 153 |
+
for descriptor in PROVIDER_CATALOG.values():
|
| 154 |
config = build_provider_config(descriptor, settings)
|
| 155 |
assert config.base_url
|
| 156 |
assert config.api_key
|
|
@@ -31,10 +31,10 @@ def client():
|
|
| 31 |
with (
|
| 32 |
patch("api.dependencies.resolve_provider", return_value=mock_provider),
|
| 33 |
patch(
|
| 34 |
-
"providers.
|
| 35 |
new_callable=AsyncMock,
|
| 36 |
),
|
| 37 |
-
patch("providers.
|
| 38 |
TestClient(app) as test_client,
|
| 39 |
):
|
| 40 |
yield test_client
|
|
|
|
| 31 |
with (
|
| 32 |
patch("api.dependencies.resolve_provider", return_value=mock_provider),
|
| 33 |
patch(
|
| 34 |
+
"providers.runtime.ProviderRuntime.validate_configured_models",
|
| 35 |
new_callable=AsyncMock,
|
| 36 |
),
|
| 37 |
+
patch("providers.runtime.ProviderRuntime.start_model_list_refresh"),
|
| 38 |
TestClient(app) as test_client,
|
| 39 |
):
|
| 40 |
yield test_client
|
|
@@ -10,7 +10,7 @@ from fastapi.testclient import TestClient
|
|
| 10 |
|
| 11 |
from config.settings import Settings
|
| 12 |
from providers.exceptions import ServiceUnavailableError
|
| 13 |
-
from providers.
|
| 14 |
|
| 15 |
_RUNTIME_EXTRAS = {
|
| 16 |
"voice_note_enabled": True,
|
|
@@ -110,9 +110,9 @@ async def test_runtime_startup_logs_admin_url_without_printed_server_banner(tmp_
|
|
| 110 |
api_runtime_mod.logging, "getLogger", return_value=uvicorn_logger
|
| 111 |
) as get_logger,
|
| 112 |
patch.object(api_runtime_mod.logger, "info") as app_info,
|
| 113 |
-
patch.object(
|
| 114 |
-
patch.object(
|
| 115 |
-
patch.object(
|
| 116 |
patch(
|
| 117 |
"messaging.platforms.factory.create_messaging_components",
|
| 118 |
return_value=None,
|
|
@@ -155,7 +155,7 @@ def test_create_app_provider_error_handler_returns_anthropic_format():
|
|
| 155 |
)
|
| 156 |
with (
|
| 157 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 158 |
-
patch.object(
|
| 159 |
):
|
| 160 |
with TestClient(app) as client:
|
| 161 |
resp = client.get("/raise_provider")
|
|
@@ -192,7 +192,7 @@ def test_create_app_provider_error_default_logs_exclude_provider_message():
|
|
| 192 |
)
|
| 193 |
with (
|
| 194 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 195 |
-
patch.object(
|
| 196 |
patch.object(api_app_mod.logger, "error") as log_err,
|
| 197 |
):
|
| 198 |
with TestClient(app) as client:
|
|
@@ -228,7 +228,7 @@ def test_create_app_general_exception_handler_returns_500():
|
|
| 228 |
)
|
| 229 |
with (
|
| 230 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 231 |
-
patch.object(
|
| 232 |
):
|
| 233 |
with TestClient(app, raise_server_exceptions=False) as client:
|
| 234 |
resp = client.get("/raise_general")
|
|
@@ -265,7 +265,7 @@ def test_create_app_general_exception_default_logs_exclude_exception_message():
|
|
| 265 |
)
|
| 266 |
with (
|
| 267 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 268 |
-
patch.object(
|
| 269 |
patch.object(api_app_mod.logger, "error") as log_err,
|
| 270 |
):
|
| 271 |
with TestClient(app, raise_server_exceptions=False) as client:
|
|
@@ -325,10 +325,10 @@ def test_app_lifespan_sets_state_and_cleans_up(tmp_path, messaging_enabled):
|
|
| 325 |
|
| 326 |
api_app_mod = importlib.import_module("api.app")
|
| 327 |
|
| 328 |
-
|
| 329 |
with (
|
| 330 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 331 |
-
patch.object(
|
| 332 |
patch(
|
| 333 |
"messaging.platforms.factory.create_messaging_components",
|
| 334 |
return_value=fake_components if messaging_enabled else None,
|
|
@@ -360,7 +360,7 @@ def test_app_lifespan_sets_state_and_cleans_up(tmp_path, messaging_enabled):
|
|
| 360 |
cli_manager.stop_all.assert_not_awaited()
|
| 361 |
assert getattr(app.state, "messaging_runtime", "missing") is None
|
| 362 |
|
| 363 |
-
|
| 364 |
|
| 365 |
|
| 366 |
def test_app_lifespan_cleanup_continues_if_platform_stop_raises(tmp_path):
|
|
@@ -396,10 +396,10 @@ def test_app_lifespan_cleanup_continues_if_platform_stop_raises(tmp_path):
|
|
| 396 |
cli_manager.stop_all = AsyncMock()
|
| 397 |
|
| 398 |
api_app_mod = importlib.import_module("api.app")
|
| 399 |
-
|
| 400 |
with (
|
| 401 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 402 |
-
patch.object(
|
| 403 |
patch(
|
| 404 |
"messaging.platforms.factory.create_messaging_components",
|
| 405 |
return_value=fake_components,
|
|
@@ -412,7 +412,7 @@ def test_app_lifespan_cleanup_continues_if_platform_stop_raises(tmp_path):
|
|
| 412 |
|
| 413 |
fake_platform.stop.assert_awaited_once()
|
| 414 |
cli_manager.stop_all.assert_awaited_once()
|
| 415 |
-
|
| 416 |
|
| 417 |
|
| 418 |
@pytest.mark.asyncio
|
|
@@ -439,8 +439,8 @@ async def test_runtime_startup_validation_failure_does_not_block_server(tmp_path
|
|
| 439 |
validation = AsyncMock(side_effect=ServiceUnavailableError("bad model"))
|
| 440 |
cleanup = AsyncMock()
|
| 441 |
with (
|
| 442 |
-
patch.object(
|
| 443 |
-
patch.object(
|
| 444 |
patch.object(api_runtime_mod.logger, "warning") as log_warning,
|
| 445 |
patch(
|
| 446 |
"messaging.platforms.factory.create_messaging_components",
|
|
@@ -450,7 +450,7 @@ async def test_runtime_startup_validation_failure_does_not_block_server(tmp_path
|
|
| 450 |
await runtime.startup()
|
| 451 |
await runtime.shutdown()
|
| 452 |
|
| 453 |
-
validation.assert_awaited_once_with(
|
| 454 |
cleanup.assert_awaited_once()
|
| 455 |
create_components.assert_called_once()
|
| 456 |
logged = " ".join(
|
|
@@ -494,8 +494,8 @@ async def test_graceful_asgi_lifespan_model_validation_failure_starts(tmp_path):
|
|
| 494 |
cleanup = AsyncMock()
|
| 495 |
with (
|
| 496 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 497 |
-
patch.object(
|
| 498 |
-
patch.object(
|
| 499 |
):
|
| 500 |
await app({"type": "lifespan"}, receive, send)
|
| 501 |
|
|
@@ -524,10 +524,10 @@ def test_app_lifespan_messaging_import_error_no_crash(tmp_path, caplog):
|
|
| 524 |
)
|
| 525 |
|
| 526 |
api_app_mod = importlib.import_module("api.app")
|
| 527 |
-
|
| 528 |
with (
|
| 529 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 530 |
-
patch.object(
|
| 531 |
patch(
|
| 532 |
"messaging.platforms.factory.create_messaging_components",
|
| 533 |
side_effect=ImportError("discord not installed"),
|
|
@@ -537,7 +537,7 @@ def test_app_lifespan_messaging_import_error_no_crash(tmp_path, caplog):
|
|
| 537 |
pass
|
| 538 |
|
| 539 |
assert getattr(app.state, "messaging_runtime", None) is None
|
| 540 |
-
|
| 541 |
|
| 542 |
|
| 543 |
def test_app_lifespan_platform_start_exception_cleanup_still_runs(tmp_path):
|
|
@@ -574,10 +574,10 @@ def test_app_lifespan_platform_start_exception_cleanup_still_runs(tmp_path):
|
|
| 574 |
cli_manager.stop_all = AsyncMock()
|
| 575 |
|
| 576 |
api_app_mod = importlib.import_module("api.app")
|
| 577 |
-
|
| 578 |
with (
|
| 579 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 580 |
-
patch.object(
|
| 581 |
patch(
|
| 582 |
"messaging.platforms.factory.create_messaging_components",
|
| 583 |
return_value=fake_components,
|
|
@@ -588,7 +588,7 @@ def test_app_lifespan_platform_start_exception_cleanup_still_runs(tmp_path):
|
|
| 588 |
):
|
| 589 |
pass
|
| 590 |
|
| 591 |
-
|
| 592 |
|
| 593 |
|
| 594 |
def test_app_lifespan_flush_pending_save_exception_warning_only(tmp_path):
|
|
@@ -626,10 +626,10 @@ def test_app_lifespan_flush_pending_save_exception_warning_only(tmp_path):
|
|
| 626 |
cli_manager.stop_all = AsyncMock()
|
| 627 |
|
| 628 |
api_app_mod = importlib.import_module("api.app")
|
| 629 |
-
|
| 630 |
with (
|
| 631 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 632 |
-
patch.object(
|
| 633 |
patch(
|
| 634 |
"messaging.platforms.factory.create_messaging_components",
|
| 635 |
return_value=fake_components,
|
|
@@ -641,7 +641,7 @@ def test_app_lifespan_flush_pending_save_exception_warning_only(tmp_path):
|
|
| 641 |
pass
|
| 642 |
|
| 643 |
session_store.flush_pending_save.assert_called_once()
|
| 644 |
-
|
| 645 |
|
| 646 |
|
| 647 |
def test_create_app_writes_server_log_under_fcc_home(monkeypatch, tmp_path):
|
|
|
|
| 10 |
|
| 11 |
from config.settings import Settings
|
| 12 |
from providers.exceptions import ServiceUnavailableError
|
| 13 |
+
from providers.runtime import ProviderRuntime
|
| 14 |
|
| 15 |
_RUNTIME_EXTRAS = {
|
| 16 |
"voice_note_enabled": True,
|
|
|
|
| 110 |
api_runtime_mod.logging, "getLogger", return_value=uvicorn_logger
|
| 111 |
) as get_logger,
|
| 112 |
patch.object(api_runtime_mod.logger, "info") as app_info,
|
| 113 |
+
patch.object(ProviderRuntime, "validate_configured_models", new=AsyncMock()),
|
| 114 |
+
patch.object(ProviderRuntime, "start_model_list_refresh"),
|
| 115 |
+
patch.object(ProviderRuntime, "cleanup", new=AsyncMock()),
|
| 116 |
patch(
|
| 117 |
"messaging.platforms.factory.create_messaging_components",
|
| 118 |
return_value=None,
|
|
|
|
| 155 |
)
|
| 156 |
with (
|
| 157 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 158 |
+
patch.object(ProviderRuntime, "cleanup", new=AsyncMock()),
|
| 159 |
):
|
| 160 |
with TestClient(app) as client:
|
| 161 |
resp = client.get("/raise_provider")
|
|
|
|
| 192 |
)
|
| 193 |
with (
|
| 194 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 195 |
+
patch.object(ProviderRuntime, "cleanup", new=AsyncMock()),
|
| 196 |
patch.object(api_app_mod.logger, "error") as log_err,
|
| 197 |
):
|
| 198 |
with TestClient(app) as client:
|
|
|
|
| 228 |
)
|
| 229 |
with (
|
| 230 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 231 |
+
patch.object(ProviderRuntime, "cleanup", new=AsyncMock()),
|
| 232 |
):
|
| 233 |
with TestClient(app, raise_server_exceptions=False) as client:
|
| 234 |
resp = client.get("/raise_general")
|
|
|
|
| 265 |
)
|
| 266 |
with (
|
| 267 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 268 |
+
patch.object(ProviderRuntime, "cleanup", new=AsyncMock()),
|
| 269 |
patch.object(api_app_mod.logger, "error") as log_err,
|
| 270 |
):
|
| 271 |
with TestClient(app, raise_server_exceptions=False) as client:
|
|
|
|
| 325 |
|
| 326 |
api_app_mod = importlib.import_module("api.app")
|
| 327 |
|
| 328 |
+
runtime_cleanup = AsyncMock()
|
| 329 |
with (
|
| 330 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 331 |
+
patch.object(ProviderRuntime, "cleanup", new=runtime_cleanup),
|
| 332 |
patch(
|
| 333 |
"messaging.platforms.factory.create_messaging_components",
|
| 334 |
return_value=fake_components if messaging_enabled else None,
|
|
|
|
| 360 |
cli_manager.stop_all.assert_not_awaited()
|
| 361 |
assert getattr(app.state, "messaging_runtime", "missing") is None
|
| 362 |
|
| 363 |
+
runtime_cleanup.assert_awaited_once()
|
| 364 |
|
| 365 |
|
| 366 |
def test_app_lifespan_cleanup_continues_if_platform_stop_raises(tmp_path):
|
|
|
|
| 396 |
cli_manager.stop_all = AsyncMock()
|
| 397 |
|
| 398 |
api_app_mod = importlib.import_module("api.app")
|
| 399 |
+
runtime_cleanup = AsyncMock()
|
| 400 |
with (
|
| 401 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 402 |
+
patch.object(ProviderRuntime, "cleanup", new=runtime_cleanup),
|
| 403 |
patch(
|
| 404 |
"messaging.platforms.factory.create_messaging_components",
|
| 405 |
return_value=fake_components,
|
|
|
|
| 412 |
|
| 413 |
fake_platform.stop.assert_awaited_once()
|
| 414 |
cli_manager.stop_all.assert_awaited_once()
|
| 415 |
+
runtime_cleanup.assert_awaited_once()
|
| 416 |
|
| 417 |
|
| 418 |
@pytest.mark.asyncio
|
|
|
|
| 439 |
validation = AsyncMock(side_effect=ServiceUnavailableError("bad model"))
|
| 440 |
cleanup = AsyncMock()
|
| 441 |
with (
|
| 442 |
+
patch.object(ProviderRuntime, "validate_configured_models", new=validation),
|
| 443 |
+
patch.object(ProviderRuntime, "cleanup", new=cleanup),
|
| 444 |
patch.object(api_runtime_mod.logger, "warning") as log_warning,
|
| 445 |
patch(
|
| 446 |
"messaging.platforms.factory.create_messaging_components",
|
|
|
|
| 450 |
await runtime.startup()
|
| 451 |
await runtime.shutdown()
|
| 452 |
|
| 453 |
+
validation.assert_awaited_once_with()
|
| 454 |
cleanup.assert_awaited_once()
|
| 455 |
create_components.assert_called_once()
|
| 456 |
logged = " ".join(
|
|
|
|
| 494 |
cleanup = AsyncMock()
|
| 495 |
with (
|
| 496 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 497 |
+
patch.object(ProviderRuntime, "validate_configured_models", new=validation),
|
| 498 |
+
patch.object(ProviderRuntime, "cleanup", new=cleanup),
|
| 499 |
):
|
| 500 |
await app({"type": "lifespan"}, receive, send)
|
| 501 |
|
|
|
|
| 524 |
)
|
| 525 |
|
| 526 |
api_app_mod = importlib.import_module("api.app")
|
| 527 |
+
runtime_cleanup = AsyncMock()
|
| 528 |
with (
|
| 529 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 530 |
+
patch.object(ProviderRuntime, "cleanup", new=runtime_cleanup),
|
| 531 |
patch(
|
| 532 |
"messaging.platforms.factory.create_messaging_components",
|
| 533 |
side_effect=ImportError("discord not installed"),
|
|
|
|
| 537 |
pass
|
| 538 |
|
| 539 |
assert getattr(app.state, "messaging_runtime", None) is None
|
| 540 |
+
runtime_cleanup.assert_awaited_once()
|
| 541 |
|
| 542 |
|
| 543 |
def test_app_lifespan_platform_start_exception_cleanup_still_runs(tmp_path):
|
|
|
|
| 574 |
cli_manager.stop_all = AsyncMock()
|
| 575 |
|
| 576 |
api_app_mod = importlib.import_module("api.app")
|
| 577 |
+
runtime_cleanup = AsyncMock()
|
| 578 |
with (
|
| 579 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 580 |
+
patch.object(ProviderRuntime, "cleanup", new=runtime_cleanup),
|
| 581 |
patch(
|
| 582 |
"messaging.platforms.factory.create_messaging_components",
|
| 583 |
return_value=fake_components,
|
|
|
|
| 588 |
):
|
| 589 |
pass
|
| 590 |
|
| 591 |
+
runtime_cleanup.assert_awaited_once()
|
| 592 |
|
| 593 |
|
| 594 |
def test_app_lifespan_flush_pending_save_exception_warning_only(tmp_path):
|
|
|
|
| 626 |
cli_manager.stop_all = AsyncMock()
|
| 627 |
|
| 628 |
api_app_mod = importlib.import_module("api.app")
|
| 629 |
+
runtime_cleanup = AsyncMock()
|
| 630 |
with (
|
| 631 |
patch.object(api_app_mod, "get_settings", return_value=settings),
|
| 632 |
+
patch.object(ProviderRuntime, "cleanup", new=runtime_cleanup),
|
| 633 |
patch(
|
| 634 |
"messaging.platforms.factory.create_messaging_components",
|
| 635 |
return_value=fake_components,
|
|
|
|
| 641 |
pass
|
| 642 |
|
| 643 |
session_store.flush_pending_save.assert_called_once()
|
| 644 |
+
runtime_cleanup.assert_awaited_once()
|
| 645 |
|
| 646 |
|
| 647 |
def test_create_app_writes_server_log_under_fcc_home(monkeypatch, tmp_path):
|
|
@@ -1,6 +1,6 @@
|
|
| 1 |
from types import SimpleNamespace
|
| 2 |
from typing import cast
|
| 3 |
-
from unittest.mock import
|
| 4 |
|
| 5 |
import pytest
|
| 6 |
from fastapi import HTTPException
|
|
@@ -8,37 +8,24 @@ from starlette.applications import Starlette
|
|
| 8 |
from starlette.datastructures import State
|
| 9 |
|
| 10 |
from api.dependencies import (
|
| 11 |
-
|
| 12 |
-
get_provider,
|
| 13 |
-
get_provider_for_type,
|
| 14 |
get_settings,
|
|
|
|
|
|
|
| 15 |
resolve_provider,
|
| 16 |
)
|
| 17 |
from config.nim import NimSettings
|
| 18 |
-
from providers.
|
| 19 |
-
from providers.codestral import CodestralProvider
|
| 20 |
-
from providers.deepseek import DeepSeekProvider
|
| 21 |
-
from providers.exceptions import ServiceUnavailableError, UnknownProviderTypeError
|
| 22 |
-
from providers.gemini import GeminiProvider
|
| 23 |
-
from providers.groq import GroqProvider
|
| 24 |
-
from providers.lmstudio import LMStudioProvider
|
| 25 |
-
from providers.mistral import MistralProvider
|
| 26 |
from providers.nvidia_nim import NvidiaNimProvider
|
| 27 |
-
from providers.
|
| 28 |
-
from providers.open_router import OpenRouterProvider
|
| 29 |
-
from providers.registry import ProviderRegistry
|
| 30 |
-
from providers.wafer import WaferProvider
|
| 31 |
|
| 32 |
|
| 33 |
def _make_mock_settings(**overrides):
|
| 34 |
-
"""Create a mock settings object with
|
| 35 |
mock = MagicMock()
|
| 36 |
mock.model = "nvidia_nim/meta/llama3"
|
| 37 |
mock.provider_type = "nvidia_nim"
|
| 38 |
mock.nvidia_nim_api_key = "test_key"
|
| 39 |
-
mock.provider_rate_limit = 40
|
| 40 |
-
mock.provider_rate_window = 60
|
| 41 |
-
mock.provider_max_concurrency = 5
|
| 42 |
mock.open_router_api_key = "test_openrouter_key"
|
| 43 |
mock.mistral_api_key = "test_mistral_key"
|
| 44 |
mock.codestral_api_key = "test_codestral_key"
|
|
@@ -56,6 +43,7 @@ def _make_mock_settings(**overrides):
|
|
| 56 |
mock.lmstudio_proxy = ""
|
| 57 |
mock.llamacpp_proxy = ""
|
| 58 |
mock.kimi_proxy = ""
|
|
|
|
| 59 |
mock.wafer_proxy = ""
|
| 60 |
mock.opencode_proxy = ""
|
| 61 |
mock.opencode_go_proxy = ""
|
|
@@ -68,624 +56,147 @@ def _make_mock_settings(**overrides):
|
|
| 68 |
mock.groq_proxy = ""
|
| 69 |
mock.cerebras_api_key = ""
|
| 70 |
mock.cerebras_proxy = ""
|
| 71 |
-
mock.
|
|
|
|
|
|
|
| 72 |
mock.http_read_timeout = 300.0
|
| 73 |
mock.http_write_timeout = 10.0
|
| 74 |
mock.http_connect_timeout = 10.0
|
| 75 |
mock.enable_model_thinking = True
|
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
for key, value in overrides.items():
|
| 77 |
setattr(mock, key, value)
|
| 78 |
return mock
|
| 79 |
|
| 80 |
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
saved = api.dependencies._providers
|
| 87 |
-
api.dependencies._providers = {}
|
| 88 |
-
yield
|
| 89 |
-
api.dependencies._providers = saved
|
| 90 |
|
| 91 |
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
|
| 97 |
-
p1 = get_provider()
|
| 98 |
-
p2 = get_provider()
|
| 99 |
|
| 100 |
-
|
| 101 |
-
assert p1 is p2
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
@pytest.mark.asyncio
|
| 105 |
-
async def test_get_settings():
|
| 106 |
settings = get_settings()
|
| 107 |
assert settings is not None
|
| 108 |
-
# Verify it calls the internal _get_settings
|
| 109 |
with patch("api.dependencies._get_settings") as mock_get:
|
| 110 |
get_settings()
|
| 111 |
mock_get.assert_called_once()
|
| 112 |
|
| 113 |
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 117 |
-
mock_settings.return_value = _make_mock_settings()
|
| 118 |
-
|
| 119 |
-
provider = get_provider()
|
| 120 |
-
assert isinstance(provider, NvidiaNimProvider)
|
| 121 |
-
provider._client = AsyncMock()
|
| 122 |
-
|
| 123 |
-
await cleanup_provider()
|
| 124 |
-
|
| 125 |
-
provider._client.close.assert_called_once()
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
@pytest.mark.asyncio
|
| 129 |
-
async def test_cleanup_provider_no_client():
|
| 130 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 131 |
-
mock_settings.return_value = _make_mock_settings()
|
| 132 |
-
|
| 133 |
-
provider = get_provider()
|
| 134 |
-
if hasattr(provider, "_client"):
|
| 135 |
-
del provider._client
|
| 136 |
-
|
| 137 |
-
await cleanup_provider()
|
| 138 |
-
# Should not raise
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
@pytest.mark.asyncio
|
| 142 |
-
async def test_get_provider_open_router():
|
| 143 |
-
"""Test that provider_type=open_router returns OpenRouterProvider."""
|
| 144 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 145 |
-
mock_settings.return_value = _make_mock_settings(provider_type="open_router")
|
| 146 |
-
|
| 147 |
-
provider = get_provider()
|
| 148 |
-
|
| 149 |
-
assert isinstance(provider, OpenRouterProvider)
|
| 150 |
-
assert provider._base_url == "https://openrouter.ai/api/v1"
|
| 151 |
-
assert provider._api_key == "test_openrouter_key"
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
@pytest.mark.asyncio
|
| 155 |
-
async def test_get_provider_lmstudio():
|
| 156 |
-
"""Test that provider_type=lmstudio returns LMStudioProvider."""
|
| 157 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 158 |
-
mock_settings.return_value = _make_mock_settings(provider_type="lmstudio")
|
| 159 |
-
|
| 160 |
-
provider = get_provider()
|
| 161 |
-
|
| 162 |
-
assert isinstance(provider, LMStudioProvider)
|
| 163 |
-
assert provider._base_url == "http://localhost:1234/v1"
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
@pytest.mark.asyncio
|
| 167 |
-
async def test_get_provider_ollama():
|
| 168 |
-
"""Test that provider_type=ollama returns OllamaProvider without an API key."""
|
| 169 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 170 |
-
mock_settings.return_value = _make_mock_settings(provider_type="ollama")
|
| 171 |
-
|
| 172 |
-
provider = get_provider()
|
| 173 |
-
|
| 174 |
-
assert isinstance(provider, OllamaProvider)
|
| 175 |
-
assert provider._base_url == "http://localhost:11434"
|
| 176 |
-
assert provider._api_key == "ollama"
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
@pytest.mark.asyncio
|
| 180 |
-
async def test_get_provider_deepseek():
|
| 181 |
-
"""Test that provider_type=deepseek returns DeepSeekProvider."""
|
| 182 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 183 |
-
mock_settings.return_value = _make_mock_settings(provider_type="deepseek")
|
| 184 |
-
|
| 185 |
-
provider = get_provider()
|
| 186 |
-
|
| 187 |
-
assert isinstance(provider, DeepSeekProvider)
|
| 188 |
-
assert provider._base_url == "https://api.deepseek.com/anthropic"
|
| 189 |
-
assert provider._api_key == "test_deepseek_key"
|
| 190 |
-
assert provider._config.enable_thinking is True
|
| 191 |
-
|
| 192 |
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
"""DeepSeek provider always uses the fixed provider base URL."""
|
| 196 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 197 |
-
mock_settings.return_value = _make_mock_settings(
|
| 198 |
-
provider_type="deepseek",
|
| 199 |
-
)
|
| 200 |
|
| 201 |
-
provider = get_provider()
|
| 202 |
|
| 203 |
-
|
| 204 |
-
|
| 205 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 206 |
|
| 207 |
-
@pytest.mark.asyncio
|
| 208 |
-
async def test_get_provider_deepseek_passes_enable_model_thinking():
|
| 209 |
-
"""DeepSeek provider receives the fallback thinking toggle."""
|
| 210 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 211 |
-
mock_settings.return_value = _make_mock_settings(
|
| 212 |
-
provider_type="deepseek",
|
| 213 |
-
enable_model_thinking=False,
|
| 214 |
-
)
|
| 215 |
-
|
| 216 |
-
provider = get_provider()
|
| 217 |
-
|
| 218 |
-
assert isinstance(provider, DeepSeekProvider)
|
| 219 |
-
assert provider._config.enable_thinking is False
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
@pytest.mark.asyncio
|
| 223 |
-
async def test_get_provider_mistral():
|
| 224 |
-
"""Test that provider_type=mistral returns MistralProvider."""
|
| 225 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 226 |
-
mock_settings.return_value = _make_mock_settings(provider_type="mistral")
|
| 227 |
-
|
| 228 |
-
provider = get_provider()
|
| 229 |
-
|
| 230 |
-
assert isinstance(provider, MistralProvider)
|
| 231 |
-
assert provider._base_url == "https://api.mistral.ai/v1"
|
| 232 |
-
assert provider._api_key == "test_mistral_key"
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
@pytest.mark.asyncio
|
| 236 |
-
async def test_get_provider_mistral_codestral():
|
| 237 |
-
"""provider_type=mistral_codestral returns CodestralProvider."""
|
| 238 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 239 |
-
mock_settings.return_value = _make_mock_settings(
|
| 240 |
-
provider_type="mistral_codestral",
|
| 241 |
-
)
|
| 242 |
-
|
| 243 |
-
provider = get_provider()
|
| 244 |
-
|
| 245 |
-
assert isinstance(provider, CodestralProvider)
|
| 246 |
-
assert provider._base_url == "https://codestral.mistral.ai/v1"
|
| 247 |
-
assert provider._api_key == "test_codestral_key"
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
@pytest.mark.asyncio
|
| 251 |
-
async def test_get_provider_gemini():
|
| 252 |
-
"""Test that provider_type=gemini returns GeminiProvider."""
|
| 253 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 254 |
-
mock_settings.return_value = _make_mock_settings(
|
| 255 |
-
provider_type="gemini",
|
| 256 |
-
gemini_api_key="secret",
|
| 257 |
-
)
|
| 258 |
-
|
| 259 |
-
provider = get_provider()
|
| 260 |
-
|
| 261 |
-
assert isinstance(provider, GeminiProvider)
|
| 262 |
-
assert provider._base_url == (
|
| 263 |
-
"https://generativelanguage.googleapis.com/v1beta/openai"
|
| 264 |
-
)
|
| 265 |
-
assert provider._api_key == "secret"
|
| 266 |
-
|
| 267 |
-
|
| 268 |
-
@pytest.mark.asyncio
|
| 269 |
-
async def test_get_provider_gemini_missing_api_key():
|
| 270 |
-
"""Gemini with empty API key raises HTTPException 503."""
|
| 271 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 272 |
-
mock_settings.return_value = _make_mock_settings(
|
| 273 |
-
provider_type="gemini",
|
| 274 |
-
gemini_api_key="",
|
| 275 |
-
)
|
| 276 |
-
|
| 277 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 278 |
-
get_provider()
|
| 279 |
-
|
| 280 |
-
assert exc_info.value.status_code == 503
|
| 281 |
-
assert "GEMINI_API_KEY" in exc_info.value.detail
|
| 282 |
-
assert "aistudio.google.com" in exc_info.value.detail
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
@pytest.mark.asyncio
|
| 286 |
-
async def test_get_provider_groq():
|
| 287 |
-
"""Test that provider_type=groq returns GroqProvider."""
|
| 288 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 289 |
-
mock_settings.return_value = _make_mock_settings(
|
| 290 |
-
provider_type="groq",
|
| 291 |
-
groq_api_key="secret",
|
| 292 |
-
)
|
| 293 |
-
|
| 294 |
-
provider = get_provider()
|
| 295 |
-
|
| 296 |
-
assert isinstance(provider, GroqProvider)
|
| 297 |
-
assert provider._base_url == "https://api.groq.com/openai/v1"
|
| 298 |
-
assert provider._api_key == "secret"
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
@pytest.mark.asyncio
|
| 302 |
-
async def test_get_provider_groq_missing_api_key():
|
| 303 |
-
"""Groq with empty API key raises HTTPException 503."""
|
| 304 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 305 |
-
mock_settings.return_value = _make_mock_settings(
|
| 306 |
-
provider_type="groq",
|
| 307 |
-
groq_api_key="",
|
| 308 |
-
)
|
| 309 |
-
|
| 310 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 311 |
-
get_provider()
|
| 312 |
-
|
| 313 |
-
assert exc_info.value.status_code == 503
|
| 314 |
-
assert "GROQ_API_KEY" in exc_info.value.detail
|
| 315 |
-
assert "console.groq.com" in exc_info.value.detail
|
| 316 |
-
|
| 317 |
-
|
| 318 |
-
@pytest.mark.asyncio
|
| 319 |
-
async def test_get_provider_cerebras():
|
| 320 |
-
"""Test that provider_type=cerebras returns CerebrasProvider."""
|
| 321 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 322 |
-
mock_settings.return_value = _make_mock_settings(
|
| 323 |
-
provider_type="cerebras",
|
| 324 |
-
cerebras_api_key="secret",
|
| 325 |
-
)
|
| 326 |
-
|
| 327 |
-
provider = get_provider()
|
| 328 |
-
|
| 329 |
-
assert isinstance(provider, CerebrasProvider)
|
| 330 |
-
assert provider._base_url == "https://api.cerebras.ai/v1"
|
| 331 |
-
assert provider._api_key == "secret"
|
| 332 |
-
|
| 333 |
-
|
| 334 |
-
@pytest.mark.asyncio
|
| 335 |
-
async def test_get_provider_cerebras_missing_api_key():
|
| 336 |
-
"""Cerebras with empty API key raises HTTPException 503."""
|
| 337 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 338 |
-
mock_settings.return_value = _make_mock_settings(
|
| 339 |
-
provider_type="cerebras",
|
| 340 |
-
cerebras_api_key="",
|
| 341 |
-
)
|
| 342 |
-
|
| 343 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 344 |
-
get_provider()
|
| 345 |
|
| 346 |
-
|
| 347 |
-
|
| 348 |
-
|
| 349 |
|
|
|
|
|
|
|
|
|
|
| 350 |
|
| 351 |
-
|
| 352 |
-
|
| 353 |
-
|
| 354 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 355 |
-
mock_settings.return_value = _make_mock_settings(provider_type="wafer")
|
| 356 |
-
|
| 357 |
-
provider = get_provider()
|
| 358 |
|
| 359 |
-
assert isinstance(provider, WaferProvider)
|
| 360 |
-
assert provider._base_url == "https://pass.wafer.ai/v1"
|
| 361 |
-
assert provider._api_key == "test_wafer_key"
|
| 362 |
|
|
|
|
|
|
|
| 363 |
|
| 364 |
-
|
| 365 |
-
|
| 366 |
-
"""LM Studio provider uses lm_studio_base_url from settings."""
|
| 367 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 368 |
-
mock_settings.return_value = _make_mock_settings(
|
| 369 |
-
provider_type="lmstudio",
|
| 370 |
-
lm_studio_base_url="http://custom:9999/v1",
|
| 371 |
-
)
|
| 372 |
|
| 373 |
-
|
|
|
|
|
|
|
| 374 |
|
| 375 |
-
assert isinstance(provider, LMStudioProvider)
|
| 376 |
-
assert provider._base_url == "http://custom:9999/v1"
|
| 377 |
|
|
|
|
|
|
|
| 378 |
|
| 379 |
-
|
| 380 |
-
|
| 381 |
-
"""Provider receives http timeouts from settings when creating client."""
|
| 382 |
-
with (
|
| 383 |
-
patch("api.dependencies.get_settings") as mock_settings,
|
| 384 |
-
patch("providers.transports.openai_chat.transport.AsyncOpenAI") as mock_openai,
|
| 385 |
-
):
|
| 386 |
-
mock_settings.return_value = _make_mock_settings(
|
| 387 |
-
http_read_timeout=600.0,
|
| 388 |
-
http_write_timeout=20.0,
|
| 389 |
-
http_connect_timeout=5.0,
|
| 390 |
-
)
|
| 391 |
-
provider = get_provider()
|
| 392 |
-
assert isinstance(provider, NvidiaNimProvider)
|
| 393 |
-
call_kwargs = mock_openai.call_args[1]
|
| 394 |
-
timeout = call_kwargs["timeout"]
|
| 395 |
-
assert timeout.read == 600.0
|
| 396 |
-
assert timeout.write == 20.0
|
| 397 |
-
assert timeout.connect == 5.0
|
| 398 |
-
|
| 399 |
-
|
| 400 |
-
@pytest.mark.asyncio
|
| 401 |
-
async def test_get_provider_passes_proxy_from_settings():
|
| 402 |
-
"""Provider receives configured proxy and builds a proxied HTTP client."""
|
| 403 |
-
with (
|
| 404 |
-
patch("api.dependencies.get_settings") as mock_settings,
|
| 405 |
-
patch(
|
| 406 |
-
"providers.transports.openai_chat.transport.httpx.AsyncClient"
|
| 407 |
-
) as mock_http_client,
|
| 408 |
-
patch("providers.transports.openai_chat.transport.AsyncOpenAI") as mock_openai,
|
| 409 |
):
|
| 410 |
-
|
| 411 |
-
nvidia_nim_proxy="http://proxy.example:8080"
|
| 412 |
-
)
|
| 413 |
|
| 414 |
-
provider = get_provider()
|
| 415 |
|
| 416 |
-
|
| 417 |
-
|
| 418 |
-
assert mock_http_client.call_args.kwargs["proxy"] == "http://proxy.example:8080"
|
| 419 |
-
assert (
|
| 420 |
-
mock_openai.call_args.kwargs["http_client"] is mock_http_client.return_value
|
| 421 |
-
)
|
| 422 |
|
|
|
|
|
|
|
| 423 |
|
| 424 |
-
@pytest.mark.asyncio
|
| 425 |
-
async def test_get_provider_ignores_non_string_proxy_value():
|
| 426 |
-
"""Mock settings without proxy attrs should not fail provider construction."""
|
| 427 |
with (
|
| 428 |
-
patch
|
| 429 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 430 |
):
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
)
|
| 434 |
-
|
| 435 |
-
provider = get_provider()
|
| 436 |
-
|
| 437 |
-
assert isinstance(provider, NvidiaNimProvider)
|
| 438 |
-
assert mock_openai.call_args.kwargs["http_client"] is None
|
| 439 |
-
|
| 440 |
-
|
| 441 |
-
@pytest.mark.asyncio
|
| 442 |
-
async def test_get_provider_nvidia_nim_missing_api_key():
|
| 443 |
-
"""NVIDIA NIM with empty API key raises HTTPException 503."""
|
| 444 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 445 |
-
mock_settings.return_value = _make_mock_settings(nvidia_nim_api_key="")
|
| 446 |
-
|
| 447 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 448 |
-
get_provider()
|
| 449 |
-
|
| 450 |
-
assert exc_info.value.status_code == 503
|
| 451 |
-
assert "NVIDIA_NIM_API_KEY" in exc_info.value.detail
|
| 452 |
-
assert "build.nvidia.com" in exc_info.value.detail
|
| 453 |
-
|
| 454 |
-
|
| 455 |
-
@pytest.mark.asyncio
|
| 456 |
-
async def test_get_provider_nvidia_nim_whitespace_only_api_key():
|
| 457 |
-
"""NVIDIA NIM with whitespace-only API key raises HTTPException 503."""
|
| 458 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 459 |
-
mock_settings.return_value = _make_mock_settings(nvidia_nim_api_key=" ")
|
| 460 |
-
|
| 461 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 462 |
-
get_provider()
|
| 463 |
-
|
| 464 |
-
assert exc_info.value.status_code == 503
|
| 465 |
-
assert "NVIDIA_NIM_API_KEY" in exc_info.value.detail
|
| 466 |
-
|
| 467 |
-
|
| 468 |
-
@pytest.mark.asyncio
|
| 469 |
-
async def test_get_provider_open_router_missing_api_key():
|
| 470 |
-
"""OpenRouter with empty API key raises HTTPException 503."""
|
| 471 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 472 |
-
mock_settings.return_value = _make_mock_settings(
|
| 473 |
-
provider_type="open_router",
|
| 474 |
-
open_router_api_key="",
|
| 475 |
-
)
|
| 476 |
-
|
| 477 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 478 |
-
get_provider()
|
| 479 |
-
|
| 480 |
-
assert exc_info.value.status_code == 503
|
| 481 |
-
assert "OPENROUTER_API_KEY" in exc_info.value.detail
|
| 482 |
-
assert "openrouter.ai" in exc_info.value.detail
|
| 483 |
-
|
| 484 |
-
|
| 485 |
-
@pytest.mark.asyncio
|
| 486 |
-
async def test_get_provider_mistral_missing_api_key():
|
| 487 |
-
"""Mistral with empty API key raises HTTPException 503."""
|
| 488 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 489 |
-
mock_settings.return_value = _make_mock_settings(
|
| 490 |
-
provider_type="mistral",
|
| 491 |
-
mistral_api_key="",
|
| 492 |
-
)
|
| 493 |
-
|
| 494 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 495 |
-
get_provider()
|
| 496 |
-
|
| 497 |
-
assert exc_info.value.status_code == 503
|
| 498 |
-
assert "MISTRAL_API_KEY" in exc_info.value.detail
|
| 499 |
-
assert "console.mistral.ai" in exc_info.value.detail
|
| 500 |
-
|
| 501 |
-
|
| 502 |
-
@pytest.mark.asyncio
|
| 503 |
-
async def test_get_provider_mistral_codestral_missing_api_key():
|
| 504 |
-
"""Mistral Codestral with empty API key raises HTTPException 503."""
|
| 505 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 506 |
-
mock_settings.return_value = _make_mock_settings(
|
| 507 |
-
provider_type="mistral_codestral",
|
| 508 |
-
codestral_api_key="",
|
| 509 |
-
)
|
| 510 |
-
|
| 511 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 512 |
-
get_provider()
|
| 513 |
-
|
| 514 |
-
assert exc_info.value.status_code == 503
|
| 515 |
-
assert "CODESTRAL_API_KEY" in exc_info.value.detail
|
| 516 |
-
assert "console.mistral.ai" in exc_info.value.detail
|
| 517 |
-
|
| 518 |
-
|
| 519 |
-
@pytest.mark.asyncio
|
| 520 |
-
async def test_get_provider_deepseek_missing_api_key():
|
| 521 |
-
"""DeepSeek with empty API key raises HTTPException 503."""
|
| 522 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 523 |
-
mock_settings.return_value = _make_mock_settings(
|
| 524 |
-
provider_type="deepseek",
|
| 525 |
-
deepseek_api_key="",
|
| 526 |
-
)
|
| 527 |
-
|
| 528 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 529 |
-
get_provider()
|
| 530 |
-
|
| 531 |
-
assert exc_info.value.status_code == 503
|
| 532 |
-
assert "DEEPSEEK_API_KEY" in exc_info.value.detail
|
| 533 |
-
assert "platform.deepseek.com" in exc_info.value.detail
|
| 534 |
-
|
| 535 |
-
|
| 536 |
-
@pytest.mark.asyncio
|
| 537 |
-
async def test_get_provider_wafer_missing_api_key():
|
| 538 |
-
"""Wafer with empty API key raises HTTPException 503."""
|
| 539 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 540 |
-
mock_settings.return_value = _make_mock_settings(
|
| 541 |
-
provider_type="wafer",
|
| 542 |
-
wafer_api_key="",
|
| 543 |
-
)
|
| 544 |
-
|
| 545 |
-
with pytest.raises(HTTPException) as exc_info:
|
| 546 |
-
get_provider()
|
| 547 |
-
|
| 548 |
-
assert exc_info.value.status_code == 503
|
| 549 |
-
assert "WAFER_API_KEY" in exc_info.value.detail
|
| 550 |
-
assert "wafer.ai" in exc_info.value.detail
|
| 551 |
-
|
| 552 |
-
|
| 553 |
-
@pytest.mark.asyncio
|
| 554 |
-
async def test_get_provider_unknown_type():
|
| 555 |
-
"""Unknown ``provider_type`` raises :exc:`~providers.exceptions.UnknownProviderTypeError`."""
|
| 556 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 557 |
-
mock_settings.return_value = _make_mock_settings(provider_type="unknown")
|
| 558 |
-
|
| 559 |
-
with pytest.raises(UnknownProviderTypeError, match="Unknown provider_type"):
|
| 560 |
-
get_provider()
|
| 561 |
-
|
| 562 |
-
|
| 563 |
-
@pytest.mark.asyncio
|
| 564 |
-
async def test_cleanup_provider_close_raises():
|
| 565 |
-
"""cleanup_provider handles close() raising an exception."""
|
| 566 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 567 |
-
mock_settings.return_value = _make_mock_settings()
|
| 568 |
-
|
| 569 |
-
provider = get_provider()
|
| 570 |
-
assert isinstance(provider, NvidiaNimProvider)
|
| 571 |
-
provider._client = AsyncMock()
|
| 572 |
-
provider._client.close = AsyncMock(side_effect=RuntimeError("cleanup failed"))
|
| 573 |
-
|
| 574 |
-
# Should propagate the error
|
| 575 |
-
with pytest.raises(RuntimeError, match="cleanup failed"):
|
| 576 |
-
await cleanup_provider()
|
| 577 |
-
|
| 578 |
-
|
| 579 |
-
# --- Provider Registry Tests ---
|
| 580 |
-
|
| 581 |
-
|
| 582 |
-
@pytest.mark.asyncio
|
| 583 |
-
async def test_get_provider_for_type_caches():
|
| 584 |
-
"""get_provider_for_type returns cached provider on second call."""
|
| 585 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 586 |
-
mock_settings.return_value = _make_mock_settings()
|
| 587 |
-
|
| 588 |
-
p1 = get_provider_for_type("nvidia_nim")
|
| 589 |
-
p2 = get_provider_for_type("nvidia_nim")
|
| 590 |
-
|
| 591 |
-
assert p1 is p2
|
| 592 |
-
assert isinstance(p1, NvidiaNimProvider)
|
| 593 |
-
|
| 594 |
-
|
| 595 |
-
@pytest.mark.asyncio
|
| 596 |
-
async def test_get_provider_for_type_different_types():
|
| 597 |
-
"""get_provider_for_type creates separate providers per type."""
|
| 598 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 599 |
-
mock_settings.return_value = _make_mock_settings()
|
| 600 |
-
|
| 601 |
-
nim = get_provider_for_type("nvidia_nim")
|
| 602 |
-
lmstudio = get_provider_for_type("lmstudio")
|
| 603 |
-
|
| 604 |
-
assert isinstance(nim, NvidiaNimProvider)
|
| 605 |
-
assert isinstance(lmstudio, LMStudioProvider)
|
| 606 |
-
assert nim is not lmstudio
|
| 607 |
-
|
| 608 |
|
| 609 |
-
@pytest.mark.asyncio
|
| 610 |
-
async def test_get_provider_for_type_missing_key_raises_503():
|
| 611 |
-
"""get_provider_for_type raises HTTPException 503 for missing API key."""
|
| 612 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 613 |
-
mock_settings.return_value = _make_mock_settings(open_router_api_key="")
|
| 614 |
|
| 615 |
-
|
| 616 |
-
|
| 617 |
|
| 618 |
-
|
| 619 |
-
assert "OPENROUTER_API_KEY" in exc_info.value.detail
|
| 620 |
|
| 621 |
|
| 622 |
-
|
| 623 |
-
|
| 624 |
-
"""cleanup_provider cleans up all providers in the registry."""
|
| 625 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 626 |
-
mock_settings.return_value = _make_mock_settings()
|
| 627 |
|
| 628 |
-
|
| 629 |
-
|
| 630 |
|
| 631 |
-
|
| 632 |
-
|
| 633 |
|
| 634 |
-
nim._client = AsyncMock()
|
| 635 |
-
lmstudio._client = AsyncMock()
|
| 636 |
|
| 637 |
-
|
|
|
|
| 638 |
|
| 639 |
-
|
| 640 |
-
lmstudio._client.aclose.assert_called_once()
|
| 641 |
|
| 642 |
|
| 643 |
-
def
|
| 644 |
-
|
| 645 |
-
|
| 646 |
-
|
| 647 |
-
|
| 648 |
-
app1 = SimpleNamespace(state=State())
|
| 649 |
-
app2 = SimpleNamespace(state=State())
|
| 650 |
-
app1.state.provider_registry = ProviderRegistry()
|
| 651 |
-
app2.state.provider_registry = ProviderRegistry()
|
| 652 |
-
p1 = resolve_provider(
|
| 653 |
-
"nvidia_nim", app=cast(Starlette, app1), settings=settings
|
| 654 |
-
)
|
| 655 |
-
p2 = resolve_provider(
|
| 656 |
-
"nvidia_nim", app=cast(Starlette, app2), settings=settings
|
| 657 |
-
)
|
| 658 |
-
assert isinstance(p1, NvidiaNimProvider)
|
| 659 |
-
assert isinstance(p2, NvidiaNimProvider)
|
| 660 |
-
assert p1 is not p2
|
| 661 |
|
|
|
|
| 662 |
|
| 663 |
-
def test_resolve_provider_missing_registry_raises_service_unavailable() -> None:
|
| 664 |
-
"""HTTP apps must install app.state.provider_registry (e.g. via AppRuntime)."""
|
| 665 |
-
with patch("api.dependencies.get_settings") as mock_settings:
|
| 666 |
-
mock_settings.return_value = _make_mock_settings()
|
| 667 |
-
settings = _make_mock_settings()
|
| 668 |
-
app = SimpleNamespace(state=State())
|
| 669 |
-
assert getattr(app.state, "provider_registry", None) is None
|
| 670 |
-
with pytest.raises(
|
| 671 |
-
ServiceUnavailableError, match="Provider registry is not configured"
|
| 672 |
-
):
|
| 673 |
-
resolve_provider("nvidia_nim", app=cast(Starlette, app), settings=settings)
|
| 674 |
|
|
|
|
|
|
|
| 675 |
|
| 676 |
-
|
| 677 |
-
|
| 678 |
-
import api.dependencies as deps
|
| 679 |
|
| 680 |
-
|
| 681 |
-
|
| 682 |
-
patch.object(
|
| 683 |
-
ProviderRegistry,
|
| 684 |
-
"get",
|
| 685 |
-
side_effect=ValueError("unrelated config"),
|
| 686 |
-
),
|
| 687 |
-
patch.object(deps.logger, "error") as log_err,
|
| 688 |
-
pytest.raises(ValueError, match="unrelated config"),
|
| 689 |
-
):
|
| 690 |
-
deps.resolve_provider("nvidia_nim", app=None, settings=_make_mock_settings())
|
| 691 |
-
log_err.assert_not_called()
|
|
|
|
| 1 |
from types import SimpleNamespace
|
| 2 |
from typing import cast
|
| 3 |
+
from unittest.mock import MagicMock, patch
|
| 4 |
|
| 5 |
import pytest
|
| 6 |
from fastapi import HTTPException
|
|
|
|
| 8 |
from starlette.datastructures import State
|
| 9 |
|
| 10 |
from api.dependencies import (
|
| 11 |
+
get_provider_runtime,
|
|
|
|
|
|
|
| 12 |
get_settings,
|
| 13 |
+
maybe_provider_runtime,
|
| 14 |
+
require_api_key,
|
| 15 |
resolve_provider,
|
| 16 |
)
|
| 17 |
from config.nim import NimSettings
|
| 18 |
+
from providers.exceptions import ServiceUnavailableError
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
from providers.nvidia_nim import NvidiaNimProvider
|
| 20 |
+
from providers.runtime import ProviderRuntime
|
|
|
|
|
|
|
|
|
|
| 21 |
|
| 22 |
|
| 23 |
def _make_mock_settings(**overrides):
|
| 24 |
+
"""Create a mock settings object with provider runtime fields."""
|
| 25 |
mock = MagicMock()
|
| 26 |
mock.model = "nvidia_nim/meta/llama3"
|
| 27 |
mock.provider_type = "nvidia_nim"
|
| 28 |
mock.nvidia_nim_api_key = "test_key"
|
|
|
|
|
|
|
|
|
|
| 29 |
mock.open_router_api_key = "test_openrouter_key"
|
| 30 |
mock.mistral_api_key = "test_mistral_key"
|
| 31 |
mock.codestral_api_key = "test_codestral_key"
|
|
|
|
| 43 |
mock.lmstudio_proxy = ""
|
| 44 |
mock.llamacpp_proxy = ""
|
| 45 |
mock.kimi_proxy = ""
|
| 46 |
+
mock.kimi_api_key = "test_kimi_key"
|
| 47 |
mock.wafer_proxy = ""
|
| 48 |
mock.opencode_proxy = ""
|
| 49 |
mock.opencode_go_proxy = ""
|
|
|
|
| 56 |
mock.groq_proxy = ""
|
| 57 |
mock.cerebras_api_key = ""
|
| 58 |
mock.cerebras_proxy = ""
|
| 59 |
+
mock.provider_rate_limit = 40
|
| 60 |
+
mock.provider_rate_window = 60
|
| 61 |
+
mock.provider_max_concurrency = 5
|
| 62 |
mock.http_read_timeout = 300.0
|
| 63 |
mock.http_write_timeout = 10.0
|
| 64 |
mock.http_connect_timeout = 10.0
|
| 65 |
mock.enable_model_thinking = True
|
| 66 |
+
mock.log_raw_sse_events = False
|
| 67 |
+
mock.log_api_error_tracebacks = False
|
| 68 |
+
mock.configured_chat_model_refs.return_value = ()
|
| 69 |
+
mock.nim = NimSettings()
|
| 70 |
for key, value in overrides.items():
|
| 71 |
setattr(mock, key, value)
|
| 72 |
return mock
|
| 73 |
|
| 74 |
|
| 75 |
+
def _app_with_runtime(settings=None):
|
| 76 |
+
app = SimpleNamespace(state=State())
|
| 77 |
+
app.state.provider_runtime = ProviderRuntime(settings or _make_mock_settings())
|
| 78 |
+
return cast(Starlette, app)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 79 |
|
| 80 |
|
| 81 |
+
def _request(headers=None, token: str = ""):
|
| 82 |
+
return SimpleNamespace(
|
| 83 |
+
headers=headers or {},
|
| 84 |
+
), SimpleNamespace(anthropic_auth_token=token)
|
| 85 |
|
|
|
|
|
|
|
| 86 |
|
| 87 |
+
def test_get_settings():
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 88 |
settings = get_settings()
|
| 89 |
assert settings is not None
|
|
|
|
| 90 |
with patch("api.dependencies._get_settings") as mock_get:
|
| 91 |
get_settings()
|
| 92 |
mock_get.assert_called_once()
|
| 93 |
|
| 94 |
|
| 95 |
+
def test_get_provider_runtime_returns_app_scoped_runtime() -> None:
|
| 96 |
+
app = _app_with_runtime()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
|
| 98 |
+
assert isinstance(get_provider_runtime(app), ProviderRuntime)
|
| 99 |
+
assert maybe_provider_runtime(app) is get_provider_runtime(app)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 100 |
|
|
|
|
| 101 |
|
| 102 |
+
def test_get_provider_runtime_missing_runtime_raises_service_unavailable() -> None:
|
| 103 |
+
app = cast(Starlette, SimpleNamespace(state=State()))
|
| 104 |
|
| 105 |
+
assert maybe_provider_runtime(app) is None
|
| 106 |
+
with pytest.raises(
|
| 107 |
+
ServiceUnavailableError, match="Provider runtime is not configured"
|
| 108 |
+
):
|
| 109 |
+
get_provider_runtime(app)
|
| 110 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 111 |
|
| 112 |
+
def test_resolve_provider_per_app_uses_separate_runtimes() -> None:
|
| 113 |
+
app1 = _app_with_runtime()
|
| 114 |
+
app2 = _app_with_runtime()
|
| 115 |
|
| 116 |
+
with patch("providers.transports.openai_chat.transport.AsyncOpenAI"):
|
| 117 |
+
p1 = resolve_provider("nvidia_nim", app=app1)
|
| 118 |
+
p2 = resolve_provider("nvidia_nim", app=app2)
|
| 119 |
|
| 120 |
+
assert isinstance(p1, NvidiaNimProvider)
|
| 121 |
+
assert isinstance(p2, NvidiaNimProvider)
|
| 122 |
+
assert p1 is not p2
|
|
|
|
|
|
|
|
|
|
|
|
|
| 123 |
|
|
|
|
|
|
|
|
|
|
| 124 |
|
| 125 |
+
def test_resolve_provider_missing_key_raises_503() -> None:
|
| 126 |
+
app = _app_with_runtime(_make_mock_settings(open_router_api_key=""))
|
| 127 |
|
| 128 |
+
with pytest.raises(HTTPException) as exc_info:
|
| 129 |
+
resolve_provider("open_router", app=app)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 130 |
|
| 131 |
+
assert exc_info.value.status_code == 503
|
| 132 |
+
assert "OPENROUTER_API_KEY" in exc_info.value.detail
|
| 133 |
+
assert "openrouter.ai" in exc_info.value.detail
|
| 134 |
|
|
|
|
|
|
|
| 135 |
|
| 136 |
+
def test_resolve_provider_missing_runtime_raises_service_unavailable() -> None:
|
| 137 |
+
app = cast(Starlette, SimpleNamespace(state=State()))
|
| 138 |
|
| 139 |
+
with pytest.raises(
|
| 140 |
+
ServiceUnavailableError, match="Provider runtime is not configured"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 141 |
):
|
| 142 |
+
resolve_provider("nvidia_nim", app=app)
|
|
|
|
|
|
|
| 143 |
|
|
|
|
| 144 |
|
| 145 |
+
def test_resolve_provider_unrelated_value_error_is_not_unknown_provider_log() -> None:
|
| 146 |
+
import api.dependencies as deps
|
|
|
|
|
|
|
|
|
|
|
|
|
| 147 |
|
| 148 |
+
app = _app_with_runtime()
|
| 149 |
+
runtime = get_provider_runtime(app)
|
| 150 |
|
|
|
|
|
|
|
|
|
|
| 151 |
with (
|
| 152 |
+
patch.object(
|
| 153 |
+
runtime,
|
| 154 |
+
"resolve_provider",
|
| 155 |
+
side_effect=ValueError("unrelated config"),
|
| 156 |
+
),
|
| 157 |
+
patch.object(deps.logger, "error") as log_err,
|
| 158 |
+
pytest.raises(ValueError, match="unrelated config"),
|
| 159 |
):
|
| 160 |
+
deps.resolve_provider("nvidia_nim", app=app)
|
| 161 |
+
log_err.assert_not_called()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 162 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 163 |
|
| 164 |
+
def test_require_api_key_allows_when_no_token_configured():
|
| 165 |
+
request, settings = _request(headers={}, token="")
|
| 166 |
|
| 167 |
+
require_api_key(request, settings)
|
|
|
|
| 168 |
|
| 169 |
|
| 170 |
+
def test_require_api_key_rejects_missing_token():
|
| 171 |
+
request, settings = _request(headers={}, token="secret")
|
|
|
|
|
|
|
|
|
|
| 172 |
|
| 173 |
+
with pytest.raises(HTTPException) as exc_info:
|
| 174 |
+
require_api_key(request, settings)
|
| 175 |
|
| 176 |
+
assert exc_info.value.status_code == 401
|
| 177 |
+
assert exc_info.value.detail == "Missing API key"
|
| 178 |
|
|
|
|
|
|
|
| 179 |
|
| 180 |
+
def test_require_api_key_accepts_x_api_key():
|
| 181 |
+
request, settings = _request(headers={"x-api-key": "secret"}, token="secret")
|
| 182 |
|
| 183 |
+
require_api_key(request, settings)
|
|
|
|
| 184 |
|
| 185 |
|
| 186 |
+
def test_require_api_key_accepts_bearer_token_and_strips_model_suffix():
|
| 187 |
+
request, settings = _request(
|
| 188 |
+
headers={"authorization": "Bearer secret:claude-sonnet"},
|
| 189 |
+
token="secret",
|
| 190 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 191 |
|
| 192 |
+
require_api_key(request, settings)
|
| 193 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 194 |
|
| 195 |
+
def test_require_api_key_rejects_invalid_token():
|
| 196 |
+
request, settings = _request(headers={"x-api-key": "wrong"}, token="secret")
|
| 197 |
|
| 198 |
+
with pytest.raises(HTTPException) as exc_info:
|
| 199 |
+
require_api_key(request, settings)
|
|
|
|
| 200 |
|
| 201 |
+
assert exc_info.value.status_code == 401
|
| 202 |
+
assert exc_info.value.detail == "Invalid API key"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -4,7 +4,7 @@ from api.app import create_app
|
|
| 4 |
from api.dependencies import get_settings
|
| 5 |
from config.settings import Settings
|
| 6 |
from providers.model_listing import ProviderModelInfo
|
| 7 |
-
from providers.
|
| 8 |
|
| 9 |
|
| 10 |
def _settings(
|
|
@@ -25,10 +25,10 @@ def _settings(
|
|
| 25 |
def test_models_list_includes_configured_refs_cached_provider_models_and_aliases():
|
| 26 |
app = create_app(lifespan_enabled=False)
|
| 27 |
settings = _settings()
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
app.state.
|
| 32 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 33 |
|
| 34 |
try:
|
|
@@ -72,16 +72,16 @@ def test_models_list_includes_configured_refs_cached_provider_models_and_aliases
|
|
| 72 |
def test_models_list_uses_openrouter_thinking_metadata_for_cached_models():
|
| 73 |
app = create_app(lifespan_enabled=False)
|
| 74 |
settings = _settings(model_opus=None)
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
"open_router",
|
| 79 |
{
|
| 80 |
ProviderModelInfo("reasoning-model", supports_thinking=True),
|
| 81 |
ProviderModelInfo("plain-model", supports_thinking=False),
|
| 82 |
},
|
| 83 |
)
|
| 84 |
-
app.state.
|
| 85 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 86 |
|
| 87 |
try:
|
|
@@ -104,12 +104,12 @@ def test_models_list_uses_cached_metadata_for_configured_openrouter_refs():
|
|
| 104 |
model_opus=None,
|
| 105 |
model_haiku=None,
|
| 106 |
)
|
| 107 |
-
|
| 108 |
-
|
| 109 |
"open_router",
|
| 110 |
{ProviderModelInfo("plain-model", supports_thinking=False)},
|
| 111 |
)
|
| 112 |
-
app.state.
|
| 113 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 114 |
|
| 115 |
try:
|
|
@@ -130,9 +130,9 @@ def test_models_list_includes_cached_wafer_models():
|
|
| 130 |
model_opus=None,
|
| 131 |
model_haiku=None,
|
| 132 |
)
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
app.state.
|
| 136 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 137 |
|
| 138 |
try:
|
|
@@ -148,7 +148,7 @@ def test_models_list_includes_cached_wafer_models():
|
|
| 148 |
assert "claude-3-freecc-no-thinking/wafer/MiniMax-M2.7" in ids
|
| 149 |
|
| 150 |
|
| 151 |
-
def
|
| 152 |
app = create_app(lifespan_enabled=False)
|
| 153 |
settings = _settings()
|
| 154 |
app.dependency_overrides[get_settings] = lambda: settings
|
|
|
|
| 4 |
from api.dependencies import get_settings
|
| 5 |
from config.settings import Settings
|
| 6 |
from providers.model_listing import ProviderModelInfo
|
| 7 |
+
from providers.runtime import ProviderRuntime
|
| 8 |
|
| 9 |
|
| 10 |
def _settings(
|
|
|
|
| 25 |
def test_models_list_includes_configured_refs_cached_provider_models_and_aliases():
|
| 26 |
app = create_app(lifespan_enabled=False)
|
| 27 |
settings = _settings()
|
| 28 |
+
runtime = ProviderRuntime(settings)
|
| 29 |
+
runtime.cache_model_ids("deepseek", {"deepseek-chat"})
|
| 30 |
+
runtime.cache_model_ids("open_router", {"meta/llama-3.3", "anthropic/claude-opus"})
|
| 31 |
+
app.state.provider_runtime = runtime
|
| 32 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 33 |
|
| 34 |
try:
|
|
|
|
| 72 |
def test_models_list_uses_openrouter_thinking_metadata_for_cached_models():
|
| 73 |
app = create_app(lifespan_enabled=False)
|
| 74 |
settings = _settings(model_opus=None)
|
| 75 |
+
runtime = ProviderRuntime(settings)
|
| 76 |
+
runtime.cache_model_ids("deepseek", {"deepseek-chat"})
|
| 77 |
+
runtime.cache_model_infos(
|
| 78 |
"open_router",
|
| 79 |
{
|
| 80 |
ProviderModelInfo("reasoning-model", supports_thinking=True),
|
| 81 |
ProviderModelInfo("plain-model", supports_thinking=False),
|
| 82 |
},
|
| 83 |
)
|
| 84 |
+
app.state.provider_runtime = runtime
|
| 85 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 86 |
|
| 87 |
try:
|
|
|
|
| 104 |
model_opus=None,
|
| 105 |
model_haiku=None,
|
| 106 |
)
|
| 107 |
+
runtime = ProviderRuntime(settings)
|
| 108 |
+
runtime.cache_model_infos(
|
| 109 |
"open_router",
|
| 110 |
{ProviderModelInfo("plain-model", supports_thinking=False)},
|
| 111 |
)
|
| 112 |
+
app.state.provider_runtime = runtime
|
| 113 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 114 |
|
| 115 |
try:
|
|
|
|
| 130 |
model_opus=None,
|
| 131 |
model_haiku=None,
|
| 132 |
)
|
| 133 |
+
runtime = ProviderRuntime(settings)
|
| 134 |
+
runtime.cache_model_ids("wafer", {"DeepSeek-V4-Pro", "MiniMax-M2.7"})
|
| 135 |
+
app.state.provider_runtime = runtime
|
| 136 |
app.dependency_overrides[get_settings] = lambda: settings
|
| 137 |
|
| 138 |
try:
|
|
|
|
| 148 |
assert "claude-3-freecc-no-thinking/wafer/MiniMax-M2.7" in ids
|
| 149 |
|
| 150 |
|
| 151 |
+
def test_models_list_works_without_provider_runtime():
|
| 152 |
app = create_app(lifespan_enabled=False)
|
| 153 |
settings = _settings()
|
| 154 |
app.dependency_overrides[get_settings] = lambda: settings
|
|
@@ -11,7 +11,7 @@ _API_ALLOWED_PROVIDER_MODULES = frozenset(
|
|
| 11 |
"providers",
|
| 12 |
"providers.base",
|
| 13 |
"providers.exceptions",
|
| 14 |
-
"providers.
|
| 15 |
}
|
| 16 |
)
|
| 17 |
|
|
@@ -56,13 +56,27 @@ def test_core_does_not_import_product_packages() -> None:
|
|
| 56 |
|
| 57 |
def test_provider_catalog_is_single_source_for_supported_ids() -> None:
|
| 58 |
from config.provider_catalog import PROVIDER_CATALOG, SUPPORTED_PROVIDER_IDS
|
| 59 |
-
from providers.
|
| 60 |
|
| 61 |
assert tuple(PROVIDER_CATALOG.keys()) == SUPPORTED_PROVIDER_IDS
|
| 62 |
-
assert PROVIDER_DESCRIPTORS is PROVIDER_CATALOG
|
| 63 |
assert set(SUPPORTED_PROVIDER_IDS) == set(PROVIDER_FACTORIES)
|
| 64 |
|
| 65 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 66 |
def test_config_does_not_import_non_config_packages() -> None:
|
| 67 |
"""Settings and env handling must not depend on transport or protocol layers."""
|
| 68 |
repo_root = Path(__file__).resolve().parents[2]
|
|
|
|
| 11 |
"providers",
|
| 12 |
"providers.base",
|
| 13 |
"providers.exceptions",
|
| 14 |
+
"providers.runtime",
|
| 15 |
}
|
| 16 |
)
|
| 17 |
|
|
|
|
| 56 |
|
| 57 |
def test_provider_catalog_is_single_source_for_supported_ids() -> None:
|
| 58 |
from config.provider_catalog import PROVIDER_CATALOG, SUPPORTED_PROVIDER_IDS
|
| 59 |
+
from providers.runtime import PROVIDER_FACTORIES
|
| 60 |
|
| 61 |
assert tuple(PROVIDER_CATALOG.keys()) == SUPPORTED_PROVIDER_IDS
|
|
|
|
| 62 |
assert set(SUPPORTED_PROVIDER_IDS) == set(PROVIDER_FACTORIES)
|
| 63 |
|
| 64 |
|
| 65 |
+
def test_provider_runtime_replaces_old_registry_module() -> None:
|
| 66 |
+
repo_root = Path(__file__).resolve().parents[2]
|
| 67 |
+
|
| 68 |
+
assert not (repo_root / "providers" / "registry.py").exists()
|
| 69 |
+
assert (repo_root / "providers" / "runtime" / "runtime.py").exists()
|
| 70 |
+
assert (repo_root / "providers" / "runtime" / "factory.py").exists()
|
| 71 |
+
assert (repo_root / "providers" / "runtime" / "discovery.py").exists()
|
| 72 |
+
|
| 73 |
+
offenders = _imports_matching(
|
| 74 |
+
[repo_root / "api", repo_root / "tests", repo_root / "smoke"],
|
| 75 |
+
forbidden_prefixes=("providers.registry",),
|
| 76 |
+
)
|
| 77 |
+
assert offenders == []
|
| 78 |
+
|
| 79 |
+
|
| 80 |
def test_config_does_not_import_non_config_packages() -> None:
|
| 81 |
"""Settings and env handling must not depend on transport or protocol layers."""
|
| 82 |
repo_root = Path(__file__).resolve().parents[2]
|
|
@@ -19,7 +19,7 @@ from providers.model_listing import ProviderModelInfo
|
|
| 19 |
from providers.nvidia_nim import NvidiaNimProvider
|
| 20 |
from providers.ollama import OllamaProvider
|
| 21 |
from providers.open_router import OpenRouterProvider
|
| 22 |
-
from providers.
|
| 23 |
from providers.wafer import WaferProvider
|
| 24 |
|
| 25 |
|
|
@@ -353,32 +353,34 @@ class FakeProvider(BaseProvider):
|
|
| 353 |
|
| 354 |
|
| 355 |
@pytest.mark.asyncio
|
| 356 |
-
async def
|
| 357 |
-
|
|
|
|
|
|
|
| 358 |
{
|
| 359 |
"nvidia_nim": FakeProvider(frozenset({"nim-model"})),
|
| 360 |
"open_router": FakeProvider(frozenset({"anthropic/claude-opus"})),
|
| 361 |
-
}
|
| 362 |
)
|
| 363 |
-
settings = _settings(model_opus="open_router/anthropic/claude-opus")
|
| 364 |
|
| 365 |
-
await
|
| 366 |
|
| 367 |
-
assert
|
| 368 |
"nvidia_nim": frozenset({"nim-model"}),
|
| 369 |
"open_router": frozenset({"anthropic/claude-opus"}),
|
| 370 |
}
|
| 371 |
|
| 372 |
|
| 373 |
@pytest.mark.asyncio
|
| 374 |
-
async def
|
| 375 |
-
registry = ProviderRegistry(
|
| 376 |
-
{"nvidia_nim": FakeProvider(frozenset({"different-model"}))}
|
| 377 |
-
)
|
| 378 |
settings = _settings(model_sonnet="nvidia_nim/nim-model")
|
|
|
|
|
|
|
|
|
|
|
|
|
| 379 |
|
| 380 |
with pytest.raises(ServiceUnavailableError) as exc_info:
|
| 381 |
-
await
|
| 382 |
|
| 383 |
message = exc_info.value.message
|
| 384 |
assert "sources=MODEL,MODEL_SONNET" in message
|
|
@@ -388,19 +390,20 @@ async def test_registry_validation_reports_missing_model_with_sources() -> None:
|
|
| 388 |
|
| 389 |
|
| 390 |
@pytest.mark.asyncio
|
| 391 |
-
async def
|
| 392 |
-
|
|
|
|
|
|
|
| 393 |
{
|
| 394 |
"nvidia_nim": FakeProvider(frozenset({"different-model"})),
|
| 395 |
"open_router": FakeProvider(
|
| 396 |
error=ModelListResponseError("bad model-list shape")
|
| 397 |
),
|
| 398 |
-
}
|
| 399 |
)
|
| 400 |
-
settings = _settings(model_opus="open_router/anthropic/claude-opus")
|
| 401 |
|
| 402 |
with pytest.raises(ServiceUnavailableError) as exc_info:
|
| 403 |
-
await
|
| 404 |
|
| 405 |
message = exc_info.value.message
|
| 406 |
assert "sources=MODEL provider=nvidia_nim model=nim-model" in message
|
|
@@ -412,10 +415,12 @@ async def test_registry_validation_aggregates_multiple_failures() -> None:
|
|
| 412 |
|
| 413 |
|
| 414 |
@pytest.mark.asyncio
|
| 415 |
-
async def
|
| 416 |
nim_started = asyncio.Event()
|
| 417 |
router_started = asyncio.Event()
|
| 418 |
-
|
|
|
|
|
|
|
| 419 |
{
|
| 420 |
"nvidia_nim": FakeProvider(
|
| 421 |
frozenset({"nim-model"}),
|
|
@@ -427,56 +432,57 @@ async def test_registry_validation_queries_providers_concurrently() -> None:
|
|
| 427 |
started=router_started,
|
| 428 |
peer_started=nim_started,
|
| 429 |
),
|
| 430 |
-
}
|
| 431 |
)
|
| 432 |
-
settings = _settings(model_opus="open_router/anthropic/claude-opus")
|
| 433 |
|
| 434 |
-
await asyncio.wait_for(
|
| 435 |
|
| 436 |
|
| 437 |
@pytest.mark.asyncio
|
| 438 |
-
async def
|
| 439 |
None
|
| 440 |
):
|
| 441 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 442 |
{
|
| 443 |
"open_router": FakeProvider(frozenset({"anthropic/claude-sonnet"})),
|
| 444 |
"lmstudio": FakeProvider(frozenset({"local-qwen"})),
|
| 445 |
"ollama": FakeProvider(frozenset({"llama3.1"})),
|
| 446 |
-
}
|
| 447 |
-
)
|
| 448 |
-
settings = _settings(
|
| 449 |
-
model="lmstudio/local-qwen",
|
| 450 |
-
open_router_api_key="open-router-key",
|
| 451 |
)
|
| 452 |
|
| 453 |
-
await
|
| 454 |
|
| 455 |
-
assert
|
| 456 |
"open_router": frozenset({"anthropic/claude-sonnet"}),
|
| 457 |
"lmstudio": frozenset({"local-qwen"}),
|
| 458 |
}
|
| 459 |
|
| 460 |
|
| 461 |
@pytest.mark.asyncio
|
| 462 |
-
async def
|
| 463 |
-
registry = ProviderRegistry(
|
| 464 |
-
{"nvidia_nim": FakeProvider(error=RuntimeError("upstream down"))}
|
| 465 |
-
)
|
| 466 |
-
registry.cache_model_ids("nvidia_nim", {"cached-model"})
|
| 467 |
settings = _settings(
|
| 468 |
model="nvidia_nim/cached-model",
|
| 469 |
nvidia_nim_api_key="nim-key",
|
| 470 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 471 |
|
| 472 |
-
await
|
| 473 |
|
| 474 |
-
assert
|
| 475 |
|
| 476 |
|
| 477 |
-
def
|
| 478 |
-
|
| 479 |
-
|
| 480 |
"open_router",
|
| 481 |
{
|
| 482 |
ProviderModelInfo("reasoning-model", supports_thinking=True),
|
|
@@ -484,39 +490,36 @@ def test_registry_metadata_cache_exposes_ids_and_prefixed_infos() -> None:
|
|
| 484 |
},
|
| 485 |
)
|
| 486 |
|
| 487 |
-
assert
|
| 488 |
"open_router": frozenset({"reasoning-model", "plain-model"})
|
| 489 |
}
|
| 490 |
assert (
|
| 491 |
-
|
| 492 |
-
is True
|
| 493 |
-
)
|
| 494 |
-
assert (
|
| 495 |
-
registry.cached_model_supports_thinking("open_router", "plain-model") is False
|
| 496 |
)
|
| 497 |
-
assert
|
|
|
|
| 498 |
ProviderModelInfo("open_router/plain-model", supports_thinking=False),
|
| 499 |
ProviderModelInfo("open_router/reasoning-model", supports_thinking=True),
|
| 500 |
)
|
| 501 |
|
| 502 |
|
| 503 |
-
def
|
| 504 |
-
|
| 505 |
-
|
| 506 |
|
| 507 |
-
assert
|
| 508 |
-
assert
|
| 509 |
-
assert
|
| 510 |
ProviderModelInfo("open_router/plain-model", supports_thinking=None),
|
| 511 |
)
|
| 512 |
|
| 513 |
|
| 514 |
-
def
|
| 515 |
-
|
| 516 |
-
|
| 517 |
-
|
| 518 |
|
| 519 |
-
assert
|
| 520 |
"open_router/a-model",
|
| 521 |
"open_router/z-model",
|
| 522 |
"deepseek/deepseek-chat",
|
|
|
|
| 19 |
from providers.nvidia_nim import NvidiaNimProvider
|
| 20 |
from providers.ollama import OllamaProvider
|
| 21 |
from providers.open_router import OpenRouterProvider
|
| 22 |
+
from providers.runtime import ProviderRuntime
|
| 23 |
from providers.wafer import WaferProvider
|
| 24 |
|
| 25 |
|
|
|
|
| 353 |
|
| 354 |
|
| 355 |
@pytest.mark.asyncio
|
| 356 |
+
async def test_runtime_validation_succeeds_for_all_configured_models() -> None:
|
| 357 |
+
settings = _settings(model_opus="open_router/anthropic/claude-opus")
|
| 358 |
+
runtime = ProviderRuntime(
|
| 359 |
+
settings,
|
| 360 |
{
|
| 361 |
"nvidia_nim": FakeProvider(frozenset({"nim-model"})),
|
| 362 |
"open_router": FakeProvider(frozenset({"anthropic/claude-opus"})),
|
| 363 |
+
},
|
| 364 |
)
|
|
|
|
| 365 |
|
| 366 |
+
await runtime.validate_configured_models()
|
| 367 |
|
| 368 |
+
assert runtime.cached_model_ids() == {
|
| 369 |
"nvidia_nim": frozenset({"nim-model"}),
|
| 370 |
"open_router": frozenset({"anthropic/claude-opus"}),
|
| 371 |
}
|
| 372 |
|
| 373 |
|
| 374 |
@pytest.mark.asyncio
|
| 375 |
+
async def test_runtime_validation_reports_missing_model_with_sources() -> None:
|
|
|
|
|
|
|
|
|
|
| 376 |
settings = _settings(model_sonnet="nvidia_nim/nim-model")
|
| 377 |
+
runtime = ProviderRuntime(
|
| 378 |
+
settings,
|
| 379 |
+
{"nvidia_nim": FakeProvider(frozenset({"different-model"}))},
|
| 380 |
+
)
|
| 381 |
|
| 382 |
with pytest.raises(ServiceUnavailableError) as exc_info:
|
| 383 |
+
await runtime.validate_configured_models()
|
| 384 |
|
| 385 |
message = exc_info.value.message
|
| 386 |
assert "sources=MODEL,MODEL_SONNET" in message
|
|
|
|
| 390 |
|
| 391 |
|
| 392 |
@pytest.mark.asyncio
|
| 393 |
+
async def test_runtime_validation_aggregates_multiple_failures() -> None:
|
| 394 |
+
settings = _settings(model_opus="open_router/anthropic/claude-opus")
|
| 395 |
+
runtime = ProviderRuntime(
|
| 396 |
+
settings,
|
| 397 |
{
|
| 398 |
"nvidia_nim": FakeProvider(frozenset({"different-model"})),
|
| 399 |
"open_router": FakeProvider(
|
| 400 |
error=ModelListResponseError("bad model-list shape")
|
| 401 |
),
|
| 402 |
+
},
|
| 403 |
)
|
|
|
|
| 404 |
|
| 405 |
with pytest.raises(ServiceUnavailableError) as exc_info:
|
| 406 |
+
await runtime.validate_configured_models()
|
| 407 |
|
| 408 |
message = exc_info.value.message
|
| 409 |
assert "sources=MODEL provider=nvidia_nim model=nim-model" in message
|
|
|
|
| 415 |
|
| 416 |
|
| 417 |
@pytest.mark.asyncio
|
| 418 |
+
async def test_runtime_validation_queries_providers_concurrently() -> None:
|
| 419 |
nim_started = asyncio.Event()
|
| 420 |
router_started = asyncio.Event()
|
| 421 |
+
settings = _settings(model_opus="open_router/anthropic/claude-opus")
|
| 422 |
+
runtime = ProviderRuntime(
|
| 423 |
+
settings,
|
| 424 |
{
|
| 425 |
"nvidia_nim": FakeProvider(
|
| 426 |
frozenset({"nim-model"}),
|
|
|
|
| 432 |
started=router_started,
|
| 433 |
peer_started=nim_started,
|
| 434 |
),
|
| 435 |
+
},
|
| 436 |
)
|
|
|
|
| 437 |
|
| 438 |
+
await asyncio.wait_for(runtime.validate_configured_models(), timeout=1.0)
|
| 439 |
|
| 440 |
|
| 441 |
@pytest.mark.asyncio
|
| 442 |
+
async def test_runtime_refresh_model_list_cache_uses_configured_remote_keys_and_referenced_local() -> (
|
| 443 |
None
|
| 444 |
):
|
| 445 |
+
settings = _settings(
|
| 446 |
+
model="lmstudio/local-qwen",
|
| 447 |
+
open_router_api_key="open-router-key",
|
| 448 |
+
)
|
| 449 |
+
runtime = ProviderRuntime(
|
| 450 |
+
settings,
|
| 451 |
{
|
| 452 |
"open_router": FakeProvider(frozenset({"anthropic/claude-sonnet"})),
|
| 453 |
"lmstudio": FakeProvider(frozenset({"local-qwen"})),
|
| 454 |
"ollama": FakeProvider(frozenset({"llama3.1"})),
|
| 455 |
+
},
|
|
|
|
|
|
|
|
|
|
|
|
|
| 456 |
)
|
| 457 |
|
| 458 |
+
await runtime.refresh_model_list_cache()
|
| 459 |
|
| 460 |
+
assert runtime.cached_model_ids() == {
|
| 461 |
"open_router": frozenset({"anthropic/claude-sonnet"}),
|
| 462 |
"lmstudio": frozenset({"local-qwen"}),
|
| 463 |
}
|
| 464 |
|
| 465 |
|
| 466 |
@pytest.mark.asyncio
|
| 467 |
+
async def test_runtime_refresh_model_list_cache_keeps_prior_cache_on_failure() -> None:
|
|
|
|
|
|
|
|
|
|
|
|
|
| 468 |
settings = _settings(
|
| 469 |
model="nvidia_nim/cached-model",
|
| 470 |
nvidia_nim_api_key="nim-key",
|
| 471 |
)
|
| 472 |
+
runtime = ProviderRuntime(
|
| 473 |
+
settings,
|
| 474 |
+
{"nvidia_nim": FakeProvider(error=RuntimeError("upstream down"))},
|
| 475 |
+
)
|
| 476 |
+
runtime.cache_model_ids("nvidia_nim", {"cached-model"})
|
| 477 |
|
| 478 |
+
await runtime.refresh_model_list_cache()
|
| 479 |
|
| 480 |
+
assert runtime.cached_model_ids() == {"nvidia_nim": frozenset({"cached-model"})}
|
| 481 |
|
| 482 |
|
| 483 |
+
def test_runtime_metadata_cache_exposes_ids_and_prefixed_infos() -> None:
|
| 484 |
+
runtime = ProviderRuntime(_settings())
|
| 485 |
+
runtime.cache_model_infos(
|
| 486 |
"open_router",
|
| 487 |
{
|
| 488 |
ProviderModelInfo("reasoning-model", supports_thinking=True),
|
|
|
|
| 490 |
},
|
| 491 |
)
|
| 492 |
|
| 493 |
+
assert runtime.cached_model_ids() == {
|
| 494 |
"open_router": frozenset({"reasoning-model", "plain-model"})
|
| 495 |
}
|
| 496 |
assert (
|
| 497 |
+
runtime.cached_model_supports_thinking("open_router", "reasoning-model") is True
|
|
|
|
|
|
|
|
|
|
|
|
|
| 498 |
)
|
| 499 |
+
assert runtime.cached_model_supports_thinking("open_router", "plain-model") is False
|
| 500 |
+
assert runtime.cached_prefixed_model_infos() == (
|
| 501 |
ProviderModelInfo("open_router/plain-model", supports_thinking=False),
|
| 502 |
ProviderModelInfo("open_router/reasoning-model", supports_thinking=True),
|
| 503 |
)
|
| 504 |
|
| 505 |
|
| 506 |
+
def test_runtime_model_id_cache_keeps_unknown_thinking_support() -> None:
|
| 507 |
+
runtime = ProviderRuntime(_settings())
|
| 508 |
+
runtime.cache_model_ids("open_router", {"plain-model"})
|
| 509 |
|
| 510 |
+
assert runtime.cached_model_ids() == {"open_router": frozenset({"plain-model"})}
|
| 511 |
+
assert runtime.cached_model_supports_thinking("open_router", "plain-model") is None
|
| 512 |
+
assert runtime.cached_prefixed_model_infos() == (
|
| 513 |
ProviderModelInfo("open_router/plain-model", supports_thinking=None),
|
| 514 |
)
|
| 515 |
|
| 516 |
|
| 517 |
+
def test_runtime_cached_prefixed_model_refs_are_deterministic() -> None:
|
| 518 |
+
runtime = ProviderRuntime(_settings())
|
| 519 |
+
runtime.cache_model_ids("deepseek", {"deepseek-chat"})
|
| 520 |
+
runtime.cache_model_ids("open_router", {"z-model", "a-model"})
|
| 521 |
|
| 522 |
+
assert runtime.cached_prefixed_model_refs() == (
|
| 523 |
"open_router/a-model",
|
| 524 |
"open_router/z-model",
|
| 525 |
"deepseek/deepseek-chat",
|
|
@@ -22,12 +22,7 @@ from providers.nvidia_nim import NvidiaNimProvider
|
|
| 22 |
from providers.ollama import OllamaProvider
|
| 23 |
from providers.open_router import OpenRouterProvider
|
| 24 |
from providers.opencode import OpenCodeProvider
|
| 25 |
-
from providers.
|
| 26 |
-
PROVIDER_DESCRIPTORS,
|
| 27 |
-
ProviderRegistry,
|
| 28 |
-
build_provider_config,
|
| 29 |
-
create_provider,
|
| 30 |
-
)
|
| 31 |
from providers.wafer import WaferProvider
|
| 32 |
from providers.zai import ZaiProvider
|
| 33 |
|
|
@@ -74,17 +69,19 @@ def _make_settings(**overrides):
|
|
| 74 |
mock.http_write_timeout = 10.0
|
| 75 |
mock.http_connect_timeout = 10.0
|
| 76 |
mock.enable_model_thinking = True
|
|
|
|
|
|
|
| 77 |
mock.nim = NimSettings()
|
| 78 |
for key, value in overrides.items():
|
| 79 |
setattr(mock, key, value)
|
| 80 |
return mock
|
| 81 |
|
| 82 |
|
| 83 |
-
def
|
| 84 |
-
"""
|
| 85 |
code = (
|
| 86 |
"import sys\n"
|
| 87 |
-
"import providers.
|
| 88 |
"assert 'providers.open_router' not in sys.modules\n"
|
| 89 |
)
|
| 90 |
proc = subprocess.run(
|
|
@@ -96,16 +93,16 @@ def test_importing_registry_does_not_eager_load_other_adapters() -> None:
|
|
| 96 |
assert proc.returncode == 0, proc.stderr or proc.stdout
|
| 97 |
|
| 98 |
|
| 99 |
-
def
|
| 100 |
-
assert set(
|
| 101 |
-
for descriptor in
|
| 102 |
assert descriptor.provider_id
|
| 103 |
assert descriptor.transport_type in {"openai_chat", "anthropic_messages"}
|
| 104 |
assert descriptor.capabilities
|
| 105 |
|
| 106 |
|
| 107 |
def test_ollama_descriptor_uses_native_anthropic_transport():
|
| 108 |
-
descriptor =
|
| 109 |
|
| 110 |
assert descriptor.transport_type == "anthropic_messages"
|
| 111 |
assert descriptor.default_base_url == "http://localhost:11434"
|
|
@@ -113,14 +110,14 @@ def test_ollama_descriptor_uses_native_anthropic_transport():
|
|
| 113 |
|
| 114 |
|
| 115 |
def test_zai_descriptor_uses_fixed_cloud_base_url():
|
| 116 |
-
descriptor =
|
| 117 |
|
| 118 |
assert descriptor.default_base_url == ZAI_DEFAULT_BASE
|
| 119 |
assert descriptor.base_url_attr is None
|
| 120 |
|
| 121 |
|
| 122 |
def test_zai_provider_config_ignores_stale_base_url_setting():
|
| 123 |
-
descriptor =
|
| 124 |
|
| 125 |
config = build_provider_config(
|
| 126 |
descriptor,
|
|
@@ -198,13 +195,12 @@ def test_create_provider_instantiates_each_builtin():
|
|
| 198 |
assert isinstance(create_provider(provider_id, settings), provider_cls)
|
| 199 |
|
| 200 |
|
| 201 |
-
def
|
| 202 |
-
|
| 203 |
-
settings = _make_settings()
|
| 204 |
|
| 205 |
with patch("providers.transports.openai_chat.transport.AsyncOpenAI"):
|
| 206 |
-
first =
|
| 207 |
-
second =
|
| 208 |
|
| 209 |
assert first is second
|
| 210 |
|
|
@@ -215,32 +211,34 @@ def test_unknown_provider_raises_unknown_provider_type_error():
|
|
| 215 |
|
| 216 |
|
| 217 |
@pytest.mark.asyncio
|
| 218 |
-
async def
|
| 219 |
"""Every provider gets cleanup; cache is cleared even when one raises."""
|
| 220 |
-
reg = ProviderRegistry()
|
| 221 |
p1 = MagicMock()
|
| 222 |
p1.cleanup = AsyncMock(side_effect=RuntimeError("first"))
|
| 223 |
p2 = MagicMock()
|
| 224 |
p2.cleanup = AsyncMock()
|
| 225 |
-
|
| 226 |
-
|
| 227 |
with pytest.raises(RuntimeError, match="first"):
|
| 228 |
-
await
|
|
|
|
| 229 |
p1.cleanup.assert_awaited_once()
|
| 230 |
p2.cleanup.assert_awaited_once()
|
| 231 |
-
assert
|
|
|
|
| 232 |
|
| 233 |
|
| 234 |
@pytest.mark.asyncio
|
| 235 |
-
async def
|
| 236 |
-
reg = ProviderRegistry()
|
| 237 |
p1 = MagicMock()
|
| 238 |
p1.cleanup = AsyncMock(side_effect=RuntimeError("a"))
|
| 239 |
p2 = MagicMock()
|
| 240 |
p2.cleanup = AsyncMock(side_effect=RuntimeError("b"))
|
| 241 |
-
|
| 242 |
-
|
| 243 |
with pytest.raises(ExceptionGroup) as exc_info:
|
| 244 |
-
await
|
|
|
|
| 245 |
assert len(exc_info.value.exceptions) == 2
|
| 246 |
-
assert
|
|
|
|
|
|
| 22 |
from providers.ollama import OllamaProvider
|
| 23 |
from providers.open_router import OpenRouterProvider
|
| 24 |
from providers.opencode import OpenCodeProvider
|
| 25 |
+
from providers.runtime import ProviderRuntime, build_provider_config, create_provider
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
from providers.wafer import WaferProvider
|
| 27 |
from providers.zai import ZaiProvider
|
| 28 |
|
|
|
|
| 69 |
mock.http_write_timeout = 10.0
|
| 70 |
mock.http_connect_timeout = 10.0
|
| 71 |
mock.enable_model_thinking = True
|
| 72 |
+
mock.log_raw_sse_events = False
|
| 73 |
+
mock.log_api_error_tracebacks = False
|
| 74 |
mock.nim = NimSettings()
|
| 75 |
for key, value in overrides.items():
|
| 76 |
setattr(mock, key, value)
|
| 77 |
return mock
|
| 78 |
|
| 79 |
|
| 80 |
+
def test_importing_runtime_does_not_eager_load_other_adapters() -> None:
|
| 81 |
+
"""Runtime metadata must not import every provider adapter up front."""
|
| 82 |
code = (
|
| 83 |
"import sys\n"
|
| 84 |
+
"import providers.runtime\n"
|
| 85 |
"assert 'providers.open_router' not in sys.modules\n"
|
| 86 |
)
|
| 87 |
proc = subprocess.run(
|
|
|
|
| 93 |
assert proc.returncode == 0, proc.stderr or proc.stdout
|
| 94 |
|
| 95 |
|
| 96 |
+
def test_provider_catalog_covers_advertised_provider_ids():
|
| 97 |
+
assert set(PROVIDER_CATALOG) == set(SUPPORTED_PROVIDER_IDS)
|
| 98 |
+
for descriptor in PROVIDER_CATALOG.values():
|
| 99 |
assert descriptor.provider_id
|
| 100 |
assert descriptor.transport_type in {"openai_chat", "anthropic_messages"}
|
| 101 |
assert descriptor.capabilities
|
| 102 |
|
| 103 |
|
| 104 |
def test_ollama_descriptor_uses_native_anthropic_transport():
|
| 105 |
+
descriptor = PROVIDER_CATALOG["ollama"]
|
| 106 |
|
| 107 |
assert descriptor.transport_type == "anthropic_messages"
|
| 108 |
assert descriptor.default_base_url == "http://localhost:11434"
|
|
|
|
| 110 |
|
| 111 |
|
| 112 |
def test_zai_descriptor_uses_fixed_cloud_base_url():
|
| 113 |
+
descriptor = PROVIDER_CATALOG["zai"]
|
| 114 |
|
| 115 |
assert descriptor.default_base_url == ZAI_DEFAULT_BASE
|
| 116 |
assert descriptor.base_url_attr is None
|
| 117 |
|
| 118 |
|
| 119 |
def test_zai_provider_config_ignores_stale_base_url_setting():
|
| 120 |
+
descriptor = PROVIDER_CATALOG["zai"]
|
| 121 |
|
| 122 |
config = build_provider_config(
|
| 123 |
descriptor,
|
|
|
|
| 195 |
assert isinstance(create_provider(provider_id, settings), provider_cls)
|
| 196 |
|
| 197 |
|
| 198 |
+
def test_provider_runtime_caches_by_provider_id():
|
| 199 |
+
runtime = ProviderRuntime(_make_settings())
|
|
|
|
| 200 |
|
| 201 |
with patch("providers.transports.openai_chat.transport.AsyncOpenAI"):
|
| 202 |
+
first = runtime.resolve_provider("nvidia_nim")
|
| 203 |
+
second = runtime.resolve_provider("nvidia_nim")
|
| 204 |
|
| 205 |
assert first is second
|
| 206 |
|
|
|
|
| 211 |
|
| 212 |
|
| 213 |
@pytest.mark.asyncio
|
| 214 |
+
async def test_provider_runtime_cleanup_runs_all_even_if_one_fails() -> None:
|
| 215 |
"""Every provider gets cleanup; cache is cleared even when one raises."""
|
|
|
|
| 216 |
p1 = MagicMock()
|
| 217 |
p1.cleanup = AsyncMock(side_effect=RuntimeError("first"))
|
| 218 |
p2 = MagicMock()
|
| 219 |
p2.cleanup = AsyncMock()
|
| 220 |
+
runtime = ProviderRuntime(_make_settings(), {"a": p1, "b": p2})
|
| 221 |
+
|
| 222 |
with pytest.raises(RuntimeError, match="first"):
|
| 223 |
+
await runtime.cleanup()
|
| 224 |
+
|
| 225 |
p1.cleanup.assert_awaited_once()
|
| 226 |
p2.cleanup.assert_awaited_once()
|
| 227 |
+
assert not runtime.is_cached("a")
|
| 228 |
+
assert not runtime.is_cached("b")
|
| 229 |
|
| 230 |
|
| 231 |
@pytest.mark.asyncio
|
| 232 |
+
async def test_provider_runtime_cleanup_exceptiongroup_on_multiple_failures() -> None:
|
|
|
|
| 233 |
p1 = MagicMock()
|
| 234 |
p1.cleanup = AsyncMock(side_effect=RuntimeError("a"))
|
| 235 |
p2 = MagicMock()
|
| 236 |
p2.cleanup = AsyncMock(side_effect=RuntimeError("b"))
|
| 237 |
+
runtime = ProviderRuntime(_make_settings(), {"x": p1, "y": p2})
|
| 238 |
+
|
| 239 |
with pytest.raises(ExceptionGroup) as exc_info:
|
| 240 |
+
await runtime.cleanup()
|
| 241 |
+
|
| 242 |
assert len(exc_info.value.exceptions) == 2
|
| 243 |
+
assert not runtime.is_cached("x")
|
| 244 |
+
assert not runtime.is_cached("y")
|
|
@@ -561,7 +561,7 @@ wheels = [
|
|
| 561 |
|
| 562 |
[[package]]
|
| 563 |
name = "free-claude-code"
|
| 564 |
-
version = "2.3.
|
| 565 |
source = { editable = "." }
|
| 566 |
dependencies = [
|
| 567 |
{ name = "aiohttp" },
|
|
|
|
| 561 |
|
| 562 |
[[package]]
|
| 563 |
name = "free-claude-code"
|
| 564 |
+
version = "2.3.17"
|
| 565 |
source = { editable = "." }
|
| 566 |
dependencies = [
|
| 567 |
{ name = "aiohttp" },
|