From 92e50099288f8e48e89a14d6bea62105e9d7b5de Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 12:47:09 +0530 Subject: [PATCH 01/14] refactor!: align ferrolabsai with the ai-gateway v1.4.5 contract The SDK read response headers the gateway never emitted (X-Ferro-Provider, X-Ferro-Latency-Ms, X-Ferro-Cost-Usd) and forwarded request fields it never decoded (route_tag/x_route_tag, template_id, template_variables). This realigns every seam with what ai-gateway v1.4.5 actually does: - Headers: X-Request-ID -> trace_id, X-Gateway-Provider -> provider (body `provider` wins), X-Gateway-Overhead-Ms -> gateway_overhead_ms. Metadata is merged into inference bodies only, never /v1/models, probes or /admin/*. - Removed: cost_usd, cache_hit, Usage.provider, latency_ms, route_tag, template_id, template_variables, the dead stream=True _request overload. - Model catalog: ModelInfo is the gateway's EnrichedModelInfo (owned_by, mode, context_window, max_output_tokens, capabilities, status, deprecated; no pricing). retrieve()/list(provider=, capability=)/search() are client-side over one GET /v1/models; retrieve() never hits /v1/models/{id}. - Streaming: Stream/AsyncStream wrappers keep the HTTP response so trace_id and provider are on the stream and every chunk; stream_options is a first-class param; terminal usage chunk is typed; mid-stream error frames raise FerroStreamError(code=...); streams are never retried. - Retries: 408/429/5xx and connect/timeout errors, capped exponential backoff with full jitter, Retry-After honoured (cap 30s). 402 -> FerroBudgetExceededError, 403 -> FerroPermissionError, FerroRateLimitError.retry_after. - Request params: max_completion_tokens, parallel_tool_calls, response_format, seed, stream_options; usage reasoning/cache counters; reasoning_content on message and delta; provider_metadata on ChatCompletion. - New surface: client.responses.{create,retrieve,delete}, capabilities(), health()/ready()/live() (503 bodies returned, not raised), rerank(), moderations.create(); admin.audit.list(), providers/plugins.catalog(), logs.list(api_key_id=), logs.stats(buckets=). - Hygiene: resources type their client via TYPE_CHECKING (no more type: ignore[no-any-return]); __version__ is a constant; mypy no longer pins python_version so each CI leg checks with its own interpreter. - Tests split by area (tests/test_{client,chat,resources,admin}.py) and extended with the good-first-issue cases (#8, #9, #10, #12, #14, #19). --- ferrolabsai/__init__.py | 20 +- ferrolabsai/_version.py | 12 +- ferrolabsai/admin/async_resource.py | 161 ++-- ferrolabsai/admin/resource.py | 271 +++--- ferrolabsai/client.py | 486 +++++----- ferrolabsai/completions/async_resource.py | 125 ++- ferrolabsai/completions/resource.py | 143 ++- ferrolabsai/embeddings/async_resource.py | 20 +- ferrolabsai/embeddings/resource.py | 38 +- ferrolabsai/exceptions/__init__.py | 41 +- ferrolabsai/images/async_resource.py | 36 +- ferrolabsai/images/resource.py | 42 +- ferrolabsai/models/async_resource.py | 30 +- ferrolabsai/models/resource.py | 99 +- ferrolabsai/moderations/__init__.py | 6 + ferrolabsai/moderations/async_resource.py | 20 + ferrolabsai/moderations/resource.py | 28 + ferrolabsai/responses/__init__.py | 6 + ferrolabsai/responses/async_resource.py | 25 + ferrolabsai/responses/resource.py | 37 + ferrolabsai/streaming.py | 125 +++ ferrolabsai/types.py | 109 ++- ferrolabsai/types_responses.py | 40 + pyproject.toml | 1 - tests/conftest.py | 115 +++ tests/test_admin.py | 269 ++++++ tests/test_chat.py | 233 +++++ tests/test_client.py | 431 +++++++++ tests/test_resources.py | 179 ++++ tests/test_sdk.py | 1038 --------------------- 30 files changed, 2431 insertions(+), 1755 deletions(-) create mode 100644 ferrolabsai/moderations/__init__.py create mode 100644 ferrolabsai/moderations/async_resource.py create mode 100644 ferrolabsai/moderations/resource.py create mode 100644 ferrolabsai/responses/__init__.py create mode 100644 ferrolabsai/responses/async_resource.py create mode 100644 ferrolabsai/responses/resource.py create mode 100644 ferrolabsai/streaming.py create mode 100644 ferrolabsai/types_responses.py create mode 100644 tests/conftest.py create mode 100644 tests/test_admin.py create mode 100644 tests/test_chat.py create mode 100644 tests/test_client.py create mode 100644 tests/test_resources.py delete mode 100644 tests/test_sdk.py diff --git a/ferrolabsai/__init__.py b/ferrolabsai/__init__.py index 795b5ca..c76f36a 100644 --- a/ferrolabsai/__init__.py +++ b/ferrolabsai/__init__.py @@ -3,18 +3,20 @@ pip install ferrolabsai +Compatibility: ferrolabsai 0.3.x ↔ ai-gateway ≥ v1.4.0. + Quick start:: from ferrolabsai import FerroClient - client = FerroClient(api_key="sk-ferro-...") + client = FerroClient(api_key="fgw_...") # OpenAI-compatible — just change base_url response = client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], ) - print(response.content) + print(response.content, response.provider, response.trace_id) # Route to ANY provider by model name — Ferro handles it response = client.chat.completions.create( @@ -26,7 +28,7 @@ from ferrolabsai import AsyncFerroClient - async with AsyncFerroClient(api_key="sk-ferro-...") as client: + async with AsyncFerroClient(api_key="fgw_...") as client: response = await client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], @@ -40,7 +42,7 @@ # After — all existing code works unchanged from ferrolabsai import FerroClient - client = FerroClient(api_key="sk-ferro-...") + client = FerroClient(api_key="fgw_...") """ from ._version import __version__ @@ -48,13 +50,16 @@ from .exceptions import ( FerroAPIError, FerroAuthError, + FerroBudgetExceededError, FerroConnectionError, FerroError, FerroNotFoundError, + FerroPermissionError, FerroRateLimitError, FerroServerError, FerroStreamError, ) +from .streaming import AsyncStream, Stream from .types import ( APIKey, ChatCompletion, @@ -69,6 +74,7 @@ ImageData, ImageResponse, ModelInfo, + Response, StreamChoice, StreamDelta, Usage, @@ -83,11 +89,16 @@ "FerroError", "FerroAPIError", "FerroAuthError", + "FerroBudgetExceededError", + "FerroPermissionError", "FerroRateLimitError", "FerroNotFoundError", "FerroServerError", "FerroConnectionError", "FerroStreamError", + # Streaming + "Stream", + "AsyncStream", # Response types "ChatCompletion", "ChatCompletionChunk", @@ -101,6 +112,7 @@ "ImageResponse", "ImageData", "ModelInfo", + "Response", # Admin types "APIKey", "CreatedAPIKey", diff --git a/ferrolabsai/_version.py b/ferrolabsai/_version.py index cc3fe54..1005a44 100644 --- a/ferrolabsai/_version.py +++ b/ferrolabsai/_version.py @@ -1,10 +1,4 @@ -"""Package version helpers.""" +"""Package version. Keep in sync with ``[project].version`` in pyproject.toml +(tests/test_client.py asserts they match).""" -from __future__ import annotations - -try: - from importlib.metadata import version - - __version__ = version("ferrolabsai") -except Exception: - __version__ = "0.2.1" +__version__ = "0.2.1" diff --git a/ferrolabsai/admin/async_resource.py b/ferrolabsai/admin/async_resource.py index 8c424ed..09f3519 100644 --- a/ferrolabsai/admin/async_resource.py +++ b/ferrolabsai/admin/async_resource.py @@ -1,41 +1,44 @@ -"""Async admin resource for /admin/*.""" +"""Async admin resource for /admin/* — see ``admin/resource.py`` for route docs.""" from __future__ import annotations import builtins -from typing import Any +from typing import TYPE_CHECKING, Any from ..types import APIKey, ConfigHistoryEntry, CreatedAPIKey, GatewayConfig +from .resource import items, query + +if TYPE_CHECKING: + from ..client import AsyncFerroClient class AsyncAdmin: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client self.keys = _AsyncKeysResource(client) self.config = _AsyncConfigResource(client) self.logs = _AsyncLogsResource(client) self.providers = _AsyncProvidersResource(client) self.plugins = _AsyncPluginsResource(client) + self.audit = _AsyncAuditResource(client) async def dashboard(self) -> dict[str, Any]: - return await self._client._request("GET", "/admin/dashboard") # type: ignore[no-any-return] + return await self._client._request("GET", "/admin/dashboard") async def health(self) -> dict[str, Any]: - return await self._client._request("GET", "/admin/health") # type: ignore[no-any-return] + return await self._client._request("GET", "/admin/health") class _AsyncKeysResource: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def list(self) -> builtins.list[APIKey]: data = await self._client._request("GET", "/admin/keys") - items = data if isinstance(data, list) else (data.get("keys") or data.get("data") or []) - return [APIKey.from_dict(k) for k in items] + return [APIKey.from_dict(k) for k in items(data, "keys", "data")] async def retrieve(self, key_id: str) -> APIKey: - data = await self._client._request("GET", f"/admin/keys/{key_id}") - return APIKey.from_dict(data) + return APIKey.from_dict(await self._client._request("GET", f"/admin/keys/{key_id}")) async def create( self, @@ -44,13 +47,10 @@ async def create( scopes: builtins.list[str] | None = None, expires_at: str | None = None, ) -> CreatedAPIKey: - body: dict[str, Any] = {"name": name} - if scopes is not None: - body["scopes"] = scopes - if expires_at is not None: - body["expires_at"] = expires_at - data = await self._client._request("POST", "/admin/keys", json=body) - return CreatedAPIKey.from_dict(data) + body = query(name=name, scopes=scopes, expires_at=expires_at) + return CreatedAPIKey.from_dict( + await self._client._request("POST", "/admin/keys", json=body) + ) async def update( self, @@ -61,15 +61,7 @@ async def update( expires_at: str | None = None, active: bool | None = None, ) -> APIKey: - body: dict[str, Any] = {} - if name is not None: - body["name"] = name - if scopes is not None: - body["scopes"] = scopes - if expires_at is not None: - body["expires_at"] = expires_at - if active is not None: - body["active"] = active + body = query(name=name, scopes=scopes, expires_at=expires_at, active=active) data = await self._client._request("PUT", f"/admin/keys/{key_id}", json=body) return APIKey.from_dict(data) @@ -92,42 +84,42 @@ async def usage( active: bool | None = None, since: str | None = None, ) -> dict[str, Any]: - params: dict[str, Any] = {"limit": limit, "offset": offset, "sort": sort} - if active is not None: - params["active"] = "true" if active else "false" - if since is not None: - params["since"] = since - return await self._client._request("GET", "/admin/keys/usage", params=params) # type: ignore[no-any-return] + params = query( + limit=limit, + offset=offset, + sort=sort, + active=None if active is None else ("true" if active else "false"), + since=since, + ) + return await self._client._request("GET", "/admin/keys/usage", params=params) class _AsyncConfigResource: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def get(self) -> GatewayConfig: - data = await self._client._request("GET", "/admin/config") - return GatewayConfig.from_dict(data) + return GatewayConfig.from_dict(await self._client._request("GET", "/admin/config")) async def create(self, config: dict[str, Any]) -> dict[str, Any]: - return await self._client._request("POST", "/admin/config", json=config) # type: ignore[no-any-return] + return await self._client._request("POST", "/admin/config", json=config) async def update(self, config: dict[str, Any]) -> dict[str, Any]: - return await self._client._request("PUT", "/admin/config", json=config) # type: ignore[no-any-return] + return await self._client._request("PUT", "/admin/config", json=config) async def delete(self) -> dict[str, Any]: - return await self._client._request("DELETE", "/admin/config") # type: ignore[no-any-return] + return await self._client._request("DELETE", "/admin/config") - async def history(self) -> list[ConfigHistoryEntry]: + async def history(self) -> builtins.list[ConfigHistoryEntry]: data = await self._client._request("GET", "/admin/config/history") - items = data.get("data") if isinstance(data, dict) else data - return [ConfigHistoryEntry.from_dict(e) for e in (items or [])] + return [ConfigHistoryEntry.from_dict(e) for e in items(data, "data")] async def rollback(self, version: int) -> dict[str, Any]: - return await self._client._request("POST", f"/admin/config/rollback/{version}") # type: ignore[no-any-return] + return await self._client._request("POST", f"/admin/config/rollback/{version}") class _AsyncLogsResource: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def list( @@ -139,30 +131,28 @@ async def list( provider: str | None = None, model: str | None = None, since: str | None = None, + api_key_id: str | None = None, ) -> dict[str, Any]: - params: dict[str, Any] = {"limit": limit, "offset": offset} - if stage is not None: - params["stage"] = stage - if provider is not None: - params["provider"] = provider - if model is not None: - params["model"] = model - if since is not None: - params["since"] = since - return await self._client._request("GET", "/admin/logs", params=params) # type: ignore[no-any-return] + params = query( + limit=limit, + offset=offset, + stage=stage, + provider=provider, + model=model, + since=since, + api_key_id=api_key_id, + ) + return await self._client._request("GET", "/admin/logs", params=params) async def stats( self, *, + buckets: int | None = None, limit: int | None = None, since: str | None = None, ) -> dict[str, Any]: - params: dict[str, Any] = {} - if limit is not None: - params["limit"] = limit - if since is not None: - params["since"] = since - return await self._client._request("GET", "/admin/logs/stats", params=params or None) # type: ignore[no-any-return] + params = query(buckets=buckets, limit=limit, since=since) + return await self._client._request("GET", "/admin/logs/stats", params=params or None) async def delete( self, @@ -170,31 +160,52 @@ async def delete( before: str | None = None, stage: str | None = None, ) -> dict[str, Any]: - params: dict[str, Any] = {} - if before is not None: - params["before"] = before - if stage is not None: - params["stage"] = stage - return await self._client._request("DELETE", "/admin/logs", params=params or None) # type: ignore[no-any-return] + params = query(before=before, stage=stage) + return await self._client._request("DELETE", "/admin/logs", params=params or None) class _AsyncProvidersResource: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def list(self) -> builtins.list[dict[str, Any]]: - data = await self._client._request("GET", "/admin/providers") - if isinstance(data, list): - return data - return data.get("data") or data.get("providers") or [] + return items(await self._client._request("GET", "/admin/providers"), "data", "providers") + + async def catalog(self) -> builtins.list[dict[str, Any]]: + return items(await self._client._request("GET", "/admin/providers/catalog"), "data") class _AsyncPluginsResource: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def list(self) -> builtins.list[dict[str, Any]]: - data = await self._client._request("GET", "/admin/plugins") - if isinstance(data, list): - return data - return data.get("data") or data.get("plugins") or [] + return items(await self._client._request("GET", "/admin/plugins"), "data", "plugins") + + async def catalog(self) -> builtins.list[dict[str, Any]]: + return items(await self._client._request("GET", "/admin/plugins/catalog"), "data") + + +class _AsyncAuditResource: + def __init__(self, client: AsyncFerroClient) -> None: + self._client = client + + async def list( + self, + *, + action: str | None = None, + actor_id: str | None = None, + outcome: str | None = None, + since: str | None = None, + limit: int = 50, + offset: int = 0, + ) -> dict[str, Any]: + params = query( + limit=limit, + offset=offset, + action=action, + actor_id=actor_id, + outcome=outcome, + since=since, + ) + return await self._client._request("GET", "/admin/audit", params=params) diff --git a/ferrolabsai/admin/resource.py b/ferrolabsai/admin/resource.py index 2bd9aad..49fe2dc 100644 --- a/ferrolabsai/admin/resource.py +++ b/ferrolabsai/admin/resource.py @@ -1,10 +1,10 @@ """ Admin resource — manages a Ferro Labs AI Gateway instance via /admin/*. -Routes mirror the OSS gateway admin API defined in -``ai-gateway/internal/admin/handlers.go`` (``Handlers.Routes``): +Routes mirror the OSS gateway admin API defined in the +``ai-gateway/internal/admin/handlers`` package (``Handlers.Routes``): -Read (read-only or admin scope): +Read (read_only or admin scope): GET /admin/dashboard GET /admin/keys GET /admin/keys/usage @@ -12,10 +12,13 @@ GET /admin/logs GET /admin/logs/stats GET /admin/providers + GET /admin/providers/catalog GET /admin/health GET /admin/plugins + GET /admin/plugins/catalog GET /admin/config GET /admin/config/history + GET /admin/audit Write (admin scope only): POST /admin/keys @@ -29,16 +32,16 @@ DELETE /admin/config POST /admin/config/rollback/{version} -These endpoints are available on any self-hosted Ferro Labs AI Gateway -instance. All requests require an API key with admin scope (or read-only -scope for read endpoints), passed via the standard -``Authorization: Bearer ...`` header set on ``FerroClient``. +Dashboard sessions (``/admin/session(s)``) are deliberately not wrapped — SDK +callers hold API keys. All requests use the ``Authorization: Bearer ...`` +header set on ``FerroClient``; a ``read_only`` key gets 403 +``insufficient_scope`` (:class:`FerroPermissionError`) on write routes. """ from __future__ import annotations import builtins -from typing import Any +from typing import TYPE_CHECKING, Any from ..types import ( APIKey, @@ -47,6 +50,25 @@ GatewayConfig, ) +if TYPE_CHECKING: + from ..client import FerroClient + + +def items(data: Any, *keys: str) -> builtins.list[dict[str, Any]]: + """Unwrap a bare array or a ``{"": [...]}`` envelope into a list.""" + if isinstance(data, list): + return data + for key in keys: + found = data.get(key) + if isinstance(found, list): + return found + return [] + + +def query(**params: Any) -> dict[str, Any]: + """Drop ``None`` values so they are not sent as ``?x=None``.""" + return {k: v for k, v in params.items() if v is not None} + class Admin: """ @@ -56,18 +78,19 @@ class Admin: keys — manage API keys (CRUD + revoke + rotate + usage) config — manage the active routing config (get/set/history/rollback) logs — query and prune the request log - providers — list registered provider plugins - plugins — list installed gateway plugins + providers — registered providers (``list``) and the full provider catalog + plugins — installed plugins (``list``) and the built-in plugin catalog + audit — admin audit trail Plus convenience methods on the namespace itself: dashboard() — high-level usage and key counts - health() — gateway health check + health() — gateway health check (admin view) Example:: # Create a key - new_key = client.admin.keys.create(name="backend-service") - print(new_key.key) # full sk-ferro-... — shown ONCE + new_key = client.admin.keys.create(name="backend-service", scopes=["admin"]) + print(new_key.key) # full fgw_... — shown ONCE # Update the active routing config (zero-downtime hot reload) client.admin.config.update({ @@ -83,21 +106,22 @@ class Admin: client.admin.config.rollback(history[-2].version) """ - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client self.keys = _KeysResource(client) self.config = _ConfigResource(client) self.logs = _LogsResource(client) self.providers = _ProvidersResource(client) self.plugins = _PluginsResource(client) + self.audit = _AuditResource(client) def dashboard(self) -> dict[str, Any]: """``GET /admin/dashboard`` — provider/key counts and request log totals.""" - return self._client._request("GET", "/admin/dashboard") # type: ignore[no-any-return] + return self._client._request("GET", "/admin/dashboard") def health(self) -> dict[str, Any]: - """``GET /admin/health`` — gateway health check.""" - return self._client._request("GET", "/admin/health") # type: ignore[no-any-return] + """``GET /admin/health`` — gateway health check (admin scope view).""" + return self._client._request("GET", "/admin/health") # ---------------------------------------------------------------------- @@ -108,19 +132,19 @@ def health(self) -> dict[str, Any]: class _KeysResource: """Manage gateway API keys via ``/admin/keys``.""" - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def list(self) -> builtins.list[APIKey]: - """``GET /admin/keys`` — list all API keys.""" - data = self._client._request("GET", "/admin/keys") - items = data if isinstance(data, list) else (data.get("keys") or data.get("data") or []) - return [APIKey.from_dict(k) for k in items] + """``GET /admin/keys`` — list all API keys (secrets masked).""" + return [ + APIKey.from_dict(k) + for k in items(self._client._request("GET", "/admin/keys"), "keys", "data") + ] def retrieve(self, key_id: str) -> APIKey: """``GET /admin/keys/{id}`` — fetch metadata for one key (key value is masked).""" - data = self._client._request("GET", f"/admin/keys/{key_id}") - return APIKey.from_dict(data) + return APIKey.from_dict(self._client._request("GET", f"/admin/keys/{key_id}")) def create( self, @@ -132,21 +156,16 @@ def create( """ ``POST /admin/keys`` — create a new API key. - The full key value (``sk-ferro-...``) is only returned in this response. + The full key value (``fgw_...``) is only returned in this response. Store it securely — it cannot be retrieved again. Args: name: Human-readable label for this key. - scopes: List of scopes (e.g. ``["admin"]``, ``["read-only"]``). + scopes: ``["admin"]`` or ``["read_only"]`` (unknown scope → 400 ``invalid_scope``). expires_at: RFC3339 expiry timestamp. ``None`` = never expires. """ - body: dict[str, Any] = {"name": name} - if scopes is not None: - body["scopes"] = scopes - if expires_at is not None: - body["expires_at"] = expires_at - data = self._client._request("POST", "/admin/keys", json=body) - return CreatedAPIKey.from_dict(data) + body = query(name=name, scopes=scopes, expires_at=expires_at) + return CreatedAPIKey.from_dict(self._client._request("POST", "/admin/keys", json=body)) def update( self, @@ -158,17 +177,8 @@ def update( active: bool | None = None, ) -> APIKey: """``PUT /admin/keys/{id}`` — update key metadata.""" - body: dict[str, Any] = {} - if name is not None: - body["name"] = name - if scopes is not None: - body["scopes"] = scopes - if expires_at is not None: - body["expires_at"] = expires_at - if active is not None: - body["active"] = active - data = self._client._request("PUT", f"/admin/keys/{key_id}", json=body) - return APIKey.from_dict(data) + body = query(name=name, scopes=scopes, expires_at=expires_at, active=active) + return APIKey.from_dict(self._client._request("PUT", f"/admin/keys/{key_id}", json=body)) def delete(self, key_id: str) -> None: """``DELETE /admin/keys/{id}`` — permanently delete a key.""" @@ -189,8 +199,9 @@ def rotate(self, key_id: str) -> CreatedAPIKey: Returns the new key. Store it securely — shown only once. """ - data = self._client._request("POST", f"/admin/keys/{key_id}/rotate") - return CreatedAPIKey.from_dict(data) + return CreatedAPIKey.from_dict( + self._client._request("POST", f"/admin/keys/{key_id}/rotate") + ) def usage( self, @@ -213,12 +224,14 @@ def usage( Returns the raw response: ``{data, summary, filters}``. """ - params: dict[str, Any] = {"limit": limit, "offset": offset, "sort": sort} - if active is not None: - params["active"] = "true" if active else "false" - if since is not None: - params["since"] = since - return self._client._request("GET", "/admin/keys/usage", params=params) # type: ignore[no-any-return] + params = query( + limit=limit, + offset=offset, + sort=sort, + active=None if active is None else ("true" if active else "false"), + since=since, + ) + return self._client._request("GET", "/admin/keys/usage", params=params) # ---------------------------------------------------------------------- @@ -232,49 +245,38 @@ class _ConfigResource: The OSS gateway has a *single* active config (not a multi-config registry). Use ``history()`` to inspect previous versions and ``rollback(version)`` to - revert. Updates are zero-downtime hot reloads. + revert. Updates are zero-downtime hot reloads. Note that ``get()`` masks + secrets and redacts free-form map keys, so its body does not round-trip + unchanged into ``update()``. """ - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def get(self) -> GatewayConfig: """``GET /admin/config`` — fetch the currently active config.""" - data = self._client._request("GET", "/admin/config") - return GatewayConfig.from_dict(data) + return GatewayConfig.from_dict(self._client._request("GET", "/admin/config")) def create(self, config: dict[str, Any]) -> dict[str, Any]: - """ - ``POST /admin/config`` — install a new config (status 201). - - ``config`` is the raw routing-config dict (``strategy``, ``targets``, - ``plugins``, ``aliases``, etc.). - """ - return self._client._request("POST", "/admin/config", json=config) # type: ignore[no-any-return] + """``POST /admin/config`` — install a new config (status 201). Unknown keys → 400.""" + return self._client._request("POST", "/admin/config", json=config) def update(self, config: dict[str, Any]) -> dict[str, Any]: - """ - ``PUT /admin/config`` — replace the active config (status 200). - - Hot-reloads in place — no gateway restart required. In-flight - requests complete with the previous config; the next request after - the update uses the new one. - """ - return self._client._request("PUT", "/admin/config", json=config) # type: ignore[no-any-return] + """``PUT /admin/config`` — replace the active config (hot reload, no restart).""" + return self._client._request("PUT", "/admin/config", json=config) def delete(self) -> dict[str, Any]: """``DELETE /admin/config`` — reset the active config to its default.""" - return self._client._request("DELETE", "/admin/config") # type: ignore[no-any-return] + return self._client._request("DELETE", "/admin/config") - def history(self) -> list[ConfigHistoryEntry]: + def history(self) -> builtins.list[ConfigHistoryEntry]: """``GET /admin/config/history`` — list all prior config versions.""" data = self._client._request("GET", "/admin/config/history") - items = data.get("data") if isinstance(data, dict) else data - return [ConfigHistoryEntry.from_dict(e) for e in (items or [])] + return [ConfigHistoryEntry.from_dict(e) for e in items(data, "data")] def rollback(self, version: int) -> dict[str, Any]: """``POST /admin/config/rollback/{version}`` — revert to a prior version.""" - return self._client._request("POST", f"/admin/config/rollback/{version}") # type: ignore[no-any-return] + return self._client._request("POST", f"/admin/config/rollback/{version}") # ---------------------------------------------------------------------- @@ -286,15 +288,11 @@ class _LogsResource: """ Query the gateway request log via ``/admin/logs``. - Replaces what was previously ``client.admin.usage.requests()`` — the OSS - gateway exposes raw per-request log entries (with trace IDs, latency, - tokens, cost, and provider routing decisions) at ``/admin/logs``. - - Note: request log storage must be enabled in the gateway (via the - ``logger`` plugin). Endpoints return HTTP 501 if it isn't. + Requires a request-log store (``REQUEST_LOG_STORE_BACKEND=sqlite|postgres`` + on the gateway); the endpoints answer 501 without one. """ - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def list( @@ -306,42 +304,43 @@ def list( provider: str | None = None, model: str | None = None, since: str | None = None, + api_key_id: str | None = None, ) -> dict[str, Any]: """ - ``GET /admin/logs`` — paginated request log entries. + ``GET /admin/logs`` — paginated request log entries (one row per request). Args: limit: Max entries to return (server caps at 200). offset: Pagination offset. - stage: Filter by lifecycle stage (e.g. ``"on_error"``). + stage: Default lists terminal rows only; ``"all"`` includes every + lifecycle stage (``before_request`` / ``after_request`` / ``on_error``). provider: Filter by provider name. model: Filter by model id. since: RFC3339 timestamp — only entries at or after this time. + api_key_id: Filter by the calling key's id (``"none"`` = master key / no key). """ - params: dict[str, Any] = {"limit": limit, "offset": offset} - if stage is not None: - params["stage"] = stage - if provider is not None: - params["provider"] = provider - if model is not None: - params["model"] = model - if since is not None: - params["since"] = since - return self._client._request("GET", "/admin/logs", params=params) # type: ignore[no-any-return] + params = query( + limit=limit, + offset=offset, + stage=stage, + provider=provider, + model=model, + since=since, + api_key_id=api_key_id, + ) + return self._client._request("GET", "/admin/logs", params=params) def stats( self, *, + buckets: int | None = None, limit: int | None = None, since: str | None = None, ) -> dict[str, Any]: - """``GET /admin/logs/stats`` — aggregate counts, latency, and cost.""" - params: dict[str, Any] = {} - if limit is not None: - params["limit"] = limit - if since is not None: - params["since"] = since - return self._client._request("GET", "/admin/logs/stats", params=params or None) # type: ignore[no-any-return] + """``GET /admin/logs/stats`` — totals, latency/TTFT percentiles, per-provider + and per-model cost; ``buckets=N`` adds an ``N``-point time series.""" + params = query(buckets=buckets, limit=limit, since=since) + return self._client._request("GET", "/admin/logs/stats", params=params or None) def delete( self, @@ -356,42 +355,70 @@ def delete( before: RFC3339 timestamp — delete entries strictly before this. stage: Restrict deletion to a single lifecycle stage. """ - params: dict[str, Any] = {} - if before is not None: - params["before"] = before - if stage is not None: - params["stage"] = stage - return self._client._request("DELETE", "/admin/logs", params=params or None) # type: ignore[no-any-return] + params = query(before=before, stage=stage) + return self._client._request("DELETE", "/admin/logs", params=params or None) # ---------------------------------------------------------------------- -# Providers / Plugins +# Providers / Plugins / Audit # ---------------------------------------------------------------------- class _ProvidersResource: - """List provider plugins via ``/admin/providers``.""" + """Providers via ``/admin/providers``.""" - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def list(self) -> builtins.list[dict[str, Any]]: """``GET /admin/providers`` — registered providers and their availability.""" - data = self._client._request("GET", "/admin/providers") - if isinstance(data, list): - return data - return data.get("data") or data.get("providers") or [] + return items(self._client._request("GET", "/admin/providers"), "data", "providers") + + def catalog(self) -> builtins.list[dict[str, Any]]: + """``GET /admin/providers/catalog`` — every provider the build knows: + ``{id, registered, catalog_models}``.""" + return items(self._client._request("GET", "/admin/providers/catalog"), "data") class _PluginsResource: - """List installed plugins via ``/admin/plugins``.""" + """Plugins via ``/admin/plugins``.""" - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def list(self) -> builtins.list[dict[str, Any]]: - """``GET /admin/plugins`` — gateway plugins (cache, logger, ratelimit, ...).""" - data = self._client._request("GET", "/admin/plugins") - if isinstance(data, list): - return data - return data.get("data") or data.get("plugins") or [] + """``GET /admin/plugins`` — configured gateway plugins.""" + return items(self._client._request("GET", "/admin/plugins"), "data", "plugins") + + def catalog(self) -> builtins.list[dict[str, Any]]: + """``GET /admin/plugins/catalog`` — built-in plugins available to configure.""" + return items(self._client._request("GET", "/admin/plugins/catalog"), "data") + + +class _AuditResource: + """Admin audit trail via ``/admin/audit``.""" + + def __init__(self, client: FerroClient) -> None: + self._client = client + + def list( + self, + *, + action: str | None = None, + actor_id: str | None = None, + outcome: str | None = None, + since: str | None = None, + limit: int = 50, + offset: int = 0, + ) -> dict[str, Any]: + """``GET /admin/audit`` — ``{data, summary, filters}`` of admin actions + (key/config writes, denied attempts).""" + params = query( + limit=limit, + offset=offset, + action=action, + actor_id=actor_id, + outcome=outcome, + since=since, + ) + return self._client._request("GET", "/admin/audit", params=params) diff --git a/ferrolabsai/client.py b/ferrolabsai/client.py index 017bd06..8b9315a 100644 --- a/ferrolabsai/client.py +++ b/ferrolabsai/client.py @@ -1,31 +1,65 @@ """ FerroClient — drop-in replacement for the OpenAI Python client. -Points at any self-hosted Ferro Labs AI Gateway instance. +Points at any self-hosted Ferro Labs AI Gateway instance (ai-gateway ≥ v1.4.0). """ from __future__ import annotations import asyncio import os +import random import time -from collections.abc import Iterator -from typing import Any, Literal, cast, overload +from typing import Any, NoReturn, cast import httpx from ._version import __version__ +from .admin.async_resource import AsyncAdmin from .admin.resource import Admin +from .completions.async_resource import AsyncCompletions from .completions.resource import Completions +from .embeddings.async_resource import AsyncEmbeddings from .embeddings.resource import Embeddings -from .exceptions import FerroAuthError, FerroConnectionError +from .exceptions import ( + FerroAPIError, + FerroAuthError, + FerroBudgetExceededError, + FerroConnectionError, + FerroNotFoundError, + FerroPermissionError, + FerroRateLimitError, + FerroServerError, +) +from .images.async_resource import AsyncImages from .images.resource import Images +from .models.async_resource import AsyncModels from .models.resource import Models +from .moderations.async_resource import AsyncModerations +from .moderations.resource import Moderations +from .responses.async_resource import AsyncResponses +from .responses.resource import Responses DEFAULT_BASE_URL: str = "http://localhost:8080" DEFAULT_TIMEOUT: float = 120.0 DEFAULT_MAX_RETRIES: int = 2 DEFAULT_RETRY_BACKOFF_BASE: float = 0.5 DEFAULT_RETRY_BACKOFF_MAX: float = 8.0 +RETRY_AFTER_MAX: float = 30.0 # same cap the gateway applies to its own upstream retries +RETRYABLE_STATUSES: frozenset[int] = frozenset({408, 429}) # plus every 5xx + +# Only inference bodies get header metadata (trace_id, provider, gateway_overhead_ms) +# merged in; catalog, probe, and admin bodies are returned untouched. +_INFERENCE_PREFIXES = ( + "/v1/chat/completions", + "/v1/completions", + "/v1/embeddings", + "/v1/images/generations", + "/v1/responses", + "/v1/rerank", + "/v1/moderations", +) +# /health and /readyz answer 503 with a JSON body when degraded; that is an answer, not an error. +_PROBE_STATUSES = (503,) def _validate_max_retries(max_retries: object) -> int: @@ -36,25 +70,69 @@ def _validate_max_retries(max_retries: object) -> int: return max_retries -def _retry_delay(attempt: int) -> float: - """Return capped exponential retry delay for a 1-based retry attempt.""" - delay = DEFAULT_RETRY_BACKOFF_BASE * float(2 ** max(attempt - 1, 0)) - return min(delay, DEFAULT_RETRY_BACKOFF_MAX) +def _resolve_credentials(api_key: str | None, base_url: str | None) -> tuple[str, str]: + key = api_key or os.environ.get("FERRO_API_KEY") or os.environ.get("OPENAI_API_KEY") + if not key: + raise FerroAuthError("No API key provided. Pass api_key=... or set FERRO_API_KEY env var.") + url = (base_url or os.environ.get("FERRO_BASE_URL") or DEFAULT_BASE_URL).rstrip("/") + return key, url + + +def _default_headers(api_key: str, extra: dict[str, str] | None) -> dict[str, str]: + return { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + "User-Agent": f"ferrolabsai-python/{__version__}", + **(extra or {}), + } + + +# ------------------------------------------------------------------ +# Retry policy (shared by sync and async; streaming is never retried) +# ------------------------------------------------------------------ + + +def _is_retryable(status: int) -> bool: + return status in RETRYABLE_STATUSES or status >= 500 + + +def _retry_after_seconds(response: httpx.Response) -> float | None: + """``Retry-After`` in seconds, or None when absent / not a number (HTTP-date form).""" + value = response.headers.get("retry-after") + try: + return float(value) if value else None + except ValueError: + return None + + +def _retry_delay(attempt: int, retry_after: float | None = None) -> float: + """Delay before the 1-based retry ``attempt``: honour Retry-After (capped), else + capped exponential backoff with full jitter.""" + if retry_after is not None: + return min(retry_after, RETRY_AFTER_MAX) + delay = min(DEFAULT_RETRY_BACKOFF_BASE * float(2 ** (attempt - 1)), DEFAULT_RETRY_BACKOFF_MAX) + return random.uniform(0, delay) + + +def _connection_error(exc: Exception, base_url: str, timeout: float) -> FerroConnectionError: + if isinstance(exc, httpx.TimeoutException): + return FerroConnectionError(f"Request timed out after {timeout}s: {exc}") + return FerroConnectionError(f"Cannot reach {base_url}. Is the gateway running? ({exc})") class FerroClient: """ Primary client for Ferro Labs AI Gateway. - Drop-in compatible with the OpenAI Python SDK for chat/embeddings/images. - Adds gateway-specific features: admin API, model catalog, prompt templates, - cost tracking, and guardrail management. + Drop-in compatible with the OpenAI Python SDK for chat/embeddings/images, + plus the gateway's own surface: model catalog, capabilities, health probes, + Responses API, rerank, moderations, and the ``/admin/*`` API. Usage:: from ferrolabsai import FerroClient - client = FerroClient(api_key="sk-ferro-...") + client = FerroClient(api_key="fgw_...") # OpenAI-compatible response = client.chat.completions.create( @@ -78,35 +156,17 @@ def __init__( default_headers: dict[str, str] | None = None, http_client: httpx.Client | None = None, ) -> None: - self.api_key = ( - api_key or os.environ.get("FERRO_API_KEY") or os.environ.get("OPENAI_API_KEY") - ) - if not self.api_key: - raise FerroAuthError( - "No API key provided. Pass api_key=... or set FERRO_API_KEY env var." - ) - - self.base_url = (base_url or os.environ.get("FERRO_BASE_URL") or DEFAULT_BASE_URL).rstrip( - "/" - ) + self.api_key, self.base_url = _resolve_credentials(api_key, base_url) self.timeout = timeout self.max_retries = _validate_max_retries(max_retries) - - self._default_headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json", - "User-Agent": f"ferrolabsai-python/{self._version()}", - **(default_headers or {}), - } + self._default_headers = _default_headers(self.api_key, default_headers) if http_client is not None: http_client.headers.update(self._default_headers) self._http = http_client else: self._http = httpx.Client( - base_url=self.base_url, - timeout=timeout, - headers=self._default_headers, + base_url=self.base_url, timeout=timeout, headers=self._default_headers ) # Resource namespaces — mirrors OpenAI SDK layout @@ -114,32 +174,52 @@ def __init__( self.embeddings = Embeddings(self) self.images = Images(self) self.models = Models(self) + self.responses = Responses(self) + self.moderations = Moderations(self) self.admin = Admin(self) # ------------------------------------------------------------------ - # Internal HTTP helpers used by all resource classes + # Gateway-level endpoints # ------------------------------------------------------------------ - @overload - def _request( - self, - method: str, - path: str, - *, - json: Any | None = ..., - params: dict[str, Any] | None = ..., - stream: Literal[False] = ..., - ) -> dict[str, Any]: ... - @overload - def _request( + def health(self) -> dict[str, Any]: + """``GET /health`` — ``{status, version, commit, built, providers}``; 503 body when + degraded (e.g. ``{"status": "no_providers"}``) is returned, not raised.""" + return self._request("GET", "/health", allow=_PROBE_STATUSES) + + def ready(self) -> dict[str, Any]: + """``GET /readyz`` — 200 ``{status: "ready", providers, targets, mcp_servers}`` or the + 503 ``{status: "not_ready", reason}`` body.""" + return self._request("GET", "/readyz", allow=_PROBE_STATUSES) + + def live(self) -> dict[str, Any]: + """``GET /livez`` — ``{"status": "ok"}``.""" + return self._request("GET", "/livez") + + def capabilities(self) -> dict[str, Any]: + """``GET /v1/capabilities`` — per-provider parameter support matrix + (``forward`` / ``translate`` / ``unsupported``) for the configured targets.""" + return self._request("GET", "/v1/capabilities") + + def rerank( self, - method: str, - path: str, *, - json: Any | None = ..., - params: dict[str, Any] | None = ..., - stream: Literal[True], - ) -> httpx.Response: ... + model: str, + query: str, + documents: list[str], + top_n: int | None = None, + **kwargs: Any, + ) -> dict[str, Any]: + """``POST /v1/rerank`` (Cohere v2 shape) — returns the provider's JSON.""" + body: dict[str, Any] = {"model": model, "query": query, "documents": documents, **kwargs} + if top_n is not None: + body["top_n"] = top_n + return self._request("POST", "/v1/rerank", json=body) + + # ------------------------------------------------------------------ + # Internal HTTP helpers used by all resource classes + # ------------------------------------------------------------------ + def _request( self, method: str, @@ -147,55 +227,39 @@ def _request( *, json: Any | None = None, params: dict[str, Any] | None = None, - stream: bool = False, - ) -> dict[str, Any] | httpx.Response: - url = path if path.startswith("http") else path + allow: tuple[int, ...] = (), + ) -> dict[str, Any]: attempt = 0 - last_exc: Exception | None = None - - while attempt <= self.max_retries: - try: - req = self._http.build_request( - method, - url, - json=json, - params=params, - ) - response = self._http.send(req, stream=stream) - response.raise_for_status() - if stream: - return response - if response.status_code == 204 or not response.content: - return {} - return _with_response_metadata(cast("dict[str, Any]", response.json()), response) - except httpx.HTTPStatusError as e: - _raise_api_error(e) - except httpx.ConnectError as e: - last_exc = FerroConnectionError( - f"Cannot reach {self.base_url}. Is the gateway running? ({e})" - ) - attempt += 1 - if attempt <= self.max_retries: - time.sleep(_retry_delay(attempt)) - except httpx.TimeoutException as e: - last_exc = FerroConnectionError(f"Request timed out after {self.timeout}s: {e}") - attempt += 1 - if attempt <= self.max_retries: - time.sleep(_retry_delay(attempt)) - - raise last_exc # type: ignore - - def _stream_request(self, path: str, json: Any) -> Iterator[str]: - """Yields raw SSE lines for streaming completions.""" - with self._http.stream("POST", path, json=json) as response: + while True: try: - response.raise_for_status() + response = self._http.request(method, path, json=json, params=params) + if response.status_code not in allow: + response.raise_for_status() + return _parse_body(response, path) except httpx.HTTPStatusError as e: - response.read() - _raise_api_error(e) - for line in response.iter_lines(): - if line: - yield line + if attempt >= self.max_retries or not _is_retryable(e.response.status_code): + _raise_api_error(e) + delay = _retry_delay(attempt + 1, _retry_after_seconds(e.response)) + except (httpx.ConnectError, httpx.TimeoutException) as e: + if attempt >= self.max_retries: + raise _connection_error(e, self.base_url, self.timeout) from e + delay = _retry_delay(attempt + 1) + attempt += 1 + time.sleep(delay) + + def _open_stream(self, path: str, json: Any) -> httpx.Response: + """POST and return the live response for SSE consumption (no retries).""" + request = self._http.build_request( + "POST", path, json=json, headers={"Accept": "text/event-stream"} + ) + response = self._http.send(request, stream=True) + try: + response.raise_for_status() + except httpx.HTTPStatusError as e: + response.read() + response.close() + _raise_api_error(e) + return response @staticmethod def _version() -> str: @@ -232,7 +296,7 @@ class AsyncFerroClient: from ferrolabsai import AsyncFerroClient - async with AsyncFerroClient(api_key="sk-ferro-...") as client: + async with AsyncFerroClient(api_key="fgw_...") as client: response = await client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], @@ -248,48 +312,58 @@ def __init__( default_headers: dict[str, str] | None = None, http_client: httpx.AsyncClient | None = None, ) -> None: - self.api_key = ( - api_key or os.environ.get("FERRO_API_KEY") or os.environ.get("OPENAI_API_KEY") - ) - if not self.api_key: - raise FerroAuthError( - "No API key provided. Pass api_key=... or set FERRO_API_KEY env var." - ) - - self.base_url = (base_url or os.environ.get("FERRO_BASE_URL") or DEFAULT_BASE_URL).rstrip( - "/" - ) + self.api_key, self.base_url = _resolve_credentials(api_key, base_url) self.timeout = timeout self.max_retries = _validate_max_retries(max_retries) - - self._default_headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json", - "User-Agent": f"ferrolabsai-python/{FerroClient._version()}", - **(default_headers or {}), - } + self._default_headers = _default_headers(self.api_key, default_headers) if http_client is not None: http_client.headers.update(self._default_headers) self._http = http_client else: self._http = httpx.AsyncClient( - base_url=self.base_url, - timeout=timeout, - headers=self._default_headers, + base_url=self.base_url, timeout=timeout, headers=self._default_headers ) - from .admin.async_resource import AsyncAdmin - from .embeddings.async_resource import AsyncEmbeddings - from .images.async_resource import AsyncImages - from .models.async_resource import AsyncModels - self.chat = _AsyncChatNamespace(self) self.embeddings = AsyncEmbeddings(self) self.images = AsyncImages(self) self.models = AsyncModels(self) + self.responses = AsyncResponses(self) + self.moderations = AsyncModerations(self) self.admin = AsyncAdmin(self) + async def health(self) -> dict[str, Any]: + """See :meth:`FerroClient.health`.""" + return await self._request("GET", "/health", allow=_PROBE_STATUSES) + + async def ready(self) -> dict[str, Any]: + """See :meth:`FerroClient.ready`.""" + return await self._request("GET", "/readyz", allow=_PROBE_STATUSES) + + async def live(self) -> dict[str, Any]: + """See :meth:`FerroClient.live`.""" + return await self._request("GET", "/livez") + + async def capabilities(self) -> dict[str, Any]: + """See :meth:`FerroClient.capabilities`.""" + return await self._request("GET", "/v1/capabilities") + + async def rerank( + self, + *, + model: str, + query: str, + documents: list[str], + top_n: int | None = None, + **kwargs: Any, + ) -> dict[str, Any]: + """See :meth:`FerroClient.rerank`.""" + body: dict[str, Any] = {"model": model, "query": query, "documents": documents, **kwargs} + if top_n is not None: + body["top_n"] = top_n + return await self._request("POST", "/v1/rerank", json=body) + async def _request( self, method: str, @@ -297,31 +371,38 @@ async def _request( *, json: Any | None = None, params: dict[str, Any] | None = None, + allow: tuple[int, ...] = (), ) -> dict[str, Any]: attempt = 0 - last_exc: Exception | None = None - while attempt <= self.max_retries: + while True: try: response = await self._http.request(method, path, json=json, params=params) - response.raise_for_status() - if response.status_code == 204 or not response.content: - return {} - return _with_response_metadata(cast("dict[str, Any]", response.json()), response) + if response.status_code not in allow: + response.raise_for_status() + return _parse_body(response, path) except httpx.HTTPStatusError as e: - _raise_api_error(e) - except httpx.ConnectError as e: - last_exc = FerroConnectionError( - f"Cannot reach {self.base_url}. Is the gateway running? ({e})" - ) - attempt += 1 - if attempt <= self.max_retries: - await asyncio.sleep(_retry_delay(attempt)) - except httpx.TimeoutException as e: - last_exc = FerroConnectionError(f"Request timed out after {self.timeout}s: {e}") - attempt += 1 - if attempt <= self.max_retries: - await asyncio.sleep(_retry_delay(attempt)) - raise last_exc # type: ignore + if attempt >= self.max_retries or not _is_retryable(e.response.status_code): + _raise_api_error(e) + delay = _retry_delay(attempt + 1, _retry_after_seconds(e.response)) + except (httpx.ConnectError, httpx.TimeoutException) as e: + if attempt >= self.max_retries: + raise _connection_error(e, self.base_url, self.timeout) from e + delay = _retry_delay(attempt + 1) + attempt += 1 + await asyncio.sleep(delay) + + async def _open_stream(self, path: str, json: Any) -> httpx.Response: + request = self._http.build_request( + "POST", path, json=json, headers={"Accept": "text/event-stream"} + ) + response = await self._http.send(request, stream=True) + try: + response.raise_for_status() + except httpx.HTTPStatusError as e: + await response.aread() + await response.aclose() + _raise_api_error(e) + return response async def close(self) -> None: await self._http.aclose() @@ -335,8 +416,6 @@ async def __aexit__(self, *_: Any) -> None: class _AsyncChatNamespace: def __init__(self, client: AsyncFerroClient) -> None: - from .completions.async_resource import AsyncCompletions - self.completions = AsyncCompletions(client) @@ -345,91 +424,70 @@ def __init__(self, client: AsyncFerroClient) -> None: # ------------------------------------------------------------------ -def _with_response_metadata(data: dict[str, Any], response: httpx.Response) -> dict[str, Any]: - """Copy gateway metadata headers into parsed response bodies. - - Successful SDK calls return dataclasses, so header-only metadata such as - X-Request-ID must be preserved before resource classes construct them. - Body fields stay authoritative when both sources are present. - """ - trace_id = ( - response.headers.get("x-request-id") - or response.headers.get("x-trace-id") - or response.headers.get("x-ferro-request-id") - ) - if trace_id and "trace_id" not in data and "x_ferro_trace_id" not in data: - data["trace_id"] = trace_id - - provider = response.headers.get("x-ferro-provider") - if provider and "provider" not in data and "x_ferro_provider" not in data: - data["provider"] = provider - usage = data.get("usage") - if isinstance(usage, dict) and "provider" not in usage: - usage["provider"] = provider - - latency_ms = _header_int(response.headers.get("x-ferro-latency-ms")) - if latency_ms is not None and "latency_ms" not in data and "x_ferro_latency_ms" not in data: - data["x_ferro_latency_ms"] = latency_ms - - cost_usd = _header_float(response.headers.get("x-ferro-cost-usd")) - if cost_usd is not None: - usage = data.get("usage") - if not isinstance(usage, dict): - usage = {} - data["usage"] = usage - if "cost_usd" not in usage: - usage["cost_usd"] = cost_usd - +def _parse_body(response: httpx.Response, path: str) -> dict[str, Any]: + if response.status_code == 204 or not response.content: + return {} + data = cast("dict[str, Any]", response.json()) + if path.startswith(_INFERENCE_PREFIXES): + return _with_response_metadata(data, response) return data -def _header_int(value: str | None) -> int | None: - if value is None or value == "": - return None - try: - return int(float(value)) - except ValueError: - return None - - def _header_float(value: str | None) -> float | None: - if value is None or value == "": - return None try: - return float(value) + return float(value) if value else None except ValueError: return None -def _raise_api_error(e: httpx.HTTPStatusError) -> None: - from .exceptions import ( - FerroAPIError, - FerroAuthError, - FerroNotFoundError, - FerroRateLimitError, - FerroServerError, - ) +def _with_response_metadata(data: dict[str, Any], response: httpx.Response) -> dict[str, Any]: + """Merge the gateway's response headers into an inference body. + ``X-Request-ID`` → ``trace_id``; ``X-Gateway-Provider`` → ``provider`` (the + non-streaming chat body carries ``provider`` itself, which wins); + ``X-Gateway-Overhead-Ms`` → ``gateway_overhead_ms``. Returns a new dict. + """ + extra: dict[str, Any] = {} + trace_id = response.headers.get("x-request-id") + if trace_id and "trace_id" not in data: + extra["trace_id"] = trace_id + provider = response.headers.get("x-gateway-provider") + if provider and "provider" not in data: + extra["provider"] = provider + overhead = _header_float(response.headers.get("x-gateway-overhead-ms")) + if overhead is not None and "gateway_overhead_ms" not in data: + extra["gateway_overhead_ms"] = overhead + return {**data, **extra} if extra else data + + +def _raise_api_error(e: httpx.HTTPStatusError) -> NoReturn: + """Map the gateway error envelope ``{"error": {message, type, code}}`` to an exception.""" status = e.response.status_code - request_id = ( - e.response.headers.get("x-request-id") - or e.response.headers.get("x-ferro-request-id") - ) + request_id = e.response.headers.get("x-request-id") try: body = e.response.json() - message = body.get("error", {}).get("message") or body.get("message") or str(e) - code = body.get("error", {}).get("code") or body.get("code") + error = body.get("error") if isinstance(body, dict) else None + error = error if isinstance(error, dict) else {} + message = error.get("message") or body.get("message") or str(e) + code = error.get("code") or body.get("code") request_id = request_id or body.get("request_id") or body.get("trace_id") except Exception: message = e.response.text or str(e) code = None + kwargs: dict[str, Any] = {"status_code": status, "code": code, "request_id": request_id} if status == 401: - raise FerroAuthError(message, request_id=request_id) from e - if status == 429: - raise FerroRateLimitError(message, request_id=request_id) from e + raise FerroAuthError(message, **kwargs) from e + if status == 402: + raise FerroBudgetExceededError(message, **kwargs) from e + if status == 403: + raise FerroPermissionError(message, **kwargs) from e if status == 404: - raise FerroNotFoundError(message, request_id=request_id) from e + raise FerroNotFoundError(message, **kwargs) from e + if status == 429: + raise FerroRateLimitError( + message, retry_after=_retry_after_seconds(e.response), **kwargs + ) from e if status >= 500: - raise FerroServerError(message, status_code=status, request_id=request_id) from e - raise FerroAPIError(message, status_code=status, code=code, request_id=request_id) from e + raise FerroServerError(message, **kwargs) from e + raise FerroAPIError(message, **kwargs) from e diff --git a/ferrolabsai/completions/async_resource.py b/ferrolabsai/completions/async_resource.py index 090e4a3..247e79e 100644 --- a/ferrolabsai/completions/async_resource.py +++ b/ferrolabsai/completions/async_resource.py @@ -2,20 +2,40 @@ from __future__ import annotations -import json -from collections.abc import AsyncIterator -from typing import Any +from typing import TYPE_CHECKING, Any, Literal, overload -import httpx +from ..streaming import AsyncStream +from ..types import ChatCompletion +from .resource import PATH, build_body -from ..exceptions import FerroStreamError -from ..types import ChatCompletion, ChatCompletionChunk +if TYPE_CHECKING: + from ..client import AsyncFerroClient class AsyncCompletions: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client + @overload + async def create( + self, + *, + model: str, + messages: list[dict[str, Any]], + stream: Literal[False] = False, + **kwargs: Any, + ) -> ChatCompletion: ... + + @overload + async def create( + self, + *, + model: str, + messages: list[dict[str, Any]], + stream: Literal[True], + **kwargs: Any, + ) -> AsyncStream: ... + async def create( self, *, @@ -24,71 +44,48 @@ async def create( stream: bool = False, temperature: float | None = None, max_tokens: int | None = None, + max_completion_tokens: int | None = None, top_p: float | None = None, frequency_penalty: float | None = None, presence_penalty: float | None = None, stop: str | list[str] | None = None, + seed: int | None = None, tools: list[dict[str, Any]] | None = None, tool_choice: Any | None = None, - template_id: str | None = None, - template_variables: dict[str, Any] | None = None, - route_tag: str | None = None, + parallel_tool_calls: bool | None = None, + response_format: dict[str, Any] | None = None, + stream_options: dict[str, Any] | None = None, user: str | None = None, **kwargs: Any, - ) -> ChatCompletion | AsyncIterator[ChatCompletionChunk]: - body: dict[str, Any] = {"model": model, "messages": messages, "stream": stream} + ) -> ChatCompletion | AsyncStream: + """Async variant of :meth:`ferrolabsai.completions.resource.Completions.create`. - if temperature is not None: - body["temperature"] = temperature - if max_tokens is not None: - body["max_tokens"] = max_tokens - if top_p is not None: - body["top_p"] = top_p - if frequency_penalty is not None: - body["frequency_penalty"] = frequency_penalty - if presence_penalty is not None: - body["presence_penalty"] = presence_penalty - if stop is not None: - body["stop"] = stop - if tools is not None: - body["tools"] = tools - if tool_choice is not None: - body["tool_choice"] = tool_choice - if user is not None: - body["user"] = user - if template_id is not None: - body["template_id"] = template_id - if template_variables is not None: - body["template_variables"] = template_variables - if route_tag is not None: - body["x_route_tag"] = route_tag - - body.update(kwargs) + With ``stream=True`` the awaited result is an :class:`~ferrolabsai.AsyncStream`:: + stream = await client.chat.completions.create(..., stream=True) + async for chunk in stream: + ... + """ + body = build_body( + model, + messages, + stream, + temperature=temperature, + max_tokens=max_tokens, + max_completion_tokens=max_completion_tokens, + top_p=top_p, + frequency_penalty=frequency_penalty, + presence_penalty=presence_penalty, + stop=stop, + seed=seed, + tools=tools, + tool_choice=tool_choice, + parallel_tool_calls=parallel_tool_calls, + response_format=response_format, + stream_options=stream_options, + user=user, + **kwargs, + ) if stream: - return self._stream("/v1/chat/completions", body) - - data = await self._client._request("POST", "/v1/chat/completions", json=body) - return ChatCompletion.from_dict(data) - - async def _stream(self, path: str, body: dict[str, Any]) -> AsyncIterator[ChatCompletionChunk]: - async with self._client._http.stream("POST", path, json=body) as response: - try: - response.raise_for_status() - except httpx.HTTPStatusError as e: - await response.aread() - from ..client import _raise_api_error - - _raise_api_error(e) - async for line in response.aiter_lines(): - if line.startswith("data: "): - payload = line[6:].strip() - if payload == "[DONE]": - return - try: - chunk_data = json.loads(payload) - except json.JSONDecodeError as e: - raise FerroStreamError( - f"Malformed SSE chunk in streaming response: {payload[:200]!r}" - ) from e - yield ChatCompletionChunk.from_dict(chunk_data) + return AsyncStream(await self._client._open_stream(PATH, body)) + return ChatCompletion.from_dict(await self._client._request("POST", PATH, json=body)) diff --git a/ferrolabsai/completions/resource.py b/ferrolabsai/completions/resource.py index 72f2a20..df9f4ea 100644 --- a/ferrolabsai/completions/resource.py +++ b/ferrolabsai/completions/resource.py @@ -2,16 +2,28 @@ from __future__ import annotations -import json -from collections.abc import Iterator -from typing import Any, Literal, overload +from typing import TYPE_CHECKING, Any, Literal, overload -from ..exceptions import FerroStreamError -from ..types import ChatCompletion, ChatCompletionChunk +from ..streaming import Stream +from ..types import ChatCompletion + +if TYPE_CHECKING: + from ..client import FerroClient + +PATH = "/v1/chat/completions" + + +def build_body( + model: str, messages: list[dict[str, Any]], stream: bool, **optional: Any +) -> dict[str, Any]: + """OpenAI-shaped request body; ``None`` optionals are omitted.""" + body: dict[str, Any] = {"model": model, "messages": messages, "stream": stream} + body.update({k: v for k, v in optional.items() if v is not None}) + return body class Completions: - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client @overload @@ -32,7 +44,7 @@ def create( messages: list[dict[str, Any]], stream: Literal[True], **kwargs: Any, - ) -> Iterator[ChatCompletionChunk]: ... + ) -> Stream: ... def create( self, @@ -42,34 +54,34 @@ def create( stream: bool = False, temperature: float | None = None, max_tokens: int | None = None, + max_completion_tokens: int | None = None, top_p: float | None = None, frequency_penalty: float | None = None, presence_penalty: float | None = None, stop: str | list[str] | None = None, + seed: int | None = None, tools: list[dict[str, Any]] | None = None, tool_choice: Any | None = None, - # Ferro-specific extras - template_id: str | None = None, - template_variables: dict[str, Any] | None = None, - route_tag: str | None = None, + parallel_tool_calls: bool | None = None, + response_format: dict[str, Any] | None = None, + stream_options: dict[str, Any] | None = None, user: str | None = None, **kwargs: Any, - ) -> ChatCompletion | Iterator[ChatCompletionChunk]: + ) -> ChatCompletion | Stream: """ - Create a chat completion. OpenAI-compatible with Ferro extras. + Create a chat completion (OpenAI-compatible). Args: - model: Model name. Ferro auto-routes to the correct provider. - E.g. "gpt-4o" → OpenAI, "claude-3-5-sonnet-20241022" → Anthropic. + model: Model name. The gateway routes it to the right provider, e.g. + "gpt-4o" → OpenAI, "claude-3-5-sonnet-20241022" → Anthropic. messages: List of message dicts with "role" and "content". - stream: If True, returns an iterator of ChatCompletionChunk objects. - template_id: Use a server-side prompt template (Ferro-specific). - template_variables: Variables for the template (Ferro-specific). - route_tag: Override routing strategy for this request (Ferro-specific). - user: End-user identifier for per-user tracking (Ferro-specific). - - Returns: - ChatCompletion or Iterator[ChatCompletionChunk] if stream=True. + stream: If True, returns a :class:`~ferrolabsai.Stream` of chunks that + also carries ``trace_id`` / ``provider`` from the response headers. + max_completion_tokens: Supersedes ``max_tokens`` (both are accepted). + stream_options: e.g. ``{"include_usage": True}`` to receive a terminal + chunk with ``usage`` (only honoured when ``stream=True``). + response_format: e.g. ``{"type": "json_object"}`` or a ``json_schema`` spec. + **kwargs: Any other OpenAI parameter is forwarded verbatim. Example:: @@ -77,69 +89,38 @@ def create( model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], ) - print(response.content) + print(response.content, response.provider, response.trace_id) # Streaming - for chunk in client.chat.completions.create( + stream = client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], stream=True, - ): + stream_options={"include_usage": True}, + ) + for chunk in stream: print(chunk.choices[0].delta.content or "", end="", flush=True) """ - body: dict[str, Any] = { - "model": model, - "messages": messages, - "stream": stream, - } - - # Standard OpenAI params - if temperature is not None: - body["temperature"] = temperature - if max_tokens is not None: - body["max_tokens"] = max_tokens - if top_p is not None: - body["top_p"] = top_p - if frequency_penalty is not None: - body["frequency_penalty"] = frequency_penalty - if presence_penalty is not None: - body["presence_penalty"] = presence_penalty - if stop is not None: - body["stop"] = stop - if tools is not None: - body["tools"] = tools - if tool_choice is not None: - body["tool_choice"] = tool_choice - if user is not None: - body["user"] = user - - # Ferro-specific - if template_id is not None: - body["template_id"] = template_id - if template_variables is not None: - body["template_variables"] = template_variables - if route_tag is not None: - body["x_route_tag"] = route_tag - - body.update(kwargs) - + body = build_body( + model, + messages, + stream, + temperature=temperature, + max_tokens=max_tokens, + max_completion_tokens=max_completion_tokens, + top_p=top_p, + frequency_penalty=frequency_penalty, + presence_penalty=presence_penalty, + stop=stop, + seed=seed, + tools=tools, + tool_choice=tool_choice, + parallel_tool_calls=parallel_tool_calls, + response_format=response_format, + stream_options=stream_options, + user=user, + **kwargs, + ) if stream: - return self._stream("/v1/chat/completions", body) - - data = self._client._request("POST", "/v1/chat/completions", json=body) - return ChatCompletion.from_dict(data) - - def _stream(self, path: str, body: dict[str, Any]) -> Iterator[ChatCompletionChunk]: - """Yields parsed ChatCompletionChunk objects from an SSE stream.""" - for line in self._client._stream_request(path, body): - if line.startswith("data: "): - payload = line[6:].strip() - if payload == "[DONE]": - return - try: - chunk_data = json.loads(payload) - except json.JSONDecodeError as e: - raise FerroStreamError( - f"Malformed SSE chunk in streaming response: {payload[:200]!r}" - ) from e - yield ChatCompletionChunk.from_dict(chunk_data) + return Stream(self._client._open_stream(PATH, body)) + return ChatCompletion.from_dict(self._client._request("POST", PATH, json=body)) diff --git a/ferrolabsai/embeddings/async_resource.py b/ferrolabsai/embeddings/async_resource.py index 9ee2b5c..8093a97 100644 --- a/ferrolabsai/embeddings/async_resource.py +++ b/ferrolabsai/embeddings/async_resource.py @@ -2,13 +2,17 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING from ..types import EmbeddingResponse +from .resource import PATH, build_body + +if TYPE_CHECKING: + from ..client import AsyncFerroClient class AsyncEmbeddings: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def create( @@ -20,12 +24,6 @@ async def create( dimensions: int | None = None, user: str | None = None, ) -> EmbeddingResponse: - body: dict[str, Any] = {"model": model, "input": input} - if encoding_format is not None: - body["encoding_format"] = encoding_format - if dimensions is not None: - body["dimensions"] = dimensions - if user is not None: - body["user"] = user - data = await self._client._request("POST", "/v1/embeddings", json=body) - return EmbeddingResponse.from_dict(data) + """See :meth:`ferrolabsai.embeddings.resource.Embeddings.create`.""" + body = build_body(model, input, encoding_format, dimensions, user) + return EmbeddingResponse.from_dict(await self._client._request("POST", PATH, json=body)) diff --git a/ferrolabsai/embeddings/resource.py b/ferrolabsai/embeddings/resource.py index c53585a..cd8fa7d 100644 --- a/ferrolabsai/embeddings/resource.py +++ b/ferrolabsai/embeddings/resource.py @@ -2,13 +2,35 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any from ..types import EmbeddingResponse +if TYPE_CHECKING: + from ..client import FerroClient + +PATH = "/v1/embeddings" + + +def build_body( + model: str, + input: str | list[str], + encoding_format: str | None, + dimensions: int | None, + user: str | None, +) -> dict[str, Any]: + body: dict[str, Any] = {"model": model, "input": input} + if encoding_format is not None: + body["encoding_format"] = encoding_format + if dimensions is not None: + body["dimensions"] = dimensions + if user is not None: + body["user"] = user + return body + class Embeddings: - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def create( @@ -38,13 +60,5 @@ def create( ) vectors = [d.embedding for d in response.data] """ - body: dict[str, Any] = {"model": model, "input": input} - if encoding_format is not None: - body["encoding_format"] = encoding_format - if dimensions is not None: - body["dimensions"] = dimensions - if user is not None: - body["user"] = user - - data = self._client._request("POST", "/v1/embeddings", json=body) - return EmbeddingResponse.from_dict(data) + body = build_body(model, input, encoding_format, dimensions, user) + return EmbeddingResponse.from_dict(self._client._request("POST", PATH, json=body)) diff --git a/ferrolabsai/exceptions/__init__.py b/ferrolabsai/exceptions/__init__.py index 60f1bee..3bbfff0 100644 --- a/ferrolabsai/exceptions/__init__.py +++ b/ferrolabsai/exceptions/__init__.py @@ -1,4 +1,12 @@ -"""Ferro Labs AI Gateway — exception hierarchy.""" +"""Ferro Labs AI Gateway — exception hierarchy. + +Status → exception mapping (ai-gateway error envelope ``{"error": {message, type, code}}``): + + 401 → FerroAuthError 402 → FerroBudgetExceededError + 403 → FerroPermissionError 404 → FerroNotFoundError + 429 → FerroRateLimitError 5xx → FerroServerError + other non-2xx → FerroAPIError +""" from __future__ import annotations @@ -35,8 +43,28 @@ class FerroAuthError(FerroAPIError): """401 — invalid or missing API key.""" +class FerroBudgetExceededError(FerroAPIError): + """402 ``insufficient_quota`` — the key's spend limit is exhausted.""" + + +class FerroPermissionError(FerroAPIError): + """403 ``insufficient_scope`` — the key lacks the scope for this route.""" + + class FerroRateLimitError(FerroAPIError): - """429 — rate limit or quota exceeded.""" + """429 — rate limit exceeded. ``retry_after`` mirrors the ``Retry-After`` header (seconds).""" + + def __init__( + self, + message: str, + *, + status_code: int | None = 429, + code: str | None = None, + request_id: str | None = None, + retry_after: float | None = None, + ) -> None: + super().__init__(message, status_code=status_code, code=code, request_id=request_id) + self.retry_after = retry_after class FerroNotFoundError(FerroAPIError): @@ -48,8 +76,13 @@ class FerroServerError(FerroAPIError): class FerroConnectionError(FerroError): - """Cannot connect to the gateway (network error, timeout, etc.).""" + """Cannot connect to the gateway (network error, timeout) after all retries.""" class FerroStreamError(FerroError): - """Error while consuming a streaming response.""" + """Error while consuming a streaming response (malformed frame or a gateway + error frame such as ``stream_error`` / ``stream_timeout``).""" + + def __init__(self, message: str, *, code: str | None = None) -> None: + super().__init__(message) + self.code = code diff --git a/ferrolabsai/images/async_resource.py b/ferrolabsai/images/async_resource.py index be4408f..cba714d 100644 --- a/ferrolabsai/images/async_resource.py +++ b/ferrolabsai/images/async_resource.py @@ -2,13 +2,17 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING from ..types import ImageResponse +from .resource import PATH, build_body + +if TYPE_CHECKING: + from ..client import AsyncFerroClient class AsyncImages: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client async def generate( @@ -23,19 +27,15 @@ async def generate( style: str | None = None, user: str | None = None, ) -> ImageResponse: - body: dict[str, Any] = {"model": model, "prompt": prompt} - if n is not None: - body["n"] = n - if size is not None: - body["size"] = size - if quality is not None: - body["quality"] = quality - if response_format is not None: - body["response_format"] = response_format - if style is not None: - body["style"] = style - if user is not None: - body["user"] = user - - data = await self._client._request("POST", "/v1/images/generations", json=body) - return ImageResponse.from_dict(data) + """See :meth:`ferrolabsai.images.resource.Images.generate`.""" + body = build_body( + model, + prompt, + n=n, + size=size, + quality=quality, + response_format=response_format, + style=style, + user=user, + ) + return ImageResponse.from_dict(await self._client._request("POST", PATH, json=body)) diff --git a/ferrolabsai/images/resource.py b/ferrolabsai/images/resource.py index cef1c12..8711192 100644 --- a/ferrolabsai/images/resource.py +++ b/ferrolabsai/images/resource.py @@ -2,13 +2,24 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any from ..types import ImageResponse +if TYPE_CHECKING: + from ..client import FerroClient + +PATH = "/v1/images/generations" + + +def build_body(model: str, prompt: str, **optional: Any) -> dict[str, Any]: + body: dict[str, Any] = {"model": model, "prompt": prompt} + body.update({k: v for k, v in optional.items() if v is not None}) + return body + class Images: - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client def generate( @@ -43,19 +54,14 @@ def generate( ) print(response.data[0].url) """ - body: dict[str, Any] = {"model": model, "prompt": prompt} - if n is not None: - body["n"] = n - if size is not None: - body["size"] = size - if quality is not None: - body["quality"] = quality - if response_format is not None: - body["response_format"] = response_format - if style is not None: - body["style"] = style - if user is not None: - body["user"] = user - - data = self._client._request("POST", "/v1/images/generations", json=body) - return ImageResponse.from_dict(data) + body = build_body( + model, + prompt, + n=n, + size=size, + quality=quality, + response_format=response_format, + style=style, + user=user, + ) + return ImageResponse.from_dict(self._client._request("POST", PATH, json=body)) diff --git a/ferrolabsai/models/async_resource.py b/ferrolabsai/models/async_resource.py index 9fd2634..1a29c51 100644 --- a/ferrolabsai/models/async_resource.py +++ b/ferrolabsai/models/async_resource.py @@ -1,38 +1,34 @@ -"""Async model catalog resource.""" +"""Async model catalog resource — see ``models/resource.py`` for the client-side lookup rules.""" from __future__ import annotations import builtins -from typing import Any +from typing import TYPE_CHECKING from ..types import ModelInfo +from .resource import PATH, filter_models, find_model, parse_catalog, search_models + +if TYPE_CHECKING: + from ..client import AsyncFerroClient class AsyncModels: - def __init__(self, client: Any) -> None: + def __init__(self, client: AsyncFerroClient) -> None: self._client = client + async def _fetch(self) -> builtins.list[ModelInfo]: + return parse_catalog(await self._client._request("GET", PATH)) + async def list( self, *, provider: str | None = None, capability: str | None = None, ) -> builtins.list[ModelInfo]: - params: dict[str, Any] = {} - if provider is not None: - params["provider"] = provider - if capability is not None: - params["capability"] = capability - - data = await self._client._request("GET", "/v1/models", params=params or None) - raw_models = data.get("data", data) if isinstance(data, dict) else data - return [ModelInfo.from_dict(m) for m in raw_models] + return filter_models(await self._fetch(), provider, capability) async def retrieve(self, model_id: str) -> ModelInfo: - data = await self._client._request("GET", f"/v1/models/{model_id}") - return ModelInfo.from_dict(data) + return find_model(await self._fetch(), model_id) async def search(self, query: str) -> builtins.list[ModelInfo]: - data = await self._client._request("GET", "/v1/models", params={"search": query}) - raw_models = data.get("data", data) if isinstance(data, dict) else data - return [ModelInfo.from_dict(m) for m in raw_models] + return search_models(await self._fetch(), query) diff --git a/ferrolabsai/models/resource.py b/ferrolabsai/models/resource.py index 8dc0c66..9ffd81d 100644 --- a/ferrolabsai/models/resource.py +++ b/ferrolabsai/models/resource.py @@ -1,17 +1,64 @@ -"""Model catalog resource — query 2,500+ models with pricing and capabilities.""" +"""Model catalog resource. + +The gateway serves exactly one catalog route, ``GET /v1/models``, and ignores +its query string. ``GET /v1/models/{id}`` is *not* a native route — it falls +through to the ``/v1/*`` pass-through and is forwarded upstream with the +operator's credential — so every lookup and filter here is client-side over +that one fetch. +""" from __future__ import annotations import builtins -from typing import Any +from typing import TYPE_CHECKING, Any +from ..exceptions import FerroNotFoundError from ..types import ModelInfo +if TYPE_CHECKING: + from ..client import FerroClient + +PATH = "/v1/models" + + +def parse_catalog(data: dict[str, Any]) -> builtins.list[ModelInfo]: + return [ModelInfo.from_dict(m) for m in data.get("data", [])] + + +def filter_models( + models: builtins.list[ModelInfo], provider: str | None, capability: str | None +) -> builtins.list[ModelInfo]: + return [ + m + for m in models + if (provider is None or m.owned_by == provider) + and (capability is None or capability in m.capabilities) + ] + + +def find_model(models: builtins.list[ModelInfo], model_id: str) -> ModelInfo: + for m in models: + if m.id == model_id: + return m + raise FerroNotFoundError( + f"Model {model_id!r} is not in the gateway catalog", + status_code=404, + code="model_not_found", + ) + + +def search_models(models: builtins.list[ModelInfo], query: str) -> builtins.list[ModelInfo]: + needle = query.lower() + return [m for m in models if needle in m.id.lower()] + class Models: - def __init__(self, client: Any) -> None: + def __init__(self, client: FerroClient) -> None: self._client = client + def _fetch(self) -> builtins.list[ModelInfo]: + return parse_catalog(self._client._request("GET", PATH)) + def list( self, *, @@ -19,59 +66,35 @@ def list( capability: str | None = None, ) -> builtins.list[ModelInfo]: """ - List all available models in the gateway's model catalog. + List the models the gateway can route to (``GET /v1/models``). Args: - provider: Filter by provider name. E.g. "openai", "anthropic", "groq". - capability: Filter by capability. E.g. "chat", "embeddings", "vision", - "function_calling". + provider: Keep only models whose ``owned_by`` matches, e.g. "openai". + capability: Keep only models whose ``capabilities`` include this, e.g. + "vision", "function_calling", "streaming", "reasoning". - Returns: - List of ModelInfo objects with pricing, context windows, and capabilities. + Filters are applied client-side — the gateway ignores query parameters. Example:: - # List all models models = client.models.list() - - # Only Anthropic models claude_models = client.models.list(provider="anthropic") - - # Only models with vision capability vision_models = client.models.list(capability="vision") """ - params: dict[str, Any] = {} - if provider is not None: - params["provider"] = provider - if capability is not None: - params["capability"] = capability - - data = self._client._request("GET", "/v1/models", params=params or None) - raw_models = data.get("data", data) if isinstance(data, dict) else data - return [ModelInfo.from_dict(m) for m in raw_models] + return filter_models(self._fetch(), provider, capability) def retrieve(self, model_id: str) -> ModelInfo: """ - Get details for a specific model by ID. + Look one model up in the catalog. Raises :class:`FerroNotFoundError` + (``code="model_not_found"``) locally; never calls ``/v1/models/{id}``. Example:: info = client.models.retrieve("gpt-4o") print(f"Context window: {info.context_window}") - print(f"Input cost: ${info.input_cost_per_token:.8f}/token") """ - data = self._client._request("GET", f"/v1/models/{model_id}") - return ModelInfo.from_dict(data) + return find_model(self._fetch(), model_id) def search(self, query: str) -> builtins.list[ModelInfo]: - """ - Search the model catalog by name or description. - - Example:: - - results = client.models.search("claude") - """ - params = {"search": query} - data = self._client._request("GET", "/v1/models", params=params) - raw_models = data.get("data", data) if isinstance(data, dict) else data - return [ModelInfo.from_dict(m) for m in raw_models] + """Case-insensitive substring match on model id, e.g. ``search("claude")``.""" + return search_models(self._fetch(), query) diff --git a/ferrolabsai/moderations/__init__.py b/ferrolabsai/moderations/__init__.py new file mode 100644 index 0000000..30bb875 --- /dev/null +++ b/ferrolabsai/moderations/__init__.py @@ -0,0 +1,6 @@ +"""Moderations resources (``/v1/moderations``).""" + +from .async_resource import AsyncModerations +from .resource import Moderations + +__all__ = ["AsyncModerations", "Moderations"] diff --git a/ferrolabsai/moderations/async_resource.py b/ferrolabsai/moderations/async_resource.py new file mode 100644 index 0000000..35b20c5 --- /dev/null +++ b/ferrolabsai/moderations/async_resource.py @@ -0,0 +1,20 @@ +"""Async moderations resource — see ``moderations/resource.py``.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from .resource import PATH, build_body + +if TYPE_CHECKING: + from ..client import AsyncFerroClient + + +class AsyncModerations: + def __init__(self, client: AsyncFerroClient) -> None: + self._client = client + + async def create( + self, *, input: str | list[str], model: str | None = None, **kwargs: Any + ) -> dict[str, Any]: + return await self._client._request("POST", PATH, json=build_body(input, model, **kwargs)) diff --git a/ferrolabsai/moderations/resource.py b/ferrolabsai/moderations/resource.py new file mode 100644 index 0000000..18edee9 --- /dev/null +++ b/ferrolabsai/moderations/resource.py @@ -0,0 +1,28 @@ +"""Moderations resource (``POST /v1/moderations``).""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from ..client import FerroClient + +PATH = "/v1/moderations" + + +def build_body(input: str | list[str], model: str | None, **kwargs: Any) -> dict[str, Any]: + body: dict[str, Any] = {"input": input, **kwargs} + if model is not None: + body["model"] = model + return body + + +class Moderations: + def __init__(self, client: FerroClient) -> None: + self._client = client + + def create( + self, *, input: str | list[str], model: str | None = None, **kwargs: Any + ) -> dict[str, Any]: + """``POST /v1/moderations`` — returns the provider's ``{id, model, results[]}`` JSON.""" + return self._client._request("POST", PATH, json=build_body(input, model, **kwargs)) diff --git a/ferrolabsai/responses/__init__.py b/ferrolabsai/responses/__init__.py new file mode 100644 index 0000000..9adbb39 --- /dev/null +++ b/ferrolabsai/responses/__init__.py @@ -0,0 +1,6 @@ +"""Responses API resources (``/v1/responses``).""" + +from .async_resource import AsyncResponses +from .resource import Responses + +__all__ = ["AsyncResponses", "Responses"] diff --git a/ferrolabsai/responses/async_resource.py b/ferrolabsai/responses/async_resource.py new file mode 100644 index 0000000..cd1b81b --- /dev/null +++ b/ferrolabsai/responses/async_resource.py @@ -0,0 +1,25 @@ +"""Async Responses API resource — see ``responses/resource.py``.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from ..types import Response +from .resource import PATH + +if TYPE_CHECKING: + from ..client import AsyncFerroClient + + +class AsyncResponses: + def __init__(self, client: AsyncFerroClient) -> None: + self._client = client + + async def create(self, **params: Any) -> Response: + return Response.from_dict(await self._client._request("POST", PATH, json=params)) + + async def retrieve(self, response_id: str) -> Response: + return Response.from_dict(await self._client._request("GET", f"{PATH}/{response_id}")) + + async def delete(self, response_id: str) -> dict[str, Any]: + return await self._client._request("DELETE", f"{PATH}/{response_id}") diff --git a/ferrolabsai/responses/resource.py b/ferrolabsai/responses/resource.py new file mode 100644 index 0000000..68ab1a9 --- /dev/null +++ b/ferrolabsai/responses/resource.py @@ -0,0 +1,37 @@ +"""Responses API resource (OpenAI-style ``/v1/responses``). + +``POST /v1/responses`` is model-routed, governed, and priced like chat. +The id sub-routes (``GET`` / ``DELETE /v1/responses/{id}``) carry no model +and pin to the gateway's ``responses_target``; without one configured the +gateway answers 501 ``not_implemented``. Streaming (``stream=True``) is not +wrapped in 0.3.0. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from ..types import Response + +if TYPE_CHECKING: + from ..client import FerroClient + +PATH = "/v1/responses" + + +class Responses: + def __init__(self, client: FerroClient) -> None: + self._client = client + + def create(self, **params: Any) -> Response: + """``POST /v1/responses`` — parameters are forwarded verbatim, e.g. + ``create(model="gpt-4o", input="Hello", instructions="Be brief")``.""" + return Response.from_dict(self._client._request("POST", PATH, json=params)) + + def retrieve(self, response_id: str) -> Response: + """``GET /v1/responses/{id}`` — 501 unless the gateway sets ``responses_target``.""" + return Response.from_dict(self._client._request("GET", f"{PATH}/{response_id}")) + + def delete(self, response_id: str) -> dict[str, Any]: + """``DELETE /v1/responses/{id}`` — 501 unless the gateway sets ``responses_target``.""" + return self._client._request("DELETE", f"{PATH}/{response_id}") diff --git a/ferrolabsai/streaming.py b/ferrolabsai/streaming.py new file mode 100644 index 0000000..e7a3226 --- /dev/null +++ b/ferrolabsai/streaming.py @@ -0,0 +1,125 @@ +"""SSE stream wrappers for ``chat.completions.create(stream=True)``. + +The wrapper keeps the HTTP response so header metadata (``X-Request-ID`` → +``trace_id``, ``X-Gateway-Provider`` → ``provider``) is available on the +stream object and copied onto every chunk. Streaming requests are never +retried: by the time the caller iterates, bytes may already have been read. + +Wire format (ai-gateway ``providers/core/chat.go``): ``data: {chunk}`` frames, +``usage`` only on the terminal chunk when ``stream_options.include_usage`` was +sent, ``data: [DONE]`` terminator, and mid-stream errors as a +``{"error": {message, type, code}}`` frame. +""" + +from __future__ import annotations + +import json +from collections.abc import AsyncIterator, Iterator +from dataclasses import replace + +import httpx + +from .exceptions import FerroStreamError +from .types import ChatCompletionChunk + +DONE = "[DONE]" + + +def _payload(line: str) -> str | None: + """Payload of one ``data:`` line, or ``None`` for comments/blank/other fields.""" + if not line.startswith("data:"): + return None + return line[5:].strip() + + +def _parse_chunk(payload: str, trace_id: str | None, provider: str | None) -> ChatCompletionChunk: + try: + data = json.loads(payload) + except json.JSONDecodeError as e: + raise FerroStreamError( + f"Malformed SSE chunk in streaming response: {payload[:200]!r}" + ) from e + error = data.get("error") if isinstance(data, dict) else None + if isinstance(error, dict): + raise FerroStreamError(error.get("message") or "stream error", code=error.get("code")) + return replace(ChatCompletionChunk.from_dict(data), trace_id=trace_id, provider=provider) + + +class Stream: + """Iterator of :class:`ChatCompletionChunk` over a live SSE response. + + Attributes: + response: The underlying ``httpx.Response`` (closed when exhausted or on ``close()``). + trace_id: ``X-Request-ID`` of the stream (32 hex chars on ai-gateway ≥ 1.4). + provider: ``X-Gateway-Provider`` when the gateway sets it (not on SSE as of v1.4.5). + """ + + def __init__(self, response: httpx.Response) -> None: + self.response = response + self.trace_id: str | None = response.headers.get("x-request-id") + self.provider: str | None = response.headers.get("x-gateway-provider") + self._chunks = self._iter() + + def __iter__(self) -> Iterator[ChatCompletionChunk]: + return self + + def __next__(self) -> ChatCompletionChunk: + return next(self._chunks) + + def _iter(self) -> Iterator[ChatCompletionChunk]: + try: + for line in self.response.iter_lines(): + payload = _payload(line) + if payload is None: + continue + if payload == DONE: + return + yield _parse_chunk(payload, self.trace_id, self.provider) + finally: + self.response.close() + + def close(self) -> None: + self.response.close() + + def __enter__(self) -> Stream: + return self + + def __exit__(self, *_: object) -> None: + self.close() + + +class AsyncStream: + """Async counterpart of :class:`Stream` (``async for chunk in stream``).""" + + def __init__(self, response: httpx.Response) -> None: + self.response = response + self.trace_id: str | None = response.headers.get("x-request-id") + self.provider: str | None = response.headers.get("x-gateway-provider") + self._chunks = self._iter() + + def __aiter__(self) -> AsyncIterator[ChatCompletionChunk]: + return self + + async def __anext__(self) -> ChatCompletionChunk: + return await self._chunks.__anext__() + + async def _iter(self) -> AsyncIterator[ChatCompletionChunk]: + try: + async for line in self.response.aiter_lines(): + payload = _payload(line) + if payload is None: + continue + if payload == DONE: + return + yield _parse_chunk(payload, self.trace_id, self.provider) + finally: + await self.response.aclose() + + async def aclose(self) -> None: + await self.response.aclose() + + async def __aenter__(self) -> AsyncStream: + return self + + async def __aexit__(self, *_: object) -> None: + await self.aclose() diff --git a/ferrolabsai/types.py b/ferrolabsai/types.py index bcc7412..a2512e2 100644 --- a/ferrolabsai/types.py +++ b/ferrolabsai/types.py @@ -1,6 +1,12 @@ """ Typed response models for the Ferro AI Gateway Python SDK. All models are dataclasses so they work without pydantic as a hard dependency. + +Field shapes follow ai-gateway v1.4.x (``providers/core/chat.go``, +``internal/handler/models.go``, ``internal/admin/model``). Gateway-specific +extensions are: body ``provider`` / ``provider_metadata`` / ``reasoning_content`` +and the extra ``usage`` token counters; ``trace_id`` and ``gateway_overhead_ms`` +come from the ``X-Request-ID`` and ``X-Gateway-Overhead-Ms`` response headers. """ from __future__ import annotations @@ -8,6 +14,28 @@ from dataclasses import dataclass, field from typing import Any +from .types_responses import Response + +__all__ = [ + "APIKey", + "ChatCompletion", + "ChatCompletionChunk", + "ChatMessage", + "Choice", + "ConfigHistoryEntry", + "CreatedAPIKey", + "EmbeddingData", + "EmbeddingResponse", + "GatewayConfig", + "ImageData", + "ImageResponse", + "ModelInfo", + "Response", + "StreamChoice", + "StreamDelta", + "Usage", +] + # ------------------------------------------------------------------ # Chat completions # ------------------------------------------------------------------ @@ -20,6 +48,7 @@ class ChatMessage: tool_calls: list[dict[str, Any]] | None = None tool_call_id: str | None = None name: str | None = None + reasoning_content: str | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> ChatMessage: @@ -29,6 +58,7 @@ def from_dict(cls, d: dict[str, Any]) -> ChatMessage: tool_calls=d.get("tool_calls"), tool_call_id=d.get("tool_call_id"), name=d.get("name"), + reasoning_content=d.get("reasoning_content"), ) @@ -37,10 +67,10 @@ class Usage: prompt_tokens: int = 0 completion_tokens: int = 0 total_tokens: int = 0 - # Ferro extras - cost_usd: float | None = None - cache_hit: bool | None = None - provider: str | None = None + # Gateway extensions (omitted by the gateway when zero). + reasoning_tokens: int | None = None + cache_read_tokens: int | None = None + cache_write_tokens: int | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> Usage: @@ -48,9 +78,9 @@ def from_dict(cls, d: dict[str, Any]) -> Usage: prompt_tokens=d.get("prompt_tokens", 0), completion_tokens=d.get("completion_tokens", 0), total_tokens=d.get("total_tokens", 0), - cost_usd=d.get("cost_usd"), - cache_hit=d.get("cache_hit"), - provider=d.get("provider"), + reasoning_tokens=d.get("reasoning_tokens"), + cache_read_tokens=d.get("cache_read_tokens"), + cache_write_tokens=d.get("cache_write_tokens"), ) @@ -79,10 +109,11 @@ class ChatCompletion: model: str choices: list[Choice] usage: Usage | None = None - # Ferro-specific extras - trace_id: str | None = None - provider: str | None = None - latency_ms: int | None = None + # Gateway extensions + trace_id: str | None = None # X-Request-ID header + provider: str | None = None # body `provider` (or X-Gateway-Provider header) + gateway_overhead_ms: float | None = None # X-Gateway-Overhead-Ms header + provider_metadata: dict[str, Any] | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> ChatCompletion: @@ -93,9 +124,10 @@ def from_dict(cls, d: dict[str, Any]) -> ChatCompletion: model=d.get("model", ""), choices=[Choice.from_dict(c) for c in d.get("choices", [])], usage=Usage.from_dict(d["usage"]) if d.get("usage") else None, - trace_id=d.get("x_ferro_trace_id") or d.get("trace_id"), - provider=d.get("x_ferro_provider") or d.get("provider"), - latency_ms=d.get("x_ferro_latency_ms") or d.get("latency_ms"), + trace_id=d.get("trace_id"), + provider=d.get("provider"), + gateway_overhead_ms=d.get("gateway_overhead_ms"), + provider_metadata=d.get("provider_metadata"), ) @property @@ -116,6 +148,7 @@ class StreamDelta: role: str | None = None content: str | None = None tool_calls: list[dict[str, Any]] | None = None + reasoning_content: str | None = None @dataclass @@ -133,6 +166,7 @@ def from_dict(cls, d: dict[str, Any]) -> StreamChoice: role=delta.get("role"), content=delta.get("content"), tool_calls=delta.get("tool_calls"), + reasoning_content=delta.get("reasoning_content"), ), finish_reason=d.get("finish_reason"), ) @@ -145,6 +179,12 @@ class ChatCompletionChunk: created: int model: str choices: list[StreamChoice] + # Only on the terminal chunk, and only when the request sent + # ``stream_options={"include_usage": True}``. + usage: Usage | None = None + # Copied from the stream's response headers onto every chunk. + trace_id: str | None = None + provider: str | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> ChatCompletionChunk: @@ -154,6 +194,7 @@ def from_dict(cls, d: dict[str, Any]) -> ChatCompletionChunk: created=d.get("created", 0), model=d.get("model", ""), choices=[StreamChoice.from_dict(c) for c in d.get("choices", [])], + usage=Usage.from_dict(d["usage"]) if d.get("usage") else None, ) @@ -183,6 +224,7 @@ class EmbeddingResponse: data: list[EmbeddingData] model: str usage: Usage | None = None + trace_id: str | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> EmbeddingResponse: @@ -191,6 +233,7 @@ def from_dict(cls, d: dict[str, Any]) -> EmbeddingResponse: data=[EmbeddingData.from_dict(e) for e in d.get("data", [])], model=d.get("model", ""), usage=Usage.from_dict(d["usage"]) if d.get("usage") else None, + trace_id=d.get("trace_id"), ) @@ -218,12 +261,14 @@ def from_dict(cls, d: dict[str, Any]) -> ImageData: class ImageResponse: created: int data: list[ImageData] + trace_id: str | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> ImageResponse: return cls( created=d.get("created", 0), data=[ImageData.from_dict(i) for i in d.get("data", [])], + trace_id=d.get("trace_id"), ) @@ -231,10 +276,9 @@ def from_dict(cls, d: dict[str, Any]) -> ImageResponse: # Admin — API Keys # ------------------------------------------------------------------ # -# Field shape matches admin.APIKey in -# ai-gateway/internal/admin/keys.go. On retrieve (GET /admin/keys/{id}) -# the gateway masks `key` to `...`. On create / rotate the -# full key is returned exactly once — captured by CreatedAPIKey below. +# Field shape matches model.APIKey in ai-gateway/internal/admin/model. +# Listings and GET /admin/keys/{id} mask `key` to `fgw_ab12...cd34`; the full +# secret is returned exactly once from create / rotate (CreatedAPIKey). @dataclass @@ -245,7 +289,7 @@ class APIKey: created_at: str = "" active: bool = True usage_count: int = 0 - key: str | None = None # masked (e.g. "fgw_abcd...") on retrieve + key: str | None = None # masked (e.g. "fgw_ab12...cd34") on retrieve expires_at: str | None = None revoked_at: str | None = None rotated_at: str | None = None @@ -342,32 +386,39 @@ def from_dict(cls, d: dict[str, Any]) -> ConfigHistoryEntry: # ------------------------------------------------------------------ -# Model catalog +# Model catalog — EnrichedModelInfo (ai-gateway internal/handler/models.go) # ------------------------------------------------------------------ @dataclass class ModelInfo: id: str - object: str - provider: str + object: str = "model" + owned_by: str = "" + created: int = 0 + mode: str | None = None # "chat", "embedding", "image", ... context_window: int | None = None max_output_tokens: int | None = None - input_cost_per_token: float | None = None - output_cost_per_token: float | None = None - capabilities: list[str] | None = None + capabilities: list[str] = field(default_factory=list) status: str | None = None + deprecated: bool = False + + @property + def provider(self) -> str: + """Alias for ``owned_by`` — the provider that serves this model.""" + return self.owned_by @classmethod def from_dict(cls, d: dict[str, Any]) -> ModelInfo: return cls( id=d.get("id", ""), object=d.get("object", "model"), - provider=d.get("owned_by") or d.get("provider", ""), + owned_by=d.get("owned_by", ""), + created=d.get("created", 0), + mode=d.get("mode"), context_window=d.get("context_window"), max_output_tokens=d.get("max_output_tokens"), - input_cost_per_token=d.get("input_cost_per_token"), - output_cost_per_token=d.get("output_cost_per_token"), - capabilities=d.get("capabilities"), + capabilities=d.get("capabilities") or [], status=d.get("status"), + deprecated=d.get("deprecated", False), ) diff --git a/ferrolabsai/types_responses.py b/ferrolabsai/types_responses.py new file mode 100644 index 0000000..82d7bd9 --- /dev/null +++ b/ferrolabsai/types_responses.py @@ -0,0 +1,40 @@ +"""Response type for the OpenAI-style Responses API (``POST /v1/responses``). + +The Responses object is large and evolving, so only the stable top-level +fields are typed; ``raw`` keeps the full body. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + + +@dataclass +class Response: + id: str + object: str = "response" + created_at: int = 0 + status: str | None = None # "completed", "in_progress", "failed", ... + model: str = "" + output: list[dict[str, Any]] = field(default_factory=list) + usage: dict[str, Any] | None = None + # Gateway extensions + trace_id: str | None = None # X-Request-ID header + provider: str | None = None # X-Gateway-Provider header + raw: dict[str, Any] = field(default_factory=dict) + + @classmethod + def from_dict(cls, d: dict[str, Any]) -> Response: + return cls( + id=d.get("id", ""), + object=d.get("object", "response"), + created_at=d.get("created_at", 0), + status=d.get("status"), + model=d.get("model", ""), + output=d.get("output") or [], + usage=d.get("usage"), + trace_id=d.get("trace_id"), + provider=d.get("provider"), + raw=d, + ) diff --git a/pyproject.toml b/pyproject.toml index 1d7496f..569132e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -61,7 +61,6 @@ target-version = "py39" select = ["E", "F", "I", "UP"] [tool.mypy] -python_version = "3.9" strict = true ignore_missing_imports = true diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..823bb61 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,115 @@ +"""Shared fixtures for the ferrolabsai unit suite. + +All HTTP is mocked with pytest-httpx — no gateway required. The contract +suite under ``tests/contract`` is the one place a real gateway is used. +""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from ferrolabsai import AsyncFerroClient, FerroClient + +BASE_URL = "http://localhost:8080" +API_KEY = "sk-ferro-testkey123" +TRACE_ID = "0af7651916cd43dd8448eb211c80319c" + +COMPLETION_RESPONSE: dict[str, Any] = { + "id": "chatcmpl-abc123", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o", + "provider": "openai", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hello from Ferro!"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, +} + +EMBEDDING_RESPONSE: dict[str, Any] = { + "object": "list", + "data": [ + {"index": 0, "object": "embedding", "embedding": [0.1, 0.2, 0.3]}, + {"index": 1, "object": "embedding", "embedding": [0.4, 0.5, 0.6]}, + ], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 8, "total_tokens": 8}, +} + +IMAGE_RESPONSE: dict[str, Any] = { + "created": 1700000000, + "data": [{"url": "https://example.com/image.png", "revised_prompt": "A polished prompt"}], +} + +# Shape of GET /v1/models: EnrichedModelInfo (ai-gateway internal/handler/models.go). +MODELS_RESPONSE: dict[str, Any] = { + "object": "list", + "data": [ + { + "id": "gpt-4o", + "object": "model", + "created": 0, + "owned_by": "openai", + "mode": "chat", + "context_window": 128000, + "max_output_tokens": 16384, + "capabilities": ["vision", "function_calling", "streaming"], + "status": "active", + }, + { + "id": "claude-3-5-sonnet-20241022", + "object": "model", + "created": 0, + "owned_by": "anthropic", + "mode": "chat", + "context_window": 200000, + "capabilities": ["function_calling", "streaming"], + "deprecated": True, + }, + { + "id": "text-embedding-3-small", + "object": "model", + "created": 0, + "owned_by": "openai", + "mode": "embedding", + }, + ], +} + + +def sse(*frames: dict[str, Any] | str) -> bytes: + """Encode dict frames (or raw payload strings) as an SSE body ending in [DONE].""" + import json + + lines = [ + f"data: {frame if isinstance(frame, str) else json.dumps(frame)}\n\n" for frame in frames + ] + return "".join(lines).encode() + b"data: [DONE]\n\n" + + +def chunk(content: str | None = None, **extra: Any) -> dict[str, Any]: + delta: dict[str, Any] = {} if content is None else {"content": content} + return { + "id": "c1", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o", + "choices": [{"index": 0, "delta": delta, "finish_reason": None}], + **extra, + } + + +@pytest.fixture +def client() -> FerroClient: + return FerroClient(api_key=API_KEY, base_url=BASE_URL, max_retries=0) + + +@pytest.fixture +def async_client() -> AsyncFerroClient: + return AsyncFerroClient(api_key=API_KEY, base_url=BASE_URL, max_retries=0) diff --git a/tests/test_admin.py b/tests/test_admin.py new file mode 100644 index 0000000..5b85e2b --- /dev/null +++ b/tests/test_admin.py @@ -0,0 +1,269 @@ +"""Admin API — /admin/* parity with ai-gateway v1.4.x.""" + +from __future__ import annotations + +import json + +import pytest +from pytest_httpx import HTTPXMock + +from .conftest import BASE_URL + +ADMIN = f"{BASE_URL}/admin" + +KEY = { + "id": "key_abc", + "name": "test-key", + "scopes": ["admin"], + "active": True, + "usage_count": 12, + "created_at": "2026-04-01T00:00:00Z", +} + + +class TestKeys: + def test_create(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=f"{ADMIN}/keys", json={**KEY, "key": "fgw_full"}) + key = client.admin.keys.create( + name="test-key", scopes=["admin"], expires_at="2027-01-01T00:00:00Z" + ) + assert key.key == "fgw_full" + assert key.scopes == ["admin"] + body = json.loads(httpx_mock.get_requests()[0].content) + assert body == { + "name": "test-key", + "scopes": ["admin"], + "expires_at": "2027-01-01T00:00:00Z", + } + + def test_list_accepts_bare_array(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=f"{ADMIN}/keys", json=[KEY]) + keys = client.admin.keys.list() + assert keys[0].id == "key_abc" + assert keys[0].usage_count == 12 + + def test_retrieve_masks_key(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/keys/key_abc", json={**KEY, "key": "fgw_ab12...cd34"} + ) + key = client.admin.keys.retrieve("key_abc") + assert key.id == "key_abc" + assert key.key == "fgw_ab12...cd34" + + def test_update(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="PUT", + url=f"{ADMIN}/keys/key_abc", + json={**KEY, "name": "renamed", "active": False}, + ) + key = client.admin.keys.update("key_abc", name="renamed", active=False) + assert key.name == "renamed" + assert key.active is False + assert json.loads(httpx_mock.get_requests()[0].content) == { + "name": "renamed", + "active": False, + } + + @pytest.mark.parametrize( + ("method", "http_method", "path"), + [("revoke", "POST", "/keys/key_abc/revoke"), ("delete", "DELETE", "/keys/key_abc")], + ) + def test_revoke_and_delete_accept_204( + self, client, httpx_mock: HTTPXMock, method, http_method, path + ): + httpx_mock.add_response( + method=http_method, url=f"{ADMIN}{path}", status_code=204, content=b"" + ) + assert getattr(client.admin.keys, method)("key_abc") is None + + def test_rotate(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{ADMIN}/keys/key_abc/rotate", json={**KEY, "key": "fgw_new"} + ) + assert client.admin.keys.rotate("key_abc").key == "fgw_new" + + def test_usage(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{ADMIN}/keys/usage?limit=10&offset=0&sort=usage&active=true", + json={"data": [], "summary": {"total_usage": 100}}, + ) + assert client.admin.keys.usage(limit=10, active=True)["summary"]["total_usage"] == 100 + + async def test_async_keys(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=f"{ADMIN}/keys", json=[KEY]) + httpx_mock.add_response(method="GET", url=f"{ADMIN}/keys/key_abc", json=KEY) + httpx_mock.add_response(method="PUT", url=f"{ADMIN}/keys/key_abc", json=KEY) + httpx_mock.add_response( + method="POST", url=f"{ADMIN}/keys/key_abc/revoke", status_code=204, content=b"" + ) + httpx_mock.add_response( + method="DELETE", url=f"{ADMIN}/keys/key_abc", status_code=204, content=b"" + ) + assert (await async_client.admin.keys.list())[0].id == "key_abc" + assert (await async_client.admin.keys.retrieve("key_abc")).name == "test-key" + assert (await async_client.admin.keys.update("key_abc", name="x")).id == "key_abc" + await async_client.admin.keys.revoke("key_abc") + await async_client.admin.keys.delete("key_abc") + + +class TestConfig: + def test_get(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{ADMIN}/config", + json={ + "strategy": {"mode": "fallback"}, + "targets": [{"virtual_key": "openai"}], + "mcp_servers": [], + }, + ) + cfg = client.admin.config.get() + assert cfg.strategy == {"mode": "fallback"} + assert cfg.targets == [{"virtual_key": "openai"}] + assert cfg.raw["mcp_servers"] == [] + + @pytest.mark.parametrize(("method", "http_method"), [("create", "POST"), ("update", "PUT")]) + def test_create_and_update(self, client, httpx_mock: HTTPXMock, method, http_method): + httpx_mock.add_response(method=http_method, url=f"{ADMIN}/config", json={"status": "ok"}) + payload = {"strategy": {"mode": "single"}, "targets": [{"virtual_key": "openai"}]} + assert getattr(client.admin.config, method)(payload) == {"status": "ok"} + assert json.loads(httpx_mock.get_requests()[0].content) == payload + + def test_delete(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="DELETE", url=f"{ADMIN}/config", json={"status": "reset"}) + assert client.admin.config.delete() == {"status": "reset"} + + def test_history_and_rollback(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{ADMIN}/config/history", + json={ + "data": [ + {"version": 1, "updated_at": "t", "config": {"strategy": {"mode": "single"}}} + ] + }, + ) + httpx_mock.add_response( + method="POST", url=f"{ADMIN}/config/rollback/1", json={"status": "rolled_back"} + ) + history = client.admin.config.history() + assert history[0].version == 1 + assert history[0].config.strategy == {"mode": "single"} + assert client.admin.config.rollback(1)["status"] == "rolled_back" + + async def test_async_history(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/config/history", json={"data": [{"version": 2}]} + ) + assert (await async_client.admin.config.history())[0].version == 2 + + +class TestLogs: + def test_list_forwards_v14_filters(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{ADMIN}/logs?limit=10&offset=0&stage=all&provider=openai&model=gpt-4o&since=t&api_key_id=none", + json={"data": [{"trace_id": "t1"}], "summary": {"total_entries": 1}}, + ) + result = client.admin.logs.list( + limit=10, stage="all", provider="openai", model="gpt-4o", since="t", api_key_id="none" + ) + assert result["data"][0]["trace_id"] == "t1" + + def test_list_rejects_trace_id_filter(self, client): + with pytest.raises(TypeError): + client.admin.logs.list(trace_id="t1") + + def test_stats_with_buckets(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/logs/stats?buckets=12", json={"total": 42, "series": []} + ) + assert client.admin.logs.stats(buckets=12)["total"] == 42 + + def test_delete(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="DELETE", + url=f"{ADMIN}/logs?before=2026-01-01T00%3A00%3A00Z", + json={"deleted": 3}, + ) + assert client.admin.logs.delete(before="2026-01-01T00:00:00Z") == {"deleted": 3} + + async def test_async_logs(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/logs?limit=50&offset=0&api_key_id=k1", json={"data": []} + ) + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/logs/stats?buckets=4", json={"total": 0} + ) + httpx_mock.add_response(method="DELETE", url=f"{ADMIN}/logs", json={"deleted": 0}) + assert await async_client.admin.logs.list(api_key_id="k1") == {"data": []} + assert await async_client.admin.logs.stats(buckets=4) == {"total": 0} + assert await async_client.admin.logs.delete() == {"deleted": 0} + + +class TestProvidersPluginsAudit: + @pytest.mark.parametrize( + "payload", + [[{"name": "cache"}], {"data": [{"name": "cache"}]}, {"plugins": [{"name": "cache"}]}], + ) + def test_plugins_list_shapes(self, client, httpx_mock: HTTPXMock, payload): + httpx_mock.add_response(method="GET", url=f"{ADMIN}/plugins", json=payload) + assert client.admin.plugins.list() == [{"name": "cache"}] + + def test_providers_list(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/providers", json={"data": [{"name": "openai"}]} + ) + assert client.admin.providers.list() == [{"name": "openai"}] + + def test_catalogs(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{ADMIN}/providers/catalog", + json=[{"id": "openai", "registered": True, "catalog_models": 90}], + ) + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/plugins/catalog", json={"data": [{"name": "budget"}]} + ) + assert client.admin.providers.catalog()[0]["id"] == "openai" + assert client.admin.plugins.catalog() == [{"name": "budget"}] + + def test_audit_list(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{ADMIN}/audit?limit=20&offset=0&action=key.create&actor_id=k1&outcome=success&since=t", + json={"data": [{"action": "key.create"}], "summary": {"total_entries": 1}}, + ) + result = client.admin.audit.list( + action="key.create", actor_id="k1", outcome="success", since="t", limit=20 + ) + assert result["data"][0]["action"] == "key.create" + + def test_dashboard_and_health(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/dashboard", json={"keys": {"active": 5}} + ) + httpx_mock.add_response(method="GET", url=f"{ADMIN}/health", json={"status": "ok"}) + assert client.admin.dashboard()["keys"]["active"] == 5 + assert client.admin.health()["status"] == "ok" + + async def test_async_parity(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/providers", json={"data": [{"name": "openai"}]} + ) + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/providers/catalog", json=[{"id": "openai"}] + ) + httpx_mock.add_response(method="GET", url=f"{ADMIN}/plugins/catalog", json={"data": []}) + httpx_mock.add_response( + method="GET", url=f"{ADMIN}/audit?limit=50&offset=0", json={"data": []} + ) + httpx_mock.add_response(method="GET", url=f"{ADMIN}/health", json={"status": "ok"}) + httpx_mock.add_response(method="GET", url=f"{ADMIN}/dashboard", json={}) + assert await async_client.admin.providers.list() == [{"name": "openai"}] + assert await async_client.admin.providers.catalog() == [{"id": "openai"}] + assert await async_client.admin.plugins.catalog() == [] + assert await async_client.admin.audit.list() == {"data": []} + assert (await async_client.admin.health())["status"] == "ok" + assert await async_client.admin.dashboard() == {} diff --git a/tests/test_chat.py b/tests/test_chat.py new file mode 100644 index 0000000..14aebc4 --- /dev/null +++ b/tests/test_chat.py @@ -0,0 +1,233 @@ +"""Chat completions: request body, non-streaming parsing, and SSE streaming.""" + +from __future__ import annotations + +import json + +import pytest +from pytest_httpx import HTTPXMock + +from ferrolabsai.exceptions import FerroAuthError, FerroRateLimitError, FerroStreamError +from ferrolabsai.streaming import AsyncStream, Stream + +from .conftest import BASE_URL, COMPLETION_RESPONSE, TRACE_ID, chunk, sse + +CHAT_URL = f"{BASE_URL}/v1/chat/completions" +MESSAGES = [{"role": "user", "content": "Hi"}] + + +class TestCreate: + def test_basic_create(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + response = client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert response.id == "chatcmpl-abc123" + assert response.model == "gpt-4o" + assert response.content == "Hello from Ferro!" + assert response.provider == "openai" + assert response.usage.total_tokens == 15 + body = json.loads(httpx_mock.get_requests()[0].content) + assert body == {"model": "gpt-4o", "messages": MESSAGES, "stream": False} + + def test_passes_optional_params(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + client.chat.completions.create( + model="gpt-4o", + messages=MESSAGES, + temperature=0.7, + max_tokens=100, + max_completion_tokens=200, + parallel_tool_calls=False, + response_format={"type": "json_object"}, + seed=42, + user="user_123", + extra_field="passthrough", + ) + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["temperature"] == 0.7 + assert body["max_tokens"] == 100 + assert body["max_completion_tokens"] == 200 + assert body["parallel_tool_calls"] is False + assert body["response_format"] == {"type": "json_object"} + assert body["seed"] == 42 + assert body["user"] == "user_123" + assert body["extra_field"] == "passthrough" + + def test_no_ferro_field_translation(self, client, httpx_mock: HTTPXMock): + # route_tag/template_* were never read by the gateway; the SDK no longer + # rewrites them (route_tag -> x_route_tag). Unknown kwargs pass through verbatim. + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + client.chat.completions.create(model="gpt-4o", messages=MESSAGES, route_tag="x") + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["route_tag"] == "x" + assert "x_route_tag" not in body + + def test_parses_gateway_body_extensions(self, client, httpx_mock: HTTPXMock): + body = json.loads(json.dumps(COMPLETION_RESPONSE)) + body["provider_metadata"] = {"openai": {"system_fingerprint": "fp_1"}} + body["choices"][0]["message"]["reasoning_content"] = "thinking..." + body["usage"].update( + {"reasoning_tokens": 3, "cache_read_tokens": 2, "cache_write_tokens": 1} + ) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=body) + response = client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert response.provider_metadata == {"openai": {"system_fingerprint": "fp_1"}} + assert response.choices[0].message.reasoning_content == "thinking..." + assert response.usage.reasoning_tokens == 3 + assert response.usage.cache_read_tokens == 2 + assert response.usage.cache_write_tokens == 1 + + async def test_async_forwards_params(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + response = await async_client.chat.completions.create( + model="gpt-4o", + messages=MESSAGES, + frequency_penalty=0.5, + presence_penalty=0.3, + stream_options={"include_usage": True}, + ) + assert response.content == "Hello from Ferro!" + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["frequency_penalty"] == 0.5 + assert body["presence_penalty"] == 0.3 + assert body["stream_options"] == {"include_usage": True} + + +class TestStreaming: + def test_sync_stream_yields_chunks_with_metadata(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + content=sse(chunk("Hello"), chunk(" world")), + headers={"X-Request-ID": TRACE_ID, "X-Gateway-Provider": "openai"}, + ) + stream = client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) + assert isinstance(stream, Stream) + assert stream.trace_id == TRACE_ID + assert stream.provider == "openai" + chunks = list(stream) + assert [c.choices[0].delta.content for c in chunks] == ["Hello", " world"] + assert chunks[0].trace_id == TRACE_ID + assert chunks[0].provider == "openai" + assert chunks[0].usage is None + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["stream"] is True + + def test_sync_stream_terminal_usage_chunk(self, client, httpx_mock: HTTPXMock): + usage = { + "prompt_tokens": 1, + "completion_tokens": 2, + "total_tokens": 3, + "reasoning_tokens": 1, + } + terminal = chunk(usage=usage) + terminal["choices"] = [] + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse(chunk("x"), terminal)) + chunks = list( + client.chat.completions.create( + model="gpt-4o", + messages=MESSAGES, + stream=True, + stream_options={"include_usage": True}, + ) + ) + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["stream_options"] == {"include_usage": True} + assert chunks[-1].choices == [] + assert chunks[-1].usage.total_tokens == 3 + assert chunks[-1].usage.reasoning_tokens == 1 + + def test_sync_stream_reasoning_content(self, client, httpx_mock: HTTPXMock): + frame = chunk() + frame["choices"][0]["delta"]["reasoning_content"] = "hmm" + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse(frame)) + chunks = list( + client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) + ) + assert chunks[0].choices[0].delta.reasoning_content == "hmm" + + def test_sync_stream_context_manager_closes_response(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse(chunk("a"), chunk("b"))) + with client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) as s: + assert next(s).choices[0].delta.content == "a" + assert s.response.is_closed + + @pytest.mark.parametrize(("status", "exc"), [(401, FerroAuthError), (429, FerroRateLimitError)]) + def test_sync_stream_http_errors_raise_before_iteration( + self, client, httpx_mock: HTTPXMock, status, exc + ): + httpx_mock.add_response( + method="POST", url=CHAT_URL, status_code=status, json={"error": {"message": "denied"}} + ) + with pytest.raises(exc, match="denied"): + client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) + + def test_sync_stream_malformed_chunk(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse("{not valid json}")) + with pytest.raises(FerroStreamError, match="Malformed SSE chunk"): + list(client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True)) + + def test_sync_stream_error_frame(self, client, httpx_mock: HTTPXMock): + error = { + "error": { + "message": "upstream hung up", + "type": "stream_error", + "code": "stream_timeout", + } + } + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse(chunk("a"), error)) + stream = client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) + assert next(stream).choices[0].delta.content == "a" + with pytest.raises(FerroStreamError, match="upstream hung up") as exc_info: + next(stream) + assert exc_info.value.code == "stream_timeout" + + async def test_async_stream_happy_path(self, async_client, httpx_mock: HTTPXMock): + usage = {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2} + terminal = chunk(usage=usage) + terminal["choices"] = [] + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + content=sse(chunk("Hel"), chunk("lo"), terminal), + headers={"X-Request-ID": TRACE_ID}, + ) + stream = await async_client.chat.completions.create( + model="gpt-4o", messages=MESSAGES, stream=True, stream_options={"include_usage": True} + ) + assert isinstance(stream, AsyncStream) + assert stream.trace_id == TRACE_ID + chunks = [c async for c in stream] + assert "".join(c.choices[0].delta.content for c in chunks[:2]) == "Hello" + assert chunks[-1].usage.total_tokens == 2 + assert chunks[0].trace_id == TRACE_ID + + async def test_async_stream_http_error(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=CHAT_URL, status_code=401, json={"error": {"message": "Invalid"}} + ) + with pytest.raises(FerroAuthError, match="Invalid"): + await async_client.chat.completions.create( + model="gpt-4o", messages=MESSAGES, stream=True + ) + + async def test_async_stream_malformed_chunk(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse("{not valid json}")) + stream = await async_client.chat.completions.create( + model="gpt-4o", messages=MESSAGES, stream=True + ) + with pytest.raises(FerroStreamError, match="Malformed SSE chunk"): + async for _ in stream: + pass + + async def test_async_stream_error_frame(self, async_client, httpx_mock: HTTPXMock): + error = {"error": {"message": "boom", "type": "stream_error", "code": "stream_error"}} + httpx_mock.add_response(method="POST", url=CHAT_URL, content=sse(error)) + stream = await async_client.chat.completions.create( + model="gpt-4o", messages=MESSAGES, stream=True + ) + with pytest.raises(FerroStreamError) as exc_info: + async with stream: + async for _ in stream: + pass + assert exc_info.value.code == "stream_error" + assert stream.response.is_closed diff --git a/tests/test_client.py b/tests/test_client.py new file mode 100644 index 0000000..ee9bd85 --- /dev/null +++ b/tests/test_client.py @@ -0,0 +1,431 @@ +"""Client construction, retries, error mapping, and response-header metadata.""" + +from __future__ import annotations + +import re +from pathlib import Path + +import httpx +import pytest +from pytest_httpx import HTTPXMock + +import ferrolabsai +from ferrolabsai import AsyncFerroClient, FerroClient +from ferrolabsai.exceptions import ( + FerroAPIError, + FerroAuthError, + FerroBudgetExceededError, + FerroConnectionError, + FerroNotFoundError, + FerroPermissionError, + FerroRateLimitError, + FerroServerError, +) + +from .conftest import API_KEY, BASE_URL, COMPLETION_RESPONSE, TRACE_ID + +CHAT_URL = f"{BASE_URL}/v1/chat/completions" + + +def _chat(client: FerroClient): + return client.chat.completions.create( + model="gpt-4o", messages=[{"role": "user", "content": "Hi"}] + ) + + +class TestClientInit: + def test_requires_api_key(self, monkeypatch): + monkeypatch.delenv("FERRO_API_KEY", raising=False) + monkeypatch.delenv("OPENAI_API_KEY", raising=False) + with pytest.raises(FerroAuthError): + FerroClient() + + def test_reads_from_env(self, monkeypatch): + monkeypatch.setenv("FERRO_API_KEY", "sk-ferro-envkey") + assert FerroClient().api_key == "sk-ferro-envkey" + + def test_falls_back_to_openai_env(self, monkeypatch): + monkeypatch.delenv("FERRO_API_KEY", raising=False) + monkeypatch.setenv("OPENAI_API_KEY", "sk-openai-compat") + assert FerroClient().api_key == "sk-openai-compat" + + def test_strips_trailing_slash(self): + c = FerroClient(api_key=API_KEY, base_url="https://localhost:8080/") + assert c.base_url == "https://localhost:8080" + + @pytest.mark.parametrize("client_cls", [FerroClient, AsyncFerroClient]) + def test_rejects_negative_max_retries(self, client_cls): + with pytest.raises(ValueError, match="max_retries must be >= 0"): + client_cls(api_key=API_KEY, max_retries=-1) + + @pytest.mark.parametrize("client_cls", [FerroClient, AsyncFerroClient]) + @pytest.mark.parametrize("invalid_value", [1.5, True]) + def test_rejects_non_integer_max_retries(self, client_cls, invalid_value): + with pytest.raises(TypeError, match="max_retries must be an integer"): + client_cls(api_key=API_KEY, max_retries=invalid_value) + + @pytest.mark.parametrize("client_cls", [FerroClient, AsyncFerroClient]) + def test_has_expected_namespaces(self, client_cls): + c = client_cls(api_key=API_KEY) + for name in ("chat", "embeddings", "images", "models", "admin", "responses", "moderations"): + assert hasattr(c, name), name + assert hasattr(c.chat, "completions") + for name in ("keys", "config", "logs", "providers", "plugins", "audit"): + assert hasattr(c.admin, name), name + + def test_sends_auth_and_user_agent(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + _chat(client) + request = httpx_mock.get_requests()[0] + assert request.headers["authorization"] == f"Bearer {API_KEY}" + assert request.headers["content-type"] == "application/json" + assert request.headers["user-agent"] == f"ferrolabsai-python/{ferrolabsai.__version__}" + + def test_version_matches_pyproject(self): + pyproject = Path(__file__).resolve().parents[1] / "pyproject.toml" + match = re.search(r'^version = "([^"]+)"', pyproject.read_text(), re.MULTILINE) + assert match and match.group(1) == ferrolabsai.__version__ + + def test_sync_byoc_merges_headers(self): + custom = httpx.Client(base_url=BASE_URL) + c = FerroClient(api_key=API_KEY, base_url=BASE_URL, http_client=custom) + assert c._http is custom + assert custom.headers["Authorization"] == f"Bearer {API_KEY}" + custom.close() + + def test_async_byoc_merges_headers(self): + custom = httpx.AsyncClient(base_url=BASE_URL) + c = AsyncFerroClient(api_key=API_KEY, base_url=BASE_URL, http_client=custom) + assert c._http is custom + assert custom.headers["Authorization"] == f"Bearer {API_KEY}" + + def test_sync_context_manager(self): + with FerroClient(api_key=API_KEY) as c: + assert c.api_key == API_KEY + + async def test_async_context_manager(self): + async with AsyncFerroClient(api_key=API_KEY) as c: + assert c.api_key == API_KEY + + +class TestResponseMetadata: + """Only the headers the gateway really sets are read (see docs/architecture.md).""" + + def test_headers_populate_inference_bodies(self, client, httpx_mock: HTTPXMock): + body = {k: v for k, v in COMPLETION_RESPONSE.items() if k != "provider"} + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + json=body, + headers={ + "X-Request-ID": TRACE_ID, + "X-Gateway-Provider": "openai", + "X-Gateway-Overhead-Ms": "1.250", + }, + ) + response = _chat(client) + assert response.trace_id == TRACE_ID + assert response.provider == "openai" + assert response.gateway_overhead_ms == 1.25 + assert not hasattr(response, "latency_ms") + assert not hasattr(response.usage, "cost_usd") + + def test_body_provider_wins_over_header(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + json=COMPLETION_RESPONSE, + headers={"X-Request-ID": TRACE_ID, "X-Gateway-Provider": "other"}, + ) + assert _chat(client).provider == "openai" + + def test_legacy_ferro_headers_are_ignored(self, client, httpx_mock: HTTPXMock): + body = {k: v for k, v in COMPLETION_RESPONSE.items() if k != "provider"} + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + json=body, + headers={"x-trace-id": "legacy", "x-request-id": TRACE_ID}, + ) + response = _chat(client) + assert response.trace_id == TRACE_ID + assert response.provider is None + + async def test_async_headers_populate_inference_bodies( + self, async_client, httpx_mock: HTTPXMock + ): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + json=COMPLETION_RESPONSE, + headers={"X-Request-ID": TRACE_ID, "X-Gateway-Overhead-Ms": "3"}, + ) + response = await async_client.chat.completions.create( + model="gpt-4o", messages=[{"role": "user", "content": "Hi"}] + ) + assert response.trace_id == TRACE_ID + assert response.gateway_overhead_ms == 3.0 + + def test_non_inference_bodies_are_not_touched(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", + url=f"{BASE_URL}/admin/dashboard", + json={"providers": {"enabled": 1}}, + headers={"X-Request-ID": TRACE_ID, "X-Gateway-Provider": "openai"}, + ) + assert client.admin.dashboard() == {"providers": {"enabled": 1}} + + +class TestErrorMapping: + @pytest.mark.parametrize( + ("status", "code", "exc"), + [ + (401, "invalid_api_key", FerroAuthError), + (402, "insufficient_quota", FerroBudgetExceededError), + (403, "insufficient_scope", FerroPermissionError), + (404, "model_not_found", FerroNotFoundError), + (429, "rate_limit_exceeded", FerroRateLimitError), + (502, "upstream_error", FerroServerError), + (400, "invalid_request", FerroAPIError), + ], + ) + def test_status_maps_to_exception(self, client, httpx_mock: HTTPXMock, status, code, exc): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=status, + json={"error": {"message": "nope", "type": "x", "code": code}}, + headers={"X-Request-ID": TRACE_ID}, + ) + with pytest.raises(exc) as exc_info: + _chat(client) + assert exc_info.value.status_code == status + assert exc_info.value.code == code + assert exc_info.value.request_id == TRACE_ID + assert "nope" in str(exc_info.value) + + def test_rate_limit_carries_retry_after(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=429, + json={"error": {"message": "slow down"}}, + headers={"Retry-After": "7"}, + ) + with pytest.raises(FerroRateLimitError) as exc_info: + _chat(client) + assert exc_info.value.retry_after == 7.0 + + def test_request_id_from_body_trace_id(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=400, + json={"error": {"message": "Bad request"}, "trace_id": "trace-xyz"}, + ) + with pytest.raises(FerroAPIError) as exc_info: + _chat(client) + assert exc_info.value.request_id == "trace-xyz" + + +class TestRetries: + @pytest.fixture(autouse=True) + def _no_jitter(self, monkeypatch): + # Full jitter picks uniformly from [0, delay]; pin it to the upper bound. + monkeypatch.setattr("ferrolabsai.client.random.uniform", lambda _lo, hi: hi) + + def test_sync_retries_connect_errors_with_backoff( + self, monkeypatch, client, httpx_mock: HTTPXMock + ): + sleeps: list[float] = [] + monkeypatch.setattr("ferrolabsai.client.time.sleep", sleeps.append) + client.max_retries = 2 + httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + assert _chat(client).id == "chatcmpl-abc123" + assert sleeps == [0.5, 1.0] + + def test_sync_retry_exhaustion_raises_connection_error( + self, monkeypatch, client, httpx_mock: HTTPXMock + ): + monkeypatch.setattr("ferrolabsai.client.time.sleep", lambda _s: None) + client.max_retries = 1 + httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) + httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) + with pytest.raises(FerroConnectionError, match="Cannot reach"): + _chat(client) + assert len(httpx_mock.get_requests()) == 2 + + @pytest.mark.parametrize("status", [408, 429, 500, 503]) + def test_sync_retries_retryable_statuses( + self, monkeypatch, client, httpx_mock: HTTPXMock, status + ): + sleeps: list[float] = [] + monkeypatch.setattr("ferrolabsai.client.time.sleep", sleeps.append) + client.max_retries = 1 + httpx_mock.add_response( + method="POST", url=CHAT_URL, status_code=status, json={"error": {"message": "x"}} + ) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + assert _chat(client).id == "chatcmpl-abc123" + assert sleeps == [0.5] + + def test_sync_honours_retry_after_capped(self, monkeypatch, client, httpx_mock: HTTPXMock): + sleeps: list[float] = [] + monkeypatch.setattr("ferrolabsai.client.time.sleep", sleeps.append) + client.max_retries = 2 + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=429, + json={"error": {"message": "x"}}, + headers={"Retry-After": "2"}, + ) + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=429, + json={"error": {"message": "x"}}, + headers={"Retry-After": "600"}, + ) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + _chat(client) + assert sleeps == [2.0, 30.0] + + def test_sync_does_not_retry_client_errors(self, client, httpx_mock: HTTPXMock): + client.max_retries = 2 + httpx_mock.add_response( + method="POST", url=CHAT_URL, status_code=400, json={"error": {"message": "bad"}} + ) + with pytest.raises(FerroAPIError): + _chat(client) + assert len(httpx_mock.get_requests()) == 1 + + def test_sync_raises_after_last_retryable_status( + self, monkeypatch, client, httpx_mock: HTTPXMock + ): + monkeypatch.setattr("ferrolabsai.client.time.sleep", lambda _s: None) + client.max_retries = 1 + for _ in range(2): + httpx_mock.add_response( + method="POST", url=CHAT_URL, status_code=503, json={"error": {"message": "x"}} + ) + with pytest.raises(FerroServerError): + _chat(client) + + def test_streaming_is_never_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 2 + httpx_mock.add_response( + method="POST", url=CHAT_URL, status_code=503, json={"error": {"message": "x"}} + ) + with pytest.raises(FerroServerError): + client.chat.completions.create( + model="gpt-4o", messages=[{"role": "user", "content": "Hi"}], stream=True + ) + assert len(httpx_mock.get_requests()) == 1 + + async def test_async_retries_with_backoff( + self, monkeypatch, async_client, httpx_mock: HTTPXMock + ): + sleeps: list[float] = [] + + async def fake_sleep(delay: float) -> None: + sleeps.append(delay) + + monkeypatch.setattr("ferrolabsai.client.asyncio.sleep", fake_sleep) + async_client.max_retries = 2 + httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=429, + json={"error": {"message": "x"}}, + headers={"Retry-After": "3"}, + ) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + response = await async_client.chat.completions.create( + model="gpt-4o", messages=[{"role": "user", "content": "Hi"}] + ) + assert response.id == "chatcmpl-abc123" + assert sleeps == [0.5, 3.0] + + async def test_async_retry_exhaustion_raises_connection_error( + self, monkeypatch, async_client, httpx_mock: HTTPXMock + ): + async def fake_sleep(_delay: float) -> None: + pass + + monkeypatch.setattr("ferrolabsai.client.asyncio.sleep", fake_sleep) + async_client.max_retries = 1 + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + with pytest.raises(FerroConnectionError, match="timed out"): + await async_client.chat.completions.create( + model="gpt-4o", messages=[{"role": "user", "content": "Hi"}] + ) + + +class TestGatewayEndpoints: + @pytest.mark.parametrize( + ("method", "path", "status", "body"), + [ + ("health", "/health", 200, {"status": "ok", "version": "1.4.5"}), + ("health", "/health", 503, {"status": "no_providers"}), + ("ready", "/readyz", 200, {"status": "ready", "targets": []}), + ("ready", "/readyz", 503, {"status": "not_ready", "reason": "x"}), + ("live", "/livez", 200, {"status": "ok"}), + ], + ) + def test_probes_return_json_on_200_and_503( + self, client, httpx_mock: HTTPXMock, method, path, status, body + ): + httpx_mock.add_response( + method="GET", url=f"{BASE_URL}{path}", status_code=status, json=body + ) + assert getattr(client, method)() == body + + async def test_async_probes(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{BASE_URL}/health", status_code=503, json={"status": "degraded"} + ) + assert await async_client.health() == {"status": "degraded"} + + def test_capabilities(self, client, httpx_mock: HTTPXMock): + caps = {"providers": {"openai": {"tools": "forward"}}, "image_response_formats": {}} + httpx_mock.add_response(method="GET", url=f"{BASE_URL}/v1/capabilities", json=caps) + assert client.capabilities() == caps + + def test_rerank(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{BASE_URL}/v1/rerank", json={"results": [{"index": 1}]} + ) + result = client.rerank(model="rerank-v3", query="q", documents=["a", "b"], top_n=1) + assert result["results"][0]["index"] == 1 + import json + + body = json.loads(httpx_mock.get_requests()[0].content) + assert body == {"model": "rerank-v3", "query": "q", "documents": ["a", "b"], "top_n": 1} + + async def test_async_rerank_and_capabilities(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=f"{BASE_URL}/v1/rerank", json={"results": []}) + httpx_mock.add_response(method="GET", url=f"{BASE_URL}/v1/capabilities", json={}) + assert await async_client.rerank(model="m", query="q", documents=[]) == {"results": []} + assert await async_client.capabilities() == {} + + def test_moderations(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{BASE_URL}/v1/moderations", json={"results": [{"flagged": False}]} + ) + result = client.moderations.create(input="hello", model="omni-moderation-latest") + assert result["results"][0]["flagged"] is False + import json + + body = json.loads(httpx_mock.get_requests()[0].content) + assert body == {"input": "hello", "model": "omni-moderation-latest"} + + async def test_async_moderations(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{BASE_URL}/v1/moderations", json={"results": []} + ) + assert await async_client.moderations.create(input=["a"]) == {"results": []} diff --git a/tests/test_resources.py b/tests/test_resources.py new file mode 100644 index 0000000..9abe46f --- /dev/null +++ b/tests/test_resources.py @@ -0,0 +1,179 @@ +"""Embeddings, images, model catalog, and the Responses API.""" + +from __future__ import annotations + +import json + +import pytest +from pytest_httpx import HTTPXMock + +from ferrolabsai.exceptions import FerroNotFoundError +from ferrolabsai.types import ModelInfo, Response + +from .conftest import BASE_URL, EMBEDDING_RESPONSE, IMAGE_RESPONSE, MODELS_RESPONSE, TRACE_ID + +MODELS_URL = f"{BASE_URL}/v1/models" + + +class TestEmbeddings: + def test_create(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=f"{BASE_URL}/v1/embeddings", + json=EMBEDDING_RESPONSE, + headers={"X-Request-ID": TRACE_ID}, + ) + response = client.embeddings.create( + model="text-embedding-3-small", input=["Hello", "world"], dimensions=3 + ) + assert len(response.data) == 2 + assert response.data[0].embedding == [0.1, 0.2, 0.3] + assert response.model == "text-embedding-3-small" + assert response.trace_id == TRACE_ID + assert json.loads(httpx_mock.get_requests()[0].content)["dimensions"] == 3 + + async def test_async_create(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{BASE_URL}/v1/embeddings", json=EMBEDDING_RESPONSE + ) + response = await async_client.embeddings.create(model="text-embedding-3-small", input="Hi") + assert response.data[1].embedding == [0.4, 0.5, 0.6] + assert json.loads(httpx_mock.get_requests()[0].content)["input"] == "Hi" + + +class TestImages: + def test_generate(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{BASE_URL}/v1/images/generations", json=IMAGE_RESPONSE + ) + image = client.images.generate(model="dall-e-3", prompt="A gateway", size="1024x1024") + assert image.data[0].url == "https://example.com/image.png" + assert image.data[0].revised_prompt == "A polished prompt" + assert json.loads(httpx_mock.get_requests()[0].content)["size"] == "1024x1024" + + async def test_async_generate(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=f"{BASE_URL}/v1/images/generations", json=IMAGE_RESPONSE + ) + image = await async_client.images.generate(model="dall-e-3", prompt="A gateway") + assert image.created == 1700000000 + + +class TestModels: + """The gateway serves only GET /v1/models and ignores its query string; + every lookup and filter is client-side over that one catalog fetch.""" + + def test_list_returns_enriched_model_info(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + models = client.models.list() + assert len(models) == 3 + gpt = models[0] + assert isinstance(gpt, ModelInfo) + assert gpt.id == "gpt-4o" + assert gpt.owned_by == "openai" + assert gpt.provider == "openai" + assert gpt.mode == "chat" + assert gpt.context_window == 128000 + assert gpt.max_output_tokens == 16384 + assert gpt.capabilities == ["vision", "function_calling", "streaming"] + assert gpt.status == "active" + assert gpt.deprecated is False + assert models[1].deprecated is True + assert models[2].capabilities == [] + assert not hasattr(gpt, "input_cost_per_token") + + def test_list_filters_client_side(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + assert [m.id for m in client.models.list(provider="anthropic")] == [ + "claude-3-5-sonnet-20241022" + ] + assert [m.id for m in client.models.list(capability="vision")] == ["gpt-4o"] + assert all(str(r.url) == MODELS_URL for r in httpx_mock.get_requests()) + + def test_retrieve_is_a_client_side_lookup(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + assert client.models.retrieve("gpt-4o").context_window == 128000 + assert [str(r.url) for r in httpx_mock.get_requests()] == [MODELS_URL] + + def test_retrieve_unknown_raises_locally(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + with pytest.raises(FerroNotFoundError) as exc_info: + client.models.retrieve("nonexistent-model") + assert exc_info.value.status_code == 404 + assert exc_info.value.code == "model_not_found" + assert [str(r.url) for r in httpx_mock.get_requests()] == [MODELS_URL] + + def test_search_is_case_insensitive_substring(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + assert [m.id for m in client.models.search("CLAUDE")] == ["claude-3-5-sonnet-20241022"] + + async def test_async_models(self, async_client, httpx_mock: HTTPXMock): + for _ in range(3): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + assert len(await async_client.models.list()) == 3 + assert (await async_client.models.retrieve("gpt-4o")).owned_by == "openai" + assert [m.id for m in await async_client.models.search("embedding")] == [ + "text-embedding-3-small" + ] + with pytest.raises(FerroNotFoundError): + httpx_mock.add_response(method="GET", url=MODELS_URL, json=MODELS_RESPONSE) + await async_client.models.retrieve("nope") + + +RESPONSE_BODY = { + "id": "resp_1", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": "gpt-4o", + "output": [{"type": "message", "content": [{"type": "output_text", "text": "hi"}]}], + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, +} + + +class TestResponses: + def test_create(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=f"{BASE_URL}/v1/responses", + json=RESPONSE_BODY, + headers={"X-Request-ID": TRACE_ID, "X-Gateway-Provider": "openai"}, + ) + response = client.responses.create(model="gpt-4o", input="hi") + assert isinstance(response, Response) + assert response.id == "resp_1" + assert response.status == "completed" + assert response.output[0]["type"] == "message" + assert response.usage == {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2} + assert response.trace_id == TRACE_ID + assert response.provider == "openai" + assert response.raw["model"] == "gpt-4o" + assert json.loads(httpx_mock.get_requests()[0].content) == { + "model": "gpt-4o", + "input": "hi", + } + + def test_retrieve_and_delete(self, client, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="GET", url=f"{BASE_URL}/v1/responses/resp_1", json=RESPONSE_BODY + ) + httpx_mock.add_response( + method="DELETE", + url=f"{BASE_URL}/v1/responses/resp_1", + json={"id": "resp_1", "object": "response", "deleted": True}, + ) + assert client.responses.retrieve("resp_1").id == "resp_1" + assert client.responses.delete("resp_1")["deleted"] is True + + async def test_async_create_retrieve_delete(self, async_client, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=f"{BASE_URL}/v1/responses", json=RESPONSE_BODY) + httpx_mock.add_response( + method="GET", url=f"{BASE_URL}/v1/responses/resp_1", json=RESPONSE_BODY + ) + httpx_mock.add_response( + method="DELETE", url=f"{BASE_URL}/v1/responses/resp_1", json={"deleted": True} + ) + assert (await async_client.responses.create(model="gpt-4o", input="hi")).id == "resp_1" + assert (await async_client.responses.retrieve("resp_1")).model == "gpt-4o" + assert await async_client.responses.delete("resp_1") == {"deleted": True} diff --git a/tests/test_sdk.py b/tests/test_sdk.py deleted file mode 100644 index e6ff7ad..0000000 --- a/tests/test_sdk.py +++ /dev/null @@ -1,1038 +0,0 @@ -""" -ferrolabsai SDK test suite. - -Tests use pytest-httpx to mock the HTTP layer — no real gateway needed. -Run: pytest tests/ -v -""" - -from __future__ import annotations - -import json - -import httpx -import pytest -from pytest_httpx import HTTPXMock - -from ferrolabsai import AsyncFerroClient, FerroClient -from ferrolabsai.exceptions import ( - FerroAPIError, - FerroAuthError, - FerroNotFoundError, - FerroRateLimitError, - FerroServerError, - FerroStreamError, -) - -BASE_URL = "http://localhost:8080" -API_KEY = "sk-ferro-testkey123" - - -@pytest.fixture -def client(): - return FerroClient(api_key=API_KEY, base_url=BASE_URL) - - -@pytest.fixture -def async_client(): - return AsyncFerroClient(api_key=API_KEY, base_url=BASE_URL) - - -# ------------------------------------------------------------------ -# Client instantiation -# ------------------------------------------------------------------ - - -class TestClientInit: - def test_requires_api_key(self): - import os - - os.environ.pop("FERRO_API_KEY", None) - os.environ.pop("OPENAI_API_KEY", None) - with pytest.raises(FerroAuthError): - FerroClient() - - def test_reads_from_env(self, monkeypatch): - monkeypatch.setenv("FERRO_API_KEY", "sk-ferro-envkey") - c = FerroClient() - assert c.api_key == "sk-ferro-envkey" - - def test_falls_back_to_openai_env(self, monkeypatch): - monkeypatch.delenv("FERRO_API_KEY", raising=False) - monkeypatch.setenv("OPENAI_API_KEY", "sk-openai-compat") - c = FerroClient() - assert c.api_key == "sk-openai-compat" - - def test_strips_trailing_slash(self): - c = FerroClient(api_key=API_KEY, base_url="https://localhost:8080/") - assert c.base_url == "https://localhost:8080" - - def test_rejects_negative_max_retries(self): - with pytest.raises(ValueError, match="max_retries must be >= 0"): - FerroClient(api_key=API_KEY, max_retries=-1) - - def test_async_rejects_negative_max_retries(self): - with pytest.raises(ValueError, match="max_retries must be >= 0"): - AsyncFerroClient(api_key=API_KEY, max_retries=-1) - - @pytest.mark.parametrize("client_cls", [FerroClient, AsyncFerroClient]) - def test_accepts_zero_max_retries(self, client_cls): - client = client_cls(api_key=API_KEY, max_retries=0) - assert client.max_retries == 0 - - @pytest.mark.parametrize("client_cls", [FerroClient, AsyncFerroClient]) - @pytest.mark.parametrize("invalid_value", [1.5, True]) - def test_rejects_non_integer_max_retries(self, client_cls, invalid_value): - with pytest.raises(TypeError, match="max_retries must be an integer"): - client_cls(api_key=API_KEY, max_retries=invalid_value) - - def test_has_expected_namespaces(self, client): - assert hasattr(client, "chat") - assert hasattr(client.chat, "completions") - assert hasattr(client, "embeddings") - assert hasattr(client, "images") - assert hasattr(client, "models") - assert hasattr(client, "admin") - # Admin sub-resources mirror the OSS gateway /admin/* surface. - assert hasattr(client.admin, "keys") - assert hasattr(client.admin, "config") - assert hasattr(client.admin, "logs") - assert hasattr(client.admin, "providers") - assert hasattr(client.admin, "plugins") - - -# ------------------------------------------------------------------ -# Chat completions -# ------------------------------------------------------------------ - -COMPLETION_RESPONSE = { - "id": "chatcmpl-abc123", - "object": "chat.completion", - "created": 1700000000, - "model": "gpt-4o", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "Hello from Ferro!"}, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15, - "cost_usd": 0.000075, - }, -} - - -class TestChatCompletions: - def test_basic_create(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - response = client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - ) - assert response.id == "chatcmpl-abc123" - assert response.model == "gpt-4o" - assert response.content == "Hello from Ferro!" - assert response.usage.cost_usd == 0.000075 - - def test_success_metadata_can_come_from_headers(self, client, httpx_mock: HTTPXMock): - body = dict(COMPLETION_RESPONSE) - body["usage"] = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=body, - headers={ - "X-Request-ID": "trace-from-header", - "x-ferro-provider": "openai", - "x-ferro-latency-ms": "42", - "x-ferro-cost-usd": "0.000075", - }, - ) - response = client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - ) - assert response.trace_id == "trace-from-header" - assert response.provider == "openai" - assert response.latency_ms == 42 - assert response.usage is not None - assert response.usage.cost_usd == 0.000075 - assert response.usage.provider == "openai" - - def test_sends_correct_headers(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - request = httpx_mock.get_requests()[0] - assert request.headers["authorization"] == f"Bearer {API_KEY}" - assert request.headers["content-type"] == "application/json" - - def test_passes_optional_params(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - temperature=0.7, - max_tokens=100, - user="user_123", - ) - body = json.loads(httpx_mock.get_requests()[0].content) - assert body["temperature"] == 0.7 - assert body["max_tokens"] == 100 - assert body["user"] == "user_123" - - def test_ferro_template_id_forwarded(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - client.chat.completions.create( - model="gpt-4o", - messages=[], - template_id="tmpl_support", - template_variables={"plan": "pro"}, - ) - body = json.loads(httpx_mock.get_requests()[0].content) - assert body["template_id"] == "tmpl_support" - assert body["template_variables"] == {"plan": "pro"} - - def test_streaming_returns_iterator(self, client, httpx_mock: HTTPXMock): - sse_data = ( - 'data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"gpt-4o",' - '"choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}\n\n' - 'data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"gpt-4o",' - '"choices":[{"index":0,"delta":{"content":" world"},"finish_reason":"stop"}]}\n\n' - "data: [DONE]\n\n" - ) - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - content=sse_data.encode(), - ) - chunks = list( - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - stream=True, - ) - ) - assert len(chunks) == 2 - assert chunks[0].choices[0].delta.content == "Hello" - assert chunks[1].choices[0].delta.content == " world" - assert chunks[1].choices[0].finish_reason == "stop" - - -# ------------------------------------------------------------------ -# Embeddings -# ------------------------------------------------------------------ - -EMBEDDING_RESPONSE = { - "object": "list", - "data": [ - {"index": 0, "object": "embedding", "embedding": [0.1, 0.2, 0.3]}, - {"index": 1, "object": "embedding", "embedding": [0.4, 0.5, 0.6]}, - ], - "model": "text-embedding-3-small", - "usage": {"prompt_tokens": 8, "total_tokens": 8}, -} - -IMAGE_RESPONSE = { - "created": 1700000000, - "data": [ - { - "url": "https://example.com/image.png", - "revised_prompt": "A polished image prompt", - } - ], -} - - -class TestEmbeddings: - def test_create(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/embeddings", - json=EMBEDDING_RESPONSE, - ) - response = client.embeddings.create( - model="text-embedding-3-small", - input=["Hello", "world"], - ) - assert len(response.data) == 2 - assert response.data[0].embedding == [0.1, 0.2, 0.3] - assert response.model == "text-embedding-3-small" - - -# ------------------------------------------------------------------ -# Models -# ------------------------------------------------------------------ - -MODELS_RESPONSE = { - "data": [ - { - "id": "gpt-4o", - "object": "model", - "owned_by": "openai", - "context_window": 128000, - "input_cost_per_token": 0.0000025, - "output_cost_per_token": 0.00001, - }, - { - "id": "claude-3-5-sonnet-20241022", - "object": "model", - "owned_by": "anthropic", - "context_window": 200000, - }, - ] -} - - -class TestModels: - def test_list(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/v1/models", - json=MODELS_RESPONSE, - ) - models = client.models.list() - assert len(models) == 2 - assert models[0].id == "gpt-4o" - assert models[0].provider == "openai" - assert models[0].context_window == 128000 - - def test_retrieve(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/v1/models/gpt-4o", - json=MODELS_RESPONSE["data"][0], - ) - model = client.models.retrieve("gpt-4o") - assert model.id == "gpt-4o" - assert model.input_cost_per_token == 0.0000025 - - -# ------------------------------------------------------------------ -# Admin — Keys -# ------------------------------------------------------------------ - - -class TestAdminKeys: - def test_create_key(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/admin/keys", - json={ - "id": "key_abc", - "name": "test-key", - "key": "fgw_fullkeyvalue", - "scopes": ["admin"], - "active": True, - "created_at": "2026-04-01T00:00:00Z", - }, - ) - key = client.admin.keys.create(name="test-key", scopes=["admin"]) - assert key.key == "fgw_fullkeyvalue" - assert key.name == "test-key" - assert key.scopes == ["admin"] - - def test_revoke_key(self, client, httpx_mock: HTTPXMock): - # OSS revoke is POST /admin/keys/{id}/revoke (record kept for audit) - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/admin/keys/key_abc/revoke", - status_code=204, - content=b"", - ) - client.admin.keys.revoke("key_abc") # should not raise - - def test_delete_key(self, client, httpx_mock: HTTPXMock): - # delete() permanently removes the key record - httpx_mock.add_response( - method="DELETE", - url=f"{BASE_URL}/admin/keys/key_abc", - status_code=204, - content=b"", - ) - client.admin.keys.delete("key_abc") # should not raise - - def test_rotate_key(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/admin/keys/key_abc/rotate", - json={ - "id": "key_abc", - "name": "prod", - "key": "fgw_newkeyvalue", - "scopes": ["admin"], - "active": True, - "created_at": "2026-04-01T00:00:00Z", - }, - ) - rotated = client.admin.keys.rotate("key_abc") - assert rotated.key == "fgw_newkeyvalue" - assert rotated.id == "key_abc" - - def test_list_keys(self, client, httpx_mock: HTTPXMock): - # OSS returns a bare JSON array from GET /admin/keys - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/keys", - json=[ - { - "id": "key_1", - "name": "prod", - "scopes": ["admin"], - "active": True, - "usage_count": 12, - "created_at": "2026-01-01T00:00:00Z", - } - ], - ) - keys = client.admin.keys.list() - assert len(keys) == 1 - assert keys[0].id == "key_1" - assert keys[0].usage_count == 12 - - def test_keys_usage(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/keys/usage?limit=10&offset=0&sort=usage", - json={ - "data": [{"id": "key_1", "name": "prod", "usage_count": 100, "active": True}], - "summary": { - "total_keys": 1, - "active_keys": 1, - "total_usage": 100, - "returned_keys": 1, - }, - "filters": {"limit": 10, "offset": 0, "sort": "usage", "active": "", "since": ""}, - }, - ) - result = client.admin.keys.usage(limit=10) - assert result["summary"]["total_usage"] == 100 - - -class TestAdminConfig: - def test_get_config(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/config", - json={ - "strategy": {"mode": "fallback"}, - "targets": [ - {"virtual_key": "openai", "weight": 1}, - {"virtual_key": "anthropic", "weight": 1}, - ], - "plugins": [], - }, - ) - cfg = client.admin.config.get() - assert cfg.strategy == {"mode": "fallback"} - assert len(cfg.targets) == 2 - - def test_update_config(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="PUT", - url=f"{BASE_URL}/admin/config", - json={"status": "updated"}, - ) - result = client.admin.config.update( - { - "strategy": {"mode": "single"}, - "targets": [{"virtual_key": "openai"}], - } - ) - assert result["status"] == "updated" - - def test_history(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/config/history", - json={ - "data": [ - { - "version": 1, - "updated_at": "2026-04-01T00:00:00Z", - "config": {"strategy": {"mode": "single"}, "targets": []}, - }, - { - "version": 2, - "updated_at": "2026-04-02T00:00:00Z", - "config": {"strategy": {"mode": "fallback"}, "targets": []}, - }, - ], - "summary": {"total_versions": 2}, - }, - ) - history = client.admin.config.history() - assert len(history) == 2 - assert history[0].version == 1 - assert history[1].config.strategy == {"mode": "fallback"} - - def test_rollback(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/admin/config/rollback/1", - json={"status": "rolled_back", "rolled_back_to": 1, "current_history_size": 3}, - ) - result = client.admin.config.rollback(1) - assert result["status"] == "rolled_back" - assert result["rolled_back_to"] == 1 - - -class TestAdminDashboard: - def test_get_dashboard(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/dashboard", - json={ - "providers": {"enabled": 3, "disabled": 1}, - "keys": {"active": 5, "revoked": 2}, - "requests": {"total": 128, "errors": 4}, - }, - ) - dashboard = client.admin.dashboard() - assert dashboard["providers"]["enabled"] == 3 - assert dashboard["keys"]["active"] == 5 - assert dashboard["requests"]["total"] == 128 - - -class TestAdminLogs: - def test_list_logs(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/logs?limit=10&offset=0", - json={ - "data": [{"trace_id": "t1", "model": "gpt-4o", "provider": "openai"}], - "summary": {"total_entries": 1, "returned_entries": 1}, - "filters": {"limit": 10, "offset": 0, "stage": "", "model": "", "provider": ""}, - }, - ) - result = client.admin.logs.list(limit=10) - assert result["summary"]["total_entries"] == 1 - assert result["data"][0]["trace_id"] == "t1" - - def test_logs_stats(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/logs/stats", - json={"total": 42, "errors": 1}, - ) - stats = client.admin.logs.stats() - assert stats["total"] == 42 - - -class TestAdminPlugins: - def test_list_plugins_accepts_bare_array(self, client, httpx_mock: HTTPXMock): - plugins = [ - {"name": "cache", "enabled": True}, - {"name": "logger", "enabled": False}, - ] - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/plugins", - json=plugins, - ) - - assert client.admin.plugins.list() == plugins - - def test_list_plugins_accepts_data_wrapper(self, client, httpx_mock: HTTPXMock): - plugins = [{"name": "ratelimit", "enabled": True}] - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/plugins", - json={"data": plugins}, - ) - - assert client.admin.plugins.list() == plugins - - def test_list_plugins_accepts_plugins_wrapper(self, client, httpx_mock: HTTPXMock): - plugins = [{"name": "logger", "enabled": True}] - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/plugins", - json={"plugins": plugins}, - ) - - assert client.admin.plugins.list() == plugins - - -# ------------------------------------------------------------------ -# Error handling -# ------------------------------------------------------------------ - - -class TestErrorHandling: - def test_401_raises_auth_error(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=401, - json={"error": {"message": "Invalid API key", "code": "invalid_api_key"}}, - ) - with pytest.raises(FerroAuthError) as exc_info: - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - assert "Invalid API key" in str(exc_info.value) - - def test_429_raises_rate_limit_error(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=429, - json={"error": {"message": "Rate limit exceeded"}}, - ) - with pytest.raises(FerroRateLimitError): - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - - def test_404_raises_not_found(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/v1/models/nonexistent-model", - status_code=404, - json={"error": {"message": "Model not found"}}, - ) - with pytest.raises(FerroNotFoundError): - client.models.retrieve("nonexistent-model") - - -# ------------------------------------------------------------------ -# Context manager -# ------------------------------------------------------------------ - - -class TestContextManager: - def test_sync_context_manager(self): - with FerroClient(api_key=API_KEY) as c: - assert c.api_key == API_KEY - - @pytest.mark.asyncio - async def test_async_context_manager(self): - async with AsyncFerroClient(api_key=API_KEY) as c: - assert c.api_key == API_KEY - - -# ------------------------------------------------------------------ -# P0-1: Streaming error handling routes through _raise_api_error -# ------------------------------------------------------------------ - - -class TestStreamingErrorHandling: - def test_sync_stream_raises_ferro_error_not_httpx(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=401, - json={"error": {"message": "Invalid API key"}}, - ) - with pytest.raises(FerroAuthError, match="Invalid API key"): - list( - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - stream=True, - ) - ) - - def test_sync_stream_429_raises_rate_limit(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=429, - json={"error": {"message": "Rate limit exceeded"}}, - ) - with pytest.raises(FerroRateLimitError): - list( - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - stream=True, - ) - ) - - def test_sync_stream_500_raises_server_error(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=500, - json={"error": {"message": "Internal error"}}, - ) - with pytest.raises(FerroServerError): - list( - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - stream=True, - ) - ) - - def test_sync_stream_malformed_chunk_raises_stream_error(self, client, httpx_mock: HTTPXMock): - sse_data = "data: {not valid json}\n\n" - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - content=sse_data.encode(), - ) - with pytest.raises(FerroStreamError, match="Malformed SSE chunk"): - list( - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - stream=True, - ) - ) - - @pytest.mark.asyncio - async def test_async_stream_malformed_chunk_raises_stream_error( - self, async_client, httpx_mock: HTTPXMock - ): - sse_data = "data: {not valid json}\n\n" - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - content=sse_data.encode(), - ) - with pytest.raises(FerroStreamError, match="Malformed SSE chunk"): - async for _ in await async_client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - stream=True, - ): - pass - - -# ------------------------------------------------------------------ -# P0-2: Async completions has parity params -# ------------------------------------------------------------------ - - -class TestAsyncCompletionsParams: - @pytest.mark.asyncio - async def test_forwards_frequency_presence_penalty(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - await async_client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - frequency_penalty=0.5, - presence_penalty=0.3, - ) - body = json.loads(httpx_mock.get_requests()[0].content) - assert body["frequency_penalty"] == 0.5 - assert body["presence_penalty"] == 0.3 - - @pytest.mark.asyncio - async def test_forwards_route_tag(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - await async_client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - route_tag="fast", - ) - body = json.loads(httpx_mock.get_requests()[0].content) - assert body["x_route_tag"] == "fast" - - @pytest.mark.asyncio - async def test_success_metadata_can_come_from_headers( - self, async_client, httpx_mock: HTTPXMock - ): - body = dict(COMPLETION_RESPONSE) - body["usage"] = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=body, - headers={ - "X-Request-ID": "async-trace-from-header", - "x-ferro-provider": "anthropic", - "x-ferro-latency-ms": "123", - "x-ferro-cost-usd": "0.001", - }, - ) - response = await async_client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - assert response.trace_id == "async-trace-from-header" - assert response.provider == "anthropic" - assert response.latency_ms == 123 - assert response.usage is not None - assert response.usage.cost_usd == 0.001 - - -# ------------------------------------------------------------------ -# P0-3: BYOC http_client merges auth headers -# ------------------------------------------------------------------ - - -class TestBYOCHttpClient: - def test_sync_byoc_merges_headers(self): - custom = httpx.Client(base_url=BASE_URL) - client = FerroClient(api_key=API_KEY, base_url=BASE_URL, http_client=custom) - assert client._http is custom - assert "Authorization" in custom.headers - assert custom.headers["Authorization"] == f"Bearer {API_KEY}" - custom.close() - - def test_async_byoc_merges_headers(self): - custom = httpx.AsyncClient(base_url=BASE_URL) - client = AsyncFerroClient(api_key=API_KEY, base_url=BASE_URL, http_client=custom) - assert client._http is custom - assert "Authorization" in custom.headers - assert custom.headers["Authorization"] == f"Bearer {API_KEY}" - - -# ------------------------------------------------------------------ -# P1-4: AsyncFerroClient has all namespaces -# ------------------------------------------------------------------ - - -class TestAsyncClientNamespaces: - def test_has_all_namespaces(self): - client = AsyncFerroClient(api_key=API_KEY) - assert hasattr(client, "chat") - assert hasattr(client.chat, "completions") - assert hasattr(client, "embeddings") - assert hasattr(client, "images") - assert hasattr(client, "models") - assert hasattr(client, "admin") - - -class TestAsyncResources: - @pytest.mark.asyncio - async def test_models_list_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/v1/models", - json=MODELS_RESPONSE, - ) - models = await async_client.models.list() - assert len(models) == 2 - assert models[0].id == "gpt-4o" - - @pytest.mark.asyncio - async def test_models_search_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/v1/models?search=claude", - json={"data": [MODELS_RESPONSE["data"][1]]}, - ) - models = await async_client.models.search("claude") - assert len(models) == 1 - assert models[0].provider == "anthropic" - - @pytest.mark.asyncio - async def test_images_generate_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/images/generations", - json=IMAGE_RESPONSE, - ) - image = await async_client.images.generate( - model="dall-e-3", - prompt="A gateway", - size="1024x1024", - ) - assert image.data[0].url == "https://example.com/image.png" - body = json.loads(httpx_mock.get_requests()[0].content) - assert body["size"] == "1024x1024" - - @pytest.mark.asyncio - async def test_admin_health_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/health", - json={"status": "ok"}, - ) - health = await async_client.admin.health() - assert health["status"] == "ok" - - @pytest.mark.asyncio - async def test_admin_keys_list_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/keys", - json=[ - { - "id": "key_1", - "name": "prod", - "scopes": ["admin"], - "active": True, - "created_at": "2026-01-01T00:00:00Z", - } - ], - ) - keys = await async_client.admin.keys.list() - assert keys[0].id == "key_1" - - @pytest.mark.asyncio - async def test_admin_keys_delete_accepts_204(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="DELETE", - url=f"{BASE_URL}/admin/keys/key_1", - status_code=204, - content=b"", - ) - await async_client.admin.keys.delete("key_1") - - @pytest.mark.asyncio - async def test_admin_keys_revoke_accepts_204(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/admin/keys/key_1/revoke", - status_code=204, - content=b"", - ) - await async_client.admin.keys.revoke("key_1") - - @pytest.mark.asyncio - async def test_admin_config_history_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/config/history", - json={ - "data": [ - { - "version": 1, - "updated_at": "2026-04-01T00:00:00Z", - "config": {"strategy": {"mode": "single"}, "targets": []}, - } - ] - }, - ) - history = await async_client.admin.config.history() - assert history[0].version == 1 - - @pytest.mark.asyncio - async def test_admin_providers_list_is_awaitable(self, async_client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="GET", - url=f"{BASE_URL}/admin/providers", - json={"data": [{"name": "openai"}]}, - ) - providers = await async_client.admin.providers.list() - assert providers == [{"name": "openai"}] - - -class TestRetryBackoff: - def test_sync_retries_use_backoff(self, monkeypatch, client, httpx_mock: HTTPXMock): - sleeps: list[float] = [] - monkeypatch.setattr("ferrolabsai.client.time.sleep", sleeps.append) - client.max_retries = 1 - httpx_mock.add_exception( - httpx.ConnectError("connection refused"), - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - ) - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - - response = client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - - assert response.id == "chatcmpl-abc123" - assert sleeps == [0.5] - - @pytest.mark.asyncio - async def test_async_retries_use_backoff( - self, - monkeypatch, - async_client, - httpx_mock: HTTPXMock, - ): - sleeps: list[float] = [] - - async def fake_sleep(delay: float) -> None: - sleeps.append(delay) - - monkeypatch.setattr("ferrolabsai.client.asyncio.sleep", fake_sleep) - async_client.max_retries = 1 - httpx_mock.add_exception( - httpx.ConnectError("connection refused"), - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - ) - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=COMPLETION_RESPONSE, - ) - - response = await async_client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - - assert response.id == "chatcmpl-abc123" - assert sleeps == [0.5] - - -# ------------------------------------------------------------------ -# P1-7: request_id populated from response headers -# ------------------------------------------------------------------ - - -class TestRequestIdPropagation: - def test_request_id_from_header(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=500, - json={"error": {"message": "Server error"}}, - headers={"x-request-id": "req-abc-123"}, - ) - with pytest.raises(FerroServerError) as exc_info: - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - assert exc_info.value.request_id == "req-abc-123" - - def test_request_id_from_body_trace_id(self, client, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - status_code=400, - json={"error": {"message": "Bad request"}, "trace_id": "trace-xyz"}, - ) - with pytest.raises(FerroAPIError) as exc_info: - client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hi"}], - ) - assert exc_info.value.request_id == "trace-xyz" From 97eb2190ebe8c879daa1224ed2fa729fa7b10c71 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 12:52:15 +0530 Subject: [PATCH 02/14] feat(langchain): langchain-ferrolabsai 0.2.0 on ferrolabsai 0.3 - response_metadata = {model, id, trace_id, provider, gateway_overhead_ms}, None-stripped; latency_ms/cost_usd/cache_hit removed (never emitted). - Drop route_tag/template_id/template_variables (never read by the gateway). - Add _agenerate/_astream on AsyncFerroClient, aembed_documents/aembed_query. - Add with_structured_output() via response_format json_schema. - trace_id on the first streamed chunk; tests mirror real gateway headers. --- .../langchain-ferrolabsai/CHANGELOG.md | 29 +- integrations/langchain-ferrolabsai/README.md | 72 +++-- .../langchain_ferrolabsai/__init__.py | 11 +- .../langchain_ferrolabsai/chat_models.py | 218 ++++++++++----- .../langchain_ferrolabsai/embeddings.py | 52 +++- .../langchain_ferrolabsai/llms.py | 3 - .../langchain-ferrolabsai/pyproject.toml | 7 +- .../langchain-ferrolabsai/tests/conftest.py | 58 ++-- .../tests/test_chat_models.py | 262 +++++++++--------- .../tests/test_embeddings.py | 30 ++ 10 files changed, 465 insertions(+), 277 deletions(-) diff --git a/integrations/langchain-ferrolabsai/CHANGELOG.md b/integrations/langchain-ferrolabsai/CHANGELOG.md index a536041..ead8372 100644 --- a/integrations/langchain-ferrolabsai/CHANGELOG.md +++ b/integrations/langchain-ferrolabsai/CHANGELOG.md @@ -9,13 +9,38 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Planned -- Async surfaces (`_agenerate`, `_astream`, `aembed_documents`, `aembed_query`). -- `with_structured_output()` helper for JSON-mode + Pydantic-schema responses. - Native multi-modal message support (image inputs) once the gateway exposes a stable contract. --- +## [0.2.0] — 2026-08-29 + +Requires `ferrolabsai >= 0.3.0` (the "truth release") and ai-gateway ≥ v1.4.0. + +### Breaking + +- `response_metadata` now contains exactly what the gateway provides: + `model`, `id`, `trace_id` (`X-Request-ID` header), `provider` (body field), + `gateway_overhead_ms` (`X-Gateway-Overhead-Ms` header). `latency_ms`, + `cost_usd`, and `cache_hit` are gone — the gateway never emitted them, so + they were always absent. +- Removed the `route_tag`, `template_id`, and `template_variables` fields from + `FerroChatModel` and `FerroLLM`. The gateway has never read them. + +### Added + +- Async surface: `FerroChatModel._agenerate` / `_astream` (so `ainvoke`, + `astream`, `abatch` run on `AsyncFerroClient` instead of a thread) and + `FerroEmbeddings.aembed_documents` / `aembed_query`. +- `FerroChatModel.with_structured_output(schema, include_raw=False)` via the + OpenAI-style `response_format={"type": "json_schema", ...}` path; accepts a + pydantic model class or a JSON-schema dict. +- Streaming: the first chunk carries `trace_id` in `response_metadata`. +- Python 3.13 classifier. + +--- + ## [0.1.0] — 2026-05-25 First functional release. Replaces the `0.0.1` placeholder. diff --git a/integrations/langchain-ferrolabsai/README.md b/integrations/langchain-ferrolabsai/README.md index b1ce926..1d6e67a 100644 --- a/integrations/langchain-ferrolabsai/README.md +++ b/integrations/langchain-ferrolabsai/README.md @@ -3,7 +3,9 @@ [![PyPI version](https://badge.fury.io/py/langchain-ferrolabsai.svg)](https://pypi.org/project/langchain-ferrolabsai/) [![License](https://img.shields.io/badge/license-Apache%202.0-blue.svg)](LICENSE) -LangChain integration for [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) — route LangChain chat, streaming, tool-calling, and embedding workloads across **30+ LLM providers** through a single OpenAI-compatible endpoint, with automatic fallback, load balancing, cost tracking, and observability. +LangChain integration for [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) — route LangChain chat, streaming, tool-calling, structured-output, and embedding workloads across **30 LLM providers** through a single OpenAI-compatible endpoint, with automatic fallback, load balancing, budgets, and observability. + +Compatibility: `langchain-ferrolabsai 0.2.x` ↔ `ferrolabsai ≥ 0.3.0` ↔ `ai-gateway ≥ v1.4.0`; `langchain-core ≥ 0.3` (tested on 1.x). --- @@ -24,23 +26,27 @@ from langchain_core.messages import HumanMessage llm = FerroChatModel( model="gpt-4o", base_url="http://localhost:8080", # any Ferro Labs AI Gateway instance - api_key="sk-ferro-...", + api_key="fgw_...", ) response = llm.invoke([HumanMessage(content="Hello, world")]) print(response.content) -print(response.response_metadata["provider"]) # which provider handled it -print(response.response_metadata["cost_usd"]) # cost for this request -print(response.response_metadata["latency_ms"]) # observed latency -print(response.response_metadata["trace_id"]) # gateway trace ID (x-trace-id) +print(response.response_metadata["provider"]) # which provider answered +print(response.response_metadata["trace_id"]) # gateway X-Request-ID +print(response.response_metadata.get("gateway_overhead_ms")) # gateway's own overhead ``` +`response_metadata` contains exactly what the gateway provides — `model`, `id`, +`trace_id`, `provider`, `gateway_overhead_ms` — with absent values stripped. +`trace_id` is the join key for `client.admin.logs.list()` and for the +gateway's observability exporters (LangSmith, Langfuse, Phoenix, …). + Swap providers without changing the model class — Ferro auto-routes by model name: ```python claude = FerroChatModel(model="claude-3-5-sonnet-20241022", base_url="...", api_key="...") -gemini = FerroChatModel(model="gemini-1.5-flash", base_url="...", api_key="...") +gemini = FerroChatModel(model="gemini-2.5-flash", base_url="...", api_key="...") ``` ### Streaming @@ -50,6 +56,15 @@ for chunk in llm.stream([HumanMessage(content="Tell me a story")]): print(chunk.content, end="", flush=True) ``` +### Async + +```python +response = await llm.ainvoke([HumanMessage(content="Hello")]) + +async for chunk in llm.astream([HumanMessage(content="Tell me a story")]): + print(chunk.content, end="", flush=True) +``` + ### Tool calling / LangGraph agents ```python @@ -65,6 +80,25 @@ response = agent_llm.invoke([HumanMessage(content="What is 4 + 7?")]) print(response.tool_calls) ``` +### Structured output + +Uses the OpenAI-style `response_format={"type": "json_schema", ...}` path, so +it works with every provider the gateway can translate it for (see +`client.capabilities()` on the core SDK). + +```python +from pydantic import BaseModel + +class Answer(BaseModel): + city: str + population: int + +structured = llm.with_structured_output(Answer) +print(structured.invoke("Largest city in France?")) # Answer(city='Paris', population=...) + +# include_raw=True → {"raw": AIMessage, "parsed": Answer | None, "parsing_error": ...} +``` + ### Embeddings ```python @@ -73,6 +107,9 @@ from langchain_ferrolabsai import FerroEmbeddings embed = FerroEmbeddings(model="text-embedding-3-small", base_url="...", api_key="...") vectors = embed.embed_documents(["hello", "world"]) query_vec = embed.embed_query("hello") + +# async +vectors = await embed.aembed_documents(["hello", "world"]) ``` ### Legacy `LLM` interface @@ -88,24 +125,23 @@ print(llm.invoke("Write a haiku about gateways")) `ChatOpenAI` pointed at a Ferro Labs gateway works as a drop-in. This package adds: -- First-class `provider`, `cost_usd`, `latency_ms`, `trace_id` exposure on `response_metadata`. -- Native support for Ferro extras: `route_tag`, `template_id`, `template_variables`. -- `trace_id` is the **join key** for the v1.2 observability bridge plugins - (LangSmith, Langfuse, Phoenix, Datadog, …) shipping from the - [`ferro-labs/ai-gateway-plugins`](https://github.com/ferro-labs) repo — - any provider's calls become visible in your existing LLMOps backend without - per-provider wiring. +- `provider`, `trace_id`, and `gateway_overhead_ms` on `response_metadata` (and + `trace_id` on the first streamed chunk) — read from the gateway's real response + headers, no guessing. +- Typed gateway errors from the core SDK: `FerroBudgetExceededError` (402), + `FerroPermissionError` (403), `FerroRateLimitError.retry_after`, with + `Retry-After`-aware retries on 429/5xx. +- Async parity (`ainvoke`, `astream`, `aembed_*`) over `AsyncFerroClient`. ## Status & roadmap -`0.1.0` is the **first functional release** of the adapter. See -[`CHANGELOG.md`](CHANGELOG.md) for what shipped and what's planned. Async -surfaces and `with_structured_output()` are the next two items. +`0.2.0` adds the async surface and `with_structured_output()` on top of +`ferrolabsai 0.3`. See [`CHANGELOG.md`](CHANGELOG.md). ## Related - [`ferrolabsai`](https://pypi.org/project/ferrolabsai/) — the core Python SDK this package wraps. -- [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) — the open-source gateway server (v1.1.0+ OTel-native). +- [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) — the open-source gateway server. - [`ai-gateway-cookbook`](https://github.com/ferro-labs/ai-gateway-cookbook) — runnable recipes (start with `python/02-langgraph-multi-provider-agent`). - [Documentation](https://docs.ferrolabs.ai) diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py index fa5bc47..9242706 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py @@ -8,10 +8,11 @@ embed = FerroEmbeddings(model="text-embedding-3-small", api_key="sk-ferro-...") legacy = FerroLLM(model="gpt-4o", api_key="sk-ferro-...") -All three classes route through a Ferro Labs AI Gateway endpoint and expose -the gateway's ``trace_id`` (frozen contract since ``ai-gateway v1.1.0``) via -``response_metadata`` — the join key for the v1.2 observability bridge plugins -(LangSmith, Langfuse, Phoenix, …). +All three classes route through a Ferro Labs AI Gateway endpoint (ai-gateway +≥ v1.4.0). Chat responses expose the gateway's ``trace_id`` (the ``X-Request-ID`` +response header), ``provider`` and ``gateway_overhead_ms`` via +``response_metadata`` — ``trace_id`` is the join key for the gateway's request +log and its observability exporters (LangSmith, Langfuse, Phoenix, …). """ from __future__ import annotations @@ -20,7 +21,7 @@ from .embeddings import FerroEmbeddings from .llms import FerroLLM -__version__ = "0.1.0" +__version__ = "0.2.0" __all__ = [ "__version__", diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py index 2830da0..0676eba 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py @@ -1,31 +1,33 @@ """FerroChatModel — LangChain ``BaseChatModel`` backed by Ferro Labs AI Gateway. -A single ``FerroChatModel`` instance can address any of the gateway's 30+ +A single ``FerroChatModel`` instance can address any of the gateway's 30 providers by name (e.g. ``"gpt-4o"``, ``"claude-3-5-sonnet-20241022"``, -``"gemini-1.5-flash"``) without changing the model class. +``"gemini-2.5-flash"``) without changing the model class. -Every response surfaces ``trace_id`` (the Ferro request ID propagated via the -``x-trace-id`` header — frozen contract since ``ai-gateway v1.1.0``) in -``response_metadata``. That value is the join key for any downstream -observability bridge plugin (LangSmith, Langfuse, Phoenix, …) that ships in -``ferro-labs/ai-gateway-plugins``. +``response_metadata`` carries exactly what the gateway provides: ``model``, +``id``, ``trace_id`` (the ``X-Request-ID`` response header — the join key for +the gateway's request log and observability exporters), ``provider`` (body +field), and ``gateway_overhead_ms`` (``X-Gateway-Overhead-Ms`` header). +Absent values are stripped. """ from __future__ import annotations import json -from collections.abc import Iterator, Sequence +from collections.abc import AsyncIterator, Iterator, Sequence +from operator import itemgetter from typing import TYPE_CHECKING, Any, cast -from ferrolabsai import ChatCompletion, FerroClient -from langchain_core.callbacks import CallbackManagerForLLMRun +from ferrolabsai import AsyncFerroClient, ChatCompletion, ChatCompletionChunk, FerroClient +from langchain_core.callbacks import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun from langchain_core.language_models import BaseChatModel, LanguageModelInput from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage +from langchain_core.output_parsers import JsonOutputParser, PydanticOutputParser from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult -from langchain_core.runnables import Runnable +from langchain_core.runnables import Runnable, RunnableMap, RunnablePassthrough from langchain_core.tools import BaseTool from langchain_core.utils.function_calling import convert_to_openai_tool -from pydantic import ConfigDict, Field, PrivateAttr, SecretStr +from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, SecretStr from ._messages import messages_to_ferro_dicts @@ -41,10 +43,10 @@ class FerroChatModel(BaseChatModel): from langchain_ferrolabsai import FerroChatModel from langchain_core.messages import HumanMessage - chat = FerroChatModel(model="gpt-4o", api_key="sk-ferro-...") + chat = FerroChatModel(model="gpt-4o", api_key="fgw_...") response = chat.invoke([HumanMessage(content="Hello")]) print(response.content) - print(response.response_metadata["trace_id"]) # Ferro request ID + print(response.response_metadata["trace_id"]) # gateway X-Request-ID """ model: str = Field(..., description="Model name routed by the gateway.") @@ -65,14 +67,6 @@ class FerroChatModel(BaseChatModel): frequency_penalty: float | None = None presence_penalty: float | None = None stop: list[str] | None = None - - # Ferro-specific extras - route_tag: str | None = Field( - default=None, - description="Override the gateway's routing strategy for this caller.", - ) - template_id: str | None = None - template_variables: dict[str, Any] | None = None user: str | None = None default_headers: dict[str, str] | None = None @@ -81,6 +75,7 @@ class FerroChatModel(BaseChatModel): model_config = ConfigDict(arbitrary_types_allowed=True, populate_by_name=True) _client_instance: FerroClient | None = PrivateAttr(default=None) + _async_client_instance: AsyncFerroClient | None = PrivateAttr(default=None) # ------------------------------------------------------------------ # LangChain identification @@ -103,17 +98,25 @@ def _identifying_params(self) -> dict[str, Any]: # Client access # ------------------------------------------------------------------ + def _client_kwargs(self) -> dict[str, Any]: + return { + "api_key": self.api_key.get_secret_value() if self.api_key else None, + "base_url": self.base_url, + "timeout": self.timeout, + "max_retries": self.max_retries, + "default_headers": self.default_headers, + } + def _get_client(self) -> FerroClient: if self._client_instance is None: - self._client_instance = FerroClient( - api_key=self.api_key.get_secret_value() if self.api_key else None, - base_url=self.base_url, - timeout=self.timeout, - max_retries=self.max_retries, - default_headers=self.default_headers, - ) + self._client_instance = FerroClient(**self._client_kwargs()) return self._client_instance + def _get_async_client(self) -> AsyncFerroClient: + if self._async_client_instance is None: + self._async_client_instance = AsyncFerroClient(**self._client_kwargs()) + return self._async_client_instance + # ------------------------------------------------------------------ # Request payload assembly # ------------------------------------------------------------------ @@ -137,12 +140,6 @@ def _build_request_params( effective_stop = stop if stop is not None else self.stop if effective_stop: params["stop"] = effective_stop - if self.route_tag is not None: - params["route_tag"] = self.route_tag - if self.template_id is not None: - params["template_id"] = self.template_id - if self.template_variables is not None: - params["template_variables"] = self.template_variables if self.user is not None: params["user"] = self.user # model_kwargs first so explicit per-call kwargs win. @@ -168,6 +165,20 @@ def _generate( ) return _completion_to_chat_result(response) + async def _agenerate( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + run_manager: AsyncCallbackManagerForLLMRun | None = None, + **kwargs: Any, + ) -> ChatResult: + params = self._build_request_params(stop, **kwargs) + response = await self._get_async_client().chat.completions.create( + messages=messages_to_ferro_dicts(messages), + **params, + ) + return _completion_to_chat_result(response) + def _stream( self, messages: list[BaseMessage], @@ -181,27 +192,45 @@ def _stream( stream=True, **params, ) + first = True for chunk in stream: - if not chunk.choices: + generation_chunk = _chunk_to_generation(chunk, first) + if generation_chunk is None: + continue + first = False + if run_manager is not None: + run_manager.on_llm_new_token( + cast("str", generation_chunk.message.content), chunk=generation_chunk + ) + yield generation_chunk + + async def _astream( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + run_manager: AsyncCallbackManagerForLLMRun | None = None, + **kwargs: Any, + ) -> AsyncIterator[ChatGenerationChunk]: + params = self._build_request_params(stop, **kwargs) + stream = await self._get_async_client().chat.completions.create( + messages=messages_to_ferro_dicts(messages), + stream=True, + **params, + ) + first = True + async for chunk in stream: + generation_chunk = _chunk_to_generation(chunk, first) + if generation_chunk is None: continue - delta = chunk.choices[0].delta - content = delta.content or "" - ai_chunk = AIMessageChunk( - content=content, - tool_call_chunks=_extract_tool_call_chunks(delta.tool_calls), - ) - generation_chunk = ChatGenerationChunk( - message=ai_chunk, - generation_info={"finish_reason": chunk.choices[0].finish_reason} - if chunk.choices[0].finish_reason - else None, - ) + first = False if run_manager is not None: - run_manager.on_llm_new_token(content, chunk=generation_chunk) + await run_manager.on_llm_new_token( + cast("str", generation_chunk.message.content), chunk=generation_chunk + ) yield generation_chunk # ------------------------------------------------------------------ - # Tool binding (LangGraph / agent support) + # Tool binding (LangGraph / agent support) and structured output # ------------------------------------------------------------------ def bind_tools( @@ -218,6 +247,47 @@ def bind_tools( bind_kwargs.update(kwargs) return cast("Runnable[LanguageModelInput, AIMessage]", super().bind(**bind_kwargs)) + def with_structured_output( + self, + schema: dict[str, Any] | type, + *, + include_raw: bool = False, + **kwargs: Any, + ) -> Runnable[LanguageModelInput, dict[str, Any] | BaseModel]: + """Constrain the model to ``schema`` via OpenAI-style + ``response_format={"type": "json_schema", ...}`` and parse the reply. + + ``schema`` is a pydantic ``BaseModel`` subclass (parsed into an instance) + or a JSON-schema dict (parsed into a dict). With ``include_raw=True`` the + output is ``{"raw": AIMessage, "parsed": ..., "parsing_error": ...}``. + """ + parser: Runnable[Any, Any] + if isinstance(schema, type) and issubclass(schema, BaseModel): + name, json_schema = schema.__name__, schema.model_json_schema() + parser = PydanticOutputParser(pydantic_object=schema) + elif isinstance(schema, dict): + name, json_schema = str(schema.get("title", "output")), schema + parser = JsonOutputParser() + else: + raise TypeError("schema must be a pydantic BaseModel subclass or a JSON-schema dict") + llm = self.bind( + response_format={ + "type": "json_schema", + "json_schema": {"name": name, "schema": json_schema}, + }, + **kwargs, + ) + if not include_raw: + return cast("Runnable[LanguageModelInput, dict[str, Any] | BaseModel]", llm | parser) + parse = RunnablePassthrough.assign( + parsed=itemgetter("raw") | parser, parsing_error=lambda _: None + ) + fallback = RunnablePassthrough.assign(parsed=lambda _: None) + chain = RunnableMap(raw=llm) | parse.with_fallbacks( + [fallback], exception_key="parsing_error" + ) + return cast("Runnable[LanguageModelInput, dict[str, Any] | BaseModel]", chain) + # --------------------------------------------------------------------------- # Response mapping @@ -226,49 +296,35 @@ def bind_tools( def _completion_to_chat_result(response: ChatCompletion) -> ChatResult: """Convert a Ferro ``ChatCompletion`` into a LangChain ``ChatResult``.""" + metadata = _response_metadata(response) if not response.choices: - empty = AIMessage( - content="", - response_metadata=_response_metadata(response), - ) + empty = AIMessage(content="", response_metadata=metadata) return ChatResult(generations=[ChatGeneration(message=empty)]) choice = response.choices[0] - tool_calls = _extract_tool_calls(choice.message.tool_calls) ai_message = AIMessage( content=choice.message.content or "", - tool_calls=tool_calls, - response_metadata=_response_metadata(response), + tool_calls=_extract_tool_calls(choice.message.tool_calls), + response_metadata=metadata, usage_metadata=_usage_metadata(response), ) generation = ChatGeneration( message=ai_message, generation_info={"finish_reason": choice.finish_reason} if choice.finish_reason else None, ) - return ChatResult( - generations=[generation], - llm_output={ - "model": response.model, - "trace_id": response.trace_id, - "provider": response.provider, - }, - ) + # LangChain merges llm_output into response_metadata; keep them identical. + return ChatResult(generations=[generation], llm_output=metadata) def _response_metadata(response: ChatCompletion) -> dict[str, Any]: - """The Ferro-specific surface every consumer (incl. v1.2 observability bridges) reads.""" + """Exactly the fields ai-gateway provides; ``None`` values are stripped.""" metadata: dict[str, Any] = { "model": response.model, "id": response.id, - # ``trace_id`` is the canonical join key. Frozen via x-trace-id since - # ai-gateway v1.1.0; mirrored by every Ferro observability bridge plugin. "trace_id": response.trace_id, "provider": response.provider, - "latency_ms": response.latency_ms, + "gateway_overhead_ms": response.gateway_overhead_ms, } - if response.usage is not None: - metadata["cost_usd"] = response.usage.cost_usd - metadata["cache_hit"] = response.usage.cache_hit return {k: v for k, v in metadata.items() if v is not None} @@ -282,6 +338,24 @@ def _usage_metadata(response: ChatCompletion) -> dict[str, int] | None: } +def _chunk_to_generation(chunk: ChatCompletionChunk, first: bool) -> ChatGenerationChunk | None: + """Map one SSE chunk to a ``ChatGenerationChunk``; ``None`` for chunks with no choices + (e.g. the terminal usage-only chunk). Stream metadata rides on the first chunk.""" + if not chunk.choices: + return None + choice = chunk.choices[0] + metadata = {k: v for k, v in (("trace_id", chunk.trace_id), ("provider", chunk.provider)) if v} + ai_chunk = AIMessageChunk( + content=choice.delta.content or "", + tool_call_chunks=_extract_tool_call_chunks(choice.delta.tool_calls), + response_metadata=metadata if first else {}, + ) + return ChatGenerationChunk( + message=ai_chunk, + generation_info={"finish_reason": choice.finish_reason} if choice.finish_reason else None, + ) + + def _extract_tool_call_chunks(raw: list[dict[str, Any]] | None) -> list[ToolCallChunk]: """Map OpenAI streaming tool-call deltas to LangChain chunk shape.""" if not raw: diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/embeddings.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/embeddings.py index 35afb3c..5ff70a8 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/embeddings.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/embeddings.py @@ -9,11 +9,17 @@ from typing import Any -from ferrolabsai import FerroClient +from ferrolabsai import AsyncFerroClient, EmbeddingResponse, FerroClient from langchain_core.embeddings import Embeddings from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, SecretStr +def _ordered_vectors(response: EmbeddingResponse) -> list[list[float]]: + # Preserve input order by sorting on `index` — the provider may return + # data out of order. + return [d.embedding for d in sorted(response.data, key=lambda d: d.index)] + + class FerroEmbeddings(BaseModel, Embeddings): """LangChain embeddings adapter for the Ferro gateway. @@ -21,7 +27,7 @@ class FerroEmbeddings(BaseModel, Embeddings): from langchain_ferrolabsai import FerroEmbeddings - embed = FerroEmbeddings(model="text-embedding-3-small", api_key="sk-ferro-...") + embed = FerroEmbeddings(model="text-embedding-3-small", api_key="fgw_...") vectors = embed.embed_documents(["hello", "world"]) query_vec = embed.embed_query("hello") """ @@ -39,18 +45,27 @@ class FerroEmbeddings(BaseModel, Embeddings): model_config = ConfigDict(arbitrary_types_allowed=True, populate_by_name=True) _client_instance: FerroClient | None = PrivateAttr(default=None) + _async_client_instance: AsyncFerroClient | None = PrivateAttr(default=None) + + def _client_kwargs(self) -> dict[str, Any]: + return { + "api_key": self.api_key.get_secret_value() if self.api_key else None, + "base_url": self.base_url, + "timeout": self.timeout, + "max_retries": self.max_retries, + "default_headers": self.default_headers, + } def _get_client(self) -> FerroClient: if self._client_instance is None: - self._client_instance = FerroClient( - api_key=self.api_key.get_secret_value() if self.api_key else None, - base_url=self.base_url, - timeout=self.timeout, - max_retries=self.max_retries, - default_headers=self.default_headers, - ) + self._client_instance = FerroClient(**self._client_kwargs()) return self._client_instance + def _get_async_client(self) -> AsyncFerroClient: + if self._async_client_instance is None: + self._async_client_instance = AsyncFerroClient(**self._client_kwargs()) + return self._async_client_instance + def _build_kwargs(self) -> dict[str, Any]: kwargs: dict[str, Any] = {"model": self.model} if self.dimensions is not None: @@ -65,13 +80,20 @@ def embed_documents(self, texts: list[str]) -> list[list[float]]: if not texts: return [] response = self._get_client().embeddings.create(input=texts, **self._build_kwargs()) - # Preserve input order by sorting on `index` — the gateway / provider - # may return data out of order. - ordered = sorted(response.data, key=lambda d: d.index) - return [d.embedding for d in ordered] + return _ordered_vectors(response) def embed_query(self, text: str) -> list[float]: response = self._get_client().embeddings.create(input=text, **self._build_kwargs()) - if not response.data: + return response.data[0].embedding if response.data else [] + + async def aembed_documents(self, texts: list[str]) -> list[list[float]]: + if not texts: return [] - return response.data[0].embedding + client = self._get_async_client() + response = await client.embeddings.create(input=texts, **self._build_kwargs()) + return _ordered_vectors(response) + + async def aembed_query(self, text: str) -> list[float]: + client = self._get_async_client() + response = await client.embeddings.create(input=text, **self._build_kwargs()) + return response.data[0].embedding if response.data else [] diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/llms.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/llms.py index 2e68c37..6f09746 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/llms.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/llms.py @@ -26,7 +26,6 @@ class FerroLLM(LLM): max_retries: int = 2 temperature: float | None = None max_tokens: int | None = None - route_tag: str | None = None default_headers: dict[str, str] | None = None model_config = ConfigDict(arbitrary_types_allowed=True, populate_by_name=True) @@ -62,8 +61,6 @@ def _call( params["max_tokens"] = self.max_tokens if stop is not None: params["stop"] = stop - if self.route_tag is not None: - params["route_tag"] = self.route_tag params.update(kwargs) response = self._get_client().chat.completions.create( diff --git a/integrations/langchain-ferrolabsai/pyproject.toml b/integrations/langchain-ferrolabsai/pyproject.toml index 70103f3..3c9fb3d 100644 --- a/integrations/langchain-ferrolabsai/pyproject.toml +++ b/integrations/langchain-ferrolabsai/pyproject.toml @@ -4,8 +4,8 @@ build-backend = "hatchling.build" [project] name = "langchain-ferrolabsai" -version = "0.1.0" -description = "LangChain integration for Ferro Labs AI Gateway — chat, streaming, embeddings, and tool calling across 30+ LLM providers via a single OpenAI-compatible endpoint" +version = "0.2.0" +description = "LangChain integration for Ferro Labs AI Gateway — chat, streaming, embeddings, and tool calling across 30 LLM providers via a single OpenAI-compatible endpoint" readme = "README.md" license = { text = "Apache-2.0" } requires-python = ">=3.9" @@ -27,11 +27,12 @@ classifiers = [ "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", "Topic :: Software Development :: Libraries :: Python Modules", "Topic :: Scientific/Engineering :: Artificial Intelligence", ] dependencies = [ - "ferrolabsai>=0.1.0", + "ferrolabsai>=0.3.0", "langchain-core>=0.3.0", "eval-type-backport>=0.2.0; python_version < '3.10'", ] diff --git a/integrations/langchain-ferrolabsai/tests/conftest.py b/integrations/langchain-ferrolabsai/tests/conftest.py index a5802a5..87a3027 100644 --- a/integrations/langchain-ferrolabsai/tests/conftest.py +++ b/integrations/langchain-ferrolabsai/tests/conftest.py @@ -1,7 +1,9 @@ """Shared test fixtures for langchain-ferrolabsai. Follows the same pytest-httpx mocking pattern as the parent ferrolabsai SDK so -no real gateway is required to run the suite. +no real gateway is required to run the suite. Payloads mirror what +ai-gateway v1.4.x returns: `provider` in the chat body, `X-Request-ID` and +`X-Gateway-Overhead-Ms` as response headers. """ from __future__ import annotations @@ -12,6 +14,8 @@ BASE_URL = "http://test-gateway:8080" API_KEY = "sk-ferro-test" +TRACE_ID = "0af7651916cd43dd8448eb211c80319c" +GATEWAY_HEADERS = {"X-Request-ID": TRACE_ID, "X-Gateway-Overhead-Ms": "1.5"} def make_chat_completion( @@ -19,12 +23,9 @@ def make_chat_completion( content: str = "Hello back", model: str = "gpt-4o", provider: str = "openai", - trace_id: str = "trace-abc-123", - latency_ms: int = 42, - cost_usd: float = 0.000123, tool_calls: list[dict[str, Any]] | None = None, ) -> dict[str, Any]: - """Build a Ferro chat-completion response payload for use with httpx_mock.""" + """Build a gateway chat-completion body for use with httpx_mock.""" message: dict[str, Any] = {"role": "assistant", "content": content} if tool_calls is not None: message["tool_calls"] = tool_calls @@ -33,23 +34,9 @@ def make_chat_completion( "object": "chat.completion", "created": 1_700_000_000, "model": model, - "choices": [ - { - "index": 0, - "message": message, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 5, - "completion_tokens": 3, - "total_tokens": 8, - "cost_usd": cost_usd, - "provider": provider, - }, - "x_ferro_trace_id": trace_id, - "x_ferro_provider": provider, - "x_ferro_latency_ms": latency_ms, + "provider": provider, + "choices": [{"index": 0, "message": message, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, } @@ -68,6 +55,33 @@ def make_embedding_response( } +def sse_chunks(*contents: str, finish: bool = True) -> bytes: + import json + + frames = [ + { + "id": "1", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o", + "choices": [{"index": 0, "delta": {"content": c}, "finish_reason": None}], + } + for c in contents + ] + if finish: + frames.append( + { + "id": "1", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o", + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + } + ) + body = "".join(f"data: {json.dumps(f)}\n\n" for f in frames) + "data: [DONE]\n\n" + return body.encode("utf-8") + + @pytest.fixture def base_url() -> str: return BASE_URL diff --git a/integrations/langchain-ferrolabsai/tests/test_chat_models.py b/integrations/langchain-ferrolabsai/tests/test_chat_models.py index 6df7ede..0a2afb1 100644 --- a/integrations/langchain-ferrolabsai/tests/test_chat_models.py +++ b/integrations/langchain-ferrolabsai/tests/test_chat_models.py @@ -7,11 +7,14 @@ import pytest from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage from langchain_core.tools import tool +from pydantic import BaseModel from pytest_httpx import HTTPXMock from langchain_ferrolabsai import FerroChatModel -from .conftest import BASE_URL, make_chat_completion +from .conftest import BASE_URL, GATEWAY_HEADERS, TRACE_ID, make_chat_completion, sse_chunks + +CHAT_URL = f"{BASE_URL}/v1/chat/completions" def _build_chat(**overrides) -> FerroChatModel: @@ -23,79 +26,44 @@ def _build_chat(**overrides) -> FerroChatModel: class TestBasicGeneration: def test_invoke_returns_ai_message_with_content(self, httpx_mock: HTTPXMock): httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(content="hi there"), + method="POST", url=CHAT_URL, json=make_chat_completion(content="hi there") ) - chat = _build_chat() - result = chat.invoke([HumanMessage(content="Hello")]) + result = _build_chat().invoke([HumanMessage(content="Hello")]) assert isinstance(result, AIMessage) assert result.content == "hi there" - def test_invoke_surfaces_trace_id_in_response_metadata(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(trace_id="my-trace-xyz"), - ) - chat = _build_chat() - result = chat.invoke([HumanMessage(content="Hello")]) - # trace_id is the join key for v1.2 observability bridges — MUST be present. - assert result.response_metadata["trace_id"] == "my-trace-xyz" - assert result.response_metadata["provider"] == "openai" - assert result.response_metadata["latency_ms"] == 42 - assert result.response_metadata["cost_usd"] == 0.000123 - - def test_invoke_surfaces_header_only_gateway_metadata(self, httpx_mock: HTTPXMock): - body = make_chat_completion(trace_id="body-trace", provider="body-provider") - body.pop("x_ferro_trace_id") - body.pop("x_ferro_provider") - body.pop("x_ferro_latency_ms") - body["usage"].pop("cost_usd") - body["usage"].pop("provider") + def test_response_metadata_is_exactly_what_the_gateway_provides(self, httpx_mock: HTTPXMock): httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=body, - headers={ - "X-Request-ID": "header-trace", - "x-ferro-provider": "openai", - "x-ferro-latency-ms": "42", - "x-ferro-cost-usd": "0.000123", - }, + method="POST", url=CHAT_URL, json=make_chat_completion(), headers=GATEWAY_HEADERS ) - chat = _build_chat() - result = chat.invoke([HumanMessage(content="Hello")]) - assert result.response_metadata["trace_id"] == "header-trace" + result = _build_chat().invoke([HumanMessage(content="Hello")]) + expected = { + "model": "gpt-4o", + "id": "cmpl-1", + "trace_id": TRACE_ID, + "provider": "openai", + "gateway_overhead_ms": 1.5, + } + # LangChain adds generation_info (finish_reason) on top; every gateway field is exact. + assert {k: result.response_metadata[k] for k in expected} == expected + + def test_response_metadata_strips_absent_fields(self, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, json=make_chat_completion()) + result = _build_chat().invoke([HumanMessage(content="Hello")]) + assert "trace_id" not in result.response_metadata + assert "gateway_overhead_ms" not in result.response_metadata assert result.response_metadata["provider"] == "openai" - assert result.response_metadata["latency_ms"] == 42 - assert result.response_metadata["cost_usd"] == 0.000123 def test_invoke_attaches_usage_metadata(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(), - ) - chat = _build_chat() - result = chat.invoke([HumanMessage(content="Hello")]) - assert result.usage_metadata == { - "input_tokens": 5, - "output_tokens": 3, - "total_tokens": 8, - } + httpx_mock.add_response(method="POST", url=CHAT_URL, json=make_chat_completion()) + result = _build_chat().invoke([HumanMessage(content="Hello")]) + assert result.usage_metadata == {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8} class TestMessageConversion: def test_system_human_messages_serialized_correctly(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(), - ) - chat = _build_chat() - chat.invoke([SystemMessage(content="be terse"), HumanMessage(content="hi")]) - + httpx_mock.add_response(method="POST", url=CHAT_URL, json=make_chat_completion()) + _build_chat().invoke([SystemMessage(content="be terse"), HumanMessage(content="hi")]) body = json.loads(httpx_mock.get_requests()[0].content) assert body["messages"] == [ {"role": "system", "content": "be terse"}, @@ -103,13 +71,8 @@ def test_system_human_messages_serialized_correctly(self, httpx_mock: HTTPXMock) ] def test_tool_messages_carry_tool_call_id(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(), - ) - chat = _build_chat() - chat.invoke( + httpx_mock.add_response(method="POST", url=CHAT_URL, json=make_chat_completion()) + _build_chat().invoke( [ HumanMessage(content="what is 1+1"), AIMessage( @@ -119,57 +82,34 @@ def test_tool_messages_carry_tool_call_id(self, httpx_mock: HTTPXMock): ToolMessage(content="2", tool_call_id="c1"), ] ) - body = json.loads(httpx_mock.get_requests()[0].content) - tool_msg = body["messages"][-1] - assert tool_msg["role"] == "tool" - assert tool_msg["tool_call_id"] == "c1" - assert tool_msg["content"] == "2" + tool_msg = json.loads(httpx_mock.get_requests()[0].content)["messages"][-1] + assert tool_msg == {"role": "tool", "tool_call_id": "c1", "content": "2"} class TestRequestParams: def test_sends_auth_header(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(), - ) - chat = _build_chat(api_key="sk-ferro-prod") - chat.invoke([HumanMessage(content="hi")]) - request = httpx_mock.get_requests()[0] - assert request.headers["Authorization"] == "Bearer sk-ferro-prod" + httpx_mock.add_response(method="POST", url=CHAT_URL, json=make_chat_completion()) + _build_chat(api_key="sk-ferro-prod").invoke([HumanMessage(content="hi")]) + assert httpx_mock.get_requests()[0].headers["Authorization"] == "Bearer sk-ferro-prod" - def test_forwards_temperature_and_max_tokens(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(), + def test_forwards_sampling_params_user_and_model_kwargs(self, httpx_mock: HTTPXMock): + httpx_mock.add_response(method="POST", url=CHAT_URL, json=make_chat_completion()) + chat = _build_chat( + temperature=0.2, max_tokens=64, user="user-123", model_kwargs={"seed": 7} ) - chat = _build_chat(temperature=0.2, max_tokens=64) chat.invoke([HumanMessage(content="hi")]) body = json.loads(httpx_mock.get_requests()[0].content) assert body["temperature"] == 0.2 assert body["max_tokens"] == 64 - - def test_forwards_ferro_extras(self, httpx_mock: HTTPXMock): - httpx_mock.add_response( - method="POST", - url=f"{BASE_URL}/v1/chat/completions", - json=make_chat_completion(), - ) - chat = _build_chat( - route_tag="premium", - template_id="customer-support", - template_variables={"tone": "friendly"}, - user="user-123", - ) - chat.invoke([HumanMessage(content="hi")]) - body = json.loads(httpx_mock.get_requests()[0].content) - # Ferro forwards `route_tag` as `x_route_tag` internally; we just - # check the field round-trips through the SDK's request builder. - assert body.get("x_route_tag") == "premium" - assert body["template_id"] == "customer-support" - assert body["template_variables"] == {"tone": "friendly"} assert body["user"] == "user-123" + assert body["seed"] == 7 + for dead in ("route_tag", "x_route_tag", "template_id", "template_variables"): + assert dead not in body + + def test_dead_gateway_fields_are_not_model_fields(self): + chat = _build_chat() + for dead in ("route_tag", "template_id", "template_variables"): + assert dead not in type(chat).model_fields class TestToolBinding: @@ -181,7 +121,7 @@ def add(a: int, b: int) -> int: httpx_mock.add_response( method="POST", - url=f"{BASE_URL}/v1/chat/completions", + url=CHAT_URL, json=make_chat_completion( content="", tool_calls=[ @@ -193,38 +133,26 @@ def add(a: int, b: int) -> int: ], ), ) - chat = _build_chat().bind_tools([add]) - result = chat.invoke([HumanMessage(content="add 1 and 2")]) - + result = _build_chat().bind_tools([add]).invoke([HumanMessage(content="add 1 and 2")]) body = json.loads(httpx_mock.get_requests()[0].content) assert body["tools"][0]["type"] == "function" assert body["tools"][0]["function"]["name"] == "add" - assert result.tool_calls == [ {"id": "call_1", "name": "add", "args": {"a": 1, "b": 2}, "type": "tool_call"} ] class TestStreaming: - def test_stream_yields_chunks(self, httpx_mock: HTTPXMock): - sse_body = ( - 'data: {"id":"1","object":"chat.completion.chunk","created":1,"model":"gpt-4o",' - '"choices":[{"index":0,"delta":{"role":"assistant","content":"Hel"},"finish_reason":null}]}\n\n' - 'data: {"id":"1","object":"chat.completion.chunk","created":1,"model":"gpt-4o",' - '"choices":[{"index":0,"delta":{"content":"lo"},"finish_reason":null}]}\n\n' - 'data: {"id":"1","object":"chat.completion.chunk","created":1,"model":"gpt-4o",' - '"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}\n\n' - "data: [DONE]\n\n" - ) + def test_stream_yields_chunks_with_trace_id(self, httpx_mock: HTTPXMock): httpx_mock.add_response( method="POST", - url=f"{BASE_URL}/v1/chat/completions", - content=sse_body.encode("utf-8"), - headers={"Content-Type": "text/event-stream"}, + url=CHAT_URL, + content=sse_chunks("Hel", "lo"), + headers={"Content-Type": "text/event-stream", **GATEWAY_HEADERS}, ) - chat = _build_chat() - chunks = list(chat.stream([HumanMessage(content="hi")])) + chunks = list(_build_chat().stream([HumanMessage(content="hi")])) assert "".join(c.content for c in chunks) == "Hello" + assert chunks[0].response_metadata["trace_id"] == TRACE_ID def test_stream_yields_tool_call_chunks(self, httpx_mock: HTTPXMock): frames = [ @@ -259,25 +187,21 @@ def test_stream_yields_tool_call_chunks(self, httpx_mock: HTTPXMock): { "index": 0, "delta": { - "tool_calls": [ - {"index": 0, "function": {"arguments": ',"b":2}'}} - ] + "tool_calls": [{"index": 0, "function": {"arguments": ',"b":2}'}}] }, "finish_reason": None, } ], }, ] - sse_body = "".join(f"data: {json.dumps(frame)}\n\n" for frame in frames) - sse_body += "data: [DONE]\n\n" + sse_body = "".join(f"data: {json.dumps(f)}\n\n" for f in frames) + "data: [DONE]\n\n" httpx_mock.add_response( method="POST", - url=f"{BASE_URL}/v1/chat/completions", + url=CHAT_URL, content=sse_body.encode("utf-8"), headers={"Content-Type": "text/event-stream"}, ) - chat = _build_chat() - chunks = list(chat.stream([HumanMessage(content="add 1 and 2")])) + chunks = list(_build_chat().stream([HumanMessage(content="add 1 and 2")])) tool_chunks = [tc for chunk in chunks for tc in chunk.tool_call_chunks] assert tool_chunks[0]["id"] == "call_1" assert tool_chunks[0]["name"] == "add" @@ -285,6 +209,71 @@ def test_stream_yields_tool_call_chunks(self, httpx_mock: HTTPXMock): assert tool_chunks[1]["args"] == ',"b":2}' +class TestAsync: + async def test_ainvoke(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + json=make_chat_completion(content="async hi"), + headers=GATEWAY_HEADERS, + ) + result = await _build_chat().ainvoke([HumanMessage(content="hi")]) + assert result.content == "async hi" + assert result.response_metadata["trace_id"] == TRACE_ID + assert result.usage_metadata["total_tokens"] == 8 + + async def test_astream(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + content=sse_chunks("a", "b", "c"), + headers={"Content-Type": "text/event-stream", **GATEWAY_HEADERS}, + ) + chunks = [c async for c in _build_chat().astream([HumanMessage(content="hi")])] + assert "".join(c.content for c in chunks) == "abc" + assert chunks[0].response_metadata["trace_id"] == TRACE_ID + assert json.loads(httpx_mock.get_requests()[0].content)["stream"] is True + + +class Answer(BaseModel): + city: str + population: int + + +class TestStructuredOutput: + def test_pydantic_schema_uses_json_schema_response_format(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + json=make_chat_completion(content='{"city": "Paris", "population": 2100000}'), + ) + structured = _build_chat().with_structured_output(Answer) + result = structured.invoke([HumanMessage(content="Biggest city in France?")]) + assert result == Answer(city="Paris", population=2100000) + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["response_format"]["type"] == "json_schema" + assert body["response_format"]["json_schema"]["name"] == "Answer" + assert "city" in body["response_format"]["json_schema"]["schema"]["properties"] + + def test_dict_schema_returns_dict(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=CHAT_URL, json=make_chat_completion(content='{"ok": true}') + ) + schema = {"title": "Flag", "type": "object", "properties": {"ok": {"type": "boolean"}}} + result = _build_chat().with_structured_output(schema).invoke("ready?") + assert result == {"ok": True} + body = json.loads(httpx_mock.get_requests()[0].content) + assert body["response_format"]["json_schema"] == {"name": "Flag", "schema": schema} + + def test_include_raw(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", url=CHAT_URL, json=make_chat_completion(content="not json") + ) + result = _build_chat().with_structured_output(Answer, include_raw=True).invoke("?") + assert isinstance(result["raw"], AIMessage) + assert result["parsed"] is None + assert result["parsing_error"] is not None + class TestIdentity: def test_llm_type(self): @@ -301,6 +290,5 @@ def test_missing_api_key_raises(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.delenv("FERRO_API_KEY", raising=False) monkeypatch.delenv("OPENAI_API_KEY", raising=False) chat = FerroChatModel(model="gpt-4o", base_url=BASE_URL) # no api_key - # FerroClient construction happens lazily on first call. with pytest.raises(Exception): chat.invoke([HumanMessage(content="hi")]) diff --git a/integrations/langchain-ferrolabsai/tests/test_embeddings.py b/integrations/langchain-ferrolabsai/tests/test_embeddings.py index 9488e97..c79dc03 100644 --- a/integrations/langchain-ferrolabsai/tests/test_embeddings.py +++ b/integrations/langchain-ferrolabsai/tests/test_embeddings.py @@ -91,3 +91,33 @@ def test_sends_single_string_input(self, httpx_mock: HTTPXMock): embed.embed_query("hello") body = json.loads(httpx_mock.get_requests()[0].content) assert body["input"] == "hello" + + +class TestAsync: + async def test_aembed_documents_preserves_order(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=f"{BASE_URL}/v1/embeddings", + json={ + "object": "list", + "model": "text-embedding-3-small", + "data": [ + {"index": 1, "embedding": [0.3], "object": "embedding"}, + {"index": 0, "embedding": [0.1], "object": "embedding"}, + ], + }, + ) + assert await _build().aembed_documents(["a", "b"]) == [[0.1], [0.3]] + + async def test_aembed_query(self, httpx_mock: HTTPXMock): + httpx_mock.add_response( + method="POST", + url=f"{BASE_URL}/v1/embeddings", + json=make_embedding_response(vectors=[[0.9]]), + ) + assert await _build().aembed_query("hello") == [0.9] + assert json.loads(httpx_mock.get_requests()[0].content)["input"] == "hello" + + async def test_aembed_documents_empty_makes_no_request(self, httpx_mock: HTTPXMock): + assert await _build().aembed_documents([]) == [] + assert httpx_mock.get_requests() == [] From 4900f1a0ed72c02c9ec82fe4f68e95635cce5eed Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 13:02:34 +0530 Subject: [PATCH 03/14] test: contract suite against a real ai-gateway + CI contract job - scripts/with-gateway.sh (adapted from gateway-cli): builds ferrogw from FERRO_GATEWAY_SOURCE, starts tests/contract/stub_upstream.py as a stdlib fake OpenAI, boots the gateway with MASTER_KEY + sqlite request log + OPENAI_BASE_URL pointing at the stub, verifies /readyz and /v1/models, runs tests/contract, always tears down. Ports 18080/18081. - tests/contract/test_contract.py (23 tests, skipped without FERRO_CONTRACT_BASE_URL): probes, capabilities, EnrichedModelInfo, client-side retrieve proven via the stub's request log, chat/streaming header+usage contract, embeddings, responses (+501 responses_not_configured), 401/403/404 envelope, admin keys/config/logs/providers/plugins/audit. - Fix: /admin/providers/catalog is wrapped as {"providers": [...]}. - ci.yml: 3.13 in the matrix, mypy + ruff format on every leg, contract job vs ai-gateway v1.4.5 (required) and main (advisory); publish needs both. - make contract; adapter publish workflow gains 3.13 and all-leg mypy. --- .github/workflows/ci.yml | 50 +++- .../publish-langchain-ferrolabsai.yml | 3 +- Makefile | 8 +- ferrolabsai/admin/async_resource.py | 4 +- ferrolabsai/admin/resource.py | 2 +- scripts/with-gateway.sh | 139 ++++++++++ tests/contract/__init__.py | 0 tests/contract/conftest.py | 68 +++++ tests/contract/stub_upstream.py | 187 +++++++++++++ tests/contract/test_contract.py | 252 ++++++++++++++++++ tests/test_admin.py | 4 +- 11 files changed, 704 insertions(+), 13 deletions(-) create mode 100755 scripts/with-gateway.sh create mode 100644 tests/contract/__init__.py create mode 100644 tests/contract/conftest.py create mode 100644 tests/contract/stub_upstream.py create mode 100644 tests/contract/test_contract.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 09adf68..be5735e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -16,7 +16,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ["3.9", "3.10", "3.11", "3.12"] + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"] steps: - uses: actions/checkout@v4 @@ -31,19 +31,57 @@ jobs: pip install -e ".[dev]" - name: Lint - run: ruff check ferrolabsai/ + run: ruff check ferrolabsai/ tests/ && ruff format --check ferrolabsai/ tests/ - name: Type check - run: mypy ferrolabsai/ --ignore-missing-imports - # Only enforce strict types on 3.11+ - if: matrix.python-version == '3.11' + run: mypy ferrolabsai/ - name: Run tests run: pytest tests/ -v --tb=short + contract: + name: Contract vs AI Gateway ${{ matrix.gateway_ref }} + runs-on: ubuntu-latest + # The pinned leg is the required signal; "main" moves upstream of this + # repo, so a failure there must only report drift, never block a merge or + # a tagged release. + continue-on-error: ${{ matrix.gateway_ref == 'main' }} + strategy: + fail-fast: false + matrix: + gateway_ref: ["v1.4.5", "main"] + steps: + - name: Check out ferrolabs-python-sdk + uses: actions/checkout@v4 + with: + path: ferrolabs-python-sdk + persist-credentials: false + - name: Check out AI Gateway + uses: actions/checkout@v4 + with: + repository: ferro-labs/ai-gateway + ref: ${{ matrix.gateway_ref }} + path: ai-gateway + persist-credentials: false + - uses: actions/setup-go@v5 + with: + go-version-file: ai-gateway/go.mod + - uses: actions/setup-python@v5 + with: + python-version: "3.12" + - name: Install the SDK + run: pip install -e ".[dev]" + working-directory: ferrolabs-python-sdk + - name: Boot the gateway and run the contract suite + run: ./scripts/with-gateway.sh + working-directory: ferrolabs-python-sdk + env: + FERRO_GATEWAY_SOURCE: ${{ github.workspace }}/ai-gateway + publish: name: Publish to PyPI - needs: test + # A release must pass the unit matrix and the pinned-gateway contract leg. + needs: [test, contract] runs-on: ubuntu-latest # Only run on semver tag pushes (see `on.push.tags` above). if: startsWith(github.ref, 'refs/tags/v') diff --git a/.github/workflows/publish-langchain-ferrolabsai.yml b/.github/workflows/publish-langchain-ferrolabsai.yml index 14ccd32..9613f0d 100644 --- a/.github/workflows/publish-langchain-ferrolabsai.yml +++ b/.github/workflows/publish-langchain-ferrolabsai.yml @@ -17,7 +17,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ["3.9", "3.10", "3.11", "3.12"] + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"] defaults: run: @@ -41,7 +41,6 @@ jobs: - name: Type check run: mypy langchain_ferrolabsai/ --ignore-missing-imports - if: matrix.python-version == '3.11' - name: Run tests run: pytest tests/ -v --tb=short diff --git a/Makefile b/Makefile index 19d4304..5981a0f 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: install test lint format build clean +.PHONY: install test lint format build contract clean install: pip install -e ".[dev]" @@ -8,6 +8,7 @@ test: lint: ruff check ferrolabsai tests + ruff format --check ferrolabsai tests mypy ferrolabsai format: @@ -16,5 +17,10 @@ format: build: python3 -m build +# Boots a real ai-gateway (FERRO_GATEWAY_SOURCE, default ../ai-gateway) against +# the stub upstream and runs tests/contract. Needs go and curl on PATH. +contract: + ./scripts/with-gateway.sh + clean: rm -rf build/ dist/ *.egg-info .pytest_cache .mypy_cache .ruff_cache diff --git a/ferrolabsai/admin/async_resource.py b/ferrolabsai/admin/async_resource.py index 09f3519..93876b6 100644 --- a/ferrolabsai/admin/async_resource.py +++ b/ferrolabsai/admin/async_resource.py @@ -172,7 +172,9 @@ async def list(self) -> builtins.list[dict[str, Any]]: return items(await self._client._request("GET", "/admin/providers"), "data", "providers") async def catalog(self) -> builtins.list[dict[str, Any]]: - return items(await self._client._request("GET", "/admin/providers/catalog"), "data") + return items( + await self._client._request("GET", "/admin/providers/catalog"), "providers", "data" + ) class _AsyncPluginsResource: diff --git a/ferrolabsai/admin/resource.py b/ferrolabsai/admin/resource.py index 49fe2dc..1641cd2 100644 --- a/ferrolabsai/admin/resource.py +++ b/ferrolabsai/admin/resource.py @@ -377,7 +377,7 @@ def list(self) -> builtins.list[dict[str, Any]]: def catalog(self) -> builtins.list[dict[str, Any]]: """``GET /admin/providers/catalog`` — every provider the build knows: ``{id, registered, catalog_models}``.""" - return items(self._client._request("GET", "/admin/providers/catalog"), "data") + return items(self._client._request("GET", "/admin/providers/catalog"), "providers", "data") class _PluginsResource: diff --git a/scripts/with-gateway.sh b/scripts/with-gateway.sh new file mode 100755 index 0000000..521a4d4 --- /dev/null +++ b/scripts/with-gateway.sh @@ -0,0 +1,139 @@ +#!/usr/bin/env bash +# Boot a real AI Gateway checkout against a stub upstream, run the SDK +# contract suite against it, and always tear both down. +# +# FERRO_GATEWAY_SOURCE=../ai-gateway ./scripts/with-gateway.sh [pytest args] +# +# Adapted from gateway-cli/scripts/with-gateway.sh. This is the one check the +# pytest-httpx unit suite cannot perform: the unit tests prove the SDK is +# correct given a gateway that behaves as documented; only this proves the +# real server agrees. A divergence found here is a contract drift — fix the +# SDK (or the docs) to match the gateway, then re-run. +# +# No provider credentials are required. tests/contract/stub_upstream.py plays +# the OpenAI API and the gateway is pointed at it via OPENAI_BASE_URL, which +# is enough to exercise probes, the catalog, chat, streaming, embeddings, +# responses, the error envelope, and every /admin/* route the SDK wraps. +# +# Env: FERRO_GATEWAY_SOURCE (default ../ai-gateway), FERRO_CONTRACT_PORT +# (default 18080), FERRO_CONTRACT_STUB_PORT (default 18081), PYTHON (default +# python3 — point it at a venv interpreter with ferrolabsai[dev] installed). +set -euo pipefail + +for tool in curl go; do + command -v "$tool" >/dev/null || { echo "$tool is required" >&2; exit 2; } +done +python="${PYTHON:-python3}" +command -v "$python" >/dev/null || { echo "$python is required (set PYTHON)" >&2; exit 2; } + +sdk="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +gateway_source="${FERRO_GATEWAY_SOURCE:-}" +if [ -z "$gateway_source" ]; then + candidate="$(cd "$sdk/.." && pwd)/ai-gateway" + if [ -f "$candidate/go.mod" ] && grep -q '^module github.com/ferro-labs/ai-gateway$' "$candidate/go.mod"; then + gateway_source="$candidate" + else + echo "FERRO_GATEWAY_SOURCE must point to an AI Gateway checkout" >&2 + exit 2 + fi +fi +gateway_source="$(cd "$gateway_source" && pwd)" +work="$(mktemp -d)" +port="${FERRO_CONTRACT_PORT:-18080}" +stub_port="${FERRO_CONTRACT_STUB_PORT:-18081}" +gw_pid="" +stub_pid="" + +# A hex master key of the shape the gateway expects (fgw_ + 32 hex chars). +key="fgw_$(od -An -tx1 -N16 /dev/urandom | tr -d ' \n')" + +cleanup() { + for pid in "$gw_pid" "$stub_pid"; do + if [ -n "$pid" ]; then + kill "$pid" 2>/dev/null || true + wait "$pid" 2>/dev/null || true + fi + done + rm -rf "$work" +} +trap cleanup EXIT + +# Poll a URL until it answers with one of the given status codes. +wait_for() { + local url="$1" pid="$2" name="$3"; shift 3 + local deadline=$((SECONDS + 30)) code + while (( SECONDS < deadline )); do + if ! kill -0 "$pid" 2>/dev/null; then + echo "$name exited during startup; log follows:" >&2 + cat "$work/$name.log" >&2 + exit 1 + fi + code="$(curl -sS --connect-timeout 1 --max-time 2 -o /dev/null -w '%{http_code}' "$url" 2>/dev/null)" || code="" + for ok in "$@"; do [ "$code" = "$ok" ] && return 0; done + sleep 0.25 + done + echo "$name did not answer $url within 30s; log follows:" >&2 + cat "$work/$name.log" >&2 + exit 1 +} + +echo "==> starting stub upstream on :$stub_port" +"$python" "$sdk/tests/contract/stub_upstream.py" --port "$stub_port" >"$work/stub.log" 2>&1 & +stub_pid=$! +wait_for "http://127.0.0.1:$stub_port/v1/models" "$stub_pid" stub 200 + +echo "==> building ferrogw from $gateway_source" +(cd "$gateway_source" && go build -o "$work/ferrogw" ./cmd/ferrogw) + +echo "==> writing a throwaway config (one openai target routed to the stub)" +# persist: true so /admin/logs and /admin/logs/stats have rows to return; the +# store itself is the SQLite file named by REQUEST_LOG_STORE_DSN below. +cat >"$work/gateway.yaml" <<'YAML' +apiVersion: v1 +strategy: + mode: single +targets: + - virtual_key: openai +plugins: + - name: request-logger + type: logging + stage: before_request + enabled: true + config: { level: info, persist: true } + - name: request-logger + type: logging + stage: after_request + enabled: true + config: { level: info, persist: true } + - name: request-logger + type: logging + stage: on_error + enabled: true + config: { level: info, persist: true } +YAML + +echo "==> starting gateway on :$port" +MASTER_KEY="$key" GATEWAY_CONFIG="$work/gateway.yaml" PORT="$port" \ + REQUEST_LOG_STORE_BACKEND=sqlite REQUEST_LOG_STORE_DSN="$work/requestlog.db" \ + OPENAI_API_KEY=stub-key OPENAI_BASE_URL="http://127.0.0.1:$stub_port/v1" \ + "$work/ferrogw" serve >"$work/gateway.log" 2>&1 & +gw_pid=$! + +# /health answers 503 when degraded (still JSON, still "up"); /readyz must be 200 +# because the whole point of the stub is a routable target. +wait_for "http://127.0.0.1:$port/health" "$gw_pid" gateway 200 503 +wait_for "http://127.0.0.1:$port/readyz" "$gw_pid" gateway 200 + +echo "==> verifying the stub's models are routable through the gateway" +if ! curl -sS -H "Authorization: Bearer $key" "http://127.0.0.1:$port/v1/models" | grep -q '"gpt-4o-mini"'; then + echo "gateway /v1/models does not list the stub's gpt-4o-mini; log follows:" >&2 + cat "$work/gateway.log" >&2 + exit 1 +fi + +echo "==> running the contract suite" +cd "$sdk" +FERRO_CONTRACT_BASE_URL="http://127.0.0.1:$port" \ + FERRO_CONTRACT_MASTER_KEY="$key" \ + FERRO_CONTRACT_STUB_URL="http://127.0.0.1:$stub_port" \ + "$python" -m pytest tests/contract -v "$@" diff --git a/tests/contract/__init__.py b/tests/contract/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/contract/conftest.py b/tests/contract/conftest.py new file mode 100644 index 0000000..906d2db --- /dev/null +++ b/tests/contract/conftest.py @@ -0,0 +1,68 @@ +"""Contract-suite fixtures. Everything here talks to a real ai-gateway. + +Set by ``scripts/with-gateway.sh``: + FERRO_CONTRACT_BASE_URL gateway URL (e.g. http://127.0.0.1:18080) + FERRO_CONTRACT_MASTER_KEY the gateway's MASTER_KEY (admin bearer) + FERRO_CONTRACT_STUB_URL the stub upstream (optional; enables the + "never reached upstream" assertions) + +Without FERRO_CONTRACT_BASE_URL the whole directory is skipped, so the plain +``pytest`` run stays offline. +""" + +from __future__ import annotations + +import os +from collections.abc import Iterator + +import httpx +import pytest + +from ferrolabsai import AsyncFerroClient, FerroClient + +BASE_URL = os.environ.get("FERRO_CONTRACT_BASE_URL") +MASTER_KEY = os.environ.get("FERRO_CONTRACT_MASTER_KEY", "") +STUB_URL = os.environ.get("FERRO_CONTRACT_STUB_URL") + +CHAT_MODEL = "gpt-4o-mini" +EMBED_MODEL = "text-embedding-3-small" + + +def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None: + if BASE_URL: + return + skip = pytest.mark.skip(reason="FERRO_CONTRACT_BASE_URL not set (run scripts/with-gateway.sh)") + for item in items: + if "tests/contract" in str(item.fspath).replace(os.sep, "/"): + item.add_marker(skip) + + +@pytest.fixture(scope="session") +def client() -> Iterator[FerroClient]: + with FerroClient(api_key=MASTER_KEY, base_url=BASE_URL, max_retries=0) as c: + yield c + + +@pytest.fixture(scope="session") +def admin_guard(client: FerroClient) -> str: + """The gateway refuses to revoke/delete the last admin *record* (409) — the + MASTER_KEY is not a record — so tests that delete admin keys need one + extra admin key parked for the whole session.""" + return client.admin.keys.create(name="contract-guard", scopes=["admin"]).id + + +@pytest.fixture +async def async_client() -> AsyncFerroClient: + return AsyncFerroClient(api_key=MASTER_KEY, base_url=BASE_URL, max_retries=0) + + +@pytest.fixture +def stub_requests(): + """Callable returning the stub upstream's request log (``["GET /v1/models", ...]``).""" + if not STUB_URL: + pytest.skip("FERRO_CONTRACT_STUB_URL not set") + + def _read() -> list[str]: + return list(httpx.get(f"{STUB_URL}/_requests", timeout=5).json()["data"]) + + return _read diff --git a/tests/contract/stub_upstream.py b/tests/contract/stub_upstream.py new file mode 100644 index 0000000..ce106e1 --- /dev/null +++ b/tests/contract/stub_upstream.py @@ -0,0 +1,187 @@ +#!/usr/bin/env python3 +"""Stub OpenAI-compatible upstream for the contract suite (stdlib only). + +The gateway is pointed at this server with ``OPENAI_BASE_URL`` so chat, +streaming, embeddings, and the Responses API can be exercised end to end +without provider credentials. It serves a fixed model list and canned +replies, and records every request at ``GET /_requests`` so tests can prove +what did (and did not) reach upstream — e.g. that ``models.retrieve()`` never +turns into a pass-through ``GET /v1/models/{id}``. + + python3 tests/contract/stub_upstream.py --port 18081 +""" + +from __future__ import annotations + +import argparse +import json +import sys +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any + +MODELS = ["gpt-4o-mini", "text-embedding-3-small"] +REPLY = ("stub", " reply") +USAGE = {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} +_requests: list[str] = [] +_lock = threading.Lock() + + +def _chunk(model: str, delta: dict[str, Any], finish: str | None) -> dict[str, Any]: + return { + "id": "chatcmpl-stub", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": model, + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}], + } + + +class Handler(BaseHTTPRequestHandler): + def log_message(self, *_: Any) -> None: # keep stdout quiet + pass + + def _record(self) -> None: + with _lock: + _requests.append(f"{self.command} {self.path}") + + def _json(self, status: int, body: dict[str, Any]) -> None: + payload = json.dumps(body).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(payload))) + self.end_headers() + self.wfile.write(payload) + + def _read_body(self) -> dict[str, Any]: + length = int(self.headers.get("Content-Length") or 0) + raw = self.rfile.read(length) if length else b"{}" + data = json.loads(raw or b"{}") + return data if isinstance(data, dict) else {} + + def _not_found(self) -> None: + self._json( + 404, + { + "error": { + "message": f"stub has no route {self.command} {self.path}", + "type": "invalid_request_error", + "code": "not_found", + } + }, + ) + + def do_GET(self) -> None: + self._record() + if self.path == "/v1/models": + data = [ + {"id": m, "object": "model", "created": 1700000000, "owned_by": "openai"} + for m in MODELS + ] + self._json(200, {"object": "list", "data": data}) + elif self.path == "/_requests": + with _lock: + self._json(200, {"data": list(_requests)}) + else: + self._not_found() + + def do_POST(self) -> None: + self._record() + body = self._read_body() + model = str(body.get("model", "gpt-4o-mini")) + if self.path == "/v1/chat/completions": + if body.get("stream"): + include_usage = bool((body.get("stream_options") or {}).get("include_usage")) + self._sse(model, include_usage) + else: + self._json( + 200, + { + "id": "chatcmpl-stub", + "object": "chat.completion", + "created": 1700000000, + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "".join(REPLY)}, + "finish_reason": "stop", + } + ], + "usage": USAGE, + }, + ) + elif self.path == "/v1/embeddings": + inputs = body.get("input", []) + count = len(inputs) if isinstance(inputs, list) else 1 + self._json( + 200, + { + "object": "list", + "model": model, + "data": [ + {"index": i, "object": "embedding", "embedding": [0.1, 0.2, 0.3]} + for i in range(count) + ], + "usage": {"prompt_tokens": count, "total_tokens": count}, + }, + ) + elif self.path == "/v1/responses": + self._json( + 200, + { + "id": "resp_stub", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": model, + "output": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "".join(REPLY)}], + } + ], + "usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}, + }, + ) + else: + self._not_found() + + def _sse(self, model: str, include_usage: bool) -> None: + frames = [_chunk(model, {"role": "assistant", "content": REPLY[0]}, None)] + frames.append(_chunk(model, {"content": REPLY[1]}, None)) + frames.append(_chunk(model, {}, "stop")) + if include_usage: + terminal = _chunk(model, {}, None) + terminal["choices"] = [] + terminal["usage"] = USAGE + frames.append(terminal) + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Cache-Control", "no-cache") + self.end_headers() + for frame in frames: + self.wfile.write(f"data: {json.dumps(frame)}\n\n".encode()) + self.wfile.write(b"data: [DONE]\n\n") + self.wfile.flush() + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--port", type=int, default=18081) + parser.add_argument("--host", default="127.0.0.1") + args = parser.parse_args() + server = ThreadingHTTPServer((args.host, args.port), Handler) + print(f"stub upstream listening on http://{args.host}:{args.port}", flush=True) + try: + server.serve_forever() + except KeyboardInterrupt: + pass + finally: + server.server_close() + sys.exit(0) + + +if __name__ == "__main__": + main() diff --git a/tests/contract/test_contract.py b/tests/contract/test_contract.py new file mode 100644 index 0000000..3eb11fa --- /dev/null +++ b/tests/contract/test_contract.py @@ -0,0 +1,252 @@ +"""Contract tests against a real ai-gateway (see conftest.py and scripts/with-gateway.sh). + +Every field the README's observability section names is asserted non-empty +here; if the gateway changes a header or body field, this suite goes red +before a user does. +""" + +from __future__ import annotations + +import re + +import pytest + +from ferrolabsai import ( + AsyncFerroClient, + FerroAPIError, + FerroAuthError, + FerroClient, + FerroNotFoundError, + FerroPermissionError, +) + +from .conftest import BASE_URL, CHAT_MODEL, EMBED_MODEL + +HEX32 = re.compile(r"^[0-9a-f]{32}$") +MESSAGES = [{"role": "user", "content": "ping"}] + + +class TestProbes: + def test_health(self, client: FerroClient): + health = client.health() + assert health["status"] + assert health["version"] + assert isinstance(health["providers"], list) + + def test_ready(self, client: FerroClient): + ready = client.ready() + assert ready["status"] == "ready", ready + assert any(t["routable"] for t in ready["targets"]) + + def test_live(self, client: FerroClient): + assert client.live() == {"status": "ok"} + + def test_capabilities(self, client: FerroClient): + caps = client.capabilities() + assert "openai" in caps["providers"] + assert set(caps["providers"]["openai"].values()) <= {"forward", "translate", "unsupported"} + + +class TestModels: + def test_list_returns_enriched_model_info(self, client: FerroClient): + models = client.models.list() + by_id = {m.id: m for m in models} + assert CHAT_MODEL in by_id and EMBED_MODEL in by_id + chat = by_id[CHAT_MODEL] + assert chat.owned_by == "openai" and chat.provider == "openai" + assert chat.object == "model" + assert isinstance(chat.created, int) + assert isinstance(chat.capabilities, list) + assert isinstance(chat.deprecated, bool) + # Catalog enrichment (internal/handler/models.go) — present for a real OpenAI id. + assert chat.mode == "chat" + assert chat.context_window and chat.context_window > 0 + + def test_filters_are_client_side(self, client: FerroClient): + assert {m.owned_by for m in client.models.list(provider="openai")} == {"openai"} + assert client.models.list(provider="no-such-provider") == [] + # The catalog is a union of live discovery and the static model catalog, + # so more embedding ids than the stub's one show up. + found = [m.id for m in client.models.search("EMBEDDING")] + assert EMBED_MODEL in found and all("embedding" in m for m in found) + assert client.models.retrieve(CHAT_MODEL).id == CHAT_MODEL + + def test_retrieve_never_reaches_upstream(self, client: FerroClient, stub_requests): + before = len(stub_requests()) + client.models.retrieve(CHAT_MODEL) + with pytest.raises(FerroNotFoundError) as exc_info: + client.models.retrieve("no-such-model") + assert exc_info.value.code == "model_not_found" + assert exc_info.value.status_code == 404 + new = stub_requests()[before:] + assert not [r for r in new if "/v1/models/" in r], new + + +class TestChat: + def test_non_streaming_carries_gateway_metadata(self, client: FerroClient): + response = client.chat.completions.create(model=CHAT_MODEL, messages=MESSAGES) + assert response.content == "stub reply" + assert response.model + # README "Observability" table — every row asserted here. + assert response.trace_id and HEX32.match(response.trace_id) + assert response.provider == "openai" + assert response.gateway_overhead_ms is not None and response.gateway_overhead_ms > 0 + assert response.usage is not None + assert response.usage.prompt_tokens == 5 + assert response.usage.completion_tokens == 2 + assert response.usage.total_tokens == 7 + + def test_streaming_carries_trace_id_and_terminal_usage(self, client: FerroClient): + stream = client.chat.completions.create( + model=CHAT_MODEL, + messages=MESSAGES, + stream=True, + stream_options={"include_usage": True}, + ) + assert stream.trace_id and HEX32.match(stream.trace_id) + chunks = list(stream) + assert chunks, "no chunks" + text = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices) + assert text == "stub reply" + assert all(c.trace_id == stream.trace_id for c in chunks) + assert chunks[-1].usage is not None and chunks[-1].usage.total_tokens == 7 + assert stream.response.is_closed + + def test_streaming_usage_chunk_policy(self, client: FerroClient): + # The gateway always requests usage upstream for metering and forwards the + # terminal usage chunk unless the client explicitly opts out + # (providers/openai/openai.go streamOptions, internal/streamwrap.Meter). + default = list( + client.chat.completions.create(model=CHAT_MODEL, messages=MESSAGES, stream=True) + ) + assert default[-1].usage is not None and default[-1].usage.total_tokens == 7 + opted_out = list( + client.chat.completions.create( + model=CHAT_MODEL, + messages=MESSAGES, + stream=True, + stream_options={"include_usage": False}, + ) + ) + assert opted_out and all(c.usage is None for c in opted_out) + + def test_embeddings(self, client: FerroClient): + response = client.embeddings.create(model=EMBED_MODEL, input=["a", "b"]) + assert [d.index for d in response.data] == [0, 1] + assert response.trace_id and HEX32.match(response.trace_id) + + def test_responses_create(self, client: FerroClient): + response = client.responses.create(model=CHAT_MODEL, input="ping") + assert response.id == "resp_stub" + assert response.status == "completed" + assert response.output[0]["content"][0]["text"] == "stub reply" + assert response.trace_id and HEX32.match(response.trace_id) + assert response.provider == "openai" + + def test_responses_id_routes_are_501_without_responses_target(self, client: FerroClient): + with pytest.raises(FerroAPIError) as exc_info: + client.responses.retrieve("resp_stub") + assert exc_info.value.status_code == 501 + assert exc_info.value.code == "responses_not_configured" + + async def test_async_chat_and_stream(self, async_client: AsyncFerroClient): + async with async_client: + response = await async_client.chat.completions.create( + model=CHAT_MODEL, messages=MESSAGES + ) + assert response.content == "stub reply" and response.provider == "openai" + stream = await async_client.chat.completions.create( + model=CHAT_MODEL, messages=MESSAGES, stream=True + ) + assert stream.trace_id and HEX32.match(stream.trace_id) + chunks = [c async for c in stream] + text = "".join(c.choices[0].delta.content or "" for c in chunks if c.choices) + assert text == "stub reply" + + +class TestErrorEnvelope: + def test_401(self): + with FerroClient(api_key="fgw_" + "0" * 32, base_url=BASE_URL, max_retries=0) as bad: + with pytest.raises(FerroAuthError) as exc_info: + bad.models.list() + assert exc_info.value.status_code == 401 + assert exc_info.value.code == "invalid_api_key" + assert exc_info.value.request_id and HEX32.match(exc_info.value.request_id) + + def test_404_unknown_model(self, client: FerroClient): + with pytest.raises(FerroNotFoundError) as exc_info: + client.chat.completions.create(model="no-such-model", messages=MESSAGES) + assert exc_info.value.code == "model_not_found" + + def test_403_read_only_scope(self, client: FerroClient): + created = client.admin.keys.create(name="contract-read-only", scopes=["read_only"]) + try: + with FerroClient(api_key=created.key, base_url=BASE_URL, max_retries=0) as ro: + assert isinstance(ro.admin.keys.list(), list) + with pytest.raises(FerroPermissionError) as exc_info: + ro.admin.keys.create(name="should-fail") + assert exc_info.value.status_code == 403 + assert exc_info.value.code == "insufficient_scope" + finally: + client.admin.keys.delete(created.id) + + +class TestAdmin: + def test_keys_lifecycle(self, client: FerroClient, admin_guard: str): + created = client.admin.keys.create(name="contract-key", scopes=["admin"]) + assert created.key.startswith("fgw_") + try: + fetched = client.admin.keys.retrieve(created.id) + assert fetched.id == created.id and fetched.active + assert fetched.key and fetched.key != created.key and "..." in fetched.key + assert ( + client.admin.keys.update(created.id, name="contract-key-2").name == "contract-key-2" + ) + rotated = client.admin.keys.rotate(created.id) + assert rotated.key.startswith("fgw_") and rotated.key != created.key + client.admin.keys.revoke(created.id) + assert client.admin.keys.retrieve(created.id).active is False + assert any(k.id == created.id for k in client.admin.keys.list()) + assert "summary" in client.admin.keys.usage(limit=5) + finally: + client.admin.keys.delete(created.id) + with pytest.raises(FerroNotFoundError): + client.admin.keys.retrieve(created.id) + + def test_config_get(self, client: FerroClient): + cfg = client.admin.config.get() + assert any(t.get("virtual_key") == "openai" for t in cfg.targets) + assert cfg.raw["strategy"] == cfg.strategy + + def test_logs_list_and_stats(self, client: FerroClient): + response = client.chat.completions.create(model=CHAT_MODEL, messages=MESSAGES) + logs = client.admin.logs.list(limit=20, model=CHAT_MODEL) + assert any(e.get("trace_id") == response.trace_id for e in logs["data"]), logs + # The master key is logged as "master-key:", not "none". + key_id = logs["data"][0]["api_key_id"] + assert key_id.startswith("master-key:") + filtered = client.admin.logs.list(limit=5, stage="all", api_key_id=key_id)["data"] + assert filtered and all(e["api_key_id"] == key_id for e in filtered) + stats = client.admin.logs.stats(buckets=4) + assert isinstance(stats, dict) and stats + + def test_providers_and_plugins(self, client: FerroClient): + assert any(p.get("name") == "openai" for p in client.admin.providers.list()) + catalog = client.admin.providers.catalog() + openai = next(p for p in catalog if p["id"] == "openai") + assert openai["registered"] is True + assert isinstance(client.admin.plugins.list(), list) + assert any(p.get("name") for p in client.admin.plugins.catalog()) + + def test_audit_list(self, client: FerroClient, admin_guard: str): + created = client.admin.keys.create(name="contract-audit", scopes=["admin"]) + client.admin.keys.delete(created.id) + audit = client.admin.audit.list(limit=50) + assert audit["summary"]["total_entries"] >= 1 + assert any(e.get("actor_id") for e in audit["data"]) + action = audit["data"][0]["action"] + assert all(e["action"] == action for e in client.admin.audit.list(action=action)["data"]) + + def test_dashboard_and_health(self, client: FerroClient): + assert client.admin.dashboard() + assert client.admin.health()["status"] diff --git a/tests/test_admin.py b/tests/test_admin.py index 5b85e2b..9d1b45a 100644 --- a/tests/test_admin.py +++ b/tests/test_admin.py @@ -221,7 +221,7 @@ def test_catalogs(self, client, httpx_mock: HTTPXMock): httpx_mock.add_response( method="GET", url=f"{ADMIN}/providers/catalog", - json=[{"id": "openai", "registered": True, "catalog_models": 90}], + json={"providers": [{"id": "openai", "registered": True, "catalog_models": 90}]}, ) httpx_mock.add_response( method="GET", url=f"{ADMIN}/plugins/catalog", json={"data": [{"name": "budget"}]} @@ -253,7 +253,7 @@ async def test_async_parity(self, async_client, httpx_mock: HTTPXMock): method="GET", url=f"{ADMIN}/providers", json={"data": [{"name": "openai"}]} ) httpx_mock.add_response( - method="GET", url=f"{ADMIN}/providers/catalog", json=[{"id": "openai"}] + method="GET", url=f"{ADMIN}/providers/catalog", json={"providers": [{"id": "openai"}]} ) httpx_mock.add_response(method="GET", url=f"{ADMIN}/plugins/catalog", json={"data": []}) httpx_mock.add_response( From c02ad0bdcd6f855ec3f8283975d681b2eba24e78 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 13:02:34 +0530 Subject: [PATCH 04/14] docs: truth pass for the 0.3.0 contract README: 30 providers, observability table lists exactly what populates and from which header/body field, dead 'templates & route tags' section removed, framework adapters section points at langchain-ferrolabsai, compatibility line (ferrolabsai 0.3.x <-> ai-gateway >= v1.4.0), retry policy, new surface, admin logs example fixed (no trace_id filter), test counts. docs/architecture.md: gateway contract tables (headers, body fields, catalog), retry policy, streaming, admin route table incl. audit/catalogs, handlers.go links -> internal/admin/handlers package. AGENT.md/CLAUDE.md: repo tree, conventions, pitfalls. SECURITY.md: 0.3.x. integrations/README.md: workflows exist. copilot-instructions refreshed. --- .github/copilot-instructions.md | 23 +- AGENT.md | 138 ++++---- README.md | 300 ++++++++--------- SECURITY.md | 7 +- docs/architecture.md | 330 +++++++++---------- integrations/README.md | 4 +- integrations/langchain-ferrolabsai/README.md | 14 +- 7 files changed, 402 insertions(+), 414 deletions(-) diff --git a/.github/copilot-instructions.md b/.github/copilot-instructions.md index 92b343c..ee445e6 100644 --- a/.github/copilot-instructions.md +++ b/.github/copilot-instructions.md @@ -3,31 +3,32 @@ ## Build, test, and lint commands - Install dev dependencies: `make install` or `pip install -e ".[dev]"` -- Run the full test suite: `make test` or `pytest tests/ -v` -- Run a single test: `pytest tests/test_sdk.py::TestChatCompletions::test_basic_create -v` -- Run a subset of tests by name: `pytest tests/test_sdk.py -k streaming -v` +- Run the full test suite: `make test` or `pytest tests/ -v` (the `tests/contract` dir auto-skips without a gateway) +- Run a single test: `pytest tests/test_chat.py::TestCreate::test_basic_create -v` +- Run a subset of tests by name: `pytest tests/test_chat.py -k streaming -v` +- Run the contract suite against a real gateway: `make contract` (builds `ferrogw` from `../ai-gateway`) - Lint and type-check: `make lint` or `ruff check ferrolabsai tests && mypy ferrolabsai` - Format: `make format` or `ruff format ferrolabsai tests` - Build the package: `make build` or `python3 -m build` ## High-level architecture -- `ferrolabsai/client.py` is the hub. It resolves API keys from `FERRO_API_KEY` with fallback to `OPENAI_API_KEY`, resolves `FERRO_BASE_URL` with a default of `http://localhost:8080`, owns the shared `httpx` client, applies retry/error handling, and wires the public namespaces. -- Public SDK namespaces intentionally mirror the OpenAI SDK and the gateway HTTP surface: `client.chat.completions`, `client.embeddings`, `client.images`, `client.models`, and `client.admin.*`. +- `ferrolabsai/client.py` is the hub. It resolves API keys from `FERRO_API_KEY` with fallback to `OPENAI_API_KEY`, resolves `FERRO_BASE_URL` with a default of `http://localhost:8080`, owns the shared `httpx` client, applies the retry policy (connect/timeout/408/429/5xx with jitter and `Retry-After`) and error mapping, merges the gateway's response headers into inference bodies, and wires the public namespaces. `ferrolabsai/streaming.py` wraps SSE responses in `Stream`/`AsyncStream`. +- Public SDK namespaces intentionally mirror the OpenAI SDK and the gateway HTTP surface: `client.chat.completions`, `client.embeddings`, `client.images`, `client.models`, `client.responses`, `client.moderations`, `client.rerank()`, the probes `client.health()/ready()/live()/capabilities()`, and `client.admin.*`. - Resource modules are thin request builders. Each `ferrolabsai//resource.py` file translates Python kwargs into the corresponding gateway request and delegates transport to `FerroClient._request(...)` or the streaming helpers instead of managing `httpx` behavior itself. -- `ferrolabsai/types.py` is the response normalization layer. Gateway JSON is converted into dataclasses there, including OpenAI-compatible fields plus Ferro-specific extras like `provider`, `trace_id`, `latency_ms`, `cost_usd`, and raw gateway config payloads. -- Admin support is a direct wrapper around `/admin/*` endpoints. The `Admin` namespace groups sub-resources for keys, config, logs, providers, and plugins, and some admin methods intentionally return raw dict payloads when the gateway response shape is still loosely structured. -- Tests in `tests/test_sdk.py` are request/response contract tests around the HTTP layer. They use `pytest-httpx` to verify headers, payloads, path routing, SSE streaming parsing, env-var fallbacks, retry validation, and admin endpoint behavior without requiring a live gateway. +- `ferrolabsai/types.py` is the response normalization layer. Gateway JSON is converted into dataclasses there, including OpenAI-compatible fields plus the gateway's real extensions: `provider`, `trace_id` (`X-Request-ID`), `gateway_overhead_ms` (`X-Gateway-Overhead-Ms`), `provider_metadata`, `reasoning_content`, the extra usage counters, and raw gateway config payloads. There is no cost, cache-hit, or latency field — the gateway does not expose them to callers. +- Admin support is a direct wrapper around `/admin/*` endpoints. The `Admin` namespace groups sub-resources for keys, config, logs, providers, plugins, and audit, and some admin methods intentionally return raw dict payloads when the gateway response shape is still loosely structured. +- Unit tests in `tests/test_{client,chat,resources,admin}.py` are request/response tests around the HTTP layer. They use `pytest-httpx` to verify headers, payloads, path routing, SSE streaming parsing, env-var fallbacks, the retry policy, and admin endpoint behavior without requiring a live gateway. `tests/contract/` runs the same claims against a real gateway booted by `scripts/with-gateway.sh`. ## Key conventions - Target Python 3.9 syntax. New modules should use `from __future__ import annotations`, public functions and methods are expected to be fully typed, and changes should stay compatible with the repo's `mypy --strict` setup. - Ruff is both the linter and formatter here. Keep changes aligned with the existing 100-character line length and prefer the repo's Ruff formatting over hand-formatting. - Preserve the OpenAI-style surface first. New capabilities should usually be exposed as additional kwargs on existing resource methods or as new resource namespaces that match gateway routes, not as a parallel custom API shape. -- Ferro-only request features are forwarded as request fields, not wrapped in separate helper abstractions. Existing examples are `template_id` and `template_variables`; sync chat completions also map `route_tag` to `x_route_tag`, but async parity is not complete yet. +- Only send request fields the gateway actually decodes (`internal/handler/chatrequest.go`). `max_completion_tokens`, `parallel_tool_calls`, `response_format`, `stream_options`, `seed` are first-class; anything else passes through `**kwargs` verbatim. Do not reintroduce `route_tag`, `template_id`, or `template_variables` — the gateway never read them. - Keep resource classes thin. Shared behavior such as retries, auth headers, connection handling, status-code mapping, and HTTP client lifecycle belongs in `client.py`, not duplicated across resource modules. - Public resource methods should return typed dataclasses parsed via `from_dict(...)` helpers in `types.py` unless the endpoint is intentionally passthrough admin data. - Error translation is centralized in `_raise_api_error(...)`. Extend the existing `Ferro*Error` hierarchy instead of leaking raw `httpx` exceptions from public SDK methods. -- Async support is added explicitly per namespace. Today `AsyncFerroClient` wires `chat.completions` and `embeddings`; if you add async support elsewhere, create the matching `async_resource.py`, register it in `AsyncFerroClient`, and add async tests. +- Async support is added explicitly per namespace. Every namespace has an `async_resource.py` that reuses the sync module's body builders and path constants; if you add a namespace, create both, register them in `FerroClient` and `AsyncFerroClient`, and add async tests. - When adding public clients, exceptions, or response types, update `ferrolabsai/__init__.py` and `__all__` so the package surface stays explicit and importable from the top level. -- Tests should mock HTTP precisely with `pytest_httpx.HTTPXMock` and assert the exact outgoing request shape, especially auth headers, routed endpoint paths, Ferro-specific fields, and SSE frames ending with `data: [DONE]`. +- Tests should mock HTTP precisely with `pytest_httpx.HTTPXMock` and assert the exact outgoing request shape, especially auth headers, routed endpoint paths, and SSE frames ending with `data: [DONE]`. Anything that claims a gateway behaviour also needs an assertion in `tests/contract/test_contract.py`. diff --git a/AGENT.md b/AGENT.md index 8cfc503..49a8a26 100644 --- a/AGENT.md +++ b/AGENT.md @@ -6,11 +6,12 @@ This document provides instructions for AI coding agents working on the `ferrola ## Project Overview -`ferrolabsai` is a drop-in replacement for the OpenAI Python SDK that routes LLM requests through the Ferro Labs AI Gateway to 29+ providers and 2,500+ models. The SDK exposes an OpenAI-compatible surface for chat completions, embeddings, images, and model catalog, plus Ferro-specific admin APIs for gateway management. +`ferrolabsai` is a drop-in replacement for the OpenAI Python SDK that routes LLM requests through the Ferro Labs AI Gateway to 30 providers and 2,500+ models. The SDK exposes an OpenAI-compatible surface for chat completions, embeddings, images, the Responses API, moderations, rerank, and the model catalog, plus gateway-specific surface: health probes, `/v1/capabilities`, and the `/admin/*` management API. - **Package name:** `ferrolabsai` -- **Version:** Defined in `pyproject.toml` under `[project].version` -- **Python support:** 3.9+ +- **Version:** `pyproject.toml` `[project].version` and `ferrolabsai/_version.py` (kept equal; a test asserts it) +- **Compatibility:** `ferrolabsai 0.3.x` ↔ `ai-gateway ≥ v1.4.0` (contract-tested against `v1.4.5`) +- **Python support:** 3.9 – 3.13 - **Only runtime dependency:** `httpx` - **License:** Apache-2.0 @@ -22,46 +23,46 @@ This document provides instructions for AI coding agents working on the `ferrola ferrolabs-python-sdk/ ├── ferrolabsai/ # Main package (publishes `ferrolabsai` to PyPI) │ ├── __init__.py # Public API surface — all exports live here -│ ├── client.py # FerroClient + AsyncFerroClient implementations -│ ├── types.py # Dataclass response models (ChatCompletion, Usage, etc.) -│ ├── completions/ # chat.completions resource (sync + async) -│ │ ├── resource.py # Completions (sync) -│ │ └── async_resource.py # AsyncCompletions -│ ├── embeddings/ # embeddings resource (sync + async) -│ │ ├── resource.py # Embeddings (sync) -│ │ └── async_resource.py # AsyncEmbeddings -│ ├── images/ # images resource -│ │ └── resource.py -│ ├── models/ # model catalog resource -│ │ └── resource.py -│ ├── admin/ # Admin API (keys, config, logs, providers, plugins) -│ │ └── resource.py +│ ├── _version.py # __version__ constant +│ ├── client.py # FerroClient + AsyncFerroClient, retry policy, error mapping +│ ├── streaming.py # Stream / AsyncStream SSE wrappers +│ ├── types.py # Dataclass response models (ChatCompletion, Usage, ModelInfo, ...) +│ ├── types_responses.py # Response dataclass (Responses API), re-exported from types +│ ├── completions/ # chat.completions (resource.py + async_resource.py) +│ ├── embeddings/ # embeddings +│ ├── images/ # images.generate +│ ├── models/ # model catalog — client-side list/retrieve/search +│ ├── responses/ # /v1/responses create/retrieve/delete +│ ├── moderations/ # /v1/moderations +│ ├── admin/ # Admin API (keys, config, logs, providers, plugins, audit) │ └── exceptions/ # Exception hierarchy -│ └── __init__.py ├── integrations/ # Sibling framework adapter packages (own pyproject.toml each) │ ├── README.md # Layout + publishing overview -│ ├── langchain-ferrolabsai/ # Publishes `langchain-ferrolabsai` to PyPI -│ │ ├── pyproject.toml -│ │ ├── README.md -│ │ ├── CHANGELOG.md -│ │ ├── LICENSE -│ │ ├── langchain_ferrolabsai/__init__.py -│ │ └── tests/test_placeholder.py -│ └── llama-index-llms-ferrolabsai/ # Publishes `llama-index-llms-ferrolabsai` to PyPI -│ ├── pyproject.toml -│ ├── README.md -│ ├── CHANGELOG.md -│ ├── LICENSE +│ ├── langchain-ferrolabsai/ # Publishes `langchain-ferrolabsai` to PyPI (0.2.0, on ferrolabsai 0.3) +│ │ ├── pyproject.toml / README.md / CHANGELOG.md / LICENSE +│ │ ├── langchain_ferrolabsai/{__init__,chat_models,embeddings,llms,_messages}.py +│ │ └── tests/ +│ └── llama-index-llms-ferrolabsai/ # Publishes `llama-index-llms-ferrolabsai` (placeholder 0.0.1) +│ ├── pyproject.toml / README.md / CHANGELOG.md / LICENSE │ ├── llama_index/llms/ferrolabsai/__init__.py # PEP 420 namespace package │ └── tests/test_placeholder.py ├── tests/ -│ └── test_sdk.py # Full test suite (pytest + pytest-httpx) +│ ├── conftest.py # Shared fixtures + canned gateway payloads +│ ├── test_client.py # Construction, retries, error mapping, header metadata, probes +│ ├── test_chat.py # Chat body, parsing, streaming (sync + async) +│ ├── test_resources.py # Embeddings, images, models, responses +│ ├── test_admin.py # /admin/* parity +│ └── contract/ # Real-gateway suite (skipped unless FERRO_CONTRACT_BASE_URL) +│ ├── conftest.py +│ ├── stub_upstream.py # stdlib fake OpenAI the gateway is pointed at +│ └── test_contract.py +├── scripts/with-gateway.sh # Builds ferrogw from ../ai-gateway, boots it + the stub, runs tests/contract ├── docs/ -│ └── architecture.md # SDK architecture, design decisions, request lifecycle +│ └── architecture.md # SDK architecture, gateway contract, request lifecycle ├── pyproject.toml # Build config, dependencies, tool settings (ferrolabsai) -├── Makefile # Dev shortcuts: install, test, lint, format, build, clean +├── Makefile # Dev shortcuts: install, test, lint, format, build, contract, clean ├── .github/workflows/ -│ ├── ci.yml # Core SDK CI + publish +│ ├── ci.yml # Unit matrix + contract job + publish │ ├── publish-langchain-ferrolabsai.yml # Test + publish on `langchain-ferrolabsai-vX.Y.Z` tags │ └── publish-llama-index-llms-ferrolabsai.yml # Test + publish on `llama-index-llms-ferrolabsai-vX.Y.Z` tags ├── README.md @@ -84,9 +85,9 @@ Framework adapter packages live under `integrations/` as independently versioned Conventions: - **Independent versioning.** Each sub-package's `pyproject.toml` version is bumped on its own cadence. Do not piggy-back on the parent SDK's `v*.*.*` tag. -- **Tag prefix pattern.** Releases are cut by pushing tags like `langchain-ferrolabsai-v0.1.0`. The publish workflow asserts the tag matches the sub-package's `pyproject.toml` version before uploading. +- **Tag prefix pattern.** Releases are cut by pushing tags like `langchain-ferrolabsai-v0.2.0`. The publish workflow asserts the tag matches the sub-package's `pyproject.toml` version before uploading. - **Trusted Publishing.** PyPI credentials use OIDC environments named `pypi-langchain-ferrolabsai` and `pypi-llama-index-llms-ferrolabsai` — provision these on PyPI before the first publish. -- **Dependency on the core SDK.** Each sub-package declares `ferrolabsai>=X.Y.Z` as a runtime dependency. +- **Dependency on the core SDK.** Each sub-package declares `ferrolabsai>=X.Y.Z` as a runtime dependency (`langchain-ferrolabsai` needs `>=0.3.0`). - **`llama_index` namespace package.** `llama-index-llms-ferrolabsai` uses the PEP 420 implicit namespace package layout (`llama_index/llms/ferrolabsai/`) so it can later be mirrored upstream into `run-llama/llama_index` with no code changes. - **Placeholder behaviour.** While a sub-package is at `0.0.x`, its `__init__.py` exposes only `__version__`; any attempt to import the planned public classes raises `NotImplementedError` with a roadmap link. Real classes land at `0.1.0`. - **Release flow.** `make build` and `make test` from the sub-folder; bump version in its `pyproject.toml` + `CHANGELOG.md`; commit + push tag with the prefix above; CI runs the full test matrix and publishes via Trusted Publishing. @@ -107,16 +108,17 @@ Dev dependencies: `pytest`, `pytest-asyncio`, `pytest-httpx`, `mypy`, `ruff`. ## Key Commands -| Command | Purpose | -| -------------- | ------------------------------------------------- | -| `make install` | Editable install with dev extras | -| `make test` | Run pytest suite | -| `make lint` | Run `ruff check` + `mypy` | -| `make format` | Run `ruff format` | -| `make build` | Build sdist + wheel into `dist/` | -| `make clean` | Remove build artifacts and tool caches | +| Command | Purpose | +| ---------------- | ------------------------------------------------------------------ | +| `make install` | Editable install with dev extras | +| `make test` | Run the unit suite (contract dir auto-skips) | +| `make lint` | Run `ruff check` + `mypy` | +| `make format` | Run `ruff format` | +| `make build` | Build sdist + wheel into `dist/` | +| `make contract` | Boot a real gateway from `../ai-gateway` and run `tests/contract` | +| `make clean` | Remove build artifacts and tool caches | -Always run `make format lint test` before committing. +Always run `make format lint test` before committing; run `make contract` when touching anything that talks to the gateway. --- @@ -124,22 +126,23 @@ Always run `make format lint test` before committing. ### Language & Style - **Python 3.9+ syntax** — use `dict[str, X]`, `X | None`, `from __future__ import annotations`. -- **Type annotations are mandatory** on every public function, method, and class attribute. `mypy --strict` must pass. +- **Type annotations are mandatory** on every public function, method, and class attribute. `mypy --strict` must pass on every CI leg (3.9 – 3.13). - **Ruff** handles linting and formatting. Line length is **100** characters. Select rules: `E`, `F`, `I`, `UP`. - Do not hand-format — run `make format`. ### Architecture Patterns - **Dataclass response models** — all types in `types.py` are `@dataclass` with a `from_dict()` classmethod. No pydantic dependency. -- **Resource pattern** — each API area (completions, embeddings, images, models, admin) lives in its own sub-package with a `resource.py` (sync) and optionally `async_resource.py`. -- **Client holds HTTP** — `FerroClient._request()` and `AsyncFerroClient._request()` are the only HTTP entry points. Resources receive the client instance and call `self._client._request(...)`. -- **Exception hierarchy** — all HTTP errors raise typed exceptions inheriting from `FerroAPIError` (which inherits from `FerroError`). Connection/timeout errors retry automatically and raise `FerroConnectionError`. -- **Immutability by default** — do not mutate arguments; return new objects. -- **Keep files small** — prefer several focused modules over one large file. +- **Resource pattern** — each API area lives in its own sub-package with a `resource.py` (sync) and `async_resource.py`. The async module imports the body builders / path constants from the sync one so the wire format is written once. +- **Client holds HTTP** — `_request()` (retry + error mapping + header metadata on inference paths) and `_open_stream()` (SSE, never retried) are the only HTTP entry points. Resources receive the typed client (`if TYPE_CHECKING: from ..client import FerroClient`) and call `self._client._request(...)`. +- **Exception hierarchy** — all HTTP errors raise typed exceptions inheriting from `FerroAPIError`. Connection/timeout errors and `408/429/5xx` retry with jittered backoff (honouring `Retry-After`), then raise `FerroConnectionError` / the mapped `FerroAPIError`. +- **Immutability by default** — do not mutate arguments; return new objects (`_with_response_metadata` returns a new dict, streaming uses `dataclasses.replace`). +- **Keep files small** — prefer several focused modules over one large file (`types_responses.py` exists for that reason). ### Public API - All public exports must be listed in `ferrolabsai/__init__.py` and the `__all__` list. -- The SDK mirrors the OpenAI SDK surface: `client.chat.completions.create()`, `client.embeddings.create()`, `client.images.generate()`, `client.models.list()`. -- Ferro-specific extras: `template_id`, `template_variables`, `route_tag` on completions; `client.admin.*` for gateway management. +- The SDK mirrors the OpenAI SDK surface: `client.chat.completions.create()`, `client.embeddings.create()`, `client.images.generate()`, `client.models.list()`, `client.responses.create()`. +- Gateway-specific surface: `client.health()/ready()/live()/capabilities()/rerank()`, `client.moderations`, `client.admin.*`. +- **Only model what the gateway really does.** Response metadata comes from `X-Request-ID`, `X-Gateway-Provider`, `X-Gateway-Overhead-Ms`, and the body fields `provider` / `provider_metadata` / `reasoning_content` / usage counters. There is no cost, cache-hit, or latency field for callers, and no `route_tag` / `template_*` request field — do not reintroduce them. `docs/architecture.md` § "Gateway Contract" is the reference; the contract suite enforces it. ### Environment Variables - `FERRO_API_KEY` — primary API key (takes precedence). @@ -150,10 +153,11 @@ Always run `make format lint test` before committing. ## Testing -- Tests live in `tests/test_sdk.py`. +- Unit tests live in `tests/test_*.py`; shared fixtures and canned gateway payloads in `tests/conftest.py`. - All HTTP is mocked using `pytest-httpx` — **no real gateway or network access is needed**. +- `tests/contract/` runs against a real gateway and is skipped unless `FERRO_CONTRACT_BASE_URL` is set; `scripts/with-gateway.sh` (or `make contract`) sets everything up. Every README observability claim is asserted there. - Async tests use `pytest-asyncio` with `asyncio_mode = "auto"`. -- Every bug fix needs a regression test. Every new feature needs unit tests. +- Every bug fix needs a regression test. Every new feature needs unit tests, and a contract assertion if it touches a gateway route. - Target 80%+ coverage on new code. - Run: `make test` or `pytest tests/ -v --tb=short`. @@ -162,9 +166,9 @@ Always run `make format lint test` before committing. ## CI / CD - **CI workflow:** `.github/workflows/ci.yml` -- Tests run on Python 3.9, 3.10, 3.11, 3.12 on `ubuntu-latest`. -- Lint and type check run as part of CI. -- **Publishing:** Triggered by semver tags (`v*.*.*`). Uses PyPI trusted publishing (OIDC). Asserts the tag matches `pyproject.toml` version. +- Unit tests, ruff, and mypy run on Python 3.9, 3.10, 3.11, 3.12, 3.13 on `ubuntu-latest`. +- **Contract job** checks out `ferro-labs/ai-gateway` at `v1.4.5` (required) and `main` (`continue-on-error`) and runs `scripts/with-gateway.sh`. Raise the pin when the SDK starts depending on newer gateway behaviour and update the README compatibility line. +- **Publishing:** Triggered by semver tags (`v*.*.*`); `needs` the unit matrix and the contract job. Uses PyPI trusted publishing (OIDC). Asserts the tag matches `pyproject.toml` version. - PRs target `development` branch; releases are cut from `main`. --- @@ -182,14 +186,14 @@ Always run `make format lint test` before committing. ## Adding a New API Resource 1. Create a new sub-package under `ferrolabsai/` (e.g., `ferrolabsai/newresource/`). -2. Add `resource.py` with a class that takes `client: FerroClient` in `__init__`. -3. Use `self._client._request(method, path, ...)` for HTTP calls. +2. Add `resource.py` with a class that takes `client: FerroClient` in `__init__` (imported under `TYPE_CHECKING`), a `PATH` constant, and a module-level body builder. +3. Use `self._client._request(method, path, ...)` for HTTP calls. If the route returns inference metadata, add its prefix to `_INFERENCE_PREFIXES` in `client.py`. 4. Return typed dataclass models — add them to `types.py` with `from_dict()`. 5. Wire the resource into `FerroClient.__init__` in `client.py`. 6. Export new types from `ferrolabsai/__init__.py` and add to `__all__`. -7. Add async variant in `async_resource.py` if needed, wire into `AsyncFerroClient`. -8. Write tests in `tests/test_sdk.py` with `pytest-httpx` mocks. -9. Update `CHANGELOG.md` under `Unreleased`. +7. Add the async variant in `async_resource.py` (reusing the sync builder), wire into `AsyncFerroClient`. +8. Write unit tests in `tests/test_.py` with `pytest-httpx` mocks, and a contract test in `tests/contract/test_contract.py` (extend `stub_upstream.py` if the route needs an upstream). +9. Update `CHANGELOG.md` under `Unreleased` and the route tables in `README.md` / `docs/architecture.md`. --- @@ -198,7 +202,9 @@ Always run `make format lint test` before committing. - **Do not add runtime dependencies** beyond `httpx` without discussion. The SDK is intentionally lightweight. - **Do not import pydantic** — response models are plain dataclasses. - **Do not hardcode secrets** in code, tests, or fixtures. -- **Ferro-specific response fields** (`trace_id`, `provider`, `latency_ms`, `cost_usd`) come from custom headers/body fields prefixed with `x_ferro_`. Handle graceful fallback when they're absent. +- **Do not call `GET /v1/models/{id}`** — it is not a native gateway route; it falls through to the `/v1/*` pass-through with the operator's provider credential. `models.retrieve()` is a client-side lookup for that reason. +- **Header metadata is for inference bodies only.** `_with_response_metadata` must never touch `/v1/models`, probe, or `/admin/*` bodies. +- **Streaming is never retried** and must keep the `httpx.Response` reachable (`Stream.response`). - **`from __future__ import annotations`** must be at the top of every module for 3.9 compatibility with `X | None` syntax. --- @@ -209,7 +215,7 @@ In-depth design docs live in `docs/`: | Document | Covers | | ----------------------------------------- | -------------------------------------------------------------------------------- | -| [`docs/architecture.md`](docs/architecture.md) | Module map, resource pattern, request lifecycle, streaming, admin API surface, error handling, Ferro-specific extensions | +| [`docs/architecture.md`](docs/architecture.md) | Module map, resource pattern, retry policy, streaming, the gateway contract (headers, body fields, catalog), admin route table, error handling | Read `architecture.md` before making structural changes to the SDK. @@ -218,4 +224,4 @@ Read `architecture.md` before making structural changes to the SDK. ## Related Repositories - [ferro-labs/ai-gateway](https://github.com/ferro-labs/ai-gateway) — The backend gateway (Go). The SDK talks to its HTTP API. -- Admin API surface is defined in `internal/admin/handlers.go` in the gateway repo. +- Admin API routes are defined in the `internal/admin/handlers` package (`server.go`, `Handlers.Routes`) in the gateway repo; the public routes in `internal/httpserver/router.go`. diff --git a/README.md b/README.md index e90cf39..5e23c89 100644 --- a/README.md +++ b/README.md @@ -13,13 +13,13 @@

-Route LLM requests across **29 providers and 2,500+ models** through a single OpenAI-compatible API. +Route LLM requests across **30 providers and 2,500+ models** through a single OpenAI-compatible API. Zero code changes to migrate from `openai`. Built on [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway). ```python from ferrolabsai import FerroClient -client = FerroClient(api_key="sk-ferro-...") +client = FerroClient(api_key="fgw_...") # Route to OpenAI response = client.chat.completions.create( @@ -34,19 +34,21 @@ response = client.chat.completions.create( ) print(response.content) -print(f"Handled by: {response.provider} in {response.latency_ms}ms") +print(f"Handled by {response.provider}, trace {response.trace_id}") ``` +**Compatibility:** `ferrolabsai 0.3.x` ↔ `ai-gateway ≥ v1.4.0`. Every claim in this README is executed against a real `ai-gateway v1.4.5` by the [contract suite](#contract-tests) on each CI run. + --- ## Why ferrolabsai -- **One API for 29 providers.** OpenAI, Anthropic, Google, Groq, Together, Mistral, Cohere, Bedrock, Vertex, Azure, and more — all via a single client. +- **One API for 30 providers.** OpenAI, Anthropic, Google, Groq, Together, Mistral, Cohere, Bedrock, Vertex, Azure, and more — all via a single client. - **Drop-in OpenAI replacement.** The surface matches the OpenAI SDK. Change two lines and keep all your existing code. -- **Smart routing built in.** Fallback chains, weighted load balancing, and per-request overrides via `route_tag`. -- **Cost and provider visibility.** Every response includes `provider`, `cost_usd`, `latency_ms`, and `trace_id` — no extra calls. +- **Smart routing built in.** Fallback chains, weighted load balancing, conditional and cost-optimized routing — configured on the gateway, invisible to callers. +- **Provider and trace visibility.** Every inference response carries `provider` and `trace_id` (the gateway's `X-Request-ID`) — no extra calls. - **Self-hostable.** Point `base_url` at any [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) instance and go. -- **Typed and async-first.** Dataclass response models, full `AsyncFerroClient`, streaming in both modes. +- **Typed and async-first.** Dataclass response models, full `AsyncFerroClient`, streaming in both modes, zero dependencies beyond `httpx`. --- @@ -55,7 +57,7 @@ print(f"Handled by: {response.provider} in {response.latency_ms}ms") - [Installation](#installation) - [Quickstart](#quickstart) - [Migrate from OpenAI](#migrate-from-openai) -- [Framework integrations](#framework-integrations) +- [Framework adapters](#framework-adapters) - [Usage](#usage) - [Chat completions](#chat-completions) - [Streaming](#streaming) @@ -63,7 +65,8 @@ print(f"Handled by: {response.provider} in {response.latency_ms}ms") - [Embeddings](#embeddings) - [Image generation](#image-generation) - [Model catalog](#model-catalog) - - [Ferro extras: templates & route tags](#ferro-extras-templates--route-tags) + - [Responses API, rerank, moderations](#responses-api-rerank-moderations) + - [Gateway probes and capabilities](#gateway-probes-and-capabilities) - [Observability](#observability) - [Configuration](#configuration) - [Error handling](#error-handling) @@ -85,13 +88,13 @@ Requires **Python 3.9+**. The only runtime dependency is [`httpx`](https://www.p ## Quickstart -You'll need a running [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) instance and an API key issued by it. +You'll need a running [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) instance and an API key issued by it (`fgw_...`, or the gateway's `MASTER_KEY`). ```python from ferrolabsai import FerroClient client = FerroClient( - api_key="sk-ferro-your-key", + api_key="fgw_your-key", base_url="http://localhost:8080", # your gateway address ) ``` @@ -99,7 +102,7 @@ client = FerroClient( ### Environment variables ```bash -export FERRO_API_KEY="sk-ferro-your-key" +export FERRO_API_KEY="fgw_your-key" export FERRO_BASE_URL="http://localhost:8080" ``` @@ -120,52 +123,30 @@ client = OpenAI(api_key="sk-openai-...") # After — all your existing code works unchanged from ferrolabsai import FerroClient -client = FerroClient(api_key="sk-ferro-...") +client = FerroClient(api_key="fgw_...") ``` Every `client.chat.completions.create(...)` call, every streaming loop, every tool call — identical API surface. Ferro routes to the right provider based on the model name. --- -## Framework integrations - -Ferro's gateway exposes an OpenAI-compatible HTTP API at `/v1/*`, so anything that speaks OpenAI works. Point the base URL at your gateway and keep your existing framework. - -### LangChain +## Framework adapters -```python -from langchain_openai import ChatOpenAI - -llm = ChatOpenAI( - api_key="sk-ferro-your-key", - base_url="http://localhost:8080/v1", - model="gpt-4o", -) -response = llm.invoke("Hello from LangChain via Ferro") -``` +The gateway exposes an OpenAI-compatible HTTP API at `/v1/*`, so `langchain_openai`, `llama_index.llms.openai`, and the Vercel AI SDK all work by pointing their base URL at your gateway. The first-party adapters go further and surface the gateway's `trace_id` / `provider` and typed errors: -### LlamaIndex +| Package | What it wraps | Status | +|---|---|---| +| [`langchain-ferrolabsai`](integrations/langchain-ferrolabsai/) | `FerroChatModel` (sync/async, streaming, tools, `with_structured_output`), `FerroEmbeddings`, `FerroLLM` | **0.2.0** — on ferrolabsai 0.3 | +| [`llama-index-llms-ferrolabsai`](integrations/llama-index-llms-ferrolabsai/) | LlamaIndex `LLM` | placeholder (0.0.1) | ```python -from llama_index.llms.openai import OpenAI +from langchain_ferrolabsai import FerroChatModel -llm = OpenAI( - api_key="sk-ferro-your-key", - api_base="http://localhost:8080/v1", - model="gpt-4o", -) +llm = FerroChatModel(model="gpt-4o", base_url="http://localhost:8080", api_key="fgw_...") +print(llm.invoke("Hello").response_metadata["trace_id"]) ``` -### Vercel AI SDK (Next.js) - -```typescript -import { createOpenAI } from '@ai-sdk/openai'; - -const ferro = createOpenAI({ - apiKey: process.env.FERRO_API_KEY, - baseURL: 'http://localhost:8080/v1', -}); -``` +See [`integrations/README.md`](integrations/README.md) for layout and publishing. --- @@ -181,24 +162,37 @@ response = client.chat.completions.create( {"role": "user", "content": "Explain LLM routing in one paragraph."}, ], temperature=0.7, - max_tokens=256, + max_completion_tokens=256, # supersedes max_tokens; both accepted + response_format={"type": "json_object"}, + seed=42, ) -print(response.content) # shortcut for choices[0].message.content -print(f"Cost: ${response.usage.cost_usd:.6f}") -print(f"Provider: {response.provider}") # which backend handled it +print(response.content) # shortcut for choices[0].message.content +print(response.provider) # which backend handled it +print(response.usage.total_tokens) +print(response.provider_metadata) # provider-specific extras, when present ``` +`tools`, `tool_choice`, `parallel_tool_calls`, `stop`, `top_p`, `frequency_penalty`, `presence_penalty`, `user` are first-class; any other OpenAI parameter passes through as `**kwargs`. + ### Streaming ```python -for chunk in client.chat.completions.create( +stream = client.chat.completions.create( model="claude-3-5-sonnet-20241022", messages=[{"role": "user", "content": "Write a haiku about Go performance."}], stream=True, -): - print(chunk.choices[0].delta.content or "", end="", flush=True) + stream_options={"include_usage": True}, # terminal chunk carries usage +) +print(stream.trace_id) # available before the first chunk +for chunk in stream: + if chunk.choices: + print(chunk.choices[0].delta.content or "", end="", flush=True) + if chunk.usage: # last chunk only + print(f"\n{chunk.usage.total_tokens} tokens") ``` +The return value is a `Stream` (an iterator that also exposes `trace_id`, `provider`, and the underlying `response`; use `with` or `close()` to release the connection early). Every chunk carries `trace_id` too. Note that ai-gateway forwards the terminal usage chunk unless you send `stream_options={"include_usage": False}`. A mid-stream gateway error frame raises `FerroStreamError` with `.code` (`stream_error`, `stream_timeout`). + ### Async ```python @@ -206,7 +200,7 @@ import asyncio from ferrolabsai import AsyncFerroClient async def main(): - async with AsyncFerroClient(api_key="sk-ferro-...") as client: + async with AsyncFerroClient(api_key="fgw_...") as client: response = await client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], @@ -220,13 +214,15 @@ Async streaming: ```python async def stream_example(): - async with AsyncFerroClient(api_key="sk-ferro-...") as client: - async for chunk in await client.chat.completions.create( + async with AsyncFerroClient(api_key="fgw_...") as client: + stream = await client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Count to 5"}], stream=True, - ): - print(chunk.choices[0].delta.content or "", end="", flush=True) + ) + async for chunk in stream: + if chunk.choices: + print(chunk.choices[0].delta.content or "", end="", flush=True) ``` ### Embeddings @@ -234,7 +230,7 @@ async def stream_example(): ```python response = client.embeddings.create( model="text-embedding-3-small", - input=["Ferro routes LLM requests", "across 29 providers"], + input=["Ferro routes LLM requests", "across 30 providers"], ) vectors = [d.embedding for d in response.data] print(f"Embedding dimensions: {len(vectors[0])}") @@ -254,80 +250,62 @@ print(response.data[0].url) ### Model catalog +`GET /v1/models` returns the gateway's enriched catalog (`ModelInfo`: `id`, `owned_by`, `mode`, `context_window`, `max_output_tokens`, `capabilities`, `status`, `deprecated`). The gateway ignores query parameters and has no `/v1/models/{id}` route, so filtering and lookup are done client-side over one fetch. + ```python -# Browse all 2,500+ models models = client.models.list() +anthropic_models = client.models.list(provider="anthropic") # matches owned_by +vision_models = client.models.list(capability="vision") # matches capabilities[] +claude = client.models.search("claude") # substring on id -# Filter by provider -anthropic_models = client.models.list(provider="anthropic") - -# Filter by capability -vision_models = client.models.list(capability="vision") - -# Pricing for a specific model -info = client.models.retrieve("gpt-4o") -print(f"Context window: {info.context_window:,} tokens") -print(f"Input: ${info.input_cost_per_token * 1_000_000:.2f}/M tokens") -print(f"Output: ${info.output_cost_per_token * 1_000_000:.2f}/M tokens") +info = client.models.retrieve("gpt-4o") # raises FerroNotFoundError locally if unknown +print(f"{info.provider}: {info.context_window:,} tokens, {info.capabilities}") ``` -### Forwarded Ferro fields: templates & route tags - -The SDK passes two Ferro-specific fields on `chat.completions.create(...)`: - -**`template_id` + `template_variables`** — forwarded in the chat completion body for gateway deployments that support server-side prompt templates: +### Responses API, rerank, moderations ```python -response = client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "I can't log in"}], - template_id="support-agent", - template_variables={ - "product": "Acme SaaS", - "plan": "Pro", - "date": "2026-04-09", - }, -) +# OpenAI-style Responses API (model-routed; governed and priced like chat) +r = client.responses.create(model="gpt-4o", input="Summarise the gateway in one line") +print(r.status, r.output, r.trace_id) +# retrieve/delete pin to the gateway's `responses_target`; 501 unless configured +client.responses.retrieve(r.id) + +# Cohere-shape rerank and OpenAI-shape moderations return the provider JSON +client.rerank(model="rerank-v3.5", query="gateway", documents=["a", "b"], top_n=1) +client.moderations.create(input="some text") ``` -**`route_tag`** — forwarded as `x_route_tag` in the chat completion body for gateway deployments that support per-request route tags: +### Gateway probes and capabilities ```python -response = client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - route_tag="low-cost", # e.g. forces fallback to cheaper providers -) +client.live() # {"status": "ok"} +client.ready() # {"status": "ready", "providers": [...], "targets": [...]} (503 body returned, not raised) +client.health() # {"status", "version", "commit", "built", "providers"} +client.capabilities() # per-provider parameter support: forward | translate | unsupported ``` -These fields are pass-through SDK fields. Confirm your gateway version supports them before relying on them for routing or template rendering. - --- ## Observability -Every `ChatCompletion` includes fields that tell you what the gateway actually did — no extra API calls, no log scraping: +Every inference response (`chat`, `embeddings`, `images`, `responses`, `rerank`, `moderations`) gets the gateway's response headers merged in. This is exactly what ai-gateway v1.4.x provides — nothing else is invented: -| Field | Type | Source | -|---|---|---| -| `response.provider` | `str` | Which upstream provider served the request (e.g. `"openai"`, `"anthropic"`) | -| `response.trace_id` | `str` | Correlates this request with gateway logs | -| `response.latency_ms` | `int` | End-to-end gateway latency | -| `response.usage.cost_usd` | `float` | Computed cost in USD | -| `response.usage.cache_hit` | `bool` | Whether the response came from the gateway's semantic cache | -| `response.usage.prompt_tokens` / `completion_tokens` / `total_tokens` | `int` | Standard OpenAI token counts | +| Field | Type | Source | Populated on | +|---|---|---|---| +| `response.trace_id` | `str` | `X-Request-ID` header (32 hex chars; equals the OTel trace id) | every response, `Stream.trace_id`, every chunk, every `FerroAPIError.request_id` | +| `response.provider` | `str` | body `provider` on chat completions; `X-Gateway-Provider` header on responses/pass-through | non-streaming chat, responses (not on SSE streams as of v1.4.5) | +| `response.gateway_overhead_ms` | `float` | `X-Gateway-Overhead-Ms` header — the gateway's own processing time, **not** end-to-end latency | non-streaming chat completions | +| `response.provider_metadata` | `dict` | body `provider_metadata` | when the provider returns extras | +| `response.usage.prompt_tokens` / `completion_tokens` / `total_tokens` | `int` | body `usage` | chat, embeddings, terminal streaming chunk | +| `response.usage.reasoning_tokens` / `cache_read_tokens` / `cache_write_tokens` | `int \| None` | body `usage` (omitted when zero) | when the provider reports them | ```python -response = client.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hello"}], -) - -print(f"trace={response.trace_id} provider={response.provider} " - f"latency={response.latency_ms}ms cost=${response.usage.cost_usd:.6f}") +response = client.chat.completions.create(model="gpt-4o", messages=[{"role": "user", "content": "Hello"}]) +print(f"trace={response.trace_id} provider={response.provider} overhead={response.gateway_overhead_ms}ms") ``` -To dig deeper into a specific request, use `client.admin.logs.list(trace_id=...)` — see [Admin API](#admin-api-oss-gateway). +Cost and cache hits are not exposed to callers — they live in the gateway's request log (`client.admin.logs.list(model=...)` joins on `trace_id`), Prometheus, and OTel spans. --- @@ -337,16 +315,16 @@ To dig deeper into a specific request, use `client.admin.logs.list(trace_id=...) ```python client = FerroClient( - api_key="sk-ferro-...", # or FERRO_API_KEY env var + api_key="fgw_...", # or FERRO_API_KEY env var base_url="http://localhost:8080", # or FERRO_BASE_URL env var timeout=120.0, # seconds (default: 120.0) - max_retries=2, # retries on connection errors (default: 2) + max_retries=2, # default: 2 default_headers={"x-env": "prod"}, # merged into every request http_client=my_httpx_client, # bring your own httpx.Client ) ``` -**Retries** are triggered only by `httpx.ConnectError` and `httpx.TimeoutException` — HTTP errors (4xx/5xx) propagate immediately as typed exceptions so you can handle them yourself. +**Retries** cover connection errors, timeouts, and HTTP `408` / `429` / `5xx` — capped exponential backoff with full jitter (0.5 s base, 8 s cap), honouring `Retry-After` when the gateway sends one (capped at 30 s, the same cap the gateway applies upstream). Other `4xx` responses and **streaming requests are never retried**. **Bring-your-own httpx client** lets you configure proxies, custom TLS, connection pool limits, or instrumentation middleware and reuse that across the SDK: @@ -354,20 +332,14 @@ client = FerroClient( import httpx pooled = httpx.Client(limits=httpx.Limits(max_connections=50)) -client = FerroClient(api_key="sk-ferro-...", http_client=pooled) +client = FerroClient(api_key="fgw_...", http_client=pooled) ``` Close the client explicitly when you're done (or use a `with` block): ```python -with FerroClient(api_key="sk-ferro-...") as client: +with FerroClient(api_key="fgw_...") as client: ... -# or -client = FerroClient(api_key="sk-ferro-...") -try: - ... -finally: - client.close() ``` --- @@ -378,6 +350,8 @@ finally: from ferrolabsai import ( FerroClient, FerroAuthError, + FerroBudgetExceededError, + FerroPermissionError, FerroRateLimitError, FerroNotFoundError, FerroServerError, @@ -389,27 +363,31 @@ try: model="gpt-4o", messages=[{"role": "user", "content": "Hello"}], ) -except FerroAuthError: +except FerroAuthError: # 401 print("Invalid API key — check FERRO_API_KEY") -except FerroRateLimitError: - print("Rate limit hit — back off and retry") -except FerroNotFoundError: - print("Model or endpoint not found") -except FerroServerError as e: - print(f"Gateway error {e.status_code} — upstream provider may be down") +except FerroBudgetExceededError: # 402 insufficient_quota + print("Spend limit reached for this key") +except FerroPermissionError: # 403 insufficient_scope + print("This key lacks the scope for that route") +except FerroRateLimitError as e: # 429 (already retried) + print(f"Rate limited — retry after {e.retry_after}s") +except FerroNotFoundError as e: # 404 model_not_found / not_found + print(f"Not found: {e.code}") +except FerroServerError as e: # 5xx (already retried) + print(f"Gateway error {e.status_code} ({e.code}) — trace {e.request_id}") except FerroConnectionError: print("Cannot reach gateway — is it running?") ``` -All HTTP-level exceptions inherit from `FerroAPIError` and expose `.status_code`, `.code`, `.message`, and `.request_id`. `FerroConnectionError` and `FerroStreamError` inherit from `FerroError` directly. +All HTTP-level exceptions inherit from `FerroAPIError` and expose `.status_code`, `.code` (the gateway's error code, e.g. `model_not_found`, `insufficient_scope`), `.message`, and `.request_id`. `FerroConnectionError` and `FerroStreamError` inherit from `FerroError` directly. --- ## Admin API (OSS gateway) -These APIs are available on any self-hosted Ferro Labs AI Gateway instance. Requires an admin-scoped API key. +These APIs are available on any self-hosted Ferro Labs AI Gateway instance. Reads need a `read_only` or `admin` key; writes need `admin` (a `read_only` key gets `FerroPermissionError`). -The admin namespace mirrors the OSS gateway's `/admin/*` HTTP surface defined in [`internal/admin/handlers.go`](https://github.com/ferro-labs/ai-gateway/blob/main/internal/admin/handlers.go). +The admin namespace mirrors the OSS gateway's `/admin/*` HTTP surface defined in the [`internal/admin/handlers`](https://github.com/ferro-labs/ai-gateway/tree/main/internal/admin/handlers) package. ### API keys @@ -417,12 +395,16 @@ The admin namespace mirrors the OSS gateway's `/admin/*` HTTP surface defined in # Create new_key = client.admin.keys.create( name="backend-service", - scopes=["admin"], + scopes=["admin"], # or ["read_only"] ) print(new_key.key) # full key value — shown ONCE, store it securely -# List +# List / retrieve (key values are masked: fgw_ab12...cd34) keys = client.admin.keys.list() +key = client.admin.keys.retrieve("key_id") + +# Update metadata +client.admin.keys.update("key_id", name="renamed", active=False) # Per-key usage counts (sorted by usage by default) usage = client.admin.keys.usage(limit=20) @@ -433,7 +415,7 @@ client.admin.keys.revoke("key_id") # Rotate — atomically invalidates old, returns new rotated = client.admin.keys.rotate("key_id") -# Permanently delete the record +# Permanently delete the record (the gateway refuses to delete the last admin key) client.admin.keys.delete("key_id") ``` @@ -442,54 +424,54 @@ client.admin.keys.delete("key_id") The OSS gateway has a single *active* routing config. Use `history()` to inspect prior versions and `rollback(version)` to revert. Updates are zero-downtime hot reloads. ```python -# Read the current config cfg = client.admin.config.get() print(cfg.strategy) # e.g. {"mode": "fallback"} print(cfg.targets) # list of {virtual_key, weight, ...} -# Replace it (PUT) — hot reload, no restart client.admin.config.update({ "strategy": {"mode": "fallback"}, "targets": [ {"virtual_key": "openai", "weight": 1}, {"virtual_key": "anthropic", "weight": 1}, - {"virtual_key": "groq", "weight": 1}, - ], - "plugins": [ - {"name": "cache", "enabled": True}, - {"name": "logger", "enabled": True}, ], }) -# Inspect history and roll back history = client.admin.config.history() client.admin.config.rollback(history[-2].version) ``` +Note: `get()` masks secrets and redacts free-form map keys, so its body does not round-trip unchanged into `update()`; unknown keys are rejected with `400`. + ### Request logs -The gateway logs every request (when the `logger` plugin is enabled). Query, aggregate, and prune via `client.admin.logs`. +The gateway records every request when a request-log store is configured (`REQUEST_LOG_STORE_BACKEND=sqlite|postgres`); the endpoints answer `501` without one. ```python -# Recent failures -errors = client.admin.logs.list(limit=20, stage="on_error") -for entry in errors["data"]: - print(entry["trace_id"], entry["model"], entry["provider"]) +# Recent entries for a model (one row per request; stage="all" shows every lifecycle stage) +entries = client.admin.logs.list(limit=20, model="gpt-4o") +for entry in entries["data"]: + print(entry["trace_id"], entry["provider"], entry["duration_ms"], entry["cost_usd"]) + +# Filter by the calling key +client.admin.logs.list(api_key_id="key_id") -# Aggregate stats -stats = client.admin.logs.stats() +# Aggregate stats with a 24-point time series +stats = client.admin.logs.stats(buckets=24) # Prune old entries client.admin.logs.delete(before="2026-01-01T00:00:00Z") ``` -### Providers, plugins, dashboard +### Providers, plugins, audit, dashboard ```python -providers = client.admin.providers.list() # registered LLM providers -plugins = client.admin.plugins.list() # installed gateway plugins -dashboard = client.admin.dashboard() # high-level counts -health = client.admin.health() # gateway health check +providers = client.admin.providers.list() # registered providers and their models +catalog = client.admin.providers.catalog() # every provider the build knows: {id, registered, catalog_models} +plugins = client.admin.plugins.list() # configured plugins +available = client.admin.plugins.catalog() # built-in plugins available to configure +audit = client.admin.audit.list(action="key.create", limit=50) # admin audit trail +dashboard = client.admin.dashboard() # high-level counts +health = client.admin.health() # gateway health check (admin view) ``` --- @@ -504,12 +486,16 @@ make test # pytest (all HTTP is mocked — no gateway needed) make lint # ruff + mypy make format # ruff format make build # build sdist + wheel into dist/ -make clean # remove artifacts +make contract # boot a real gateway from ../ai-gateway and run tests/contract ``` -All 30 tests run in under a second against `pytest-httpx` fixtures, so no network or running gateway is required. +The 113 unit tests run in a few seconds against `pytest-httpx` fixtures, so no network or running gateway is required. + +### Contract tests + +`tests/contract/` is skipped unless `FERRO_CONTRACT_BASE_URL` is set. `scripts/with-gateway.sh` builds `ferrogw` from an [ai-gateway](https://github.com/ferro-labs/ai-gateway) checkout (`FERRO_GATEWAY_SOURCE`, default `../ai-gateway`), points it at a stdlib stub upstream (`tests/contract/stub_upstream.py`), and runs the 23 contract tests: probes, catalog, chat, streaming, embeddings, responses, the error envelope (401/403/404/501), and every admin route the SDK wraps. CI runs it against the pinned `v1.4.5` (required) and `main` (advisory). -See [CHANGELOG.md](CHANGELOG.md) for release history. +See [CHANGELOG.md](CHANGELOG.md) for release history and [docs/architecture.md](docs/architecture.md) for the design. --- diff --git a/SECURITY.md b/SECURITY.md index bd4fada..c2e4b1d 100644 --- a/SECURITY.md +++ b/SECURITY.md @@ -8,8 +8,8 @@ We currently ship security fixes for the latest minor release on PyPI. Older ver | Version | Supported | | ------- | --------- | -| 0.1.x | Yes | -| < 0.1 | No | +| 0.3.x | Yes | +| < 0.3 | No | ## Reporting a Vulnerability @@ -59,7 +59,8 @@ If you are integrating `ferrolabsai` into your own application: - **Never hardcode API keys.** Load them from environment variables or a secret manager. - **Use HTTPS** for any non-localhost `base_url`. - **Rotate keys** periodically and immediately if you suspect exposure. -- **Validate `trace_id` and cost fields** on responses if you store them — treat them as untrusted data at your application boundary. +- **Validate `trace_id` and `provider`** on responses if you store them — treat them as untrusted data at your application boundary. +- **Scope your keys.** Use `read_only` keys wherever an integration only reads `/admin/*`; the SDK surfaces scope violations as `FerroPermissionError`. - **Pin the SDK version** in production and review the changelog before upgrading. ## Questions diff --git a/docs/architecture.md b/docs/architecture.md index 9136ef3..a778d2b 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -1,18 +1,18 @@ # Architecture — ferrolabsai Python SDK -This document describes the internal architecture of the `ferrolabsai` SDK, how the pieces fit together, and the design decisions behind them. +This document describes the internal architecture of the `ferrolabsai` SDK, how the pieces fit together, and the design decisions behind them. It is written against **ai-gateway v1.4.5**; the contract suite (`tests/contract/`) keeps it honest. --- ## High-Level Overview -The SDK acts as a thin HTTP client that sits between application code and a running [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) instance. It provides an OpenAI-compatible surface so users can switch from `openai.OpenAI` to `ferrolabsai.FerroClient` with minimal code changes, while gaining access to 29+ LLM providers, smart routing, and gateway management APIs. +The SDK acts as a thin HTTP client that sits between application code and a running [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) instance. It provides an OpenAI-compatible surface so users can switch from `openai.OpenAI` to `ferrolabsai.FerroClient` with minimal code changes, while gaining access to 30 LLM providers, smart routing, and gateway management APIs. ``` ┌─────────────────────────────────────────────────────────────┐ │ Application Code │ │ │ -│ client.chat.completions.create(model="gpt-4o", ...) │ +│ client.chat.completions.create(model="gpt-4o", ...) │ │ client.embeddings.create(model="text-embedding-3-small") │ │ client.admin.config.update({...}) │ └──────────────────────────┬──────────────────────────────────┘ @@ -25,26 +25,30 @@ The SDK acts as a thin HTTP client that sits between application code and a runn │ ├── embeddings │ │ ├── images │ │ ├── models │ + │ ├── responses │ + │ ├── moderations │ + │ ├── rerank() / probes │ │ └── admin │ │ ├── keys │ │ ├── config │ │ ├── logs │ │ ├── providers │ - │ └── plugins │ + │ ├── plugins │ + │ └── audit │ └────────────┬────────────────┘ │ HTTP (httpx) ┌────────────▼────────────────┐ │ Ferro Labs AI Gateway │ - │ /v1/* /admin/* │ + │ /v1/* /admin/* /health │ │ │ │ Routing · Fallback · Cache │ - │ Rate limiting · Logging │ + │ Budgets · Logging · OTel │ └────────────┬────────────────┘ │ ┌───────────────┼───────────────┐ │ │ │ ┌────▼────┐ ┌────▼─────┐ ┌────▼────┐ - │ OpenAI │ │Anthropic │ │ Groq │ ... 29+ providers + │ OpenAI │ │Anthropic │ │ Groq │ ... 30 providers └─────────┘ └──────────┘ └─────────┘ ``` @@ -55,25 +59,24 @@ The SDK acts as a thin HTTP client that sits between application code and a runn ``` ferrolabsai/ ├── __init__.py # Public API surface & __all__ -├── client.py # FerroClient, AsyncFerroClient, _raise_api_error -├── types.py # All dataclass response models -├── exceptions/ -│ └── __init__.py # FerroError hierarchy -├── completions/ -│ ├── resource.py # Completions (sync + streaming) -│ └── async_resource.py # AsyncCompletions (async + streaming) -├── embeddings/ -│ ├── resource.py # Embeddings (sync) -│ └── async_resource.py # AsyncEmbeddings -├── images/ -│ └── resource.py # Images (sync) -├── models/ -│ └── resource.py # Models catalog (sync) -└── admin/ - └── resource.py # Admin, _KeysResource, _ConfigResource, - # _LogsResource, _ProvidersResource, _PluginsResource +├── _version.py # __version__ constant (kept in sync with pyproject.toml) +├── client.py # FerroClient, AsyncFerroClient, retry policy, _raise_api_error +├── streaming.py # Stream / AsyncStream SSE wrappers +├── types.py # Dataclass response models +├── types_responses.py # Response (Responses API), re-exported from types +├── exceptions/__init__.py # FerroError hierarchy +├── completions/ # chat.completions (resource.py + async_resource.py) +├── embeddings/ # embeddings +├── images/ # images.generate +├── models/ # model catalog (client-side lookup/filter) +├── responses/ # Responses API +├── moderations/ # moderations +└── admin/ # Admin, _KeysResource, _ConfigResource, _LogsResource, + # _ProvidersResource, _PluginsResource, _AuditResource ``` +Every resource sub-package has a sync `resource.py` and an `async_resource.py`; the async module imports the request-body builders and path constants from the sync one so the wire format is defined once. + --- ## Core Design Decisions @@ -89,83 +92,90 @@ class ChatCompletion: model: str choices: list[Choice] usage: Usage | None = None - # Ferro extras - trace_id: str | None = None - provider: str | None = None - latency_ms: int | None = None - - @classmethod - def from_dict(cls, d: dict[str, Any]) -> ChatCompletion: - ... + # Gateway extensions + trace_id: str | None = None # X-Request-ID header + provider: str | None = None # body `provider` / X-Gateway-Provider header + gateway_overhead_ms: float | None = None # X-Gateway-Overhead-Ms header + provider_metadata: dict[str, Any] | None = None ``` ### 2. Resource Pattern -Each API surface (completions, embeddings, images, models, admin) is a **resource class** that: - -1. Receives the client instance via `__init__(self, client)`. -2. Calls `self._client._request(method, path, ...)` for all HTTP traffic. -3. Returns typed dataclass models. - -This mirrors the OpenAI SDK's structure (`client.chat.completions`, `client.embeddings`, etc.) and keeps individual files focused and small. +Each API surface is a **resource class** that receives the client via `__init__(self, client: FerroClient)` (typed under `TYPE_CHECKING` to avoid an import cycle), calls `self._client._request(method, path, ...)` for HTTP, and returns typed dataclasses (or the raw dict where the gateway's shape is loosely structured — admin lists, rerank, moderations). ``` FerroClient - ├── _http: httpx.Client # connection pool + auth headers - ├── _request(method, path, ...) # central HTTP with retry logic + ├── _http: httpx.Client # connection pool + auth headers + ├── _request(method, path, ...) # central HTTP with retry + error mapping + ├── _open_stream(path, body) # SSE: no retry, returns the live response + ├── health() / ready() / live() # → /health, /readyz, /livez + ├── capabilities() # → /v1/capabilities + ├── rerank(...) # → /v1/rerank │ - ├── chat: _ChatNamespace - │ └── completions: Completions # → /v1/chat/completions - ├── embeddings: Embeddings # → /v1/embeddings - ├── images: Images # → /v1/images/generations - ├── models: Models # → /v1/models - └── admin: Admin # → /admin/* - ├── keys: _KeysResource # → /admin/keys - ├── config: _ConfigResource # → /admin/config - ├── logs: _LogsResource # → /admin/logs - ├── providers: _ProvidersResource # → /admin/providers - └── plugins: _PluginsResource # → /admin/plugins + ├── chat.completions: Completions # → /v1/chat/completions + ├── embeddings: Embeddings # → /v1/embeddings + ├── images: Images # → /v1/images/generations + ├── models: Models # → /v1/models (client-side filter/lookup) + ├── responses: Responses # → /v1/responses[/{id}] + ├── moderations: Moderations # → /v1/moderations + └── admin: Admin # → /admin/* + ├── keys / config / logs / providers / plugins / audit ``` ### 3. Single HTTP Entry Point -All HTTP traffic flows through `FerroClient._request()` (sync) or `AsyncFerroClient._request()` (async). This centralizes: +All non-streaming traffic flows through `_request()` (sync) or `AsyncFerroClient._request()`. It centralizes: + +- **Authentication** — `Authorization: Bearer {api_key}` header on the `httpx` client. +- **Retry policy** — see below. +- **Error mapping** — `_raise_api_error()` translates the gateway's error envelope into typed exceptions. +- **Response parsing** — JSON, 204 handling, and header metadata injection for inference paths only (`/v1/chat/completions`, `/v1/completions`, `/v1/embeddings`, `/v1/images/generations`, `/v1/responses`, `/v1/rerank`, `/v1/moderations`). Catalog, probe, and admin bodies are returned untouched. +- **Probe semantics** — `/health` and `/readyz` answer `503` with a JSON body when degraded; `_request(..., allow=(503,))` returns that body instead of raising. + +### 4. Retry Policy + +Shared by sync and async; streaming requests are never retried. -- **Authentication** — `Authorization: Bearer {api_key}` header injected by `httpx.Client`. -- **Retry logic** — retries on `httpx.ConnectError` and `httpx.TimeoutException` only (not HTTP errors). -- **Error mapping** — `_raise_api_error()` translates HTTP status codes into typed exceptions. -- **Response parsing** — JSON deserialization, 204 handling. +| Trigger | Retried? | +|---|---| +| `httpx.ConnectError`, `httpx.TimeoutException` | yes | +| HTTP `408`, `429`, `5xx` | yes | +| any other `4xx` | no — raised immediately | +| streaming (`_open_stream`) | never | -Streaming is the exception: `_stream_request()` returns an iterator of raw SSE lines, and the resource class handles SSE parsing (`data: ...` / `[DONE]`). +Delay before retry *n*: `Retry-After` seconds when present (capped at 30 s — the same cap the gateway applies to its own upstream retries), else `uniform(0, min(0.5 · 2^(n-1), 8))` (full jitter). `max_retries` defaults to 2 and is validated at construction. -### 4. Typed Exception Hierarchy +### 5. Typed Exception Hierarchy ``` FerroError (base) -├── FerroAPIError (any non-2xx HTTP response) -│ ├── FerroAuthError (401) -│ ├── FerroRateLimitError (429) -│ ├── FerroNotFoundError (404) -│ └── FerroServerError (5xx) -├── FerroConnectionError (network / timeout — retried first) -└── FerroStreamError (SSE parse failure) +├── FerroAPIError (any non-2xx HTTP response; .status_code .code .message .request_id) +│ ├── FerroAuthError (401) +│ ├── FerroBudgetExceededError (402 insufficient_quota) +│ ├── FerroPermissionError (403 insufficient_scope) +│ ├── FerroNotFoundError (404 model_not_found / not_found — also raised locally by models.retrieve) +│ ├── FerroRateLimitError (429; .retry_after from the Retry-After header) +│ └── FerroServerError (5xx) +├── FerroConnectionError (network / timeout after all retries) +└── FerroStreamError (malformed SSE frame, or a gateway error frame; .code) ``` -`FerroAPIError` carries `.status_code`, `.code`, `.message`, and `.request_id`. Connection and stream errors inherit from `FerroError` directly since they have no HTTP status. +`.request_id` is the `X-Request-ID` of the failed response (falls back to `request_id` / `trace_id` in the body). -### 5. Sync + Async Duality +### 6. Streaming -The SDK provides both `FerroClient` (synchronous, `httpx.Client`) and `AsyncFerroClient` (asynchronous, `httpx.AsyncClient`). +`chat.completions.create(stream=True)` returns a `Stream` (async: `AsyncStream`) rather than a bare generator so the HTTP response — and therefore its headers — stays reachable: -- Sync resources live in `resource.py` within each sub-package. -- Async resources live in `async_resource.py`. -- Resources that don't yet have an async variant (images, models, admin) only have `resource.py`. +- `stream.trace_id` (`X-Request-ID`) and `stream.provider` (`X-Gateway-Provider`, not set on SSE as of v1.4.5) are available before the first chunk and copied onto every `ChatCompletionChunk`. +- Frames are parsed line-wise: `data: {...}` → `ChatCompletionChunk`; `data: [DONE]` ends the stream; `{"error": {...}}` → `FerroStreamError(code=...)`; anything unparsable → `FerroStreamError("Malformed SSE chunk ...")`. +- `usage` appears on the terminal chunk. The gateway always asks the upstream for usage (for metering) and forwards that chunk unless the client sent `stream_options={"include_usage": False}`. +- The response is closed when the stream is exhausted, on error, on `close()`, or when the `with` block exits. -The async client currently supports completions and embeddings. Other resources can be added following the same pattern. +### 7. Sync + Async Duality -### 6. OpenAI Compatibility Layer +`FerroClient` (`httpx.Client`) and `AsyncFerroClient` (`httpx.AsyncClient`) expose the same namespaces; every resource has an async twin. The async client's `create(stream=True)` is awaited once (opening the stream, so HTTP errors surface there) and returns an `AsyncStream`. -The SDK intentionally mirrors OpenAI SDK ergonomics: +### 8. OpenAI Compatibility Layer | OpenAI | ferrolabsai | | --------------------------------- | ------------------------------------------ | @@ -173,6 +183,8 @@ The SDK intentionally mirrors OpenAI SDK ergonomics: | `client.chat.completions.create` | `client.chat.completions.create` | | `client.embeddings.create` | `client.embeddings.create` | | `client.images.generate` | `client.images.generate` | +| `client.models.list / retrieve` | `client.models.list / retrieve` (client-side) | +| `client.responses.create` | `client.responses.create` | | `OPENAI_API_KEY` | `FERRO_API_KEY` (falls back to `OPENAI_API_KEY`) | The `_ChatNamespace` class exists solely to provide the `client.chat.completions` accessor, matching OpenAI's nested layout. @@ -185,109 +197,90 @@ The `_ChatNamespace` class exists solely to provide the `client.chat.completions 1. User calls: client.chat.completions.create(model="gpt-4o", messages=[...]) │ 2. Completions.create() - ├── Builds request body (model, messages, temperature, Ferro extras...) - ├── Non-streaming: calls self._client._request("POST", "/v1/chat/completions", json=body) - └── Streaming: calls self._client._stream_request(path, body) → yields SSE lines + ├── build_body(): model, messages, stream + non-None optionals + **kwargs + ├── Non-streaming: self._client._request("POST", "/v1/chat/completions", json=body) + └── Streaming: Stream(self._client._open_stream(path, body)) │ 3. FerroClient._request() - ├── Builds httpx.Request with auth headers ├── Sends via self._http (httpx.Client) - ├── On success: returns parsed JSON dict - ├── On HTTP error: _raise_api_error() → typed FerroXxxError - └── On connection/timeout error: retries up to max_retries, then FerroConnectionError + ├── 2xx: parse JSON; inference path → merge X-Request-ID / X-Gateway-Provider / + │ X-Gateway-Overhead-Ms into the dict (body fields win) + ├── 408/429/5xx or connect/timeout: sleep (Retry-After or jittered backoff), retry + └── Other non-2xx or retries exhausted: _raise_api_error() → typed FerroXxxError │ -4. Completions.create() (continued) - ├── Non-streaming: ChatCompletion.from_dict(data) → typed dataclass - └── Streaming: parses "data: {json}" lines → yields ChatCompletionChunk +4. Completions.create() → ChatCompletion.from_dict(data) │ 5. User receives ChatCompletion with: - ├── .content → shortcut to first choice - ├── .provider → which backend handled it (Ferro extra) - ├── .trace_id → correlation ID for gateway logs - ├── .latency_ms → end-to-end gateway latency - └── .usage.cost_usd → computed cost in USD + ├── .content → shortcut to first choice + ├── .provider → which backend answered (body `provider`) + ├── .trace_id → X-Request-ID; joins with admin.logs / OTel + ├── .gateway_overhead_ms → gateway's own processing time + ├── .provider_metadata → provider-specific extras + └── .usage → tokens (+ reasoning / cache counters when reported) ``` --- -## Streaming Architecture +## Gateway Contract -The SDK supports server-sent events (SSE) for real-time token streaming. +### Response headers read by the SDK -### Sync Streaming -```python -for chunk in client.chat.completions.create(model="gpt-4o", messages=[...], stream=True): - print(chunk.choices[0].delta.content, end="") -``` +| Header | Set by the gateway on | SDK field | +|---|---|---| +| `X-Request-ID` | every response (32 lowercase hex, equals the OTel trace id) | `trace_id`, `Stream.trace_id`, chunk `trace_id`, `FerroAPIError.request_id` | +| `X-Gateway-Provider` | `/v1/responses`, `/v1/*` pass-through, legacy `/v1/completions` — **not** non-streaming chat (body `provider` instead), **not** SSE | `provider` (only when the body has none) | +| `X-Gateway-Overhead-Ms` | non-streaming `/v1/chat/completions`, when > 0 | `gateway_overhead_ms` | +| `Retry-After` | every gateway-originated `429` (`"1"`), upstream `429` (upstream value) | retry delay; `FerroRateLimitError.retry_after` | -Internally: -1. `Completions.create(stream=True)` calls `self._stream()`. -2. `_stream()` calls `FerroClient._stream_request()` which opens a streaming `httpx` response. -3. Lines are iterated via `response.iter_lines()`. -4. Each `data: {...}` line is parsed into a `ChatCompletionChunk` and yielded. -5. `data: [DONE]` terminates the iterator. +No `X-Ferro-*`, cost, or cache-hit header exists at any gateway version; the SDK reads none. -### Async Streaming -```python -async for chunk in await client.chat.completions.create(model="gpt-4o", messages=[...], stream=True): - print(chunk.choices[0].delta.content, end="") -``` +### Body extensions read by the SDK + +| Field | Where | SDK field | +|---|---|---| +| `provider`, `provider_metadata` | non-streaming chat completion | `ChatCompletion.provider` / `.provider_metadata` | +| `usage.reasoning_tokens`, `cache_read_tokens`, `cache_write_tokens` | chat usage (omitted when zero) | `Usage.*` | +| `message.reasoning_content`, `delta.reasoning_content` | chat message / stream delta | `ChatMessage.reasoning_content`, `StreamDelta.reasoning_content` | +| `{"error": {message, type, code}}` | every non-2xx and mid-stream error frame | exception `.message` / `.code` | + +### Model catalog -Uses `response.aiter_lines()` within an `async with self._client._http.stream(...)` context. +`GET /v1/models` returns `EnrichedModelInfo` (`ai-gateway/internal/handler/models.go`): `id`, `object`, `created`, `owned_by`, plus `mode`, `context_window`, `max_output_tokens`, `capabilities[]`, `status`, `deprecated` when the catalog knows the model. Query parameters are ignored and `GET /v1/models/{id}` is not a native route (it would fall through to the `/v1/*` pass-through with the operator's credential), so `list(provider=, capability=)`, `search()`, and `retrieve()` all work client-side over one fetch. --- ## Admin API Surface -The admin namespace exposes gateway management operations that map 1:1 to the OSS gateway's `/admin/*` HTTP routes (defined in `ai-gateway/internal/admin/handlers.go`). +The admin namespace maps 1:1 to the gateway's `/admin/*` routes (`ai-gateway/internal/admin/handlers/server.go`, `Handlers.Routes`). | SDK Method | HTTP Route | Scope | | --------------------------------------- | ----------------------------------- | ---------- | -| `admin.dashboard()` | `GET /admin/dashboard` | read-only | -| `admin.health()` | `GET /admin/health` | read-only | -| `admin.keys.list()` | `GET /admin/keys` | read-only | -| `admin.keys.retrieve(id)` | `GET /admin/keys/{id}` | read-only | -| `admin.keys.create(name=...)` | `POST /admin/keys` | admin | -| `admin.keys.update(id, ...)` | `PUT /admin/keys/{id}` | admin | -| `admin.keys.delete(id)` | `DELETE /admin/keys/{id}` | admin | -| `admin.keys.revoke(id)` | `POST /admin/keys/{id}/revoke` | admin | -| `admin.keys.rotate(id)` | `POST /admin/keys/{id}/rotate` | admin | -| `admin.keys.usage(limit=...)` | `GET /admin/keys/usage` | read-only | -| `admin.config.get()` | `GET /admin/config` | read-only | -| `admin.config.create(config)` | `POST /admin/config` | admin | -| `admin.config.update(config)` | `PUT /admin/config` | admin | -| `admin.config.delete()` | `DELETE /admin/config` | admin | -| `admin.config.history()` | `GET /admin/config/history` | read-only | -| `admin.config.rollback(version)` | `POST /admin/config/rollback/{v}` | admin | -| `admin.logs.list(limit=...)` | `GET /admin/logs` | read-only | -| `admin.logs.stats()` | `GET /admin/logs/stats` | read-only | +| `admin.dashboard()` | `GET /admin/dashboard` | read_only | +| `admin.health()` | `GET /admin/health` | read_only | +| `admin.keys.list()` | `GET /admin/keys` | read_only | +| `admin.keys.retrieve(id)` | `GET /admin/keys/{id}` | read_only | +| `admin.keys.create(name=...)` | `POST /admin/keys` | admin | +| `admin.keys.update(id, ...)` | `PUT /admin/keys/{id}` | admin | +| `admin.keys.delete(id)` | `DELETE /admin/keys/{id}` | admin | +| `admin.keys.revoke(id)` | `POST /admin/keys/{id}/revoke` | admin | +| `admin.keys.rotate(id)` | `POST /admin/keys/{id}/rotate` | admin | +| `admin.keys.usage(limit=...)` | `GET /admin/keys/usage` | read_only | +| `admin.config.get()` | `GET /admin/config` | read_only | +| `admin.config.create(config)` | `POST /admin/config` | admin | +| `admin.config.update(config)` | `PUT /admin/config` | admin | +| `admin.config.delete()` | `DELETE /admin/config` | admin | +| `admin.config.history()` | `GET /admin/config/history` | read_only | +| `admin.config.rollback(version)` | `POST /admin/config/rollback/{v}` | admin | +| `admin.logs.list(stage=, api_key_id=, ...)` | `GET /admin/logs` | read_only | +| `admin.logs.stats(buckets=)` | `GET /admin/logs/stats` | read_only | | `admin.logs.delete(before=...)` | `DELETE /admin/logs` | admin | -| `admin.providers.list()` | `GET /admin/providers` | read-only | -| `admin.plugins.list()` | `GET /admin/plugins` | read-only | +| `admin.providers.list()` | `GET /admin/providers` | read_only | +| `admin.providers.catalog()` | `GET /admin/providers/catalog` | read_only | +| `admin.plugins.list()` | `GET /admin/plugins` | read_only | +| `admin.plugins.catalog()` | `GET /admin/plugins/catalog` | read_only | +| `admin.audit.list(action=, actor_id=, outcome=, since=)` | `GET /admin/audit` | read_only | ---- - -## Ferro-Specific Extensions - -The SDK passes through fields that the standard OpenAI API doesn't know about. These are safe — any OpenAI-compatible backend that doesn't recognize them silently ignores them. - -### Request Extensions (on `chat.completions.create`) - -| Parameter | Wire Field | Purpose | -| -------------------- | --------------------- | ------------------------------------------------------------------ | -| `template_id` | `template_id` | Render a server-side prompt template (Go `text/template` syntax) | -| `template_variables` | `template_variables` | Variables injected into the template | -| `route_tag` | `x_route_tag` | Override the routing strategy for this single request | - -### Response Extensions (on `ChatCompletion`) - -| Field | Source | Purpose | -| ------------- | ------------------------------------------- | ------------------------------------ | -| `trace_id` | `x_ferro_trace_id` or `trace_id` in body | Correlates with gateway logs | -| `provider` | `x_ferro_provider` or `provider` in body | Which upstream served the request | -| `latency_ms` | `x_ferro_latency_ms` in body | End-to-end gateway latency | -| `cost_usd` | `usage.cost_usd` in body | Computed cost in USD | -| `cache_hit` | `usage.cache_hit` in body | Whether semantic cache was used | +Not wrapped on purpose: `/admin/session(s)` (dashboard-only), `/metrics`, `/debug/vars`. Gateway rules worth knowing: the last admin key record cannot be revoked or deleted (`409`), `read_only` writes get `403 insufficient_scope`, and `GET /admin/config` masks secrets so it does not round-trip into `PUT`. --- @@ -297,19 +290,22 @@ The SDK passes through fields that the standard OpenAI API doesn't know about. T HTTP response received │ ├── 2xx → parse JSON → return dict / dataclass + ├── 503 on /health, /readyz → return the JSON body (degraded is an answer) + │ + ├── 408 / 429 / 5xx → retry (Retry-After or jittered backoff) … then map as below │ ├── 401 → FerroAuthError + ├── 402 → FerroBudgetExceededError + ├── 403 → FerroPermissionError ├── 404 → FerroNotFoundError - ├── 429 → FerroRateLimitError + ├── 429 → FerroRateLimitError(retry_after=...) ├── 5xx → FerroServerError - ├── other 4xx → FerroAPIError + ├── other 4xx (400, 405, 409, 413, 501, …) → FerroAPIError(code=...) │ - ├── ConnectError → retry up to max_retries → FerroConnectionError - └── TimeoutException → retry up to max_retries → FerroConnectionError + ├── ConnectError / TimeoutException → retry up to max_retries → FerroConnectionError + └── SSE: HTTP error before the first byte → same mapping; error frame → FerroStreamError ``` -HTTP errors (4xx/5xx) are **never retried** — they propagate immediately so the caller can handle them. Only connection and timeout errors trigger the retry loop. - --- ## Configuration & Auth @@ -319,18 +315,13 @@ FerroClient( api_key="...", # or FERRO_API_KEY / OPENAI_API_KEY env var base_url="...", # or FERRO_BASE_URL env var (default: http://localhost:8080) timeout=120.0, # httpx timeout in seconds - max_retries=2, # connection error retries (default: 2) + max_retries=2, # retries for connect/timeout/408/429/5xx (default: 2) default_headers={...}, # merged into every request http_client=my_httpx, # bring your own httpx.Client ) ``` -Auth resolution order: -1. `api_key` parameter -2. `FERRO_API_KEY` environment variable -3. `OPENAI_API_KEY` environment variable (migration fallback) - -If none is found, `FerroAuthError` is raised at construction time. +Auth resolution order: `api_key` parameter → `FERRO_API_KEY` → `OPENAI_API_KEY` (migration fallback). If none is found, `FerroAuthError` is raised at construction time. --- @@ -348,13 +339,10 @@ Dev only: └── ruff ≥ 0.1.0 ``` -The SDK intentionally keeps zero additional runtime dependencies to minimize install footprint and avoid version conflicts in user projects. - --- ## Testing Strategy -- All tests live in `tests/test_sdk.py`. -- HTTP is fully mocked via `pytest-httpx` — no network access, no running gateway. +- **Unit** (`tests/test_client.py`, `test_chat.py`, `test_resources.py`, `test_admin.py`): HTTP fully mocked via `pytest-httpx`; covers construction, auth resolution, error mapping, the retry policy, header metadata, streaming (happy path, usage chunk, error frames), every resource and admin route, sync and async. +- **Contract** (`tests/contract/`): runs only when `FERRO_CONTRACT_BASE_URL` is set. `scripts/with-gateway.sh` builds `ferrogw` from an ai-gateway checkout, starts `tests/contract/stub_upstream.py` as a fake OpenAI, points the gateway at it, and asserts the header/body/error contract described above against the real server. CI runs it against the pinned gateway tag (required) and `main` (advisory). - Async tests use `pytest-asyncio` with `asyncio_mode = "auto"`. -- Tests cover: client construction, auth resolution, error mapping, retries, response parsing, streaming, admin CRUD, and resource wiring. diff --git a/integrations/README.md b/integrations/README.md index c3ea0cd..c436eae 100644 --- a/integrations/README.md +++ b/integrations/README.md @@ -23,7 +23,9 @@ python -m build twine upload dist/* ``` -CI workflows for each package live under `.github/workflows/publish-.yml` (to be added). +CI workflows for each package live under `.github/workflows/publish-.yml` +(`publish-langchain-ferrolabsai.yml`, `publish-llama-index-llms-ferrolabsai.yml`): they run the +sub-package's test matrix on PRs touching its folder and publish via Trusted Publishing on a matching tag. ## Upstream mirroring diff --git a/integrations/langchain-ferrolabsai/README.md b/integrations/langchain-ferrolabsai/README.md index 1d6e67a..b4d70bb 100644 --- a/integrations/langchain-ferrolabsai/README.md +++ b/integrations/langchain-ferrolabsai/README.md @@ -25,14 +25,14 @@ from langchain_core.messages import HumanMessage llm = FerroChatModel( model="gpt-4o", - base_url="http://localhost:8080", # any Ferro Labs AI Gateway instance + base_url="http://localhost:8080", # any Ferro Labs AI Gateway instance api_key="fgw_...", ) response = llm.invoke([HumanMessage(content="Hello, world")]) print(response.content) -print(response.response_metadata["provider"]) # which provider answered -print(response.response_metadata["trace_id"]) # gateway X-Request-ID +print(response.response_metadata["provider"]) # which provider answered +print(response.response_metadata["trace_id"]) # gateway X-Request-ID print(response.response_metadata.get("gateway_overhead_ms")) # gateway's own overhead ``` @@ -46,7 +46,7 @@ name: ```python claude = FerroChatModel(model="claude-3-5-sonnet-20241022", base_url="...", api_key="...") -gemini = FerroChatModel(model="gemini-2.5-flash", base_url="...", api_key="...") +gemini = FerroChatModel(model="gemini-2.5-flash", base_url="...", api_key="...") ``` ### Streaming @@ -70,11 +70,13 @@ async for chunk in llm.astream([HumanMessage(content="Tell me a story")]): ```python from langchain_core.tools import tool + @tool def add(a: int, b: int) -> int: """Add two integers.""" return a + b + agent_llm = llm.bind_tools([add]) response = agent_llm.invoke([HumanMessage(content="What is 4 + 7?")]) print(response.tool_calls) @@ -89,12 +91,14 @@ it works with every provider the gateway can translate it for (see ```python from pydantic import BaseModel + class Answer(BaseModel): city: str population: int + structured = llm.with_structured_output(Answer) -print(structured.invoke("Largest city in France?")) # Answer(city='Paris', population=...) +print(structured.invoke("Largest city in France?")) # Answer(city='Paris', population=...) # include_raw=True → {"raw": AIMessage, "parsed": Answer | None, "parsing_error": ...} ``` From d7abcd9d539a7a88ca9f4bb8233213e02a96ffc3 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 13:04:23 +0530 Subject: [PATCH 05/14] chore: release 0.3.0 - pyproject.toml / ferrolabsai/_version.py -> 0.3.0; description says 30 providers; Python 3.13 classifier. - CHANGELOG 0.3.0 with Breaking / Added / Fixed / Removed. --- .github/copilot-instructions.md | 2 +- AGENT.md | 2 +- CHANGELOG.md | 103 ++++++++++++++++++++++++++++++++ docs/architecture.md | 2 +- ferrolabsai/_version.py | 2 +- pyproject.toml | 5 +- 6 files changed, 110 insertions(+), 6 deletions(-) diff --git a/.github/copilot-instructions.md b/.github/copilot-instructions.md index ee445e6..4a4353f 100644 --- a/.github/copilot-instructions.md +++ b/.github/copilot-instructions.md @@ -25,7 +25,7 @@ - Target Python 3.9 syntax. New modules should use `from __future__ import annotations`, public functions and methods are expected to be fully typed, and changes should stay compatible with the repo's `mypy --strict` setup. - Ruff is both the linter and formatter here. Keep changes aligned with the existing 100-character line length and prefer the repo's Ruff formatting over hand-formatting. - Preserve the OpenAI-style surface first. New capabilities should usually be exposed as additional kwargs on existing resource methods or as new resource namespaces that match gateway routes, not as a parallel custom API shape. -- Only send request fields the gateway actually decodes (`internal/handler/chatrequest.go`). `max_completion_tokens`, `parallel_tool_calls`, `response_format`, `stream_options`, `seed` are first-class; anything else passes through `**kwargs` verbatim. Do not reintroduce `route_tag`, `template_id`, or `template_variables` — the gateway never read them. +- Only send request fields the gateway actually decodes (`internal/handler/chatrequest.go`). `max_completion_tokens`, `parallel_tool_calls`, `response_format`, `stream_options`, `seed` are first-class; anything else passes through `**kwargs` verbatim. Do not reintroduce the old routing-tag / prompt-template request fields — the gateway never read them. - Keep resource classes thin. Shared behavior such as retries, auth headers, connection handling, status-code mapping, and HTTP client lifecycle belongs in `client.py`, not duplicated across resource modules. - Public resource methods should return typed dataclasses parsed via `from_dict(...)` helpers in `types.py` unless the endpoint is intentionally passthrough admin data. - Error translation is centralized in `_raise_api_error(...)`. Extend the existing `Ferro*Error` hierarchy instead of leaking raw `httpx` exceptions from public SDK methods. diff --git a/AGENT.md b/AGENT.md index 49a8a26..5ebe423 100644 --- a/AGENT.md +++ b/AGENT.md @@ -142,7 +142,7 @@ Always run `make format lint test` before committing; run `make contract` when t - All public exports must be listed in `ferrolabsai/__init__.py` and the `__all__` list. - The SDK mirrors the OpenAI SDK surface: `client.chat.completions.create()`, `client.embeddings.create()`, `client.images.generate()`, `client.models.list()`, `client.responses.create()`. - Gateway-specific surface: `client.health()/ready()/live()/capabilities()/rerank()`, `client.moderations`, `client.admin.*`. -- **Only model what the gateway really does.** Response metadata comes from `X-Request-ID`, `X-Gateway-Provider`, `X-Gateway-Overhead-Ms`, and the body fields `provider` / `provider_metadata` / `reasoning_content` / usage counters. There is no cost, cache-hit, or latency field for callers, and no `route_tag` / `template_*` request field — do not reintroduce them. `docs/architecture.md` § "Gateway Contract" is the reference; the contract suite enforces it. +- **Only model what the gateway really does.** Response metadata comes from `X-Request-ID`, `X-Gateway-Provider`, `X-Gateway-Overhead-Ms`, and the body fields `provider` / `provider_metadata` / `reasoning_content` / usage counters. There is no cost, cache-hit, or latency field for callers, and no per-request routing-tag or prompt-template request field — do not reintroduce them. `docs/architecture.md` § "Gateway Contract" is the reference; the contract suite enforces it. ### Environment Variables - `FERRO_API_KEY` — primary API key (takes precedence). diff --git a/CHANGELOG.md b/CHANGELOG.md index 754572a..c0047bf 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,109 @@ This project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.htm --- +## [0.3.0] — 2026-08-29 + +The "truth release": every claim the SDK makes now matches what +[ai-gateway](https://github.com/ferro-labs/ai-gateway) actually does, and a +contract suite (`tests/contract/`, run in CI against `v1.4.5` and `main`) +keeps it that way. Compatibility: `ferrolabsai 0.3.x` ↔ `ai-gateway ≥ v1.4.0`. + +### Breaking +- **Response metadata reads the gateway's real headers.** `trace_id` comes + from `X-Request-ID`, `provider` from the chat body's `provider` field (or + `X-Gateway-Provider` on responses/pass-through), and the new + `gateway_overhead_ms` from `X-Gateway-Overhead-Ms`. The legacy + `x-trace-id` / `x-ferro-request-id` header fallbacks and the `x_ferro_*` + body keys are gone. +- **Removed `ChatCompletion.latency_ms`, `Usage.cost_usd`, `Usage.cache_hit`, + `Usage.provider`.** They were read from `X-Ferro-Latency-Ms` / + `X-Ferro-Cost-Usd` / `X-Ferro-Provider` headers that the gateway has never + emitted at any version, so they were always `None`. `gateway_overhead_ms` + is *not* a rename of `latency_ms`: it is the gateway's own processing time, + not end-to-end latency. Cost and cache hits remain available in the request + log (`admin.logs.list()`), Prometheus, and OTel — not to callers. +- **Removed `route_tag` / `x_route_tag`, `template_id`, `template_variables`** + from `chat.completions.create()`. The gateway never decoded them + (`internal/handler/chatrequest.go`); unknown kwargs still pass through + verbatim. +- **`ModelInfo` is the gateway's `EnrichedModelInfo`.** `provider` is now a + read-only alias of the new `owned_by`; `input_cost_per_token` / + `output_cost_per_token` are removed (never served); added `created`, `mode`, + `max_output_tokens`, `deprecated`; `capabilities` defaults to `[]`. +- **`models.retrieve(id)` no longer calls `GET /v1/models/{id}`.** That path is + not a native gateway route — it fell through to the `/v1/*` pass-through and + was forwarded upstream with the operator's provider credential. It is now a + client-side lookup over `GET /v1/models` and raises `FerroNotFoundError` + (`code="model_not_found"`) locally. `models.list(provider=, capability=)` + and `models.search()` filter client-side too (the gateway ignores query + parameters); `provider` matches `owned_by`, `capability` matches + `capabilities[]`, `search` is a case-insensitive substring on `id`. +- **Streaming returns a `Stream` / `AsyncStream` object** instead of a bare + generator. It iterates exactly as before and additionally exposes + `trace_id`, `provider`, `response`, `close()`, and context-manager support. + HTTP errors on a stream now raise when `create(stream=True)` is called (async: + when awaited), not on first iteration. +- **Retries now cover HTTP 408/429/5xx** (previously connect/timeout only), + with full-jitter backoff and `Retry-After` honoured (capped at 30 s). + `max_retries=0` disables them. Streaming is never retried. +- `FerroRateLimitError.__init__` gained keyword-only `retry_after`; + `FerroStreamError.__init__` gained keyword-only `code`. +- `_request(stream=True)` (private) is removed; `_open_stream()` replaces it. + +### Added +- `stream_options` parameter (e.g. `{"include_usage": True}`); + `ChatCompletionChunk.usage`, `.trace_id`, `.provider`; `Usage.reasoning_tokens`, + `.cache_read_tokens`, `.cache_write_tokens`; `ChatMessage.reasoning_content`, + `StreamDelta.reasoning_content`; `ChatCompletion.provider_metadata`. + Mid-stream `{"error": ...}` frames raise `FerroStreamError` with `.code` + (`stream_error`, `stream_timeout`). +- Request params `max_completion_tokens`, `parallel_tool_calls`, + `response_format`, `seed` on `chat.completions.create()` (sync + async). +- `FerroBudgetExceededError` (402 `insufficient_quota`) and + `FerroPermissionError` (403 `insufficient_scope`); every `FerroAPIError` + now carries `status_code` and the gateway's `code`. +- `client.responses.create() / retrieve(id) / delete(id)` → `/v1/responses` + (`Response` dataclass; id routes answer 501 `responses_not_configured` + unless the gateway sets `responses_target`). +- `client.capabilities()` → `/v1/capabilities`; `client.health()`, + `client.ready()`, `client.live()` → `/health`, `/readyz`, `/livez` (503 + bodies are returned, not raised); `client.rerank()` → `/v1/rerank`; + `client.moderations.create()` → `/v1/moderations`. +- Admin parity with ai-gateway 1.4: `admin.audit.list()`, + `admin.providers.catalog()`, `admin.plugins.catalog()`, + `admin.logs.list(api_key_id=)` (and `stage="all"`), + `admin.logs.stats(buckets=)`. +- `EmbeddingResponse.trace_id`, `ImageResponse.trace_id`. +- Python 3.13 in the CI matrix and classifiers; mypy runs on every leg. +- Contract suite: `scripts/with-gateway.sh` + `tests/contract/` (stub + upstream, 23 assertions against a real gateway); `make contract`. +- Unit tests for the previously open good-first-issues (#8 keys + retrieve/update, #9 config create/delete, #10 logs delete, #12 sync + models.search/images.generate/admin.health/providers.list, #14 retry + exhaustion → `FerroConnectionError`, #19 async streaming + embeddings). + +### Fixed +- Header metadata is merged into inference bodies only; `/v1/models`, + probes, and `/admin/*` bodies are returned untouched. +- `admin.providers.catalog()` unwraps the gateway's `{"providers": [...]}` + envelope. +- `__version__` is a constant kept in sync with `pyproject.toml` (the + installed-distribution lookup reported a stale number in editable checkouts). +- Resource classes type their `client` (no more `Any` / + `# type: ignore[no-any-return]`). +- README: "29 providers" → 30 ; `admin.logs.list(trace_id=...)` example (no + such parameter) replaced; dead `internal/admin/handlers.go` links → the + `internal/admin/handlers` package; observability section rewritten to list + exactly what populates and from where; framework section now points at + `langchain-ferrolabsai`. + +### Removed +- Everything listed under **Breaking**: `latency_ms`, `cost_usd`, + `cache_hit`, `Usage.provider`, `route_tag`, `template_id`, + `template_variables`, `ModelInfo` pricing fields, `x_ferro_*` / + `X-Ferro-*` / `x-trace-id` handling. None of it was ever provided or read + by the gateway. + ## [0.2.1] — 2026-06-13 ### Fixed diff --git a/docs/architecture.md b/docs/architecture.md index a778d2b..b23041d 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -232,7 +232,7 @@ The `_ChatNamespace` class exists solely to provide the `client.chat.completions | `X-Gateway-Overhead-Ms` | non-streaming `/v1/chat/completions`, when > 0 | `gateway_overhead_ms` | | `Retry-After` | every gateway-originated `429` (`"1"`), upstream `429` (upstream value) | retry delay; `FerroRateLimitError.retry_after` | -No `X-Ferro-*`, cost, or cache-hit header exists at any gateway version; the SDK reads none. +No Ferro-branded, cost, latency, or cache-hit response header exists at any gateway version; the SDK reads none. ### Body extensions read by the SDK diff --git a/ferrolabsai/_version.py b/ferrolabsai/_version.py index 1005a44..d192e13 100644 --- a/ferrolabsai/_version.py +++ b/ferrolabsai/_version.py @@ -1,4 +1,4 @@ """Package version. Keep in sync with ``[project].version`` in pyproject.toml (tests/test_client.py asserts they match).""" -__version__ = "0.2.1" +__version__ = "0.3.0" diff --git a/pyproject.toml b/pyproject.toml index 569132e..e5d1e63 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,8 +4,8 @@ build-backend = "hatchling.build" [project] name = "ferrolabsai" -version = "0.2.1" -description = "Official Python SDK for Ferro Labs AI Gateway — route LLM requests across 29 providers with a single OpenAI-compatible API" +version = "0.3.0" +description = "Official Python SDK for Ferro Labs AI Gateway — route LLM requests across 30 providers with a single OpenAI-compatible API" readme = "README.md" license = { text = "Apache-2.0" } requires-python = ">=3.9" @@ -26,6 +26,7 @@ classifiers = [ "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", "Topic :: Software Development :: Libraries :: Python Modules", "Topic :: Scientific/Engineering :: Artificial Intelligence", ] From 3aca88c67abda56a1bd0120a021f44a98f7468ff Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 13:31:43 +0530 Subject: [PATCH 06/14] ci(contract): disable the per-IP limiter in the gateway harness The contract suite fires ~60 requests in well under a second on a fast runner and tripped the gateway's default per-IP bucket (20 rps / burst 40) with a 429 the maxRetries=0 fixture cannot absorb. RATE_LIMIT_RPS=0 removes that limiter (the /admin/session limiter is separate and stays); the suite tests the API contract, not the limiter. --- scripts/with-gateway.sh | 1 + 1 file changed, 1 insertion(+) diff --git a/scripts/with-gateway.sh b/scripts/with-gateway.sh index 521a4d4..22ac6ce 100755 --- a/scripts/with-gateway.sh +++ b/scripts/with-gateway.sh @@ -116,6 +116,7 @@ echo "==> starting gateway on :$port" MASTER_KEY="$key" GATEWAY_CONFIG="$work/gateway.yaml" PORT="$port" \ REQUEST_LOG_STORE_BACKEND=sqlite REQUEST_LOG_STORE_DSN="$work/requestlog.db" \ OPENAI_API_KEY=stub-key OPENAI_BASE_URL="http://127.0.0.1:$stub_port/v1" \ + RATE_LIMIT_RPS=0 \ "$work/ferrogw" serve >"$work/gateway.log" 2>&1 & gw_pid=$! From f8a4f5ef429b3fde961de32a6585f3fc7295d216 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:12:54 +0530 Subject: [PATCH 07/14] fix(retry): retry 408/5xx and read timeouts only for idempotent methods A POST that answers 408/5xx or times out mid-flight may already have been processed by the gateway, so re-sending it can double-charge or duplicate side effects. Status retries for 408/5xx and read/write/pool timeouts are now limited to GET/HEAD/PUT/DELETE/OPTIONS. 429 (the gateway did not process the request) and connect errors / connect timeouts (the request never left) are still retried for every method. The decision lives in one pure helper, _should_retry(method, status=|exc=), shared by the sync and async loops. Backoff, jitter and Retry-After are unchanged. README, architecture doc, agent guide and the 0.3.0 changelog describe the policy; the changelog also calls out that POST read timeouts are no longer retried (0.2.x retried every timeout). --- AGENT.md | 2 +- CHANGELOG.md | 11 ++- README.md | 2 +- docs/architecture.md | 8 ++- ferrolabsai/client.py | 32 +++++++-- tests/test_client.py | 29 ++++---- tests/test_retry.py | 154 ++++++++++++++++++++++++++++++++++++++++++ 7 files changed, 212 insertions(+), 26 deletions(-) create mode 100644 tests/test_retry.py diff --git a/AGENT.md b/AGENT.md index 5ebe423..9e0fa5f 100644 --- a/AGENT.md +++ b/AGENT.md @@ -134,7 +134,7 @@ Always run `make format lint test` before committing; run `make contract` when t - **Dataclass response models** — all types in `types.py` are `@dataclass` with a `from_dict()` classmethod. No pydantic dependency. - **Resource pattern** — each API area lives in its own sub-package with a `resource.py` (sync) and `async_resource.py`. The async module imports the body builders / path constants from the sync one so the wire format is written once. - **Client holds HTTP** — `_request()` (retry + error mapping + header metadata on inference paths) and `_open_stream()` (SSE, never retried) are the only HTTP entry points. Resources receive the typed client (`if TYPE_CHECKING: from ..client import FerroClient`) and call `self._client._request(...)`. -- **Exception hierarchy** — all HTTP errors raise typed exceptions inheriting from `FerroAPIError`. Connection/timeout errors and `408/429/5xx` retry with jittered backoff (honouring `Retry-After`), then raise `FerroConnectionError` / the mapped `FerroAPIError`. +- **Exception hierarchy** — all HTTP errors raise typed exceptions inheriting from `FerroAPIError`. `429`, connection errors, and connect timeouts retry for every method; `408/5xx` and read/write/pool timeouts retry only for idempotent methods (`_should_retry`). Jittered backoff honours `Retry-After`; exhaustion raises `FerroConnectionError` / the mapped `FerroAPIError`. - **Immutability by default** — do not mutate arguments; return new objects (`_with_response_metadata` returns a new dict, streaming uses `dataclasses.replace`). - **Keep files small** — prefer several focused modules over one large file (`types_responses.py` exists for that reason). diff --git a/CHANGELOG.md b/CHANGELOG.md index c0047bf..da6875e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -49,9 +49,14 @@ keeps it that way. Compatibility: `ferrolabsai 0.3.x` ↔ `ai-gateway ≥ v1.4.0 `trace_id`, `provider`, `response`, `close()`, and context-manager support. HTTP errors on a stream now raise when `create(stream=True)` is called (async: when awaited), not on first iteration. -- **Retries now cover HTTP 408/429/5xx** (previously connect/timeout only), - with full-jitter backoff and `Retry-After` honoured (capped at 30 s). - `max_retries=0` disables them. Streaming is never retried. +- **Retries are idempotent-aware.** HTTP `429`, connection errors, and + connect timeouts are retried for every method; HTTP `408` / `5xx` and + read / write / pool timeouts are retried only for `GET` / `HEAD` / `PUT` / + `DELETE` / `OPTIONS`. A `POST` that hits a read timeout is **no longer + retried** (0.2.x retried every timeout regardless of method) — it may + already have been processed. Full-jitter backoff, `Retry-After` honoured + (capped at 30 s). + `max_retries=0` disables retries. Streaming is never retried. - `FerroRateLimitError.__init__` gained keyword-only `retry_after`; `FerroStreamError.__init__` gained keyword-only `code`. - `_request(stream=True)` (private) is removed; `_open_stream()` replaces it. diff --git a/README.md b/README.md index 5e23c89..f3e9327 100644 --- a/README.md +++ b/README.md @@ -324,7 +324,7 @@ client = FerroClient( ) ``` -**Retries** cover connection errors, timeouts, and HTTP `408` / `429` / `5xx` — capped exponential backoff with full jitter (0.5 s base, 8 s cap), honouring `Retry-After` when the gateway sends one (capped at 30 s, the same cap the gateway applies upstream). Other `4xx` responses and **streaming requests are never retried**. +**Retries** are idempotent-aware. HTTP `429` is retried for every method (the gateway did not process the request), as are connection errors and connect timeouts (the request never left). HTTP `408` / `5xx` and read / write / pool timeouts are retried **only for idempotent methods** (`GET`, `HEAD`, `PUT`, `DELETE`, `OPTIONS`) — a `POST` that timed out mid-flight may already have been processed, so it is raised as-is. Delays use capped exponential backoff with full jitter (0.5 s base, 8 s cap), honouring `Retry-After` when the gateway sends one (capped at 30 s, the same cap the gateway applies upstream). Other `4xx` responses and **streaming requests are never retried**. **Bring-your-own httpx client** lets you configure proxies, custom TLS, connection pool limits, or instrumentation middleware and reuse that across the SDK: diff --git a/docs/architecture.md b/docs/architecture.md index b23041d..fe197f5 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -138,11 +138,15 @@ Shared by sync and async; streaming requests are never retried. | Trigger | Retried? | |---|---| -| `httpx.ConnectError`, `httpx.TimeoutException` | yes | -| HTTP `408`, `429`, `5xx` | yes | +| HTTP `429` | yes, every method (the gateway did not process the request) | +| `httpx.ConnectError`, `httpx.ConnectTimeout` | yes, every method (the request never left) | +| HTTP `408`, `5xx` | idempotent methods only (`GET`, `HEAD`, `PUT`, `DELETE`, `OPTIONS`) | +| other `httpx.TimeoutException` (read / write / pool) | idempotent methods only — a `POST` may already have been processed | | any other `4xx` | no — raised immediately | | streaming (`_open_stream`) | never | +The decision lives in `_should_retry(method, status=…)` / `_should_retry(method, exc=…)`, shared by both loops. + Delay before retry *n*: `Retry-After` seconds when present (capped at 30 s — the same cap the gateway applies to its own upstream retries), else `uniform(0, min(0.5 · 2^(n-1), 8))` (full jitter). `max_retries` defaults to 2 and is validated at construction. ### 5. Typed Exception Hierarchy diff --git a/ferrolabsai/client.py b/ferrolabsai/client.py index 8b9315a..5e5ef41 100644 --- a/ferrolabsai/client.py +++ b/ferrolabsai/client.py @@ -46,6 +46,7 @@ DEFAULT_RETRY_BACKOFF_MAX: float = 8.0 RETRY_AFTER_MAX: float = 30.0 # same cap the gateway applies to its own upstream retries RETRYABLE_STATUSES: frozenset[int] = frozenset({408, 429}) # plus every 5xx +IDEMPOTENT_METHODS: frozenset[str] = frozenset({"GET", "HEAD", "PUT", "DELETE", "OPTIONS"}) # Only inference bodies get header metadata (trace_id, provider, gateway_overhead_ms) # merged in; catalog, probe, and admin bodies are returned untouched. @@ -92,8 +93,23 @@ def _default_headers(api_key: str, extra: dict[str, str] | None) -> dict[str, st # ------------------------------------------------------------------ -def _is_retryable(status: int) -> bool: - return status in RETRYABLE_STATUSES or status >= 500 +def _should_retry(method: str, *, status: int | None = None, exc: Exception | None = None) -> bool: + """Whether a failed attempt may be re-sent. + + A 429 means the gateway did not process the request, and a connect error or + connect timeout means the request never left, so those retry for every + method. Any other retryable status (408/5xx) or timeout (read/write/pool) + leaves a non-idempotent request ambiguous — a POST may already have been + processed — so those retry only for idempotent methods. + """ + idempotent = method.upper() in IDEMPOTENT_METHODS + if status is not None: + if status == 429: + return True + return idempotent and (status in RETRYABLE_STATUSES or status >= 500) + if isinstance(exc, (httpx.ConnectError, httpx.ConnectTimeout)): + return True + return idempotent and isinstance(exc, httpx.TimeoutException) def _retry_after_seconds(response: httpx.Response) -> float | None: @@ -237,11 +253,13 @@ def _request( response.raise_for_status() return _parse_body(response, path) except httpx.HTTPStatusError as e: - if attempt >= self.max_retries or not _is_retryable(e.response.status_code): + if attempt >= self.max_retries or not _should_retry( + method, status=e.response.status_code + ): _raise_api_error(e) delay = _retry_delay(attempt + 1, _retry_after_seconds(e.response)) except (httpx.ConnectError, httpx.TimeoutException) as e: - if attempt >= self.max_retries: + if attempt >= self.max_retries or not _should_retry(method, exc=e): raise _connection_error(e, self.base_url, self.timeout) from e delay = _retry_delay(attempt + 1) attempt += 1 @@ -381,11 +399,13 @@ async def _request( response.raise_for_status() return _parse_body(response, path) except httpx.HTTPStatusError as e: - if attempt >= self.max_retries or not _is_retryable(e.response.status_code): + if attempt >= self.max_retries or not _should_retry( + method, status=e.response.status_code + ): _raise_api_error(e) delay = _retry_delay(attempt + 1, _retry_after_seconds(e.response)) except (httpx.ConnectError, httpx.TimeoutException) as e: - if attempt >= self.max_retries: + if attempt >= self.max_retries or not _should_retry(method, exc=e): raise _connection_error(e, self.base_url, self.timeout) from e delay = _retry_delay(attempt + 1) attempt += 1 diff --git a/tests/test_client.py b/tests/test_client.py index ee9bd85..34ef779 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -25,6 +25,7 @@ from .conftest import API_KEY, BASE_URL, COMPLETION_RESPONSE, TRACE_ID CHAT_URL = f"{BASE_URL}/v1/chat/completions" +CAPABILITIES_URL = f"{BASE_URL}/v1/capabilities" def _chat(client: FerroClient): @@ -241,7 +242,7 @@ def test_sync_retries_connect_errors_with_backoff( monkeypatch.setattr("ferrolabsai.client.time.sleep", sleeps.append) client.max_retries = 2 httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) - httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + httpx_mock.add_exception(httpx.ConnectTimeout("slow"), method="POST", url=CHAT_URL) httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) assert _chat(client).id == "chatcmpl-abc123" assert sleeps == [0.5, 1.0] @@ -265,10 +266,10 @@ def test_sync_retries_retryable_statuses( monkeypatch.setattr("ferrolabsai.client.time.sleep", sleeps.append) client.max_retries = 1 httpx_mock.add_response( - method="POST", url=CHAT_URL, status_code=status, json={"error": {"message": "x"}} + method="GET", url=CAPABILITIES_URL, status_code=status, json={"error": {"message": "x"}} ) - httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) - assert _chat(client).id == "chatcmpl-abc123" + httpx_mock.add_response(method="GET", url=CAPABILITIES_URL, json={"providers": {}}) + assert client.capabilities() == {"providers": {}} assert sleeps == [0.5] def test_sync_honours_retry_after_capped(self, monkeypatch, client, httpx_mock: HTTPXMock): @@ -309,17 +310,20 @@ def test_sync_raises_after_last_retryable_status( client.max_retries = 1 for _ in range(2): httpx_mock.add_response( - method="POST", url=CHAT_URL, status_code=503, json={"error": {"message": "x"}} + method="GET", + url=CAPABILITIES_URL, + status_code=503, + json={"error": {"message": "x"}}, ) with pytest.raises(FerroServerError): - _chat(client) + client.capabilities() def test_streaming_is_never_retried(self, client, httpx_mock: HTTPXMock): client.max_retries = 2 httpx_mock.add_response( - method="POST", url=CHAT_URL, status_code=503, json={"error": {"message": "x"}} + method="POST", url=CHAT_URL, status_code=429, json={"error": {"message": "x"}} ) - with pytest.raises(FerroServerError): + with pytest.raises(FerroRateLimitError): client.chat.completions.create( model="gpt-4o", messages=[{"role": "user", "content": "Hi"}], stream=True ) @@ -358,12 +362,11 @@ async def fake_sleep(_delay: float) -> None: monkeypatch.setattr("ferrolabsai.client.asyncio.sleep", fake_sleep) async_client.max_retries = 1 - httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) - httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="GET", url=CAPABILITIES_URL) + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="GET", url=CAPABILITIES_URL) with pytest.raises(FerroConnectionError, match="timed out"): - await async_client.chat.completions.create( - model="gpt-4o", messages=[{"role": "user", "content": "Hi"}] - ) + await async_client.capabilities() + assert len(httpx_mock.get_requests()) == 2 class TestGatewayEndpoints: diff --git a/tests/test_retry.py b/tests/test_retry.py new file mode 100644 index 0000000..60a5bf3 --- /dev/null +++ b/tests/test_retry.py @@ -0,0 +1,154 @@ +"""Retry policy: status/exception retries are idempotent-only except 429 and connect +failures.""" + +from __future__ import annotations + +import httpx +import pytest +from pytest_httpx import HTTPXMock + +from ferrolabsai.client import _should_retry +from ferrolabsai.exceptions import FerroConnectionError, FerroServerError + +from .conftest import BASE_URL, COMPLETION_RESPONSE + +CHAT_URL = f"{BASE_URL}/v1/chat/completions" +CAPABILITIES_URL = f"{BASE_URL}/v1/capabilities" +KEY_URL = f"{BASE_URL}/admin/keys/k1" +ERROR_BODY = {"error": {"message": "x"}} +MESSAGES = [{"role": "user", "content": "Hi"}] + + +class TestShouldRetry: + @pytest.mark.parametrize( + ("method", "status", "expected"), + [ + ("POST", 500, False), + ("POST", 408, False), + ("POST", 429, True), + ("GET", 500, True), + ("PUT", 503, True), + ("DELETE", 502, True), + ("GET", 400, False), + ("get", 500, True), + ], + ) + def test_status(self, method, status, expected): + assert _should_retry(method, status=status) is expected + + @pytest.mark.parametrize( + ("method", "exc", "expected"), + [ + ("POST", httpx.ConnectError("refused"), True), + ("POST", httpx.ConnectTimeout("slow"), True), + ("POST", httpx.ReadTimeout("slow"), False), + ("POST", httpx.WriteTimeout("slow"), False), + ("POST", httpx.PoolTimeout("slow"), False), + ("GET", httpx.ReadTimeout("slow"), True), + ("GET", httpx.PoolTimeout("slow"), True), + ("GET", RuntimeError("other"), False), + ], + ) + def test_exception(self, method, exc, expected): + assert _should_retry(method, exc=exc) is expected + + +class TestRequestRetryPolicy: + @pytest.fixture(autouse=True) + def _no_sleep(self, monkeypatch): + monkeypatch.setattr("ferrolabsai.client.time.sleep", lambda _s: None) + + async def fake_sleep(_delay: float) -> None: + pass + + monkeypatch.setattr("ferrolabsai.client.asyncio.sleep", fake_sleep) + + def test_post_500_is_not_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 2 + httpx_mock.add_response(method="POST", url=CHAT_URL, status_code=500, json=ERROR_BODY) + with pytest.raises(FerroServerError): + client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert len(httpx_mock.get_requests()) == 1 + + def test_post_429_is_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 1 + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=429, + json=ERROR_BODY, + headers={"Retry-After": "0"}, + ) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + response = client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert response.id == "chatcmpl-abc123" + assert len(httpx_mock.get_requests()) == 2 + + def test_get_500_is_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 1 + httpx_mock.add_response( + method="GET", url=CAPABILITIES_URL, status_code=500, json=ERROR_BODY + ) + httpx_mock.add_response(method="GET", url=CAPABILITIES_URL, json={"providers": {}}) + assert client.capabilities() == {"providers": {}} + assert len(httpx_mock.get_requests()) == 2 + + @pytest.mark.parametrize(("method", "status"), [("PUT", 503), ("DELETE", 502)]) + def test_put_and_delete_5xx_are_retried(self, client, httpx_mock: HTTPXMock, method, status): + client.max_retries = 1 + httpx_mock.add_response(method=method, url=KEY_URL, status_code=status, json=ERROR_BODY) + httpx_mock.add_response(method=method, url=KEY_URL, json={"ok": True}) + assert client._request(method, "/admin/keys/k1") == {"ok": True} + assert len(httpx_mock.get_requests()) == 2 + + def test_post_read_timeout_is_not_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 2 + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + with pytest.raises(FerroConnectionError, match="timed out"): + client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert len(httpx_mock.get_requests()) == 1 + + def test_get_read_timeout_is_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 1 + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="GET", url=CAPABILITIES_URL) + httpx_mock.add_response(method="GET", url=CAPABILITIES_URL, json={"providers": {}}) + assert client.capabilities() == {"providers": {}} + assert len(httpx_mock.get_requests()) == 2 + + def test_post_connect_error_is_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 1 + httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + response = client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert response.id == "chatcmpl-abc123" + assert len(httpx_mock.get_requests()) == 2 + + async def test_async_post_500_is_not_retried(self, async_client, httpx_mock: HTTPXMock): + async_client.max_retries = 2 + httpx_mock.add_response(method="POST", url=CHAT_URL, status_code=500, json=ERROR_BODY) + with pytest.raises(FerroServerError): + await async_client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert len(httpx_mock.get_requests()) == 1 + + async def test_async_post_429_is_retried(self, async_client, httpx_mock: HTTPXMock): + async_client.max_retries = 1 + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + status_code=429, + json=ERROR_BODY, + headers={"Retry-After": "0"}, + ) + httpx_mock.add_response(method="POST", url=CHAT_URL, json=COMPLETION_RESPONSE) + response = await async_client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert response.id == "chatcmpl-abc123" + assert len(httpx_mock.get_requests()) == 2 + + async def test_async_post_read_timeout_is_not_retried( + self, async_client, httpx_mock: HTTPXMock + ): + async_client.max_retries = 2 + httpx_mock.add_exception(httpx.ReadTimeout("slow"), method="POST", url=CHAT_URL) + with pytest.raises(FerroConnectionError, match="timed out"): + await async_client.chat.completions.create(model="gpt-4o", messages=MESSAGES) + assert len(httpx_mock.get_requests()) == 1 From 415ca3c025c5d4f27ae8f61a5d48735da4d660f0 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:13:22 +0530 Subject: [PATCH 08/14] fix(stream): map transport errors on stream open to FerroConnectionError _open_stream (sync and async) let httpx.ConnectError / TimeoutException escape raw while _request mapped them. Wrap send() with the same _connection_error() so callers get one exception type either way. Streams are still never retried. --- CHANGELOG.md | 3 ++- docs/architecture.md | 2 +- ferrolabsai/client.py | 10 ++++++++-- tests/test_retry.py | 35 +++++++++++++++++++++++++++++++++-- 4 files changed, 44 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index da6875e..d46e996 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -56,7 +56,8 @@ keeps it that way. Compatibility: `ferrolabsai 0.3.x` ↔ `ai-gateway ≥ v1.4.0 retried** (0.2.x retried every timeout regardless of method) — it may already have been processed. Full-jitter backoff, `Retry-After` honoured (capped at 30 s). - `max_retries=0` disables retries. Streaming is never retried. + `max_retries=0` disables retries. Streaming is never retried, and its + transport failures now raise `FerroConnectionError` like `_request` does. - `FerroRateLimitError.__init__` gained keyword-only `retry_after`; `FerroStreamError.__init__` gained keyword-only `code`. - `_request(stream=True)` (private) is removed; `_open_stream()` replaces it. diff --git a/docs/architecture.md b/docs/architecture.md index fe197f5..167dd67 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -143,7 +143,7 @@ Shared by sync and async; streaming requests are never retried. | HTTP `408`, `5xx` | idempotent methods only (`GET`, `HEAD`, `PUT`, `DELETE`, `OPTIONS`) | | other `httpx.TimeoutException` (read / write / pool) | idempotent methods only — a `POST` may already have been processed | | any other `4xx` | no — raised immediately | -| streaming (`_open_stream`) | never | +| streaming (`_open_stream`) | never; transport errors map to `FerroConnectionError` | The decision lives in `_should_retry(method, status=…)` / `_should_retry(method, exc=…)`, shared by both loops. diff --git a/ferrolabsai/client.py b/ferrolabsai/client.py index 5e5ef41..5baea57 100644 --- a/ferrolabsai/client.py +++ b/ferrolabsai/client.py @@ -270,7 +270,10 @@ def _open_stream(self, path: str, json: Any) -> httpx.Response: request = self._http.build_request( "POST", path, json=json, headers={"Accept": "text/event-stream"} ) - response = self._http.send(request, stream=True) + try: + response = self._http.send(request, stream=True) + except (httpx.ConnectError, httpx.TimeoutException) as e: + raise _connection_error(e, self.base_url, self.timeout) from e try: response.raise_for_status() except httpx.HTTPStatusError as e: @@ -415,7 +418,10 @@ async def _open_stream(self, path: str, json: Any) -> httpx.Response: request = self._http.build_request( "POST", path, json=json, headers={"Accept": "text/event-stream"} ) - response = await self._http.send(request, stream=True) + try: + response = await self._http.send(request, stream=True) + except (httpx.ConnectError, httpx.TimeoutException) as e: + raise _connection_error(e, self.base_url, self.timeout) from e try: response.raise_for_status() except httpx.HTTPStatusError as e: diff --git a/tests/test_retry.py b/tests/test_retry.py index 60a5bf3..5e7629c 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -1,5 +1,5 @@ """Retry policy: status/exception retries are idempotent-only except 429 and connect -failures.""" +failures; streaming maps transport errors but never retries.""" from __future__ import annotations @@ -8,7 +8,7 @@ from pytest_httpx import HTTPXMock from ferrolabsai.client import _should_retry -from ferrolabsai.exceptions import FerroConnectionError, FerroServerError +from ferrolabsai.exceptions import FerroConnectionError, FerroRateLimitError, FerroServerError from .conftest import BASE_URL, COMPLETION_RESPONSE @@ -152,3 +152,34 @@ async def test_async_post_read_timeout_is_not_retried( with pytest.raises(FerroConnectionError, match="timed out"): await async_client.chat.completions.create(model="gpt-4o", messages=MESSAGES) assert len(httpx_mock.get_requests()) == 1 + + +class TestStreamTransportErrors: + """``_open_stream`` maps transport failures like ``_request`` does, and never retries.""" + + @pytest.mark.parametrize( + ("exc", "match"), + [(httpx.ConnectError("refused"), "Cannot reach"), (httpx.ReadTimeout("slow"), "timed out")], + ) + def test_sync_stream_maps_transport_errors(self, client, httpx_mock: HTTPXMock, exc, match): + client.max_retries = 2 + httpx_mock.add_exception(exc, method="POST", url=CHAT_URL) + with pytest.raises(FerroConnectionError, match=match): + client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) + assert len(httpx_mock.get_requests()) == 1 + + async def test_async_stream_maps_transport_errors(self, async_client, httpx_mock: HTTPXMock): + async_client.max_retries = 2 + httpx_mock.add_exception(httpx.ConnectError("refused"), method="POST", url=CHAT_URL) + with pytest.raises(FerroConnectionError, match="Cannot reach"): + await async_client.chat.completions.create( + model="gpt-4o", messages=MESSAGES, stream=True + ) + assert len(httpx_mock.get_requests()) == 1 + + def test_stream_429_is_not_retried(self, client, httpx_mock: HTTPXMock): + client.max_retries = 2 + httpx_mock.add_response(method="POST", url=CHAT_URL, status_code=429, json=ERROR_BODY) + with pytest.raises(FerroRateLimitError): + client.chat.completions.create(model="gpt-4o", messages=MESSAGES, stream=True) + assert len(httpx_mock.get_requests()) == 1 From 7ab13359d8af86314b44586d3a14a2e1f685370b Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:13:24 +0530 Subject: [PATCH 09/14] fix(retry): ignore negative, NaN and infinite Retry-After values float() accepts "nan", "inf" and negative numbers, which then reached time.sleep()/asyncio.sleep() and FerroRateLimitError.retry_after. _retry_after_seconds now returns None unless the value is finite and non-negative, so such headers fall back to jittered backoff. --- CHANGELOG.md | 2 +- ferrolabsai/client.py | 8 +++++--- tests/test_retry.py | 14 ++++++++++++-- 3 files changed, 18 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index d46e996..0dcd11b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -55,7 +55,7 @@ keeps it that way. Compatibility: `ferrolabsai 0.3.x` ↔ `ai-gateway ≥ v1.4.0 `DELETE` / `OPTIONS`. A `POST` that hits a read timeout is **no longer retried** (0.2.x retried every timeout regardless of method) — it may already have been processed. Full-jitter backoff, `Retry-After` honoured - (capped at 30 s). + (capped at 30 s; negative, `NaN`, and infinite values are ignored). `max_retries=0` disables retries. Streaming is never retried, and its transport failures now raise `FerroConnectionError` like `_request` does. - `FerroRateLimitError.__init__` gained keyword-only `retry_after`; diff --git a/ferrolabsai/client.py b/ferrolabsai/client.py index 5baea57..21f181c 100644 --- a/ferrolabsai/client.py +++ b/ferrolabsai/client.py @@ -6,6 +6,7 @@ from __future__ import annotations import asyncio +import math import os import random import time @@ -113,12 +114,13 @@ def _should_retry(method: str, *, status: int | None = None, exc: Exception | No def _retry_after_seconds(response: httpx.Response) -> float | None: - """``Retry-After`` in seconds, or None when absent / not a number (HTTP-date form).""" - value = response.headers.get("retry-after") + """``Retry-After`` in seconds, or None when absent, not a number (HTTP-date form), + negative, NaN, or infinite.""" try: - return float(value) if value else None + seconds = float(response.headers.get("retry-after", "")) except ValueError: return None + return seconds if math.isfinite(seconds) and seconds >= 0 else None def _retry_delay(attempt: int, retry_after: float | None = None) -> float: diff --git a/tests/test_retry.py b/tests/test_retry.py index 5e7629c..61c9a1f 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -1,5 +1,5 @@ """Retry policy: status/exception retries are idempotent-only except 429 and connect -failures; streaming maps transport errors but never retries.""" +failures; streaming maps transport errors but never retries; Retry-After parsing.""" from __future__ import annotations @@ -7,7 +7,7 @@ import pytest from pytest_httpx import HTTPXMock -from ferrolabsai.client import _should_retry +from ferrolabsai.client import _retry_after_seconds, _should_retry from ferrolabsai.exceptions import FerroConnectionError, FerroRateLimitError, FerroServerError from .conftest import BASE_URL, COMPLETION_RESPONSE @@ -53,6 +53,16 @@ def test_exception(self, method, exc, expected): assert _should_retry(method, exc=exc) is expected +class TestRetryAfterSeconds: + @pytest.mark.parametrize( + ("value", "expected"), + [("-1", None), ("nan", None), ("inf", None), ("abc", None), ("2", 2.0), (None, None)], + ) + def test_parses_only_finite_non_negative_numbers(self, value, expected): + headers = {"Retry-After": value} if value is not None else {} + assert _retry_after_seconds(httpx.Response(429, headers=headers)) == expected + + class TestRequestRetryPolicy: @pytest.fixture(autouse=True) def _no_sleep(self, monkeypatch): From 891bf635f2a84a036bc88ed9be69dea665fde5af Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:13:24 +0530 Subject: [PATCH 10/14] test(contract): delete the parked admin guard key at session end admin_guard is now a generator fixture that removes the key it created in a finally block; a failure there is swallowed so teardown never masks a test failure. --- tests/contract/conftest.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/tests/contract/conftest.py b/tests/contract/conftest.py index 906d2db..f7b3880 100644 --- a/tests/contract/conftest.py +++ b/tests/contract/conftest.py @@ -44,11 +44,18 @@ def client() -> Iterator[FerroClient]: @pytest.fixture(scope="session") -def admin_guard(client: FerroClient) -> str: +def admin_guard(client: FerroClient) -> Iterator[str]: """The gateway refuses to revoke/delete the last admin *record* (409) — the MASTER_KEY is not a record — so tests that delete admin keys need one extra admin key parked for the whole session.""" - return client.admin.keys.create(name="contract-guard", scopes=["admin"]).id + key_id = client.admin.keys.create(name="contract-guard", scopes=["admin"]).id + try: + yield key_id + finally: + try: + client.admin.keys.delete(key_id) + except Exception: + pass # teardown must never mask a test failure @pytest.fixture From 905f0eb90671677d42ebeb56d15fdcf73d770918 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:13:46 +0530 Subject: [PATCH 11/14] fix(langchain): surface streamed usage_metadata from the terminal usage chunk _chunk_to_generation dropped every chunk without choices, including the usage-only chunk the gateway sends with stream_options include_usage. It now yields an empty AIMessageChunk carrying usage_metadata so the token counts survive chunk aggregation; chunks with neither choices nor usage still yield None. --- .../langchain_ferrolabsai/chat_models.py | 24 ++++++++++------- .../tests/test_chat_models.py | 27 +++++++++++++++++++ 2 files changed, 42 insertions(+), 9 deletions(-) diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py index 0676eba..e3a699d 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py @@ -22,6 +22,7 @@ from langchain_core.callbacks import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun from langchain_core.language_models import BaseChatModel, LanguageModelInput from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage +from langchain_core.messages.ai import UsageMetadata from langchain_core.output_parsers import JsonOutputParser, PydanticOutputParser from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult from langchain_core.runnables import Runnable, RunnableMap, RunnablePassthrough @@ -328,21 +329,26 @@ def _response_metadata(response: ChatCompletion) -> dict[str, Any]: return {k: v for k, v in metadata.items() if v is not None} -def _usage_metadata(response: ChatCompletion) -> dict[str, int] | None: +def _usage_metadata(response: ChatCompletion | ChatCompletionChunk) -> UsageMetadata | None: if response.usage is None: return None - return { - "input_tokens": response.usage.prompt_tokens, - "output_tokens": response.usage.completion_tokens, - "total_tokens": response.usage.total_tokens, - } + return UsageMetadata( + input_tokens=response.usage.prompt_tokens, + output_tokens=response.usage.completion_tokens, + total_tokens=response.usage.total_tokens, + ) def _chunk_to_generation(chunk: ChatCompletionChunk, first: bool) -> ChatGenerationChunk | None: - """Map one SSE chunk to a ``ChatGenerationChunk``; ``None`` for chunks with no choices - (e.g. the terminal usage-only chunk). Stream metadata rides on the first chunk.""" + """Map one SSE chunk to a ``ChatGenerationChunk``. The terminal usage-only chunk + becomes an empty message carrying ``usage_metadata`` (so it survives chunk + aggregation); a chunk with neither choices nor usage yields ``None``. Stream + metadata rides on the first chunk.""" if not chunk.choices: - return None + usage = _usage_metadata(chunk) + if usage is None: + return None + return ChatGenerationChunk(message=AIMessageChunk(content="", usage_metadata=usage)) choice = chunk.choices[0] metadata = {k: v for k, v in (("trace_id", chunk.trace_id), ("provider", chunk.provider)) if v} ai_chunk = AIMessageChunk( diff --git a/integrations/langchain-ferrolabsai/tests/test_chat_models.py b/integrations/langchain-ferrolabsai/tests/test_chat_models.py index 0a2afb1..4704af8 100644 --- a/integrations/langchain-ferrolabsai/tests/test_chat_models.py +++ b/integrations/langchain-ferrolabsai/tests/test_chat_models.py @@ -154,6 +154,33 @@ def test_stream_yields_chunks_with_trace_id(self, httpx_mock: HTTPXMock): assert "".join(c.content for c in chunks) == "Hello" assert chunks[0].response_metadata["trace_id"] == TRACE_ID + def test_stream_terminal_usage_chunk_survives_aggregation(self, httpx_mock: HTTPXMock): + usage_frame = { + "id": "1", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o", + "choices": [], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + body = sse_chunks("Hel", "lo").replace( + b"data: [DONE]", f"data: {json.dumps(usage_frame)}\n\ndata: [DONE]".encode() + ) + httpx_mock.add_response( + method="POST", + url=CHAT_URL, + content=body, + headers={"Content-Type": "text/event-stream", **GATEWAY_HEADERS}, + ) + chunks = list(_build_chat().stream([HumanMessage(content="hi")])) + total = chunks[0] + for chunk in chunks[1:]: + total = total + chunk + assert total.content == "Hello" + assert total.usage_metadata is not None + assert total.usage_metadata["total_tokens"] == 8 + assert total.response_metadata["trace_id"] == TRACE_ID + def test_stream_yields_tool_call_chunks(self, httpx_mock: HTTPXMock): frames = [ { From 7634300794b270dac65df4ef2c1ea84b3de2cec2 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:13:47 +0530 Subject: [PATCH 12/14] docs(langchain): state the gateway compatibility and response_metadata claims precisely Compatibility now reads "requires ai-gateway >= v1.4.0; contract-tested against v1.4.5". The {model, id, trace_id, provider, gateway_overhead_ms} set is described as the gateway-derived fields on response_metadata rather than the complete map, since LangChain adds finish_reason and friends. The 0.2.0 changelog also notes streamed usage_metadata. --- integrations/langchain-ferrolabsai/CHANGELOG.md | 14 +++++++++----- integrations/langchain-ferrolabsai/README.md | 7 ++++--- .../langchain_ferrolabsai/__init__.py | 11 ++++++----- .../langchain_ferrolabsai/chat_models.py | 8 ++++---- 4 files changed, 23 insertions(+), 17 deletions(-) diff --git a/integrations/langchain-ferrolabsai/CHANGELOG.md b/integrations/langchain-ferrolabsai/CHANGELOG.md index ead8372..7c3ef6b 100644 --- a/integrations/langchain-ferrolabsai/CHANGELOG.md +++ b/integrations/langchain-ferrolabsai/CHANGELOG.md @@ -16,13 +16,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [0.2.0] — 2026-08-29 -Requires `ferrolabsai >= 0.3.0` (the "truth release") and ai-gateway ≥ v1.4.0. +Requires `ferrolabsai >= 0.3.0` (the "truth release") and ai-gateway ≥ v1.4.0 +(contract-tested against v1.4.5). ### Breaking -- `response_metadata` now contains exactly what the gateway provides: - `model`, `id`, `trace_id` (`X-Request-ID` header), `provider` (body field), - `gateway_overhead_ms` (`X-Gateway-Overhead-Ms` header). `latency_ms`, +- The gateway-derived fields on `response_metadata` are now exactly what the + gateway provides: `model`, `id`, `trace_id` (`X-Request-ID` header), + `provider` (body field), `gateway_overhead_ms` (`X-Gateway-Overhead-Ms` + header); LangChain adds its own (e.g. `finish_reason`). `latency_ms`, `cost_usd`, and `cache_hit` are gone — the gateway never emitted them, so they were always absent. - Removed the `route_tag`, `template_id`, and `template_variables` fields from @@ -36,7 +38,9 @@ Requires `ferrolabsai >= 0.3.0` (the "truth release") and ai-gateway ≥ v1.4.0. - `FerroChatModel.with_structured_output(schema, include_raw=False)` via the OpenAI-style `response_format={"type": "json_schema", ...}` path; accepts a pydantic model class or a JSON-schema dict. -- Streaming: the first chunk carries `trace_id` in `response_metadata`. +- Streaming: the first chunk carries `trace_id` in `response_metadata`, and the + terminal usage-only chunk surfaces `usage_metadata` (send + `stream_options={"include_usage": True}`). - Python 3.13 classifier. --- diff --git a/integrations/langchain-ferrolabsai/README.md b/integrations/langchain-ferrolabsai/README.md index b4d70bb..a7b9549 100644 --- a/integrations/langchain-ferrolabsai/README.md +++ b/integrations/langchain-ferrolabsai/README.md @@ -5,7 +5,7 @@ LangChain integration for [Ferro Labs AI Gateway](https://github.com/ferro-labs/ai-gateway) — route LangChain chat, streaming, tool-calling, structured-output, and embedding workloads across **30 LLM providers** through a single OpenAI-compatible endpoint, with automatic fallback, load balancing, budgets, and observability. -Compatibility: `langchain-ferrolabsai 0.2.x` ↔ `ferrolabsai ≥ 0.3.0` ↔ `ai-gateway ≥ v1.4.0`; `langchain-core ≥ 0.3` (tested on 1.x). +Compatibility: `langchain-ferrolabsai 0.2.x` ↔ `ferrolabsai ≥ 0.3.0`; requires `ai-gateway ≥ v1.4.0` (contract-tested against `v1.4.5`); `langchain-core ≥ 0.3` (tested on 1.x). --- @@ -36,8 +36,9 @@ print(response.response_metadata["trace_id"]) # gateway X-Request-ID print(response.response_metadata.get("gateway_overhead_ms")) # gateway's own overhead ``` -`response_metadata` contains exactly what the gateway provides — `model`, `id`, -`trace_id`, `provider`, `gateway_overhead_ms` — with absent values stripped. +The gateway-derived fields on `response_metadata` are `model`, `id`, +`trace_id`, `provider`, and `gateway_overhead_ms`, with absent values stripped +(LangChain adds its own, such as `finish_reason`). `trace_id` is the join key for `client.admin.logs.list()` and for the gateway's observability exporters (LangSmith, Langfuse, Phoenix, …). diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py index 9242706..9bacfb5 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/__init__.py @@ -8,11 +8,12 @@ embed = FerroEmbeddings(model="text-embedding-3-small", api_key="sk-ferro-...") legacy = FerroLLM(model="gpt-4o", api_key="sk-ferro-...") -All three classes route through a Ferro Labs AI Gateway endpoint (ai-gateway -≥ v1.4.0). Chat responses expose the gateway's ``trace_id`` (the ``X-Request-ID`` -response header), ``provider`` and ``gateway_overhead_ms`` via -``response_metadata`` — ``trace_id`` is the join key for the gateway's request -log and its observability exporters (LangSmith, Langfuse, Phoenix, …). +All three classes route through a Ferro Labs AI Gateway endpoint (requires +ai-gateway ≥ v1.4.0; contract-tested against v1.4.5). Chat responses expose +the gateway's ``trace_id`` (the ``X-Request-ID`` response header), +``provider`` and ``gateway_overhead_ms`` via ``response_metadata`` — +``trace_id`` is the join key for the gateway's request log and its +observability exporters (LangSmith, Langfuse, Phoenix, …). """ from __future__ import annotations diff --git a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py index e3a699d..56d24f2 100644 --- a/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py +++ b/integrations/langchain-ferrolabsai/langchain_ferrolabsai/chat_models.py @@ -4,11 +4,11 @@ providers by name (e.g. ``"gpt-4o"``, ``"claude-3-5-sonnet-20241022"``, ``"gemini-2.5-flash"``) without changing the model class. -``response_metadata`` carries exactly what the gateway provides: ``model``, -``id``, ``trace_id`` (the ``X-Request-ID`` response header — the join key for -the gateway's request log and observability exporters), ``provider`` (body +The gateway-derived fields on ``response_metadata`` are ``model``, ``id``, +``trace_id`` (the ``X-Request-ID`` response header — the join key for the +gateway's request log and observability exporters), ``provider`` (body field), and ``gateway_overhead_ms`` (``X-Gateway-Overhead-Ms`` header). -Absent values are stripped. +Absent values are stripped; LangChain adds its own (e.g. ``finish_reason``). """ from __future__ import annotations From 483cb3e0c275bbe9905902bea4c34a30094bb044 Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:13:47 +0530 Subject: [PATCH 13/14] ci: read-only default token, no persisted credentials, SHA-pinned gateway leg - top-level permissions: contents: read (publish keeps its job-level grant) - persist-credentials: false on every actions/checkout step - the contract matrix pins ai-gateway to e8e4e26ddbd1dcf734722d82f02fabb50ce50037 (the commit behind v1.4.5) via a matrix include with a label so the check name stays "Contract vs AI Gateway v1.4.5"; the main leg keeps continue-on-error --- .github/workflows/ci.yml | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index be5735e..65d04fc 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -9,6 +9,9 @@ on: pull_request: branches: [main, development] +permissions: + contents: read + jobs: test: name: Test (Python ${{ matrix.python-version }}) @@ -20,6 +23,8 @@ jobs: steps: - uses: actions/checkout@v4 + with: + persist-credentials: false - name: Set up Python ${{ matrix.python-version }} uses: actions/setup-python@v5 @@ -40,7 +45,7 @@ jobs: run: pytest tests/ -v --tb=short contract: - name: Contract vs AI Gateway ${{ matrix.gateway_ref }} + name: Contract vs AI Gateway ${{ matrix.label }} runs-on: ubuntu-latest # The pinned leg is the required signal; "main" moves upstream of this # repo, so a failure there must only report drift, never block a merge or @@ -49,7 +54,15 @@ jobs: strategy: fail-fast: false matrix: - gateway_ref: ["v1.4.5", "main"] + include: + # The pin is the commit SHA behind the tag so a moved tag cannot change + # what a release is verified against; `label` keeps the check name + # stable ("Contract vs AI Gateway v1.4.5"), which the publish job and + # the branch protection rules reference. + - gateway_ref: e8e4e26ddbd1dcf734722d82f02fabb50ce50037 # v1.4.5 + label: v1.4.5 + - gateway_ref: main + label: main steps: - name: Check out ferrolabs-python-sdk uses: actions/checkout@v4 @@ -100,6 +113,8 @@ jobs: steps: - uses: actions/checkout@v4 + with: + persist-credentials: false - name: Set up Python uses: actions/setup-python@v5 From 4aed1d5102014ed3bf0eb22800d93f05c93db94c Mon Sep 17 00:00:00 2001 From: Mitul Shah Date: Sat, 29 Aug 2026 14:14:22 +0530 Subject: [PATCH 14/14] chore(llama-index): apply ruff format to the placeholder package Pre-existing drift; both the root and the sub-package ruff configs use line-length 100, so this is a no-op for the sub-package's own lint job. --- integrations/llama-index-llms-ferrolabsai/README.md | 2 +- .../llama_index/llms/ferrolabsai/__init__.py | 4 +--- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/integrations/llama-index-llms-ferrolabsai/README.md b/integrations/llama-index-llms-ferrolabsai/README.md index 7afc32c..c373a50 100644 --- a/integrations/llama-index-llms-ferrolabsai/README.md +++ b/integrations/llama-index-llms-ferrolabsai/README.md @@ -22,7 +22,7 @@ from llama_index.llms.ferrolabsai import FerroLabsAI llm = FerroLabsAI( model="gpt-4o", - base_url="http://localhost:8080", # any Ferro Labs AI Gateway instance + base_url="http://localhost:8080", # any Ferro Labs AI Gateway instance api_key="sk-ferro-...", ) diff --git a/integrations/llama-index-llms-ferrolabsai/llama_index/llms/ferrolabsai/__init__.py b/integrations/llama-index-llms-ferrolabsai/llama_index/llms/ferrolabsai/__init__.py index 8110fcc..603937c 100644 --- a/integrations/llama-index-llms-ferrolabsai/llama_index/llms/ferrolabsai/__init__.py +++ b/integrations/llama-index-llms-ferrolabsai/llama_index/llms/ferrolabsai/__init__.py @@ -28,6 +28,4 @@ def _not_implemented(name: str) -> None: def __getattr__(name: str) -> object: if name == "FerroLabsAI": _not_implemented(name) - raise AttributeError( - f"module 'llama_index.llms.ferrolabsai' has no attribute {name!r}" - ) + raise AttributeError(f"module 'llama_index.llms.ferrolabsai' has no attribute {name!r}")