From cd1b781f5b0a91b78c688fb24565bc5d263b6902 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Sun, 20 Sep 2026 23:00:40 -0300 Subject: [PATCH 01/11] feat(providers): add FriendliAI provider (Model APIs, Dedicated, Container) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Implement FriendliProvider covering all three FriendliAI inference surfaces from one [providers.friendli] section, selected with endpoint_type: "serverless" (Model APIs, the hosted pay-per-token catalog), "dedicated" (the model field is the endpoint ID, or ID:ADAPTER_ROUTE for Multi-LoRA), and "container" (self-hosted Friendli Engine; base_url required, API key optional). The chat endpoint is OpenAI-compatible, plus the Friendli extensions: reasoning controls (reasoning_effort incl. the Friendli-only "ultracode" tier, reasoning_budget, parse_reasoning, include_reasoning), the chat-template switches enable_thinking / clear_thinking folded into chat_template_kwargs, Friendli Engine sampling (top_k, min_p, min_tokens, repetition_penalty, eos_token, XTC), regex-constrained structured output, cache-aware usage, and the exact /tokenize endpoint. Mutually exclusive body fields (tools vs min_tokens/response_format) are dropped with a warning instead of 422-ing. Transport is selectable via `backend` and auto-resolves openai -> httpx -> sdk. The vendor `friendli` SDK is supported but ranked last on purpose: its generated response models ignore unknown fields, so reasoning_content and reasoning are silently dropped, and it offers no extra_body escape hatch. The provider warns at startup when backend="sdk" meets parse_reasoning, and filters kwargs the SDK cannot type rather than surfacing a TypeError from inside the vendor package. Beyond chat: rich catalog discovery (context, pricing, modalities, reasoning options) cached and primed by warm_up(), tokenize/detokenize/ render_chat, text_completion, transcribe_audio, and — gated to dedicated/container — create_embeddings and generate_image. get_team_cost() and get_team_usage() read the Friendli Suite billing APIs for the configured team, which is also sent as X-Friendli-Team on every request. Token counting stays local (tiktoken) by default: llmcore counts tokens every turn and Model APIs rate limits are tier-based, so exact /tokenize counts are opt-in via native_token_count. Register the provider in ProviderManager with friendliai / friendli_ai aliases, add the llmcore[friendli] extra, the [providers.friendli] config section, and the provider_friendli confy schema section. Validated live against api.friendli.ai: catalog discovery, exact tokenization, chat on all three backends, SSE streaming with reasoning deltas, tool-call round trip, team cost/usage, and 401/429 mapping. Co-Authored-By: Claude Opus 5 (1M context) --- .github/workflows/ci.yml | 15 +- pyproject.toml | 8 + src/llmcore/config/default_config.toml | 122 ++ src/llmcore/providers/friendli_provider.py | 2118 ++++++++++++++++++++ src/llmcore/providers/manager.py | 8 + tests/providers/test_friendli_provider.py | 1208 +++++++++++ tools/llmcore.confy-schema.json | 245 +++ 7 files changed, 3718 insertions(+), 6 deletions(-) create mode 100644 src/llmcore/providers/friendli_provider.py create mode 100644 tests/providers/test_friendli_provider.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f5bfb9ff..b64b9ce7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -59,12 +59,15 @@ jobs: # grimoire (the control plane) is also GitHub-only. pip install "grimoire @ git+https://github.com/araray/grimoire.git@v0.4.0" # dev pulls in test (pytest, pytest-asyncio, respx, ...) and sandbox - # deps; the extras below are [all] minus [zai]: provider SDKs - # (openai, anthropic, ...), chromadb, and the bridge deps - # (grpcio/protobuf/starlette) that unit tests import at collection - # or fixture time. zai-sdk is deliberately NOT installed: the zai - # unit tests mock the openai/httpx fallback transport, and the SDK - # transport would bypass those mocks and hit the live API. + # deps; the extras below are [all] minus [zai] and [friendli]: + # provider SDKs (openai, anthropic, ...), chromadb, and the bridge + # deps (grpcio/protobuf/starlette) that unit tests import at + # collection or fixture time. zai-sdk is deliberately NOT installed: + # the zai unit tests mock the openai/httpx fallback transport, and + # the SDK transport would bypass those mocks and hit the live API. + # The friendli SDK is likewise skipped: the provider auto-resolves to + # the openai backend (already installed) and its tests pin the + # backend explicitly, so the vendor SDK adds nothing offline. # numpy is used by embedding-shaped unit tests. pip install -e ".[dev,bridge,openai,anthropic,gemini,ollama,deepinfra,deepgram,typesafe,brightdata,serper,serpapi,semanticscholar,postgres,chromadb]" pip install numpy diff --git a/pyproject.toml b/pyproject.toml index 91d74c50..2c7ae2ab 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -70,6 +70,13 @@ deepgram = ["deepgram-sdk>=7.0.0", "websockets>=12.0"] # mode) and/or direct 'httpx' calls (native media endpoints: image/audio/OCR/ # video generation, web search). Install with: pip install llmcore[zai] zai = ["zai-sdk>=0.2.0", "openai>=2.31.0", "httpx>=0.27.0"] +# FriendliAI (Model APIs / Dedicated Endpoints / Container). The provider +# prefers direct transports — the 'openai' SDK in OpenAI-compatibility mode +# (default) or raw 'httpx' — because the official 'friendli' SDK's generated +# response models drop fields outside the published schema (notably +# reasoning_content). Install the vendor SDK too if you want backend = "sdk". +# Install with: pip install llmcore[friendli] +friendli = ["openai>=2.31.0", "httpx>=0.27.0", "friendli>=0.15.1"] # TypeSafe.ai (System One typed judgments: noul/choice/score questions with # calibrated probabilities). NOT a chat API. The provider talks to the two REST # endpoints (/v1/systemone, /v1/models) directly over httpx — no vendor SDK is @@ -179,6 +186,7 @@ all = [ "llmcore[deepinfra]", "llmcore[deepgram]", "llmcore[zai]", + "llmcore[friendli]", "llmcore[typesafe]", "llmcore[brightdata]", "llmcore[serper]", diff --git a/src/llmcore/config/default_config.toml b/src/llmcore/config/default_config.toml index 64eeab86..1291376d 100644 --- a/src/llmcore/config/default_config.toml +++ b/src/llmcore/config/default_config.toml @@ -501,6 +501,128 @@ max_retries = 2 # Override the endpoint chosen by `region`. # base_url = "https://api.z.ai/api/paas/v4" + # --- FriendliAI Provider --- + # Native provider for FriendliAI's three inference surfaces: + # "serverless" — Friendli Model APIs, the hosted pay-per-token catalog + # (zai-org/GLM-5.3, deepseek-ai/DeepSeek-V3.2, + # google/gemma-4-31B-it, MiniMaxAI/MiniMax-M2.5, ...) + # "dedicated" — your own GPU deployments; `model` is the ENDPOINT ID + # (or "ENDPOINT_ID:ADAPTER_ROUTE" for Multi-LoRA) + # "container" — a self-hosted Friendli Engine; `base_url` is required + # + # The chat endpoint is OpenAI-compatible, plus Friendli extensions: + # reasoning_effort / reasoning_budget / parse_reasoning / include_reasoning, + # chat_template_kwargs (enable_thinking, clear_thinking), Friendli Engine + # sampling (top_k, min_p, min_tokens, repetition_penalty, eos_token, XTC), + # regex-constrained structured output, and an exact /tokenize endpoint. + # + # Transport is selectable via `backend`: + # "openai" — openai SDK pointed at the Friendli base URL (DEFAULT) + # "httpx" — direct REST calls + # "sdk" — the official `friendli` Python SDK + # Omit/"auto" to auto-detect in order openai -> httpx -> sdk. The vendor SDK + # is last on purpose: its generated response models drop fields outside the + # published schema, so `reasoning_content` is lost on that backend and there + # is no extra_body escape hatch. + # Install with: pip install llmcore[friendli] + # + # Docs: https://friendli.ai/docs/llms.txt + # Pricing: https://friendli.ai/docs/guides/model-apis/pricing + [providers.friendli] + # type = "friendli" # auto-detected from section name + # Aliases that map to this provider class: friendliai, friendli_ai. + + # API Key: Set via environment variable for security. A Friendli Personal + # API key looks like "flp_...". The provider checks, in order: + # api_key -> api_key_env_var -> FRIENDLI_TOKEN -> FRIENDLIAI_API_KEY + # -> FRIENDLI_API_KEY + # (FRIENDLI_TOKEN is the official SDK's variable; FRIENDLIAI_API_KEY is the + # spelling used in friendli.ai's own documentation examples.) + # Not required when endpoint_type = "container" and the container has no auth. + # api_key = "flp_..." + # api_key_env_var = "FRIENDLI_TOKEN" + + # Team scoping: sent as the X-Friendli-Team header on every request and used + # by get_team_cost() / get_team_usage(). Falls back to FRIENDLI_TEAM_ID then + # FRIENDLIAI_TEAM_ID. Omit to use the default team in Friendli Suite. + # team_id = "..." + # team_id_env_var = "FRIENDLI_TEAM_ID" + + # Which Friendli surface to talk to. + endpoint_type = "serverless" # "serverless" | "dedicated" | "container" + + # Transport backend: "openai" (default) | "httpx" | "sdk". Omit for auto. + # backend = "openai" + + # Default model. On Model APIs this is a catalog model ID; on Dedicated + # Endpoints it must be the endpoint ID. + # Current Model APIs catalog (September 2026): + # zai-org/GLM-5.3 — flagship; 1M context, reasoning + tools + # zai-org/GLM-5.3-Flash — fast + cheap; 1M context, text/image/video in + # zai-org/GLM-5.2 — 1M context, reasoning toggle + # zai-org/GLM-5.1 — 198K context + # google/gemma-4-31B-it — 256K context, text/image in + # deepseek-ai/DeepSeek-V3.2 — 160K context + # MiniMaxAI/MiniMax-M2.5 — 192K context + default_model = "zai-org/GLM-5.3" + + # Request timeout in seconds. 1M-context reasoning responses can be slow. + timeout = 300 + + # --- Base URLs --- + # Inference root. Defaults per endpoint_type: + # serverless -> https://api.friendli.ai/serverless/v1 + # dedicated -> https://api.friendli.ai/dedicated/v1 + # container -> (no default; REQUIRED, e.g. http://localhost:8000/v1) + # base_url = "https://api.friendli.ai/serverless/v1" + # + # Friendli Suite root used by get_team_cost() / get_team_usage(). + # suite_base_url = "https://api.friendli.ai/v1" + + # --- Reasoning controls --- + # Friendli splits reasoning control across four request-body fields plus the + # chat template. Per-request overrides go through chat_completion(...). + # + # reasoning_effort: minimal | low | medium | high | xhigh | max | ultracode + # The tiers a model accepts are advertised in its /models entry + # (reasoning_options); unset leaves the model's own default in place. + # reasoning_effort = "high" + # + # reasoning_budget: hard cap (in tokens) on the chain of thought. + # reasoning_budget = 10000 + # + # parse_reasoning: split the chain of thought out of `content` into + # `reasoning_content` (surfaced by extract_reasoning_content() and + # extract_delta_reasoning_content()). Requires backend "openai"/"httpx". + parse_reasoning = true + # + # include_reasoning: when parsing is on, include the parsed reasoning in the + # response (Friendli defaults this to true). + # include_reasoning = true + # + # enable_thinking: default chat_template_kwargs.enable_thinking for + # *controllable* reasoning models (e.g. zai-org/GLM-5.2). Unset leaves it + # to the model's template. `clear_thinking` is accepted per-request. + # enable_thinking = true + + # --- Token counting --- + # false (default): count locally with tiktoken cl100k_base. + # true: use the model's own tokenizer via POST /tokenize — exact, but one API + # request per count, which consumes the Model APIs rate-limit budget. + # provider.tokenize() / detokenize() are always available regardless. + native_token_count = false + + # Context window used when neither live discovery nor a model card knows the + # model (Dedicated Endpoints and Containers have no catalog). + fallback_context_length = 131072 + + # --- Endpoint-type notes --- + # Embeddings (create_embeddings) and image generation (generate_image) are + # served by Dedicated Endpoints / Container only — Model APIs has neither. + # Audio transcription (transcribe_audio) works on all three surfaces. + # /detokenize and /chat/render are documented but currently return 404 on + # Model APIs; they are available on Dedicated Endpoints / Container. + # --- Ollama Provider --- # For locally running models via Ollama. [providers.ollama] diff --git a/src/llmcore/providers/friendli_provider.py b/src/llmcore/providers/friendli_provider.py new file mode 100644 index 00000000..2bdd04fc --- /dev/null +++ b/src/llmcore/providers/friendli_provider.py @@ -0,0 +1,2118 @@ +# src/llmcore/providers/friendli_provider.py +""" +FriendliAI provider implementation for the LLMCore library. + +Handles interactions with all three FriendliAI inference surfaces: + +- **Model APIs** (``serverless``) — the hosted, pay-per-token catalog + (``zai-org/GLM-5.3``, ``deepseek-ai/DeepSeek-V3.2``, ``google/gemma-4-31B-it``, + ``MiniMaxAI/MiniMax-M2.5``, …) at ``https://api.friendli.ai/serverless/v1``. +- **Dedicated Endpoints** (``dedicated``) — your own GPU deployments at + ``https://api.friendli.ai/dedicated/v1``; the ``model`` field is the + *endpoint ID* (optionally ``ENDPOINT_ID:ADAPTER_ROUTE`` for Multi-LoRA). +- **Friendli Container** (``container``) — a self-hosted Friendli Engine; set + ``base_url`` to your container (e.g. ``http://localhost:8000/v1``). + +The chat endpoint is OpenAI-compatible (``POST /chat/completions``) with a set +of Friendli-specific extensions handled here: + +- **Reasoning controls** — ``reasoning_effort`` + (``minimal|low|medium|high|xhigh|max|ultracode``), ``reasoning_budget`` + (token cap on the chain of thought), ``parse_reasoning`` (split the chain of + thought out of ``content`` into ``reasoning_content``), and + ``include_reasoning``. +- **Chat-template kwargs** — ``chat_template_kwargs`` carries the per-model + template switches documented by FriendliAI, notably ``enable_thinking`` + (controllable reasoning models) and ``clear_thinking``. Both are also + accepted as flat kwargs and folded into ``chat_template_kwargs`` for you. +- **Friendli Engine sampling** — ``top_k``, ``min_p``, ``min_tokens``, + ``repetition_penalty``, ``eos_token``, and XTC sampling + (``xtc_threshold`` / ``xtc_probability``). +- **Structured output** — ``response_format`` of ``json_schema``, + ``json_object``, ``regex`` (Friendli-specific), or ``text``. +- **Exact token counting** via the native ``POST /tokenize`` endpoint, with a + tiktoken/heuristic fallback; ``POST /detokenize`` and ``POST /chat/render`` + are exposed as auxiliary helpers. +- **Cache-aware usage** — ``usage.prompt_tokens_details.cached_tokens``. +- **Team scoping** — every request carries the ``X-Friendli-Team`` header when + a team ID is configured, and :meth:`get_team_cost` / :meth:`get_team_usage` + read the Friendli Suite billing APIs for that team. + +Transport (selectable via the ``backend`` config key) +----------------------------------------------------- + +- ``"openai"`` — the ``openai`` SDK (``AsyncOpenAI``) pointed at the Friendli + base URL. Native async, full SSE handling; **the default**. +- ``"httpx"`` — direct async REST calls against the documented endpoints. +- ``"sdk"`` — the official ``friendli`` Python SDK (``AsyncFriendli``). + +The backend governs **chat completions**. Endpoints outside the OpenAI-compatible +surface — the model catalog, ``/tokenize``, ``/detokenize``, ``/chat/render``, +``/completions``, ``/embeddings``, ``/images/generations``, +``/audio/transcriptions`` and the Suite billing reads — always travel over the +provider's own ``httpx`` client (or the vendor SDK when ``backend = "sdk"``), +because the ``openai`` SDK either does not model them or would discard the extra +fields Friendli returns. + +When unset, the backend auto-resolves ``openai`` → ``httpx`` → ``sdk`` based on +which libraries are installed. The vendor SDK is deliberately *last*: its +generated response models are strict (``extra`` is ignored), so fields Friendli +adds outside the published schema — including ``reasoning_content`` and +``reasoning`` on assistant messages — are silently dropped, and it offers no +``extra_body`` escape hatch. Request ``backend = "sdk"`` explicitly if you want +it anyway; this provider logs a warning when reasoning parsing is combined with +the SDK backend. + +References: + - https://friendli.ai/docs/llms.txt (documentation index) + - https://friendli.ai/docs/guides/openai-compatibility + - https://friendli.ai/docs/openapi/model-apis/chat-completions + - https://friendli.ai/docs/guides/capabilities/reasoning + - https://friendli.ai/docs/guides/structured-outputs + - Friendli Python SDK (``friendli``) +""" + +from __future__ import annotations + +import asyncio +import inspect +import json +import logging +import os +from collections.abc import AsyncGenerator +from typing import Any, Literal + +# --- Optional official Friendli SDK (``friendli``) --- +# Native async client. NOTE: its generated response models drop unknown +# fields, so reasoning_content is lost on this backend (see module docstring). +try: + from friendli import AsyncFriendli + from friendli.models import FriendliCoreError + + friendli_sdk_available = True +except ImportError: + friendli_sdk_available = False + AsyncFriendli = None # type: ignore + FriendliCoreError = Exception # type: ignore + +# --- Optional OpenAI SDK (OpenAI-compatibility backend) --- +try: + from openai import AsyncOpenAI + from openai._exceptions import ( + APIConnectionError as OpenAIAPIConnectionError, + ) + from openai._exceptions import ( + APIError as OpenAIAPIError, + ) + from openai._exceptions import ( + APIStatusError as OpenAIAPIStatusError, + ) + from openai._exceptions import ( + APITimeoutError as OpenAIAPITimeoutError, + ) + from openai._exceptions import ( + OpenAIError, + ) + + openai_available = True +except ImportError: + openai_available = False + AsyncOpenAI = None # type: ignore + OpenAIError = Exception # type: ignore + OpenAIAPIError = Exception # type: ignore + OpenAIAPIStatusError = Exception # type: ignore + OpenAIAPIConnectionError = Exception # type: ignore + OpenAIAPITimeoutError = Exception # type: ignore + +try: + import httpx + + httpx_available = True +except ImportError: # pragma: no cover - httpx is a hard dep of openai + httpx_available = False + httpx = None # type: ignore + +try: + import tiktoken + + tiktoken_available = True +except ImportError: + tiktoken_available = False + tiktoken = None # type: ignore + +from ..exceptions import ConfigError, ContextLengthError, ProviderError +from ..model_cards.registry import get_model_card_registry +from ..models import Message, ModelDetails, Tool, ToolCall +from ..models import Role as LLMCoreRole +from ..models_multimodal import ( + GeneratedImage, + ImageGenerationResult, + TranscriptionResult, +) +from ..tokens import EstimateCounter as _EstimateCounter +from .base import BaseProvider, ContextPayload + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +#: Inference base URLs per endpoint type. ``container`` has no default — it is +#: self-hosted, so ``base_url`` is mandatory there. +_BASE_URLS: dict[str, str] = { + "serverless": "https://api.friendli.ai/serverless/v1", + "dedicated": "https://api.friendli.ai/dedicated/v1", +} + +#: Friendli Suite (team billing/usage) API root — distinct from inference. +_SUITE_BASE_URL = "https://api.friendli.ai/v1" + +#: Default model when none is configured (Model APIs flagship). +_DEFAULT_MODEL = "zai-org/GLM-5.3" + +#: Endpoint types this provider understands. +EndpointType = Literal["serverless", "dedicated", "container"] + +#: Transport backends, in auto-resolution preference order. +_BACKEND_ORDER: tuple[str, ...] = ("openai", "httpx", "sdk") + +#: Reasoning-effort tiers accepted by the Friendli chat API. The tiers a given +#: model actually supports are advertised per-model in ``/models`` +#: (``reasoning_options``); unsupported tiers are rejected by the server. +_VALID_EFFORTS: frozenset[str] = frozenset( + {"minimal", "low", "medium", "high", "xhigh", "max", "ultracode"} +) + +#: Static context-length fallback for Model APIs models, used only when live +#: discovery and the model-card registry are both unavailable. +_CONTEXT_LENGTHS: dict[str, int] = { + "zai-org/GLM-5.3": 1_048_576, + "zai-org/GLM-5.3-Flash": 1_048_576, + "zai-org/GLM-5.2": 1_048_576, + "zai-org/GLM-5.1": 202_752, + "google/gemma-4-31B-it": 262_144, + "deepseek-ai/DeepSeek-V3.2": 163_840, + "MiniMaxAI/MiniMax-M2.5": 196_608, +} + +#: Conservative context length when nothing else is known. +_FALLBACK_CONTEXT_LENGTH = 131_072 + +#: Environment variables checked (in order) for the Friendli API key. +#: ``FRIENDLI_TOKEN`` is the official SDK convention; ``FRIENDLIAI_API_KEY`` is +#: the spelling used throughout friendli.ai's own documentation examples. +_API_KEY_ENV_VARS: tuple[str, ...] = ( + "FRIENDLI_TOKEN", + "FRIENDLIAI_API_KEY", + "FRIENDLI_API_KEY", +) + +#: Environment variables checked (in order) for the Friendli team ID. +_TEAM_ID_ENV_VARS: tuple[str, ...] = ("FRIENDLI_TEAM_ID", "FRIENDLIAI_TEAM_ID") + +#: Placeholder key for self-hosted containers launched without auth (the +#: ``openai`` SDK refuses an empty API key). +_CONTAINER_PLACEHOLDER_KEY = "EMPTY" + +#: Request-body keys that are Friendli-specific, i.e. not native ``openai`` +#: chat-completion parameters. They travel via ``extra_body`` on the openai +#: backend and as plain body keys everywhere else. +_EXTRA_BODY_KEYS: frozenset[str] = frozenset( + { + "chat_template_kwargs", + "eos_token", + "include_reasoning", + "min_p", + "min_tokens", + "parse_reasoning", + "reasoning_budget", + "reasoning_effort", + "repetition_penalty", + "top_k", + "xtc_probability", + "xtc_threshold", + } +) + +#: Flat kwargs folded into ``chat_template_kwargs`` (Friendli passes these to +#: the model's chat template rather than treating them as body parameters). +_TEMPLATE_KWARG_KEYS: frozenset[str] = frozenset({"enable_thinking", "clear_thinking"}) + + +def _looks_like_context_overflow(message: str) -> bool: + """Heuristic: does *message* describe a context/length overflow?""" + low = message.lower() + if "max_tokens" in low and "exceed" in low: + return True + return ("context" in low or "prompt" in low or "input" in low) and ( + "too long" in low or "length" in low or "exceed" in low + ) + + +class FriendliProvider(BaseProvider): + """First-class FriendliAI provider (Model APIs / Dedicated / Container). + + Configuration keys (under ``[providers.friendli]``): + + - ``api_key`` / ``api_key_env_var`` — Friendli Personal API key (``flp_…``). + Resolved from the config, then ``api_key_env_var``, then + ``FRIENDLI_TOKEN`` / ``FRIENDLIAI_API_KEY`` / ``FRIENDLI_API_KEY``. + Optional when ``endpoint_type = "container"``. + - ``team_id`` / ``team_id_env_var`` — team to run requests as, sent as the + ``X-Friendli-Team`` header. Falls back to ``FRIENDLI_TEAM_ID`` / + ``FRIENDLIAI_TEAM_ID``. + - ``endpoint_type`` — ``"serverless"`` (default), ``"dedicated"``, or + ``"container"``. + - ``base_url`` — override the inference root. Required for ``container``. + - ``suite_base_url`` — override the Suite billing root + (default ``https://api.friendli.ai/v1``). + - ``backend`` — ``"openai"`` (default), ``"httpx"``, or ``"sdk"``. Omit or + use ``"auto"`` to resolve openai → httpx → sdk. + - ``default_model`` — default model ID (Model APIs) or endpoint ID + (Dedicated). Default: ``zai-org/GLM-5.3``. + - ``timeout`` — HTTP request timeout in seconds (default: 300). + - ``reasoning_effort`` — default effort tier, or unset to leave it to the + model's own default. + - ``reasoning_budget`` — default cap on reasoning tokens. + - ``parse_reasoning`` — split reasoning into ``reasoning_content`` + (default: ``true``). + - ``include_reasoning`` — include parsed reasoning in the response. + - ``enable_thinking`` — default ``chat_template_kwargs.enable_thinking`` + for controllable reasoning models. + - ``native_token_count`` — count tokens with the model's own tokenizer via + ``POST /tokenize`` instead of locally (default: ``false``; each count is + an extra API request). + - ``fallback_context_length`` — context window used when neither live + discovery nor a model card knows the model (default: 131072). + """ + + default_model: str + _backend: str # "openai" | "httpx" | "sdk" + _endpoint_type: str + _api_key: str + _base_url: str + _suite_base_url: str + _team_id: str | None + _timeout: float + _client: Any # AsyncOpenAI | None + _sdk_client: Any # AsyncFriendli | None + _http: Any # httpx.AsyncClient | None + _encoding: Any # tiktoken.Encoding | None + _catalog: dict[str, dict[str, Any]] | None + + def __init__(self, config: dict[str, Any], log_raw_payloads: bool = False): + """Initialize the FriendliAI provider. + + Args: + config: Provider configuration dict from ``[providers.friendli]``. + log_raw_payloads: Whether to log raw request/response payloads. + + Raises: + ConfigError: If no transport library is installed, no API key is + available for a hosted endpoint type, or ``base_url`` is + missing for ``endpoint_type = "container"``. + """ + super().__init__(config, log_raw_payloads) + + if not (openai_available or httpx_available or friendli_sdk_available): + raise ConfigError( + "The Friendli provider requires one of: the 'openai' SDK " + "(compatibility mode, preferred), 'httpx' (direct REST), or the " + "official 'friendli' SDK. Install with: pip install llmcore[friendli]" + ) + + # --- Endpoint type --- + endpoint_type = str(config.get("endpoint_type", "serverless")).lower() + if endpoint_type not in ("serverless", "dedicated", "container"): + logger.warning( + "Invalid Friendli endpoint_type '%s'; defaulting to 'serverless'.", + endpoint_type, + ) + endpoint_type = "serverless" + self._endpoint_type = endpoint_type + + # --- Endpoint URL --- + base_url = config.get("base_url") + if not base_url: + if endpoint_type == "container": + raise ConfigError( + "Friendli Container requires an explicit base_url " + '(e.g. providers.friendli.base_url = "http://localhost:8000/v1").' + ) + base_url = _BASE_URLS[endpoint_type] + self._base_url = str(base_url).rstrip("/") + self._suite_base_url = str(config.get("suite_base_url", _SUITE_BASE_URL)).rstrip("/") + + # --- API key --- + api_key = self._resolve_api_key(config) + if not api_key: + if endpoint_type == "container": + # Containers are frequently launched without auth. + api_key = _CONTAINER_PLACEHOLDER_KEY + else: + raise ConfigError( + "Friendli API key not found. Set FRIENDLI_TOKEN (or " + "FRIENDLIAI_API_KEY) or configure providers.friendli.api_key / " + "api_key_env_var. Create a key at " + "https://friendli.ai/suite/~/setting/keys." + ) + self._api_key = api_key + + # --- Team scoping (X-Friendli-Team) --- + self._team_id = self._resolve_team_id(config) + + # --- Model / timeout --- + self.default_model = config.get("default_model", _DEFAULT_MODEL) + self._timeout = float(config.get("timeout", 300)) + self._fallback_context_length = int( + config.get("fallback_context_length", _FALLBACK_CONTEXT_LENGTH) + ) + # Exact per-model token counts cost one API round-trip each and consume + # the Model APIs request budget, so count_tokens()/count_message_tokens() + # stay local by default. tokenize()/detokenize() are always available. + self._native_token_count = bool(config.get("native_token_count", False)) + + # --- Reasoning defaults --- + effort_raw = config.get("reasoning_effort") + self._default_reasoning_effort: str | None = None + if effort_raw is not None: + effort = str(effort_raw).lower() + if effort in _VALID_EFFORTS: + self._default_reasoning_effort = effort + else: + logger.warning( + "Invalid Friendli reasoning_effort '%s'; leaving it to the model " + "default. Valid tiers: %s.", + effort_raw, + ", ".join(sorted(_VALID_EFFORTS)), + ) + budget_raw = config.get("reasoning_budget") + self._default_reasoning_budget: int | None = ( + int(budget_raw) if budget_raw is not None else None + ) + parse_raw = config.get("parse_reasoning", True) + self._default_parse_reasoning: bool | None = ( + bool(parse_raw) if parse_raw is not None else None + ) + include_raw = config.get("include_reasoning") + self._default_include_reasoning: bool | None = ( + bool(include_raw) if include_raw is not None else None + ) + thinking_raw = config.get("enable_thinking") + self._default_enable_thinking: bool | None = ( + bool(thinking_raw) if thinking_raw is not None else None + ) + + # --- Transport --- + self._client = None + self._sdk_client = None + self._http = None + self._catalog = None + self._backend = self._resolve_backend(config.get("backend")) + + if self._backend == "sdk" and self._default_parse_reasoning: + logger.warning( + "Friendli backend 'sdk' is active with parse_reasoning enabled: the " + "official SDK's response models drop reasoning_content. Use " + 'backend = "openai" or "httpx" to receive parsed reasoning.' + ) + + try: + if self._backend == "openai": + self._client = AsyncOpenAI( + api_key=self._api_key, + base_url=self._base_url, + timeout=self._timeout, + default_headers=self._team_headers(), + ) + elif self._backend == "sdk": + self._sdk_client = AsyncFriendli( + token=self._api_key, + server_url=self._sdk_server_url(), + timeout_ms=int(self._timeout * 1000), + x_friendli_team=self._team_id, + ) + # The "httpx" backend lazily creates its client on first use. + logger.debug( + "Friendli client initialized (backend=%s, endpoint_type=%s, " + "base_url=%s, default_model=%s, team=%s).", + self._backend, + self._endpoint_type, + self._base_url, + self.default_model, + self._team_id or "", + ) + except Exception as e: + raise ConfigError(f"Friendli client initialization failed: {e}") + + # --- Tokenizer fallback (the native /tokenize endpoint is preferred) --- + self._encoding = None + if tiktoken_available: + try: + self._encoding = tiktoken.get_encoding("cl100k_base") + except Exception as e: + logger.warning("Failed to load tiktoken for Friendli: %s", e) + + # ========================================================================= + # Configuration helpers + # ========================================================================= + + @staticmethod + def _resolve_api_key(config: dict[str, Any]) -> str | None: + """Resolve the Friendli API key from config or the environment.""" + api_key = config.get("api_key") + if api_key: + return str(api_key) + env_var = config.get("api_key_env_var") + if env_var: + value = os.environ.get(str(env_var)) + if value: + return value + for name in _API_KEY_ENV_VARS: + value = os.environ.get(name) + if value: + return value + return None + + @staticmethod + def _resolve_team_id(config: dict[str, Any]) -> str | None: + """Resolve the Friendli team ID from config or the environment.""" + team_id = config.get("team_id") + if team_id: + return str(team_id) + env_var = config.get("team_id_env_var") + if env_var: + value = os.environ.get(str(env_var)) + if value: + return value + for name in _TEAM_ID_ENV_VARS: + value = os.environ.get(name) + if value: + return value + return None + + @staticmethod + def _resolve_backend(requested: str | None) -> str: + """Resolve the transport backend, honoring library availability. + + Preference order when unset/``"auto"``: ``openai`` → ``httpx`` → ``sdk`` + (see the module docstring for why the vendor SDK is last). An + explicitly requested backend that is unavailable falls back through the + same chain with a warning. + """ + available = { + "openai": openai_available, + "httpx": httpx_available, + "sdk": friendli_sdk_available, + } + + req = (requested or "auto").lower() + if req not in ("auto", *_BACKEND_ORDER): + logger.warning("Unknown Friendli backend '%s'; using auto-detection.", req) + req = "auto" + + if req != "auto": + if available.get(req): + return req + logger.warning( + "Requested Friendli backend '%s' is unavailable; falling back. " + "Install with: pip install llmcore[friendli]", + req, + ) + + for backend in _BACKEND_ORDER: + if available[backend]: + if req != "auto" and backend != req: + logger.info("Friendli backend resolved to '%s'.", backend) + return backend + # Unreachable given the __init__ guard. + raise ConfigError("No usable Friendli transport backend is installed.") + + def _team_headers(self, extra: dict[str, str] | None = None) -> dict[str, str]: + """Return the default headers, including team scoping when configured.""" + headers: dict[str, str] = {} + if self._team_id: + headers["X-Friendli-Team"] = self._team_id + if extra: + headers.update(extra) + return headers + + def _sdk_server_url(self) -> str | None: + """Return the ``server_url`` the official SDK should target. + + The SDK's generated operations already carry the ``/serverless/v1`` and + ``/dedicated/v1`` path prefixes, so the hosted types need only the API + origin. Containers pass their own root through unchanged. + """ + if self._endpoint_type == "container": + return self._base_url + default_root = _BASE_URLS[self._endpoint_type] + if self._base_url == default_root: + return None # SDK default server (https://api.friendli.ai) + # A custom hosted root: strip the operation path prefix if present. + suffix = default_root.rsplit("/", 2)[-2:] # e.g. ["serverless", "v1"] + trimmed = self._base_url + for part in reversed(suffix): + if trimmed.endswith(f"/{part}"): + trimmed = trimmed[: -(len(part) + 1)] + return trimmed or None + + async def _run_sdk(self, fn: Any) -> Any: + """Await an SDK coroutine factory (kept for symmetry/testability).""" + return await fn() + + def _sdk_namespace(self, resource: str) -> Any: + """Return the SDK sub-client for *resource* on the active endpoint type. + + Args: + resource: ``"chat"``, ``"token"``, ``"completions"``, ``"audio"``, + ``"image"``, ``"chat_render"``, or ``"embeddings"``. + + Raises: + ProviderError: If the SDK backend is not active, or the resource is + not available for this endpoint type. + """ + if self._sdk_client is None: + raise ProviderError(self.get_name(), "Friendli SDK client not initialized.") + root = getattr(self._sdk_client, self._endpoint_type, None) + if root is None: + raise ProviderError( + self.get_name(), + f"Friendli SDK has no '{self._endpoint_type}' namespace.", + ) + namespace = getattr(root, resource, None) + if namespace is None: + raise ProviderError( + self.get_name(), + f"Friendli SDK '{self._endpoint_type}' endpoints do not expose '{resource}'.", + ) + return namespace + + # ========================================================================= + # BaseProvider interface + # ========================================================================= + + def get_name(self) -> str: + """Return the configured instance name (default: ``"friendli"``).""" + return self._provider_instance_name or "friendli" + + async def warm_up(self) -> None: + """Prime the Model APIs catalog so context lookups are exact. + + Only ``serverless`` exposes a catalog; for ``dedicated`` and + ``container`` this is a cheap no-op that logs the resolved endpoint. + """ + if self._endpoint_type != "serverless": + logger.debug( + "Friendli provider ready (instance=%s, endpoint_type=%s, base_url=%s, model=%s).", + self.get_name(), + self._endpoint_type, + self._base_url, + self.default_model, + ) + return + try: + catalog = await self._fetch_catalog() + logger.debug("Friendli catalog warmed: %d models.", len(catalog)) + except Exception as e: + logger.warning("Friendli catalog warm-up failed: %s", e) + + async def _fetch_catalog(self, force: bool = False) -> dict[str, dict[str, Any]]: + """Fetch and cache the Model APIs catalog, keyed by model ID. + + Args: + force: Re-fetch even when a cached catalog is present. + + Returns: + Mapping of model ID to the raw catalog entry. Empty for endpoint + types that expose no catalog. + """ + if self._catalog is not None and not force: + return self._catalog + if self._endpoint_type != "serverless": + self._catalog = {} + return self._catalog + + raw: list[dict[str, Any]] = [] + if self._backend == "sdk": + resp = await self._run_sdk(lambda: self._sdk_client.serverless.model.models()) + raw = [self._normalize_obj(m) for m in (getattr(resp, "data", None) or [])] + else: + # Both the openai and httpx backends read the catalog over raw HTTP: + # the Friendli listing is a superset of the OpenAI /models shape and + # the openai SDK's typed Model objects would discard the extras. + resp = await self._raw_get("/models") + raw = resp.json().get("data", []) or [] + + self._catalog = {m["id"]: m for m in raw if isinstance(m, dict) and m.get("id")} + return self._catalog + + async def get_models_details(self) -> list[ModelDetails]: + """Discover available models. + + ``serverless`` uses the rich Friendli catalog (``GET /models``), which + reports context length, modalities, reasoning options, and per-token + pricing. ``dedicated`` and ``container`` have no catalog, so the + configured default model is described from the model-card registry. + """ + try: + registry = get_model_card_registry() + except Exception: + registry = None + + if self._endpoint_type == "serverless": + try: + catalog = await self._fetch_catalog() + details = [ + self._build_model_details(mid, registry, entry) + for mid, entry in catalog.items() + ] + if details: + return details + except Exception as e: + logger.warning( + "Friendli model listing failed (%s); falling back to the static table.", + e, + ) + return [self._build_model_details(mid, registry, None) for mid in _CONTEXT_LENGTHS] + + # Dedicated endpoints / containers serve exactly one deployment. + return [self._build_model_details(self.default_model, registry, None)] + + def _build_model_details( + self, + model_id: str, + registry: Any | None, + entry: dict[str, Any] | None, + ) -> ModelDetails: + """Build :class:`ModelDetails` from a catalog entry and/or a model card.""" + provider = self.get_name() + functionality = (entry or {}).get("functionality") or {} + input_modalities = (entry or {}).get("input_modalities") or [] + output_modalities = (entry or {}).get("output_modalities") or [] + + supports_tools = bool(functionality.get("tool_call", True)) + supports_vision = "image" in input_modalities + supports_reasoning = bool((entry or {}).get("reasoning", False)) + max_output = (entry or {}).get("max_completion_tokens") + + if entry is None and registry: + card = registry.get(provider, model_id) + if card is not None and card.capabilities: + supports_tools = card.capabilities.tool_use or card.capabilities.function_calling + supports_vision = card.capabilities.vision + supports_reasoning = card.capabilities.reasoning + + return ModelDetails( + id=model_id, + provider_name=provider, + display_name=(entry or {}).get("name") or None, + context_length=self.get_max_context_length(model_id), + max_output_tokens=max_output, + supports_streaming=True, + supports_tools=supports_tools, + supports_vision=supports_vision, + supports_reasoning=supports_reasoning, + model_type=(entry or {}).get("mode") or "chat", + metadata={ + "endpoint_type": self._endpoint_type, + "base_model": (entry or {}).get("base_model"), + "description": (entry or {}).get("description"), + "input_modalities": input_modalities, + "output_modalities": output_modalities, + "reasoning_options": (entry or {}).get("reasoning_options"), + "interleaved": (entry or {}).get("interleaved"), + "pricing": (entry or {}).get("pricing"), + "functionality": functionality or None, + "default_params": (entry or {}).get("default_params"), + "deprecation_date": (entry or {}).get("deprecation_date"), + }, + ) + + def get_supported_parameters(self, model: str | None = None) -> dict[str, Any]: + """Return the inference parameters accepted by the Friendli chat API.""" + return { + # --- OpenAI-compatible sampling --- + "temperature": {"type": "number", "minimum": 0.0}, + "top_p": {"type": "number", "minimum": 0.0, "maximum": 1.0}, + "max_tokens": {"type": "integer", "minimum": 1}, + "n": {"type": "integer", "minimum": 1}, + "seed": {"type": ["integer", "array"]}, + "stop": {"type": "array", "items": {"type": "string"}}, + "frequency_penalty": {"type": "number", "minimum": -2.0, "maximum": 2.0}, + "presence_penalty": {"type": "number", "minimum": -2.0, "maximum": 2.0}, + "logit_bias": {"type": "object"}, + "logprobs": {"type": "boolean"}, + "top_logprobs": {"type": "integer", "minimum": 0}, + "parallel_tool_calls": {"type": "boolean"}, + "response_format": {"type": "object"}, + "stream_options": {"type": "object"}, + # --- Friendli Engine sampling --- + "top_k": {"type": "integer", "minimum": 0}, + "min_p": {"type": "number", "minimum": 0.0, "maximum": 1.0}, + "min_tokens": {"type": "integer", "minimum": 0}, + "repetition_penalty": {"type": "number", "exclusiveMinimum": 0.0}, + "eos_token": {"type": "array", "items": {"type": "integer"}}, + "xtc_threshold": {"type": "number", "minimum": 0.0, "maximum": 1.0}, + "xtc_probability": {"type": "number", "minimum": 0.0, "maximum": 1.0}, + # --- Reasoning controls --- + "reasoning_effort": {"type": "string", "enum": sorted(_VALID_EFFORTS)}, + "reasoning_budget": {"type": "integer"}, + "parse_reasoning": {"type": "boolean"}, + "include_reasoning": {"type": "boolean"}, + # --- Chat-template switches --- + "chat_template_kwargs": {"type": "object"}, + "enable_thinking": { + "type": "boolean", + "description": "Folded into chat_template_kwargs.enable_thinking.", + }, + "clear_thinking": { + "type": "boolean", + "description": "Folded into chat_template_kwargs.clear_thinking.", + }, + } + + def get_max_context_length(self, model: str | None = None) -> int: + """Return the maximum context length (tokens) for a Friendli model. + + Resolution order: the live catalog (when warmed), the static table, the + model-card registry, then the configured ``fallback_context_length``. + """ + model_name = model or self.default_model + + if self._catalog: + entry = self._catalog.get(model_name) + if entry and entry.get("context_length"): + return int(entry["context_length"]) + + limit = _CONTEXT_LENGTHS.get(model_name) + if limit is not None: + return limit + + try: + registry = get_model_card_registry() + card = registry.get(self.get_name(), model_name) + if card is not None: + return card.get_context_length() + except Exception: + pass + + logger.warning( + "Unknown context length for Friendli model '%s'. Falling back to %d.", + model_name, + self._fallback_context_length, + ) + return self._fallback_context_length + + # ========================================================================= + # Request construction + # ========================================================================= + + @staticmethod + def _normalize_media_part(item: Any, kind: str) -> dict[str, Any] | None: + """Normalize an inline media entry into a Friendli content part. + + Accepted *item* forms: + - ``str`` — an HTTPS URL or a base64 data URI. + - ``dict`` — a ready-made content part (passed through), a + ``{"": {...}}`` wrapper, or ``{"url": "..."}``. + + Args: + item: The raw inline entry. + kind: ``"image_url"``, ``"audio_url"``, or ``"video_url"``. + + Returns: + A content-part dict, or ``None`` if the entry is unrecognized. + """ + if isinstance(item, str): + return {"type": kind, kind: {"url": item}} + if isinstance(item, dict): + if item.get("type") in ("image_url", "audio_url", "video_url"): + return item + if kind in item: + inner = item[kind] + if isinstance(inner, str): + inner = {"url": inner} + return {"type": kind, kind: inner} + if "url" in item: + return {"type": kind, kind: {"url": item["url"]}} + logger.warning("Skipping unrecognized inline %s entry: %r", kind, item) + return None + + def _build_message_payload(self, msg: Message) -> dict[str, Any]: + """Build a Friendli-format message dict from an llmcore Message. + + Handles multimodal user content (``inline_images`` / ``inline_audio`` / + ``inline_videos`` / ``content_parts`` in ``metadata``), assistant + ``tool_calls`` with ``reasoning_content`` preservation, and tool-result + messages (``tool_call_id``). + """ + role_str = msg.role.value if hasattr(msg.role, "value") else str(msg.role) + metadata = msg.metadata or {} + + content: Any = msg.content + if "content_parts" in metadata: + content = metadata["content_parts"] + else: + inline_images = metadata.get("inline_images") or [] + inline_audio = metadata.get("inline_audio") or [] + inline_videos = metadata.get("inline_videos") or [] + if inline_images or inline_audio or inline_videos: + parts: list[dict[str, Any]] = [] + for item, kind in ( + *((i, "image_url") for i in inline_images), + *((a, "audio_url") for a in inline_audio), + *((v, "video_url") for v in inline_videos), + ): + part = self._normalize_media_part(item, kind) + if part: + parts.append(part) + if msg.content: + parts.append({"type": "text", "text": msg.content}) + content = parts + + msg_dict: dict[str, Any] = {"role": role_str, "content": content} + + if msg.role == LLMCoreRole.TOOL and msg.tool_call_id: + msg_dict["tool_call_id"] = msg.tool_call_id + + if role_str == "assistant": + # First-class Message.tool_calls (R-2) takes precedence over the + # legacy metadata channel so native tool-role results pair up. + tool_calls = getattr(msg, "tool_calls", None) or metadata.get("tool_calls") + if tool_calls: + msg_dict["tool_calls"] = tool_calls + if not msg.content: + # Friendli's AssistantMessage allows null content when + # tool_calls is present. + msg_dict["content"] = None + if "reasoning_content" in metadata: + msg_dict["reasoning_content"] = metadata["reasoning_content"] + + name = metadata.get("name") + if name: + msg_dict["name"] = name + + return msg_dict + + def _resolve_request_params( + self, kwargs: dict[str, Any] + ) -> tuple[dict[str, Any], dict[str, Any]]: + """Split validated kwargs into native and Friendli-specific body params. + + Configured reasoning defaults are applied for any key the caller did not + supply, and the flat ``enable_thinking`` / ``clear_thinking`` switches + are folded into ``chat_template_kwargs``. + + Returns: + ``(native, extras)`` — ``native`` holds OpenAI-compatible chat + parameters, ``extras`` holds the Friendli-specific ones. + """ + native: dict[str, Any] = {} + extras: dict[str, Any] = {} + template_kwargs: dict[str, Any] = dict(kwargs.pop("chat_template_kwargs", None) or {}) + + for key in _TEMPLATE_KWARG_KEYS: + if key in kwargs: + template_kwargs[key] = bool(kwargs.pop(key)) + if "enable_thinking" not in template_kwargs and self._default_enable_thinking is not None: + template_kwargs["enable_thinking"] = self._default_enable_thinking + + effort = kwargs.pop("reasoning_effort", None) + if effort is None: + effort = self._default_reasoning_effort + elif str(effort).lower() not in _VALID_EFFORTS: + logger.warning( + "Ignoring invalid Friendli reasoning_effort '%s'. Valid tiers: %s.", + effort, + ", ".join(sorted(_VALID_EFFORTS)), + ) + effort = self._default_reasoning_effort + else: + effort = str(effort).lower() + if effort is not None: + extras["reasoning_effort"] = effort + + budget = kwargs.pop("reasoning_budget", self._default_reasoning_budget) + if budget is not None: + extras["reasoning_budget"] = int(budget) + + parse_reasoning = kwargs.pop("parse_reasoning", self._default_parse_reasoning) + if parse_reasoning is not None: + extras["parse_reasoning"] = bool(parse_reasoning) + + include_reasoning = kwargs.pop("include_reasoning", self._default_include_reasoning) + if include_reasoning is not None: + extras["include_reasoning"] = bool(include_reasoning) + + for key, value in kwargs.items(): + if key in _EXTRA_BODY_KEYS: + extras[key] = value + else: + native[key] = value + + if template_kwargs: + extras["chat_template_kwargs"] = template_kwargs + + # ``seed`` may be a list (one per generation with ``n``), which the + # openai SDK does not type; route those through the Friendli extras. + if isinstance(native.get("seed"), (list, tuple)): + extras["seed"] = list(native.pop("seed")) + + return native, extras + + @staticmethod + def _apply_mutual_exclusions( + native: dict[str, Any], extras: dict[str, Any], has_tools: bool + ) -> None: + """Drop body fields Friendli rejects in combination, with a warning. + + The API rejects ``min_tokens`` and ``response_format`` when ``tools`` is + set, and ``min_tokens`` when ``response_format`` is set. + """ + if has_tools: + for key, bucket in (("min_tokens", extras), ("response_format", native)): + if key in bucket: + bucket.pop(key) + logger.warning("Friendli rejects '%s' together with 'tools'; dropping it.", key) + elif "response_format" in native and "min_tokens" in extras: + extras.pop("min_tokens") + logger.warning( + "Friendli rejects 'min_tokens' together with 'response_format'; dropping it." + ) + + # ========================================================================= + # Chat completion + # ========================================================================= + + async def chat_completion( + self, + context: ContextPayload, + model: str | None = None, + stream: bool = False, + tools: list[Tool] | None = None, + tool_choice: str | None = None, + **kwargs: Any, + ) -> dict[str, Any] | AsyncGenerator[dict[str, Any], None]: + """Perform a chat completion against the Friendli API. + + Dispatches to the active backend (``openai`` / ``httpx`` / ``sdk``). + Across all backends it provides: + + - Reasoning controls (``reasoning_effort``, ``reasoning_budget``, + ``parse_reasoning``, ``include_reasoning``) with configured defaults. + - ``chat_template_kwargs`` assembly from ``enable_thinking`` / + ``clear_thinking``. + - Friendli Engine sampling parameters (``top_k``, ``min_p``, + ``repetition_penalty``, ``min_tokens``, XTC, ``eos_token``). + - Multimodal (image / audio / video) content assembly. + - ``reasoning_content`` in both streaming and non-streaming responses + (``openai`` and ``httpx`` backends). + + Args: + context: The conversation as a list of ``Message`` objects. + model: Model ID (Model APIs) or endpoint ID (Dedicated Endpoints). + stream: Return an async generator of raw chunk dicts. + tools: Tools the model may call. + tool_choice: ``"none"``, ``"auto"``, ``"required"``, or a + ``{"type": "function", "function": {"name": ...}}`` dict. + **kwargs: Any parameter from :meth:`get_supported_parameters`. + + Returns: + A response dict (``stream=False``) or an async generator of chunk + dicts (``stream=True``). + + Raises: + ValueError: On an unsupported parameter name. + ProviderError: On API, auth, or transport failures. + ContextLengthError: When the prompt exceeds the model's window. + """ + model_name = model or self.default_model + + supported = self.get_supported_parameters(model_name) + for key in kwargs: + if key not in supported: + raise ValueError(f"Unsupported parameter '{key}' for Friendli provider.") + + if not (isinstance(context, list) and all(isinstance(m, Message) for m in context)): + raise ProviderError(self.get_name(), "Context must be list[Message].") + + messages_payload = [self._build_message_payload(m) for m in context] + if not messages_payload: + raise ProviderError(self.get_name(), "No valid messages.") + + tools_payload = None + if tools: + tools_payload = [{"type": "function", "function": t.model_dump()} for t in tools] + + native, extras = self._resolve_request_params(dict(kwargs)) + self._apply_mutual_exclusions(native, extras, has_tools=bool(tools_payload)) + + if tool_choice: + native["tool_choice"] = tool_choice + if stream: + native.setdefault("stream_options", {"include_usage": True}) + + if self.log_raw_payloads_enabled and logger.isEnabledFor(logging.DEBUG): + logger.debug( + "RAW FRIENDLI REQUEST (backend=%s, endpoint_type=%s): %s", + self._backend, + self._endpoint_type, + json.dumps( + { + "model": model_name, + "messages": messages_payload, + "stream": stream, + "tools": tools_payload, + **native, + **extras, + }, + indent=2, + default=str, + ), + ) + + try: + if self._backend == "openai": + return await self._chat_via_openai( + model_name, messages_payload, stream, tools_payload, native, extras + ) + if self._backend == "sdk": + return await self._chat_via_sdk( + model_name, messages_payload, stream, tools_payload, native, extras + ) + return await self._chat_via_httpx( + model_name, messages_payload, stream, tools_payload, native, extras + ) + except (ProviderError, ContextLengthError, ValueError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover - _raise_error always raises + + async def _chat_via_openai( + self, + model_name: str, + messages: list[dict[str, Any]], + stream: bool, + tools_payload: list[dict[str, Any]] | None, + native: dict[str, Any], + extras: dict[str, Any], + ) -> dict[str, Any] | AsyncGenerator[dict[str, Any], None]: + """Chat via ``AsyncOpenAI`` pointed at Friendli (extras via extra_body).""" + api_kwargs: dict[str, Any] = dict(native) + if extras: + api_kwargs["extra_body"] = dict(extras) + if tools_payload: + api_kwargs["tools"] = tools_payload + + resp = await self._client.chat.completions.create( + model=model_name, + messages=messages, + stream=stream, + **api_kwargs, + ) # type: ignore[arg-type] + + if stream: + + async def stream_wrapper() -> AsyncGenerator[dict[str, Any], None]: + async for chunk in resp: # type: ignore[union-attr] + chunk_dict = self._normalize_obj(chunk) + if self.log_raw_payloads_enabled and logger.isEnabledFor(logging.DEBUG): + logger.debug( + "RAW FRIENDLI STREAM CHUNK: %s", json.dumps(chunk_dict, default=str) + ) + yield chunk_dict + + return stream_wrapper() + + response_dict = self._normalize_obj(resp) + if self.log_raw_payloads_enabled and logger.isEnabledFor(logging.DEBUG): + logger.debug( + "RAW FRIENDLI RESPONSE: %s", json.dumps(response_dict, indent=2, default=str) + ) + return response_dict + + async def _chat_via_httpx( + self, + model_name: str, + messages: list[dict[str, Any]], + stream: bool, + tools_payload: list[dict[str, Any]] | None, + native: dict[str, Any], + extras: dict[str, Any], + ) -> dict[str, Any] | AsyncGenerator[dict[str, Any], None]: + """Chat via direct httpx calls against ``POST /chat/completions``.""" + body: dict[str, Any] = { + "model": model_name, + "messages": messages, + "stream": stream, + **native, + **extras, + } + if tools_payload: + body["tools"] = tools_payload + + if stream: + return self._httpx_sse_stream(body) + resp = await self._raw_post("/chat/completions", json_body=body) + return resp.json() + + async def _httpx_sse_stream(self, body: dict[str, Any]) -> AsyncGenerator[dict[str, Any], None]: + """Stream ``/chat/completions`` over httpx, parsing SSE ``data:`` lines.""" + client = self._get_http() + model_name = body.get("model", "") + try: + async with client.stream("POST", "/chat/completions", json=body) as resp: + if resp.status_code >= 400: + await resp.aread() + self._raise_status_error(resp.status_code, resp.text, model_name) + async for line in resp.aiter_lines(): + if not line or not line.startswith("data:"): + continue + data = line[len("data:") :].strip() + if not data or data == "[DONE]": + continue + try: + chunk = json.loads(data) + except json.JSONDecodeError: + continue + if self.log_raw_payloads_enabled and logger.isEnabledFor(logging.DEBUG): + logger.debug( + "RAW FRIENDLI STREAM CHUNK: %s", json.dumps(chunk, default=str) + ) + yield chunk + except (ProviderError, ContextLengthError): + raise + except httpx.HTTPError as e: + raise ProviderError(self.get_name(), f"Streaming error: {e}", model_name=model_name) + + async def _chat_via_sdk( + self, + model_name: str, + messages: list[dict[str, Any]], + stream: bool, + tools_payload: list[dict[str, Any]] | None, + native: dict[str, Any], + extras: dict[str, Any], + ) -> dict[str, Any] | AsyncGenerator[dict[str, Any], None]: + """Chat via the official ``friendli`` SDK. + + The SDK types every parameter explicitly, so there is no ``extra_body`` + escape hatch: any key it does not declare is dropped with a warning + rather than silently changing the request. + """ + chat = self._sdk_namespace("chat") + call_kwargs: dict[str, Any] = {**native, **extras} + if tools_payload: + call_kwargs["tools"] = tools_payload + if self._endpoint_type == "container": + call_kwargs["server_url"] = self._base_url + + method = chat.stream if stream else chat.complete + call_kwargs = self._filter_sdk_kwargs(method, call_kwargs) + + if stream: + event_stream = await self._run_sdk( + lambda: chat.stream(model=model_name, messages=messages, **call_kwargs) + ) + return self._bridge_sdk_stream(event_stream) + + resp = await self._run_sdk( + lambda: chat.complete(model=model_name, messages=messages, **call_kwargs) + ) + return self._normalize_obj(resp) + + @staticmethod + def _filter_sdk_kwargs(method: Any, call_kwargs: dict[str, Any]) -> dict[str, Any]: + """Drop kwargs the SDK method does not declare, with a warning. + + The generated SDK types every parameter explicitly and has no + ``extra_body`` escape hatch, so an undeclared key would raise an opaque + ``TypeError`` deep inside the vendor package. Dropping it here keeps + the failure legible and the request valid; the ``openai`` and ``httpx`` + backends pass the same key through untouched. + + Args: + method: The bound SDK operation about to be called. + call_kwargs: The assembled request parameters. + + Returns: + ``call_kwargs`` restricted to the parameters *method* accepts. + """ + try: + accepted = set(inspect.signature(method).parameters) + except (TypeError, ValueError): # pragma: no cover - defensive + return call_kwargs + kept = {k: v for k, v in call_kwargs.items() if k in accepted} + dropped = sorted(set(call_kwargs) - set(kept)) + if dropped: + logger.warning( + "Friendli SDK backend does not accept %s; dropping. Use " + 'backend = "openai" or "httpx" to send them.', + ", ".join(dropped), + ) + return kept + + async def _bridge_sdk_stream(self, event_stream: Any) -> AsyncGenerator[dict[str, Any], None]: + """Yield normalized chunk dicts from an SDK ``EventStreamAsync``.""" + try: + async for chunk in event_stream: + chunk_dict = self._normalize_obj(chunk) + if self.log_raw_payloads_enabled and logger.isEnabledFor(logging.DEBUG): + logger.debug( + "RAW FRIENDLI SDK STREAM CHUNK: %s", json.dumps(chunk_dict, default=str) + ) + yield chunk_dict + finally: + close = getattr(event_stream, "close", None) + if callable(close): + try: + await close() + except Exception as e: # pragma: no cover - best effort + logger.debug("Error closing Friendli SDK stream: %s", e) + + @staticmethod + def _normalize_obj(obj: Any) -> dict[str, Any]: + """Normalize an SDK/OpenAI pydantic response into a plain dict.""" + if hasattr(obj, "model_dump"): + return obj.model_dump(exclude_none=True) + if isinstance(obj, dict): + return obj + return dict(obj) + + # ========================================================================= + # Error mapping + # ========================================================================= + + def _raise_status_error(self, status: int, body: str, model_name: str) -> None: + """Map an HTTP status + body to an llmcore exception. + + Raises: + ContextLengthError: For context/length overflow (HTTP 400/422). + ProviderError: For every other failure. + """ + logger.error("Friendli status error (%s): %s", status, body) + if status in (400, 422) and _looks_like_context_overflow(body): + raise ContextLengthError( + model_name=model_name, + limit=self.get_max_context_length(model_name), + actual=0, + message=body, + ) + if status == 401: + raise ProviderError( + self.get_name(), + f"Friendli authentication failed. Verify FRIENDLI_TOKEN / " + f"FRIENDLIAI_API_KEY (Personal API keys start with 'flp_'). " + f"Error: {body}", + model_name=model_name, + status_code=status, + ) + if status == 403: + raise ProviderError( + self.get_name(), + f"Friendli request forbidden. Check the team scope " + f"(X-Friendli-Team = {self._team_id or ''}) and that the key " + f"may use this endpoint. Error: {body}", + model_name=model_name, + status_code=status, + ) + if status == 404: + raise ProviderError( + self.get_name(), + f"Friendli returned 404 on the '{self._endpoint_type}' surface " + f"({self._base_url}) for model/endpoint '{model_name}'. Either the " + f"model/endpoint ID is wrong (Dedicated Endpoints take the endpoint ID, " + f"not a model name), or that route is not served on this surface - " + f"/detokenize and /chat/render are documented but currently 404 on " + f"Model APIs. Error: {body}", + model_name=model_name, + status_code=status, + ) + if status == 429: + raise ProviderError( + self.get_name(), + f"Friendli rate limit exceeded. Model APIs limits scale with your " + f"usage tier (https://friendli.ai/docs/guides/model-apis/rate-limits). " + f"Error: {body}", + model_name=model_name, + status_code=status, + ) + raise ProviderError( + self.get_name(), + f"API Error ({status}): {body}", + model_name=model_name, + status_code=status, + ) + + def _raise_error(self, e: Exception, model_name: str) -> None: + """Map a backend exception to an llmcore exception. + + Raises: + ProviderError or ContextLengthError: Always. + """ + status = getattr(e, "status_code", None) + if status is None: + response = getattr(e, "response", None) + status = getattr(response, "status_code", None) + body = getattr(e, "body", None) or str(e) + if status is not None: + self._raise_status_error(int(status), str(body), model_name) + + if isinstance(e, OpenAIAPITimeoutError): + raise ProviderError( + self.get_name(), f"Timeout: {e}", model_name=model_name, retryable=True + ) + if isinstance(e, OpenAIAPIConnectionError): + raise ProviderError( + self.get_name(), f"Connection error: {e}", model_name=model_name, retryable=True + ) + if httpx_available and isinstance(e, httpx.TimeoutException): + raise ProviderError( + self.get_name(), f"Timeout: {e}", model_name=model_name, retryable=True + ) + if httpx_available and isinstance(e, httpx.HTTPError): + raise ProviderError( + self.get_name(), f"HTTP error: {e}", model_name=model_name, retryable=True + ) + logger.error("Unexpected Friendli error: %s", e, exc_info=True) + raise ProviderError( + self.get_name(), f"Error: {e}", model_name=model_name, original_exception=e + ) + + # ========================================================================= + # Response extraction + # ========================================================================= + + def extract_response_content(self, response: dict[str, Any]) -> str: + """Extract the final text content from a non-streaming response.""" + try: + choices = response.get("choices", []) + if not choices: + return "" + return choices[0].get("message", {}).get("content") or "" + except (KeyError, IndexError, TypeError) as e: + logger.warning("Failed to extract Friendli content: %s", e) + return "" + + def extract_delta_content(self, chunk: dict[str, Any]) -> str: + """Extract the text delta from a streaming chunk.""" + try: + choices = chunk.get("choices", []) + if not choices: + return "" + return choices[0].get("delta", {}).get("content") or "" + except (KeyError, IndexError, TypeError): + return "" + + def extract_reasoning_content(self, response: dict[str, Any]) -> str | None: + """Extract parsed reasoning from a non-streaming response. + + Friendli returns the chain of thought as ``reasoning_content`` (and the + compatibility alias ``reasoning``) when ``parse_reasoning`` is on and + the model supports it. + """ + try: + choices = response.get("choices", []) + if not choices: + return None + message = choices[0].get("message", {}) + return message.get("reasoning_content") or message.get("reasoning") + except (KeyError, IndexError, TypeError): + return None + + def extract_delta_reasoning_content(self, chunk: dict[str, Any]) -> str | None: + """Extract the reasoning delta from a streaming chunk.""" + try: + choices = chunk.get("choices", []) + if not choices: + return None + delta = choices[0].get("delta", {}) + return delta.get("reasoning_content") or delta.get("reasoning") + except (KeyError, IndexError, TypeError): + return None + + def extract_tool_calls(self, response: dict[str, Any]) -> list[ToolCall]: + """Extract tool calls from a Friendli response.""" + out: list[ToolCall] = [] + try: + choices = response.get("choices", []) + if not choices: + return out + raw_calls = choices[0].get("message", {}).get("tool_calls") + if not raw_calls: + return out + for tc in raw_calls: + if tc.get("type", "function") != "function": + continue + func = tc.get("function", {}) + args_str = func.get("arguments", "{}") + try: + args_dict = json.loads(args_str) + except (json.JSONDecodeError, TypeError): + args_dict = {"_raw": args_str} + out.append( + ToolCall( + id=tc.get("id", ""), + name=func.get("name", ""), + arguments=args_dict, + ) + ) + except (KeyError, IndexError, TypeError) as e: + logger.warning("Failed to extract Friendli tool calls: %s", e) + return out + + def extract_usage_details(self, response: dict[str, Any]) -> dict[str, Any]: + """Extract usage, including Friendli's cached-prompt accounting.""" + usage = response.get("usage") or {} + if not usage: + return {} + result: dict[str, Any] = { + "prompt_tokens": usage.get("prompt_tokens"), + "completion_tokens": usage.get("completion_tokens"), + "total_tokens": usage.get("total_tokens"), + } + details = usage.get("prompt_tokens_details") or {} + if details.get("cached_tokens") is not None: + result["cached_tokens"] = details["cached_tokens"] + return result + + def extract_finish_reason(self, response: dict[str, Any]) -> str | None: + """Extract the finish reason (``stop`` / ``length`` / ``tool_calls``).""" + try: + choices = response.get("choices", []) + if not choices: + return None + return choices[0].get("finish_reason") + except (KeyError, IndexError, TypeError): + return None + + # ========================================================================= + # Token counting (native /tokenize, with local fallback) + # ========================================================================= + + async def tokenize(self, text: str, model: str | None = None) -> list[int]: + """Tokenize *text* with the model's own tokenizer (``POST /tokenize``). + + Args: + text: The prompt text to tokenize. + model: Model/endpoint ID; defaults to the configured model. + + Returns: + The list of token IDs. + + Raises: + ProviderError: On API or transport failures. + """ + model_name = model or self.default_model + body = {"model": model_name, "prompt": text} + try: + if self._backend == "sdk": + resp = await self._run_sdk( + lambda: self._sdk_namespace("token").tokenize( + model=model_name, + prompt=text, + **self._sdk_server_kwargs(), + ) + ) + return list(self._normalize_obj(resp).get("tokens", [])) + resp = await self._raw_post("/tokenize", json_body=body) + return list(resp.json().get("tokens", [])) + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + async def detokenize(self, tokens: list[int], model: str | None = None) -> str: + """Convert token IDs back into text (``POST /detokenize``). + + Args: + tokens: The token IDs to decode. + model: Model/endpoint ID; defaults to the configured model. + + Returns: + The decoded text. + """ + model_name = model or self.default_model + body = {"model": model_name, "tokens": list(tokens)} + try: + if self._backend == "sdk": + resp = await self._run_sdk( + lambda: self._sdk_namespace("token").detokenize( + model=model_name, + tokens=list(tokens), + **self._sdk_server_kwargs(), + ) + ) + return str(self._normalize_obj(resp).get("text", "")) + resp = await self._raw_post("/detokenize", json_body=body) + return str(resp.json().get("text", "")) + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + async def render_chat( + self, + messages: list[Message], + model: str | None = None, + tools: list[Tool] | None = None, + chat_template_kwargs: dict[str, Any] | None = None, + ) -> str: + """Render messages into the exact prompt text sent to the model. + + Useful for debugging chat templates and for computing an exact prompt + token count (render, then :meth:`tokenize`). + + Args: + messages: The conversation to render. + model: Model/endpoint ID; defaults to the configured model. + tools: Tools to include in the rendered template. + chat_template_kwargs: Template switches (e.g. ``enable_thinking``). + + Returns: + The rendered prompt text. + """ + model_name = model or self.default_model + payload = [self._build_message_payload(m) for m in messages] + body: dict[str, Any] = {"model": model_name, "messages": payload} + if tools: + body["tools"] = [{"type": "function", "function": t.model_dump()} for t in tools] + if chat_template_kwargs: + body["chat_template_kwargs"] = chat_template_kwargs + try: + if self._backend == "sdk": + resp = await self._run_sdk( + lambda: self._sdk_namespace("chat_render").render( + **body, + **self._sdk_server_kwargs(), + ) + ) + return str(self._normalize_obj(resp).get("text", "")) + resp = await self._raw_post("/chat/render", json_body=body) + return str(resp.json().get("text", "")) + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + async def count_tokens(self, text: str, model: str | None = None) -> int: + """Count tokens for *text* using the model's own tokenizer. + + Counting is local (tiktoken ``cl100k_base``, then a character-ratio + estimate) unless ``native_token_count = true`` is configured, in which + case the model's own tokenizer is used via ``POST /tokenize``. Each + native count is an API request against the Model APIs rate-limit + budget, so it is opt-in; a failed native call falls back locally. + """ + if not text: + return 0 + if self._native_token_count: + try: + return len(await self.tokenize(text, model)) + except Exception as e: + logger.debug("Friendli /tokenize failed (%s); estimating locally.", e) + return self._count_tokens_locally(text) + + def _count_tokens_locally(self, text: str) -> int: + """Count tokens with tiktoken, or a character-ratio estimate.""" + if not text: + return 0 + if self._encoding is None: + return _EstimateCounter().count(text) + try: + return len(self._encoding.encode(text)) + except Exception: + return _EstimateCounter().count(text) + + async def count_message_tokens(self, messages: list[Message], model: str | None = None) -> int: + """Estimate the total prompt tokens for *messages*. + + The role-tagged conversation is counted as one block plus a small + per-message overhead for the chat template's delimiters. With + ``native_token_count = true`` the block is tokenized by the model's own + tokenizer in a single ``/tokenize`` call; otherwise it is counted + locally. For an exact, template-accurate count use ``render_chat()`` + followed by ``tokenize()``. + """ + if not messages: + return 0 + joined = "\n".join( + f"{m.role.value if hasattr(m.role, 'value') else m.role}: {m.content or ''}" + for m in messages + ) + overhead = 4 * len(messages) + 3 + if self._native_token_count: + try: + return len(await self.tokenize(joined, model)) + overhead + except Exception as e: + logger.debug("Friendli /tokenize failed for message counting (%s); estimating.", e) + return self._count_tokens_locally(joined) + overhead + + # ========================================================================= + # Auxiliary inference endpoints + # ========================================================================= + + def _sdk_server_kwargs(self) -> dict[str, Any]: + """Per-call SDK kwargs that pin a self-hosted container's URL.""" + if self._endpoint_type == "container": + return {"server_url": self._base_url} + return {} + + async def text_completion( + self, + prompt: str, + *, + model: str | None = None, + stream: bool = False, + **kwargs: Any, + ) -> dict[str, Any] | AsyncGenerator[dict[str, Any], None]: + """Generate a raw text completion (``POST /completions``). + + This is the prompt-based (non-chat) surface; the chat template is *not* + applied, so the prompt is sent to the model verbatim. It always goes + over the raw HTTP client, independent of the configured ``backend``. + + Args: + prompt: The raw prompt text. + model: Model/endpoint ID; defaults to the configured model. + stream: Return an async generator of raw chunk dicts. + **kwargs: Additional generation parameters (``max_tokens``, + ``temperature``, ``top_k``, ``min_tokens``, ``stop``, …). + + Returns: + A response dict, or an async generator of chunk dicts when + ``stream=True``. + """ + model_name = model or self.default_model + body: dict[str, Any] = {"model": model_name, "prompt": prompt, "stream": stream, **kwargs} + try: + if stream: + return self._httpx_sse_stream_path("/completions", body) + resp = await self._raw_post("/completions", json_body=body) + return resp.json() + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + async def create_embeddings( + self, + input_texts: str | list[str], + *, + model: str | None = None, + encoding_format: str | None = None, + **kwargs: Any, + ) -> dict[str, Any]: + """Create text embeddings (``POST /embeddings``). + + Embeddings are served by Dedicated Endpoints and Friendli Container + only — the hosted Model APIs catalog has no embedding surface. + + Args: + input_texts: A string or list of strings to embed. + model: Endpoint ID; defaults to the configured model. + encoding_format: ``"float"`` (default) or ``"base64"``. + **kwargs: Additional body parameters. + + Returns: + The raw API response dict with ``data``, ``model``, and ``usage``. + + Raises: + ProviderError: On the ``serverless`` endpoint type, or on API errors. + """ + if self._endpoint_type == "serverless": + raise ProviderError( + self.get_name(), + "Friendli Model APIs do not expose an embeddings endpoint. Deploy an " + "embedding model on a Dedicated Endpoint (or Container) and set " + "providers.friendli.endpoint_type accordingly.", + ) + model_name = model or self.default_model + body: dict[str, Any] = {"model": model_name, "input": input_texts, **kwargs} + if encoding_format is not None: + body["encoding_format"] = encoding_format + try: + resp = await self._raw_post("/embeddings", json_body=body) + return resp.json() + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + async def generate_image( + self, + prompt: str, + *, + model: str | None = None, + n: int = 1, + size: str | None = None, + quality: str | None = None, + response_format: str = "url", + style: str | None = None, + **kwargs: Any, + ) -> ImageGenerationResult: + """Generate images (``POST /images/generations``). + + Image generation is served by Dedicated Endpoints and Friendli Container + only. ``n``, ``size``, ``quality``, and ``style`` are part of the + llmcore signature but have no Friendli equivalent and are ignored; + use ``num_inference_steps``, ``guidance_scale``, ``seed``, and + ``control_images`` instead. + + Args: + prompt: A text description of the desired image. + model: Endpoint ID; defaults to the configured model. + n: Ignored (Friendli returns a single image per request). + size: Ignored. + quality: Ignored. + response_format: ``"url"`` (default), ``"raw"``, ``"png"``, + ``"jpeg"``, or ``"jpg"``. + style: Ignored. + **kwargs: Friendli parameters (``num_inference_steps``, + ``guidance_scale``, ``seed``, ``control_images``, + ``controlnet_weights``). + + Returns: + An :class:`ImageGenerationResult`. + + Raises: + ProviderError: On the ``serverless`` endpoint type, or on API errors. + """ + if self._endpoint_type == "serverless": + raise ProviderError( + self.get_name(), + "Friendli Model APIs do not expose an image-generation endpoint. " + "Deploy an image model on a Dedicated Endpoint (or Container) and set " + "providers.friendli.endpoint_type accordingly.", + ) + for ignored, value in ( + ("n", n if n != 1 else None), + ("size", size), + ("quality", quality), + ("style", style), + ): + if value is not None: + logger.debug("Friendli image generation ignores '%s'.", ignored) + + model_name = model or self.default_model + body: dict[str, Any] = { + "model": model_name, + "prompt": prompt, + "response_format": response_format, + **kwargs, + } + try: + resp = await self._raw_post("/images/generations", json_body=body) + payload = resp.json() + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + images: list[GeneratedImage] = [] + for item in payload.get("data", []) or []: + fmt = item.get("response_format") or response_format + images.append( + GeneratedImage( + url=item.get("url"), + data=item.get("b64_json") or item.get("image"), + format="png" if fmt in ("url", "raw") else str(fmt), + ) + ) + return ImageGenerationResult( + images=images, + model=model_name, + metadata={"raw": payload, "endpoint_type": self._endpoint_type}, + ) + + async def transcribe_audio( + self, + audio_data: bytes | str, + *, + model: str | None = None, + language: str | None = None, + prompt: str | None = None, + response_format: str = "json", + temperature: float | None = None, + timestamp_granularities: list[str] | None = None, + **kwargs: Any, + ) -> TranscriptionResult: + """Transcribe audio to text (``POST /audio/transcriptions``). + + Args: + audio_data: Raw audio bytes or a path to an audio file. + model: Transcription model (e.g. ``openai/whisper-large-v3``) or + endpoint ID; defaults to the configured model. + language: ISO-639-1 hint (e.g. ``"en"``) to improve accuracy. + prompt: Unsupported by Friendli; ignored. + response_format: Unsupported by Friendli; ignored (JSON is always + returned). + temperature: Sampling temperature between 0 and 1. + timestamp_granularities: Unsupported by Friendli; ignored. + **kwargs: Additional form fields (e.g. ``chunking_strategy``). + + Returns: + A :class:`TranscriptionResult`. + """ + for name, value in ( + ("prompt", prompt), + ("timestamp_granularities", timestamp_granularities), + ): + if value is not None: + logger.debug("Friendli audio transcription ignores '%s'.", name) + if response_format not in ("json", "verbose_json"): + logger.debug( + "Friendli audio transcription always returns JSON; ignoring response_format='%s'.", + response_format, + ) + + model_name = model or self.default_model + if isinstance(audio_data, str): + with open(audio_data, "rb") as fh: + payload_bytes = fh.read() + filename = os.path.basename(audio_data) + else: + payload_bytes = audio_data + filename = "audio.wav" + + data: dict[str, Any] = {"model": model_name} + if language: + data["language"] = language + if temperature is not None: + data["temperature"] = str(temperature) + for key, value in kwargs.items(): + data[key] = value if isinstance(value, str) else json.dumps(value) + + try: + resp = await self._raw_post( + "/audio/transcriptions", + data=data, + files={"file": (filename, payload_bytes)}, + ) + payload = resp.json() + except (ProviderError, ContextLengthError): + raise + except Exception as e: + self._raise_error(e, model_name) + raise # pragma: no cover + + usage = payload.get("usage") or {} + duration_ms = usage.get("input_audio_length_ms") + return TranscriptionResult( + text=payload.get("text", ""), + language=language, + duration_seconds=(duration_ms / 1000.0) if duration_ms else None, + model=model_name, + metadata={"usage": usage, "raw": payload}, + ) + + # ========================================================================= + # Friendli Suite (team billing / usage) + # ========================================================================= + + async def get_team_cost( + self, + start_time: str, + end_time: str, + *, + bucket_width: str | None = None, + limit: int | None = None, + page: str | None = None, + group_by: str | None = None, + ) -> dict[str, Any]: + """Read team cost buckets from the Friendli Suite API. + + Args: + start_time: RFC 3339 UTC timestamp with a zeroed time portion + (e.g. ``"2026-09-01T00:00:00Z"``); no earlier than one year ago. + end_time: RFC 3339 UTC timestamp with a zeroed time portion. + bucket_width: Bucket size; only ``"1d"`` is currently supported. + limit: Number of buckets to return (1-35, default 7). + page: Pagination cursor from a previous ``next_page``. + group_by: Currently only ``"line_item"``. + + Returns: + The raw API response dict. + + Note: + Friendli asks callers to wait at least five minutes between repeated + calls; usage takes a short while to appear in cost. + """ + params: dict[str, Any] = {"start_time": start_time, "end_time": end_time} + for key, value in ( + ("bucket_width", bucket_width), + ("limit", limit), + ("page", page), + ("group_by", group_by), + ): + if value is not None: + params[key] = value + return await self._suite_get("/team/cost", params) + + async def get_team_usage( + self, + start_time: str, + end_time: str, + *, + bucket_width: str | None = None, + limit: int | None = None, + page: str | None = None, + group_by: str | None = None, + ) -> dict[str, Any]: + """Read team usage buckets from the Friendli Suite API. + + Args: + start_time: RFC 3339 UTC timestamp with a zeroed time portion. + end_time: RFC 3339 UTC timestamp with a zeroed time portion. + bucket_width: Bucket size (e.g. ``"1d"``). + limit: Number of buckets to return. + page: Pagination cursor from a previous ``next_page``. + group_by: Grouping dimension supported by the API. + + Returns: + The raw API response dict. + """ + params: dict[str, Any] = {"start_time": start_time, "end_time": end_time} + for key, value in ( + ("bucket_width", bucket_width), + ("limit", limit), + ("page", page), + ("group_by", group_by), + ): + if value is not None: + params[key] = value + return await self._suite_get("/team/usage", params) + + async def _suite_get(self, path: str, params: dict[str, Any]) -> dict[str, Any]: + """GET a Friendli Suite endpoint (separate root from inference).""" + if not httpx_available: + raise ProviderError( + self.get_name(), "The 'httpx' package is required for the Friendli Suite API." + ) + url = f"{self._suite_base_url}{path}" + headers = self._team_headers({"Authorization": f"Bearer {self._api_key}"}) + try: + async with httpx.AsyncClient(timeout=self._timeout) as client: + resp = await client.get(url, params=params, headers=headers) + if resp.status_code >= 400: + self._raise_status_error(resp.status_code, resp.text, "") + return resp.json() + except (ProviderError, ContextLengthError): + raise + except httpx.HTTPError as e: + raise ProviderError(self.get_name(), f"HTTP error: {e}", retryable=True) + + # ========================================================================= + # HTTP plumbing (httpx backend + endpoints with no SDK/OpenAI surface) + # ========================================================================= + + def _get_http(self) -> Any: + """Return (lazily creating) the raw httpx client for Friendli REST calls.""" + if not httpx_available: + raise ProviderError( + self.get_name(), + "The 'httpx' package is required for the Friendli REST endpoints. " + "Install with: pip install llmcore[friendli]", + ) + if self._http is None: + self._http = httpx.AsyncClient( + base_url=self._base_url, + headers=self._team_headers({"Authorization": f"Bearer {self._api_key}"}), + timeout=self._timeout, + ) + return self._http + + async def _raw_post( + self, + path: str, + *, + json_body: dict[str, Any] | None = None, + data: dict[str, Any] | None = None, + files: Any | None = None, + ) -> Any: + """POST to a Friendli endpoint and return the httpx response.""" + client = self._get_http() + try: + resp = await client.post(path, json=json_body, data=data, files=files) + except httpx.HTTPError as e: + raise ProviderError(self.get_name(), f"HTTP error: {e}", retryable=True) + if resp.status_code >= 400: + model_name = (json_body or {}).get("model") or (data or {}).get("model") or "" + self._raise_status_error(resp.status_code, resp.text, str(model_name)) + return resp + + async def _raw_get(self, path: str, params: dict[str, Any] | None = None) -> Any: + """GET a Friendli endpoint and return the httpx response.""" + client = self._get_http() + try: + resp = await client.get(path, params=params) + except httpx.HTTPError as e: + raise ProviderError(self.get_name(), f"HTTP error: {e}", retryable=True) + if resp.status_code >= 400: + self._raise_status_error(resp.status_code, resp.text, "") + return resp + + async def _httpx_sse_stream_path( + self, path: str, body: dict[str, Any] + ) -> AsyncGenerator[dict[str, Any], None]: + """Stream any SSE endpoint over httpx, parsing ``data:`` lines.""" + client = self._get_http() + model_name = body.get("model", "") + try: + async with client.stream("POST", path, json=body) as resp: + if resp.status_code >= 400: + await resp.aread() + self._raise_status_error(resp.status_code, resp.text, model_name) + async for line in resp.aiter_lines(): + if not line or not line.startswith("data:"): + continue + payload = line[len("data:") :].strip() + if not payload or payload == "[DONE]": + continue + try: + yield json.loads(payload) + except json.JSONDecodeError: + continue + except (ProviderError, ContextLengthError): + raise + except httpx.HTTPError as e: + raise ProviderError(self.get_name(), f"Streaming error: {e}", model_name=model_name) + + # ========================================================================= + # Resource cleanup + # ========================================================================= + + async def close(self) -> None: + """Close the OpenAI / Friendli SDK / raw HTTP clients (best effort).""" + if self._client is not None: + try: + await self._client.close() + except Exception as e: + logger.error("Error closing Friendli OpenAI client: %s", e) + self._client = None + if self._sdk_client is not None: + try: + closer = getattr(self._sdk_client, "close", None) + if callable(closer): + result = closer() + if asyncio.iscoroutine(result): + await result + except Exception as e: + logger.error("Error closing Friendli SDK client: %s", e) + self._sdk_client = None + if self._http is not None: + try: + await self._http.aclose() + except Exception as e: + logger.error("Error closing Friendli HTTP client: %s", e) + self._http = None + logger.info("FriendliProvider closed.") diff --git a/src/llmcore/providers/manager.py b/src/llmcore/providers/manager.py index 879afdbc..5ac407da 100644 --- a/src/llmcore/providers/manager.py +++ b/src/llmcore/providers/manager.py @@ -28,6 +28,7 @@ from .deepgram_provider import DeepgramProvider from .deepinfra_provider import DeepInfraProvider from .deepseek_provider import DeepSeekProvider +from .friendli_provider import FriendliProvider from .gemini_provider import GeminiProvider from .huggingface_provider import HuggingFaceProvider from .kimi_provider import KimiProvider @@ -65,6 +66,8 @@ "deepinfra": DeepInfraProvider, # Z.ai (Zhipu AI) — GLM family of models. "zai": ZaiProvider, + # FriendliAI: Model APIs (serverless), Dedicated Endpoints, and Container. + "friendli": FriendliProvider, # Deepgram: speech/audio provider (STT/TTS/Voice Agent) — native SDK. "deepgram": DeepgramProvider, # TypeSafe.ai: System One typed-judgment provider (noul/choice/score) — @@ -74,6 +77,9 @@ "jev": TypeSafeProvider, # Alias: moonshot → kimi (Moonshot AI is the vendor; Kimi is the brand). "moonshot": KimiProvider, + # Aliases for FriendliAI: friendliai (brand) / friendli_ai. + "friendliai": FriendliProvider, + "friendli_ai": FriendliProvider, # Aliases for Z.ai: glm (brand) and zhipu/zhipuai/bigmodel (vendor). "glm": ZaiProvider, "zhipu": ZaiProvider, @@ -101,6 +107,8 @@ "zhipuai": "zai", "bigmodel": "zai", "jev": "typesafe", + "friendliai": "friendli", + "friendli_ai": "friendli", } # Well-known defaults for providers that reuse OpenAIProvider. diff --git a/tests/providers/test_friendli_provider.py b/tests/providers/test_friendli_provider.py new file mode 100644 index 00000000..e33414a7 --- /dev/null +++ b/tests/providers/test_friendli_provider.py @@ -0,0 +1,1208 @@ +# tests/providers/test_friendli_provider.py +""" +Tests for the FriendliAI provider implementation. + +Covers: +- Initialization: API key / team ID resolution, endpoint types, base URLs +- Backend resolution (openai -> httpx -> sdk) and explicit selection +- Request parameter splitting (native vs Friendli extras, chat_template_kwargs) +- Mutually-exclusive body fields (tools vs min_tokens / response_format) +- Message payload building (multimodal, tool_calls, reasoning_content) +- Chat completion on the ``openai`` backend (mocked AsyncOpenAI) +- Chat completion + SSE streaming on the ``httpx`` backend (respx) +- Response extraction (content, reasoning, tool calls, usage, finish reason) +- Model discovery from the rich Friendli catalog +- Context length resolution and token counting (local + native) +- Error mapping (401/403/404/429/context overflow) +- Endpoint-type gating for embeddings and image generation +- Provider registration and aliases + +Backends are always selected explicitly so the suite is deterministic whether +or not the optional ``friendli`` SDK is installed. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +import respx + +from llmcore.exceptions import ConfigError, ContextLengthError, ProviderError +from llmcore.models import Message, Role, Tool +from llmcore.providers.friendli_provider import ( + FriendliProvider, + friendli_sdk_available, +) + +SERVERLESS_URL = "https://api.friendli.ai/serverless/v1" +DEDICATED_URL = "https://api.friendli.ai/dedicated/v1" + +BASE_CONFIG: dict[str, Any] = { + "api_key": "flp_test_key", + "default_model": "zai-org/GLM-5.3", + "timeout": 30, +} + +CATALOG_ENTRY: dict[str, Any] = { + "id": "zai-org/GLM-5.3-Flash", + "name": "zai-org/GLM-5.3-Flash", + "created": 1787929200, + "context_length": 1048576, + "max_completion_tokens": 1048576, + "pricing": { + "input": "0.00000015", + "output": "0.0000005", + "prompt": "0.00000015", + "completion": "0.0000005", + "input_cache_read": "0.00000003", + }, + "functionality": { + "tool_call": True, + "parallel_tool_call": True, + "structured_output": True, + "tool_choice": True, + "system_messages": True, + }, + "description": "Fast GLM model", + "deprecation_date": None, + "reasoning": True, + "reasoning_options": [{"type": "effort", "values": ["low", "high", "max"]}], + "input_modalities": ["text", "image", "video"], + "output_modalities": ["text"], + "interleaved": "reasoning_content", + "base_model": "zhipuai/glm-5.3-flash", + "mode": "chat", + "default_params": {"temperature": 1.0, "top_p": 1.0, "top_k": 0, "min_p": 0.0}, +} + + +def _clean_env(monkeypatch) -> None: + """Remove every Friendli environment variable the provider consults.""" + for name in ( + "FRIENDLI_TOKEN", + "FRIENDLIAI_API_KEY", + "FRIENDLI_API_KEY", + "FRIENDLI_TEAM_ID", + "FRIENDLIAI_TEAM_ID", + ): + monkeypatch.delenv(name, raising=False) + + +@pytest.fixture(autouse=True) +def _isolated_env(monkeypatch): + """Keep ambient Friendli credentials out of every test.""" + _clean_env(monkeypatch) + yield + + +@pytest.fixture +def provider(): + """A FriendliProvider on the mocked ``openai`` backend.""" + with patch("llmcore.providers.friendli_provider.AsyncOpenAI") as mock_cls: + mock_cls.return_value = MagicMock() + p = FriendliProvider({**BASE_CONFIG, "backend": "openai"}, log_raw_payloads=False) + return p + + +@pytest.fixture +def httpx_provider(): + """A FriendliProvider on the ``httpx`` backend (exercised with respx).""" + return FriendliProvider({**BASE_CONFIG, "backend": "httpx"}) + + +def _mock_openai_response(payload: dict[str, Any]) -> MagicMock: + resp = MagicMock() + resp.model_dump.return_value = payload + return resp + + +# --------------------------------------------------------------------------- +# Initialization +# --------------------------------------------------------------------------- + + +class TestInitialization: + def test_basic_init(self, provider): + assert provider.get_name() == "friendli" + assert provider.default_model == "zai-org/GLM-5.3" + assert provider._endpoint_type == "serverless" + assert provider._base_url == SERVERLESS_URL + assert provider._backend == "openai" + + def test_instance_name_override(self): + with patch("llmcore.providers.friendli_provider.AsyncOpenAI"): + p = FriendliProvider({**BASE_CONFIG, "backend": "openai", "_instance_name": "my-friendli"}) + assert p.get_name() == "my-friendli" + + def test_dedicated_base_url(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "endpoint_type": "dedicated"}) + assert p._base_url == DEDICATED_URL + + def test_container_requires_base_url(self): + with pytest.raises(ConfigError, match="requires an explicit base_url"): + FriendliProvider({**BASE_CONFIG, "backend": "httpx", "endpoint_type": "container"}) + + def test_container_without_key_uses_placeholder(self): + p = FriendliProvider( + {"backend": "httpx", "endpoint_type": "container", "base_url": "http://localhost:8000/v1"} + ) + assert p._api_key == "EMPTY" + assert p._base_url == "http://localhost:8000/v1" + + def test_missing_key_raises_for_hosted(self): + with pytest.raises(ConfigError, match="Friendli API key not found"): + FriendliProvider({"backend": "httpx", "default_model": "m"}) + + def test_invalid_endpoint_type_falls_back(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "endpoint_type": "bogus"}) + assert p._endpoint_type == "serverless" + + def test_base_url_trailing_slash_stripped(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "base_url": f"{SERVERLESS_URL}/"}) + assert p._base_url == SERVERLESS_URL + + def test_invalid_reasoning_effort_ignored(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "reasoning_effort": "turbo"}) + assert p._default_reasoning_effort is None + + def test_valid_reasoning_effort_kept(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "reasoning_effort": "MAX"}) + assert p._default_reasoning_effort == "max" + + +class TestCredentialResolution: + """API key and team ID resolve from config then the documented env vars.""" + + @pytest.mark.parametrize( + "env_var", ["FRIENDLI_TOKEN", "FRIENDLIAI_API_KEY", "FRIENDLI_API_KEY"] + ) + def test_api_key_env_vars(self, monkeypatch, env_var): + monkeypatch.setenv(env_var, f"flp_from_{env_var}") + p = FriendliProvider({"backend": "httpx"}) + assert p._api_key == f"flp_from_{env_var}" + + def test_api_key_env_var_indirection(self, monkeypatch): + monkeypatch.setenv("MY_CUSTOM_KEY", "flp_custom") + p = FriendliProvider({"backend": "httpx", "api_key_env_var": "MY_CUSTOM_KEY"}) + assert p._api_key == "flp_custom" + + def test_explicit_key_wins(self, monkeypatch): + monkeypatch.setenv("FRIENDLI_TOKEN", "flp_env") + p = FriendliProvider({"backend": "httpx", "api_key": "flp_explicit"}) + assert p._api_key == "flp_explicit" + + def test_token_preferred_over_api_key_env(self, monkeypatch): + monkeypatch.setenv("FRIENDLI_TOKEN", "flp_token") + monkeypatch.setenv("FRIENDLIAI_API_KEY", "flp_aikey") + p = FriendliProvider({"backend": "httpx"}) + assert p._api_key == "flp_token" + + @pytest.mark.parametrize("env_var", ["FRIENDLI_TEAM_ID", "FRIENDLIAI_TEAM_ID"]) + def test_team_id_env_vars(self, monkeypatch, env_var): + monkeypatch.setenv(env_var, "team-123") + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx"}) + assert p._team_id == "team-123" + assert p._team_headers()["X-Friendli-Team"] == "team-123" + + def test_team_id_absent(self, provider): + assert provider._team_id is None + assert "X-Friendli-Team" not in provider._team_headers() + + def test_explicit_team_id_wins(self, monkeypatch): + monkeypatch.setenv("FRIENDLI_TEAM_ID", "env-team") + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "team_id": "cfg-team"}) + assert p._team_id == "cfg-team" + + +class TestBackendResolution: + def test_auto_prefers_openai(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=True, + httpx_available=True, + friendli_sdk_available=True, + ): + assert FriendliProvider._resolve_backend(None) == "openai" + assert FriendliProvider._resolve_backend("auto") == "openai" + + def test_auto_falls_back_to_httpx(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=False, + httpx_available=True, + friendli_sdk_available=True, + ): + assert FriendliProvider._resolve_backend(None) == "httpx" + + def test_auto_falls_back_to_sdk(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=False, + httpx_available=False, + friendli_sdk_available=True, + ): + assert FriendliProvider._resolve_backend(None) == "sdk" + + def test_explicit_backend_honored(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=True, + httpx_available=True, + friendli_sdk_available=True, + ): + assert FriendliProvider._resolve_backend("httpx") == "httpx" + assert FriendliProvider._resolve_backend("sdk") == "sdk" + + def test_unavailable_backend_falls_back(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=True, + httpx_available=True, + friendli_sdk_available=False, + ): + assert FriendliProvider._resolve_backend("sdk") == "openai" + + def test_unknown_backend_name_auto_detects(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=True, + httpx_available=True, + friendli_sdk_available=False, + ): + assert FriendliProvider._resolve_backend("grpc") == "openai" + + def test_no_transport_raises(self): + with patch.multiple( + "llmcore.providers.friendli_provider", + openai_available=False, + httpx_available=False, + friendli_sdk_available=False, + ): + with pytest.raises(ConfigError, match="requires one of"): + FriendliProvider(dict(BASE_CONFIG)) + + +# --------------------------------------------------------------------------- +# Request parameter handling +# --------------------------------------------------------------------------- + + +class TestRequestParams: + def test_native_vs_extras_split(self, provider): + native, extras = provider._resolve_request_params( + {"temperature": 0.5, "max_tokens": 64, "top_k": 40, "repetition_penalty": 1.1} + ) + assert native == {"temperature": 0.5, "max_tokens": 64} + assert extras["top_k"] == 40 + assert extras["repetition_penalty"] == 1.1 + + def test_parse_reasoning_default_on(self, provider): + _, extras = provider._resolve_request_params({}) + assert extras["parse_reasoning"] is True + + def test_parse_reasoning_can_be_disabled_per_request(self, provider): + _, extras = provider._resolve_request_params({"parse_reasoning": False}) + assert extras["parse_reasoning"] is False + + def test_reasoning_effort_default_applied(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "reasoning_effort": "high"}) + _, extras = p._resolve_request_params({}) + assert extras["reasoning_effort"] == "high" + + def test_reasoning_effort_per_request_override(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "reasoning_effort": "high"}) + _, extras = p._resolve_request_params({"reasoning_effort": "max"}) + assert extras["reasoning_effort"] == "max" + + def test_invalid_effort_falls_back_to_default(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "reasoning_effort": "low"}) + _, extras = p._resolve_request_params({"reasoning_effort": "nope"}) + assert extras["reasoning_effort"] == "low" + + def test_ultracode_effort_accepted(self, provider): + _, extras = provider._resolve_request_params({"reasoning_effort": "ultracode"}) + assert extras["reasoning_effort"] == "ultracode" + + def test_enable_thinking_folded_into_template_kwargs(self, provider): + _, extras = provider._resolve_request_params({"enable_thinking": True}) + assert extras["chat_template_kwargs"] == {"enable_thinking": True} + + def test_clear_thinking_folded_into_template_kwargs(self, provider): + _, extras = provider._resolve_request_params({"clear_thinking": True}) + assert extras["chat_template_kwargs"]["clear_thinking"] is True + + def test_explicit_template_kwargs_merge(self, provider): + _, extras = provider._resolve_request_params( + {"chat_template_kwargs": {"custom": 1}, "enable_thinking": False} + ) + assert extras["chat_template_kwargs"] == {"custom": 1, "enable_thinking": False} + + def test_config_enable_thinking_default(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "enable_thinking": True}) + _, extras = p._resolve_request_params({}) + assert extras["chat_template_kwargs"]["enable_thinking"] is True + + def test_list_seed_routed_to_extras(self, provider): + native, extras = provider._resolve_request_params({"seed": [1, 2, 3]}) + assert "seed" not in native + assert extras["seed"] == [1, 2, 3] + + def test_int_seed_stays_native(self, provider): + native, extras = provider._resolve_request_params({"seed": 7}) + assert native["seed"] == 7 + assert "seed" not in extras + + def test_reasoning_budget(self, provider): + _, extras = provider._resolve_request_params({"reasoning_budget": 1024}) + assert extras["reasoning_budget"] == 1024 + + def test_tools_drop_min_tokens_and_response_format(self, provider): + native = {"response_format": {"type": "json_object"}} + extras = {"min_tokens": 5} + provider._apply_mutual_exclusions(native, extras, has_tools=True) + assert "response_format" not in native + assert "min_tokens" not in extras + + def test_response_format_drops_min_tokens(self, provider): + native = {"response_format": {"type": "json_object"}} + extras = {"min_tokens": 5} + provider._apply_mutual_exclusions(native, extras, has_tools=False) + assert native["response_format"] == {"type": "json_object"} + assert "min_tokens" not in extras + + def test_no_exclusions_when_unrelated(self, provider): + native: dict[str, Any] = {"temperature": 0.2} + extras: dict[str, Any] = {"min_tokens": 5} + provider._apply_mutual_exclusions(native, extras, has_tools=False) + assert extras["min_tokens"] == 5 + + +# --------------------------------------------------------------------------- +# Message payload building +# --------------------------------------------------------------------------- + + +class TestMessagePayload: + def test_basic_user_message(self, provider): + payload = provider._build_message_payload(Message(role=Role.USER, content="Hello")) + assert payload == {"role": "user", "content": "Hello"} + + def test_tool_message(self, provider): + msg = Message(role=Role.TOOL, content='{"temp": 22}', tool_call_id="call_abc") + payload = provider._build_message_payload(msg) + assert payload["role"] == "tool" + assert payload["tool_call_id"] == "call_abc" + + def test_assistant_first_class_tool_calls(self, provider): + calls = [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}}] + msg = Message(role=Role.ASSISTANT, content="", tool_calls=calls) + payload = provider._build_message_payload(msg) + assert payload["tool_calls"] == calls + assert payload["content"] is None + + def test_first_class_tool_calls_beat_metadata(self, provider): + first_class = [{"id": "new", "type": "function", "function": {"name": "f", "arguments": "{}"}}] + legacy = [{"id": "old", "type": "function", "function": {"name": "g", "arguments": "{}"}}] + msg = Message( + role=Role.ASSISTANT, content="", tool_calls=first_class, metadata={"tool_calls": legacy} + ) + assert provider._build_message_payload(msg)["tool_calls"] == first_class + + def test_assistant_reasoning_content_preserved(self, provider): + msg = Message( + role=Role.ASSISTANT, + content="42", + metadata={"reasoning_content": "thinking..."}, + ) + assert provider._build_message_payload(msg)["reasoning_content"] == "thinking..." + + def test_inline_images(self, provider): + msg = Message( + role=Role.USER, + content="What is this?", + metadata={"inline_images": ["https://example.com/a.png"]}, + ) + parts = provider._build_message_payload(msg)["content"] + assert parts[0] == {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}} + assert parts[-1] == {"type": "text", "text": "What is this?"} + + def test_inline_audio_and_video(self, provider): + msg = Message( + role=Role.USER, + content="Describe", + metadata={ + "inline_audio": ["data:audio/wav;base64,AAA"], + "inline_videos": [{"url": "https://example.com/v.mp4"}], + }, + ) + kinds = [p["type"] for p in provider._build_message_payload(msg)["content"]] + assert kinds == ["audio_url", "video_url", "text"] + + def test_content_parts_passthrough(self, provider): + parts = [{"type": "text", "text": "hi"}] + msg = Message(role=Role.USER, content="ignored", metadata={"content_parts": parts}) + assert provider._build_message_payload(msg)["content"] == parts + + def test_unrecognized_media_entry_skipped(self, provider): + msg = Message(role=Role.USER, content="x", metadata={"inline_images": [42]}) + assert provider._build_message_payload(msg)["content"] == [{"type": "text", "text": "x"}] + + def test_name_field(self, provider): + msg = Message(role=Role.USER, content="hi", metadata={"name": "alice"}) + assert provider._build_message_payload(msg)["name"] == "alice" + + +# --------------------------------------------------------------------------- +# Chat completion — openai backend +# --------------------------------------------------------------------------- + + +class TestChatCompletionOpenAIBackend: + async def test_basic_completion(self, provider): + provider._client.chat.completions.create = AsyncMock( + return_value=_mock_openai_response( + {"choices": [{"message": {"content": "Paris"}}], "usage": {}} + ) + ) + result = await provider.chat_completion([Message(role=Role.USER, content="Capital?")]) + assert result["choices"][0]["message"]["content"] == "Paris" + + kwargs = provider._client.chat.completions.create.call_args.kwargs + assert kwargs["model"] == "zai-org/GLM-5.3" + assert kwargs["extra_body"]["parse_reasoning"] is True + + async def test_friendli_params_routed_to_extra_body(self, provider): + provider._client.chat.completions.create = AsyncMock( + return_value=_mock_openai_response({"choices": [], "usage": {}}) + ) + await provider.chat_completion( + [Message(role=Role.USER, content="Hi")], + temperature=0.3, + top_k=20, + min_p=0.05, + repetition_penalty=1.05, + reasoning_effort="max", + ) + kwargs = provider._client.chat.completions.create.call_args.kwargs + extra = kwargs["extra_body"] + assert kwargs["temperature"] == 0.3 + assert extra["top_k"] == 20 + assert extra["min_p"] == 0.05 + assert extra["repetition_penalty"] == 1.05 + assert extra["reasoning_effort"] == "max" + for key in ("top_k", "min_p", "repetition_penalty", "reasoning_effort"): + assert key not in kwargs + + async def test_tools_payload(self, provider): + provider._client.chat.completions.create = AsyncMock( + return_value=_mock_openai_response({"choices": [], "usage": {}}) + ) + tool = Tool( + name="get_weather", + description="Get weather", + parameters={"type": "object", "properties": {"city": {"type": "string"}}}, + ) + await provider.chat_completion( + [Message(role=Role.USER, content="Weather?")], tools=[tool], tool_choice="required" + ) + kwargs = provider._client.chat.completions.create.call_args.kwargs + assert kwargs["tools"][0]["function"]["name"] == "get_weather" + assert kwargs["tool_choice"] == "required" + + async def test_stream_sets_include_usage(self, provider): + async def _agen(): + yield _mock_openai_response({"choices": [{"delta": {"content": "Hi"}}]}) + + provider._client.chat.completions.create = AsyncMock(return_value=_agen()) + gen = await provider.chat_completion([Message(role=Role.USER, content="Hi")], stream=True) + chunks = [c async for c in gen] + assert chunks[0]["choices"][0]["delta"]["content"] == "Hi" + kwargs = provider._client.chat.completions.create.call_args.kwargs + assert kwargs["stream_options"] == {"include_usage": True} + + async def test_unsupported_param_raises(self, provider): + with pytest.raises(ValueError, match="Unsupported parameter"): + await provider.chat_completion([Message(role=Role.USER, content="Hi")], bogus=1) + + async def test_non_message_context_raises(self, provider): + with pytest.raises(ProviderError, match="list\\[Message\\]"): + await provider.chat_completion(["not a message"]) # type: ignore[list-item] + + async def test_empty_context_raises(self, provider): + with pytest.raises(ProviderError, match="No valid messages"): + await provider.chat_completion([]) + + +# --------------------------------------------------------------------------- +# Chat completion — httpx backend +# --------------------------------------------------------------------------- + + +class TestChatCompletionHttpxBackend: + @respx.mock + async def test_non_streaming(self, httpx_provider): + route = respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response( + 200, + json={ + "choices": [{"index": 0, "message": {"content": "Hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}, + }, + ) + ) + result = await httpx_provider.chat_completion([Message(role=Role.USER, content="Hi")]) + assert result["choices"][0]["message"]["content"] == "Hi" + + body = route.calls[0].request.content.decode() + assert '"parse_reasoning":true' in body.replace(" ", "") + await httpx_provider.close() + + @respx.mock + async def test_team_header_sent(self, monkeypatch): + monkeypatch.setenv("FRIENDLI_TEAM_ID", "team-xyz") + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx"}) + route = respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(200, json={"choices": [], "usage": {}}) + ) + await p.chat_completion([Message(role=Role.USER, content="Hi")]) + assert route.calls[0].request.headers["X-Friendli-Team"] == "team-xyz" + await p.close() + + @respx.mock + async def test_streaming_sse(self, httpx_provider): + sse = ( + 'data: {"choices":[{"delta":{"reasoning_content":"think"}}]}\n\n' + 'data: {"choices":[{"delta":{"content":"He"}}]}\n\n' + 'data: {"choices":[{"delta":{"content":"llo"}}]}\n\n' + 'data: {"choices":[],"usage":{"total_tokens":9}}\n\n' + "data: [DONE]\n\n" + ) + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response( + 200, text=sse, headers={"Content-Type": "text/event-stream"} + ) + ) + gen = await httpx_provider.chat_completion( + [Message(role=Role.USER, content="Hi")], stream=True + ) + text, reasoning, usage = "", "", None + async for chunk in gen: + text += httpx_provider.extract_delta_content(chunk) + reasoning += httpx_provider.extract_delta_reasoning_content(chunk) or "" + if chunk.get("usage"): + usage = chunk["usage"] + assert text == "Hello" + assert reasoning == "think" + assert usage == {"total_tokens": 9} + await httpx_provider.close() + + @respx.mock + async def test_streaming_error_maps_to_provider_error(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(429, json={"message": "Rate limit exceeded"}) + ) + gen = await httpx_provider.chat_completion( + [Message(role=Role.USER, content="Hi")], stream=True + ) + with pytest.raises(ProviderError, match="rate limit"): + async for _ in gen: + pass + await httpx_provider.close() + + +# --------------------------------------------------------------------------- +# Error mapping +# --------------------------------------------------------------------------- + + +class TestErrorMapping: + @respx.mock + async def test_401_auth_message(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(401, json={"detail": "Unauthorized"}) + ) + with pytest.raises(ProviderError, match="authentication failed"): + await httpx_provider.chat_completion([Message(role=Role.USER, content="Hi")]) + await httpx_provider.close() + + @respx.mock + async def test_403_mentions_team(self, monkeypatch): + monkeypatch.setenv("FRIENDLI_TEAM_ID", "team-abc") + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx"}) + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(403, json={"detail": "Forbidden"}) + ) + with pytest.raises(ProviderError, match="team-abc"): + await p.chat_completion([Message(role=Role.USER, content="Hi")]) + await p.close() + + @respx.mock + async def test_404_message(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(404, json={"detail": "Not Found"}) + ) + with pytest.raises(ProviderError, match="404"): + await httpx_provider.chat_completion([Message(role=Role.USER, content="Hi")]) + await httpx_provider.close() + + @respx.mock + async def test_429_retryable(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(429, json={"message": "Rate limit exceeded"}) + ) + with pytest.raises(ProviderError) as exc: + await httpx_provider.chat_completion([Message(role=Role.USER, content="Hi")]) + assert exc.value.status_code == 429 + assert exc.value.retryable is True + await httpx_provider.close() + + @respx.mock + async def test_context_overflow_maps_to_context_length_error(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response( + 400, json={"detail": "input is too long: prompt length exceeds context"} + ) + ) + with pytest.raises(ContextLengthError): + await httpx_provider.chat_completion([Message(role=Role.USER, content="Hi")]) + await httpx_provider.close() + + @respx.mock + async def test_plain_400_is_provider_error(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/completions").mock( + return_value=httpx.Response(400, json={"detail": "bad parameter"}) + ) + with pytest.raises(ProviderError) as exc: + await httpx_provider.chat_completion([Message(role=Role.USER, content="Hi")]) + assert not isinstance(exc.value, ContextLengthError) + await httpx_provider.close() + + +# --------------------------------------------------------------------------- +# Response extraction +# --------------------------------------------------------------------------- + + +class TestResponseExtraction: + SAMPLE: dict[str, Any] = { + "id": "chatcmpl-1", + "model": "zai-org/GLM-5.3", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": "Paris.", + "reasoning_content": "The user asks about France.", + }, + } + ], + "usage": { + "prompt_tokens": 18, + "completion_tokens": 4, + "total_tokens": 22, + "prompt_tokens_details": {"cached_tokens": 12}, + }, + } + + def test_extract_content(self, provider): + assert provider.extract_response_content(self.SAMPLE) == "Paris." + + def test_extract_reasoning_content(self, provider): + assert provider.extract_reasoning_content(self.SAMPLE) == "The user asks about France." + + def test_extract_reasoning_alias_field(self, provider): + resp = {"choices": [{"message": {"content": "x", "reasoning": "alias"}}]} + assert provider.extract_reasoning_content(resp) == "alias" + + def test_extract_reasoning_absent(self, provider): + assert provider.extract_reasoning_content({"choices": [{"message": {"content": "x"}}]}) is None + + def test_extract_usage_with_cache(self, provider): + usage = provider.extract_usage_details(self.SAMPLE) + assert usage == { + "prompt_tokens": 18, + "completion_tokens": 4, + "total_tokens": 22, + "cached_tokens": 12, + } + + def test_extract_usage_empty(self, provider): + assert provider.extract_usage_details({}) == {} + + def test_extract_finish_reason(self, provider): + assert provider.extract_finish_reason(self.SAMPLE) == "stop" + + def test_extract_tool_calls(self, provider): + resp = { + "choices": [ + { + "message": { + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"city":"Lisbon"}'}, + } + ] + } + } + ] + } + calls = provider.extract_tool_calls(resp) + assert len(calls) == 1 + assert calls[0].name == "get_weather" + assert calls[0].arguments == {"city": "Lisbon"} + + def test_extract_tool_calls_invalid_json(self, provider): + resp = { + "choices": [ + { + "message": { + "tool_calls": [ + {"id": "c", "type": "function", "function": {"name": "f", "arguments": "{"}} + ] + } + } + ] + } + assert provider.extract_tool_calls(resp)[0].arguments == {"_raw": "{"} + + def test_extract_delta_content(self, provider): + assert provider.extract_delta_content({"choices": [{"delta": {"content": "Hi"}}]}) == "Hi" + + def test_extract_delta_reasoning(self, provider): + chunk = {"choices": [{"delta": {"reasoning_content": "hm"}}]} + assert provider.extract_delta_reasoning_content(chunk) == "hm" + + def test_extraction_is_defensive(self, provider): + assert provider.extract_response_content({}) == "" + assert provider.extract_delta_content({"choices": []}) == "" + assert provider.extract_finish_reason({}) is None + assert provider.extract_tool_calls({}) == [] + + +# --------------------------------------------------------------------------- +# Model discovery and context length +# --------------------------------------------------------------------------- + + +class TestModelDiscovery: + @respx.mock + async def test_get_models_details_from_catalog(self, httpx_provider): + respx.get(f"{SERVERLESS_URL}/models").mock( + return_value=httpx.Response(200, json={"data": [CATALOG_ENTRY]}) + ) + details = await httpx_provider.get_models_details() + assert len(details) == 1 + d = details[0] + assert d.id == "zai-org/GLM-5.3-Flash" + assert d.context_length == 1_048_576 + assert d.max_output_tokens == 1_048_576 + assert d.supports_tools is True + assert d.supports_vision is True + assert d.supports_reasoning is True + assert d.metadata["base_model"] == "zhipuai/glm-5.3-flash" + assert d.metadata["endpoint_type"] == "serverless" + await httpx_provider.close() + + @respx.mock + async def test_catalog_is_cached(self, httpx_provider): + route = respx.get(f"{SERVERLESS_URL}/models").mock( + return_value=httpx.Response(200, json={"data": [CATALOG_ENTRY]}) + ) + await httpx_provider.get_models_details() + await httpx_provider.get_models_details() + assert route.call_count == 1 + await httpx_provider.close() + + @respx.mock + async def test_catalog_failure_falls_back_to_static_table(self, httpx_provider): + respx.get(f"{SERVERLESS_URL}/models").mock(return_value=httpx.Response(500, text="boom")) + details = await httpx_provider.get_models_details() + ids = {d.id for d in details} + assert "zai-org/GLM-5.3" in ids + await httpx_provider.close() + + @respx.mock + async def test_warm_up_primes_context_lengths(self, httpx_provider): + respx.get(f"{SERVERLESS_URL}/models").mock( + return_value=httpx.Response(200, json={"data": [CATALOG_ENTRY]}) + ) + await httpx_provider.warm_up() + assert httpx_provider.get_max_context_length("zai-org/GLM-5.3-Flash") == 1_048_576 + await httpx_provider.close() + + @respx.mock + async def test_warm_up_survives_failure(self, httpx_provider): + respx.get(f"{SERVERLESS_URL}/models").mock(side_effect=httpx.ConnectError("down")) + await httpx_provider.warm_up() # must not raise + await httpx_provider.close() + + async def test_dedicated_has_no_catalog(self): + p = FriendliProvider( + {**BASE_CONFIG, "backend": "httpx", "endpoint_type": "dedicated", "default_model": "ep-1"} + ) + details = await p.get_models_details() + assert [d.id for d in details] == ["ep-1"] + await p.close() + + +class TestContextLength: + def test_known_model(self, provider): + assert provider.get_max_context_length("zai-org/GLM-5.1") == 202_752 + + def test_default_model(self, provider): + assert provider.get_max_context_length() == 1_048_576 + + def test_model_card_lookup(self, provider): + # google/gemma-4-31B-it is absent from the static table but has a card. + assert provider.get_max_context_length("google/gemma-4-31B-it") == 262_144 + + def test_unknown_model_uses_fallback(self, provider): + assert provider.get_max_context_length("nobody/nothing") == 131_072 + + def test_configured_fallback_is_honored(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "fallback_context_length": 8192}) + assert p.get_max_context_length("nobody/nothing") == 8192 + + +class TestSupportedParameters: + def test_includes_reasoning_controls(self, provider): + params = provider.get_supported_parameters() + for key in ( + "reasoning_effort", + "reasoning_budget", + "parse_reasoning", + "include_reasoning", + ): + assert key in params + + def test_includes_friendli_engine_sampling(self, provider): + params = provider.get_supported_parameters() + for key in ("top_k", "min_p", "min_tokens", "repetition_penalty", "eos_token", + "xtc_threshold", "xtc_probability"): + assert key in params + + def test_effort_enum_matches_api(self, provider): + efforts = provider.get_supported_parameters()["reasoning_effort"]["enum"] + assert set(efforts) == { + "minimal", "low", "medium", "high", "xhigh", "max", "ultracode" + } + + +# --------------------------------------------------------------------------- +# Tokenization +# --------------------------------------------------------------------------- + + +class TestTokenization: + @respx.mock + async def test_tokenize(self, httpx_provider): + route = respx.post(f"{SERVERLESS_URL}/tokenize").mock( + return_value=httpx.Response(200, json={"tokens": [1, 2, 3]}) + ) + assert await httpx_provider.tokenize("hello") == [1, 2, 3] + assert route.call_count == 1 + await httpx_provider.close() + + @respx.mock + async def test_detokenize(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/detokenize").mock( + return_value=httpx.Response(200, json={"text": "hello"}) + ) + assert await httpx_provider.detokenize([1, 2, 3]) == "hello" + await httpx_provider.close() + + @respx.mock + async def test_render_chat(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/chat/render").mock( + return_value=httpx.Response(200, json={"text": "<|user|>Hi"}) + ) + rendered = await httpx_provider.render_chat([Message(role=Role.USER, content="Hi")]) + assert rendered == "<|user|>Hi" + await httpx_provider.close() + + @respx.mock + async def test_count_tokens_is_local_by_default(self, httpx_provider): + route = respx.post(f"{SERVERLESS_URL}/tokenize") + count = await httpx_provider.count_tokens("hello world") + assert count > 0 + assert route.call_count == 0 + await httpx_provider.close() + + @respx.mock + async def test_count_tokens_native_when_enabled(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "native_token_count": True}) + respx.post(f"{SERVERLESS_URL}/tokenize").mock( + return_value=httpx.Response(200, json={"tokens": [1, 2, 3, 4]}) + ) + assert await p.count_tokens("hello world") == 4 + await p.close() + + @respx.mock + async def test_native_count_falls_back_on_error(self): + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx", "native_token_count": True}) + respx.post(f"{SERVERLESS_URL}/tokenize").mock( + return_value=httpx.Response(404, json={"detail": "Not Found"}) + ) + assert await p.count_tokens("hello world") > 0 + await p.close() + + async def test_count_tokens_empty(self, httpx_provider): + assert await httpx_provider.count_tokens("") == 0 + await httpx_provider.close() + + async def test_count_message_tokens_empty(self, httpx_provider): + assert await httpx_provider.count_message_tokens([]) == 0 + await httpx_provider.close() + + @respx.mock + async def test_count_message_tokens_local(self, httpx_provider): + route = respx.post(f"{SERVERLESS_URL}/tokenize") + messages = [ + Message(role=Role.SYSTEM, content="Be terse."), + Message(role=Role.USER, content="Hi"), + ] + count = await httpx_provider.count_message_tokens(messages) + assert count > 2 * 4 + assert route.call_count == 0 + await httpx_provider.close() + + +# --------------------------------------------------------------------------- +# Endpoint-type gating for the optional media/embedding surfaces +# --------------------------------------------------------------------------- + + +class TestEndpointTypeGating: + async def test_embeddings_rejected_on_serverless(self, httpx_provider): + with pytest.raises(ProviderError, match="do not expose an embeddings endpoint"): + await httpx_provider.create_embeddings(["hello"]) + await httpx_provider.close() + + async def test_image_generation_rejected_on_serverless(self, httpx_provider): + with pytest.raises(ProviderError, match="do not expose an image-generation endpoint"): + await httpx_provider.generate_image("a cat") + await httpx_provider.close() + + @respx.mock + async def test_embeddings_on_dedicated(self): + p = FriendliProvider( + {**BASE_CONFIG, "backend": "httpx", "endpoint_type": "dedicated", "default_model": "ep-1"} + ) + respx.post(f"{DEDICATED_URL}/embeddings").mock( + return_value=httpx.Response(200, json={"data": [{"embedding": [0.1, 0.2]}]}) + ) + out = await p.create_embeddings(["hello"], encoding_format="float") + assert out["data"][0]["embedding"] == [0.1, 0.2] + await p.close() + + @respx.mock + async def test_image_generation_on_dedicated(self): + p = FriendliProvider( + {**BASE_CONFIG, "backend": "httpx", "endpoint_type": "dedicated", "default_model": "ep-1"} + ) + respx.post(f"{DEDICATED_URL}/images/generations").mock( + return_value=httpx.Response( + 200, json={"data": [{"url": "https://cdn/img.png", "seed": 7}]} + ) + ) + result = await p.generate_image("a cat", num_inference_steps=10) + assert result.images[0].url == "https://cdn/img.png" + assert result.model == "ep-1" + await p.close() + + @respx.mock + async def test_transcribe_audio(self, httpx_provider): + respx.post(f"{SERVERLESS_URL}/audio/transcriptions").mock( + return_value=httpx.Response( + 200, json={"text": "hello there", "usage": {"input_audio_length_ms": 2000}} + ) + ) + result = await httpx_provider.transcribe_audio( + b"RIFFfake", model="openai/whisper-large-v3", language="en" + ) + assert result.text == "hello there" + assert result.duration_seconds == 2.0 + await httpx_provider.close() + + +# --------------------------------------------------------------------------- +# Friendli Suite (team billing / usage) +# --------------------------------------------------------------------------- + + +class TestSuiteApis: + @respx.mock + async def test_get_team_cost(self, monkeypatch): + monkeypatch.setenv("FRIENDLI_TEAM_ID", "team-1") + p = FriendliProvider({**BASE_CONFIG, "backend": "httpx"}) + route = respx.get("https://api.friendli.ai/v1/team/cost").mock( + return_value=httpx.Response(200, json={"data": [], "has_more": False}) + ) + out = await p.get_team_cost("2026-09-01T00:00:00Z", "2026-09-02T00:00:00Z", limit=3) + assert out == {"data": [], "has_more": False} + request = route.calls[0].request + assert request.headers["X-Friendli-Team"] == "team-1" + assert "limit=3" in str(request.url) + await p.close() + + @respx.mock + async def test_get_team_usage(self, httpx_provider): + respx.get("https://api.friendli.ai/v1/team/usage").mock( + return_value=httpx.Response(200, json={"data": [{"start_time": "x"}]}) + ) + out = await httpx_provider.get_team_usage("2026-09-01T00:00:00Z", "2026-09-02T00:00:00Z") + assert out["data"][0]["start_time"] == "x" + await httpx_provider.close() + + @respx.mock + async def test_suite_error_maps(self, httpx_provider): + respx.get("https://api.friendli.ai/v1/team/cost").mock( + return_value=httpx.Response(403, json={"detail": "Forbidden"}) + ) + with pytest.raises(ProviderError): + await httpx_provider.get_team_cost("2026-09-01T00:00:00Z", "2026-09-02T00:00:00Z") + await httpx_provider.close() + + +# --------------------------------------------------------------------------- +# Official SDK backend (skipped when the vendor SDK is not installed) +# --------------------------------------------------------------------------- + + +@pytest.mark.skipif(not friendli_sdk_available, reason="the 'friendli' SDK is not installed") +class TestSdkBackend: + """The vendor SDK is optional, so these only run when it is present.""" + + @staticmethod + def _sdk_provider(**overrides: Any) -> FriendliProvider: + return FriendliProvider( + {**BASE_CONFIG, "backend": "sdk", "parse_reasoning": False, **overrides} + ) + + def test_namespaces_resolve(self): + p = self._sdk_provider() + # The generated SDK spells this resource chat_render, not chatrender. + for resource in ("chat", "token", "chat_render"): + assert p._sdk_namespace(resource) is not None + + def test_dedicated_exposes_embeddings(self): + p = self._sdk_provider(endpoint_type="dedicated", default_model="ep-1") + assert p._sdk_namespace("embeddings") is not None + + def test_unknown_resource_raises_actionable_error(self): + p = self._sdk_provider() + with pytest.raises(ProviderError, match="do not expose 'nope'"): + p._sdk_namespace("nope") + + def test_container_has_no_chat_render(self): + p = self._sdk_provider( + endpoint_type="container", base_url="http://localhost:8000/v1" + ) + with pytest.raises(ProviderError, match="do not expose 'chat_render'"): + p._sdk_namespace("chat_render") + + def test_undeclared_kwargs_are_dropped(self): + p = self._sdk_provider() + complete = p._sdk_namespace("chat").complete + kept = p._filter_sdk_kwargs( + complete, {"temperature": 0.5, "top_k": 10, "parse_reasoning": True, "bogus": 1} + ) + assert kept == {"temperature": 0.5, "top_k": 10, "parse_reasoning": True} + + def test_filter_is_a_noop_for_unintrospectable_callables(self): + p = self._sdk_provider() + payload = {"a": 1} + assert p._filter_sdk_kwargs(object(), payload) == payload + + def test_sdk_warns_when_reasoning_parsing_requested(self, caplog): + import logging as _logging + + with caplog.at_level(_logging.WARNING, logger="llmcore.providers.friendli_provider"): + FriendliProvider({**BASE_CONFIG, "backend": "sdk", "parse_reasoning": True}) + assert any("drop reasoning_content" in r.message for r in caplog.records) + + async def test_chat_via_sdk_filters_and_dispatches(self): + p = self._sdk_provider() + chat = MagicMock() + chat.complete = AsyncMock( + return_value=_mock_openai_response({"choices": [{"message": {"content": "ok"}}]}) + ) + chat.complete.__signature__ = __import__("inspect").Signature( + [ + __import__("inspect").Parameter(name, kind=__import__("inspect").Parameter.KEYWORD_ONLY) + for name in ("model", "messages", "temperature") + ] + ) + with patch.object(FriendliProvider, "_sdk_namespace", return_value=chat): + result = await p.chat_completion( + [Message(role=Role.USER, content="Hi")], temperature=0.4, top_k=5 + ) + assert result["choices"][0]["message"]["content"] == "ok" + kwargs = chat.complete.call_args.kwargs + assert kwargs["temperature"] == 0.4 + assert "top_k" not in kwargs # dropped: not in the (stubbed) signature + + +# --------------------------------------------------------------------------- +# Provider registration +# --------------------------------------------------------------------------- + + +class TestProviderRegistration: + def test_registered_in_provider_map(self): + from llmcore.providers.manager import PROVIDER_MAP + + assert PROVIDER_MAP["friendli"] is FriendliProvider + + def test_aliases_registered(self): + from llmcore.providers.manager import _PROVIDER_INSTANCE_ALIASES, PROVIDER_MAP + + for alias in ("friendliai", "friendli_ai"): + assert PROVIDER_MAP[alias] is FriendliProvider + assert _PROVIDER_INSTANCE_ALIASES[alias] == "friendli" + + def test_default_config_section_exists(self): + import tomllib + from pathlib import Path + + import llmcore + + path = Path(llmcore.__file__).parent / "config" / "default_config.toml" + with open(path, "rb") as fh: + cfg = tomllib.load(fh) + friendli = cfg["providers"]["friendli"] + assert friendli["default_model"] + assert friendli["endpoint_type"] == "serverless" + + def test_model_cards_are_registered(self): + from llmcore.model_cards.registry import get_model_card_registry + + registry = get_model_card_registry() + card = registry.get("friendli", "zai-org/GLM-5.3") + assert card is not None + assert card.get_context_length() == 1_048_576 + + +# --------------------------------------------------------------------------- +# Cleanup +# --------------------------------------------------------------------------- + + +class TestClose: + async def test_close_is_idempotent(self, provider): + provider._client.close = AsyncMock() + await provider.close() + await provider.close() + assert provider._client is None + + async def test_close_swallows_errors(self, provider): + provider._client.close = AsyncMock(side_effect=RuntimeError("boom")) + await provider.close() + assert provider._client is None diff --git a/tools/llmcore.confy-schema.json b/tools/llmcore.confy-schema.json index e1f50286..6da6c564 100644 --- a/tools/llmcore.confy-schema.json +++ b/tools/llmcore.confy-schema.json @@ -234,6 +234,11 @@ "value": "poe", "label": "Poe", "description": "Gateway to models and community bots" + }, + { + "value": "friendli", + "label": "FriendliAI", + "description": "Model APIs catalog, Dedicated Endpoints, Container" } ], "see_also": [ @@ -2563,6 +2568,246 @@ } ] }, + { + "id": "provider_friendli", + "title": "Provider: FriendliAI", + "description": "FriendliAI provider covering all three inference surfaces through one section: Friendli Model APIs (serverless, pay-per-token catalog), Friendli Dedicated Endpoints (your own GPU deployments; the model field is the endpoint ID), and Friendli Container (self-hosted Friendli Engine; base_url required). The chat endpoint is OpenAI-compatible plus Friendli extensions: reasoning_effort / reasoning_budget / parse_reasoning / include_reasoning, chat_template_kwargs (enable_thinking, clear_thinking), Friendli Engine sampling (top_k, min_p, min_tokens, repetition_penalty, eos_token, XTC), regex-constrained structured output, and an exact /tokenize endpoint. Transport is selectable via 'backend': the openai SDK in compatibility mode (default), direct httpx, or the official friendli SDK. Aliases: friendliai, friendli_ai. Install with: pip install llmcore[friendli]. Docs: https://friendli.ai/docs/llms.txt", + "icon": "⚡", + "order": 14, + "fields": [ + { + "key": "providers.friendli.api_key", + "title": "Api Key", + "type": "secret", + "description": "Friendli Personal API key (starts with 'flp_'); create one at https://friendli.ai/suite/~/setting/keys. Strongly recommended to set via an environment variable. Optional when endpoint_type = \"container\" and the container runs without auth.", + "placeholder": "flp_...", + "recommendation": "Set via FRIENDLI_TOKEN, FRIENDLIAI_API_KEY, or LLMCORE_PROVIDERS__FRIENDLI__API_KEY environment variable.", + "tags": [ + "security", + "credentials" + ], + "required": false + }, + { + "key": "providers.friendli.api_key_env_var", + "title": "Api Key Env Var", + "type": "string", + "required": false, + "description": "Environment variable holding the Friendli API key. When unset the provider checks FRIENDLI_TOKEN (the official SDK's variable), then FRIENDLIAI_API_KEY (the spelling used in friendli.ai's documentation examples), then FRIENDLI_API_KEY. type = \"friendli\" (auto-detected from section name).", + "default": "FRIENDLI_TOKEN" + }, + { + "key": "providers.friendli.team_id", + "title": "Team Id", + "type": "secret", + "description": "Friendli team to run requests as, sent as the X-Friendli-Team header and used by get_team_cost() / get_team_usage(). Falls back to FRIENDLI_TEAM_ID then FRIENDLIAI_TEAM_ID. Leave empty to use the default team in Friendli Suite.", + "placeholder": "...", + "recommendation": "Set via FRIENDLI_TEAM_ID or FRIENDLIAI_TEAM_ID environment variable.", + "tags": [ + "security", + "credentials" + ], + "required": false + }, + { + "key": "providers.friendli.team_id_env_var", + "title": "Team Id Env Var", + "type": "string", + "required": false, + "description": "Environment variable holding the Friendli team ID.", + "default": "FRIENDLI_TEAM_ID" + }, + { + "key": "providers.friendli.endpoint_type", + "title": "Endpoint Type", + "type": "enum", + "required": false, + "description": "Which Friendli surface to talk to. Selects the default base_url and determines what the 'model' field means. Embeddings and image generation are served by dedicated/container only.", + "default": "serverless", + "options": [ + { + "value": "serverless", + "label": "Model APIs (serverless catalog)" + }, + { + "value": "dedicated", + "label": "Dedicated Endpoints (model = endpoint ID)" + }, + { + "value": "container", + "label": "Friendli Container (base_url required)" + } + ] + }, + { + "key": "providers.friendli.backend", + "title": "Transport Backend", + "type": "enum", + "required": false, + "description": "Transport: \"openai\" (openai SDK pointed at the Friendli base URL; the default), \"httpx\" (direct REST), or \"sdk\" (the official friendli SDK). Leave empty for auto-detection in order openai -> httpx -> sdk. The vendor SDK is last on purpose: its generated response models drop fields outside the published schema, so reasoning_content is lost on that backend.", + "default": "", + "options": [ + { + "value": "", + "label": "Auto (openai -> httpx -> sdk)" + }, + { + "value": "openai", + "label": "OpenAI-compatibility mode (preferred)" + }, + { + "value": "httpx", + "label": "Direct REST (httpx)" + }, + { + "value": "sdk", + "label": "Official friendli SDK (drops reasoning_content)" + } + ] + }, + { + "key": "providers.friendli.base_url", + "title": "Base Url", + "type": "url", + "description": "Inference root. Leave empty to use the default for endpoint_type: https://api.friendli.ai/serverless/v1 or https://api.friendli.ai/dedicated/v1. REQUIRED for endpoint_type = \"container\" (e.g. http://localhost:8000/v1).", + "placeholder": "https://api.friendli.ai/serverless/v1", + "validation": { + "schemes": [ + "https", + "http" + ] + }, + "required": false + }, + { + "key": "providers.friendli.suite_base_url", + "title": "Suite Base Url", + "type": "url", + "description": "Friendli Suite API root used by get_team_cost() / get_team_usage(). Distinct from the inference root. Leave empty for https://api.friendli.ai/v1.", + "placeholder": "https://api.friendli.ai/v1", + "validation": { + "schemes": [ + "https", + "http" + ] + }, + "required": false + }, + { + "key": "providers.friendli.default_model", + "title": "Default Model", + "type": "string", + "required": false, + "description": "Default model. On Model APIs this is a catalog model ID (zai-org/GLM-5.3, zai-org/GLM-5.3-Flash, zai-org/GLM-5.2, zai-org/GLM-5.1, google/gemma-4-31B-it, deepseek-ai/DeepSeek-V3.2, MiniMaxAI/MiniMax-M2.5). On Dedicated Endpoints it must be the endpoint ID (or ENDPOINT_ID:ADAPTER_ROUTE for Multi-LoRA). See the live catalog at https://friendli.ai/docs/guides/model-apis/pricing.", + "default": "zai-org/GLM-5.3" + }, + { + "key": "providers.friendli.timeout", + "title": "Timeout", + "type": "integer", + "required": false, + "description": "Seconds per HTTP operation. 1M-context reasoning responses can take a while, so the default is generous.", + "default": 300, + "validation": { + "min": 1, + "max": 3600 + } + }, + { + "key": "providers.friendli.reasoning_effort", + "title": "Reasoning Effort", + "type": "enum", + "required": false, + "description": "Default reasoning effort tier. Leave empty to use each model's own default. The tiers a model accepts are advertised per-model in its /models entry (reasoning_options) and on its model card; unsupported tiers are rejected by the API.", + "default": "", + "options": [ + { + "value": "", + "label": "Model default (unset)" + }, + { + "value": "minimal", + "label": "minimal" + }, + { + "value": "low", + "label": "low" + }, + { + "value": "medium", + "label": "medium" + }, + { + "value": "high", + "label": "high" + }, + { + "value": "xhigh", + "label": "xhigh" + }, + { + "value": "max", + "label": "max" + }, + { + "value": "ultracode", + "label": "ultracode" + } + ] + }, + { + "key": "providers.friendli.reasoning_budget", + "title": "Reasoning Budget", + "type": "integer", + "required": false, + "description": "Default hard cap (in tokens) on the chain of thought. Leave unset for no cap. Only effective for reasoning models.", + "validation": { + "min": 1 + } + }, + { + "key": "providers.friendli.parse_reasoning", + "title": "Parse Reasoning", + "type": "boolean", + "required": false, + "description": "Split the chain of thought out of 'content' into 'reasoning_content', surfaced by extract_reasoning_content() and extract_delta_reasoning_content(). Requires the openai or httpx backend (the vendor SDK drops the field).", + "default": true + }, + { + "key": "providers.friendli.include_reasoning", + "title": "Include Reasoning", + "type": "boolean", + "required": false, + "description": "When parsing is on, include the parsed reasoning in the response. Leave unset to use the Friendli default (true)." + }, + { + "key": "providers.friendli.enable_thinking", + "title": "Enable Thinking", + "type": "boolean", + "required": false, + "description": "Default chat_template_kwargs.enable_thinking for controllable reasoning models (e.g. zai-org/GLM-5.2). Leave unset to let the model's chat template decide. 'clear_thinking' is accepted as a per-request kwarg." + }, + { + "key": "providers.friendli.native_token_count", + "title": "Native Token Count", + "type": "boolean", + "required": false, + "description": "Count tokens with the model's own tokenizer via POST /tokenize instead of locally (tiktoken cl100k_base). Exact, but one API request per count, which consumes the Model APIs rate-limit budget - llmcore counts tokens on every turn. provider.tokenize() / detokenize() are available regardless.", + "default": false + }, + { + "key": "providers.friendli.fallback_context_length", + "title": "Fallback Context Length", + "type": "integer", + "required": false, + "description": "Context window used when neither live catalog discovery nor a model card knows the model. Dedicated Endpoints and Containers expose no catalog, so this is the usual source there unless a model card matches.", + "default": 131072, + "validation": { + "min": 1024 + } + } + ] + }, { "id": "embedding", "title": "Embedding Providers", From e43280ed6ae51e8cebfaac52da02d3df17345c8d Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Sun, 20 Sep 2026 23:00:48 -0300 Subject: [PATCH 02/11] feat(cardctl): add FriendliAI model-card adapter MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Friendli's GET /serverless/v1/models is a rich catalog rather than the minimal OpenAI /models shape: context length, max completion tokens, per-token pricing (input/output/cache-read/cache-write/audio-minute), a functionality capability block, input/output modalities, reasoning support with the available reasoning_options, the canonical models.dev base_model, the serving mode, and the deprecation date. FriendliAdapter derives nearly every card field from that live data, so the friendli.toml enrichment overlay only carries what the API cannot know: architecture family/type for the open-weight checkpoints Friendli hosts, short display names, and aliases. Pricing is deliberately not pinned in the overlay — the live catalog is authoritative and Friendli adjusts rates. Only the hosted Model APIs catalog is discoverable; Dedicated Endpoints and Container serve a single deployment each and expose no listing endpoint. The adapter also accepts every documented key spelling (FRIENDLI_TOKEN, FRIENDLIAI_API_KEY, FRIENDLI_API_KEY), matching the provider. Registered as "friendli" with a "friendliai" alias. Co-Authored-By: Claude Opus 5 (1M context) --- tools/cardctl/adapters/__init__.py | 3 + tools/cardctl/adapters/friendli_adapter.py | 198 +++++++++++++++++++++ tools/cardctl/enrichments/friendli.toml | 46 +++++ 3 files changed, 247 insertions(+) create mode 100644 tools/cardctl/adapters/friendli_adapter.py create mode 100644 tools/cardctl/enrichments/friendli.toml diff --git a/tools/cardctl/adapters/__init__.py b/tools/cardctl/adapters/__init__.py index 86707853..dc584723 100644 --- a/tools/cardctl/adapters/__init__.py +++ b/tools/cardctl/adapters/__init__.py @@ -20,6 +20,9 @@ "kimi": "kimi_adapter.KimiAdapter", # Back-compat alias: "moonshot" → Kimi adapter (canonical key is "kimi"). "moonshot": "kimi_adapter.KimiAdapter", + "friendli": "friendli_adapter.FriendliAdapter", + # Alias: "friendliai" (brand spelling) -> Friendli. + "friendliai": "friendli_adapter.FriendliAdapter", "zai": "zai_adapter.ZaiAdapter", # Aliases: "glm" (brand) and "zhipu"/"zhipuai"/"bigmodel" (vendor) → Z.ai. "glm": "zai_adapter.ZaiAdapter", diff --git a/tools/cardctl/adapters/friendli_adapter.py b/tools/cardctl/adapters/friendli_adapter.py new file mode 100644 index 00000000..43d69040 --- /dev/null +++ b/tools/cardctl/adapters/friendli_adapter.py @@ -0,0 +1,198 @@ +# tools/cardctl/adapters/friendli_adapter.py +"""FriendliAI (Model APIs) model discovery adapter. + +Unlike most OpenAI-compatible services, Friendli's ``GET /serverless/v1/models`` +returns a *rich* catalog: context length, max completion tokens, per-token +pricing (input / output / cache read / cache write / audio minute), a +``functionality`` capability block, input/output modalities, reasoning support +with the available ``reasoning_options``, the canonical ``base_model`` id from +models.dev, the serving ``mode``, and the deprecation date. + +Almost every card field is therefore derived live; the enrichment overlay +(``enrichments/friendli.toml``) only carries curation that the API cannot know +(architecture family/type, display names, extra aliases). + +Only the hosted Model APIs catalog is discoverable. Dedicated Endpoints and +Friendli Container serve a single deployment each and expose no listing +endpoint, so cards for those are hand-written. + +Service docs: https://friendli.ai/docs/llms.txt +Catalog: https://friendli.ai/docs/guides/model-apis/pricing + +Canonical provider key: ``friendli`` (default_cards/friendli/). Also reachable +via the ``friendliai`` alias in the registry. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from .base import NormalizedModel +from .openai_compat import OpenAICompatAdapter + +logger = logging.getLogger(__name__) + +#: ``mode`` values in the catalog mapped to llmcore model-card types. +_MODE_TO_MODEL_TYPE: dict[str, str] = { + "chat": "chat", + "completion": "completion", + "embedding": "embedding", +} + +#: Vendor prefix on a Friendli model id mapped to an architecture family. +#: Friendli serves other vendors' open-weight checkpoints, so the family comes +#: from the checkpoint owner rather than from Friendli itself. +_FAMILY_BY_OWNER: dict[str, str] = { + "deepseek-ai": "DeepSeek", + "google": "Gemma", + "meta-llama": "Llama", + "minimaxai": "MiniMax", + "mistralai": "Mistral", + "moonshotai": "Kimi", + "nvidia": "Nemotron", + "openai": "GPT-OSS", + "qwen": "Qwen", + "zai-org": "GLM", +} + + +def _price_per_million(value: Any) -> float | None: + """Convert a Friendli per-token price string to USD per million tokens.""" + if value in (None, ""): + return None + try: + return round(float(value) * 1_000_000, 6) + except (TypeError, ValueError): + return None + + +class FriendliAdapter(OpenAICompatAdapter): + """Discover the Friendli Model APIs catalog and map it onto model cards.""" + + provider_name = "friendli" + api_key_env_var = "FRIENDLI_TOKEN" + base_url = "https://api.friendli.ai/serverless/v1" + + def get_api_key(self) -> str | None: + """Resolve the key, accepting every documented Friendli variable. + + ``FRIENDLI_TOKEN`` is the official SDK's variable and + ``FRIENDLIAI_API_KEY`` is the spelling used throughout friendli.ai's + own documentation examples; both are honored, as is the provider-side + ``FRIENDLI_API_KEY``. + """ + import os + + key = super().get_api_key() + if key: + return key + for name in ("FRIENDLIAI_API_KEY", "FRIENDLI_API_KEY"): + value = os.environ.get(name) + if value: + return value + return None + + def _include_model(self, model: dict[str, Any]) -> bool: + """Include every catalog entry that carries an id.""" + return bool(model.get("id")) + + def _enrich_model(self, normalized: NormalizedModel, raw: dict[str, Any]) -> NormalizedModel: + """Map the rich Friendli catalog entry onto the normalized model.""" + functionality = raw.get("functionality") or {} + inputs = raw.get("input_modalities") or [] + outputs = raw.get("output_modalities") or [] + + normalized.display_name = raw.get("name") or normalized.display_name + normalized.description = raw.get("description") + normalized.model_type = _MODE_TO_MODEL_TYPE.get(raw.get("mode", "chat"), "chat") + normalized.context_length = raw.get("context_length") + normalized.max_output_tokens = raw.get("max_completion_tokens") + + normalized.supports_tools = bool(functionality.get("tool_call")) + normalized.supports_structured_output = bool(functionality.get("structured_output")) + # Friendli enforces response_format on every structured-output model, + # which subsumes plain JSON mode. + normalized.supports_json_mode = normalized.supports_structured_output + normalized.supports_reasoning = bool(raw.get("reasoning")) + normalized.supports_vision = "image" in inputs + normalized.supports_video_input = "video" in inputs + normalized.supports_audio_input = "audio" in inputs + normalized.supports_audio_output = "audio" in outputs + + # Friendli serves open-weight checkpoints hosted on Hugging Face. + normalized.open_weights = True + owner = normalized.model_id.split("/", 1)[0].lower() if "/" in normalized.model_id else "" + family = _FAMILY_BY_OWNER.get(owner) + if family: + normalized.architecture_family = family + normalized.owned_by = ( + normalized.model_id.split("/", 1)[0] if "/" in normalized.model_id else None + ) + + deprecation = raw.get("deprecation_date") + if deprecation: + normalized.deprecation_date = str(deprecation)[:10] + normalized.is_deprecated = True + + # --- Pricing (per-token USD strings -> per-million USD floats) --- + pricing = raw.get("pricing") or {} + input_price = _price_per_million(pricing.get("input") or pricing.get("prompt")) + output_price = _price_per_million(pricing.get("output") or pricing.get("completion")) + if input_price is not None or output_price is not None: + block: dict[str, Any] = { + "input": input_price or 0.0, + "output": output_price or 0.0, + } + cached = _price_per_million(pricing.get("input_cache_read")) + if cached is not None: + block["cached_input"] = cached + normalized.raw_api_data["_pricing"] = block + + # --- Provider extension (Friendli-specific serving metadata) --- + extension: dict[str, Any] = {"endpoint_type": "serverless"} + if raw.get("base_model"): + extension["base_model"] = raw["base_model"] + if raw.get("mode"): + extension["mode"] = raw["mode"] + if raw.get("interleaved") is not None: + extension["interleaved"] = raw["interleaved"] + if raw.get("reasoning_options"): + extension["reasoning_options"] = raw["reasoning_options"] + efforts = [ + opt.get("values") + for opt in raw["reasoning_options"] + if isinstance(opt, dict) and opt.get("type") == "effort" + ] + if efforts and efforts[0]: + extension["reasoning_effort_levels"] = efforts[0] + extension["reasoning_toggle"] = any( + isinstance(opt, dict) and opt.get("type") == "toggle" + for opt in raw["reasoning_options"] + ) + if functionality: + extension["functionality"] = functionality + if raw.get("default_params"): + extension["default_params"] = raw["default_params"] + cache_write = _price_per_million(pricing.get("cache_write")) + if cache_write is not None: + extension["cache_write_per_million"] = cache_write + audio_minute = pricing.get("audio_minute") + if audio_minute not in (None, ""): + extension["audio_minute_usd"] = audio_minute + normalized.raw_api_data["_extension"] = extension + + # --- Tags --- + tags: list[str] = [f"input:{m}" for m in inputs] + tags += [f"output:{m}" for m in outputs] + if normalized.supports_reasoning: + tags.append("reasoning") + if normalized.supports_tools: + tags.append("tools") + if normalized.supports_structured_output: + tags.append("structured-output") + if (normalized.context_length or 0) >= 1_000_000: + tags.append("million-context") + normalized.tags.extend(tags) + + return normalized diff --git a/tools/cardctl/enrichments/friendli.toml b/tools/cardctl/enrichments/friendli.toml new file mode 100644 index 00000000..182a66df --- /dev/null +++ b/tools/cardctl/enrichments/friendli.toml @@ -0,0 +1,46 @@ +# tools/cardctl/enrichments/friendli.toml +# Manual enrichment overlay for the FriendliAI provider (Model APIs catalog). +# +# Friendli's GET /serverless/v1/models is unusually rich: context length, max +# completion tokens, per-token pricing, a capability block, modalities, +# reasoning options, base model, serving mode, and deprecation date all come +# from the API and are mapped by FriendliAdapter. This overlay therefore only +# carries what the API cannot know: +# +# - architecture family / type for the open-weight checkpoints Friendli hosts +# - short display names (the API repeats the full "owner/Model" id) +# - extra aliases +# +# Pricing is deliberately NOT pinned here: Friendli adjusts per-token rates and +# the live catalog is authoritative. Add a [pricing] entry only to correct the +# API (enrichment wins over live data). +# +# Canonical provider key: "friendli" (default_cards/friendli/). +# Service docs: https://friendli.ai/docs/llms.txt +# Catalog/pricing: https://friendli.ai/docs/guides/model-apis/pricing +# Last updated: 2026-09-20 + +[architecture] +"zai-org/GLM-5.3" = { family = "GLM", type = "moe" } +"zai-org/GLM-5.3-Flash" = { family = "GLM", type = "moe" } +"zai-org/GLM-5.2" = { family = "GLM", type = "moe" } +"zai-org/GLM-5.1" = { family = "GLM", type = "moe" } +"deepseek-ai/DeepSeek-V3.2" = { family = "DeepSeek", type = "moe" } +"MiniMaxAI/MiniMax-M2.5" = { family = "MiniMax", type = "moe" } +"google/gemma-4-31B-it" = { family = "Gemma", type = "dense", parameter_count = "31B" } + +[overrides] +"zai-org/GLM-5.3" = { display_name = "GLM-5.3 (Friendli)" } +"zai-org/GLM-5.3-Flash" = { display_name = "GLM-5.3-Flash (Friendli)" } +"zai-org/GLM-5.2" = { display_name = "GLM-5.2 (Friendli)" } +"zai-org/GLM-5.1" = { display_name = "GLM-5.1 (Friendli)" } +"deepseek-ai/DeepSeek-V3.2" = { display_name = "DeepSeek-V3.2 (Friendli)" } +"MiniMaxAI/MiniMax-M2.5" = { display_name = "MiniMax-M2.5 (Friendli)" } +"google/gemma-4-31B-it" = { display_name = "Gemma 4 31B Instruct (Friendli)" } + +[aliases] +# Friendli model ids are case-sensitive on the wire; these aliases let callers +# use the shorter brand spelling with llm.chat(model_name=...). +"zai-org/GLM-5.3" = ["friendli/glm-5.3"] +"zai-org/GLM-5.3-Flash" = ["friendli/glm-5.3-flash"] +"zai-org/GLM-5.2" = ["friendli/glm-5.2"] From 749482f453634b983fc63190fbcbd35f5562ab6b Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Sun, 20 Sep 2026 23:01:03 -0300 Subject: [PATCH 03/11] modelcards update: friendli Seven cards generated from the live Friendli Model APIs catalog with `python -m tools.cardctl generate friendli` (2026-09-20): zai-org/GLM-5.3 1048576 ctx $1.26 / $3.96 per 1M zai-org/GLM-5.3-Flash 1048576 ctx $0.15 / $0.50 (text+image+video) zai-org/GLM-5.2 1048576 ctx $1.40 / $4.40 zai-org/GLM-5.1 202752 ctx $1.40 / $4.40 google/gemma-4-31B-it 262144 ctx $0.14 / $0.40 (text+image) deepseek-ai/DeepSeek-V3.2 163840 ctx $0.50 / $1.50 MiniMaxAI/MiniMax-M2.5 196608 ctx $0.30 / $1.20 Context, pricing, capabilities, modalities and per-model reasoning options come straight from the API. `cardctl diff friendli` reports no differences and all seven validate. Data only. Co-Authored-By: Claude Opus 5 (1M context) --- .../friendli/MiniMaxAI--MiniMax-M2.5.json | 78 +++++++++++++++ .../default_cards/friendli/__init__.py | 1 + .../friendli/deepseek-ai--DeepSeek-V3.2.json | 80 ++++++++++++++++ .../friendli/google--gemma-4-31B-it.json | 82 ++++++++++++++++ .../friendli/zai-org--GLM-5.1.json | 81 ++++++++++++++++ .../friendli/zai-org--GLM-5.2.json | 95 ++++++++++++++++++ .../friendli/zai-org--GLM-5.3-Flash.json | 96 +++++++++++++++++++ .../friendli/zai-org--GLM-5.3.json | 94 ++++++++++++++++++ 8 files changed, 607 insertions(+) create mode 100644 src/llmcore/model_cards/default_cards/friendli/MiniMaxAI--MiniMax-M2.5.json create mode 100644 src/llmcore/model_cards/default_cards/friendli/__init__.py create mode 100644 src/llmcore/model_cards/default_cards/friendli/deepseek-ai--DeepSeek-V3.2.json create mode 100644 src/llmcore/model_cards/default_cards/friendli/google--gemma-4-31B-it.json create mode 100644 src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.1.json create mode 100644 src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.2.json create mode 100644 src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3-Flash.json create mode 100644 src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3.json diff --git a/src/llmcore/model_cards/default_cards/friendli/MiniMaxAI--MiniMax-M2.5.json b/src/llmcore/model_cards/default_cards/friendli/MiniMaxAI--MiniMax-M2.5.json new file mode 100644 index 00000000..3320f7b8 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/MiniMaxAI--MiniMax-M2.5.json @@ -0,0 +1,78 @@ +{ + "model_id": "MiniMaxAI/MiniMax-M2.5", + "display_name": "MiniMax-M2.5 (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "MiniMax", + "architecture_type": "moe" + }, + "context": { + "max_input_tokens": 196608, + "max_output_tokens": 196608 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": false, + "audio_input": false, + "audio_output": false, + "video_input": false, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 0.3, + "output": 1.2, + "cached_input": 0.06 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-02-19" + }, + "license": null, + "open_weights": true, + "aliases": [], + "description": "Prior MiniMax coding model for agent workflows, office edits, and automation", + "tags": [ + "input:text", + "output:text", + "reasoning", + "tools", + "structured-output" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "base_model": "minimax/minimax-m2.5", + "mode": "chat", + "interleaved": "reasoning_content", + "reasoning_options": [ + { + "type": "budget_tokens", + "min": -1, + "max": 196608 + } + ], + "reasoning_toggle": false, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} diff --git a/src/llmcore/model_cards/default_cards/friendli/__init__.py b/src/llmcore/model_cards/default_cards/friendli/__init__.py new file mode 100644 index 00000000..b87b95b3 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/__init__.py @@ -0,0 +1 @@ +"""Auto-generated friendli model cards.""" diff --git a/src/llmcore/model_cards/default_cards/friendli/deepseek-ai--DeepSeek-V3.2.json b/src/llmcore/model_cards/default_cards/friendli/deepseek-ai--DeepSeek-V3.2.json new file mode 100644 index 00000000..39b91a7d --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/deepseek-ai--DeepSeek-V3.2.json @@ -0,0 +1,80 @@ +{ + "model_id": "deepseek-ai/DeepSeek-V3.2", + "display_name": "DeepSeek-V3.2 (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "DeepSeek", + "architecture_type": "moe" + }, + "context": { + "max_input_tokens": 163840, + "max_output_tokens": 163840 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": false, + "audio_input": false, + "audio_output": false, + "video_input": false, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 0.5, + "output": 1.5, + "cached_input": 0.25 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-03-07" + }, + "license": null, + "open_weights": true, + "aliases": [], + "description": "DeepSeek chat model for instruction following, coding, and analysis", + "tags": [ + "input:text", + "output:text", + "reasoning", + "tools", + "structured-output" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "mode": "chat", + "interleaved": false, + "reasoning_options": [ + { + "type": "toggle" + }, + { + "type": "budget_tokens", + "min": -1, + "max": 163840 + } + ], + "reasoning_toggle": true, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} diff --git a/src/llmcore/model_cards/default_cards/friendli/google--gemma-4-31B-it.json b/src/llmcore/model_cards/default_cards/friendli/google--gemma-4-31B-it.json new file mode 100644 index 00000000..3d26d478 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/google--gemma-4-31B-it.json @@ -0,0 +1,82 @@ +{ + "model_id": "google/gemma-4-31B-it", + "display_name": "Gemma 4 31B Instruct (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "Gemma", + "architecture_type": "dense", + "parameter_count": "31B" + }, + "context": { + "max_input_tokens": 262144, + "max_output_tokens": 262144 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": true, + "audio_input": false, + "audio_output": false, + "video_input": false, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 0.14, + "output": 0.4 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-05-01" + }, + "license": null, + "open_weights": true, + "aliases": [], + "description": "Largest Gemma 4 instruction model for open, self-hosted chat and reasoning", + "tags": [ + "input:text", + "input:image", + "output:text", + "reasoning", + "tools", + "structured-output" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "base_model": "google/gemma-4-31b-it", + "mode": "chat", + "interleaved": "reasoning_content", + "reasoning_options": [ + { + "type": "toggle" + }, + { + "type": "budget_tokens", + "min": -1, + "max": 262144 + } + ], + "reasoning_toggle": true, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} diff --git a/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.1.json b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.1.json new file mode 100644 index 00000000..c3f59610 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.1.json @@ -0,0 +1,81 @@ +{ + "model_id": "zai-org/GLM-5.1", + "display_name": "GLM-5.1 (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "GLM", + "architecture_type": "moe" + }, + "context": { + "max_input_tokens": 202752, + "max_output_tokens": 202752 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": false, + "audio_input": false, + "audio_output": false, + "video_input": false, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 1.4, + "output": 4.4, + "cached_input": 0.26 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-04-07" + }, + "license": null, + "open_weights": true, + "aliases": [], + "description": "Strong GLM coding model for agentic engineering, terminals, and repository generation", + "tags": [ + "input:text", + "output:text", + "reasoning", + "tools", + "structured-output" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "base_model": "zhipuai/glm-5.1", + "mode": "chat", + "interleaved": "reasoning_content", + "reasoning_options": [ + { + "type": "toggle" + }, + { + "type": "budget_tokens", + "min": -1, + "max": 202752 + } + ], + "reasoning_toggle": true, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} diff --git a/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.2.json b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.2.json new file mode 100644 index 00000000..e53d7512 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.2.json @@ -0,0 +1,95 @@ +{ + "model_id": "zai-org/GLM-5.2", + "display_name": "GLM-5.2 (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "GLM", + "architecture_type": "moe" + }, + "context": { + "max_input_tokens": 1048576, + "max_output_tokens": 1048576 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": false, + "audio_input": false, + "audio_output": false, + "video_input": false, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 1.4, + "output": 4.4, + "cached_input": 0.26 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-06-16" + }, + "license": null, + "open_weights": true, + "aliases": [ + "friendli/glm-5.2" + ], + "description": "Open flagship GLM for long-horizon coding agents and million-token context work", + "tags": [ + "input:text", + "output:text", + "reasoning", + "tools", + "structured-output", + "million-context" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "base_model": "zhipuai/glm-5.2", + "mode": "chat", + "interleaved": "reasoning_content", + "reasoning_options": [ + { + "type": "toggle" + }, + { + "type": "effort", + "values": [ + "high", + "max" + ] + }, + { + "type": "budget_tokens", + "min": -1, + "max": 1048576 + } + ], + "reasoning_effort_levels": [ + "high", + "max" + ], + "reasoning_toggle": true, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} diff --git a/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3-Flash.json b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3-Flash.json new file mode 100644 index 00000000..25586c90 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3-Flash.json @@ -0,0 +1,96 @@ +{ + "model_id": "zai-org/GLM-5.3-Flash", + "display_name": "GLM-5.3-Flash (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "GLM", + "architecture_type": "moe" + }, + "context": { + "max_input_tokens": 1048576, + "max_output_tokens": 1048576 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": true, + "audio_input": false, + "audio_output": false, + "video_input": true, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 0.15, + "output": 0.5, + "cached_input": 0.03 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-08-28" + }, + "license": null, + "open_weights": true, + "aliases": [ + "friendli/glm-5.3-flash" + ], + "description": "Native multimodal GLM model for efficient coding and long-horizon agent tasks", + "tags": [ + "input:text", + "input:image", + "input:video", + "output:text", + "reasoning", + "tools", + "structured-output", + "million-context" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "base_model": "zhipuai/glm-5.3-flash", + "mode": "chat", + "interleaved": "reasoning_content", + "reasoning_options": [ + { + "type": "effort", + "values": [ + "low", + "high", + "max" + ] + }, + { + "type": "budget_tokens", + "min": -1, + "max": 1048576 + } + ], + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], + "reasoning_toggle": false, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} diff --git a/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3.json b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3.json new file mode 100644 index 00000000..985afe52 --- /dev/null +++ b/src/llmcore/model_cards/default_cards/friendli/zai-org--GLM-5.3.json @@ -0,0 +1,94 @@ +{ + "model_id": "zai-org/GLM-5.3", + "display_name": "GLM-5.3 (Friendli)", + "provider": "friendli", + "model_type": "chat", + "architecture": { + "family": "GLM", + "architecture_type": "moe" + }, + "context": { + "max_input_tokens": 1048576, + "max_output_tokens": 1048576 + }, + "capabilities": { + "streaming": true, + "function_calling": true, + "tool_use": true, + "json_mode": true, + "structured_output": true, + "vision": false, + "audio_input": false, + "audio_output": false, + "video_input": false, + "reasoning": true + }, + "pricing": { + "currency": "USD", + "per_million_tokens": { + "input": 1.26, + "output": 3.96, + "cached_input": 0.234 + } + }, + "lifecycle": { + "status": "active", + "release_date": "2026-08-28" + }, + "license": null, + "open_weights": true, + "aliases": [ + "friendli/glm-5.3" + ], + "description": "Flagship GLM model for long-horizon coding, agents, and complex project delivery", + "tags": [ + "input:text", + "output:text", + "reasoning", + "tools", + "structured-output", + "million-context" + ], + "source": "generated", + "provider_extension": { + "endpoint_type": "serverless", + "base_model": "zhipuai/glm-5.3", + "mode": "chat", + "interleaved": "reasoning_content", + "reasoning_options": [ + { + "type": "effort", + "values": [ + "low", + "high", + "max" + ] + }, + { + "type": "budget_tokens", + "min": -1, + "max": 1048576 + } + ], + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], + "reasoning_toggle": false, + "functionality": { + "tool_call": true, + "parallel_tool_call": true, + "structured_output": true, + "tool_choice": true, + "system_messages": true + }, + "default_params": { + "repetition_penalty": 1.0, + "temperature": 1.0, + "top_p": 1.0, + "min_p": 0.0, + "top_k": 0 + } + } +} From 97aa1e61cd6885147eec1735da86f6f05148c8b6 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Sun, 20 Sep 2026 23:01:14 -0300 Subject: [PATCH 04/11] fix(providers): construct ContextLengthError with its real signature MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ContextLengthError takes (model_name, limit, actual, message), but the OpenAI, DeepSeek and Z.ai providers each constructed it with a keyword set it has never accepted (provider_name / model / max_tokens / requested_tokens). The raise statement therefore blew up inside __init__: TypeError: ContextLengthError.__init__() got an unexpected keyword argument 'provider_name' Every context-overflow response produced an opaque TypeError carrying no model and no limit, and nothing catching ContextLengthError — including llmcore's own context-management and agent retry paths — ever saw it. Fixing OpenAIProvider also fixes its subclasses (DeepInfra, vLLM, Poe, OpenRouter). Anthropic, Mistral, Gemini, Kimi and Friendli already used the documented signature. actual=0 is the faithful translation of the requested_tokens=None all three were passing. The defect survived because no test touched those branches, so add tests/providers/test_context_length_error_mapping.py with two independent guards: - A static AST check over src/llmcore asserting that every ContextLengthError(...) call site uses keywords the constructor accepts. It is import-free, so it covers providers with no error-path tests and any added later — this is the guard that would have caught the bug. - Behavioural tests driving the real chat_completion() failure path of each fixed provider, asserting the mapped exception carries the model name and the model's context limit, plus negative cases (a plain 400, a 401) that must not become ContextLengthError. Both guards were verified to fail against the pre-fix code. The behavioural tests skip with an explicit reason when tests/providers/test_openai_provider.py has already replaced the openai package in sys.modules with MagicMock placeholders — in that case the provider binds a mock exception class no except clause can match. The static check still runs unconditionally. Co-Authored-By: Claude Opus 5 (1M context) --- src/llmcore/providers/deepseek_provider.py | 7 +- src/llmcore/providers/openai_provider.py | 7 +- src/llmcore/providers/zai_provider.py | 7 +- .../test_context_length_error_mapping.py | 258 ++++++++++++++++++ 4 files changed, 267 insertions(+), 12 deletions(-) create mode 100644 tests/providers/test_context_length_error_mapping.py diff --git a/src/llmcore/providers/deepseek_provider.py b/src/llmcore/providers/deepseek_provider.py index 9a49f904..49ee3f42 100644 --- a/src/llmcore/providers/deepseek_provider.py +++ b/src/llmcore/providers/deepseek_provider.py @@ -691,10 +691,9 @@ async def stream_wrapper() -> AsyncGenerator[dict[str, Any], None]: if status == 400 and "context_length" in msg.lower(): raise ContextLengthError( - provider_name=self.get_name(), - model=model_name, - max_tokens=self.get_max_context_length(model_name), - requested_tokens=None, + model_name=model_name, + limit=self.get_max_context_length(model_name), + actual=0, message=msg, ) if status == 400 and any( diff --git a/src/llmcore/providers/openai_provider.py b/src/llmcore/providers/openai_provider.py index e79a376d..d8a0d620 100644 --- a/src/llmcore/providers/openai_provider.py +++ b/src/llmcore/providers/openai_provider.py @@ -554,10 +554,9 @@ async def stream_wrapper(): logger.error(f"OpenAI status error ({status}): {msg}", exc_info=True) if status == 400 and "context_length" in msg.lower(): raise ContextLengthError( - provider_name=self.get_name(), - model=model_name, - max_tokens=self.get_max_context_length(model_name), - requested_tokens=None, + model_name=model_name, + limit=self.get_max_context_length(model_name), + actual=0, message=msg, ) # Detect model-not-found errors and provide actionable info. diff --git a/src/llmcore/providers/zai_provider.py b/src/llmcore/providers/zai_provider.py index 8de9c521..f7f9d5ac 100644 --- a/src/llmcore/providers/zai_provider.py +++ b/src/llmcore/providers/zai_provider.py @@ -933,10 +933,9 @@ def _raise_chat_error(self, e: Exception, model_name: str) -> None: logger.error("Z.ai status error (%s): %s", status, msg) if status == 400 and "context" in msg.lower() and "length" in msg.lower(): raise ContextLengthError( - provider_name=self.get_name(), - model=model_name, - max_tokens=self.get_max_context_length(model_name), - requested_tokens=None, + model_name=model_name, + limit=self.get_max_context_length(model_name), + actual=0, message=msg, ) if status in (401, 403): diff --git a/tests/providers/test_context_length_error_mapping.py b/tests/providers/test_context_length_error_mapping.py new file mode 100644 index 00000000..678e9f4f --- /dev/null +++ b/tests/providers/test_context_length_error_mapping.py @@ -0,0 +1,258 @@ +# tests/providers/test_context_length_error_mapping.py +""" +Regression tests for provider -> :class:`ContextLengthError` mapping. + +``ContextLengthError`` takes ``(model_name, limit, actual, message)``. The +OpenAI, DeepSeek and Z.ai providers used to construct it with a different, +never-supported keyword set (``provider_name`` / ``model`` / ``max_tokens`` / +``requested_tokens``), so every context-overflow response raised an opaque +``TypeError`` from inside the exception constructor instead of the +``ContextLengthError`` callers catch. The bug survived because no test drove +those error branches. + +This module locks the contract down two ways: + +1. A static check over the whole package: every ``ContextLengthError(...)`` + call site must use keywords the constructor actually accepts. This covers + providers that are not exercised below, and any added later. +2. Behavioural tests that drive the real ``chat_completion()`` error path of + each previously-broken provider and assert the mapped exception carries the + model name and the model's context limit. +""" + +from __future__ import annotations + +import ast +import inspect +import pathlib +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from llmcore.exceptions import ContextLengthError +from llmcore.models import Message, Role + +SRC_ROOT = pathlib.Path(__file__).resolve().parents[2] / "src" / "llmcore" + +#: Keywords ``ContextLengthError.__init__`` actually accepts. +ACCEPTED_KWARGS = frozenset(inspect.signature(ContextLengthError).parameters) + + +# --------------------------------------------------------------------------- +# 1. Static contract: no call site may pass unsupported keywords +# --------------------------------------------------------------------------- + + +def _context_length_error_call_sites() -> list[tuple[pathlib.Path, int, set[str]]]: + """Return every ``ContextLengthError(...)`` call in the package. + + Returns: + ``(path, lineno, keyword_names)`` for each call site found. + """ + sites: list[tuple[pathlib.Path, int, set[str]]] = [] + for path in sorted(SRC_ROOT.rglob("*.py")): + try: + tree = ast.parse(path.read_text(encoding="utf-8")) + except SyntaxError: # pragma: no cover - defensive + continue + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + func = node.func + name = func.id if isinstance(func, ast.Name) else getattr(func, "attr", None) + if name != "ContextLengthError": + continue + sites.append( + (path, node.lineno, {kw.arg for kw in node.keywords if kw.arg is not None}) + ) + return sites + + +class TestContextLengthErrorCallSites: + def test_call_sites_were_found(self): + """Guard the guard: the AST scan must actually find the call sites.""" + sites = _context_length_error_call_sites() + assert len(sites) >= 5, f"expected several call sites, found {len(sites)}" + + def test_every_call_site_uses_supported_keywords(self): + """No provider may pass keywords ``ContextLengthError`` does not accept. + + This is the exact defect that broke OpenAI, DeepSeek and Z.ai: the + unsupported keywords raised ``TypeError`` at the ``raise`` statement. + """ + offenders = [ + f"{path.relative_to(SRC_ROOT)}:{lineno} -> " + f"{sorted(kwargs - ACCEPTED_KWARGS)}" + for path, lineno, kwargs in _context_length_error_call_sites() + if not kwargs <= ACCEPTED_KWARGS + ] + assert not offenders, ( + "ContextLengthError call sites with unsupported keywords " + f"(accepted: {sorted(ACCEPTED_KWARGS)}):\n " + "\n ".join(offenders) + ) + + def test_constructor_rejects_the_legacy_keywords(self): + """Pin the failure mode so the signature cannot silently absorb them.""" + with pytest.raises(TypeError): + ContextLengthError( + provider_name="openai", # type: ignore[call-arg] + model="gpt-4o", + max_tokens=128_000, + requested_tokens=None, + message="context_length exceeded", + ) + + def test_supported_keywords_populate_the_exception(self): + err = ContextLengthError( + model_name="gpt-4o", limit=128_000, actual=0, message="context_length exceeded" + ) + assert err.model_name == "gpt-4o" + assert err.limit == 128_000 + assert err.actual == 0 + assert "gpt-4o" in str(err) + + +# --------------------------------------------------------------------------- +# 2. Behavioural: drive each previously-broken provider's error path +# --------------------------------------------------------------------------- + + +def _require_real_sdk_exception(module: Any) -> type[BaseException]: + """Return *module*'s bound ``APIStatusError``, skipping if it was stubbed. + + ``tests/providers/test_openai_provider.py`` installs ``MagicMock`` + placeholders for the ``openai`` package into ``sys.modules`` at import + time. If that happens before the provider under test is first imported, + the provider binds a mock instead of the real exception class and no + ``except`` clause can ever match it. Skip explicitly in that case instead + of failing on another module's import-time side effect — the static call + site check above still guards the actual defect unconditionally. + """ + exc_cls = getattr(module, "OpenAIAPIStatusError", None) + if not (isinstance(exc_cls, type) and issubclass(exc_cls, BaseException)): + pytest.skip( + f"{module.__name__} bound a stubbed 'openai' SDK in this process; " + "the real APIStatusError class is required to drive its error path" + ) + return exc_cls + + +def _api_status_error(module: Any, message: str, status: int = 400) -> Exception: + """Build the ``APIStatusError`` instance *module*'s handler catches. + + The providers bind the class at import time, so instantiating the bound + attribute keeps the ``except`` clause matching whatever is installed. + """ + exc_cls = _require_real_sdk_exception(module) + request = httpx.Request("POST", "https://example.invalid/v1/chat/completions") + response = httpx.Response(status, request=request, json={"error": {"message": message}}) + return exc_cls(message, response=response, body=None) + + +CONTEXT_OVERFLOW_MESSAGES = { + # Each provider sniffs the body with its own phrasing test. + "openai": "This model's maximum context_length is 128000 tokens.", + "deepseek": "This model's maximum context_length is 131072 tokens.", + "zai": "Input context length exceeds the model limit.", +} + +USER_TURN = [Message(role=Role.USER, content="hello" * 10)] + + +class TestOpenAIProviderMapping: + @pytest.fixture + def provider(self): + from llmcore.providers import openai_provider + + with patch.object(openai_provider, "AsyncOpenAI") as mock_cls: + mock_cls.return_value = MagicMock() + return openai_provider.OpenAIProvider( + {"api_key": "sk-test", "default_model": "gpt-4o"} + ) + + async def test_context_overflow_maps_to_context_length_error(self, provider): + from llmcore.providers import openai_provider + + provider._client.chat.completions.create = AsyncMock( + side_effect=_api_status_error(openai_provider, CONTEXT_OVERFLOW_MESSAGES["openai"]) + ) + with pytest.raises(ContextLengthError) as exc: + await provider.chat_completion(USER_TURN) + + assert exc.value.model_name == "gpt-4o" + assert exc.value.limit == provider.get_max_context_length("gpt-4o") + assert "context_length" in str(exc.value) + + async def test_other_400s_are_not_context_length_errors(self, provider): + from llmcore.exceptions import ProviderError + from llmcore.providers import openai_provider + + provider._client.chat.completions.create = AsyncMock( + side_effect=_api_status_error(openai_provider, "Invalid value for 'temperature'.") + ) + with pytest.raises(ProviderError) as exc: + await provider.chat_completion(USER_TURN) + assert not isinstance(exc.value, ContextLengthError) + + +class TestDeepSeekProviderMapping: + @pytest.fixture + def provider(self): + from llmcore.providers import deepseek_provider + + with patch.object(deepseek_provider, "AsyncOpenAI") as mock_cls: + mock_cls.return_value = MagicMock() + return deepseek_provider.DeepSeekProvider( + {"api_key": "sk-test", "default_model": "deepseek-v4-pro"} + ) + + async def test_context_overflow_maps_to_context_length_error(self, provider): + from llmcore.providers import deepseek_provider + + provider._client.chat.completions.create = AsyncMock( + side_effect=_api_status_error( + deepseek_provider, CONTEXT_OVERFLOW_MESSAGES["deepseek"] + ) + ) + with pytest.raises(ContextLengthError) as exc: + await provider.chat_completion(USER_TURN) + + assert exc.value.model_name == "deepseek-v4-pro" + assert exc.value.limit == provider.get_max_context_length("deepseek-v4-pro") + + +class TestZaiProviderMapping: + @pytest.fixture + def provider(self): + from llmcore.providers import zai_provider + + with patch.object(zai_provider, "AsyncOpenAI") as mock_cls: + mock_cls.return_value = MagicMock() + return zai_provider.ZaiProvider( + {"api_key": "test-key", "default_model": "glm-5.2", "backend": "openai"} + ) + + async def test_context_overflow_maps_to_context_length_error(self, provider): + from llmcore.providers import zai_provider + + provider._client.chat.completions.create = AsyncMock( + side_effect=_api_status_error(zai_provider, CONTEXT_OVERFLOW_MESSAGES["zai"]) + ) + with pytest.raises(ContextLengthError) as exc: + await provider.chat_completion(USER_TURN) + + assert exc.value.model_name == "glm-5.2" + assert exc.value.limit == provider.get_max_context_length("glm-5.2") + + async def test_auth_failure_is_still_a_provider_error(self, provider): + from llmcore.exceptions import ProviderError + from llmcore.providers import zai_provider + + provider._client.chat.completions.create = AsyncMock( + side_effect=_api_status_error(zai_provider, "invalid api key", status=401) + ) + with pytest.raises(ProviderError) as exc: + await provider.chat_completion(USER_TURN) + assert not isinstance(exc.value, ContextLengthError) From c22b088b6ab63ed1325d543755634bfd273a1da6 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Sun, 20 Sep 2026 23:01:24 -0300 Subject: [PATCH 05/11] docs(friendli): usage guide, config reference, example and changelog MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add docs/Friendli_provider_usage.md covering the three endpoint types, the transport backends (including the measured reason the vendor SDK is not the default — its response models drop reasoning_content), reasoning controls, tool calling and regex structured output, multimodal input, catalog and cardctl workflow, the auxiliary endpoints, token-counting trade-offs, and error/rate-limit mapping. Record the Friendli-only "ultracode" reasoning tier in docs/model_cards.md alongside the canonical vocabulary, add Friendli to the per-provider wire mapping table, and note that it has no "none" tier (reasoning is turned off through chat_template_kwargs.enable_thinking instead). Also note two verified Friendli behaviours callers will hit: /detokenize and /chat/render are documented but currently 404 on Model APIs (they work on Dedicated Endpoints and Container), and tier-0 rate limits are adaptive and in practice allow only a couple of requests per minute — which is why native token counting is opt-in and the example paces its calls. Co-Authored-By: Claude Opus 5 (1M context) --- CHANGELOG.md | 99 +++++++++ README.md | 17 +- docs/CONFIG_REFERENCE.md | 59 ++++++ docs/Friendli_provider_usage.md | 357 ++++++++++++++++++++++++++++++++ docs/model_cards.md | 36 +++- examples/README.md | 4 + examples/friendli_example.py | 185 +++++++++++++++++ 7 files changed, 747 insertions(+), 10 deletions(-) create mode 100644 docs/Friendli_provider_usage.md create mode 100644 examples/friendli_example.py diff --git a/CHANGELOG.md b/CHANGELOG.md index ea1b7fb6..e6b4a813 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,105 @@ All notable changes to **llmcore** are documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## Unreleased + +### Added — FriendliAI provider + +- **FriendliAI provider**: first-class `FriendliProvider` covering all three + Friendli inference surfaces through one `[providers.friendli]` section, + selected with `endpoint_type`: + `"serverless"` (Friendli Model APIs — the hosted pay-per-token catalog), + `"dedicated"` (Dedicated Endpoints; the `model` field is the **endpoint ID**, + or `ENDPOINT_ID:ADAPTER_ROUTE` for Multi-LoRA), and `"container"` + (self-hosted Friendli Engine; `base_url` required, API key optional). +- **Dual transport**: `backend = "openai" | "httpx" | "sdk"`, auto-resolving + **openai → httpx → sdk**. The vendor `friendli` SDK is supported but ranked + last on purpose: its generated response models ignore unknown fields, so + `reasoning_content` / `reasoning` are silently dropped and there is no + `extra_body` escape hatch. The provider warns at startup when `backend = + "sdk"` is combined with `parse_reasoning`. +- **Reasoning controls**: `reasoning_effort` + (`minimal|low|medium|high|xhigh|max|ultracode`), `reasoning_budget`, + `parse_reasoning`, `include_reasoning`, plus the chat-template switches + `enable_thinking` / `clear_thinking` folded into `chat_template_kwargs`. + Parsed chains of thought are surfaced by `extract_reasoning_content()` and + `extract_delta_reasoning_content()` in both streaming and non-streaming mode. +- **Friendli Engine sampling**: `top_k`, `min_p`, `min_tokens`, + `repetition_penalty`, `eos_token`, and XTC (`xtc_threshold` / + `xtc_probability`) routed through `extra_body`; mutually exclusive body + fields (`tools` vs `min_tokens`/`response_format`) are dropped with a warning + instead of 422-ing. +- **Structured output** including Friendli's `regex` `response_format`, tool + calling with first-class `Message.tool_calls` (R-2), and multimodal input via + `metadata["inline_images"|"inline_audio"|"inline_videos"|"content_parts"]`. +- **Rich model discovery**: `GET /models` reports context length, max completion + tokens, per-token pricing, a `functionality` block, modalities, reasoning + options, `base_model` and `mode`; the catalog is cached, primed by + `warm_up()`, and drives `get_max_context_length()`. +- **Auxiliary surfaces**: `tokenize()` / `detokenize()` / `render_chat()`, + `text_completion()`, `transcribe_audio()`, and — on dedicated/container only + — `create_embeddings()` / `generate_image()`, each gated with an actionable + error on the wrong endpoint type. `get_team_cost()` / `get_team_usage()` read + the Friendli Suite billing APIs for the configured team. +- **Team scoping**: `team_id` (or `FRIENDLI_TEAM_ID` / `FRIENDLIAI_TEAM_ID`) is + sent as `X-Friendli-Team` on every request. +- **Token counting**: local by default (tiktoken `cl100k_base`, then a + character-ratio estimate); `native_token_count = true` routes counts through + Friendli's exact `/tokenize` endpoint at the cost of one API request per + count, with a local fallback when that call fails. +- **Config**: new `[providers.friendli]` section in `default_config.toml` + (`api_key`/`api_key_env_var` → `FRIENDLI_TOKEN` → `FRIENDLIAI_API_KEY` → + `FRIENDLI_API_KEY`, `team_id`/`team_id_env_var`, `endpoint_type`, `backend`, + `base_url`, `suite_base_url`, `default_model`, `timeout`, the reasoning + defaults, `native_token_count`, `fallback_context_length`) and a matching + `provider_friendli` section in the confy schema. +- **Model cards**: `FriendliAdapter` for cardctl derives context, pricing, + capabilities, modalities and reasoning options straight from the live + catalog, with a `friendli.toml` enrichment overlay for architecture and + display names; seven generated cards under + `model_cards/default_cards/friendli/`. +- **Packaging**: `llmcore[friendli]` extra (`openai`, `httpx`, and the optional + `friendli` SDK), included in `llmcore[all]`; `friendli` registered in + `ProviderManager` with the `friendliai` / `friendli_ai` aliases. +- **Errors**: 401/403/404/429 map to actionable `ProviderError`s (429 marked + retryable so `chat_completion_with_retry` applies), and context-overflow + wording on 400/422 maps to `ContextLengthError`. +- **Docs, examples & tests**: `docs/Friendli_provider_usage.md`, + `examples/friendli_example.py`, README updates, and a 120-test offline suite + (`tests/providers/test_friendli_provider.py`) covering credential/backend + resolution, parameter splitting, payload building, both direct backends, + discovery, tokenization, endpoint gating, Suite APIs and error mapping. + +### Fixed — `ContextLengthError` construction in three providers + +- **OpenAI, DeepSeek and Z.ai raised `TypeError` instead of + `ContextLengthError` on every context overflow.** All three constructed the + exception with a keyword set it has never accepted + (`provider_name` / `model` / `max_tokens` / `requested_tokens`) rather than + its real signature `(model_name, limit, actual, message)`, so the `raise` + statement itself blew up inside `__init__`. Callers catching + `ContextLengthError` — including llmcore's own context-management and agent + retry paths — never saw it, and the user got an opaque `TypeError` with no + model or limit attached. Fixing `OpenAIProvider` also fixes its subclasses + (DeepInfra, vLLM, Poe, OpenRouter). Anthropic, Mistral, Gemini, Kimi and the + new Friendli provider already used the correct signature. +- **Regression coverage** (`tests/providers/test_context_length_error_mapping.py`): + a static AST check asserts that *every* `ContextLengthError(...)` call site + in `src/llmcore` uses keywords the constructor accepts — covering providers + with no error-path tests and any added later — plus behavioural tests that + drive the real `chat_completion()` failure path of each fixed provider and + assert the mapped exception carries the model name and context limit. Both + guards were verified to fail against the pre-fix code. + +### Notes + +- Friendli's documented `/detokenize` and `/chat/render` routes currently + return 404 on Model APIs (verified 2026-09-20); they are implemented and work + on Dedicated Endpoints / Container, and every caller degrades gracefully. +- Model APIs rate limits are tier-based; tier 0 is "adaptive" and in practice + allows only a couple of requests per minute, which is why native token + counting is opt-in and `examples/friendli_example.py` paces its calls. + ## v0.53.0 ### Added — TypeSafe.ai (System One) provider diff --git a/README.md b/README.md index 3b922232..faf80618 100644 --- a/README.md +++ b/README.md @@ -34,7 +34,7 @@ | Category | Features | |----------|----------| -| **🔌 Multi-Provider Support** | OpenAI, Anthropic, Google Gemini, Ollama, DeepSeek, Z.ai (GLM), Mistral, Qwen, xAI, vLLM, DeepInfra, Deepgram, TypeSafe.ai (System One typed judgments) | +| **🔌 Multi-Provider Support** | OpenAI, Anthropic, Google Gemini, Ollama, DeepSeek, Z.ai (GLM), FriendliAI, Mistral, Qwen, xAI, vLLM, DeepInfra, Deepgram, TypeSafe.ai (System One typed judgments) | | **💬 Chat Interface** | Unified `chat()` API, streaming responses, tool/function calling, per-call usage via `chat_with_usage()` | | **📦 Session Management** | Persistent conversations, SQLite/PostgreSQL backends, transient sessions | | **🔍 RAG System** | ChromaDB/pgvector storage, semantic search, context injection | @@ -178,6 +178,9 @@ pip install llmcore[deepgram] # Z.ai (GLM) support pip install llmcore[zai] +# FriendliAI support (Model APIs, Dedicated Endpoints, Container) +pip install llmcore[friendli] + # TypeSafe.ai System One typed judgments (noul/choice/score; httpx only) pip install llmcore[typesafe] @@ -331,6 +334,15 @@ thinking = "enabled" # "enabled" | "disabled" reasoning_effort = "high" # none|minimal|low|medium|high|xhigh|max timeout = 300 +[providers.friendli] +# FriendliAI. API key via FRIENDLI_TOKEN (FRIENDLIAI_API_KEY also accepted). +# endpoint_type = "serverless" # "serverless" | "dedicated" | "container" +# backend = "openai" # "openai" (default) | "httpx" | "sdk" +default_model = "zai-org/GLM-5.3" +parse_reasoning = true # split reasoning into reasoning_content +# reasoning_effort = "high" # minimal|low|medium|high|xhigh|max|ultracode +timeout = 300 + [providers.typesafe] # TypeSafe.ai System One (typed judgments, NOT chat). API key via TYPESAFE_API_KEY. # Use provider.system_one(state, questions) or llm.chat(..., provider_name="typesafe", questions={...}). @@ -403,6 +415,7 @@ LLMCore supports multiple LLM providers through a unified interface: | **Ollama** | Llama 3.2/3.3, Gemma 3, Phi-3, Mistral | Streaming, Local | | **DeepSeek** | DeepSeek-R1, DeepSeek-V3.2, DeepSeek-Chat | Streaming, Reasoning | | **Z.ai (GLM)** | GLM-5.2, GLM-5.1, GLM-4.7, GLM-4.6V, CogView, CogVideoX, GLM-TTS/ASR/OCR, Embedding-3 | Streaming, Tools, Reasoning, Vision, Embeddings, Image, Video, TTS, STT, OCR, Web Search | +| **FriendliAI** | Model APIs catalog (GLM-5.3/5.3-Flash/5.2/5.1, DeepSeek-V3.2, Gemma 4 31B, MiniMax-M2.5) + your own Dedicated Endpoints / Container | Streaming, Tools, Reasoning (effort/budget/parse), Vision, Structured output incl. regex, Exact tokenizer, Embeddings & Images (dedicated), STT | | **Mistral** | Mistral Large 3 | Streaming, Tools | | **Qwen** | Qwen 3 Max, Qwen3-Coder-480B | Streaming, Tools | | **xAI** | Grok-4, Grok-4-Heavy | Streaming, Tools | @@ -804,6 +817,7 @@ Built-in model cards for: - **Ollama**: Llama 3.2/3.3, Gemma 3, Phi-3, Mistral, CodeLlama - **DeepSeek**: DeepSeek-R1, DeepSeek-V3.2 - **Z.ai (GLM)**: GLM-5.2, GLM-5.1, GLM-4.7, GLM-4.6V (vision), Embedding-3 +- **FriendliAI**: GLM-5.3, GLM-5.3-Flash, GLM-5.2, GLM-5.1, DeepSeek-V3.2, Gemma 4 31B, MiniMax-M2.5 (context, pricing, and reasoning options generated live from the Friendli catalog) - **Mistral**: Mistral Large 3 - **Qwen**: Qwen 3 Max, Qwen3-Coder - **xAI**: Grok-4, Grok-4-Heavy @@ -1062,6 +1076,7 @@ from llmcore import ( - [Search providers usage](docs/Search_providers_usage.md) - [Search providers rationale](docs/Search_providers_rationale.md) - [Deepgram provider usage](docs/Deepgram_provider_usage.md) +- [FriendliAI provider usage](docs/Friendli_provider_usage.md) - [TypeSafe.ai provider usage](docs/TypeSafe_provider_usage.md) - [`chat_with_usage` guide](docs/USAGE_chat_with_usage.md) - [Model cards](docs/model_cards.md) diff --git a/docs/CONFIG_REFERENCE.md b/docs/CONFIG_REFERENCE.md index 1ed7e750..20a114c6 100644 --- a/docs/CONFIG_REFERENCE.md +++ b/docs/CONFIG_REFERENCE.md @@ -34,6 +34,7 @@ Specifies which provider instance (from [providers] below) to use by default whe - `mistral`: Mistral AI — Mistral Large, Codestral, Magistral - `vllm`: vLLM (self-hosted) — Self-hosted vLLM inference server - `poe`: Poe — Gateway to models and community bots +- `friendli`: FriendliAI — Model APIs catalog, Dedicated Endpoints, Container ### `llmcore.default_embedding_model` @@ -709,6 +710,64 @@ The official SDK defaults to 10 s; 30 s leaves room for large states near the 64 Retries honour `Retry-After` (seconds or HTTP-date) and `retry-after-ms`; `0` disables retries. `429` (rate limit) and `529` (overloaded) are the statuses you will actually see. +## ⚡ Provider: FriendliAI + +FriendliAI provider covering all three inference surfaces through one section: **Friendli Model APIs** (serverless, pay-per-token catalog), **Friendli Dedicated Endpoints** (your own GPU deployments; the `model` field is the endpoint ID), and **Friendli Container** (self-hosted Friendli Engine; `base_url` required). The chat endpoint is OpenAI-compatible plus Friendli extensions: `reasoning_effort` / `reasoning_budget` / `parse_reasoning` / `include_reasoning`, `chat_template_kwargs` (`enable_thinking`, `clear_thinking`), Friendli Engine sampling (`top_k`, `min_p`, `min_tokens`, `repetition_penalty`, `eos_token`, XTC), regex-constrained structured output, and an exact `/tokenize` endpoint. Transport is selectable via `backend`. Aliases: `friendliai`, `friendli_ai`. Install with `pip install llmcore[friendli]`. See [FriendliAI provider usage](Friendli_provider_usage.md). + +| Key | Type | Required | Default | Description | +|-----|------|----------|---------|-------------| +| `providers.friendli.api_key` | secret | | — | Friendli Personal API key (starts with `flp_`). Strongly recommended to set via `FRIENDLI_TOKEN` or `FRIENDLIAI_API_KEY`. Optional for `endpoint_type = "container"` without auth. | +| `providers.friendli.api_key_env_var` | string | | `FRIENDLI_TOKEN` | Environment variable holding the API key. When unset the provider checks `FRIENDLI_TOKEN`, `FRIENDLIAI_API_KEY`, then `FRIENDLI_API_KEY`. | +| `providers.friendli.team_id` | secret | | — | Team to run requests as, sent as the `X-Friendli-Team` header and used by the Suite billing reads. Falls back to `FRIENDLI_TEAM_ID` then `FRIENDLIAI_TEAM_ID`. | +| `providers.friendli.team_id_env_var` | string | | `FRIENDLI_TEAM_ID` | Environment variable holding the Friendli team ID. | +| `providers.friendli.endpoint_type` | enum | | `serverless` | Which Friendli surface to talk to: `serverless`, `dedicated`, or `container`. | +| `providers.friendli.backend` | enum | | — | Transport: `openai` (default), `httpx`, or `sdk`. Empty auto-detects openai → httpx → sdk. | +| `providers.friendli.base_url` | url | | — | Inference root. Empty uses the default for `endpoint_type`; **required** for `container`. | +| `providers.friendli.suite_base_url` | url | | — | Friendli Suite API root for `get_team_cost()` / `get_team_usage()`. Empty uses `https://api.friendli.ai/v1`. | +| `providers.friendli.default_model` | string | | `zai-org/GLM-5.3` | Catalog model ID (Model APIs) or endpoint ID (Dedicated Endpoints). | +| `providers.friendli.timeout` | integer | | `300` | Seconds per HTTP operation. | +| `providers.friendli.reasoning_effort` | enum | | — | Default effort tier: `minimal`, `low`, `medium`, `high`, `xhigh`, `max`, `ultracode`. Empty leaves the model default. | +| `providers.friendli.reasoning_budget` | integer | | — | Default hard cap (tokens) on the chain of thought. | +| `providers.friendli.parse_reasoning` | boolean | | `True` | Split reasoning out of `content` into `reasoning_content`. | +| `providers.friendli.include_reasoning` | boolean | | — | With parsing on, include the parsed reasoning in the response. | +| `providers.friendli.enable_thinking` | boolean | | — | Default `chat_template_kwargs.enable_thinking` for controllable reasoning models. | +| `providers.friendli.native_token_count` | boolean | | `False` | Count tokens via the model's own tokenizer (`POST /tokenize`) instead of locally. | +| `providers.friendli.fallback_context_length` | integer | | `131072` | Context window used when neither the live catalog nor a model card knows the model. | + +### `providers.friendli.endpoint_type` + +Selects the default `base_url` and determines what the `model` field means. + +**Options:** + +- `serverless`: Friendli Model APIs — the hosted catalog; `model` is a catalog ID (default) +- `dedicated`: Dedicated Endpoints — `model` is the **endpoint ID** (or `ENDPOINT_ID:ADAPTER_ROUTE` for Multi-LoRA) +- `container`: Friendli Container — self-hosted; `base_url` is required, the API key is optional + +Embeddings (`create_embeddings`) and image generation (`generate_image`) are served by `dedicated` / `container` only. Audio transcription works on all three. + +### `providers.friendli.backend` + +**Options:** + +- `openai`: the `openai` SDK pointed at the Friendli base URL (default) +- `httpx`: direct REST calls +- `sdk`: the official `friendli` SDK + +The vendor SDK is ranked last on purpose: its generated response models ignore unknown fields, so `reasoning_content` and `reasoning` are silently dropped, and it has no `extra_body` escape hatch. The provider warns at startup when `backend = "sdk"` is combined with `parse_reasoning`. + +### `providers.friendli.default_model` + +Current Model APIs catalog: `zai-org/GLM-5.3`, `zai-org/GLM-5.3-Flash`, `zai-org/GLM-5.2`, `zai-org/GLM-5.1`, `google/gemma-4-31B-it`, `deepseek-ai/DeepSeek-V3.2`, `MiniMaxAI/MiniMax-M2.5`. The live list (with context, pricing and reasoning options) is at and is mirrored into model cards by `python -m tools.cardctl generate friendli`. + +### `providers.friendli.reasoning_effort` + +The tiers a model accepts are advertised per-model in its `/models` entry (`reasoning_options`) and on its model card (`provider_extension.reasoning_effort_levels`); unsupported tiers are rejected by the API. Override per request with `chat_completion(reasoning_effort=...)`. + +### `providers.friendli.native_token_count` + +Exact, but one API request per count — and llmcore counts tokens on every turn for context budgeting. Model APIs rate limits are tier-based (tier 0 allows only a couple of requests per minute), so this stays off by default and a failed native count falls back locally. `provider.tokenize()` / `detokenize()` are available regardless. + ## 💾 Storage: Session & Vector Persistence backends for conversation history and RAG documents. diff --git a/docs/Friendli_provider_usage.md b/docs/Friendli_provider_usage.md new file mode 100644 index 00000000..db5db475 --- /dev/null +++ b/docs/Friendli_provider_usage.md @@ -0,0 +1,357 @@ +# FriendliAI provider — usage guide + +`llmcore` ships a first-class provider for [FriendliAI](https://friendli.ai/docs), +covering all three of its inference surfaces through one configuration section: + +| `endpoint_type` | Surface | Base URL | `model` field | +|---|---|---|---| +| `serverless` (default) | **Friendli Model APIs** — the hosted, pay-per-token catalog | `https://api.friendli.ai/serverless/v1` | catalog model ID (`zai-org/GLM-5.3`) | +| `dedicated` | **Friendli Dedicated Endpoints** — your own GPU deployments | `https://api.friendli.ai/dedicated/v1` | **endpoint ID** (or `ENDPOINT_ID:ADAPTER_ROUTE` for Multi-LoRA) | +| `container` | **Friendli Container** — self-hosted Friendli Engine | your URL (required) | whatever the container serves | + +The chat endpoint is OpenAI-compatible, so everything llmcore already does — +streaming, tool calling, structured output, sessions, RAG, agents — works +unchanged. On top of that the provider exposes Friendli's own extensions: +reasoning controls, chat-template switches, Friendli Engine sampling, +regex-constrained output, an exact tokenizer endpoint, and the Suite billing API. + +--- + +## 1. Install & configure + +```bash +pip install "llmcore[friendli]" # openai + httpx (+ the optional vendor SDK) +export FRIENDLI_TOKEN="flp_..." # https://friendli.ai/suite/~/setting/keys +export FRIENDLI_TEAM_ID="..." # optional; scopes requests + billing reads +``` + +### Environment variables + +The provider accepts every spelling in circulation, checked in this order: + +| Purpose | Variables (in order) | +|---|---| +| API key | `api_key` → `api_key_env_var` → `FRIENDLI_TOKEN` → `FRIENDLIAI_API_KEY` → `FRIENDLI_API_KEY` | +| Team ID | `team_id` → `team_id_env_var` → `FRIENDLI_TEAM_ID` → `FRIENDLIAI_TEAM_ID` | + +`FRIENDLI_TOKEN` is the variable the official SDK reads; `FRIENDLIAI_API_KEY` is +the spelling used throughout friendli.ai's own documentation examples. Both work, +so no renaming is required if you already have one of them set. + +### `[providers.friendli]` + +```toml +[providers.friendli] +# api_key = "flp_..." # prefer FRIENDLI_TOKEN / FRIENDLIAI_API_KEY +# team_id = "..." # prefer FRIENDLI_TEAM_ID / FRIENDLIAI_TEAM_ID +endpoint_type = "serverless" # "serverless" | "dedicated" | "container" +# backend = "openai" # "openai" (default) | "httpx" | "sdk" +default_model = "zai-org/GLM-5.3" +timeout = 300 + +# --- Reasoning --- +# reasoning_effort = "high" # minimal|low|medium|high|xhigh|max|ultracode +# reasoning_budget = 10000 # hard cap on chain-of-thought tokens +parse_reasoning = true # split reasoning out of `content` +# include_reasoning = true +# enable_thinking = true # chat_template_kwargs.enable_thinking + +# --- Token counting --- +native_token_count = false # true => exact, but 1 API request per count +fallback_context_length = 131072 +``` + +`friendliai` and `friendli_ai` are accepted as aliases for the provider name +(`get_provider("friendliai")`). + +--- + +## 2. Transport backends — and why the vendor SDK is not the default + +`backend` selects how requests reach Friendli: + +| Backend | Library | Notes | +|---|---|---| +| `openai` | `openai` (`AsyncOpenAI`) | **Default.** Native async, battle-tested SSE, Friendli extensions travel in `extra_body`. | +| `httpx` | `httpx` | Direct REST. Same wire format, no SDK in the path. | +| `sdk` | `friendli` (`AsyncFriendli`) | The official SDK. Fully supported, but lossy — see below. | + +With `backend` unset (or `"auto"`) the provider resolves **openai → httpx → sdk** +based on what is installed. + +The vendor SDK is last on purpose. Its generated response models are strict +(`extra` is ignored), so any field Friendli returns outside the published schema +is silently dropped — including `reasoning_content` and `reasoning` on assistant +messages — and it offers no `extra_body` escape hatch for request fields it does +not declare. Verified directly: + +```python +from friendli.models import ServerlessChatCompleteSuccess +ServerlessChatCompleteSuccess.model_validate({ + ..., "choices": [{"index": 0, "finish_reason": "stop", + "message": {"role": "assistant", "content": "hi", + "reasoning_content": "THINK"}}], +}).model_dump() +# -> message == {"role": "assistant", "content": "hi"} # reasoning_content gone +``` + +The `openai` and `httpx` backends both preserve it. If you set `backend = "sdk"` +while `parse_reasoning` is on, the provider logs a warning at startup. + +--- + +## 3. Chat + +```python +from llmcore import LLMCore + +CONFIG = { + "llmcore": {"default_provider": "friendli"}, + "providers": {"friendli": {"default_model": "zai-org/GLM-5.3-Flash"}}, +} + +async with await LLMCore.create(config_overrides=CONFIG) as llm: + print(await llm.chat("What makes the Friendli Engine fast?")) + + # Streaming + async for chunk in await llm.chat("Explain speculative decoding.", stream=True): + print(chunk, end="", flush=True) +``` + +Per-request Friendli parameters go straight through `chat()` / `chat_completion()`: + +```python +await llm.chat( + "Plan a migration.", + provider_name="friendli", + reasoning_effort="max", # minimal|low|medium|high|xhigh|max|ultracode + reasoning_budget=8000, # cap the chain of thought + top_k=40, min_p=0.05, # Friendli Engine sampling + repetition_penalty=1.05, +) +``` + +--- + +## 4. Reasoning + +Friendli splits reasoning control across four body fields plus the chat template. +The provider applies your configured defaults and lets any call override them. + +| Parameter | Type | Meaning | +|---|---|---| +| `reasoning_effort` | str | How hard the model thinks. Accepted tiers vary per model — see `reasoning_options` in the model's catalog entry. | +| `reasoning_budget` | int | Hard cap (tokens) on the chain of thought. | +| `parse_reasoning` | bool | Split the chain of thought out of `content` into `reasoning_content`. | +| `include_reasoning` | bool | With parsing on, include the parsed reasoning in the response. | +| `enable_thinking` | bool | Chat-template switch for *controllable* reasoning models (e.g. `zai-org/GLM-5.2`). Folded into `chat_template_kwargs`. | +| `clear_thinking` | bool | Chat-template switch; drop prior reasoning from the context window. | + +Reading the parsed reasoning back: + +```python +provider = llm._provider_manager.get_provider("friendli") + +resp = await provider.chat_completion( + [Message(role=Role.USER, content="Is 8191 prime? Think it through.")], + reasoning_effort="high", +) +print(provider.extract_response_content(resp)) # the answer +print(provider.extract_reasoning_content(resp)) # the chain of thought + +# Streaming: reasoning arrives as its own delta field +async for chunk in await provider.chat_completion(msgs, stream=True): + text = provider.extract_delta_content(chunk) + think = provider.extract_delta_reasoning_content(chunk) +``` + +A model's supported tiers are discoverable — `get_models_details()` puts the raw +`reasoning_options` in `ModelDetails.metadata`, and the generated model cards +carry `provider_extension.reasoning_effort_levels`. + +--- + +## 5. Tool calling and structured output + +Tool calling is OpenAI-shaped and works through llmcore's unified `Tool` / +`ToolCall` types, including the full assistant-`tool_calls` → `role="tool"` +round trip. + +```python +tool = Tool(name="get_weather", description="Get the weather for a city.", + parameters={"type": "object", "properties": {"city": {"type": "string"}}, + "required": ["city"]}) + +resp = await provider.chat_completion(msgs, tools=[tool], tool_choice="required") +calls = provider.extract_tool_calls(resp) # [ToolCall(name='get_weather', ...)] +``` + +`response_format` accepts `json_schema`, `json_object`, `text`, and Friendli's +**`regex`** variant: + +```python +await provider.chat_completion( + msgs, + response_format={"type": "regex", "schema": r"[A-Z][a-z]+, [A-Z]{2}"}, +) +``` + +Friendli rejects some combinations outright — `min_tokens` and `response_format` +alongside `tools`, and `min_tokens` alongside `response_format`. The provider +drops the offending field with a warning rather than letting the request 422. + +--- + +## 6. Multimodal input + +Vision/audio/video models take content parts. Supply them through message +metadata and the provider assembles the payload: + +```python +Message( + role=Role.USER, + content="What is in this image?", + metadata={"inline_images": ["https://example.com/photo.png"]}, +) +# also: inline_audio, inline_videos, or a ready-made content_parts list +``` + +Entries may be HTTPS URLs or base64 data URIs. Which modalities a model accepts +is in its catalog entry (`input_modalities`) and on its model card +(`capabilities.vision` / `audio_input` / `video_input`). + +--- + +## 7. Model discovery and cost data + +Model APIs exposes an unusually rich catalog, which the provider maps onto +`ModelDetails` and `cardctl` turns into model cards: + +```python +for d in await provider.get_models_details(): + print(d.id, d.context_length, d.supports_reasoning, d.metadata["pricing"]) +``` + +Regenerate the bundled cards after Friendli changes the catalog: + +```bash +python -m tools.cardctl generate friendli # context, pricing, capabilities, reasoning options +python -m tools.cardctl diff friendli # read-only comparison vs the live API +python -m tools.cardctl validate friendli +``` + +Pricing comes straight from the API (per-token USD, converted to per-million), +so cards stay accurate without manual curation. The overlay in +`tools/cardctl/enrichments/friendli.toml` only carries what the API cannot know: +architecture family/type, display names, aliases. + +Dedicated Endpoints and Containers serve a single deployment each and have no +catalog, so `get_models_details()` describes the configured model from its card. + +--- + +## 8. Auxiliary endpoints + +```python +provider = llm._provider_manager.get_provider("friendli") + +# Exact tokenization with the model's own tokenizer +tokens = await provider.tokenize("What is generative AI?") # [3838, 374, ...] + +# Raw (non-chat) completions — the chat template is NOT applied +resp = await provider.text_completion("Once upon a time", max_tokens=64) + +# Audio transcription (all three endpoint types) +result = await provider.transcribe_audio( + "/path/to/audio.mp3", model="openai/whisper-large-v3", language="en" +) + +# Embeddings and image generation (Dedicated Endpoints / Container only) +vectors = await provider.create_embeddings(["hello"]) +image = await provider.generate_image("an orange Lamborghini", num_inference_steps=10) + +# Friendli Suite billing (uses the team ID / X-Friendli-Team header) +cost = await provider.get_team_cost("2026-09-01T00:00:00Z", "2026-09-20T00:00:00Z") +usage = await provider.get_team_usage("2026-09-01T00:00:00Z", "2026-09-20T00:00:00Z") +``` + +`detokenize()` and `render_chat()` are implemented against the documented +`/detokenize` and `/chat/render` routes, but those currently return **404 on +Model APIs** (verified 2026-09-20); use them on Dedicated Endpoints or a +Container. Everything that depends on them degrades gracefully. + +--- + +## 9. Token counting + +`count_tokens()` and `count_message_tokens()` are **local by default** +(tiktoken `cl100k_base`, then a character-ratio estimate). Set +`native_token_count = true` to route them through Friendli's `/tokenize` +endpoint for exact, per-model counts. + +The trade-off is rate limit, not accuracy: llmcore counts tokens on every turn +for context budgeting, and each native count is one API request. Model APIs +limits scale with your usage tier — tier 0 is only a couple of requests per +minute — so native counting is opt-in. `tokenize()` / `detokenize()` remain +available regardless of the setting, and a failed native count falls back +locally rather than raising. + +For an exact, template-accurate prompt count on a surface that serves +`/chat/render`: + +```python +exact = len(await provider.tokenize(await provider.render_chat(messages))) +``` + +--- + +## 10. Dedicated Endpoints and Container + +```toml +[providers.friendli_dedicated] +type = "friendli" +endpoint_type = "dedicated" +default_model = "YOUR_ENDPOINT_ID" # not a model name +# default_model = "YOUR_ENDPOINT_ID:adapter-route" # Multi-LoRA + +[providers.friendli_local] +type = "friendli" +endpoint_type = "container" +base_url = "http://localhost:8000/v1" # REQUIRED +# api_key is optional when the container has no auth +default_model = "meta-llama/Llama-3.1-8B-Instruct" +``` + +Both sections coexist with `[providers.friendli]`; pick one per call with +`provider_name=`. + +--- + +## 11. Errors and rate limits + +| Status | Mapped to | Message hints | +|---|---|---| +| 400/422 with overflow wording | `ContextLengthError` | carries the model's context limit | +| 401 | `ProviderError` | which env vars to check; keys start with `flp_` | +| 403 | `ProviderError` | names the active `X-Friendli-Team` scope | +| 404 | `ProviderError` | endpoint-ID vs model-name confusion, and the routes that 404 on Model APIs | +| 429 | `ProviderError` (`retryable=True`) | links the rate-limit tiers | +| other | `ProviderError` | status + body | + +Model APIs rate limits are tier-based and rise with lifetime spend; tier 0 gets +"adaptive" limits that in practice are a couple of requests per minute, and the +responses carry `x-ratelimit-limit-requests` / `x-ratelimit-remaining-requests` / +`x-ratelimit-reset-requests`. Since 429 is marked retryable, llmcore's +`chat_completion_with_retry` policy applies on top. + +--- + +## 12. References + +- Documentation index: +- OpenAI compatibility: +- Chat completions API: +- Reasoning: +- Structured outputs: +- Models & pricing: +- Rate limits: diff --git a/docs/model_cards.md b/docs/model_cards.md index 1a217cb5..43377b24 100644 --- a/docs/model_cards.md +++ b/docs/model_cards.md @@ -151,6 +151,9 @@ src/llmcore/model_cards/default_cards/ │ ├── claude-sonnet-4-5-20250929.json │ └── ... ├── deepseek/ +├── friendli/ +│ ├── zai-org--GLM-5.3.json # generated live from the Friendli catalog +│ └── ... ├── google/ ├── mistral/ ├── ollama/ @@ -597,6 +600,11 @@ providers, ordered weakest to strongest: none | minimal | low | medium | high | xhigh | max ``` +FriendliAI additionally exposes a coding-specialised tier, `ultracode`, above +`max`. It sits outside the ordered scale above (it is a *mode*, not a rung), so +it is accepted verbatim by the Friendli provider and is not mapped onto — or +from — the canonical tiers by any other provider. + Model cards, config files, and per-call `reasoning_effort` kwargs always use canonical spellings. Each provider is responsible for translating canonical values to whatever its wire protocol expects — **inside the provider module, @@ -608,15 +616,15 @@ the tiers OpenAI-family reasoning models expose. #### Per-provider wire mapping (as implemented) -| Canonical | OpenAI (`openai_provider.py`) | DeepSeek (`deepseek_provider.py`) | Z.ai (`zai_provider.py`) | -|-----------|-------------------------------|-----------------------------------|--------------------------| -| `none` | — (not advertised) | `high` (fold-to-default) | `none` | -| `minimal` | — (not advertised) | `high` (fold-to-default) | `minimal` | -| `low` | `low` | `high` | `low` | -| `medium` | `medium` | `high` | `medium` | -| `high` | `high` | `high` | `high` | -| `xhigh` | `xhigh` (unchanged on wire) | `max` | `xhigh` | -| `max` | — (not advertised) | `max` | `max` | +| Canonical | OpenAI (`openai_provider.py`) | DeepSeek (`deepseek_provider.py`) | Z.ai (`zai_provider.py`) | Friendli (`friendli_provider.py`) | +|-----------|-------------------------------|-----------------------------------|--------------------------|-----------------------------------| +| `none` | — (not advertised) | `high` (fold-to-default) | `none` | — (omit the field instead) | +| `minimal` | — (not advertised) | `high` (fold-to-default) | `minimal` | `minimal` | +| `low` | `low` | `high` | `low` | `low` | +| `medium` | `medium` | `high` | `medium` | `medium` | +| `high` | `high` | `high` | `high` | `high` | +| `xhigh` | `xhigh` (unchanged on wire) | `max` | `xhigh` | `xhigh` | +| `max` | — (not advertised) | `max` | `max` | `max` | Provider notes (verified against the code): @@ -634,6 +642,16 @@ Provider notes (verified against the code): lowercases input, and passes it through **verbatim** (GLM-5.2+). Invalid values fall back to the configured default (`high`). Only sent when thinking mode is enabled. +- **FriendliAI** — accepts `minimal | low | medium | high | xhigh | max` + plus the Friendli-only `ultracode`, lowercases input, and passes it through + **verbatim** via `extra_body`. There is no `none`: omitting + `reasoning_effort` (the default when neither the config nor the call sets + one) leaves the model's own default in place, and reasoning is turned off + through `chat_template_kwargs.enable_thinking` on controllable models. + Invalid values fall back to the configured default. Which tiers a given + model actually accepts is advertised per-model in the catalog + (`reasoning_options`) and mirrored onto its card as + `provider_extension.reasoning_effort_levels`. #### Adding a new provider diff --git a/examples/README.md b/examples/README.md index f86638fd..73f5b7d8 100644 --- a/examples/README.md +++ b/examples/README.md @@ -43,6 +43,10 @@ small calls and skip providers whose required environment is incomplete. - `hosted_providers_example.py` - Poe, OpenRouter, DeepSeek, Kimi, DeepInfra, and Mistral from one LLMCore instance. - `zai_example.py` - Z.ai (GLM) chat, thinking mode, and streaming. +- `friendli_example.py` - FriendliAI Model APIs: catalog discovery, chat, + reasoning controls with parsed chain of thought, streaming, a tool-calling + round trip, the native tokenizer, and Suite team usage. Requires + `FRIENDLI_TOKEN` (or `FRIENDLIAI_API_KEY`). - `typesafe_example.py` - TypeSafe.ai System One typed judgments (`system_one()` with Noul/Choice/Score questions, confidence-gated routing, and the `llm.chat(..., questions=...)` bridge). Requires `TYPESAFE_API_KEY`. diff --git a/examples/friendli_example.py b/examples/friendli_example.py new file mode 100644 index 00000000..94531f88 --- /dev/null +++ b/examples/friendli_example.py @@ -0,0 +1,185 @@ +# examples/friendli_example.py +""" +Example demonstrating the FriendliAI provider in LLMCore. + +This script shows how to: +1. Initialize LLMCore with a Friendli provider configured at runtime. +2. Send a chat request to a Friendli Model APIs model. +3. Control reasoning (``reasoning_effort`` / ``reasoning_budget``) and read the + parsed chain of thought back out of the response. +4. Stream a response token-by-token. +5. Call a tool and feed the result back for a second turn. +6. Use the Friendli-specific surfaces: catalog discovery, exact tokenization, + and the Suite team-usage API. + +Transport: the provider prefers direct transports — the ``openai`` SDK pointed at +the Friendli base URL (default) or raw ``httpx`` — over the official ``friendli`` +SDK, whose generated response models drop ``reasoning_content``. Pick one +explicitly with ``backend = "openai" | "httpx" | "sdk"``. + +To run this example: +- Install with the Friendli extra: ``pip install llmcore[friendli]`` +- Set a Friendli Personal API key (https://friendli.ai/suite/~/setting/keys): + export FRIENDLI_TOKEN='flp_...' # FRIENDLIAI_API_KEY also works +- Optionally scope requests and billing reads to a team: + export FRIENDLI_TEAM_ID='...' # FRIENDLIAI_TEAM_ID also works + +NOTE: Friendli Model APIs rate limits are tier-based and tier 0 allows only a +couple of requests per minute, so this example paces its calls. + +Docs: https://friendli.ai/docs/llms.txt +""" + +import asyncio +import json +import logging + +from llmcore import ConfigError, LLMCore, LLMCoreError, ProviderError +from llmcore.models import Message, Role, Tool + +logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s") +logger = logging.getLogger(__name__) + +# Seconds to wait between chat calls so a low usage tier does not 429. +PACE_SECONDS = 35 + +# Configure the Friendli provider at runtime. The API key is picked up from +# FRIENDLI_TOKEN / FRIENDLIAI_API_KEY by the provider itself. +CONFIG_OVERRIDES = { + "llmcore": {"default_provider": "friendli"}, + "providers": { + "friendli": { + # Model APIs catalog id. For a Dedicated Endpoint use the endpoint + # ID here and set endpoint_type = "dedicated". + "default_model": "zai-org/GLM-5.3-Flash", + "endpoint_type": "serverless", + "parse_reasoning": True, + # "backend": "openai", # "openai" (default) | "httpx" | "sdk" + # "reasoning_effort": "high", + } + }, +} + + +async def main() -> None: + """Run the Friendli examples.""" + llm = None + try: + logger.info("Initializing LLMCore with the Friendli provider...") + llm = await LLMCore.create(config_overrides=CONFIG_OVERRIDES) + friendli = llm._provider_manager.get_provider("friendli") + + # --- Example 1: Catalog discovery (no generation, no rate-limit cost) --- + logger.info("\n--- Model APIs catalog ---") + for details in await friendli.get_models_details(): + pricing = details.metadata.get("pricing") or {} + logger.info( + " %-28s ctx=%-9s reasoning=%-5s in=%s", + details.id, + details.context_length, + details.supports_reasoning, + pricing.get("input"), + ) + + # --- Example 2: Standard chat --- + prompt1 = "In one sentence, what makes the Friendli Engine fast?" + logger.info(f"\n--- Prompt 1: '{prompt1}' ---") + response1 = await llm.chat(prompt1, provider_name="friendli") + logger.info(f"Friendli Response 1:\n{response1}") + + # --- Example 3: Reasoning controls + parsed chain of thought --- + await asyncio.sleep(PACE_SECONDS) + logger.info("\n--- Prompt 2 (reasoning_effort=high, budget capped) ---") + resp2 = await friendli.chat_completion( + [Message(role=Role.USER, content="Is 8191 prime? Think it through.")], + reasoning_effort="high", + reasoning_budget=2000, + max_tokens=1024, + ) + logger.info(f"Answer:\n{friendli.extract_response_content(resp2)}") + reasoning = friendli.extract_reasoning_content(resp2) or "" + logger.info(f"Reasoning ({len(reasoning)} chars):\n{reasoning[:400]}") + logger.info(f"Usage: {friendli.extract_usage_details(resp2)}") + + # --- Example 4: Streaming --- + await asyncio.sleep(PACE_SECONDS) + prompt3 = "Explain speculative decoding in two sentences." + logger.info(f"\n--- Prompt 3 (streaming): '{prompt3}' ---") + async for chunk in await llm.chat(prompt3, provider_name="friendli", stream=True): + print(chunk, end="", flush=True) + print() + + # --- Example 5: Tool calling round-trip --- + await asyncio.sleep(PACE_SECONDS) + logger.info("\n--- Tool calling ---") + weather = Tool( + name="get_weather", + description="Get the current weather for a city.", + parameters={ + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + ) + messages = [Message(role=Role.USER, content="What's the weather in Lisbon?")] + tool_resp = await friendli.chat_completion( + messages, tools=[weather], tool_choice="required", max_tokens=512 + ) + calls = friendli.extract_tool_calls(tool_resp) + logger.info(f"Tool calls: {[(c.name, c.arguments) for c in calls]}") + + if calls: + call = calls[0] + await asyncio.sleep(PACE_SECONDS) + follow_up = [ + *messages, + Message( + role=Role.ASSISTANT, + content="", + tool_calls=[ + { + "id": call.id, + "type": "function", + "function": { + "name": call.name, + "arguments": json.dumps(call.arguments), + }, + } + ], + ), + Message(role=Role.TOOL, content="18C and sunny", tool_call_id=call.id), + ] + final = await friendli.chat_completion(follow_up, max_tokens=256) + logger.info(f"After the tool result:\n{friendli.extract_response_content(final)}") + + # --- Example 6: Exact tokenization with the model's own tokenizer --- + logger.info("\n--- Native tokenizer ---") + tokens = await friendli.tokenize("What is generative AI?") + logger.info(f"Token IDs: {tokens} ({len(tokens)} tokens)") + + # --- Example 7: Friendli Suite team usage (needs a team ID) --- + logger.info("\n--- Team usage (Friendli Suite) ---") + try: + usage = await friendli.get_team_usage( + "2026-09-01T00:00:00Z", "2026-09-20T00:00:00Z", limit=3 + ) + logger.info(f"Usage buckets: {len(usage.get('data', []))}") + except ProviderError as e: + logger.warning(f"Team usage unavailable: {e}") + + except ConfigError as e: + logger.error(f"Configuration error: {e}") + except ProviderError as e: + logger.error(f"Friendli provider error (is FRIENDLI_TOKEN set?): {e}") + except LLMCoreError as e: + logger.error(f"An LLMCore error occurred: {e}") + except Exception as e: + logger.exception(f"An unexpected error occurred: {e}") + finally: + if llm: + logger.info("Closing LLMCore resources...") + await llm.close() + + +if __name__ == "__main__": + asyncio.run(main()) From f5a55fa7926e8441d76371cd854dce48f6f23d39 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Tue, 29 Sep 2026 23:21:57 -0300 Subject: [PATCH 06/11] docs(providers): add provider support matrix and modernization plan MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Audit every curated provider against its upstream SDK/API and record the result in two tracking documents. docs/PROVIDER_SUPPORT_MATRIX.md is the ongoing tracker: per provider, the vendor SDK clone in /av/avalon/xrepos with its tag, commit and date, our pyproject pin, the installed version, the transport shape, and a capability matrix (chat/stream/tools/structured/reasoning/vision/audio/image/video/ embeddings/OCR/search/tokenizer) extracted from the provider classes rather than assumed. Section 6 is a runnable refresh procedure so the document can be regenerated per release. docs/PROVIDER_MODERNIZATION_PLAN.md turns the gaps into a phased program, starting from the dual-transport and one-contract principles. Findings that drove the plan: - openai (2.31 pin vs 3.22.1), anthropic (0.94 vs 1.9.0) and google-genai (1.72 vs 2.25.0) are each a MAJOR version behind. - openai 3.x and anthropic 1.x moved to httpx2 and no longer install httpx. Verified we never hand httpx objects to those clients and that no respx test routes traffic through a vendor SDK, so the port is packaging-only — but six providers (mistral, kimi, poe, openrouter, vllm, huggingface) import httpx with no extra of their own and would fail at import. - httpx2 verifies against the OS trust store, not certifi: a deployment risk worth documenting. - Only 4 of 16 providers implement the full extractor contract; ollama and gemini surface reasoning under provider-specific names that callers cannot use polymorphically. - anthropic's thinking_budget_tokens config key is now rejected with a 400 on every current Claude model; adaptive thinking + output_config.effort is the current API. Model defaults are several generations stale (openai gpt-4o, anthropic claude-sonnet-4-6, ollama llama3). - Gemini's media surface (Imagen, Veo, native TTS, Live API, embeddings) is entirely unexposed; xai/groq/together now ship native SDKs we do not use; mistralai v3.0.0 sits unused while the provider is httpx-only. - OpenAI deprecated the Sora video APIs in 3.1 — recorded so we don't add them. Vendor SDK clones under /av/avalon/xrepos were fast-forwarded as part of this audit, and xai-sdk-python, groq-python and together-python were cloned (they were missing). No llmcore code changes yet — the phases are the follow-up. Co-Authored-By: Claude Opus 5 (1M context) --- README.md | 2 + docs/PROVIDER_MODERNIZATION_PLAN.md | 329 ++++++++++++++++++++++++++++ docs/PROVIDER_SUPPORT_MATRIX.md | 311 ++++++++++++++++++++++++++ 3 files changed, 642 insertions(+) create mode 100644 docs/PROVIDER_MODERNIZATION_PLAN.md create mode 100644 docs/PROVIDER_SUPPORT_MATRIX.md diff --git a/README.md b/README.md index faf80618..564b9421 100644 --- a/README.md +++ b/README.md @@ -1077,6 +1077,8 @@ from llmcore import ( - [Search providers rationale](docs/Search_providers_rationale.md) - [Deepgram provider usage](docs/Deepgram_provider_usage.md) - [FriendliAI provider usage](docs/Friendli_provider_usage.md) +- [Provider support matrix](docs/PROVIDER_SUPPORT_MATRIX.md) — SDK/API versions we track per provider, plus the capability matrix +- [Provider modernization plan](docs/PROVIDER_MODERNIZATION_PLAN.md) — phased plan to close the gaps in that matrix - [TypeSafe.ai provider usage](docs/TypeSafe_provider_usage.md) - [`chat_with_usage` guide](docs/USAGE_chat_with_usage.md) - [Model cards](docs/model_cards.md) diff --git a/docs/PROVIDER_MODERNIZATION_PLAN.md b/docs/PROVIDER_MODERNIZATION_PLAN.md new file mode 100644 index 00000000..45d3f424 --- /dev/null +++ b/docs/PROVIDER_MODERNIZATION_PLAN.md @@ -0,0 +1,329 @@ +# Provider Modernization Plan + +The program that closes the gaps recorded in +[`PROVIDER_SUPPORT_MATRIX.md`](PROVIDER_SUPPORT_MATRIX.md), so that `llmcore` +stays at the bleeding edge of the APIs it curates. + +- **Written:** 2026-09-29, from the audit of the same date +- **Scope:** all 16 providers + the 3 OpenAI-compatible aliases (xai, groq, together) +- **Status:** plan only — no phase has landed yet + +--- + +## 1. Principles + +These are the rules the phases below are designed against. They are also the +review checklist for any new provider. + +1. **Dual transport, always.** Every provider should reach its API by at least + two independent paths: direct REST from llmcore (`httpx`) and the vendor SDK + (or the `openai` SDK in compatibility mode). One `backend` config key + selects; unset auto-resolves. Rationale: vendor SDKs lag their own APIs, + drop fields outside their published schema (proven for `friendli`), and add + breaking dependency churn (proven for `openai` 3.x / `anthropic` 1.x). +2. **Direct-first where the SDK is lossy.** Auto-resolution should prefer + whichever path preserves the most response data. Record the reason in the + provider docstring, as `friendli_provider.py` does. +3. **One contract, every provider.** The `BaseProvider` extractor surface + (`extract_reasoning_content`, `extract_delta_reasoning_content`, + `extract_tool_calls`, `extract_usage_details`, `extract_finish_reason`) is + not optional-in-practice. Provider-specific names like + `extract_thinking_content()` are bugs in disguise: callers cannot use them + polymorphically. +4. **Capabilities are declared, not guessed.** Anything a provider can do must + be discoverable from `get_models_details()` and from its model cards, so + routing and validation work without hardcoded model lists. +5. **Live-validated.** No provider change lands without a real call against the + real API using the keys in `/av/data/dbs/.env`, plus offline tests that mock + the transport. +6. **Additive.** Config keys and methods are added with safe defaults; existing + callers keep working. + +--- + +## 2. Phase 0 — Unblock the dependency upgrade (P0, breaking) + +**Why first:** `openai>=3` and `anthropic>=1` cannot be adopted until this is +done, and every later phase needs those SDKs. Nothing else in the plan is +blocked by anything but this. + +### 2.1 Declare `httpx` where it is actually used + +`openai` 3.x no longer installs `httpx` transitively. Six providers import it +with no extra of their own and would fail at import (matrix §2.1). Add extras: + +```toml +mistral = ["httpx>=0.27.0"] +kimi = ["openai>=3.0.0", "httpx>=0.27.0", "tiktoken>=0.9.0"] +poe = ["openai>=3.0.0", "httpx>=0.27.0"] # + fastapi-poe optional +openrouter = ["openai>=3.0.0", "httpx>=0.27.0"] # + openrouter optional +vllm = ["openai>=3.0.0", "httpx>=0.27.0"] +huggingface = ["huggingface-hub>=1.12.0"] # currently unpinned +``` + +Add all six to `[all]`, and to the CI install list. + +### 2.2 Bump the pins + +| Extra | From | To | Notes | +|---|---|---|---| +| `openai` | `openai>=2.31.0` | `openai>=3.0.0,<4` | pulls `httpx2`, drops `httpx`/`certifi` | +| `anthropic` | `anthropic>=0.94.0` | `anthropic>=1,<2` | min Python 3.10 — llmcore is already ≥3.11 | +| `gemini` | `google-genai>=1.72.0` | `google-genai>=2,<3` | `GenerateContent` unaffected by the 2.0 break | +| `ollama` | `>=0.6.0` | `>=0.6.3` | | +| `deepgram` | `>=7.0.0` | `>=7.11.0` | | +| `friendli` | `>=0.15.1` | `>=0.15.2` | | + +### 2.3 Port the breaking changes + +- **openai 3.x** — verified: llmcore never passes `httpx` objects into the + client, so there is *no code* to port. Only packaging and docs change. +- **anthropic 1.x** — audit for: removed deprecated request params, removed type + aliases/exports, `.with_raw_response` shape (async now awaited), byte-valued + headers, and the removed legacy Text Completions API. llmcore uses none of + these as far as the audit could tell — **confirm with a type-checker pass**, + which the upstream guide recommends as the checklist. +- The bundled `claude-api` skill has a `/claude-api upgrade python` subcommand + and a `python/claude-api/sdk-upgrade.md` guide written for exactly this + migration — use it rather than improvising. + +### 2.4 Document the TLS change + +httpx2 verifies against the **OS trust store**, not `certifi`. Add a note to +`CONFIG_REFERENCE.md` and the provider guides: minimal containers and +TLS-inspecting corporate proxies need CA certs installed, or +`SSL_CERT_FILE` / `SSL_CERT_DIR` set. + +### 2.5 Verification + +- Install `[all]` into a scratch venv and import every provider module. +- Full unit suite green. +- One live call per SDK-backed provider. +- Confirm `respx` still intercepts llmcore's own clients (it should — the audit + found no test routes traffic through a vendor SDK). + +**Deliverable:** one PR, packaging + docs + any anthropic 1.x fixes. No new features. + +--- + +## 3. Phase 1 — Contract consistency (P1, high value / low risk) + +Cheap, purely additive, and it makes every later phase easier. + +### 3.1 Unify the extractor contract + +Implement the full five-method surface on every chat provider (matrix §4.1 — +only 4 of 16 have it today). Specifically: + +- `ollama`: rename/alias `extract_thinking_content()` → + `extract_reasoning_content()` (keep the old name as a deprecated alias). +- `gemini`: expose the `thinking` parts it already parses through + `extract_reasoning_content()` / `extract_delta_reasoning_content()` instead of + only `message["thinking"]`. +- `anthropic`: expose `thinking` content blocks through the same contract. +- `mistral`: same for Magistral reasoning. +- `openai`: add reasoning extraction for the reasoning-model families. +- All: add `extract_usage_details()` and `extract_finish_reason()` where missing. + +Add a **cross-provider contract test** — parametrized over every registered +provider — asserting each method exists, is defensive against `{}`, and returns +the documented type. This is the same shape as the static guard added in +`tests/providers/test_context_length_error_mapping.py`, which caught a bug three +providers had silently carried. + +### 3.2 Provider-level embeddings + +`create_embeddings()` is missing on `openai`, `gemini`, `ollama`, `vllm`, +`openrouter`, `deepseek`, `kimi`. The separate `[embedding.*]` subsystem covers +some of these, but the provider surface should be complete and consistent so +callers can embed through whichever provider they already hold. + +### 3.3 New `BaseProvider` surfaces + +Add, defaulting to `NotImplementedError` so nothing existing changes: + +| Method | Rationale | +|---|---| +| `generate_video()` | only `zai` has one today, under a provider-specific name; Gemini Veo and others need a home | +| `rerank()` | vLLM, HF and several hosted providers expose rerank/score endpoints | +| `create_batch()` / `retrieve_batch()` | OpenAI + Anthropic + Mistral all have batch APIs at ~50% cost | +| `upload_file()` / `list_files()` | OpenAI, Anthropic, Gemini, Kimi all have Files APIs; `kimi` already has a bespoke `upload_file()` | +| `count_tokens_native()` | distinguish exact provider tokenizers from local estimates; `friendli` and `kimi` already have them under bespoke names | + +--- + +## 4. Phase 2 — Anthropic (P1, largest single gap) + +The audit found the widest divergence here. Split into reviewable PRs: + +1. **Models + cards.** Refresh the lineup to Opus 5.5 / Opus 5 / Sonnet 5.5 / + Sonnet 5 / Haiku 4.5 / Fable 5.x, update the default from + `claude-sonnet-4-6`, regenerate cards, and record per-model thinking rules. +2. **Thinking and effort — correctness fix, not a feature.** `budget_tokens` is + **rejected with a 400** on Opus 5.x / Sonnet 5.x / Fable 5.x, yet + `thinking_budget_tokens` is still a documented llmcore config key. Move to + `thinking: {type: "adaptive"}` + `output_config.effort` + (`low|medium|high|xhigh|max`), mapped from llmcore's canonical + reasoning-effort vocabulary (`model_cards.md` §Reasoning-Effort Vocabulary). + Keep `budget_tokens` only for the pre-4.6 models that still accept it. +3. **Refusals.** Handle `stop_reason: "refusal"` + `stop_details` (currently + unhandled — a refusal looks like an empty response), and expose the + server-side `fallbacks` parameter. +4. **Server tools.** `web_search_20260209`, `web_fetch_20260209`, + `code_execution_20260521`, tool search — and wire `web_search` into + `supports_native_search()`, which today returns `False` for Anthropic. +5. **Batches + Files + count_tokens + Models API.** `messages.count_tokens` + replaces llmcore's local estimate for Anthropic; the Models API gives live + context/capability discovery for `get_models_details()`. +6. **Context lifecycle.** Compaction, context editing, mid-conversation system + messages — these interact directly with llmcore's own context manager, and + **preserved thinking** means llmcore's history rewriting can invalidate + thinking blocks. Treat as a design task, not a passthrough. +7. **Optional/lower priority.** Fast mode, task budgets, memory tool, citations, + Bedrock/Vertex/Foundry client variants, Admin API. + +> Anthropic publishes no separate REST client, so "dual approach" here means +> SDK + llmcore's own `httpx` path against `/v1/messages`. Worth building for +> the same reason as Friendli: it removes the SDK from the critical path. + +--- + +## 5. Phase 3 — OpenAI (P1) + +1. **Provider-level `create_embeddings()`** (currently absent). +2. **Responses API** as a second chat surface alongside Chat Completions, + selected per call/config. This is where OpenAI ships new capability first. +3. **Reasoning extraction** for the reasoning-model families, via the Phase 1 + contract. +4. **Batch, Files, Vector Stores, Containers.** +5. **Refresh the default model** off `gpt-4o`, refresh cards. +6. **Do not add Sora video** — deprecated upstream in `openai` 3.1. +7. Evaluate the "Ultrafast" tier and structured MCP errors (both v3.1). +8. **Add an `httpx` direct path** to `OpenAIProvider`. It is the base class for + `deepinfra`, `vllm`, `poe`, `openrouter`, so one direct backend gives five + providers a dual approach at once. Highest leverage item in this phase. + +--- + +## 6. Phase 4 — Gemini media (P2) + +Gemini has the largest *unexposed* media surface of any provider llmcore +curates: + +| Capability | Target | +|---|---| +| Image generation (Imagen family) | `generate_image()` | +| Video generation (Veo family) | `generate_video()` (new, Phase 1.3) | +| Native TTS | `generate_speech()` | +| Audio understanding | `transcribe_audio()` | +| Embeddings | `create_embeddings()` | +| Live API (bidirectional realtime) | new streaming surface, mirroring the Deepgram voice-agent socket pattern | +| Multimodal file search | evaluate | + +Also: port off the deprecated `response_format` to the v2 polymorphic field, and +verify whether the v2.0 "interactions" surface is worth adopting (its breaking +changes do not affect `GenerateContent`, which is what llmcore uses). + +--- + +## 7. Phase 5 — Native SDKs for the compatibility aliases (P2) + +`xai`, `groq` and `together` are bare `OpenAIProvider` + `base_url`. Each now +has a native SDK. Promote each to a first-class provider with the dual-transport +shape, which also unlocks vendor-specific surfaces (xAI Live Search is already +referenced by `supports_native_search()` with no native path behind it). + +Same phase: give **Mistral** a dual approach. llmcore is httpx-only there while +`mistralai` v3.0.0 sits unused — and v3.0's breaking changes are undocumented in +the repo, so review the SDK source before pinning. + +--- + +## 8. Phase 6 — Long tail (P3) + +- **vLLM**: `/v1/embeddings`, `/pooling`, `/score`, `/rerank`; refresh the clone (5 months stale). +- **OpenRouter**: adopt the v1.x SDK as the `sdk` backend; expose routing/ZDR/caching controls. +- **Ollama**: provider-level embeddings; audit the newer `/api/*` surface. +- **Hugging Face**: pin the dep; rerank; audit Inference Providers routing. +- **Poe**: surface media bots through the media APIs. +- **Deepgram**: review 7.3.1 → 7.11.0 for new surface. +- **Z.ai**: install `zai-sdk` in CI so the preferred SDK backend is exercised. + +--- + +## 9. Testing and validation strategy + +Per provider touched: + +1. **Offline** — mock at the transport boundary (`respx` for llmcore's own + `httpx` clients; `AsyncMock` on the SDK client object for SDK paths). Pin the + `backend` explicitly in tests so results do not depend on what is installed — + the pattern established in `tests/providers/test_friendli_provider.py`. +2. **Contract** — the parametrized cross-provider tests from Phase 1. +3. **Static guards** — AST checks for whole-package invariants, in the style of + `tests/providers/test_context_length_error_mapping.py`. Candidates: every + provider registered in `PROVIDER_MAP` has a `[providers.*]` config section, a + confy schema section, and a matrix row. +4. **Live** — a key-gated smoke per provider under `tests/integration/`, + skipped when the key is absent (the `test_typesafe_live.py` pattern). +5. **CI** — keep vendor SDKs whose presence would bypass mocks out of the + install list, and say why in the workflow comment. + +--- + +## 10. Risks and open questions + +| Risk | Mitigation | +|---|---| +| httpx2's OS trust store breaks deployments in minimal containers | document `SSL_CERT_FILE`/`SSL_CERT_DIR`; call it out in release notes | +| `anthropic` 1.x removals not caught by the audit | run mypy/pyright over the anthropic provider after bumping — upstream recommends exactly this | +| `mistralai` v3.0 breaking changes undocumented | read the SDK source before pinning; keep httpx as the default backend | +| Phase 2.6 (compaction / context editing / preserved thinking) touches llmcore's own context manager | design task with its own review, not a passthrough PR | +| Provider APIs drift again | §6 of the matrix is a runnable refresh procedure; re-run it per release | + +**Open questions for the maintainer:** + +1. **Phase order** — this plan front-loads Phase 0 (blocking) then Anthropic + (biggest gap). Prefer breadth-first instead (every provider to parity before + any provider gets new capabilities)? +2. **Minimum SDK floors** — pin to `>=3,<4` style ranges as proposed, or track + exact versions for reproducibility? +3. **`openai` direct backend** (§5.8) is the highest-leverage single item — five + providers gain a dual approach at once. Promote it ahead of Anthropic? +4. Should the compatibility aliases (`xai`/`groq`/`together`) become first-class + providers, or stay thin aliases with a documented native-SDK escape hatch? + +--- + +## Appendix A — capability extraction script + +Regenerates matrix §4/§4.1 from the source rather than by hand: + +```python +import ast, pathlib + +BASE_OPTIONAL = ["generate_speech", "transcribe_audio", "generate_image", "ocr", + "create_embeddings", "warm_up", "supports_native_search"] +EXTRACTORS = ["extract_reasoning_content", "extract_delta_reasoning_content", + "extract_tool_calls", "extract_usage_details", "extract_finish_reason"] +CORE = {"get_name", "get_models_details", "get_supported_parameters", + "get_max_context_length", "chat_completion", "count_tokens", + "count_message_tokens", "extract_response_content", + "extract_delta_content", "close"} + +for f in sorted(pathlib.Path("src/llmcore/providers").glob("*_provider.py")): + tree = ast.parse(f.read_text()) + cls = next((n for n in tree.body if isinstance(n, ast.ClassDef)), None) + if cls is None: + continue + methods = {n.name for n in cls.body + if isinstance(n, (ast.FunctionDef, ast.AsyncFunctionDef))} + extra = sorted(m for m in methods if not m.startswith("_") + and m not in BASE_OPTIONAL and m not in EXTRACTORS and m not in CORE) + print(f"\n## {f.stem.replace('_provider', '')} ({cls.name} <- " + f"{', '.join(ast.unparse(b) for b in cls.bases)})") + print(" media/opt :", ", ".join(m for m in BASE_OPTIONAL if m in methods) or "-") + print(" extractors:", ", ".join(e.replace("extract_", "") + for e in EXTRACTORS if e in methods) or "-") + print(" extra API :", ", ".join(extra) or "-") +``` diff --git a/docs/PROVIDER_SUPPORT_MATRIX.md b/docs/PROVIDER_SUPPORT_MATRIX.md new file mode 100644 index 00000000..6df94a07 --- /dev/null +++ b/docs/PROVIDER_SUPPORT_MATRIX.md @@ -0,0 +1,311 @@ +# Provider Support Matrix + +Tracking document for every provider `llmcore` curates: which upstream API/SDK +version we are known-good against, how we talk to it, and which of the +provider's capabilities we actually expose. + +**This file is the source of truth for "are we current?".** Update it in the +same commit as any provider change — see [Refreshing this document](#refreshing-this-document). + +- **Last full audit:** 2026-09-29 +- **llmcore version at audit:** 0.53.0 +- **Vendor SDK clones:** `/av/avalon/xrepos/` (paths in the table below) + +--- + +## 1. Legend + +| Mark | Meaning | +|:---:|---| +| ✅ | Implemented in llmcore and covered by tests | +| 🟡 | Partially implemented (see the note) | +| ❌ | Provider supports it; llmcore does **not** | +| — | Provider does not offer it (nothing to do) | +| ❓ | Provider surface not yet verified against upstream docs — audit pending | + +**Transport** column values: + +- `sdk` — vendor SDK only +- `openai` — the `openai` SDK pointed at the vendor's OpenAI-compatible base URL +- `httpx` — direct REST calls from llmcore +- Multiple values joined by `→` are the runtime auto-resolution order (the + "dual approach"): first available wins, and each is individually selectable + via the provider's `backend` config key. + +--- + +## 2. SDK / API version tracking + +Upstream columns are the **vendor SDK clone in `/av/avalon/xrepos`** as of the +audit date. `llmcore pin` is from `pyproject.toml`; `installed` is the shared +dev venv. A pin that trails the upstream **major** version is a red flag. + +| Provider | SDK package | Clone (`/av/avalon/xrepos/…`) | Upstream tag | Upstream commit | Tag date | llmcore pin | Installed | Status | +|---|---|---|---|---|---|---|---|:---:| +| OpenAI | `openai` | `openai-python` | **v3.22.1** | `58aca1dcfd8d` | 2026-09-30 | `>=2.31.0` | 2.32.0 | 🔴 major behind | +| Anthropic | `anthropic` | `anthropic-sdk-python` | **v1.9.0** | `a7285e919ab7` | 2026-09-28 | `>=0.94.0` | 0.97.0 | 🔴 major behind | +| Google Gemini | `google-genai` | `python-genai` | **v2.25.0** | `f15d1482d747` | 2026-09-29 | `>=1.72.0` | 1.73.1 | 🔴 major behind | +| Mistral | `mistralai` | `mistral-client-python` | **v3.0.0** | `e8dfa1c8a2d0` | 2026-09-28 | *(none — httpx only)* | not installed | 🟠 SDK unused | +| OpenRouter | `openrouter` | `openrouter_python_sdk` | **v1.3.9** | `fd5ffce2995d` | 2026-09-30 | *(optional backend)* | not installed | 🟠 major behind + absent | +| Ollama | `ollama` | `ollama-python` | v0.6.3 | `8785556559ec` | 2026-09-28 | `>=0.6.0` | 0.6.1 | 🟡 patch behind | +| Deepgram | `deepgram-sdk` | `deepgram-python-sdk` | v7.11.0 | `a379a7f37b11` | 2026-09-28 | `>=7.0.0` | 7.3.1 | 🟡 minor behind | +| Hugging Face | `huggingface-hub` | `huggingface_hub` | `main` @ v0.9.0.rc1 | `1092497a9b65` | 2026-09-29 | *(none)* | 1.12.0 | 🟠 unpinned | +| Z.ai (GLM) | `zai-sdk` | `z-ai-sdk-python` | v0.2.3 | `ca5109c0aa9b` | 2026-06-16 | `>=0.2.0` | not installed | 🟡 SDK path untested | +| FriendliAI | `friendli` | `friendli-python` | v0.15.1 (pyproject 0.15.2) | `f3039e22ec0d` | 2026-09-28 | `>=0.15.1` | 0.15.1 | ✅ current | +| TypeSafe.ai | `typesafe-sdk` | `typesafe-sdk-python` | v0.7.2 | `f078f1e208a0` | 2026-09-26 | *(none — httpx only)* | not installed | ✅ by design | +| Poe | `fastapi_poe` | `fastapi_poe` | 0.0.83 | `41ffd02e16f2` | 2026-01-21 | *(optional backend)* | not installed | 🟡 SDK path untested | +| vLLM (self-hosted) | `vllm` (server) | `vllm` | v0.19.1rc0 | `219bb5b8c0dc` | 2026-04-16 | *(server, not a client dep)* | n/a | 🟡 clone stale | +| DeepSeek | *(none published)* | — | — | — | — | uses `openai` | — | ✅ by design | +| Kimi (Moonshot) | *(none published)* | — | — | — | — | uses `openai` | — | ✅ by design | +| DeepInfra | *(none published)* | — | — | — | — | uses `openai` | — | ✅ by design | +| xAI (Grok) | `xai-sdk` | `xai-sdk-python` | **v1.20.0** | `1d9e1dffc9a0` | 2026-09-28 | *(none)* | not installed | 🔴 native SDK unused | +| Groq | `groq` | `groq-python` | **v1.7.0** | `55066d94acca` | 2026-09-04 | *(none)* | not installed | 🔴 native SDK unused | +| Together | `together` | `together-python` | **v1.5.35** | `cc9f25369987` | 2026-03-18 | *(none)* | not installed | 🔴 native SDK unused | + +### 2.1 Cross-cutting: the HTTPX2 migration + +`openai` 3.0.0 (2026-08-12) and `anthropic` 1.0.0 (2026-08-20) both moved their +HTTP layer from `httpx` to [`httpx2`](https://httpx2.pydantic.dev/) (Pydantic's +maintained fork) — see `openai-python/httpx2.md` and +`anthropic-sdk-python/MIGRATION.md` in the clones. + +What this means for llmcore, **verified against the source**: + +| Concern | Status | +|---|---| +| Does llmcore pass `httpx` objects *into* the OpenAI/Anthropic clients (`http_client=`, `httpx.Timeout`)? | **No.** Only numeric timeouts are passed. Nothing to port. | +| Does llmcore use `httpx` for its *own* clients? | **Yes** — 9 providers (`deepinfra`, `poe`, `mistral`, `friendli`, `kimi`, `openrouter`, `vllm`, `zai`, `typesafe`) plus all search providers. These keep using `httpx` and are unaffected. | +| Do `respx`-based tests break? | **No.** Every `respx` mock in `tests/` targets llmcore's own `httpx` clients, never traffic routed through a vendor SDK. (Upstream vendored `tests/respx2` for their own suite; we don't need it.) | +| **`httpx` is no longer installed transitively by `openai`** | 🔴 **Action required** — see below. | +| `certifi` is no longer installed by `openai`; httpx2 uses the **OS trust store** | 🔴 Deployment risk to document (minimal containers, TLS-inspecting proxies). Mitigate with `SSL_CERT_FILE` / `SSL_CERT_DIR`. | + +**The concrete break:** six providers import `httpx` but have **no extra of +their own**, so today they only work because `openai` happened to install +`httpx` transitively. Under `openai>=3` they fail at import: + +| Provider | Imports `httpx` | Own extra | +|---|:---:|:---:| +| `mistral` | yes | ❌ missing | +| `kimi` | yes | ❌ missing | +| `poe` | yes | ❌ missing | +| `openrouter` | yes | ❌ missing | +| `vllm` | yes | ❌ missing | +| `huggingface` | — | ❌ missing | + +(`deepinfra`, `zai`, `friendli`, `typesafe` already declare `httpx` explicitly.) + +--- + +## 3. Default models configured in `default_config.toml` + +Stale defaults are a correctness problem, not cosmetics — they drive model-card +lookups, context budgets and cost estimates. + +| Provider | Configured default | Assessment | +|---|---|---| +| openai | `gpt-4o` | 🔴 several generations stale | +| anthropic | `claude-sonnet-4-6` | 🔴 current lineup is Opus 5.5 / Sonnet 5.5 (see §5.2) | +| gemini | `gemini-3.1-flash-lite-preview` | ❓ verify against current lineup | +| deepseek | `deepseek-v4-pro` | ✅ current at last provider audit | +| zai | `glm-5.2` | 🟡 GLM-5.3 is served (Friendli lists it) | +| friendli | `zai-org/GLM-5.3` | ✅ verified live 2026-09-20 | +| ollama | `llama3` | 🔴 very stale | +| openrouter | `openai/gpt-4o-mini` | 🔴 stale | +| poe | `GPT-4o-Mini` | 🔴 stale | +| vllm | `meta-llama/Llama-3.1-8B-Instruct` | 🟡 example value, self-hosted | +| mistral | `mistral-large-latest` | ✅ alias, self-updating | +| huggingface | `meta-llama/Llama-3.3-70B-Instruct` | 🟡 stale example | +| deepinfra | `deepseek-ai/DeepSeek-V3` | 🟡 V3.2 is served | +| kimi | `kimi-k2.6` | ✅ current at last provider audit | +| typesafe | `jev-latest` | ✅ alias, self-updating | + +--- + +## 4. Capability matrix + +What llmcore **exposes today**, extracted from the provider classes (not from +vendor docs). Columns are the `BaseProvider` surface plus the media APIs. + +| Provider | Transport | Chat | Stream | Tools | Structured out | Reasoning extract | Vision in | Audio in (STT) | Audio out (TTS) | Image gen | Video gen | Embeddings | OCR | Native search | Exact tokenizer | +|---|---|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:|:-:| +| openai | `sdk` | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | — | ✅ | ✅ | +| anthropic | `sdk` | ✅ | ✅ | ✅ | ✅ | 🟡 | ✅ | — | — | — | — | — | — | ❌ | ❌ | +| gemini | `sdk` | ✅ | ✅ | ✅ | ✅ | 🟡 | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | — | ✅ | ❌ | +| deepseek | `sdk`(openai) | ✅ | ✅ | ✅ | ✅ | ✅ | — | — | — | — | — | ❌ | — | — | ❌ | +| zai | `sdk → openai → httpx` | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | +| friendli | `openai → httpx → sdk` | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | — | ✅ | — | ✅ | — | — | ✅ | +| mistral | `httpx` | ✅ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | — | — | ✅ | ✅ | — | ❌ | +| kimi | `openai + httpx` | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | — | — | — | — | ❌ | — | — | ✅ | +| deepinfra | `openai + httpx` | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | — | ✅ | — | — | ❌ | +| ollama | `sdk` | ✅ | ✅ | ✅ | ✅ | 🟡 | ✅ | — | — | — | — | ❌ | — | — | 🟡 | +| huggingface | `sdk` | ✅ | ✅ | ✅ | 🟡 | ❌ | ✅ | ✅ | ✅ | ✅ | — | ✅ | — | — | ❌ | +| openrouter | `openai (+sdk)` | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | — | — | — | — | ❌ | — | ❓ | ❌ | +| poe | `openai (+native)` | ✅ | ✅ | ✅ | 🟡 | ❌ | ✅ | ❓ | ❓ | ❓ | ❓ | — | — | — | ❌ | +| vllm | `openai + httpx` | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | — | — | — | — | ❌ | — | — | ❌ | +| deepgram | `sdk` | — | — | — | — | — | — | ✅ | ✅ | — | — | — | — | — | — | +| typesafe | `httpx` | — | — | — | ✅ | — | — | — | — | — | — | — | — | — | — | + +Notes on the 🟡 cells: + +- **anthropic / gemini / ollama / mistral reasoning** — reasoning *is* surfaced, + but through provider-specific channels (`message["thinking"]`, + `extract_thinking_content()`) instead of the + `extract_reasoning_content()` / `extract_delta_reasoning_content()` contract + that deepseek, zai, kimi and friendli implement. **Naming should be unified.** +- **ollama exact tokenizer** — tiktoken approximation, not the model's tokenizer. +- **huggingface / poe structured output** — passthrough only, unvalidated. +- **deepgram** is a voice/audio provider by design; `chat_completion()` + intentionally raises. +- **typesafe** is a typed-judgment API, not a chat model, by design. + +### 4.1 Extractor contract coverage + +| Provider | `reasoning_content` | `delta_reasoning_content` | `tool_calls` | `usage_details` | `finish_reason` | +|---|:-:|:-:|:-:|:-:|:-:| +| deepseek, zai, kimi, friendli | ✅ | ✅ | ✅ | ✅ | ✅ | +| openai, mistral | ❌ | ❌ | ✅ | ✅ | ❌ | +| anthropic, gemini, ollama, huggingface | ❌ | ❌ | ✅ | ❌ | ❌ | +| openrouter, poe, vllm, deepinfra | ❌ | ❌ | inherited | inherited | ❌ | + +Only four of sixteen providers implement the full extractor contract. This is +the single biggest consistency gap in the provider layer. + +--- + +## 5. Known capability gaps per provider + +Vendor-side surfaces that exist upstream and are **not** in llmcore. Items +marked ❓ still need verification against current vendor docs. + +### 5.1 OpenAI (`openai` v3.x) + +- ❌ **Embeddings** at provider level (`create_embeddings()` is not overridden — + only the separate `[embedding.openai]` subsystem covers it) +- ❌ Responses API (llmcore uses Chat Completions only) +- ❌ Batch API, Files API, Vector Stores, Containers +- ❌ Realtime / WebSocket sessions (v3.1 added WebSocket stream IDs) +- ❌ `reasoning_content` extractor for the reasoning-model families +- ❌ "Ultrafast" service tier (added v3.1) +- ⚠️ **Sora video APIs were deprecated in v3.1** — do *not* add them +- ❓ Structured MCP / separate websocket error events (v3.1) + +### 5.2 Anthropic (`anthropic` v1.x) + +Verified against the bundled `claude-api` skill reference (2026-09-25 cache): + +- 🔴 **Model lineup is stale.** Current: `claude-opus-5-5` (default, + 1M ctx / 128K out), `claude-opus-5`, `claude-opus-4-8/4-7/4-6`, + `claude-sonnet-5-5`, `claude-sonnet-5`, `claude-sonnet-4-6`, + `claude-haiku-4-5`, `claude-fable-5-1`/`claude-fable-5`. + llmcore defaults to `claude-sonnet-4-6` and its cards stop at the 4.x family. +- 🔴 **`budget_tokens` is rejected (400) on Opus 5.x / Sonnet 5.x / Fable 5.x.** + llmcore's `thinking_budget_tokens` config would hard-fail on current models. + `thinking: {type: "adaptive"}` + `output_config.effort` is the current API. +- ❌ `output_config.effort` (`low|medium|high|xhigh|max`) not plumbed +- ❌ `stop_reason: "refusal"` / `stop_details` handling +- ❌ Server tools: `web_search_20260209`, `web_fetch_20260209`, + `code_execution_20260521`, tool search +- ❌ Batches, Files, Skills, Models API, `messages.count_tokens` +- ❌ Citations on document blocks +- ❌ Compaction, context editing, mid-conversation system messages +- ❌ Fast mode, task budgets, memory tool, Tool Runner +- ❌ Bedrock / Vertex / Foundry provider clients +- ❌ Preserved-thinking (history-editing) compliance — relevant to llmcore's + context-management rewrites + +### 5.3 Google Gemini (`google-genai` v2.x) + +- ❌ Image generation (Imagen family) +- ❌ Video generation (Veo family) +- ❌ Native TTS / audio output +- ❌ Live API (bidirectional realtime) — a `*-live-preview` model is already in + the context table but unused +- ❌ Provider-level `create_embeddings()` +- ❌ Multimodal file search (added v1.75) +- ❌ Interactions API (v2.0's breaking surface; `GenerateContent` unaffected) +- ⚠️ Legacy `response_format` deprecated for a new polymorphic field (v2.0) +- 🟡 Reasoning exposed as `thinking`, not via the extractor contract + +### 5.4 Mistral (`mistralai` v3.0.0) + +- 🟠 llmcore is httpx-only; the **v3 SDK is not used at all** → no dual approach +- ❓ v3.0 breaking changes unknown (no CHANGELOG in the repo) — needs review +- ❌ Agents / conversations API ❓ +- ❌ Batch API ❓ + +### 5.5 xAI / Groq / Together + +All three are wired as bare `OpenAIProvider` + `base_url`. Each now ships a +**native Python SDK** (`xai-sdk` v1.20.0, `groq` v1.7.0, `together` v1.5.35), +so none has a dual approach and none exposes vendor-specific surfaces +(xAI Live Search is referenced by `supports_native_search()` but there is no +native-SDK path). + +### 5.6 Others + +| Provider | Gaps | +|---|---| +| ollama | ❌ provider-level embeddings (`/api/embed`); 🟡 tokenizer is approximate; ❓ newer `/api/*` surface | +| openrouter | 🟠 v1.x SDK unused; ❌ embeddings; ❓ provider routing/ZDR/prompt-caching controls | +| poe | ❓ media bots (image/video/audio) not surfaced through the media APIs | +| vllm | ❌ `/v1/embeddings`, `/pooling`, `/score`, `/rerank`; clone is 5 months stale | +| huggingface | 🟠 unpinned dep; ❌ provider-level rerank; ❓ Inference Providers routing surface | +| deepgram | 🟡 SDK 7.3.1 vs 7.11.0 — review new surface | +| zai | 🟡 `zai-sdk` never installed, so the preferred SDK backend is never exercised in CI | +| friendli | ⚠️ `/detokenize` + `/chat/render` 404 on Model APIs (upstream gap, tracked) | + +--- + +## 6. Refreshing this document + +Run from the repo root. **Never print secret values.** + +```bash +# 1. Sync every vendor SDK clone (fast-forward only) +cd /av/avalon/xrepos +for d in openai-python anthropic-sdk-python python-genai ollama-python \ + mistral-client-python huggingface_hub z-ai-sdk-python \ + deepgram-python-sdk typesafe-sdk-python friendli-python \ + openrouter_python_sdk fastapi_poe xai-sdk-python groq-python \ + together-python vllm; do + git -C "$d" pull --ff-only --tags origin >/dev/null 2>&1 + printf "%-26s %-12s %s %s\n" "$d" \ + "$(git -C $d describe --tags --abbrev=0 2>/dev/null)" \ + "$(git -C $d rev-parse --short=12 HEAD)" \ + "$(git -C $d log -1 --format=%cs)" +done + +# 2. Compare against llmcore's pins and the installed venv +cd /media/araray/kilgharrah/repos/llmcore +./venv/bin/python -c " +import importlib.metadata as md, tomllib, pathlib +ex = tomllib.loads(pathlib.Path('pyproject.toml').read_text())['project']['optional-dependencies'] +print({k: v for k, v in ex.items() if k in ('openai','anthropic','gemini','ollama','zai','deepgram','friendli')}) +for p in ('openai','anthropic','google-genai','ollama','deepgram-sdk','friendli','httpx','httpx2'): + try: print(p, md.version(p)) + except Exception: print(p, 'NOT INSTALLED') +" + +# 3. Re-extract the implemented-capability matrix (§4) from the source +# (the AST walk used for the audit lives in the plan doc's appendix) + +# 4. Refresh model cards for every provider with an adapter +./venv/bin/python -m tools.cardctl generate # no --force +./venv/bin/python -m tools.cardctl diff +``` + +Then update §2, §3, §4 and the audit date at the top **in the same commit** as +any provider change. + +--- + +## 7. Related documents + +- [`PROVIDER_MODERNIZATION_PLAN.md`](PROVIDER_MODERNIZATION_PLAN.md) — the phased + plan that closes the gaps listed in §5 +- [`model_cards.md`](model_cards.md) — card schema, the canonical + reasoning-effort vocabulary, and the cardctl workflow +- [`CONFIG_REFERENCE.md`](CONFIG_REFERENCE.md) — every provider config key +- Per-provider guides: [`Friendli_provider_usage.md`](Friendli_provider_usage.md), + [`Deepgram_provider_usage.md`](Deepgram_provider_usage.md), + [`TypeSafe_provider_usage.md`](TypeSafe_provider_usage.md) From 6b71bc74eadce531b980e9f4b25a5937ff5244d8 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Tue, 29 Sep 2026 23:32:16 -0300 Subject: [PATCH 07/11] chore(deps)!: adopt openai 3.x, anthropic 1.x and google-genai 2.x MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase 0 of docs/PROVIDER_MODERNIZATION_PLAN.md: unblock the SDK upgrade so the later capability phases have something current to build on. BREAKING: minimum SDK versions move across a major boundary. openai >=2.31.0 -> >=3.0.0,<4 anthropic >=0.94.0 -> >=1,<2 google-genai >=1.72.0 -> >=2,<3 ollama >=0.6.0 -> >=0.6.3 deepgram-sdk >=7.0.0 -> >=7.11.0 zai-sdk >=0.2.0 -> >=0.2.3 The three majors land together because openai 3.x and anthropic 1.x share one breaking change: their HTTP layer moved from httpx to httpx2 (Pydantic's maintained fork), which is now installed in place of httpx and certifi. Two consequences, both handled here: 1. Six providers (mistral, kimi, poe, openrouter, vllm, huggingface) import httpx but had no extra of their own — they worked only because openai installed httpx transitively. Under openai>=3 they fail at import. Each now has an extra declaring what it actually needs, and all six are in [all]. 2. httpx2 verifies TLS against the OS trust store rather than certifi, which can break minimal containers and TLS-inspecting proxies. Documented in CONFIG_REFERENCE.md with the SSL_CERT_FILE / SSL_CERT_DIR escape hatches. No provider code needed porting: llmcore only ever passes numeric timeouts to the vendor clients (never httpx objects), and no respx test routes traffic through a vendor SDK. Both were verified before bumping rather than assumed. Also fixes a latent test-isolation bug this upgrade exposed. Installing zai-sdk flipped ZaiProvider's backend auto-resolution from "openai" to "sdk", bypassing the AsyncOpenAI mocks in 21 tests — the exact hazard the old CI comment described when it deliberately left zai-sdk uninstalled. The tests now pin `backend` explicitly and patch the availability flags for resolution assertions, so they no longer depend on what happens to be installed. That removes the carve-out, and Z.ai's preferred SDK transport is exercised in CI and validated live for the first time. CI now installs .[dev,all] rather than a hand-maintained extras subset, so a new extra is covered the moment it is added to pyproject. The friendli pin stays at >=0.15.1: the vendor repo's pyproject reads 0.15.2 but that version is not published on PyPI. Verified: all 17 provider modules import; full unit suite green (5164 passed, 29 skipped); live calls through OpenAI 3.22.1, Google Gemini 2.25.0 (47 models discovered), Z.ai on the native SDK backend, and DeepSeek. NOT verified live: anthropic 1.9.0 — import- and test-clean, but no ANTHROPIC_API_KEY is available in this environment. Co-Authored-By: Claude Opus 5 (1M context) --- .github/workflows/ci.yml | 26 +++++++------ CHANGELOG.md | 41 +++++++++++++++++++++ docs/CONFIG_REFERENCE.md | 30 +++++++++++++++ docs/PROVIDER_SUPPORT_MATRIX.md | 33 +++++++++-------- pyproject.toml | 55 +++++++++++++++++++++++----- tests/providers/test_zai_provider.py | 50 +++++++++++++++++++++---- 6 files changed, 192 insertions(+), 43 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b64b9ce7..54c360f6 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -59,17 +59,21 @@ jobs: # grimoire (the control plane) is also GitHub-only. pip install "grimoire @ git+https://github.com/araray/grimoire.git@v0.4.0" # dev pulls in test (pytest, pytest-asyncio, respx, ...) and sandbox - # deps; the extras below are [all] minus [zai] and [friendli]: - # provider SDKs (openai, anthropic, ...), chromadb, and the bridge - # deps (grpcio/protobuf/starlette) that unit tests import at - # collection or fixture time. zai-sdk is deliberately NOT installed: - # the zai unit tests mock the openai/httpx fallback transport, and - # the SDK transport would bypass those mocks and hit the live API. - # The friendli SDK is likewise skipped: the provider auto-resolves to - # the openai backend (already installed) and its tests pin the - # backend explicitly, so the vendor SDK adds nothing offline. - # numpy is used by embedding-shaped unit tests. - pip install -e ".[dev,bridge,openai,anthropic,gemini,ollama,deepinfra,deepgram,typesafe,brightdata,serper,serpapi,semanticscholar,postgres,chromadb]" + # deps; [all] pulls every provider/search/storage extra plus the + # bridge deps (grpcio/protobuf/starlette) that unit tests import at + # collection or fixture time. Installing [all] rather than a + # hand-maintained subset means a new extra is exercised in CI the + # moment it is added to pyproject. + # + # Optional vendor SDKs (zai-sdk, friendli, ...) are installed on + # purpose: every provider test pins its `backend` explicitly, so the + # presence of a vendor SDK can no longer bypass the mocks and reach a + # live API. See docs/PROVIDER_MODERNIZATION_PLAN.md §9. + # + # NOTE: openai>=3 and anthropic>=1 are built on httpx2 and no longer + # install httpx/certifi transitively; the provider extras declare + # httpx explicitly where llmcore's own clients need it. + pip install -e ".[dev,all]" pip install numpy - name: Version drift guard diff --git a/CHANGELOG.md b/CHANGELOG.md index e6b4a813..810215ef 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,47 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## Unreleased +### Changed — provider SDK majors (Phase 0 of the modernization plan) + +- **`openai` `>=3.0.0,<4`** (was `>=2.31.0`), **`anthropic` `>=1,<2`** (was + `>=0.94.0`), **`google-genai` `>=2,<3`** (was `>=1.72.0`), plus + `ollama>=0.6.3`, `deepgram-sdk>=7.11.0`, `zai-sdk>=0.2.3`. All three majors + were adopted at once because `openai` 3.x and `anthropic` 1.x share the same + breaking change: their HTTP layer moved from `httpx` to **httpx2**. +- **New extras for six providers that had none**: `mistral`, `kimi`, `poe`, + `openrouter`, `vllm`, `huggingface`. These import `httpx` but previously + worked only because `openai` installed it transitively — under `openai>=3` + they would have failed at import. All six are in `[all]`. +- **TLS behaviour change documented**: httpx2 verifies against the OS trust + store, not `certifi`, which can break minimal containers and TLS-inspecting + proxies. `CONFIG_REFERENCE.md` gained an "HTTP transport and TLS" section + with the `SSL_CERT_FILE` / `SSL_CERT_DIR` escape hatches. +- **No provider code changes were needed for httpx2**: llmcore only ever passes + numeric timeouts to the vendor clients, never `httpx` objects, and no `respx` + test routes traffic through a vendor SDK. Verified before bumping. +- **Z.ai's native SDK backend is now exercised.** `zai-sdk` was never installed, + so the provider's preferred transport was dead code in CI. Its tests now pin + `backend` explicitly (the pattern the plan's §9 mandates) instead of depending + on what happens to be installed, so the SDK can be installed safely — and the + SDK backend is validated live for the first time. +- **CI installs `.[dev,all]`** instead of a hand-maintained extras subset, so a + new extra is exercised the moment it is added. The previous `zai-sdk` + carve-out is gone, since tests can no longer be bypassed by an installed SDK. + +Live-validated after the upgrade: OpenAI (3.22.1), Google Gemini (2.25.0, 47 +models discovered), Z.ai (SDK backend), DeepSeek. Full unit suite green (5164 +passed). **Anthropic 1.9.0 is import- and test-verified but not live-validated — +no `ANTHROPIC_API_KEY` is available in this environment.** + +### Added — provider audit documents + +- `docs/PROVIDER_SUPPORT_MATRIX.md` — the ongoing tracker: per provider, the + vendor SDK clone with tag/commit/date, our pin, the installed version, the + transport shape, and a capability matrix extracted from the provider classes. + Section 6 is a runnable refresh procedure. +- `docs/PROVIDER_MODERNIZATION_PLAN.md` — the phased program that closes the + gaps, built on the dual-transport and one-contract principles. + ### Added — FriendliAI provider - **FriendliAI provider**: first-class `FriendliProvider` covering all three diff --git a/docs/CONFIG_REFERENCE.md b/docs/CONFIG_REFERENCE.md index 20a114c6..4548179c 100644 --- a/docs/CONFIG_REFERENCE.md +++ b/docs/CONFIG_REFERENCE.md @@ -8,6 +8,36 @@ Self-hosted LLM provider abstraction layer with agentic capabilities, RAG, embed --- +## 🔐 HTTP transport and TLS (applies to every provider) + +`openai` 3.x and `anthropic` 1.x build their HTTP layer on +[`httpx2`](https://httpx2.pydantic.dev/) (Pydantic's maintained fork of `httpx`) +and **no longer install `httpx` or `certifi` transitively**. Two consequences +for deployments: + +1. **TLS trust store.** httpx2 verifies certificates against the **operating + system** trust store rather than the `certifi` bundle. This can break + certificate verification in minimal container images without system CA + certificates, behind TLS-inspecting corporate proxies, or where a customised + `certifi` bundle was relied on. Fix by installing CA certificates into the OS + trust store, or point the SDKs at an explicit bundle: + + ```bash + export SSL_CERT_FILE=/path/to/ca-bundle.pem # a single bundle + export SSL_CERT_DIR=/path/to/ca-directory # or a directory of CAs + ``` + +2. **`httpx` is an explicit dependency now.** llmcore's own REST transports + (mistral, kimi, poe, openrouter, vllm, deepinfra, zai, friendli, typesafe and + all search providers) still use `httpx`, and every extra that needs it + declares it. If you install llmcore without extras and call those providers, + install `httpx` yourself. + +llmcore never hands `httpx` objects to the vendor SDK clients (only numeric +timeouts), so there is no per-provider transport configuration to migrate. + +--- + ## ⚙️ Core Settings Fundamental llmcore configuration — provider selection, embedding model, and diagnostics. diff --git a/docs/PROVIDER_SUPPORT_MATRIX.md b/docs/PROVIDER_SUPPORT_MATRIX.md index 6df94a07..d2b5c4e0 100644 --- a/docs/PROVIDER_SUPPORT_MATRIX.md +++ b/docs/PROVIDER_SUPPORT_MATRIX.md @@ -8,6 +8,7 @@ provider's capabilities we actually expose. same commit as any provider change — see [Refreshing this document](#refreshing-this-document). - **Last full audit:** 2026-09-29 +- **Phase 0 landed:** 2026-09-29 — pins bumped to the majors below, extras corrected, live-validated - **llmcore version at audit:** 0.53.0 - **Vendor SDK clones:** `/av/avalon/xrepos/` (paths in the table below) @@ -42,16 +43,16 @@ dev venv. A pin that trails the upstream **major** version is a red flag. | Provider | SDK package | Clone (`/av/avalon/xrepos/…`) | Upstream tag | Upstream commit | Tag date | llmcore pin | Installed | Status | |---|---|---|---|---|---|---|---|:---:| -| OpenAI | `openai` | `openai-python` | **v3.22.1** | `58aca1dcfd8d` | 2026-09-30 | `>=2.31.0` | 2.32.0 | 🔴 major behind | -| Anthropic | `anthropic` | `anthropic-sdk-python` | **v1.9.0** | `a7285e919ab7` | 2026-09-28 | `>=0.94.0` | 0.97.0 | 🔴 major behind | -| Google Gemini | `google-genai` | `python-genai` | **v2.25.0** | `f15d1482d747` | 2026-09-29 | `>=1.72.0` | 1.73.1 | 🔴 major behind | +| OpenAI | `openai` | `openai-python` | **v3.22.1** | `58aca1dcfd8d` | 2026-09-30 | `>=3.0.0,<4` | 3.22.1 | ✅ current (live ✓) | +| Anthropic | `anthropic` | `anthropic-sdk-python` | **v1.9.0** | `a7285e919ab7` | 2026-09-28 | `>=1,<2` | 1.9.0 | 🟡 current, **not live-validated** (no API key) | +| Google Gemini | `google-genai` | `python-genai` | **v2.25.0** | `f15d1482d747` | 2026-09-29 | `>=2,<3` | 2.25.0 | ✅ current (live ✓, 47 models) | | Mistral | `mistralai` | `mistral-client-python` | **v3.0.0** | `e8dfa1c8a2d0` | 2026-09-28 | *(none — httpx only)* | not installed | 🟠 SDK unused | | OpenRouter | `openrouter` | `openrouter_python_sdk` | **v1.3.9** | `fd5ffce2995d` | 2026-09-30 | *(optional backend)* | not installed | 🟠 major behind + absent | -| Ollama | `ollama` | `ollama-python` | v0.6.3 | `8785556559ec` | 2026-09-28 | `>=0.6.0` | 0.6.1 | 🟡 patch behind | -| Deepgram | `deepgram-sdk` | `deepgram-python-sdk` | v7.11.0 | `a379a7f37b11` | 2026-09-28 | `>=7.0.0` | 7.3.1 | 🟡 minor behind | -| Hugging Face | `huggingface-hub` | `huggingface_hub` | `main` @ v0.9.0.rc1 | `1092497a9b65` | 2026-09-29 | *(none)* | 1.12.0 | 🟠 unpinned | -| Z.ai (GLM) | `zai-sdk` | `z-ai-sdk-python` | v0.2.3 | `ca5109c0aa9b` | 2026-06-16 | `>=0.2.0` | not installed | 🟡 SDK path untested | -| FriendliAI | `friendli` | `friendli-python` | v0.15.1 (pyproject 0.15.2) | `f3039e22ec0d` | 2026-09-28 | `>=0.15.1` | 0.15.1 | ✅ current | +| Ollama | `ollama` | `ollama-python` | v0.6.3 | `8785556559ec` | 2026-09-28 | `>=0.6.3` | 0.6.3 | ✅ current | +| Deepgram | `deepgram-sdk` | `deepgram-python-sdk` | v7.11.0 | `a379a7f37b11` | 2026-09-28 | `>=7.11.0` | 7.11.0 | ✅ current | +| Hugging Face | `huggingface-hub` | `huggingface_hub` | `main` @ v0.9.0.rc1 | `1092497a9b65` | 2026-09-29 | `>=1.12.0` | 1.12.0 | ✅ pinned | +| Z.ai (GLM) | `zai-sdk` | `z-ai-sdk-python` | v0.2.3 | `ca5109c0aa9b` | 2026-06-16 | `>=0.2.3` | 0.2.3 | ✅ current (SDK backend live ✓) | +| FriendliAI | `friendli` | `friendli-python` | v0.15.1 (repo pyproject reads 0.15.2, unreleased) | `f3039e22ec0d` | 2026-09-28 | `>=0.15.1` | 0.15.1 | ✅ current | | TypeSafe.ai | `typesafe-sdk` | `typesafe-sdk-python` | v0.7.2 | `f078f1e208a0` | 2026-09-26 | *(none — httpx only)* | not installed | ✅ by design | | Poe | `fastapi_poe` | `fastapi_poe` | 0.0.83 | `41ffd02e16f2` | 2026-01-21 | *(optional backend)* | not installed | 🟡 SDK path untested | | vLLM (self-hosted) | `vllm` (server) | `vllm` | v0.19.1rc0 | `219bb5b8c0dc` | 2026-04-16 | *(server, not a client dep)* | n/a | 🟡 clone stale | @@ -77,7 +78,7 @@ What this means for llmcore, **verified against the source**: | Does llmcore use `httpx` for its *own* clients? | **Yes** — 9 providers (`deepinfra`, `poe`, `mistral`, `friendli`, `kimi`, `openrouter`, `vllm`, `zai`, `typesafe`) plus all search providers. These keep using `httpx` and are unaffected. | | Do `respx`-based tests break? | **No.** Every `respx` mock in `tests/` targets llmcore's own `httpx` clients, never traffic routed through a vendor SDK. (Upstream vendored `tests/respx2` for their own suite; we don't need it.) | | **`httpx` is no longer installed transitively by `openai`** | 🔴 **Action required** — see below. | -| `certifi` is no longer installed by `openai`; httpx2 uses the **OS trust store** | 🔴 Deployment risk to document (minimal containers, TLS-inspecting proxies). Mitigate with `SSL_CERT_FILE` / `SSL_CERT_DIR`. | +| `certifi` is no longer installed by `openai`; httpx2 uses the **OS trust store** | ✅ Documented in `CONFIG_REFERENCE.md` § HTTP transport and TLS. | **The concrete break:** six providers import `httpx` but have **no extra of their own**, so today they only work because `openai` happened to install @@ -85,12 +86,12 @@ their own**, so today they only work because `openai` happened to install | Provider | Imports `httpx` | Own extra | |---|:---:|:---:| -| `mistral` | yes | ❌ missing | -| `kimi` | yes | ❌ missing | -| `poe` | yes | ❌ missing | -| `openrouter` | yes | ❌ missing | -| `vllm` | yes | ❌ missing | -| `huggingface` | — | ❌ missing | +| `mistral` | yes | ✅ added (Phase 0) | +| `kimi` | yes | ✅ added (Phase 0) | +| `poe` | yes | ✅ added (Phase 0) | +| `openrouter` | yes | ✅ added (Phase 0) | +| `vllm` | yes | ✅ added (Phase 0) | +| `huggingface` | — | ✅ added (Phase 0) | (`deepinfra`, `zai`, `friendli`, `typesafe` already declare `httpx` explicitly.) @@ -251,7 +252,7 @@ native-SDK path). | vllm | ❌ `/v1/embeddings`, `/pooling`, `/score`, `/rerank`; clone is 5 months stale | | huggingface | 🟠 unpinned dep; ❌ provider-level rerank; ❓ Inference Providers routing surface | | deepgram | 🟡 SDK 7.3.1 vs 7.11.0 — review new surface | -| zai | 🟡 `zai-sdk` never installed, so the preferred SDK backend is never exercised in CI | +| zai | ✅ Phase 0 installed `zai-sdk` and made the tests backend-hermetic; the SDK backend is now exercised in CI and validated live | | friendli | ⚠️ `/detokenize` + `/chat/render` 404 on Model APIs (upstream gap, tracked) | --- diff --git a/pyproject.toml b/pyproject.toml index 2c7ae2ab..530fca76 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,30 +53,61 @@ dependencies = [ ] [project.optional-dependencies] -# Dependencies for specific LLM providers -openai = ["openai>=2.31.0"] -anthropic = ["anthropic>=0.94.0"] -gemini = ["google-genai>=1.72.0", "google-api-core>=2.30.0"] -ollama = ["ollama>=0.6.0"] +# ---------------------------------------------------------------------------- +# LLM providers +# ---------------------------------------------------------------------------- +# NOTE on httpx: `openai` 3.x and `anthropic` 1.x moved their HTTP layer to +# httpx2 (https://httpx2.pydantic.dev/) and NO LONGER install `httpx` or +# `certifi` transitively. Every extra whose provider imports `httpx` directly +# must therefore declare it explicitly — previously they got it for free from +# `openai`. httpx2 also verifies TLS against the OS trust store rather than +# certifi; see docs/CONFIG_REFERENCE.md if you deploy into minimal containers +# or behind a TLS-inspecting proxy. +openai = ["openai>=3.0.0,<4"] +anthropic = ["anthropic>=1,<2"] +gemini = ["google-genai>=2,<3", "google-api-core>=2.30.0"] +ollama = ["ollama>=0.6.3"] # DeepInfra reuses the OpenAI SDK (OpenAI-compatible API) plus httpx for the # native audio (/v1/audio/*) and model-discovery (/v1/models) endpoints. -deepinfra = ["openai>=2.31.0", "httpx>=0.27.0"] +deepinfra = ["openai>=3.0.0,<4", "httpx>=0.27.0"] +# Mistral: the provider talks to the REST API over httpx directly (no vendor +# SDK today — `mistralai` v3.x is a candidate second backend, see +# docs/PROVIDER_MODERNIZATION_PLAN.md). +mistral = ["httpx>=0.27.0"] +# Kimi (Moonshot AI): OpenAI-compatible chat plus httpx for the Moonshot-only +# REST endpoints (token estimate, balance, file upload) and tiktoken for the +# local token-count fallback. Moonshot publishes no Python SDK. +kimi = ["openai>=3.0.0,<4", "httpx>=0.27.0", "tiktoken>=0.9.0"] +# Poe: OpenAI-compatible gateway plus httpx; `fastapi-poe` enables the optional +# native SSE backend (backend = "native"). +poe = ["openai>=3.0.0,<4", "httpx>=0.27.0"] +# OpenRouter: OpenAI-compatible gateway plus httpx; the `openrouter` SDK +# enables the optional native backend (backend = "sdk"). +openrouter = ["openai>=3.0.0,<4", "httpx>=0.27.0"] +# vLLM (self-hosted server): OpenAI-compatible surface plus httpx for +# /v1/models discovery of max_model_len. The `vllm` server itself is not a +# client dependency. +vllm = ["openai>=3.0.0,<4", "httpx>=0.27.0"] +# Hugging Face Inference API / Inference Providers routing. +huggingface = ["huggingface-hub>=1.12.0"] # Deepgram is a real-time voice/audio provider. The official deepgram-sdk is # WebSocket-native and pulls in the `websockets` library for streaming STT/TTS # and the Voice Agent. Install with: pip install llmcore[deepgram] -deepgram = ["deepgram-sdk>=7.0.0", "websockets>=12.0"] +deepgram = ["deepgram-sdk>=7.11.0", "websockets>=12.0"] # Z.ai (Zhipu AI / GLM). The provider prefers the official 'zai-sdk' (SDK # backend, default), and falls back to the 'openai' SDK (OpenAI-compatibility # mode) and/or direct 'httpx' calls (native media endpoints: image/audio/OCR/ # video generation, web search). Install with: pip install llmcore[zai] -zai = ["zai-sdk>=0.2.0", "openai>=2.31.0", "httpx>=0.27.0"] +zai = ["zai-sdk>=0.2.3", "openai>=3.0.0,<4", "httpx>=0.27.0"] # FriendliAI (Model APIs / Dedicated Endpoints / Container). The provider # prefers direct transports — the 'openai' SDK in OpenAI-compatibility mode # (default) or raw 'httpx' — because the official 'friendli' SDK's generated # response models drop fields outside the published schema (notably # reasoning_content). Install the vendor SDK too if you want backend = "sdk". # Install with: pip install llmcore[friendli] -friendli = ["openai>=2.31.0", "httpx>=0.27.0", "friendli>=0.15.1"] +# NOTE: the friendli-python repo's pyproject reads 0.15.2 but the newest +# release on PyPI is 0.15.1 — keep the pin at what is installable. +friendli = ["openai>=3.0.0,<4", "httpx>=0.27.0", "friendli>=0.15.1"] # TypeSafe.ai (System One typed judgments: noul/choice/score questions with # calibrated probabilities). NOT a chat API. The provider talks to the two REST # endpoints (/v1/systemone, /v1/models) directly over httpx — no vendor SDK is @@ -188,6 +219,12 @@ all = [ "llmcore[zai]", "llmcore[friendli]", "llmcore[typesafe]", + "llmcore[mistral]", + "llmcore[kimi]", + "llmcore[poe]", + "llmcore[openrouter]", + "llmcore[vllm]", + "llmcore[huggingface]", "llmcore[brightdata]", "llmcore[serper]", "llmcore[serpapi]", diff --git a/tests/providers/test_zai_provider.py b/tests/providers/test_zai_provider.py index e15f545d..45aed4e7 100644 --- a/tests/providers/test_zai_provider.py +++ b/tests/providers/test_zai_provider.py @@ -29,12 +29,18 @@ # Fixtures # --------------------------------------------------------------------------- +# ``backend`` is pinned so these tests exercise the AsyncOpenAI transport they +# mock, whether or not the optional ``zai-sdk`` happens to be installed. Without +# the pin, installing the SDK flips _resolve_backend() to "sdk" and every mock +# below is bypassed. Backend resolution itself is tested separately, with the +# availability flags patched explicitly. MINIMAL_CONFIG: dict[str, Any] = { "api_key": "test-zai-key-000", "default_model": "glm-5.2", "timeout": 30, "thinking": "enabled", "reasoning_effort": "high", + "backend": "openai", } @@ -615,28 +621,58 @@ def _make_sdk_provider(): return p, sdk_client +def _availability(*, sdk: bool, openai: bool = True, httpx: bool = True): + """Patch the module-level SDK availability flags for backend resolution. + + Resolution order is sdk -> openai -> httpx, so the outcome depends on what + is importable. Patching makes these assertions independent of whether the + optional ``zai-sdk`` is installed in the running environment. + """ + return patch.multiple( + "llmcore.providers.zai_provider", + zai_sdk_available=sdk, + openai_available=openai, + httpx_available=httpx, + ) + + class TestBackendResolution: - def test_auto_prefers_available(self): + def test_auto_prefers_sdk_when_available(self): + from llmcore.providers.zai_provider import ZaiProvider + + with _availability(sdk=True): + assert ZaiProvider._resolve_backend("auto") == "sdk" + assert ZaiProvider._resolve_backend(None) == "sdk" + + def test_auto_prefers_openai_without_sdk(self): + from llmcore.providers.zai_provider import ZaiProvider + + with _availability(sdk=False): + assert ZaiProvider._resolve_backend("auto") == "openai" + + def test_auto_falls_through_to_httpx(self): from llmcore.providers.zai_provider import ZaiProvider - # In this environment zai-sdk is absent, openai present → "openai". - assert ZaiProvider._resolve_backend("auto") == "openai" + with _availability(sdk=False, openai=False): + assert ZaiProvider._resolve_backend("auto") == "httpx" def test_explicit_httpx(self): from llmcore.providers.zai_provider import ZaiProvider - assert ZaiProvider._resolve_backend("httpx") == "httpx" + with _availability(sdk=True): + assert ZaiProvider._resolve_backend("httpx") == "httpx" def test_unavailable_sdk_falls_back(self): from llmcore.providers.zai_provider import ZaiProvider - # sdk unavailable here → falls through to openai - assert ZaiProvider._resolve_backend("sdk") == "openai" + with _availability(sdk=False): + assert ZaiProvider._resolve_backend("sdk") == "openai" def test_unknown_backend_uses_auto(self): from llmcore.providers.zai_provider import ZaiProvider - assert ZaiProvider._resolve_backend("bogus") == "openai" + with _availability(sdk=False): + assert ZaiProvider._resolve_backend("bogus") == "openai" class TestSDKBackend: From 22557234fafd85889727013e2f0da403d9c801cd Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Tue, 29 Sep 2026 23:56:50 -0300 Subject: [PATCH 08/11] docs: specify the media subsystem and remote GPU runtimes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two design+specification documents for the capability programs that add subsystems rather than extending providers. Specification only — no code. docs/MEDIA_SUBSYSTEM_SPEC.md turns the provider survey in /av/data/repos/docs/llmcore/researches into an llmcore-side design: a first-class llmcore.media subsystem with MediaArtifact / MediaUsage / MediaJob, capability Protocols per modality, and — the core distinction — three execution classes rather than one, because image generation is request/response, TTS is a byte stream and video is a long-running job. Capability metadata lives in model cards (extended with media/sourcing/policy blocks) instead of provider-specific tables in code; aggregators get the four-way split between who we call, whose weights, whose licence and whose AUP, and "uncensored" is represented honestly as supports_custom_weights plus provider_policy_applies rather than a boolean no vendor actually offers. It also records what the research could not see: - OpenAI's Sora video APIs were DEPRECATED in openai 3.1.0 (confirmed in the vendor CHANGELOG), so the survey's P0 "add Sora" item is dropped; frontier video comes from Veo and fal-hosted models. - llmcore already returns SpeechResult / ImageGenerationResult / OCRResult from seven providers, so models_multimodal types become views over MediaArtifact rather than being replaced. - Deepgram's surface is 12 public methods including a bidirectional voice agent, so the "refactor behind protocols" step is bigger than it looks — and it is the right first migration precisely because it exercises batch, realtime WebSocket and voice agent. - fal's own env var is FAL_KEY while the configured key is FAL_API_KEY; accept both, as the Friendli provider does for its three spellings. - Artifacts carry expires_at and checksum_sha256 from day one: every aggregator returns short-lived URLs, so storing a URI instead of bytes yields dead links. - Long media jobs get idempotency keys so a retried submit cannot double-bill. docs/COLAB_RUNTIME_SPEC.md designs llmcore.runtimes, a remote-compute abstraction with Colab as the first backend, after studying agent-lens's implemented design (391-line spec + ~3,750 lines across 13 modules) and the official google-colab-cli. The factoring argument: agent-lens already ends its bootstrap by registering the endpoint as an llmcore provider, which means the capability is being built on top of llmcore by a consumer and every other consumer must rebuild it. llmcore should own provisioning/bootstrap/tunnel/ lifecycle; agent-lens keeps its CLI and heuristics and deletes the duplication. Because the endpoint vLLM exposes is OpenAI-compatible, no new provider class is required — only dynamic instance registration in ProviderManager, which is also the one capability the media program needs. The safety model is the part that differs from every other provider: a Colab runtime bills per minute from assignment, not per request. Hence explicit-action-only provisioning (LLMCore.create() must never boot a VM), fail-closed bootstrap that releases the VM on any error, orphan detection so an unmonitored VM is visible, and a new max_lifetime_minutes hard cap on top of the reference idle reaper — an idle reaper does not protect against a runtime that is busy in a loop. Also records this session's live validation in the support matrix, including the two results that are not green: Anthropic 1.9.0 authenticates and maps errors correctly through the new major but every request returns "credit balance is too low", so no completion was validated; and Mistral's refreshed key works (46 models, open-mistral-nemo verified) but the configured default mistral-large-latest returns 403 — not in the account's tier. Co-Authored-By: Claude Opus 5 (1M context) --- CHANGELOG.md | 26 ++ README.md | 2 + docs/COLAB_RUNTIME_SPEC.md | 334 +++++++++++++++++++++ docs/MEDIA_SUBSYSTEM_SPEC.md | 443 ++++++++++++++++++++++++++++ docs/PROVIDER_MODERNIZATION_PLAN.md | 20 ++ docs/PROVIDER_SUPPORT_MATRIX.md | 28 +- 6 files changed, 849 insertions(+), 4 deletions(-) create mode 100644 docs/COLAB_RUNTIME_SPEC.md create mode 100644 docs/MEDIA_SUBSYSTEM_SPEC.md diff --git a/CHANGELOG.md b/CHANGELOG.md index 810215ef..7a028c3d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,32 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## Unreleased +### Added — media and remote-runtime specifications + +- `docs/MEDIA_SUBSYSTEM_SPEC.md` — design and specification for a first-class + `llmcore.media` subsystem: `MediaArtifact` / `MediaUsage` / `MediaJob`, + capability `Protocol`s per modality, the three execution classes + (request/response, byte stream, async job), capability-oriented model cards + with the aggregator sourcing split, generic webhooks with polling fallback, + an artifact store, a selection policy, and a nine-phase vendor rollout + (Deepgram refactor → OpenAI → Google/Veo → fal → ElevenLabs → Replicate → HF + Endpoints → direct specialists). Includes a backward-compatibility path that + keeps `BaseProvider`'s five media methods and the `models_multimodal` types + working. Notably corrects the source research: **OpenAI's Sora video APIs + were deprecated in `openai` 3.1**, so frontier video comes from Veo and + fal-hosted models instead. +- `docs/COLAB_RUNTIME_SPEC.md` — design and specification for a + `llmcore.runtimes` subsystem that provisions and controls remote GPU runtimes + (Google Colab first) and attaches the resulting OpenAI-compatible endpoint as + a provider instance, so a remotely served model is reachable through the + normal `llm.chat(provider_name=...)` path. Built on the study of agent-lens's + implemented design and BellaVox's process; proposes that llmcore own the + abstraction and agent-lens delegate to it. Adds a spend-ceiling + (`max_lifetime_minutes`) on top of the reference idle reaper, since an idle + reaper does not protect against a busy runaway runtime. + +Both documents are specification only — no implementation. + ### Changed — provider SDK majors (Phase 0 of the modernization plan) - **`openai` `>=3.0.0,<4`** (was `>=2.31.0`), **`anthropic` `>=1,<2`** (was diff --git a/README.md b/README.md index 564b9421..961ffe5f 100644 --- a/README.md +++ b/README.md @@ -1079,6 +1079,8 @@ from llmcore import ( - [FriendliAI provider usage](docs/Friendli_provider_usage.md) - [Provider support matrix](docs/PROVIDER_SUPPORT_MATRIX.md) — SDK/API versions we track per provider, plus the capability matrix - [Provider modernization plan](docs/PROVIDER_MODERNIZATION_PLAN.md) — phased plan to close the gaps in that matrix +- [Media subsystem spec](docs/MEDIA_SUBSYSTEM_SPEC.md) — design for first-class image/audio/video generation +- [Remote runtime spec](docs/COLAB_RUNTIME_SPEC.md) — design for serving models on remote GPUs (Colab first) - [TypeSafe.ai provider usage](docs/TypeSafe_provider_usage.md) - [`chat_with_usage` guide](docs/USAGE_chat_with_usage.md) - [Model cards](docs/model_cards.md) diff --git a/docs/COLAB_RUNTIME_SPEC.md b/docs/COLAB_RUNTIME_SPEC.md new file mode 100644 index 00000000..3cff5002 --- /dev/null +++ b/docs/COLAB_RUNTIME_SPEC.md @@ -0,0 +1,334 @@ +# Remote Compute Runtimes — Design & Specification (Colab first) + +Let llmcore provision, control and serve models on remote GPU runtimes — Google +Colab first — so a model running on someone else's GPU is just another provider. + +- **Status:** design + specification. Nothing implemented. +- **Written:** 2026-09-29 +- **Reference implementation studied:** `/av/repos/agent-lens` + (`docs/colab-design.md`, 391 lines + ~3,750 lines across 13 modules) and + `/av/repos/BellaVox` (the process agent-lens generalized) +- **Upstream CLI studied:** `/av/avalon/xrepos/google-colab-cli` (official) + +--- + +## 1. Problem statement and the factoring decision + +`agent-lens` already does this, and does it well: size a Hugging Face model, +pick a GPU SKU, boot a Colab VM, restore a cached venv + weights from Drive, +start vLLM on the VM, tunnel it to a local port over Colab's authenticated +WebSocket, **and register it as an llmcore provider**. + +That last step is the tell. The capability is being built *on top of* llmcore by +a consumer, so every other llmcore consumer has to rebuild it. Meanwhile llmcore +already owns everything the endpoint needs once it exists: the `vllm` provider +speaks the OpenAI-compatible surface vLLM exposes, `ProviderManager` owns +registration and lifecycle, and the model-card registry owns capability +metadata. + +**Decision: llmcore owns the runtime abstraction; agent-lens becomes a consumer +of it.** llmcore gains a `llmcore.runtimes` subsystem. `agent-lens colab` keeps +its CLI and its opinions (sizing heuristics, the "lens" adversarial plan review, +its dashboard) but delegates provisioning/bootstrap/tunnel/lifecycle to llmcore, +deleting most of its ~3,750 lines. + +Generalize one level, not two: the abstraction is **"remote compute runtime"**, +with Colab as the first backend and RunPod / Modal / Lambda / plain SSH as +plausible later ones. The interface is designed for that, but only Colab is +implemented. + +### Non-goals (v1) + +Inherited from the reference design, and correct: + +- No multi-tenant serving, queueing or cross-VM routing. +- No training or fine-tuning — inference only. +- No browser automation. We use the official CLI; if auth breaks, it fails loudly. +- No attempt to circumvent platform limits (no fake-activity tricks). + +--- + +## 2. Safety model — this one is different from every other provider + +Every other llmcore provider is stateless and costs money per request. **A Colab +runtime costs money per minute from the moment it is assigned, whether or not +anyone calls it.** The safety rules follow from that, and they are +non-negotiable design constraints rather than polish: + +1. **No implicit spend.** Nothing provisions or *keeps* a runtime without an + explicit caller action. `LLMCore.create()` must never boot a VM, no matter + what is in the config. Read-only operations are read-only. +2. **No implicit persistence of spend.** A runtime that llmcore started is + recorded in inspectable state so it can always be found and killed — a leaked + VM is a leaked credit card. +3. **Bounded by default.** Idle reaping is **on** by default (45 min), because + stopping saves money and the Drive cache makes restart cheap. Warn before the + axe falls. +4. **Fail closed.** Any bootstrap failure releases the VM. Never leave one + burning after an error. Keep the Drive cache. +5. **No implicit secrets.** HF tokens and Colab auth come from explicit config / + existing CLI state; never in argv, never in logs, never persisted on the VM. + +> **Improvement over the reference:** add a **spend ceiling**. A runtime carries +> `max_lifetime_minutes` (hard kill) and optionally `max_compute_units`. The +> reaper enforces both. The reference design has an idle reaper and notes +> Colab's ~12 h horizon but has no user-set hard cap; an idle reaper does not +> protect against a runtime that is *busy* in a loop. + +--- + +## 3. Architecture + +``` +LLMCore + ├── ProviderManager (existing) + │ └── dynamically registered instance → the runtime's endpoint + └── RuntimeManager (new) + ├── ColabRuntime backend: official colab CLI + SSH + tunnel + ├── (future) SSHRuntime / RunPodRuntime / ModalRuntime + ├── Sizer HF metadata → Plan (SKU, quant, ctx, VRAM) + ├── ServerRecipe vllm | llamacpp | tgi | custom + ├── CacheStore Drive (Colab) / volume (others): env + weights + ├── Tunnel local port ⇄ remote 127.0.0.1:port + ├── Keepalive + Reaper liveness, idle kill, hard lifetime cap + └── RuntimeState inspectable JSON under ~/.llmcore/runtimes/ +``` + +### 3.1 The runtime protocol + +```python +@runtime_checkable +class ComputeRuntime(Protocol): + name: str + + async def estimate(self, spec: ModelSpec) -> Plan: ... + async def up(self, plan: Plan, *, name: str) -> RuntimeHandle: ... + async def status(self, name: str | None = None) -> list[RuntimeStatus]: ... + async def logs(self, name: str, *, component: str, tail: int) -> AsyncIterator[str]: ... + async def down(self, name: str, *, release: bool = True) -> None: ... + async def adopt(self, external_id: str, *, name: str) -> RuntimeHandle: ... +``` + +`RuntimeHandle` carries what the provider layer needs: + +```python +@dataclass(slots=True) +class RuntimeHandle: + name: str + runtime: str # "colab" + external_id: str # colab session id + base_url: str # http://127.0.0.1:/v1 + served_model: str # the HF repo id vLLM was told to serve + api_style: str # "openai" → which llmcore provider to attach + recipe: str # "vllm" + sku: str # "L4" / "A100-40" / ... + started_at: datetime + idle_deadline: datetime | None + hard_deadline: datetime | None + state_path: Path +``` + +### 3.2 Provider attachment — the key integration + +Once a handle exists, its endpoint is OpenAI-compatible, so **no new provider +class is needed**. `RuntimeManager.attach()` registers a `vllm`-type provider +instance into the live `ProviderManager` under the runtime's name: + +```python +rt = await llm.runtimes.up("Qwen/Qwen3-30B-A3B-Instruct-2507", name="qwen30") +# -> provider instance "qwen30" now exists, type=vllm, base_url=the tunnel +answer = await llm.chat("Explain GQA briefly.", provider_name="qwen30") +await llm.runtimes.down("qwen30") # unregisters, then releases the VM +``` + +This is why the abstraction is cheap: llmcore already has the client. Three +small additions are needed to `ProviderManager`: + +- `register_instance(name, type, config)` / `unregister_instance(name)` — + dynamic registration at runtime (today providers are only built in `__init__`). +- Instances marked ephemeral so `close_all()` tears down runtimes it owns. +- `get_provider()` raising a clear error when a runtime-backed instance exists + but its runtime is `DEGRADED`. + +> **Improvement:** because the handle records `api_style`, a future recipe that +> speaks a different protocol (TGI, llama.cpp server) attaches a different +> provider type without touching the runtime layer. + +### 3.3 Sizing engine + +Ported from `agent-lens/colab/sizing.py`, which is already well-specified: + +1. **Metadata** — HF Hub API `safetensors` parameter map (exact params + dtypes); + fall back to file sizes × dtypes; fall back to the local `model_cards` + registry (offline). +2. **Weights** — Σ params × bytes/param, adjusted for quantization (explicit + `--quant`, else detected from the repo name: `-AWQ`, `-GPTQ`, `IQ4_XS`; GGUF + repos switch to the llama.cpp recipe). +3. **KV cache** — `2 × layers × kv_heads × head_dim × bytes × ctx`, GQA-aware, + per architecture family. Default ctx = `min(model_max, 32k)`. +4. **Runtime overhead** — ~2–3 GB activation/fragmentation, target + `gpu_memory_utilization = 0.90`. +5. **SKU ladder** — cheapest that fits with ≥15% headroom: + `T4 16GB → L4 24GB → G4 24GB → A100 40GB → A100 80GB → H100`, pruning SKUs + the account cannot get (a 400 from `colab new` prunes interactively). If it + does not fit even quantized, **refuse with a concrete smaller suggestion**. +6. Every number is printed with its arithmetic. `estimate` is a pure dry run. + +> **Improvement:** the sizer should write its `Plan` into the model card +> registry as a `runtime_hint` for that `(model, sku, quant, ctx)` tuple, so +> repeat launches skip the HF round trip and the ladder is learned rather than +> recomputed. + +### 3.4 Bootstrap sequence (Colab) + +Faithful to the reference, which is battle-tested: + +``` +1. estimate → Plan; print summary + burn rate; require explicit confirmation +2. colab new -s --gpu [--high-mem] +3. guard: session appears in assignments before any SSH ("never ssh into the void") +4. ssh master up (ProxyCommand = colab ssh --proxy-mode; isolated ed25519 key; + ControlMaster + ControlPersist) +5. push a versioned bootstrap bundle, then on the VM: +6. mount Drive +7. ensure pinned Python (runtime pin file in Drive) +8. ensure env: tar-restore from Drive, else pip build with PIP_CACHE_DIR on + Drive, then re-tar to Drive on a miss +9. ensure weights: local? Drive? else snapshot_download → copy to Drive + (sentinel files; allow_patterns = safetensors/config/tokenizer) +10. setsid serve --host 127.0.0.1 --port 8000 ... > server.log 2>&1 & + ready marker once /v1/models answers +11. local: ssh -N -L 127.0.0.1::127.0.0.1:8000 over the master +12. keepalive on; write state; attach provider (§3.2) +``` + +Any failure → log pointer + **release the VM**, keep the Drive cache. +The HF token is forwarded over SSH **stdin** to the download step only. + +A `bake` command pre-builds environment tars on a *CPU* VM so GPU minutes are +never spent on `pip install` — carried over from the reference and worth keeping. + +### 3.5 Keepalive, liveness, reaping + +- **Keepalive** — a kernel-side loop via `colab exec` stdin holds the kernel + active, which holds the VM; PID tracked locally. The server itself runs + `setsid`-detached so it survives kernel churn. +- **Liveness** — poll `GET /v1/models` through the tunnel every 60 s; three + consecutive failures → `DEGRADED` with a restart hint. Restarting the *server* + over the existing SSH master is cheap and safe; restarting the *VM* is not + automatic. +- **Idle reaper** — vLLM exposes no last-request metric, so count traffic on the + local tunnel side. After `idle_minutes` (default 45) → `down` + release. Warn + 10 minutes ahead in `status`. +- **Hard deadline** — `max_lifetime_minutes` (new, §2). Colab's own ~12 h + horizon is surfaced as an ETA, never circumvented. + +### 3.6 State + +`~/.llmcore/runtimes/.json` — small, human-readable, deletable: +the plan, the handle, PIDs (keepalive, tunnel, SSH master), timestamps, +deadlines, and the bootstrap log path. `status` reconciles state against +`colab ls`: sessions llmcore knows that are gone → mark stale; sessions that +exist but llmcore does not know → show as **orphans** with an `adopt` hint. +Orphan detection is a safety feature, not a nicety: an unknown running VM is +unmonitored spend. + +--- + +## 4. Public API and CLI + +```python +llm.runtimes.estimate("Qwen/Qwen3-30B-A3B-Instruct-2507", ctx=32768) +handle = await llm.runtimes.up(repo, name="qwen30", gpu="L4", idle_minutes=45) +await llm.runtimes.status() # list[RuntimeStatus] +await llm.runtimes.logs("qwen30", component="server", follow=True) +await llm.runtimes.down("qwen30") # unregister + release +await llm.runtimes.adopt("", name="rescued") +``` + +CLI (`llmcore-runtimes`, mirroring the proven agent-lens surface): + +``` +llmcore-runtimes estimate [--rev] [--ctx N] [--quant Q] +llmcore-runtimes up [--name] [--gpu SKU] [--ctx N] [--quant] + [--idle-min N] [--max-lifetime-min N] [--recipe] +llmcore-runtimes status [NAME] [-f] | ps [--json] | logs [NAME] [--component] +llmcore-runtimes keepalive on|off [NAME] +llmcore-runtimes bake [--recipe vllm] | cache ls|gc +llmcore-runtimes down [NAME] [--all] | adopt --name NAME +``` + +--- + +## 5. Configuration sketch + +```toml +[runtimes] +enabled = true # gate the subsystem; NEVER auto-provisions +default_backend = "colab" +state_dir = "~/.llmcore/runtimes" + +[runtimes.defaults] +recipe = "vllm" +idle_minutes = 45 # 0 disables the reaper +max_lifetime_minutes = 240 # hard kill regardless of activity +confirm_spend = true # require explicit confirmation before `up` +gpu_memory_utilization = 0.90 +headroom_fraction = 0.15 + +[runtimes.colab] +cli_path = "colab" # auto-discovered; install hint if missing +drive_cache_dir = "/content/drive/MyDrive/.llmcore-cache" +sku_ladder = ["T4", "L4", "G4", "A100-40", "A100-80", "H100"] +# hf_token_env_var = "HF_TOKEN" +``` + +--- + +## 6. Implementation plan + +| Phase | Scope | Gate | +|---|---|---| +| **R1** | `llmcore.runtimes` core: `ComputeRuntime` protocol, `ModelSpec`/`Plan`/`RuntimeHandle`/`RuntimeStatus`, `RuntimeState`, config section. `ProviderManager.register_instance()` / `unregister_instance()` + ephemeral teardown. A `FakeRuntime` for tests. **No network.** | Dynamic provider registration works and is covered | +| **R2** | `Sizer` — HF metadata, KV math, quant detection, SKU ladder, `estimate`. GET-only, no spend. | Sizing verified against several known models | +| **R3** | `ColabRuntime` — CLI discovery, `new`, assignment guard, SSH master, bundle push, Drive cache, vLLM recipe, tunnel, ready marker. `up`/`down`/`status`/`logs`. | One real model served end-to-end and reachable through `llm.chat()` | +| **R4** | Keepalive, liveness probe, idle reaper, hard deadline, orphan detection + `adopt`. | A leaked VM is impossible to create accidentally | +| **R5** | `bake`, Drive cache inventory + `cache gc`, `llamacpp`/GGUF recipe. | Cold start on a cached model is seconds | +| **R6** | CLI + docs; **agent-lens migration guide** so it delegates here. | agent-lens can drop its duplicated modules | + +**R1–R2 involve no spend at all** and are worth landing early: they are pure +computation and unlock `estimate` as a useful standalone tool. + +--- + +## 7. Risks + +| Risk | Mitigation | +|---|---| +| **Runaway spend** — the defining risk | §2: explicit-action-only, idle reaper on, hard lifetime cap, orphan detection, state always inspectable | +| Colab CLI is **not installed** on this machine (`colab` not on PATH) and is Linux/macOS only | Discover at call time, fail with the `uv tool install google-colab-cli` hint; never a hard dependency of llmcore | +| Colab auth/quota errors (400/412) | Surface verbatim in `status` with the next action; prune the SKU ladder interactively | +| Upstream CLI is young; flags may move | Pin a tested CLI version range in the docs; parse `--json` output where offered, never scrape human text | +| Platform limits (~12 h, ~90 min idle) | Surface as ETAs; never circumvent | +| Tunnel dies silently | Liveness probe + `DEGRADED` state + cheap server-only restart | +| Drive cache corruption | Sentinel files per artifact; `cache gc`; re-download on sentinel mismatch | + +--- + +## 8. Open questions + +1. **Does the runtimes subsystem belong in llmcore core, or in an extra?** + *Recommendation: `llmcore[runtimes]`* — it needs no new hard dependency, but + the Colab backend shells out to an external CLI, which is unusual for llmcore + and should be opt-in. +2. **Should `agent-lens` migrate in the same cycle,** or should llmcore ship R1–R4 + and let agent-lens migrate when convenient? The duplicated logic will drift. +3. **Where do the sizing heuristics live?** agent-lens has an opinionated, + working sizer with an optional LLM plan review. Move the arithmetic to + llmcore and leave the "lens" adversarial review in agent-lens? +4. **Multi-runtime routing** is a non-goal for v1, but should `RuntimeHandle` + carry enough (cost/sku/latency) for a future router? *Recommendation: yes — + the fields are free now and retrofitting them is not.* +5. **Colab auth** — needs a Google account OAuth via the CLI, not an API key. + Confirm the intended account before R3, since it is the account that gets + billed in compute units. diff --git a/docs/MEDIA_SUBSYSTEM_SPEC.md b/docs/MEDIA_SUBSYSTEM_SPEC.md new file mode 100644 index 00000000..1648e00b --- /dev/null +++ b/docs/MEDIA_SUBSYSTEM_SPEC.md @@ -0,0 +1,443 @@ +# Media Subsystem — Design & Specification + +Generative image, audio and video as a first-class `llmcore` subsystem, plus the +provider adapters that sit behind it. + +- **Status:** design + specification. Nothing implemented. +- **Written:** 2026-09-29 +- **Primary input:** `/av/data/repos/docs/llmcore/researches/image-audio-video_providers_2026september.md` + (the provider survey and priority matrix; this document is the llmcore-side design) +- **Related:** [`PROVIDER_SUPPORT_MATRIX.md`](PROVIDER_SUPPORT_MATRIX.md), + [`PROVIDER_MODERNIZATION_PLAN.md`](PROVIDER_MODERNIZATION_PLAN.md) + +--- + +## 1. Problem statement + +llmcore already reaches media APIs, but through a surface that does not +generalize: + +- `BaseProvider` carries five optional media methods (`generate_speech`, + `transcribe_audio`, `generate_image`, `ocr`, `create_embeddings`), each + defaulting to `NotImplementedError`. +- Seven providers implement some subset: `zai` (the most complete — image, TTS, + STT, OCR, video, web search), `deepinfra`, `mistral`, `huggingface`, + `openai`, `friendli`, `deepgram`. +- `deepgram` is the extreme case: 12 provider-specific public methods + (`open_voice_agent`, `transcribe_stream_flux`, `stream_speech`, …) because + the facade has nowhere to put realtime audio. +- `zai.generate_video()` is the only video generation in the codebase, under a + provider-specific name with no shared job model. + +Three structural problems follow: + +1. **No shared execution model.** Image generation is request/response, TTS is a + byte stream, video is a long-running job. The current surface assumes + request/response and forces the other two into provider-private methods. +2. **No capability discovery.** A caller cannot ask "which configured provider + can do video interpolation?" without hardcoding provider names. +3. **No normalized results.** `models_multimodal.py` has `SpeechResult`, + `TranscriptionResult`, `ImageGenerationResult`, `GeneratedImage`, `OCRResult` + — but nothing for video, nothing for jobs, and no common artifact type. + +**The fix is not more provider classes.** It is a media subsystem the providers +plug into. + +--- + +## 2. Design + +### 2.1 Shape + +``` +LLMCore + ├── ProviderManager (existing — chat) + └── MediaManager (new) + ├── AudioRouter tts · asr · music · sfx · realtime + ├── ImageRouter generate · edit · upscale · variate + ├── VideoRouter generate · edit · interpolate · reframe · upscale + ├── MediaJobManager poll · webhook · cancel · resume + └── ArtifactStore bytes/URI persistence + checksums +``` + +Routers resolve `(capability, model?, provider?)` → an adapter, using model +cards for capability metadata (§2.5) and a selection policy (§2.7). They do not +contain vendor logic. + +### 2.2 Three execution classes — the core distinction + +| Class | Examples | Surface | +|---|---|---| +| **Request/response** | image generation & edit, batch ASR, one-shot TTS | `await media.images.generate(...) -> MediaResult` | +| **Byte / event stream** | streaming TTS, realtime ASR, voice agents | `async for chunk in media.audio.stream_tts(...)` / async session object | +| **Long-running job** | Veo, fal queue, Replicate predictions, Luma generations | `job = await media.video.generate(...)` → `MediaJob`, then poll or webhook | + +Modelling streaming as a job (or a job as request/response) is the main +abstraction mistake to avoid. Deepgram already proves the streaming semantics +are irreducible. + +### 2.3 Core types + +New module `src/llmcore/media/models.py`. These are the provider-independent +contract; adapters translate to and from vendor shapes. + +```python +class MediaKind(StrEnum): + AUDIO = "audio"; IMAGE = "image"; VIDEO = "video"; TEXT = "text" + +class MediaCapability(StrEnum): + # audio + TTS = "tts"; TTS_STREAM = "tts_stream" + ASR = "asr"; ASR_STREAM = "asr_stream" + VOICE_AGENT = "voice_agent" + MUSIC = "music"; SFX = "sfx"; VOICE_DESIGN = "voice_design" + # image + IMAGE_GENERATE = "image_generate"; IMAGE_EDIT = "image_edit" + IMAGE_UPSCALE = "image_upscale"; IMAGE_VARIATE = "image_variate" + OCR = "ocr" + # video + VIDEO_GENERATE = "video_generate"; VIDEO_EDIT = "video_edit" + VIDEO_INTERPOLATE = "video_interpolate"; VIDEO_REFRAME = "video_reframe" + VIDEO_UPSCALE = "video_upscale"; VIDEO_EXTEND = "video_extend" + +class MediaExecution(StrEnum): + REQUEST_RESPONSE = "request_response"; STREAM = "stream"; ASYNC_JOB = "async_job" + +class MediaJobStatus(StrEnum): + QUEUED = "queued"; RUNNING = "running"; SUCCEEDED = "succeeded" + FAILED = "failed"; CANCELED = "canceled"; EXPIRED = "expired" +``` + +`MediaArtifact` — one produced asset: + +```python +@dataclass(frozen=True, slots=True) +class MediaArtifact: + kind: MediaKind + uri: str | None = None # provider URL, or artifact-store URI + data: bytes | None = None # inline bytes (small results) + mime_type: str | None = None + width: int | None = None + height: int | None = None + duration_seconds: float | None = None + sample_rate_hz: int | None = None + fps: float | None = None + frame_count: int | None = None + checksum_sha256: str | None = None + expires_at: datetime | None = None # provider URLs are usually temporary + provenance: MediaProvenance | None = None + provider_metadata: Mapping[str, Any] = field(default_factory=dict) +``` + +`MediaUsage` — normalized billing units (§2.6), `MediaJob` — the async handle, +`MediaResult` — the request/response wrapper (`artifacts`, `usage`, `model`, +`provider`, `raw`). + +> **Addition beyond the research doc:** `expires_at` and `checksum_sha256` are +> load-bearing, not nice-to-have. Every aggregator returns short-lived URLs; a +> caller that stores the URI instead of the bytes gets a dead link hours later. +> The `ArtifactStore` uses both to decide what to materialize and to dedupe. + +### 2.4 Capability protocols + +`typing.Protocol` classes in `src/llmcore/media/protocols.py`, one per +capability group. Adapters implement only what they support; routers check with +`isinstance(adapter, ImageGenerationProvider)` (runtime-checkable). + +```python +@runtime_checkable +class ImageGenerationProvider(Protocol): + async def generate_image_media( + self, prompt: str, *, model: str | None = None, n: int = 1, + size: str | None = None, seed: int | None = None, + reference_images: Sequence[MediaRef] | None = None, **kwargs: Any, + ) -> MediaResult | MediaJob: ... + +@runtime_checkable +class StreamingTTSProvider(Protocol): + async def stream_tts( + self, text: str, *, model: str | None = None, voice: str | None = None, + sample_rate_hz: int | None = None, **kwargs: Any, + ) -> AsyncIterator[bytes]: ... + +@runtime_checkable +class VideoGenerationProvider(Protocol): + async def generate_video( + self, prompt: str, *, model: str | None = None, + first_frame: MediaRef | None = None, last_frame: MediaRef | None = None, + duration_seconds: float | None = None, resolution: str | None = None, + with_audio: bool | None = None, **kwargs: Any, + ) -> MediaJob: ... +``` + +`MediaRef` is the input counterpart of `MediaArtifact`: a URL, local path, raw +bytes, or a previously produced `MediaArtifact`. Adapters upload/inline as their +API requires — callers never hand-roll base64. + +> Note the method name `generate_image_media`, not `generate_image`: the latter +> is already taken on `BaseProvider` with a different return type. §4 covers the +> migration; the new names are temporary and collapse at the 1.0 boundary. + +### 2.5 Model cards carry the capabilities + +Extend the existing card schema rather than adding provider-specific tables in +code. New optional `media` block: + +```json +{ + "model_id": "fal-ai/film/video", + "provider": "fal", + "model_type": "video", + "media": { + "kind": "video", + "capabilities": ["video_interpolate"], + "execution": "async_job", + "supports_webhooks": true, + "inputs": {"video": true, "image": false, "text": false}, + "outputs": {"video": true}, + "max_duration_seconds": 10, + "resolutions": ["720p", "1080p"] + }, + "sourcing": { + "model_owner": "google-research", + "model_family": "film", + "model_license": "apache-2.0", + "hosting_provider": "fal", + "hosting_policy": "fal-aup" + }, + "policy": { + "supports_custom_weights": false, + "provider_policy_applies": true, + "commercial_use": "model_specific" + } +} +``` + +Two deliberate choices from the research, both adopted: + +- **Aggregators need four separate concepts**, not one `provider` field: + `provider` (who we call) vs `model_owner` / `model_family` / + `model_license` (whose weights) vs `hosting_policy` (whose AUP). Without this, + "can I use this commercially?" is unanswerable for fal/Replicate/HF. +- **There is no "uncensored" boolean.** The survey found no mainstream managed + API that credibly promises policy-free generation. The honest representation + is two orthogonal facts: `supports_custom_weights` (can modified/ablated + weights run?) and `provider_policy_applies` (does the host still enforce an + AUP?). Replicate, HF Endpoints and BFL's open-weight route are `true`/`true`. + +`cardctl` gains a `--kind media` mode and per-provider media adapters. + +### 2.6 Cost normalization without pretending + +Media vendors bill in incompatible units: per-image, per-megapixel, per-second +of video, per-audio-minute, per-character, per-token, per-compute-second. +`MediaUsage` keeps whichever units the vendor reported *and* an +`estimated_cost_usd` with a `pricing_as_of` stamp and a `basis` label. Callers +that need exactness read the native fields; dashboards read the estimate and can +see how stale the pricing is. Never synthesize a token count for a video. + +### 2.7 Provider selection policy + +`media..(...)` resolves in this order: + +1. Explicit `provider=` → use it, or raise if it lacks the capability. +2. Explicit `model=` → resolve the owning provider from cards. +3. Configured `[media.routing]` preference list for that capability. +4. The research doc's workload defaults (§"Final provider choices") as built-in + fallbacks — ElevenLabs for TTS, Deepgram for ASR, OpenAI/Google for frontier + image, fal for breadth, etc. +5. Otherwise raise `MediaCapabilityError` listing which configured providers + *could* satisfy it if enabled. + +Constraint filters compose with all of the above: `require_commercial_use=True`, +`require_custom_weights=True`, `max_cost_usd=...`, `exclude_providers=[...]`. + +### 2.8 Webhooks as a core facility + +Async jobs need a callback path. A generic receiver — not per-provider: + +- `MediaJobManager` issues a signed, single-use callback token per job. +- An optional ASGI app (`llmcore.media.webhooks:app`) mounts at a configured + path; the existing `llmcore[bridge]` server can host it. +- Providers that support webhooks register the URL; providers that don't are + polled with capped exponential backoff. +- **Polling is always the fallback**, so llmcore stays usable with no public + ingress — the common case for local development. + +### 2.9 Artifact store + +Reuse llmcore's storage layer rather than inventing one. `ArtifactStore` +persists bytes to a configured backend (filesystem by default, with the existing +SQLite/Postgres metadata store for the index), keyed by +`checksum_sha256`. Policy: `materialize = "always" | "on_expiry" | "never"`. +Default `on_expiry` — fetch and store before the provider URL dies. + +--- + +## 3. Public API + +```python +async with await LLMCore.create() as llm: + # request/response + img = await llm.media.images.generate("an orange tabby", size="1024x1024") + img.artifacts[0].uri + + # edit with a reference + edited = await llm.media.images.edit( + "make it night-time", image=MediaRef.from_path("cat.png") + ) + + # streaming TTS + async for chunk in llm.media.audio.stream_tts("Hello there", voice="rachel"): + speaker.write(chunk) + + # long-running job + job = await llm.media.video.generate("a drone shot over dunes", + duration_seconds=8, with_audio=True) + job = await llm.media.jobs.wait(job, timeout=600) # poll or webhook + await llm.media.artifacts.download(job.artifacts[0], "dunes.mp4") + + # capability discovery + llm.media.capabilities() # {capability: [providers]} + llm.media.who_can(MediaCapability.VIDEO_INTERPOLATE) +``` + +--- + +## 4. Migration and backward compatibility + +Non-negotiable: **no existing call site breaks.** + +1. **Phase M1 adds only new code.** `llmcore.media` lands with the types, + protocols, routers and job manager, plus an in-repo fake adapter for tests. +2. **`BaseProvider`'s five media methods stay**, and keep their current return + types (`SpeechResult`, `TranscriptionResult`, `ImageGenerationResult`, + `OCRResult`). They become thin shims that call the new subsystem when the + provider has an adapter, and keep their current implementation otherwise. +3. **`models_multimodal.py` types become views over `MediaArtifact`.** They are + public API today (returned by seven providers), so they get + `from_artifact()` / `to_artifact()` and a deprecation note, not deletion. +4. **Deepgram is the reference migration** (research doc's second merge, adopted). + Its 12 provider-specific methods stay as the compatibility surface while the + streaming protocols are implemented behind them. It is the only integration + that already exercises batch + realtime WebSocket + voice agent, so it + validates the hard parts of the abstraction before any new vendor lands. +5. **`zai.generate_video()`** is renamed into the protocol with an alias kept. + +--- + +## 5. Implementation plan + +Order follows the research doc's rollout, with llmcore-specific gates. + +| Phase | Scope | Gate | +|---|---|---| +| **M1** | `llmcore.media` core: types, protocols, routers, `MediaJobManager` (poll only), `ArtifactStore`, card schema `media`/`sourcing`/`policy` blocks, config section, fake adapter + tests | No provider work lands before this | +| **M2** | Refactor **Deepgram** behind the audio protocols; keep its public methods | Realtime + batch both pass through the abstraction | +| **M3** | **OpenAI** media: images generate/edit, TTS, ASR, realtime audio | Lowest marginal cost — adapter already exists | +| **M4** | **Google** media: Imagen/Nano-Banana images, **Veo** video (async job), native TTS | First true async-job provider; validates §2.8 | +| **M5** | **fal** — queue/webhook lifecycle, URL inputs, video, SFX, FILM interpolation | The provider-neutrality test: if the abstraction bends here, fix the abstraction | +| **M6** | **ElevenLabs** — batch + realtime STT, TTS, SFX, music, voice design | Consent/provenance as first-class metadata | +| **M7** | **Replicate** — one generic prediction adapter + model-schema descriptors | Explicitly *not* a class per model | +| **M8** | **Hugging Face Inference Endpoints** — configurable endpoint/schema adapter | Custom weights / private repos path | +| **M9** | Webhook receiver, then direct specialists (BFL, Luma, Stability) when justified: lower unit cost, first-party-only feature, data contract, or pre-aggregator access | Otherwise fal/Replicate already cover it | + +Runway stays on the watchlist — the research run could not verify its current +API contract, and the doc is explicit about not freezing a guessed model id. + +--- + +## 6. Corrections and additions to the research + +Flagging these because the research snapshot and the code disagree, or the +research could not see llmcore internals: + +1. **OpenAI Sora is deprecated.** The survey lists Sora video as a P0 reason to + extend the OpenAI adapter. But `openai` 3.1.0 (2026-08-14) shipped + *"**api:** deprecate Sora video APIs"* — confirmed in + `/av/avalon/xrepos/openai-python/CHANGELOG.md`. **Do not build the Sora + adapter.** Frontier video comes from **Veo (M4)** and fal-hosted models (M5). + This is the single most important correction: M3 shrinks to images + speech. +2. **Reuse `models_multimodal.py`, don't parallel it.** The research proposes + fresh types without knowing llmcore already returns `SpeechResult` / + `ImageGenerationResult` from seven providers. §4.3 keeps them as views. +3. **Deepgram's surface is bigger than the survey assumes** — 12 public methods + including a bidirectional voice agent. The M2 refactor is a larger job than + "refactor behind protocols" implies; budget for it. +4. **`zai` is already the most media-complete provider** (image, TTS, STT, OCR, + video, web search). It is a better second migration target than the survey + suggests, and a free second data point on whether the protocols generalize. +5. **Env var naming.** fal's own convention is `FAL_KEY`; the key in + `/av/data/dbs/.env` is `FAL_API_KEY`. Adapters should accept both, in the + order `config → api_key_env_var → FAL_KEY → FAL_API_KEY` — the same + multi-spelling tolerance the Friendli provider uses. +6. **Provenance is emerging as a real field.** Several vendors now emit C2PA / + content-credential metadata. `MediaProvenance` is in the artifact from day + one so it is not retrofitted later. +7. **Idempotency.** Long-running media jobs are expensive; a retried submit + must not double-bill. `MediaJobManager` sends a client-generated + idempotency key where the vendor supports one, and always records the + submitted key so a resumed process re-attaches instead of resubmitting. + +--- + +## 7. Testing + +- **Offline**: `respx` for httpx adapters, `AsyncMock` on vendor SDK clients; + pin transport `backend` explicitly (the rule in the modernization plan §9). +- **Fake adapter**: an in-repo provider implementing every protocol, driving + router/selection/job tests with no network. +- **Contract tests**: parametrized over all registered media adapters — + protocol conformance, artifact normalization, usage population, error mapping. +- **Job lifecycle**: simulated queue→running→succeeded, failure, cancel, + expiry, webhook-vs-poll equivalence. +- **Live smokes**: key-gated under `tests/integration/`, one per vendor, using + the cheapest model and smallest output. + +--- + +## 8. Configuration sketch + +```toml +[media] +default_image_provider = "openai" +default_audio_provider = "elevenlabs" +default_video_provider = "google" +artifact_materialize = "on_expiry" # always | on_expiry | never +artifact_path = "~/.llmcore/media" + +[media.routing] +tts = ["elevenlabs", "openai", "deepgram"] +asr = ["deepgram", "openai", "elevenlabs"] +image_generate = ["openai", "google", "fal"] +video_generate = ["google", "fal", "replicate"] +video_interpolate = ["fal", "replicate"] + +[media.jobs] +poll_initial_seconds = 2 +poll_max_seconds = 30 +job_timeout_seconds = 1800 +webhook_base_url = "" # empty = poll only + +[media.providers.fal] +# api_key_env_var = "FAL_KEY" # FAL_API_KEY also accepted +[media.providers.elevenlabs] +# api_key_env_var = "ELEVENLABS_API_KEY" +``` + +--- + +## 9. Open questions + +1. **`[media.providers.*]` vs `[providers.*]`** — a separate section (clean + separation, some duplication for OpenAI/Google which appear in both), or + extend the existing provider sections (single source of credentials, but + `[providers.fal]` would be a chat provider that cannot chat)? *Recommendation: + extend `[providers.*]`, with `[media.routing]` separate — one credential per + vendor, and the capability matrix already tolerates non-chat providers + (`deepgram`, `typesafe`).* +2. **Artifact retention default** — `on_expiry` costs disk silently. Acceptable? +3. **Does the bridge expose media?** The gRPC/HTTP bridge would need artifact + streaming; defer to a later phase or design in now? +4. **Keys still needed**: `REPLICATE_API_TOKEN` (M7), `BFL_API_KEY` / + `LUMA_API_KEY` (M9). `ELEVENLABS_API_KEY` and `FAL_API_KEY` are available. diff --git a/docs/PROVIDER_MODERNIZATION_PLAN.md b/docs/PROVIDER_MODERNIZATION_PLAN.md index 45d3f424..a2f8fc6e 100644 --- a/docs/PROVIDER_MODERNIZATION_PLAN.md +++ b/docs/PROVIDER_MODERNIZATION_PLAN.md @@ -295,6 +295,26 @@ Per provider touched: --- +## 11. Companion specifications + +Two capability programs are specified separately because they add subsystems +rather than extending providers: + +- [`MEDIA_SUBSYSTEM_SPEC.md`](MEDIA_SUBSYSTEM_SPEC.md) — a first-class + `llmcore.media` subsystem (image/audio/video), the capability protocols + providers implement, async `MediaJob` lifecycle, capability-oriented model + cards, and the vendor rollout (Deepgram refactor → OpenAI → Google/Veo → fal → + ElevenLabs → Replicate → HF Endpoints). This **supersedes** Phase 4 of this + plan (Gemini media) and Phase 1.3's `generate_video()` placeholder, which + become M4 and part of M1 there. +- [`COLAB_RUNTIME_SPEC.md`](COLAB_RUNTIME_SPEC.md) — a `llmcore.runtimes` + subsystem that provisions and controls remote GPU runtimes (Colab first) and + attaches the resulting OpenAI-compatible endpoint as a provider instance. + Requires the only new `ProviderManager` capability either program needs: + dynamic instance registration. + +--- + ## Appendix A — capability extraction script Regenerates matrix §4/§4.1 from the source rather than by hand: diff --git a/docs/PROVIDER_SUPPORT_MATRIX.md b/docs/PROVIDER_SUPPORT_MATRIX.md index d2b5c4e0..2d219728 100644 --- a/docs/PROVIDER_SUPPORT_MATRIX.md +++ b/docs/PROVIDER_SUPPORT_MATRIX.md @@ -44,9 +44,9 @@ dev venv. A pin that trails the upstream **major** version is a red flag. | Provider | SDK package | Clone (`/av/avalon/xrepos/…`) | Upstream tag | Upstream commit | Tag date | llmcore pin | Installed | Status | |---|---|---|---|---|---|---|---|:---:| | OpenAI | `openai` | `openai-python` | **v3.22.1** | `58aca1dcfd8d` | 2026-09-30 | `>=3.0.0,<4` | 3.22.1 | ✅ current (live ✓) | -| Anthropic | `anthropic` | `anthropic-sdk-python` | **v1.9.0** | `a7285e919ab7` | 2026-09-28 | `>=1,<2` | 1.9.0 | 🟡 current, **not live-validated** (no API key) | +| Anthropic | `anthropic` | `anthropic-sdk-python` | **v1.9.0** | `a7285e919ab7` | 2026-09-28 | `>=1,<2` | 1.9.0 | 🟡 transport ✓, **no completion** — account credit balance is zero | | Google Gemini | `google-genai` | `python-genai` | **v2.25.0** | `f15d1482d747` | 2026-09-29 | `>=2,<3` | 2.25.0 | ✅ current (live ✓, 47 models) | -| Mistral | `mistralai` | `mistral-client-python` | **v3.0.0** | `e8dfa1c8a2d0` | 2026-09-28 | *(none — httpx only)* | not installed | 🟠 SDK unused | +| Mistral | `mistralai` | `mistral-client-python` | **v3.0.0** | `e8dfa1c8a2d0` | 2026-09-28 | *(none — httpx only)* | not installed | 🟠 SDK unused (httpx path live ✓, 46 models) | | OpenRouter | `openrouter` | `openrouter_python_sdk` | **v1.3.9** | `fd5ffce2995d` | 2026-09-30 | *(optional backend)* | not installed | 🟠 major behind + absent | | Ollama | `ollama` | `ollama-python` | v0.6.3 | `8785556559ec` | 2026-09-28 | `>=0.6.3` | 0.6.3 | ✅ current | | Deepgram | `deepgram-sdk` | `deepgram-python-sdk` | v7.11.0 | `a379a7f37b11` | 2026-09-28 | `>=7.11.0` | 7.11.0 | ✅ current | @@ -114,7 +114,7 @@ lookups, context budgets and cost estimates. | openrouter | `openai/gpt-4o-mini` | 🔴 stale | | poe | `GPT-4o-Mini` | 🔴 stale | | vllm | `meta-llama/Llama-3.1-8B-Instruct` | 🟡 example value, self-hosted | -| mistral | `mistral-large-latest` | ✅ alias, self-updating | +| mistral | `mistral-large-latest` | 🔴 alias resolves, but **403 — not in this account's tier**; `open-mistral-nemo` verified working | | huggingface | `meta-llama/Llama-3.3-70B-Instruct` | 🟡 stale example | | deepinfra | `deepseek-ai/DeepSeek-V3` | 🟡 V3.2 is served | | kimi | `kimi-k2.6` | ✅ current at last provider audit | @@ -300,13 +300,33 @@ any provider change. --- -## 7. Related documents +## 7. Live validation log + +Recorded per audit so "current" always means "we called it". + +| Date | Provider | SDK | Result | +|---|---|---|---| +| 2026-09-29 | OpenAI | `openai` 3.22.1 | ✅ completion + usage | +| 2026-09-29 | Google Gemini | `google-genai` 2.25.0 | ✅ completion; 47 models discovered | +| 2026-09-29 | Z.ai | `zai-sdk` 0.2.3 (**native SDK backend**) | ✅ completion + reasoning tokens | +| 2026-09-29 | DeepSeek | via `openai` 3.22.1 | ✅ completion + cache/reasoning usage | +| 2026-09-29 | Mistral | httpx path | ✅ 46 models; `open-mistral-nemo` completion. `mistral-large-latest` → 403 (tier) | +| 2026-09-29 | Anthropic | `anthropic` 1.9.0 | 🟡 auth + error mapping ✓ through the new major, but every request returns `invalid_request_error: credit balance is too low` — **no completion validated** | +| 2026-09-20 | FriendliAI | `openai`/`httpx`/`friendli` | ✅ all three backends, streaming, tools, team billing | + +--- + +## 8. Related documents - [`PROVIDER_MODERNIZATION_PLAN.md`](PROVIDER_MODERNIZATION_PLAN.md) — the phased plan that closes the gaps listed in §5 - [`model_cards.md`](model_cards.md) — card schema, the canonical reasoning-effort vocabulary, and the cardctl workflow - [`CONFIG_REFERENCE.md`](CONFIG_REFERENCE.md) — every provider config key +- [`MEDIA_SUBSYSTEM_SPEC.md`](MEDIA_SUBSYSTEM_SPEC.md) — design/spec for the + image/audio/video subsystem and its provider adapters +- [`COLAB_RUNTIME_SPEC.md`](COLAB_RUNTIME_SPEC.md) — design/spec for remote GPU + runtimes (Colab first), so a remotely served model is just another provider - Per-provider guides: [`Friendli_provider_usage.md`](Friendli_provider_usage.md), [`Deepgram_provider_usage.md`](Deepgram_provider_usage.md), [`TypeSafe_provider_usage.md`](TypeSafe_provider_usage.md) From c890cdedee2989e59cb9b3e769511448b4c7b562 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Wed, 30 Sep 2026 02:00:02 -0300 Subject: [PATCH 09/11] feat(media): implement the media subsystem core (M1) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase M1 of docs/MEDIA_SUBSYSTEM_SPEC.md: the llmcore.media subsystem, reached through llm.media, plus the one ProviderManager capability that both the media and remote-runtime programs need. No vendor adapters — that is the gate the spec requires before any provider work lands, so that each implementation does not establish its own incompatible conventions. The central design choice is three execution classes rather than one. MediaResult covers request/response (image generation, batch ASR), AsyncIterator[bytes] covers byte streams (TTS, realtime ASR), and MediaJob covers long-running work (video, queue-based vendors). Deepgram's existing 12-method provider-private surface is what happens without that distinction. Core pieces: - models.py — MediaKind, 19 MediaCapability values, MediaExecution, MediaJobStatus, MediaRef, MediaArtifact, MediaProvenance, MediaUsage, MediaResult, MediaJob. MediaRef accepts url/path/bytes/artifact so callers never hand-roll base64. MediaArtifact carries expires_at and checksum_sha256 because every aggregator returns short-lived URLs and a caller who stores the URI gets a dead link hours later. MediaUsage keeps the vendor's own billing units (images, megapixels, seconds, audio minutes, characters, compute seconds) with a stamped estimate, rather than synthesizing a token count for a video. - protocols.py — runtime-checkable Protocols per capability group, plus a CAPABILITY_PROTOCOLS table so adding a capability is one entry rather than an edit to routing code. A capability an adapter declares but does not back is dropped with a warning: better a missing capability than a confident AttributeError at call time. - manager.py — MediaManager with per-modality routers, capability discovery, and the documented resolution order (explicit provider, explicit model, [media.routing], built-in defaults, any capable adapter). Adapters ARE the chat providers: any [providers.*] instance implementing the protocols becomes one, so there is one credential per vendor and no parallel [media.providers.*] tree to keep in sync. - jobs.py — MediaJobManager owns polling, capped backoff with jitter, timeouts and cancellation so no adapter writes its own loop. A timeout raises WITHOUT cancelling the job: the handle stays valid and can be waited on again, because an expensive generation must not be thrown away over a client-side deadline. That is also why the deadline is an explicit parameter rather than asyncio.timeout, which would cancel the task. - artifacts.py — content-addressed store, sharded by SHA-256, atomic publish via a .part rename, with always/on_expiry/never policies. The byte fetcher is injected, so the store carries no hard httpx dependency. - testing.py — FakeMediaProvider implements every protocol and ships inside the package, so downstream projects building adapters can use it too. All 134 new tests run against it: no network, no vendor account. ProviderManager gains register_instance() / unregister_instance() with ephemeral tracking. Providers were previously only constructible during __init__, but a subsystem that creates an endpoint has to add one afterwards — a booted Colab VM's OpenAI-compatible endpoint registers as a vllm instance (docs/COLAB_RUNTIME_SPEC.md §3.2). Guards: a name collision raises unless replace=True so a live provider is never silently swapped out from under its callers; the configured default cannot be unregistered; construction failures surface as ConfigError; a failing close() is logged rather than blocking teardown. Also fixes a real bug the new tests caught: the polling backoff computed 2 ** attempt, which stops converting to float past ~1024 polls and killed the wait loop with OverflowError. A multi-hour video job polled every few seconds reaches that. The exponent is now capped and regression tested at 10,000 polls. Backward compatible throughout: BaseProvider's five media methods and the models_multimodal result types are untouched, so the seven providers using them keep working. Until M2 migrates them, llm.media.adapter_names is empty and says so honestly rather than pretending to route. Verified: 134 new tests; full unit suite 5298 passed, 29 skipped; ruff CI gate clean; llm.media exercised end-to-end through LLMCore.create(). Co-Authored-By: Claude Opus 5 (1M context) --- CHANGELOG.md | 66 +++ docs/MEDIA_SUBSYSTEM_SPEC.md | 5 +- src/llmcore/api.py | 47 ++ src/llmcore/config/default_config.toml | 65 +++ src/llmcore/exceptions.py | 112 +++++ src/llmcore/media/__init__.py | 102 +++++ src/llmcore/media/artifacts.py | 272 +++++++++++ src/llmcore/media/jobs.py | 278 +++++++++++ src/llmcore/media/manager.py | 387 ++++++++++++++++ src/llmcore/media/models.py | 612 +++++++++++++++++++++++++ src/llmcore/media/protocols.py | 432 +++++++++++++++++ src/llmcore/media/routers.py | 407 ++++++++++++++++ src/llmcore/media/testing.py | 285 ++++++++++++ src/llmcore/providers/manager.py | 119 +++++ tests/media/__init__.py | 0 tests/media/test_media_integration.py | 224 +++++++++ tests/media/test_media_models.py | 228 +++++++++ tests/media/test_media_subsystem.py | 566 +++++++++++++++++++++++ 18 files changed, 4205 insertions(+), 2 deletions(-) create mode 100644 src/llmcore/media/__init__.py create mode 100644 src/llmcore/media/artifacts.py create mode 100644 src/llmcore/media/jobs.py create mode 100644 src/llmcore/media/manager.py create mode 100644 src/llmcore/media/models.py create mode 100644 src/llmcore/media/protocols.py create mode 100644 src/llmcore/media/routers.py create mode 100644 src/llmcore/media/testing.py create mode 100644 tests/media/__init__.py create mode 100644 tests/media/test_media_integration.py create mode 100644 tests/media/test_media_models.py create mode 100644 tests/media/test_media_subsystem.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 7a028c3d..57091360 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,72 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## Unreleased +### Added — media subsystem core (M1) + +- **`llmcore.media`**, reached through `llm.media`: a sibling subsystem to chat + providers and search providers for generative image, audio and video. + Implements phase M1 of `docs/MEDIA_SUBSYSTEM_SPEC.md`; no vendor adapters yet, + which is the gate the spec requires before any provider work lands. +- **Three execution classes, not one** — `MediaResult` for request/response, + `AsyncIterator[bytes]` for streams, `MediaJob` for long-running work. Image + generation, TTS and video generation genuinely differ, and collapsing them + into one shape is the modelling mistake the design avoids. +- **Types**: `MediaKind`, `MediaCapability` (19 capabilities), `MediaExecution`, + `MediaJobStatus`, `MediaRef` (url/path/bytes/artifact inputs, so callers never + hand-roll base64), `MediaArtifact` (with `expires_at` + `checksum_sha256`, + because every aggregator returns short-lived URLs), `MediaProvenance`, + `MediaUsage` (keeps the vendor's native billing units rather than inventing a + token count), `MediaResult` and `MediaJob`. +- **Capability protocols** (`typing.Protocol`, runtime-checkable) — routers + discover what an adapter can do with `isinstance`, so routing logic never + names a provider. A capability declared but not backed by its protocol is + dropped with a warning rather than failing at call time. +- **`MediaManager`** with per-modality routers (`images`, `audio`, `video`), + capability discovery (`capabilities()`, `who_can()`), and a documented + resolution order: explicit provider → explicit model → `[media.routing]` → + built-in defaults → any capable adapter. **Adapters are the chat providers**: + any `[providers.*]` instance implementing the protocols becomes a media + adapter, so there is one credential per vendor and nothing to duplicate. +- **`MediaJobManager`** owns polling, capped exponential backoff with jitter, + timeouts and cancellation, so no adapter writes its own poll loop. A timeout + raises **without cancelling the job** — the handle stays valid and can be + waited on again, because an expensive video generation must not be discarded + over a client-side deadline. +- **`ArtifactStore`** — content-addressed, sharded by SHA-256, atomic publish, + with `always` / `on_expiry` (default) / `never` materialization policies. The + byte fetcher is injected, so the store has no hard dependency on `httpx`. +- **`FakeMediaProvider`** ships inside the package (not under `tests/`) so + downstream projects building adapters can use it too. It implements every + protocol, and is what the 134 new tests run against — no network, no account. +- **`[media]` config section** (artifact policy/path, `[media.routing]` + preferences per capability, `[media.jobs]` poll/timeout policy). Entirely + optional: omit it and `llm.media` still works on built-in defaults. +- **Backward compatible**: `BaseProvider`'s five media methods and the + `models_multimodal` result types are untouched. Providers gain routing when + they are migrated in M2 onward; until then `llm.media.adapter_names` is empty + and reports so honestly. + +### Added — dynamic provider registration + +- **`ProviderManager.register_instance()` / `unregister_instance()`** plus + `is_ephemeral()` / `ephemeral_instances`. Providers were previously only + constructible during `__init__`; subsystems that *create* endpoints need to + add one afterwards. This is the single capability shared by the media program + and the remote-runtime program (`docs/COLAB_RUNTIME_SPEC.md`), where a Colab + VM's OpenAI-compatible endpoint is registered as a `vllm` instance. +- Guards that matter: a name collision raises unless `replace=True` (so a live + provider is never silently swapped out from under its callers), the configured + default provider cannot be unregistered, construction failures surface as + `ConfigError`, and a failing `close()` during unregister is logged rather than + blocking teardown. + +### Fixed + +- **Job-polling backoff could overflow.** `2 ** attempt` stops converting to + float past ~1024 polls, so the wait loop would die with `OverflowError` — a + multi-hour video job polled every few seconds would actually reach that. The + exponent is now capped; regression tested at 10,000 polls. + ### Added — media and remote-runtime specifications - `docs/MEDIA_SUBSYSTEM_SPEC.md` — design and specification for a first-class diff --git a/docs/MEDIA_SUBSYSTEM_SPEC.md b/docs/MEDIA_SUBSYSTEM_SPEC.md index 1648e00b..40a7be9e 100644 --- a/docs/MEDIA_SUBSYSTEM_SPEC.md +++ b/docs/MEDIA_SUBSYSTEM_SPEC.md @@ -3,7 +3,8 @@ Generative image, audio and video as a first-class `llmcore` subsystem, plus the provider adapters that sit behind it. -- **Status:** design + specification. Nothing implemented. +- **Status:** **M1 implemented** (core types, protocols, routers, job manager, + artifact store, config, fake adapter). M2 onward not started — see §5. - **Written:** 2026-09-29 - **Primary input:** `/av/data/repos/docs/llmcore/researches/image-audio-video_providers_2026september.md` (the provider survey and priority matrix; this document is the llmcore-side design) @@ -332,7 +333,7 @@ Order follows the research doc's rollout, with llmcore-specific gates. | Phase | Scope | Gate | |---|---|---| -| **M1** | `llmcore.media` core: types, protocols, routers, `MediaJobManager` (poll only), `ArtifactStore`, card schema `media`/`sourcing`/`policy` blocks, config section, fake adapter + tests | No provider work lands before this | +| **M1** ✅ | `llmcore.media` core: types, protocols, routers, `MediaJobManager` (poll only), `ArtifactStore`, config section, fake adapter + tests. *Card schema blocks deferred to M2, where the first real adapter needs them.* | Landed 2026-09-30, 134 tests | | **M2** | Refactor **Deepgram** behind the audio protocols; keep its public methods | Realtime + batch both pass through the abstraction | | **M3** | **OpenAI** media: images generate/edit, TTS, ASR, realtime audio | Lowest marginal cost — adapter already exists | | **M4** | **Google** media: Imagen/Nano-Banana images, **Veo** video (async job), native TTS | First true async-job provider; validates §2.8 | diff --git a/src/llmcore/api.py b/src/llmcore/api.py index 854f5a14..7065bb0d 100644 --- a/src/llmcore/api.py +++ b/src/llmcore/api.py @@ -217,6 +217,7 @@ def __init__(self): self._runtime_config_dirty = False self._original_config_dict = {} self._observability = None + self._media_manager: Any | None = None # Grimoire control plane (populated by _initialize_from_config) self._grimoire: Any | None = None self._grimoire_config: Any | None = None @@ -236,6 +237,35 @@ def prompt_registry(self) -> Any: """The instance-level prompt registry (grimoire-backed adapter).""" return self._prompt_registry + @property + def media(self) -> Any: + """The media subsystem: generative image, audio and video. + + Routers hang off this accessor:: + + await llm.media.images.generate("an orange tabby") + async for chunk in llm.media.audio.stream_tts("hello"): + ... + job = await llm.media.video.generate("a drone shot over dunes") + result = await llm.media.wait(job) + + Adapters are the configured chat providers that implement the media + protocols, so no separate credentials are needed. Capability discovery + (``llm.media.capabilities()`` / ``llm.media.who_can(...)``) reports what + the current configuration can actually do. + + Returns: + The :class:`~llmcore.media.MediaManager` for this instance. + + Raises: + ConfigError: If accessed before ``LLMCore.create()`` finished. + """ + if self._media_manager is None: + raise ConfigError( + "The media subsystem is not initialized. Use 'await LLMCore.create()'." + ) + return self._media_manager + @classmethod async def create( cls, @@ -516,6 +546,17 @@ async def _initialize_from_config( ) await self._search_provider_manager.initialize() + logger.debug("Initializing MediaManager...") + # Media is an optional capability layered on the SAME provider + # instances as chat: any provider implementing the media protocols + # becomes an adapter, so there is one credential per vendor. Never + # fails when [media] is absent, so existing configs are unaffected. + from llmcore.media import MediaManager + + self._media_manager = MediaManager.from_provider_manager( + self._provider_manager, self.config.get + ) + logger.debug("Initializing SessionManager...") self._session_manager = SessionManager(self._storage_manager.session_storage) @@ -602,6 +643,12 @@ async def close(self) -> None: """ logger.info("Closing LLMCore instance...") try: + # Media holds no connections of its own (adapters are the chat + # providers), but it logs any job left running so an expensive + # generation is not silently abandoned. + media_mgr = getattr(self, "_media_manager", None) + if media_mgr is not None: + await media_mgr.close() await self._provider_manager.close_all() # Search manager may not exist if init failed very early; guard it. search_mgr = getattr(self, "_search_provider_manager", None) diff --git a/src/llmcore/config/default_config.toml b/src/llmcore/config/default_config.toml index 1291376d..8919c713 100644 --- a/src/llmcore/config/default_config.toml +++ b/src/llmcore/config/default_config.toml @@ -1239,6 +1239,71 @@ max_retries = 2 +# ============================================================================== +# [media] - Generative Media (image / audio / video) +# ============================================================================== +# The media subsystem is the sibling of [providers] (chat) and +# [search_providers] (web/data search). It is reached through `llm.media`: +# +# await llm.media.images.generate("an orange tabby") +# async for chunk in llm.media.audio.stream_tts("hello"): ... +# job = await llm.media.video.generate("a drone shot over dunes") +# result = await llm.media.wait(job) +# +# IMPORTANT: media adapters ARE the chat providers above. Any provider +# instance that implements the media protocols becomes an adapter, so there is +# ONE credential per vendor and nothing to duplicate here. This section only +# configures routing, artifact handling and job policy. +# +# Media is entirely optional: omit this section and `llm.media` still works +# with built-in defaults, reporting whatever the configured providers can do +# (`llm.media.capabilities()` / `llm.media.who_can("video_generate")`). +# +# See docs/MEDIA_SUBSYSTEM_SPEC.md for the design. +[media] + + # --- Artifact handling --- + # Media vendors return SHORT-LIVED URLs. Storing a URI instead of the bytes + # yields a dead link hours later, so artifacts can be materialized locally: + # "on_expiry" (default) - fetch only artifacts whose URL carries an expiry + # "always" - fetch every remote artifact + # "never" - never fetch; you handle URIs yourself + artifact_materialize = "on_expiry" + + # Content-addressed store for materialized bytes (sharded by SHA-256). + artifact_path = "~/.llmcore/media" + + # --- Routing --- + # Preference order per capability. The first CONFIGURED provider that + # supports the capability wins; unconfigured names are skipped, so a list can + # name providers you have not set up yet. Omit a capability to use the + # built-in default order (see llmcore/media/manager.py: DEFAULT_ROUTING). + # + # Capability names: tts, tts_stream, asr, asr_stream, voice_agent, music, + # sfx, voice_design, image_generate, image_edit, image_upscale, + # image_variate, ocr, video_generate, video_edit, video_interpolate, + # video_reframe, video_upscale, video_extend. + [media.routing] + # tts = ["elevenlabs", "openai", "deepgram"] + # asr = ["deepgram", "openai", "elevenlabs"] + # image_generate = ["openai", "gemini", "fal"] + # video_generate = ["gemini", "fal", "replicate"] + # video_interpolate = ["fal", "replicate"] + + # --- Long-running jobs --- + # Video generation and queue-based vendors are asynchronous: submit, then + # wait. Polling is always available; webhooks (a later phase) only ever + # short-circuit the wait, so llmcore stays usable with no public ingress. + [media.jobs] + # First sleep after submission, in seconds. + poll_initial_seconds = 2 + # Ceiling for the exponential backoff, in seconds. + poll_max_seconds = 30 + # Default wall-clock budget for `llm.media.wait(job)`. A timeout does NOT + # cancel the job - the handle stays valid and can be waited on again, so an + # expensive generation is never discarded over a client-side deadline. + job_timeout_seconds = 1800 + # ============================================================================== # [search_providers] - Web / Data Search Provider Configurations # ============================================================================== diff --git a/src/llmcore/exceptions.py b/src/llmcore/exceptions.py index 782f404e..d534dbd4 100644 --- a/src/llmcore/exceptions.py +++ b/src/llmcore/exceptions.py @@ -268,6 +268,118 @@ def __init__( super().__init__(f"Error with search provider '{provider_name}': {message}{detail}") +# ============================================================================= +# MEDIA SUBSYSTEM EXCEPTIONS +# ============================================================================= + + +class MediaError(LLMCoreError): + """Base class for errors raised by the media subsystem. + + The media-side analogue of :class:`ProviderError`. Used by + :mod:`llmcore.media` for transport faults, unsupported operations, and + job-lifecycle failures. + + Attributes: + provider_name: Media provider instance involved, when known. + capability: The media capability being attempted, when known. + """ + + def __init__( + self, + message: str = "Media error.", + *, + provider_name: str | None = None, + capability: str | None = None, + ): + self.provider_name = provider_name + self.capability = capability + bits = [] + if provider_name: + bits.append(f"provider='{provider_name}'") + if capability: + bits.append(f"capability='{capability}'") + detail = f" ({', '.join(bits)})" if bits else "" + super().__init__(f"{message}{detail}") + + +class MediaCapabilityError(MediaError): + """Raised when no configured provider can satisfy a media capability. + + Carries the providers that *could* satisfy it if enabled, so the message can + tell the caller what to configure rather than just what failed. + + Attributes: + candidates: Providers known to support the capability but not currently + usable (unconfigured, missing dependency, or excluded by a filter). + """ + + def __init__( + self, + message: str = "No provider satisfies this media capability.", + *, + capability: str | None = None, + candidates: list[str] | None = None, + ): + self.candidates = candidates or [] + if self.candidates: + message = f"{message} Providers that could satisfy it: {', '.join(self.candidates)}." + super().__init__(message, capability=capability) + + +class MediaJobError(MediaError): + """Raised when a long-running media job fails or is unrecoverable. + + Attributes: + job_id: The llmcore-local job id. + status: Terminal status the job reached. + """ + + def __init__( + self, + message: str = "Media job failed.", + *, + job_id: str | None = None, + status: str | None = None, + provider_name: str | None = None, + capability: str | None = None, + ): + self.job_id = job_id + self.status = status + bits = [] + if job_id: + bits.append(f"job={job_id}") + if status: + bits.append(f"status={status}") + prefix = f"[{', '.join(bits)}] " if bits else "" + super().__init__( + f"{prefix}{message}", provider_name=provider_name, capability=capability + ) + + +class MediaJobTimeoutError(MediaJobError): + """Raised when a media job does not reach a terminal state in time. + + The job itself is *not* cancelled — the handle remains valid and can be + waited on again, so an expensive video generation is never thrown away just + because a client-side deadline elapsed. + """ + + def __init__( + self, + message: str = "Media job did not finish before the timeout.", + *, + job_id: str | None = None, + status: str | None = None, + timeout_seconds: float | None = None, + provider_name: str | None = None, + ): + self.timeout_seconds = timeout_seconds + if timeout_seconds is not None: + message = f"{message} Waited {timeout_seconds:g}s; the job is still live." + super().__init__(message, job_id=job_id, status=status, provider_name=provider_name) + + # ============================================================================= # MODERATION EXCEPTIONS (plan SF-1 / DDS-06) # ============================================================================= diff --git a/src/llmcore/media/__init__.py b/src/llmcore/media/__init__.py new file mode 100644 index 00000000..27ca1898 --- /dev/null +++ b/src/llmcore/media/__init__.py @@ -0,0 +1,102 @@ +# src/llmcore/media/__init__.py +"""Generative media (image, audio, video) for LLMCore. + +A sibling subsystem to chat providers and search providers, reached through +:attr:`llmcore.LLMCore.media`. See ``docs/MEDIA_SUBSYSTEM_SPEC.md`` for the +design, and :mod:`llmcore.media.protocols` for what a provider adapter must +implement. + +Quick shape:: + + result = await llm.media.images.generate("an orange tabby") + async for chunk in llm.media.audio.stream_tts("hello"): + ... + job = await llm.media.video.generate("a drone shot over dunes") + result = await llm.media.wait(job) +""" + +from .artifacts import ArtifactStore, MaterializePolicy +from .jobs import JobPolicy, MediaJobManager +from .manager import MediaManager +from .models import ( + AUDIO_CAPABILITIES, + CAPABILITY_KINDS, + IMAGE_CAPABILITIES, + TERMINAL_JOB_STATUSES, + VIDEO_CAPABILITIES, + MediaArtifact, + MediaCapability, + MediaExecution, + MediaJob, + MediaJobStatus, + MediaKind, + MediaProvenance, + MediaRef, + MediaResult, + MediaUsage, +) +from .protocols import ( + CAPABILITY_PROTOCOLS, + ASRProvider, + ImageEditProvider, + ImageGenerationProvider, + ImageUpscaleProvider, + MediaCapableProvider, + MediaJobPoller, + MusicProvider, + OCRMediaProvider, + SFXProvider, + StreamingASRProvider, + StreamingTTSProvider, + TTSProvider, + VideoEditProvider, + VideoGenerationProvider, + VideoInterpolationProvider, +) +from .routers import AudioRouter, ImageRouter, VideoRouter + +__all__ = [ + # manager + routers + "MediaManager", + "ImageRouter", + "AudioRouter", + "VideoRouter", + # models + "MediaKind", + "MediaCapability", + "MediaExecution", + "MediaJobStatus", + "MediaRef", + "MediaArtifact", + "MediaProvenance", + "MediaUsage", + "MediaResult", + "MediaJob", + "AUDIO_CAPABILITIES", + "IMAGE_CAPABILITIES", + "VIDEO_CAPABILITIES", + "CAPABILITY_KINDS", + "TERMINAL_JOB_STATUSES", + # jobs + artifacts + "MediaJobManager", + "JobPolicy", + "ArtifactStore", + "MaterializePolicy", + # protocols + "MediaCapableProvider", + "MediaJobPoller", + "ImageGenerationProvider", + "ImageEditProvider", + "ImageUpscaleProvider", + "OCRMediaProvider", + "TTSProvider", + "StreamingTTSProvider", + "ASRProvider", + "StreamingASRProvider", + "MusicProvider", + "SFXProvider", + "VideoGenerationProvider", + "VideoEditProvider", + "VideoInterpolationProvider", + "CAPABILITY_PROTOCOLS", +] diff --git a/src/llmcore/media/artifacts.py b/src/llmcore/media/artifacts.py new file mode 100644 index 00000000..3373b4e6 --- /dev/null +++ b/src/llmcore/media/artifacts.py @@ -0,0 +1,272 @@ +# src/llmcore/media/artifacts.py +"""Persistence for produced media assets. + +Every media aggregator returns **short-lived URLs**. A caller that stores the +URI instead of the bytes gets a dead link hours later, which is why +:class:`~llmcore.media.models.MediaArtifact` carries ``expires_at`` and +``checksum_sha256`` and why this store exists. + +:class:`ArtifactStore` is content-addressed: bytes land at a path derived from +their SHA-256, so re-materializing the same asset is free and two providers +returning identical output cost one file. The fetch itself is injected rather +than hard-wired, so the store has no opinion about HTTP clients and no hard +dependency on ``httpx``. +""" + +from __future__ import annotations + +import hashlib +import logging +from collections.abc import Awaitable, Callable +from enum import StrEnum +from pathlib import Path + +from ..exceptions import MediaError +from .models import MediaArtifact + +logger = logging.getLogger(__name__) + +__all__ = ["DEFAULT_ARTIFACT_PATH", "ArtifactStore", "MaterializePolicy"] + +DEFAULT_ARTIFACT_PATH = "~/.llmcore/media" + +#: Signature of a byte fetcher: URL in, bytes out. +Fetcher = Callable[[str], Awaitable[bytes]] + + +class MaterializePolicy(StrEnum): + """When to fetch and persist remote artifact bytes. + + Attributes: + ALWAYS: Fetch every remote artifact. + ON_EXPIRY: Fetch only artifacts whose URI carries an expiry — the + default, because it protects against dead links without + downloading assets the caller may never read. + NEVER: Never fetch; callers handle URIs themselves. + """ + + ALWAYS = "always" + ON_EXPIRY = "on_expiry" + NEVER = "never" + + +class ArtifactStore: + """Content-addressed local store for media artifacts. + + Args: + base_path: Root directory; ``~`` is expanded. Created lazily on first + write, so constructing a store never touches the filesystem. + policy: When :meth:`materialize` should fetch remote bytes. + fetcher: Async callable that downloads a URL. Required only for + materialization of remote artifacts. + """ + + def __init__( + self, + base_path: str | Path = DEFAULT_ARTIFACT_PATH, + *, + policy: MaterializePolicy | str = MaterializePolicy.ON_EXPIRY, + fetcher: Fetcher | None = None, + ) -> None: + self.base_path = Path(base_path).expanduser() + self.policy = MaterializePolicy(policy) + self._fetcher = fetcher + + # --- addressing --- + + def path_for(self, checksum: str, *, suffix: str = "") -> Path: + """Return the on-disk path for a content hash. + + Sharded two levels by hash prefix so a large store stays navigable. + + Args: + checksum: Hex SHA-256 digest. + suffix: Optional file extension, including the dot. + """ + return self.base_path / checksum[:2] / checksum[2:4] / f"{checksum}{suffix}" + + def has(self, checksum: str, *, suffix: str = "") -> bool: + """Whether the content hash is already stored.""" + return self.path_for(checksum, suffix=suffix).exists() + + # --- writing --- + + def put(self, data: bytes, *, suffix: str = "") -> tuple[str, Path]: + """Store *data* and return its ``(checksum, path)``. + + Idempotent: storing identical bytes twice writes once. + """ + checksum = hashlib.sha256(data).hexdigest() + path = self.path_for(checksum, suffix=suffix) + if not path.exists(): + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".part") + tmp.write_bytes(data) + tmp.replace(path) # atomic publish, so a crash never leaves a partial file + logger.debug("Stored media artifact %s (%d bytes)", checksum[:12], len(data)) + return checksum, path + + def should_materialize(self, artifact: MediaArtifact) -> bool: + """Whether *artifact* should be fetched under the active policy.""" + if artifact.data is not None or not artifact.uri: + return False + if self.policy is MaterializePolicy.NEVER: + return False + if self.policy is MaterializePolicy.ALWAYS: + return True + return artifact.expires_at is not None + + async def materialize( + self, + artifact: MediaArtifact, + *, + fetcher: Fetcher | None = None, + force: bool = False, + ) -> MediaArtifact: + """Fetch and persist *artifact*'s bytes when the policy calls for it. + + Args: + artifact: The artifact to materialize. + fetcher: Overrides the store's fetcher for this call. + force: Materialize regardless of policy (but never re-fetch bytes + the artifact already carries). + + Returns: + An artifact carrying ``data``, ``checksum_sha256`` and a local + ``file://`` URI — or the original, unchanged, when the policy says + not to fetch. + + Raises: + MediaError: If materialization is required but no fetcher is + available, or the download fails. + """ + if artifact.data is not None: + return artifact + if not (force or self.should_materialize(artifact)): + return artifact + + fetch = fetcher or self._fetcher + if fetch is None: + raise MediaError( + "Artifact materialization requires a fetcher; none was configured. " + "Install httpx or pass fetcher= explicitly." + ) + if not artifact.uri: + raise MediaError("Artifact has no uri to materialize.") + + try: + data = await fetch(artifact.uri) + except MediaError: + raise + except Exception as e: + raise MediaError(f"Failed to materialize artifact from {artifact.uri}: {e}") from e + + suffix = _suffix_for(artifact) + checksum, path = self.put(data, suffix=suffix) + stored = artifact.with_data(data) + # Re-point the URI at the local copy; the provider URL is preserved in + # provider_metadata so provenance is not lost. + return MediaArtifact( + **{ + **{f: getattr(stored, f) for f in stored.__slots__}, + "uri": path.as_uri(), + "checksum_sha256": checksum, + "expires_at": None, + "provider_metadata": { + **dict(stored.provider_metadata), + "source_uri": artifact.uri, + }, + } + ) + + # --- reading --- + + async def download( + self, + artifact: MediaArtifact, + destination: str | Path, + *, + fetcher: Fetcher | None = None, + ) -> Path: + """Write *artifact* to *destination* and return the path. + + Uses inline bytes when present, the local store when already + materialized, and the fetcher otherwise. + """ + dest = Path(destination).expanduser() + dest.parent.mkdir(parents=True, exist_ok=True) + + if artifact.data is not None: + dest.write_bytes(artifact.data) + return dest + + materialized = await self.materialize(artifact, fetcher=fetcher, force=True) + if materialized.data is None: # pragma: no cover - defensive + raise MediaError("Materialization produced no bytes.") + dest.write_bytes(materialized.data) + return dest + + def gc(self, *, keep_checksums: set[str] | None = None) -> int: + """Delete stored files not in *keep_checksums*. + + Args: + keep_checksums: Digests to retain; everything else is removed. When + ``None``, nothing is deleted (a no-op, so a mistaken call cannot + wipe the store). + + Returns: + Number of files removed. + """ + if keep_checksums is None: + return 0 + removed = 0 + if not self.base_path.exists(): + return 0 + for path in self.base_path.rglob("*"): + if path.is_file() and path.stem not in keep_checksums: + path.unlink() + removed += 1 + return removed + + +def _suffix_for(artifact: MediaArtifact) -> str: + """Best-effort file extension for an artifact.""" + import mimetypes + + if artifact.mime_type: + ext = mimetypes.guess_extension(artifact.mime_type) + if ext: + return ext + if artifact.uri: + tail = Path(artifact.uri.split("?", 1)[0]).suffix + if tail: + return tail + return "" + + +def default_fetcher(timeout: float = 120.0) -> Fetcher: + """Return an ``httpx``-backed fetcher, if httpx is installed. + + Args: + timeout: Per-request timeout in seconds. + + Returns: + An async fetcher suitable for :class:`ArtifactStore`. + + Raises: + MediaError: If ``httpx`` is not available. + """ + try: + import httpx + except ImportError as e: # pragma: no cover - httpx is present in practice + raise MediaError( + "The 'httpx' package is required to download media artifacts." + ) from e + + async def _fetch(url: str) -> bytes: + async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client: + resp = await client.get(url) + resp.raise_for_status() + return resp.content + + return _fetch diff --git a/src/llmcore/media/jobs.py b/src/llmcore/media/jobs.py new file mode 100644 index 00000000..679f79a8 --- /dev/null +++ b/src/llmcore/media/jobs.py @@ -0,0 +1,278 @@ +# src/llmcore/media/jobs.py +"""Lifecycle management for long-running media jobs. + +Video generation, queue-based image vendors and aggregator predictions are all +asynchronous: submit, then wait. :class:`MediaJobManager` owns that waiting — +the backoff schedule, the timeout policy, cancellation and the job registry — so +that no provider adapter implements its own polling loop. Adapters only supply +:meth:`~llmcore.media.protocols.MediaJobPoller.poll_media_job`. + +Polling is always available. Webhooks (a later phase) are an optimization that +short-circuits the wait; they never become a requirement, because the common +development case has no public ingress. +""" + +from __future__ import annotations + +import asyncio +import logging +import random +import time +from collections.abc import Callable, Iterable +from typing import Any + +from ..exceptions import MediaJobError, MediaJobTimeoutError +from .models import MediaJob, MediaJobStatus +from .protocols import MediaJobPoller + +logger = logging.getLogger(__name__) + +__all__ = ["JobPolicy", "MediaJobManager"] + +#: Defaults chosen so a short image job feels responsive while a 10-minute video +#: job does not generate hundreds of requests. +DEFAULT_POLL_INITIAL_SECONDS = 2.0 +DEFAULT_POLL_MAX_SECONDS = 30.0 +DEFAULT_JOB_TIMEOUT_SECONDS = 1800.0 + + +class JobPolicy: + """Backoff and timeout policy for job polling. + + Attributes: + poll_initial_seconds: First sleep after submission. + poll_max_seconds: Ceiling for the exponential backoff. + job_timeout_seconds: Default wall-clock budget for :meth:`MediaJobManager.wait`. + jitter: Fractional jitter applied to each sleep, to avoid thundering + herds when many jobs are submitted together. + """ + + __slots__ = ("jitter", "job_timeout_seconds", "poll_initial_seconds", "poll_max_seconds") + + + def __init__( + self, + poll_initial_seconds: float = DEFAULT_POLL_INITIAL_SECONDS, + poll_max_seconds: float = DEFAULT_POLL_MAX_SECONDS, + job_timeout_seconds: float = DEFAULT_JOB_TIMEOUT_SECONDS, + jitter: float = 0.1, + ) -> None: + self.poll_initial_seconds = max(0.0, float(poll_initial_seconds)) + self.poll_max_seconds = max(self.poll_initial_seconds, float(poll_max_seconds)) + self.job_timeout_seconds = float(job_timeout_seconds) + self.jitter = max(0.0, min(1.0, float(jitter))) + + #: Exponent ceiling for the backoff doubling. Any realistic + #: ``poll_initial_seconds`` reaches ``poll_max_seconds`` long before 2**32, + #: and without the cap a job polled a few thousand times overflows the float + #: conversion of ``2 ** attempt`` — which a multi-hour video job would hit. + _MAX_BACKOFF_SHIFT = 32 + + def delay_for(self, attempt: int) -> float: + """Return the sleep before poll *attempt* (1-based), with jitter.""" + shift = min(max(0, attempt - 1), self._MAX_BACKOFF_SHIFT) + base = min(self.poll_initial_seconds * (2**shift), self.poll_max_seconds) + if not self.jitter: + return base + return base * (1.0 + random.uniform(-self.jitter, self.jitter)) + + @classmethod + def from_config(cls, get: Callable[[str, Any], Any]) -> JobPolicy: + """Build a policy from an llmcore config accessor. + + Args: + get: A ``config.get(key, default)``-style callable. + + Returns: + The configured policy. + """ + return cls( + poll_initial_seconds=get("media.jobs.poll_initial_seconds", DEFAULT_POLL_INITIAL_SECONDS), + poll_max_seconds=get("media.jobs.poll_max_seconds", DEFAULT_POLL_MAX_SECONDS), + job_timeout_seconds=get("media.jobs.job_timeout_seconds", DEFAULT_JOB_TIMEOUT_SECONDS), + ) + + +class MediaJobManager: + """Tracks and drives long-running media jobs. + + The manager keeps every job it has seen in an in-memory registry so callers + can enumerate outstanding work, and resolves a job back to its provider + adapter through the resolver supplied at construction (normally + :class:`~llmcore.media.manager.MediaManager`'s adapter lookup). + + Args: + resolver: Maps a provider instance name to its media adapter. + policy: Backoff/timeout policy; defaults are used when omitted. + """ + + def __init__( + self, + resolver: Callable[[str], Any], + policy: JobPolicy | None = None, + ) -> None: + self._resolver = resolver + self._policy = policy or JobPolicy() + self._jobs: dict[str, MediaJob] = {} + + # --- registry --- + + def track(self, job: MediaJob) -> MediaJob: + """Record *job* in the registry and return it. + + Called by routers immediately after an adapter returns a job handle, so + an expensive submission is never lost to a dropped reference. + """ + self._jobs[job.id] = job + logger.debug( + "Tracking media job %s (%s on %s/%s)", + job.id, + job.capability, + job.provider, + job.model, + ) + return job + + def get(self, job_id: str) -> MediaJob | None: + """Return the tracked job with *job_id*, if any.""" + return self._jobs.get(job_id) + + def list(self, *, active_only: bool = False) -> list[MediaJob]: + """Return tracked jobs, newest first. + + Args: + active_only: Exclude jobs that have reached a terminal state. + """ + jobs = sorted(self._jobs.values(), key=lambda j: j.created_at, reverse=True) + return [j for j in jobs if not j.is_terminal] if active_only else jobs + + def forget(self, job_id: str) -> None: + """Drop *job_id* from the registry.""" + self._jobs.pop(job_id, None) + + # --- driving --- + + def _poller_for(self, job: MediaJob) -> MediaJobPoller: + """Resolve the adapter that can poll *job*. + + Raises: + MediaJobError: If the provider is gone or cannot poll. + """ + adapter = self._resolver(job.provider) + if adapter is None: + raise MediaJobError( + "Provider is no longer configured, so the job cannot be polled.", + job_id=job.id, + status=job.status, + provider_name=job.provider, + capability=job.capability, + ) + if not isinstance(adapter, MediaJobPoller): + raise MediaJobError( + "Provider does not implement media job polling.", + job_id=job.id, + status=job.status, + provider_name=job.provider, + capability=job.capability, + ) + return adapter + + async def poll(self, job: MediaJob) -> MediaJob: + """Refresh *job* once and return the updated handle. + + Terminal jobs are returned unchanged without contacting the vendor. + """ + if job.is_terminal: + return job + updated = await self._poller_for(job).poll_media_job(job) + updated.touch() + self._jobs[updated.id] = updated + return updated + + async def wait( + self, + job: MediaJob, + *, + timeout: float | None = None, + raise_on_failure: bool = True, + ) -> MediaJob: + """Poll *job* until it reaches a terminal state. + + Args: + job: The handle returned by a router. + timeout: Wall-clock budget in seconds; the configured default when + omitted. ``None`` in config means wait indefinitely. + raise_on_failure: Raise :class:`MediaJobError` when the job ends + ``FAILED``/``CANCELED``/``EXPIRED`` instead of returning it. + + Returns: + The terminal job handle. + + Raises: + MediaJobTimeoutError: If the budget elapses first. **The job is not + cancelled** — the handle stays valid and can be waited on again, + so an expensive generation is never discarded over a client-side + deadline. This is why the deadline is an explicit ``timeout`` + parameter rather than ``asyncio.timeout``: cancelling the task + would be exactly the wrong behaviour here. + MediaJobError: On terminal failure when *raise_on_failure*. + """ + budget = self._policy.job_timeout_seconds if timeout is None else timeout + started = time.monotonic() + attempt = 0 + + while not job.is_terminal: + attempt += 1 + delay = self._policy.delay_for(attempt) + if budget is not None and budget >= 0: + remaining = budget - (time.monotonic() - started) + if remaining <= 0: + raise MediaJobTimeoutError( + job_id=job.id, + status=job.status, + timeout_seconds=budget, + provider_name=job.provider, + ) + delay = min(delay, remaining) + await asyncio.sleep(delay) + job = await self.poll(job) + logger.debug( + "Polled media job %s: status=%s progress=%s", job.id, job.status, job.progress + ) + + if raise_on_failure and job.status is not MediaJobStatus.SUCCEEDED: + raise MediaJobError( + job.error or "Job did not succeed.", + job_id=job.id, + status=job.status, + provider_name=job.provider, + capability=job.capability, + ) + return job + + async def cancel(self, job: MediaJob) -> MediaJob: + """Ask the vendor to cancel *job*. + + Terminal jobs are returned unchanged. + """ + if job.is_terminal: + return job + updated = await self._poller_for(job).cancel_media_job(job) + updated.touch() + self._jobs[updated.id] = updated + return updated + + async def cancel_all(self, jobs: Iterable[MediaJob] | None = None) -> list[MediaJob]: + """Cancel *jobs* (default: every active tracked job), best effort. + + Failures are logged, not raised: shutdown should not be blocked by a + vendor that will not answer. + """ + targets = list(jobs) if jobs is not None else self.list(active_only=True) + results: list[MediaJob] = [] + for job in targets: + try: + results.append(await self.cancel(job)) + except Exception as e: + logger.warning("Failed to cancel media job %s: %s", job.id, e) + return results diff --git a/src/llmcore/media/manager.py b/src/llmcore/media/manager.py new file mode 100644 index 00000000..4f7c44c4 --- /dev/null +++ b/src/llmcore/media/manager.py @@ -0,0 +1,387 @@ +# src/llmcore/media/manager.py +"""The media subsystem entry point. + +:class:`MediaManager` is the sibling of +:class:`~llmcore.providers.manager.ProviderManager` and +:class:`~llmcore.search.manager.SearchProviderManager`: it discovers media +adapters, resolves capabilities to providers, and owns the job manager and +artifact store. + +**Adapters are the chat providers.** Rather than a parallel +``[media.providers.*]`` credential tree, the manager scans the already-loaded +``[providers.*]`` instances and treats any that implements +:class:`~llmcore.media.protocols.MediaCapableProvider` as a media adapter. One +credential per vendor, one place to configure it, and the capability matrix +already tolerates providers that cannot chat (``deepgram``, ``typesafe``). +``[media.routing]`` then expresses preference order per capability. +""" + +from __future__ import annotations + +import logging +from collections.abc import Callable +from typing import Any + +from ..exceptions import MediaCapabilityError +from .artifacts import ArtifactStore, MaterializePolicy, default_fetcher +from .jobs import JobPolicy, MediaJobManager +from .models import MediaCapability, MediaJob, MediaResult +from .protocols import CAPABILITY_PROTOCOLS, MediaCapableProvider +from .routers import AudioRouter, ImageRouter, VideoRouter + +logger = logging.getLogger(__name__) + +__all__ = ["MediaManager"] + +#: Built-in preference order per capability, used when ``[media.routing]`` is +#: silent. Sourced from the workload recommendations in the provider survey +#: (see ``docs/MEDIA_SUBSYSTEM_SPEC.md`` §2.7); a provider that is not +#: configured is skipped, so these are hints rather than requirements. +DEFAULT_ROUTING: dict[MediaCapability, tuple[str, ...]] = { + MediaCapability.TTS: ("elevenlabs", "openai", "deepgram", "zai", "mistral", "deepinfra"), + MediaCapability.TTS_STREAM: ("elevenlabs", "deepgram", "openai"), + MediaCapability.ASR: ("deepgram", "openai", "elevenlabs", "mistral", "zai", "friendli"), + MediaCapability.ASR_STREAM: ("deepgram", "openai", "elevenlabs"), + MediaCapability.VOICE_AGENT: ("deepgram", "openai"), + MediaCapability.MUSIC: ("elevenlabs", "fal", "replicate"), + MediaCapability.SFX: ("elevenlabs", "fal", "replicate"), + MediaCapability.IMAGE_GENERATE: ("openai", "gemini", "fal", "zai", "replicate", "deepinfra"), + MediaCapability.IMAGE_EDIT: ("openai", "gemini", "bfl", "fal"), + MediaCapability.IMAGE_UPSCALE: ("fal", "replicate", "luma"), + MediaCapability.OCR: ("mistral", "zai", "gemini"), + MediaCapability.VIDEO_GENERATE: ("gemini", "fal", "replicate", "luma", "zai"), + MediaCapability.VIDEO_EDIT: ("luma", "fal", "replicate"), + MediaCapability.VIDEO_INTERPOLATE: ("fal", "replicate"), +} + + +class MediaManager: + """Discovers media adapters and routes capability requests to them. + + Args: + adapters: Mapping of provider instance name to media adapter. + routing: Per-capability preference order, overriding :data:`DEFAULT_ROUTING`. + artifact_store: Store for produced assets; a default is built when omitted. + job_policy: Polling/timeout policy for long-running jobs. + """ + + def __init__( + self, + adapters: dict[str, MediaCapableProvider] | None = None, + *, + routing: dict[MediaCapability, tuple[str, ...]] | None = None, + artifact_store: ArtifactStore | None = None, + job_policy: JobPolicy | None = None, + ) -> None: + self._adapters: dict[str, MediaCapableProvider] = dict(adapters or {}) + self._routing: dict[MediaCapability, tuple[str, ...]] = { + **DEFAULT_ROUTING, + **(routing or {}), + } + self.artifacts = artifact_store or ArtifactStore() + self.jobs = MediaJobManager(self._adapters.get, policy=job_policy) + + self.images = ImageRouter(self) + self.audio = AudioRouter(self) + self.video = VideoRouter(self) + + # ------------------------------------------------------------------ + # Construction from config + # ------------------------------------------------------------------ + + @classmethod + def from_provider_manager( + cls, + provider_manager: Any, + config_get: Callable[[str, Any], Any] | None = None, + ) -> MediaManager: + """Build a manager from the loaded chat providers. + + Any provider instance implementing + :class:`~llmcore.media.protocols.MediaCapableProvider` becomes a media + adapter. Providers that only implement the legacy + :class:`~llmcore.providers.base.BaseProvider` media methods are *not* + adapters — they keep working through their own methods, and gain routing + when they are migrated. + + Args: + provider_manager: The loaded :class:`ProviderManager`. + config_get: A ``config.get(key, default)`` accessor; defaults are + used when omitted. + + Returns: + A configured manager. Never raises on an absent ``[media]`` section. + """ + get = config_get or (lambda _k, d=None: d) + + adapters: dict[str, MediaCapableProvider] = {} + for name in provider_manager.get_available_providers(): + try: + provider = provider_manager.get_provider(name) + except Exception as e: + logger.debug("Skipping provider '%s' during media discovery: %s", name, e) + continue + if isinstance(provider, MediaCapableProvider): + adapters[name] = provider + + routing = cls._routing_from_config(get) + + store = ArtifactStore( + get("media.artifact_path", "~/.llmcore/media"), + policy=_coerce_policy(get("media.artifact_materialize", MaterializePolicy.ON_EXPIRY)), + fetcher=_safe_default_fetcher(), + ) + + manager = cls( + adapters, + routing=routing, + artifact_store=store, + job_policy=JobPolicy.from_config(get), + ) + logger.debug( + "MediaManager initialized with %d adapter(s): %s", + len(adapters), + ", ".join(sorted(adapters)) or "none", + ) + return manager + + @staticmethod + def _routing_from_config( + get: Callable[[str, Any], Any], + ) -> dict[MediaCapability, tuple[str, ...]]: + """Read ``[media.routing]`` into a capability→preference map.""" + section = get("media.routing", {}) or {} + if not isinstance(section, dict): + logger.warning("[media.routing] is not a table; ignoring it.") + return {} + routing: dict[MediaCapability, tuple[str, ...]] = {} + for key, value in section.items(): + try: + capability = MediaCapability(str(key).lower()) + except ValueError: + logger.warning("Unknown media capability '%s' in [media.routing]; ignoring.", key) + continue + if isinstance(value, str): + value = [value] + if not isinstance(value, (list, tuple)): + logger.warning("[media.routing].%s must be a list of provider names.", key) + continue + routing[capability] = tuple(str(v).lower() for v in value) + return routing + + # ------------------------------------------------------------------ + # Adapter registry + # ------------------------------------------------------------------ + + def register_adapter(self, name: str, adapter: MediaCapableProvider) -> None: + """Register (or replace) a media adapter under *name*.""" + self._adapters[name.lower()] = adapter + logger.debug("Registered media adapter '%s'.", name) + + def unregister_adapter(self, name: str) -> None: + """Remove the adapter registered under *name*, if present.""" + if self._adapters.pop(name.lower(), None) is not None: + logger.debug("Unregistered media adapter '%s'.", name) + + @property + def adapter_names(self) -> list[str]: + """Names of every registered media adapter.""" + return sorted(self._adapters) + + def has_adapters(self) -> bool: + """Whether any media adapter is available.""" + return bool(self._adapters) + + # ------------------------------------------------------------------ + # Discovery + # ------------------------------------------------------------------ + + def capabilities(self) -> dict[MediaCapability, list[str]]: + """Return every available capability mapped to the providers offering it.""" + out: dict[MediaCapability, list[str]] = {} + for name, adapter in self._adapters.items(): + for capability in self._capabilities_of(adapter): + out.setdefault(capability, []).append(name) + return {c: sorted(v) for c, v in sorted(out.items())} + + def who_can(self, capability: MediaCapability | str) -> list[str]: + """Return providers that can serve *capability*, in preference order.""" + cap = MediaCapability(capability) + available = { + name + for name, adapter in self._adapters.items() + if cap in self._capabilities_of(adapter) + } + preferred = [p for p in self._routing.get(cap, ()) if p in available] + rest = sorted(available - set(preferred)) + return preferred + rest + + def _capabilities_of(self, adapter: MediaCapableProvider) -> frozenset[MediaCapability]: + """Capabilities *adapter* declares, filtered by protocol conformance. + + A declared capability whose protocol the adapter does not actually + implement is dropped with a warning: better a missing capability than a + confident ``AttributeError`` at call time. + """ + try: + declared = adapter.media_capabilities() + except Exception as e: + logger.warning( + "Adapter '%s' failed to report capabilities: %s", _name_of(adapter), e + ) + return frozenset() + + usable: set[MediaCapability] = set() + for capability in declared: + protocol = CAPABILITY_PROTOCOLS.get(capability) + if protocol is None or isinstance(adapter, protocol): + usable.add(capability) + else: + logger.warning( + "Adapter '%s' declares %s but does not implement %s; skipping it.", + _name_of(adapter), + capability, + protocol.__name__, + ) + return frozenset(usable) + + # ------------------------------------------------------------------ + # Resolution + # ------------------------------------------------------------------ + + def resolve( + self, + capability: MediaCapability | str, + *, + provider: str | None = None, + model: str | None = None, + ) -> MediaCapableProvider: + """Choose the adapter that will serve *capability*. + + Resolution order (spec §2.7): + + 1. Explicit ``provider`` — used, or an error if it lacks the capability. + 2. Explicit ``model`` — the owning provider, resolved from model cards. + 3. ``[media.routing]`` preference for the capability. + 4. Built-in :data:`DEFAULT_ROUTING` preference. + 5. Any remaining adapter that can do it. + + Raises: + MediaCapabilityError: When nothing can serve it, naming the + providers that could if they were configured. + """ + cap = MediaCapability(capability) + + if provider: + adapter = self._adapters.get(provider.lower()) + if adapter is None: + raise MediaCapabilityError( + f"Media provider '{provider}' is not configured.", + capability=cap, + candidates=self.adapter_names, + ) + if cap not in self._capabilities_of(adapter): + raise MediaCapabilityError( + f"Media provider '{provider}' does not support this capability.", + capability=cap, + candidates=self.who_can(cap), + ) + return adapter + + if model: + owner = self._provider_for_model(model, cap) + if owner is not None: + return self._adapters[owner] + + candidates = self.who_can(cap) + if not candidates: + raise MediaCapabilityError( + "No configured media provider offers this capability.", + capability=cap, + candidates=list(self._routing.get(cap, ())), + ) + return self._adapters[candidates[0]] + + def _provider_for_model(self, model: str, capability: MediaCapability) -> str | None: + """Resolve *model* to a configured adapter via the model-card registry.""" + try: + from ..model_cards.registry import get_model_card_registry + + registry = get_model_card_registry() + except Exception: + return None + + for name, adapter in self._adapters.items(): + if capability not in self._capabilities_of(adapter): + continue + try: + if registry.get(name, model) is not None: + return name + except Exception: + continue + return None + + # ------------------------------------------------------------------ + # Post-processing + # ------------------------------------------------------------------ + + def _finalize(self, outcome: MediaResult | MediaJob) -> MediaResult | MediaJob: + """Track jobs so a submitted job is never lost to a dropped reference.""" + if isinstance(outcome, MediaJob): + return self.jobs.track(outcome) + return outcome + + async def wait( + self, job: MediaJob, *, timeout: float | None = None + ) -> MediaResult: + """Wait for *job* and return its result. + + Convenience over ``manager.jobs.wait(...)`` plus ``job.to_result()``, + which is what almost every caller wants. + """ + finished = await self.jobs.wait(job, timeout=timeout) + return finished.to_result() + + async def close(self) -> None: + """Release media resources. Adapters are owned by ``ProviderManager``. + + Active jobs are deliberately **not** cancelled: a long video generation + is already paid for, and killing it on shutdown would waste it. Callers + that want cancellation call ``manager.jobs.cancel_all()`` explicitly. + """ + active = self.jobs.list(active_only=True) + if active: + logger.info( + "MediaManager closing with %d active job(s) left running: %s", + len(active), + ", ".join(j.id for j in active), + ) + + +def _name_of(adapter: Any) -> str: + """Best-effort adapter name for log messages.""" + try: + return adapter.get_name() + except Exception: + return type(adapter).__name__ + + +def _coerce_policy(value: Any) -> MaterializePolicy: + """Coerce a config value to a materialize policy, defaulting safely.""" + try: + return MaterializePolicy(str(value)) + except ValueError: + logger.warning("Unknown media.artifact_materialize '%s'; using 'on_expiry'.", value) + return MaterializePolicy.ON_EXPIRY + + +def _safe_default_fetcher() -> Any: + """Return an httpx fetcher when available, else ``None``. + + Materialization then raises a clear error only if it is actually attempted, + rather than making httpx a hard requirement of the subsystem. + """ + try: + return default_fetcher() + except Exception: + return None diff --git a/src/llmcore/media/models.py b/src/llmcore/media/models.py new file mode 100644 index 00000000..25e042fb --- /dev/null +++ b/src/llmcore/media/models.py @@ -0,0 +1,612 @@ +# src/llmcore/media/models.py +"""Provider-agnostic types for the LLMCore media subsystem. + +These are the contract between callers and media provider adapters, for +generative image, audio and video work. They are deliberately decoupled from +any vendor SDK so that switching between OpenAI, Google, fal, ElevenLabs, +Replicate or a Hugging Face endpoint does not change how results are read — +mirroring how :mod:`llmcore.search.models` normalizes search results and +:class:`~llmcore.providers.base.BaseProvider` normalizes chat responses. + +Design notes +------------ +* **Three execution classes, not one.** Image generation is usually + request/response, speech is often a byte stream, and video is almost always a + long-running job. :class:`MediaResult` covers the first, + ``AsyncIterator[bytes]`` the second, and :class:`MediaJob` the third. Forcing + all three into one shape is the main modelling mistake to avoid. +* **Inputs are refs, not bytes.** :class:`MediaRef` lets a caller pass a URL, a + path, raw bytes, or a previously produced :class:`MediaArtifact`; the adapter + decides whether its API wants an upload, a URL or inline base64. Callers never + hand-roll base64. +* **Artifacts know when they die.** Every aggregator returns short-lived URLs, + so :attr:`MediaArtifact.expires_at` and :attr:`MediaArtifact.checksum_sha256` + are first-class: they let :class:`~llmcore.media.artifacts.ArtifactStore` + decide what to materialize, and let callers dedupe. +* **Usage keeps native units.** Media vendors bill per image, per megapixel, + per second of video, per audio-minute, per character or per compute-second. + :class:`MediaUsage` records whichever units the vendor reported plus an + explicitly-stamped cost estimate, rather than inventing a token count. +* ``raw`` / ``provider_metadata`` always preserve the vendor payload so power + users can reach fields the normalizer does not surface. + +See ``docs/MEDIA_SUBSYSTEM_SPEC.md`` for the full design. +""" + +from __future__ import annotations + +import base64 +import hashlib +import mimetypes +import uuid +from collections.abc import Mapping, Sequence +from dataclasses import asdict, dataclass, field +from datetime import UTC, datetime +from enum import StrEnum +from pathlib import Path +from typing import Any + +__all__ = [ + "AUDIO_CAPABILITIES", + "CAPABILITY_KINDS", + "IMAGE_CAPABILITIES", + "TERMINAL_JOB_STATUSES", + "VIDEO_CAPABILITIES", + "MediaArtifact", + "MediaCapability", + "MediaExecution", + "MediaJob", + "MediaJobStatus", + "MediaKind", + "MediaProvenance", + "MediaRef", + "MediaResult", + "MediaUsage", +] + + +# --------------------------------------------------------------------------- +# Enumerations +# --------------------------------------------------------------------------- + + +class MediaKind(StrEnum): + """The modality of a media input or output.""" + + AUDIO = "audio" + IMAGE = "image" + VIDEO = "video" + TEXT = "text" + + +class MediaCapability(StrEnum): + """A specific operation a media provider can perform. + + Capability is the unit of routing and of model-card declaration: callers ask + for a capability, and the router resolves it to a provider that implements + the matching protocol. Provider names never appear in routing logic. + """ + + # --- audio --- + TTS = "tts" + TTS_STREAM = "tts_stream" + ASR = "asr" + ASR_STREAM = "asr_stream" + VOICE_AGENT = "voice_agent" + MUSIC = "music" + SFX = "sfx" + VOICE_DESIGN = "voice_design" + + # --- image --- + IMAGE_GENERATE = "image_generate" + IMAGE_EDIT = "image_edit" + IMAGE_UPSCALE = "image_upscale" + IMAGE_VARIATE = "image_variate" + OCR = "ocr" + + # --- video --- + VIDEO_GENERATE = "video_generate" + VIDEO_EDIT = "video_edit" + VIDEO_INTERPOLATE = "video_interpolate" + VIDEO_REFRAME = "video_reframe" + VIDEO_UPSCALE = "video_upscale" + VIDEO_EXTEND = "video_extend" + + +class MediaExecution(StrEnum): + """How an operation completes. + + Declared per model in its card so callers can tell, before submitting, + whether to expect a result, a stream or a job handle. + """ + + REQUEST_RESPONSE = "request_response" + STREAM = "stream" + ASYNC_JOB = "async_job" + + +class MediaJobStatus(StrEnum): + """Lifecycle state of a long-running media job.""" + + QUEUED = "queued" + RUNNING = "running" + SUCCEEDED = "succeeded" + FAILED = "failed" + CANCELED = "canceled" + EXPIRED = "expired" + + +#: Capabilities grouped by the router that owns them. +AUDIO_CAPABILITIES: frozenset[MediaCapability] = frozenset( + { + MediaCapability.TTS, + MediaCapability.TTS_STREAM, + MediaCapability.ASR, + MediaCapability.ASR_STREAM, + MediaCapability.VOICE_AGENT, + MediaCapability.MUSIC, + MediaCapability.SFX, + MediaCapability.VOICE_DESIGN, + } +) + +IMAGE_CAPABILITIES: frozenset[MediaCapability] = frozenset( + { + MediaCapability.IMAGE_GENERATE, + MediaCapability.IMAGE_EDIT, + MediaCapability.IMAGE_UPSCALE, + MediaCapability.IMAGE_VARIATE, + MediaCapability.OCR, + } +) + +VIDEO_CAPABILITIES: frozenset[MediaCapability] = frozenset( + { + MediaCapability.VIDEO_GENERATE, + MediaCapability.VIDEO_EDIT, + MediaCapability.VIDEO_INTERPOLATE, + MediaCapability.VIDEO_REFRAME, + MediaCapability.VIDEO_UPSCALE, + MediaCapability.VIDEO_EXTEND, + } +) + +#: The modality each capability produces. OCR is the one audio/image capability +#: whose *output* modality differs from its router's modality. +CAPABILITY_KINDS: dict[MediaCapability, MediaKind] = { + **dict.fromkeys(AUDIO_CAPABILITIES, MediaKind.AUDIO), + **dict.fromkeys(IMAGE_CAPABILITIES, MediaKind.IMAGE), + **dict.fromkeys(VIDEO_CAPABILITIES, MediaKind.VIDEO), + MediaCapability.ASR: MediaKind.TEXT, + MediaCapability.ASR_STREAM: MediaKind.TEXT, + MediaCapability.OCR: MediaKind.TEXT, +} + +#: Job statuses from which no further transition occurs. +TERMINAL_JOB_STATUSES: frozenset[MediaJobStatus] = frozenset( + { + MediaJobStatus.SUCCEEDED, + MediaJobStatus.FAILED, + MediaJobStatus.CANCELED, + MediaJobStatus.EXPIRED, + } +) + + +def _utcnow() -> datetime: + """Return an aware UTC timestamp.""" + return datetime.now(UTC) + + +def _serialize(value: Any) -> Any: + """Recursively make a value JSON-compatible. + + ``datetime`` becomes ISO-8601, ``bytes`` becomes a byte count marker rather + than inline base64 (media payloads are far too large to serialize by + accident), and enums become their values. + + Args: + value: Any value that may contain nested datetimes, bytes or enums. + + Returns: + A structurally identical value that ``json.dumps`` accepts. + """ + if isinstance(value, datetime): + return value.isoformat() + if isinstance(value, StrEnum): + return value.value + if isinstance(value, bytes): + return f"<{len(value)} bytes>" + if isinstance(value, Path): + return str(value) + if isinstance(value, dict): + return {k: _serialize(v) for k, v in value.items()} + if isinstance(value, (list, tuple, set, frozenset)): + return [_serialize(v) for v in value] + return value + + +# --------------------------------------------------------------------------- +# Inputs +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True, slots=True) +class MediaRef: + """A reference to media supplied *to* a provider. + + Exactly one of :attr:`url`, :attr:`path` or :attr:`data` is set. Adapters + call :meth:`read_bytes` or :meth:`as_data_uri` when their API needs inline + content, and use :attr:`url` directly when it accepts a remote reference — + so callers never encode base64 themselves. + + Attributes: + url: A remote URL the provider can fetch. + path: A local filesystem path. + data: Raw bytes. + mime_type: Media type; inferred from ``path``/``url`` when omitted. + filename: Preferred upload filename. + """ + + url: str | None = None + path: Path | None = None + data: bytes | None = None + mime_type: str | None = None + filename: str | None = None + + def __post_init__(self) -> None: + provided = [x is not None for x in (self.url, self.path, self.data)] + if sum(provided) != 1: + raise ValueError("MediaRef requires exactly one of url, path or data.") + + # --- constructors --- + + @classmethod + def from_url(cls, url: str, *, mime_type: str | None = None) -> MediaRef: + """Build a ref from a remote URL.""" + return cls(url=url, mime_type=mime_type or _guess_mime(url)) + + @classmethod + def from_path(cls, path: str | Path, *, mime_type: str | None = None) -> MediaRef: + """Build a ref from a local file path.""" + p = Path(path) + return cls( + path=p, + mime_type=mime_type or _guess_mime(p.name), + filename=p.name, + ) + + @classmethod + def from_bytes( + cls, data: bytes, *, mime_type: str | None = None, filename: str | None = None + ) -> MediaRef: + """Build a ref from raw bytes.""" + return cls(data=data, mime_type=mime_type, filename=filename) + + @classmethod + def from_artifact(cls, artifact: MediaArtifact) -> MediaRef: + """Chain a previously produced artifact back in as an input. + + Prefers inline bytes when the artifact carries them, otherwise its URI. + + Raises: + ValueError: If the artifact has neither ``data`` nor ``uri``. + """ + if artifact.data is not None: + return cls(data=artifact.data, mime_type=artifact.mime_type) + if artifact.uri: + return cls(url=artifact.uri, mime_type=artifact.mime_type) + raise ValueError("MediaArtifact carries neither data nor uri; cannot build a ref.") + + # --- accessors --- + + def read_bytes(self) -> bytes: + """Return the referenced content as bytes. + + Returns: + The raw content. + + Raises: + ValueError: If this ref is a URL (the adapter must fetch it, since + only the adapter knows which client and auth to use). + """ + if self.data is not None: + return self.data + if self.path is not None: + return self.path.read_bytes() + raise ValueError( + "MediaRef holds a URL; the provider adapter must fetch it with its own client." + ) + + def as_data_uri(self) -> str: + """Return the content as a ``data:`` URI for inline APIs.""" + payload = base64.b64encode(self.read_bytes()).decode("ascii") + return f"data:{self.mime_type or 'application/octet-stream'};base64,{payload}" + + @property + def is_remote(self) -> bool: + """Whether this ref points at a URL the provider must fetch.""" + return self.url is not None + + +def _guess_mime(name: str) -> str | None: + """Best-effort media type from a filename or URL.""" + return mimetypes.guess_type(name)[0] + + +# --------------------------------------------------------------------------- +# Outputs +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True, slots=True) +class MediaProvenance: + """Content-credential / provenance metadata attached to generated media. + + Several vendors now emit C2PA-style credentials or SynthID-style watermarks. + Modelled from day one so it is not retrofitted later. + + Attributes: + watermarked: Whether the provider states the output is watermarked. + c2pa_manifest: Raw C2PA manifest, when supplied. + generator: Model or system credited with producing the asset. + provider_declared: Whether these facts come from the provider (``True``) + or were inferred locally (``False``). + """ + + watermarked: bool | None = None + c2pa_manifest: Mapping[str, Any] | None = None + generator: str | None = None + provider_declared: bool = True + + +@dataclass(frozen=True, slots=True) +class MediaArtifact: + """One produced media asset. + + Attributes: + kind: Modality of this asset. + uri: Provider URL or artifact-store URI. + data: Inline bytes, for small results. + mime_type: Media type of the asset. + width: Pixel width (image/video). + height: Pixel height (image/video). + duration_seconds: Duration (audio/video). + sample_rate_hz: Sample rate (audio). + fps: Frames per second (video). + frame_count: Total frames (video). + text: Transcribed or extracted text (ASR/OCR outputs). + checksum_sha256: Hex digest of the bytes, when known. + expires_at: When ``uri`` stops resolving, when the provider says so. + provenance: Content-credential metadata. + provider_metadata: Untouched vendor fields. + """ + + kind: MediaKind + uri: str | None = None + data: bytes | None = None + mime_type: str | None = None + width: int | None = None + height: int | None = None + duration_seconds: float | None = None + sample_rate_hz: int | None = None + fps: float | None = None + frame_count: int | None = None + text: str | None = None + checksum_sha256: str | None = None + expires_at: datetime | None = None + provenance: MediaProvenance | None = None + provider_metadata: Mapping[str, Any] = field(default_factory=dict) + + @property + def is_expired(self) -> bool: + """Whether :attr:`expires_at` is in the past.""" + return self.expires_at is not None and self.expires_at <= _utcnow() + + @property + def needs_materialization(self) -> bool: + """Whether the bytes should be fetched before the URI dies. + + ``True`` when the asset exists only as a URI that carries an expiry. + """ + return self.data is None and self.uri is not None and self.expires_at is not None + + def with_data(self, data: bytes) -> MediaArtifact: + """Return a copy carrying *data* and its checksum. + + Args: + data: The materialized bytes. + + Returns: + A new artifact with ``data`` and ``checksum_sha256`` populated. + """ + return MediaArtifact( + **{ + **{f: getattr(self, f) for f in self.__slots__}, + "data": data, + "checksum_sha256": hashlib.sha256(data).hexdigest(), + } + ) + + def to_dict(self) -> dict[str, Any]: + """JSON-compatible view; inline bytes are summarized, never embedded.""" + return _serialize(asdict(self)) + + +@dataclass(frozen=True, slots=True) +class MediaUsage: + """Normalized-but-honest billing units for one media operation. + + Media vendors bill in incompatible units, so every native unit the provider + reported is kept alongside an explicitly stamped estimate. Never synthesize + a unit the vendor did not report. + + Attributes: + provider: Provider instance name. + model: Model or endpoint that ran. + basis: Vendor's billing basis (e.g. ``"per_image"``, ``"per_second"``). + images: Number of images produced. + megapixels: Megapixels produced. + seconds: Seconds of output media. + audio_minutes: Minutes of audio processed or produced. + characters: Characters consumed (common for TTS). + input_tokens: Input tokens, when the vendor bills tokens. + output_tokens: Output tokens, when the vendor bills tokens. + compute_seconds: Billed compute time (aggregators). + estimated_cost_usd: Best-effort cost estimate. + pricing_as_of: Date stamp of the pricing used for the estimate. + raw: Untouched vendor usage payload. + """ + + provider: str + model: str + basis: str | None = None + images: int | None = None + megapixels: float | None = None + seconds: float | None = None + audio_minutes: float | None = None + characters: int | None = None + input_tokens: int | None = None + output_tokens: int | None = None + compute_seconds: float | None = None + estimated_cost_usd: float | None = None + pricing_as_of: str | None = None + raw: Mapping[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + """JSON-compatible view.""" + return _serialize(asdict(self)) + + +@dataclass(frozen=True, slots=True) +class MediaResult: + """The outcome of a request/response media operation. + + Attributes: + capability: Which operation produced this. + provider: Provider instance name. + model: Model or endpoint that ran. + artifacts: Produced assets, in provider order. + usage: Billing units, when reported. + raw: Untouched vendor response. + """ + + capability: MediaCapability + provider: str + model: str + artifacts: Sequence[MediaArtifact] = field(default_factory=tuple) + usage: MediaUsage | None = None + raw: Mapping[str, Any] = field(default_factory=dict) + + @property + def artifact(self) -> MediaArtifact: + """The first artifact. + + Raises: + IndexError: If the operation produced nothing. + """ + if not self.artifacts: + raise IndexError( + f"{self.capability} on {self.provider}/{self.model} produced no artifacts." + ) + return self.artifacts[0] + + @property + def text(self) -> str | None: + """Concatenated text across artifacts, for ASR/OCR results.""" + parts = [a.text for a in self.artifacts if a.text] + return "\n".join(parts) if parts else None + + def to_dict(self) -> dict[str, Any]: + """JSON-compatible view.""" + return _serialize(asdict(self)) + + +@dataclass(slots=True) +class MediaJob: + """Handle for a long-running media operation. + + Normalizes fal queue requests, Replicate predictions, Veo operations and + Luma generations without pretending their native APIs are identical: the + vendor's own identifiers and poll targets live in :attr:`provider_job_id` + and :attr:`poll_url`, and the whole payload stays in + :attr:`provider_metadata`. + + Attributes: + id: llmcore-local job id, stable across resumes. + capability: Which operation this job performs. + provider: Provider instance name. + model: Model or endpoint that runs it. + status: Current lifecycle state. + provider_job_id: The vendor's identifier. + poll_url: Vendor status URL, when it supplies one. + artifacts: Produced assets once succeeded. + usage: Billing units, when reported. + error: Failure message when ``status`` is ``FAILED``. + progress: Fractional progress in ``[0, 1]`` when the vendor reports it. + idempotency_key: Client-generated key so a retried submit cannot + double-bill, and a resumed process re-attaches instead of + resubmitting. + created_at: When llmcore submitted the job. + updated_at: Last status refresh. + queue_position: Position in the vendor queue, when reported. + provider_metadata: Untouched vendor payload. + """ + + capability: MediaCapability + provider: str + model: str + status: MediaJobStatus = MediaJobStatus.QUEUED + id: str = field(default_factory=lambda: f"mj_{uuid.uuid4().hex[:16]}") + provider_job_id: str | None = None + poll_url: str | None = None + artifacts: list[MediaArtifact] = field(default_factory=list) + usage: MediaUsage | None = None + error: str | None = None + progress: float | None = None + idempotency_key: str | None = None + created_at: datetime = field(default_factory=_utcnow) + updated_at: datetime = field(default_factory=_utcnow) + queue_position: int | None = None + provider_metadata: dict[str, Any] = field(default_factory=dict) + + @property + def is_terminal(self) -> bool: + """Whether no further status transition will occur.""" + return self.status in TERMINAL_JOB_STATUSES + + @property + def succeeded(self) -> bool: + """Whether the job completed successfully.""" + return self.status is MediaJobStatus.SUCCEEDED + + def to_result(self) -> MediaResult: + """Project a succeeded job into a :class:`MediaResult`. + + Returns: + The equivalent request/response result, so callers can treat both + execution classes uniformly once a job finishes. + + Raises: + ValueError: If the job has not succeeded. + """ + if not self.succeeded: + raise ValueError( + f"Job {self.id} is {self.status}, not succeeded; cannot project to a result." + ) + return MediaResult( + capability=self.capability, + provider=self.provider, + model=self.model, + artifacts=tuple(self.artifacts), + usage=self.usage, + raw=dict(self.provider_metadata), + ) + + def touch(self) -> None: + """Mark the job as just refreshed.""" + self.updated_at = _utcnow() + + def to_dict(self) -> dict[str, Any]: + """JSON-compatible view.""" + return _serialize(asdict(self)) diff --git a/src/llmcore/media/protocols.py b/src/llmcore/media/protocols.py new file mode 100644 index 00000000..8b73e02b --- /dev/null +++ b/src/llmcore/media/protocols.py @@ -0,0 +1,432 @@ +# src/llmcore/media/protocols.py +"""Capability protocols implemented by LLMCore media provider adapters. + +Each protocol describes one coherent group of media operations. An adapter +implements only what its vendor supports, and the routers in +:mod:`llmcore.media.routers` discover that with ``isinstance`` — so routing +logic never names a provider. + +Three method-shape conventions, matching the three execution classes in +``docs/MEDIA_SUBSYSTEM_SPEC.md``: + +* **Request/response** returns :class:`~llmcore.media.models.MediaResult`. +* **Byte stream** returns ``AsyncIterator[bytes]`` (or an async session object). +* **Long-running job** returns :class:`~llmcore.media.models.MediaJob`. + +An operation whose execution class varies by *model* (image generation is +request/response on OpenAI but a queued job on fal) is annotated +``MediaResult | MediaJob``; callers either branch on the returned type or ask +:meth:`MediaCapableProvider.media_execution` up front. + +Method names carry a ``_media`` suffix where they would otherwise collide with +the legacy :class:`~llmcore.providers.base.BaseProvider` media methods +(``generate_image``, ``transcribe_audio``, ``generate_speech``, ``ocr``), which +keep their current signatures and return types for backward compatibility. The +suffix disappears at the next major version. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Sequence +from typing import Any, Protocol, runtime_checkable + +from .models import ( + MediaCapability, + MediaExecution, + MediaJob, + MediaRef, + MediaResult, +) + +__all__ = [ + "CAPABILITY_PROTOCOLS", + "ASRProvider", + "ImageEditProvider", + "ImageGenerationProvider", + "ImageUpscaleProvider", + "MediaCapableProvider", + "MediaJobPoller", + "MusicProvider", + "OCRMediaProvider", + "SFXProvider", + "StreamingASRProvider", + "StreamingTTSProvider", + "TTSProvider", + "VideoEditProvider", + "VideoGenerationProvider", + "VideoInterpolationProvider", +] + + +# --------------------------------------------------------------------------- +# Base +# --------------------------------------------------------------------------- + + +@runtime_checkable +class MediaCapableProvider(Protocol): + """Minimum surface every media adapter exposes. + + Implemented by every adapter regardless of which capability protocols it + also satisfies, so the manager can enumerate and describe providers without + knowing what they can do. + """ + + def get_name(self) -> str: + """Return the provider instance name.""" + ... + + def media_capabilities(self) -> frozenset[MediaCapability]: + """Return every capability this adapter can currently serve. + + Should reflect configuration (credentials present, optional dependency + installed), not just what the vendor theoretically offers. + """ + ... + + def media_execution( + self, capability: MediaCapability, model: str | None = None + ) -> MediaExecution: + """Return how *capability* completes for *model* on this provider. + + Lets a caller know before submitting whether to expect a result, a + stream or a job handle. + """ + ... + + +@runtime_checkable +class MediaJobPoller(Protocol): + """Implemented by adapters that return :class:`MediaJob` handles. + + :class:`~llmcore.media.jobs.MediaJobManager` drives these; adapters never + implement their own polling loop, backoff or timeout policy. + """ + + async def poll_media_job(self, job: MediaJob) -> MediaJob: + """Refresh *job* against the vendor and return the updated handle. + + Must be safe to call on a terminal job (returning it unchanged) so the + manager does not need to pre-check. + """ + ... + + async def cancel_media_job(self, job: MediaJob) -> MediaJob: + """Ask the vendor to cancel *job* and return the updated handle.""" + ... + + +# --------------------------------------------------------------------------- +# Image +# --------------------------------------------------------------------------- + + +@runtime_checkable +class ImageGenerationProvider(Protocol): + """Text-to-image generation.""" + + async def generate_image_media( + self, + prompt: str, + *, + model: str | None = None, + n: int = 1, + size: str | None = None, + seed: int | None = None, + negative_prompt: str | None = None, + reference_images: Sequence[MediaRef] | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Generate *n* images from *prompt*. + + Args: + prompt: Text description of the desired image. + model: Model id; the provider default when omitted. + n: Number of images to produce. + size: Vendor-specific size token (e.g. ``"1024x1024"``). + seed: Sampling seed for reproducibility, where supported. + negative_prompt: What to avoid, where supported. + reference_images: Style or subject references, where supported. + **kwargs: Vendor-specific parameters, passed through. + + Returns: + A result, or a job handle for queue-based vendors. + """ + ... + + +@runtime_checkable +class ImageEditProvider(Protocol): + """Instruction-guided editing of an existing image.""" + + async def edit_image_media( + self, + prompt: str, + *, + image: MediaRef, + mask: MediaRef | None = None, + model: str | None = None, + n: int = 1, + size: str | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Edit *image* according to *prompt*, optionally restricted by *mask*.""" + ... + + +@runtime_checkable +class ImageUpscaleProvider(Protocol): + """Resolution enhancement of an existing image.""" + + async def upscale_image_media( + self, + *, + image: MediaRef, + model: str | None = None, + scale: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Upscale *image* by *scale*, or to the model's native target.""" + ... + + +@runtime_checkable +class OCRMediaProvider(Protocol): + """Document/image text extraction returning normalized artifacts.""" + + async def ocr_media( + self, + *, + document: MediaRef, + model: str | None = None, + pages: Sequence[int] | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Extract text (and optionally layout) from *document*.""" + ... + + +# --------------------------------------------------------------------------- +# Audio +# --------------------------------------------------------------------------- + + +@runtime_checkable +class TTSProvider(Protocol): + """One-shot text-to-speech.""" + + async def synthesize_speech_media( + self, + text: str, + *, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + speed: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Synthesize *text* into a single audio artifact.""" + ... + + +@runtime_checkable +class StreamingTTSProvider(Protocol): + """Chunked text-to-speech for low-latency playback. + + Deliberately *not* a job: the value is receiving the first bytes before + synthesis finishes, which a job handle cannot express. + """ + + def stream_speech_media( + self, + text: str, + *, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> AsyncIterator[bytes]: + """Yield audio chunks as they are synthesized. + + Note this is a plain method returning an async iterator, not a coroutine + — callers write ``async for chunk in provider.stream_speech_media(...)``. + """ + ... + + +@runtime_checkable +class ASRProvider(Protocol): + """Batch speech-to-text.""" + + async def transcribe_media( + self, + *, + audio: MediaRef, + model: str | None = None, + language: str | None = None, + diarize: bool | None = None, + timestamps: bool | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Transcribe *audio*; text lands on the artifact's ``text`` field.""" + ... + + +@runtime_checkable +class StreamingASRProvider(Protocol): + """Realtime speech-to-text over a bidirectional session. + + Returns an opaque session object rather than an iterator because realtime + ASR is duplex: the caller pushes audio *and* consumes events. Deepgram's + existing sockets are the reference shape. + """ + + async def open_transcription_session( + self, + *, + model: str | None = None, + language: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> Any: + """Open a realtime transcription session.""" + ... + + +@runtime_checkable +class MusicProvider(Protocol): + """Text-to-music generation.""" + + async def generate_music_media( + self, + prompt: str, + *, + model: str | None = None, + duration_seconds: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Generate music from *prompt*.""" + ... + + +@runtime_checkable +class SFXProvider(Protocol): + """Sound-effect generation, optionally conditioned on a video.""" + + async def generate_sfx_media( + self, + prompt: str | None = None, + *, + model: str | None = None, + video: MediaRef | None = None, + duration_seconds: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Generate sound effects from *prompt* and/or *video* (foley).""" + ... + + +# --------------------------------------------------------------------------- +# Video +# --------------------------------------------------------------------------- + + +@runtime_checkable +class VideoGenerationProvider(Protocol): + """Text/image-conditioned video generation. + + Almost always a long-running job, hence the :class:`MediaJob` return. + """ + + async def generate_video_media( + self, + prompt: str, + *, + model: str | None = None, + first_frame: MediaRef | None = None, + last_frame: MediaRef | None = None, + reference_images: Sequence[MediaRef] | None = None, + duration_seconds: float | None = None, + resolution: str | None = None, + aspect_ratio: str | None = None, + fps: float | None = None, + with_audio: bool | None = None, + seed: int | None = None, + **kwargs: Any, + ) -> MediaJob: + """Generate a video. + + ``first_frame`` / ``last_frame`` express generative transitions, which + are semantically distinct from frame interpolation — see + :class:`VideoInterpolationProvider`. + """ + ... + + +@runtime_checkable +class VideoEditProvider(Protocol): + """Instruction-guided modification of an existing video.""" + + async def edit_video_media( + self, + prompt: str, + *, + video: MediaRef, + model: str | None = None, + **kwargs: Any, + ) -> MediaJob: + """Edit *video* according to *prompt*.""" + ... + + +@runtime_checkable +class VideoInterpolationProvider(Protocol): + """Frame interpolation / FPS increase. + + Distinct from a generative first/last-frame transition: interpolation fills + between *existing* frames rather than inventing new content. + """ + + async def interpolate_video_media( + self, + *, + video: MediaRef | None = None, + frames: Sequence[MediaRef] | None = None, + model: str | None = None, + target_fps: float | None = None, + **kwargs: Any, + ) -> MediaJob: + """Interpolate *video*, or between the supplied *frames*.""" + ... + + +# --------------------------------------------------------------------------- +# Capability → protocol mapping +# --------------------------------------------------------------------------- + +#: Which protocol an adapter must implement to serve each capability. The +#: routers use this for discovery, so adding a capability means adding one entry +#: here rather than editing routing code. +CAPABILITY_PROTOCOLS: dict[MediaCapability, type] = { + MediaCapability.IMAGE_GENERATE: ImageGenerationProvider, + MediaCapability.IMAGE_EDIT: ImageEditProvider, + MediaCapability.IMAGE_UPSCALE: ImageUpscaleProvider, + MediaCapability.IMAGE_VARIATE: ImageGenerationProvider, + MediaCapability.OCR: OCRMediaProvider, + MediaCapability.TTS: TTSProvider, + MediaCapability.TTS_STREAM: StreamingTTSProvider, + MediaCapability.ASR: ASRProvider, + MediaCapability.ASR_STREAM: StreamingASRProvider, + MediaCapability.VOICE_AGENT: StreamingASRProvider, + MediaCapability.MUSIC: MusicProvider, + MediaCapability.SFX: SFXProvider, + MediaCapability.VIDEO_GENERATE: VideoGenerationProvider, + MediaCapability.VIDEO_EDIT: VideoEditProvider, + MediaCapability.VIDEO_INTERPOLATE: VideoInterpolationProvider, + MediaCapability.VIDEO_EXTEND: VideoEditProvider, + MediaCapability.VIDEO_REFRAME: VideoEditProvider, + MediaCapability.VIDEO_UPSCALE: VideoEditProvider, + MediaCapability.VOICE_DESIGN: TTSProvider, +} diff --git a/src/llmcore/media/routers.py b/src/llmcore/media/routers.py new file mode 100644 index 00000000..dfc840d9 --- /dev/null +++ b/src/llmcore/media/routers.py @@ -0,0 +1,407 @@ +# src/llmcore/media/routers.py +"""Per-modality routers for the LLMCore media subsystem. + +A router is the caller-facing surface for one modality. Its only jobs are to +resolve a capability to an adapter, call the protocol method, and hand any +returned job to the job manager so an expensive submission is never lost. All +vendor logic lives in the adapters; all selection logic lives in +:class:`~llmcore.media.manager.MediaManager`. + +Every router method accepts ``provider=`` and ``model=`` to pin the route +explicitly, and forwards unknown keyword arguments to the adapter so +vendor-specific parameters need no plumbing here. +""" + +from __future__ import annotations + +import logging +from collections.abc import AsyncIterator, Sequence +from typing import TYPE_CHECKING, Any + +from .models import MediaCapability, MediaJob, MediaRef, MediaResult + +if TYPE_CHECKING: # pragma: no cover + from .manager import MediaManager + +logger = logging.getLogger(__name__) + +__all__ = ["AudioRouter", "ImageRouter", "VideoRouter"] + + +class _BaseRouter: + """Shared dispatch for the modality routers. + + Args: + manager: The owning media manager, used for resolution and job tracking. + """ + + def __init__(self, manager: MediaManager) -> None: + self._manager = manager + + async def _dispatch( + self, + capability: MediaCapability, + method: str, + *args: Any, + provider: str | None = None, + model: str | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Resolve *capability* and invoke *method* on the chosen adapter. + + Args: + capability: The operation being requested. + method: Protocol method name to call on the adapter. + *args: Positional arguments for the adapter method. + provider: Pin the provider instance. + model: Pin the model; also used to resolve the provider. + **kwargs: Forwarded to the adapter, vendor parameters included. + + Returns: + The adapter's result, with any job handle registered for tracking. + """ + adapter = self._manager.resolve(capability, provider=provider, model=model) + fn = getattr(adapter, method) + outcome = await fn(*args, model=model, **kwargs) + return self._manager._finalize(outcome) + + def _stream( + self, + capability: MediaCapability, + method: str, + *args: Any, + provider: str | None = None, + model: str | None = None, + **kwargs: Any, + ) -> AsyncIterator[bytes]: + """Resolve *capability* and return the adapter's byte stream. + + Not a coroutine: streaming methods return an async iterator directly, so + callers write ``async for chunk in router.stream_...(...)``. + """ + adapter = self._manager.resolve(capability, provider=provider, model=model) + return getattr(adapter, method)(*args, model=model, **kwargs) + + +class ImageRouter(_BaseRouter): + """Image generation, editing, upscaling and OCR.""" + + async def generate( + self, + prompt: str, + *, + provider: str | None = None, + model: str | None = None, + n: int = 1, + size: str | None = None, + seed: int | None = None, + negative_prompt: str | None = None, + reference_images: Sequence[MediaRef] | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Generate images from a text prompt.""" + return await self._dispatch( + MediaCapability.IMAGE_GENERATE, + "generate_image_media", + prompt, + provider=provider, + model=model, + n=n, + size=size, + seed=seed, + negative_prompt=negative_prompt, + reference_images=reference_images, + **kwargs, + ) + + async def edit( + self, + prompt: str, + *, + image: MediaRef, + mask: MediaRef | None = None, + provider: str | None = None, + model: str | None = None, + n: int = 1, + size: str | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Edit an existing image according to a prompt.""" + return await self._dispatch( + MediaCapability.IMAGE_EDIT, + "edit_image_media", + prompt, + provider=provider, + model=model, + image=image, + mask=mask, + n=n, + size=size, + **kwargs, + ) + + async def upscale( + self, + *, + image: MediaRef, + provider: str | None = None, + model: str | None = None, + scale: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Increase an image's resolution.""" + return await self._dispatch( + MediaCapability.IMAGE_UPSCALE, + "upscale_image_media", + provider=provider, + model=model, + image=image, + scale=scale, + **kwargs, + ) + + async def ocr( + self, + *, + document: MediaRef, + provider: str | None = None, + model: str | None = None, + pages: Sequence[int] | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Extract text from a document or image.""" + return await self._dispatch( + MediaCapability.OCR, + "ocr_media", + provider=provider, + model=model, + document=document, + pages=pages, + **kwargs, + ) + + +class AudioRouter(_BaseRouter): + """Speech synthesis, transcription, music and sound effects.""" + + async def speak( + self, + text: str, + *, + provider: str | None = None, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + speed: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Synthesize speech from text in one shot.""" + return await self._dispatch( + MediaCapability.TTS, + "synthesize_speech_media", + text, + provider=provider, + model=model, + voice=voice, + audio_format=audio_format, + sample_rate_hz=sample_rate_hz, + speed=speed, + **kwargs, + ) + + def stream_tts( + self, + text: str, + *, + provider: str | None = None, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> AsyncIterator[bytes]: + """Stream synthesized speech as it is produced.""" + return self._stream( + MediaCapability.TTS_STREAM, + "stream_speech_media", + text, + provider=provider, + model=model, + voice=voice, + audio_format=audio_format, + sample_rate_hz=sample_rate_hz, + **kwargs, + ) + + async def transcribe( + self, + *, + audio: MediaRef, + provider: str | None = None, + model: str | None = None, + language: str | None = None, + diarize: bool | None = None, + timestamps: bool | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Transcribe audio to text.""" + return await self._dispatch( + MediaCapability.ASR, + "transcribe_media", + provider=provider, + model=model, + audio=audio, + language=language, + diarize=diarize, + timestamps=timestamps, + **kwargs, + ) + + async def open_transcription_session( + self, + *, + provider: str | None = None, + model: str | None = None, + language: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> Any: + """Open a realtime, bidirectional transcription session.""" + adapter = self._manager.resolve( + MediaCapability.ASR_STREAM, provider=provider, model=model + ) + return await adapter.open_transcription_session( + model=model, language=language, sample_rate_hz=sample_rate_hz, **kwargs + ) + + async def music( + self, + prompt: str, + *, + provider: str | None = None, + model: str | None = None, + duration_seconds: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Generate music from a text prompt.""" + return await self._dispatch( + MediaCapability.MUSIC, + "generate_music_media", + prompt, + provider=provider, + model=model, + duration_seconds=duration_seconds, + **kwargs, + ) + + async def sfx( + self, + prompt: str | None = None, + *, + provider: str | None = None, + model: str | None = None, + video: MediaRef | None = None, + duration_seconds: float | None = None, + **kwargs: Any, + ) -> MediaResult | MediaJob: + """Generate sound effects from a prompt and/or a video (foley).""" + return await self._dispatch( + MediaCapability.SFX, + "generate_sfx_media", + prompt, + provider=provider, + model=model, + video=video, + duration_seconds=duration_seconds, + **kwargs, + ) + + +class VideoRouter(_BaseRouter): + """Video generation, editing and frame interpolation.""" + + async def generate( + self, + prompt: str, + *, + provider: str | None = None, + model: str | None = None, + first_frame: MediaRef | None = None, + last_frame: MediaRef | None = None, + reference_images: Sequence[MediaRef] | None = None, + duration_seconds: float | None = None, + resolution: str | None = None, + aspect_ratio: str | None = None, + fps: float | None = None, + with_audio: bool | None = None, + seed: int | None = None, + **kwargs: Any, + ) -> MediaJob: + """Generate a video; returns a job handle. + + ``first_frame``/``last_frame`` request a *generative* transition, which + is distinct from :meth:`interpolate`. + """ + outcome = await self._dispatch( + MediaCapability.VIDEO_GENERATE, + "generate_video_media", + prompt, + provider=provider, + model=model, + first_frame=first_frame, + last_frame=last_frame, + reference_images=reference_images, + duration_seconds=duration_seconds, + resolution=resolution, + aspect_ratio=aspect_ratio, + fps=fps, + with_audio=with_audio, + seed=seed, + **kwargs, + ) + return outcome # type: ignore[return-value] + + async def edit( + self, + prompt: str, + *, + video: MediaRef, + provider: str | None = None, + model: str | None = None, + **kwargs: Any, + ) -> MediaJob: + """Edit an existing video according to a prompt.""" + outcome = await self._dispatch( + MediaCapability.VIDEO_EDIT, + "edit_video_media", + prompt, + provider=provider, + model=model, + video=video, + **kwargs, + ) + return outcome # type: ignore[return-value] + + async def interpolate( + self, + *, + video: MediaRef | None = None, + frames: Sequence[MediaRef] | None = None, + provider: str | None = None, + model: str | None = None, + target_fps: float | None = None, + **kwargs: Any, + ) -> MediaJob: + """Interpolate frames to raise FPS or blend between supplied frames.""" + outcome = await self._dispatch( + MediaCapability.VIDEO_INTERPOLATE, + "interpolate_video_media", + provider=provider, + model=model, + video=video, + frames=frames, + target_fps=target_fps, + **kwargs, + ) + return outcome # type: ignore[return-value] diff --git a/src/llmcore/media/testing.py b/src/llmcore/media/testing.py new file mode 100644 index 00000000..e49bafbb --- /dev/null +++ b/src/llmcore/media/testing.py @@ -0,0 +1,285 @@ +# src/llmcore/media/testing.py +"""In-repo fake media adapter, for exercising the subsystem without a network. + +:class:`FakeMediaProvider` implements every capability protocol, so router +dispatch, capability discovery, selection policy, job lifecycle and artifact +materialization can all be tested without a vendor account. It is also the +reference for what a real adapter must implement. + +Shipped inside the package (not under ``tests/``) so that downstream projects +building their own media adapters can use it in their suites too. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import AsyncIterator, Sequence +from datetime import UTC, datetime, timedelta +from typing import Any + +from .models import ( + MediaArtifact, + MediaCapability, + MediaExecution, + MediaJob, + MediaJobStatus, + MediaKind, + MediaRef, + MediaResult, + MediaUsage, +) + +__all__ = ["FakeMediaProvider"] + +#: Capabilities the fake serves as long-running jobs rather than immediately. +_JOB_CAPABILITIES: frozenset[MediaCapability] = frozenset( + { + MediaCapability.VIDEO_GENERATE, + MediaCapability.VIDEO_EDIT, + MediaCapability.VIDEO_INTERPOLATE, + } +) + + +class FakeMediaProvider: + """A deterministic, offline media adapter. + + Args: + name: Provider instance name reported by :meth:`get_name`. + capabilities: Capabilities to declare; every capability by default. + poll_count: How many polls a job needs before succeeding. ``0`` makes + jobs succeed on submission. + fail_jobs: Make every job terminate ``FAILED``, for error-path tests. + declare_only: Capabilities to declare **without** implementing, to test + that the manager drops unbacked declarations. + """ + + def __init__( + self, + name: str = "fake", + *, + capabilities: Sequence[MediaCapability] | None = None, + poll_count: int = 1, + fail_jobs: bool = False, + declare_only: Sequence[MediaCapability] | None = None, + ) -> None: + self._name = name + self._capabilities = frozenset(capabilities) if capabilities is not None else frozenset( + MediaCapability + ) + self._declare_only = frozenset(declare_only or ()) + self._poll_count = max(0, int(poll_count)) + self._fail_jobs = fail_jobs + self._polls: dict[str, int] = {} + #: Every call recorded as ``(method, kwargs)``, for assertions. + self.calls: list[tuple[str, dict[str, Any]]] = [] + + # --- MediaCapableProvider --- + + def get_name(self) -> str: + """Return the instance name.""" + return self._name + + def media_capabilities(self) -> frozenset[MediaCapability]: + """Return declared capabilities, including any declare-only ones.""" + return self._capabilities | self._declare_only + + def media_execution( + self, capability: MediaCapability, model: str | None = None + ) -> MediaExecution: + """Return the execution class the fake uses for *capability*.""" + if capability in _JOB_CAPABILITIES: + return MediaExecution.ASYNC_JOB + if capability in (MediaCapability.TTS_STREAM, MediaCapability.ASR_STREAM): + return MediaExecution.STREAM + return MediaExecution.REQUEST_RESPONSE + + # --- helpers --- + + def _record(self, method: str, kwargs: dict[str, Any]) -> None: + self.calls.append((method, {k: v for k, v in kwargs.items() if v is not None})) + + def _artifact(self, kind: MediaKind, payload: str, *, expires: bool = False) -> MediaArtifact: + data = payload.encode() + return MediaArtifact( + kind=kind, + uri=f"https://fake.invalid/{hashlib.sha256(data).hexdigest()[:12]}", + mime_type={"image": "image/png", "audio": "audio/mpeg", "video": "video/mp4"}.get( + kind.value, "text/plain" + ), + checksum_sha256=hashlib.sha256(data).hexdigest(), + expires_at=(datetime.now(UTC) + timedelta(hours=1)) if expires else None, + provider_metadata={"fake": True}, + ) + + def _result( + self, capability: MediaCapability, payload: str, *, model: str | None, n: int = 1, + kind: MediaKind | None = None, text: str | None = None, + ) -> MediaResult: + resolved_kind = kind or { + True: MediaKind.IMAGE, + }.get(False, MediaKind.IMAGE) + artifacts = [self._artifact(resolved_kind, f"{payload}-{i}") for i in range(n)] + if text is not None: + artifacts = [ + MediaArtifact(**{**{f: getattr(a, f) for f in a.__slots__}, "text": text}) + for a in artifacts + ] + return MediaResult( + capability=capability, + provider=self._name, + model=model or "fake-model", + artifacts=tuple(artifacts), + usage=MediaUsage( + provider=self._name, model=model or "fake-model", basis="per_call", images=n + ), + raw={"fake": True}, + ) + + def _job(self, capability: MediaCapability, *, model: str | None) -> MediaJob: + job = MediaJob( + capability=capability, + provider=self._name, + model=model or "fake-model", + provider_job_id=f"fake-{len(self._polls) + 1}", + idempotency_key=f"idem-{len(self._polls) + 1}", + ) + self._polls[job.id] = 0 + if self._poll_count == 0: + self._complete(job) + return job + + def _complete(self, job: MediaJob) -> None: + if self._fail_jobs: + job.status = MediaJobStatus.FAILED + job.error = "fake failure" + return + job.status = MediaJobStatus.SUCCEEDED + job.progress = 1.0 + job.artifacts = [self._artifact(MediaKind.VIDEO, job.id, expires=True)] + job.usage = MediaUsage(provider=self._name, model=job.model, basis="per_second", seconds=4.0) + + # --- MediaJobPoller --- + + async def poll_media_job(self, job: MediaJob) -> MediaJob: + """Advance *job* one step toward completion.""" + self._record("poll_media_job", {"job": job.id}) + if job.is_terminal: + return job + self._polls[job.id] = self._polls.get(job.id, 0) + 1 + if self._polls[job.id] >= self._poll_count: + self._complete(job) + else: + job.status = MediaJobStatus.RUNNING + job.progress = self._polls[job.id] / self._poll_count + return job + + async def cancel_media_job(self, job: MediaJob) -> MediaJob: + """Mark *job* cancelled.""" + self._record("cancel_media_job", {"job": job.id}) + job.status = MediaJobStatus.CANCELED + return job + + # --- image --- + + async def generate_image_media( + self, prompt: str, *, model: str | None = None, n: int = 1, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce *n* fake images.""" + self._record("generate_image_media", {"prompt": prompt, "model": model, "n": n, **kwargs}) + return self._result(MediaCapability.IMAGE_GENERATE, prompt, model=model, n=n) + + async def edit_image_media( + self, prompt: str, *, image: MediaRef, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce a fake edited image.""" + self._record("edit_image_media", {"prompt": prompt, "model": model, **kwargs}) + return self._result(MediaCapability.IMAGE_EDIT, prompt, model=model) + + async def upscale_image_media( + self, *, image: MediaRef, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce a fake upscaled image.""" + self._record("upscale_image_media", {"model": model, **kwargs}) + return self._result(MediaCapability.IMAGE_UPSCALE, "upscaled", model=model) + + async def ocr_media( + self, *, document: MediaRef, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce fake extracted text.""" + self._record("ocr_media", {"model": model, **kwargs}) + return self._result( + MediaCapability.OCR, "ocr", model=model, kind=MediaKind.TEXT, text="fake ocr text" + ) + + # --- audio --- + + async def synthesize_speech_media( + self, text: str, *, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce a fake audio artifact.""" + self._record("synthesize_speech_media", {"text": text, "model": model, **kwargs}) + return self._result(MediaCapability.TTS, text, model=model, kind=MediaKind.AUDIO) + + def stream_speech_media( + self, text: str, *, model: str | None = None, **kwargs: Any + ) -> AsyncIterator[bytes]: + """Yield the text back as fake audio chunks, one word at a time.""" + self._record("stream_speech_media", {"text": text, "model": model, **kwargs}) + + async def _gen() -> AsyncIterator[bytes]: + for word in text.split(): + yield word.encode() + + return _gen() + + async def transcribe_media( + self, *, audio: MediaRef, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce a fake transcript.""" + self._record("transcribe_media", {"model": model, **kwargs}) + return self._result( + MediaCapability.ASR, "asr", model=model, kind=MediaKind.TEXT, text="fake transcript" + ) + + async def open_transcription_session(self, *, model: str | None = None, **kwargs: Any) -> Any: + """Return a trivial stand-in session object.""" + self._record("open_transcription_session", {"model": model, **kwargs}) + return {"session": "fake", "model": model} + + async def generate_music_media( + self, prompt: str, *, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce a fake music artifact.""" + self._record("generate_music_media", {"prompt": prompt, "model": model, **kwargs}) + return self._result(MediaCapability.MUSIC, prompt, model=model, kind=MediaKind.AUDIO) + + async def generate_sfx_media( + self, prompt: str | None = None, *, model: str | None = None, **kwargs: Any + ) -> MediaResult | MediaJob: + """Produce a fake SFX artifact.""" + self._record("generate_sfx_media", {"prompt": prompt, "model": model, **kwargs}) + return self._result(MediaCapability.SFX, prompt or "sfx", model=model, kind=MediaKind.AUDIO) + + # --- video --- + + async def generate_video_media( + self, prompt: str, *, model: str | None = None, **kwargs: Any + ) -> MediaJob: + """Submit a fake video job.""" + self._record("generate_video_media", {"prompt": prompt, "model": model, **kwargs}) + return self._job(MediaCapability.VIDEO_GENERATE, model=model) + + async def edit_video_media( + self, prompt: str, *, video: MediaRef, model: str | None = None, **kwargs: Any + ) -> MediaJob: + """Submit a fake video-edit job.""" + self._record("edit_video_media", {"prompt": prompt, "model": model, **kwargs}) + return self._job(MediaCapability.VIDEO_EDIT, model=model) + + async def interpolate_video_media( + self, *, video: MediaRef | None = None, model: str | None = None, **kwargs: Any + ) -> MediaJob: + """Submit a fake interpolation job.""" + self._record("interpolate_video_media", {"model": model, **kwargs}) + return self._job(MediaCapability.VIDEO_INTERPOLATE, model=model) diff --git a/src/llmcore/providers/manager.py b/src/llmcore/providers/manager.py index 5ac407da..f9727aa4 100644 --- a/src/llmcore/providers/manager.py +++ b/src/llmcore/providers/manager.py @@ -145,6 +145,7 @@ class ProviderManager: """ _providers: dict[str, BaseProvider] + _ephemeral_instances: set[str] _config: ConfyConfig _default_provider_name: str _event_logger: Any | None @@ -175,6 +176,9 @@ def __init__( """ self._config = config self._providers = {} + # Instances registered at runtime by a subsystem (media adapters, remote + # compute runtimes) rather than loaded from [providers.*]. + self._ephemeral_instances: set[str] = set() self._log_raw_payloads_override = log_raw_payloads self._event_logger = event_logger self._initialized = False @@ -715,6 +719,121 @@ def _load_configured_providers(self) -> None: if not self._providers: logger.warning("No provider instances were successfully loaded.") + def register_instance( + self, + name: str, + provider_type: str, + config: dict[str, Any], + *, + ephemeral: bool = False, + replace: bool = False, + ) -> BaseProvider: + """Build and register a provider instance at runtime. + + Providers are normally constructed during ``__init__`` from + ``[providers.*]``. This adds one afterwards, which is what subsystems + that *create* endpoints need: a remote GPU runtime that has just booted + (see ``docs/COLAB_RUNTIME_SPEC.md``) or a media adapter discovered + dynamically. + + Args: + name: Instance name callers will pass to ``get_provider()``. + provider_type: A key in :data:`PROVIDER_MAP`. + config: Provider configuration, as a ``[providers.]`` section + would supply it. + ephemeral: Mark the instance as owned by a subsystem, so + ``close_providers()`` knows it was not user-configured. Purely + informational today; consumed by the runtimes subsystem. + replace: Allow replacing an existing instance of the same name. + Without it, a collision raises rather than silently swapping a + live provider out from under its callers. + + Returns: + The constructed provider instance. + + Raises: + ConfigError: If the type is unknown, the name collides without + *replace*, or construction fails. + """ + key = name.lower() + provider_cls = PROVIDER_MAP.get(provider_type.lower()) + if provider_cls is None: + raise ConfigError( + f"Cannot register provider '{key}': type '{provider_type}' is not supported. " + f"Known types: {', '.join(sorted(PROVIDER_MAP))}" + ) + if key in self._providers and not replace: + raise ConfigError( + f"Provider instance '{key}' already exists. Pass replace=True to swap it." + ) + + instance_config = dict(config) + instance_config["_instance_name"] = key + log_raw = ( + self._log_raw_payloads_override + if self._log_raw_payloads_override + else self._config.get("llmcore.log_raw_payloads", False) + ) + try: + provider = provider_cls(instance_config, log_raw_payloads=log_raw) + except Exception as e: + raise ConfigError(f"Failed to register provider '{key}': {e}") from e + + previous = self._providers.get(key) + self._providers[key] = provider + if ephemeral: + self._ephemeral_instances.add(key) + else: + self._ephemeral_instances.discard(key) + logger.info( + "Registered provider instance '%s' (type=%s, ephemeral=%s)%s.", + key, + provider_type, + ephemeral, + " replacing an existing instance" if previous is not None else "", + ) + return provider + + async def unregister_instance(self, name: str, *, close: bool = True) -> bool: + """Remove a dynamically registered provider instance. + + Args: + name: The instance name. + close: Await the provider's ``close()`` before dropping it. + + Returns: + ``True`` if an instance was removed, ``False`` if none existed. + + Raises: + ConfigError: If *name* is the configured default provider, since + removing it would leave the manager unable to serve a default. + """ + key = name.lower() + if key == self._default_provider_name: + raise ConfigError( + f"Refusing to unregister '{key}': it is the configured default provider." + ) + provider = self._providers.pop(key, None) + self._ephemeral_instances.discard(key) + if provider is None: + return False + if close: + try: + await provider.close() + except Exception as e: # noqa: BLE001 - teardown must not raise + logger.error("Error closing provider '%s' during unregister: %s", key, e) + logger.info("Unregistered provider instance '%s'.", key) + return True + + def is_ephemeral(self, name: str) -> bool: + """Whether *name* was registered at runtime by a subsystem.""" + return name.lower() in self._ephemeral_instances + + @property + def ephemeral_instances(self) -> list[str]: + """Names of every runtime-registered instance.""" + return sorted(self._ephemeral_instances) + def get_provider(self, name: str | None = None) -> BaseProvider: """ Gets a provider instance by its configured name, or the default provider. diff --git a/tests/media/__init__.py b/tests/media/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/media/test_media_integration.py b/tests/media/test_media_integration.py new file mode 100644 index 00000000..9183d997 --- /dev/null +++ b/tests/media/test_media_integration.py @@ -0,0 +1,224 @@ +# tests/media/test_media_integration.py +"""Integration of the media subsystem with the facade and the provider manager. + +Covers ``llm.media`` wiring, the ``[media]`` config section, and the dynamic +provider registration that both the media and remote-runtime programs need. +""" + +from __future__ import annotations + +import tomllib +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +import llmcore +from llmcore.exceptions import ConfigError +from llmcore.media import MediaCapability, MediaManager +from llmcore.media.testing import FakeMediaProvider +from llmcore.providers.manager import ProviderManager + +DEFAULT_CONFIG = Path(llmcore.__file__).parent / "config" / "default_config.toml" + + +def _config(providers: dict, default: str = "vllm"): + """A mock ConfyConfig exposing ``.get(key, default)``.""" + store = { + "llmcore.default_provider": default, + "llmcore.log_raw_payloads": False, + "providers": providers, + } + cfg = MagicMock() + cfg.get = lambda key, d=None: store.get(key, d) + return cfg + + +VLLM = {"base_url": "http://localhost:8000/v1", "default_model": "m"} + + +# --------------------------------------------------------------------------- +# Config section +# --------------------------------------------------------------------------- + + +class TestMediaConfigSection: + @pytest.fixture(scope="class") + def cfg(self) -> dict: + with open(DEFAULT_CONFIG, "rb") as fh: + return tomllib.load(fh) + + def test_media_section_exists(self, cfg): + assert "media" in cfg + + def test_artifact_defaults(self, cfg): + assert cfg["media"]["artifact_materialize"] == "on_expiry" + assert cfg["media"]["artifact_path"] + + def test_job_defaults(self, cfg): + jobs = cfg["media"]["jobs"] + assert jobs["poll_initial_seconds"] < jobs["poll_max_seconds"] + assert jobs["job_timeout_seconds"] >= jobs["poll_max_seconds"] + + def test_routing_table_present_and_commented_out(self, cfg): + # Present so the table exists, empty so built-in defaults apply. + assert cfg["media"]["routing"] == {} + + def test_no_media_provider_credentials_section(self, cfg): + """Media adapters are the chat providers; credentials are not duplicated.""" + assert "providers" not in cfg["media"] + + +# --------------------------------------------------------------------------- +# Facade wiring +# --------------------------------------------------------------------------- + + +class TestFacadeWiring: + def test_media_property_exists(self): + assert isinstance(getattr(llmcore.LLMCore, "media", None), property) + + def test_media_before_create_raises(self): + # __init__ is private; a bare instance has no initialized subsystems. + bare = llmcore.LLMCore() + with pytest.raises(ConfigError, match="not initialized"): + _ = bare.media + + def test_media_manager_is_exported(self): + from llmcore.media import MediaManager as Exported + + assert Exported is MediaManager + + +# --------------------------------------------------------------------------- +# Dynamic provider registration (shared prerequisite) +# --------------------------------------------------------------------------- + + +class TestDynamicProviderRegistration: + @pytest.fixture + def pm(self) -> ProviderManager: + return ProviderManager(_config({"vllm": dict(VLLM)})) + + def test_baseline(self, pm): + assert pm.get_available_providers() == ["vllm"] + assert pm.ephemeral_instances == [] + + def test_register_instance(self, pm): + provider = pm.register_instance( + "colab-qwen", "vllm", {**VLLM, "base_url": "http://127.0.0.1:19001/v1"} + ) + assert provider.get_name() == "colab-qwen" + assert pm.get_provider("colab-qwen") is provider + assert "colab-qwen" in pm.get_available_providers() + + def test_register_marks_ephemeral(self, pm): + pm.register_instance("rt", "vllm", dict(VLLM), ephemeral=True) + assert pm.is_ephemeral("rt") is True + assert pm.ephemeral_instances == ["rt"] + assert pm.is_ephemeral("vllm") is False + + def test_name_is_case_insensitive(self, pm): + pm.register_instance("MixedCase", "vllm", dict(VLLM)) + assert pm.get_provider("mixedcase").get_name() == "mixedcase" + + def test_collision_raises_without_replace(self, pm): + pm.register_instance("rt", "vllm", dict(VLLM)) + with pytest.raises(ConfigError, match="already exists"): + pm.register_instance("rt", "vllm", dict(VLLM)) + + def test_replace_swaps_the_instance(self, pm): + first = pm.register_instance("rt", "vllm", dict(VLLM)) + second = pm.register_instance("rt", "vllm", dict(VLLM), replace=True) + assert first is not second + assert pm.get_provider("rt") is second + + def test_replace_clears_ephemeral_when_not_requested(self, pm): + pm.register_instance("rt", "vllm", dict(VLLM), ephemeral=True) + pm.register_instance("rt", "vllm", dict(VLLM), replace=True) + assert pm.is_ephemeral("rt") is False + + def test_unknown_type_raises(self, pm): + with pytest.raises(ConfigError, match="not supported"): + pm.register_instance("x", "no-such-provider", {}) + + def test_construction_failure_raises_config_error(self, pm): + # vllm requires a base_url; omitting it fails construction. + with pytest.raises(ConfigError, match="Failed to register"): + pm.register_instance("broken", "vllm", {}) + + async def test_unregister(self, pm): + pm.register_instance("rt", "vllm", dict(VLLM), ephemeral=True) + assert await pm.unregister_instance("rt") is True + assert "rt" not in pm.get_available_providers() + assert pm.ephemeral_instances == [] + + async def test_unregister_missing_returns_false(self, pm): + assert await pm.unregister_instance("ghost") is False + + async def test_unregister_refuses_the_default(self, pm): + with pytest.raises(ConfigError, match="default provider"): + await pm.unregister_instance("vllm") + + async def test_unregister_closes_by_default(self, pm): + provider = pm.register_instance("rt", "vllm", dict(VLLM)) + closed = False + + async def _close(): + nonlocal closed + closed = True + + provider.close = _close + await pm.unregister_instance("rt") + assert closed is True + + async def test_unregister_can_skip_close(self, pm): + provider = pm.register_instance("rt", "vllm", dict(VLLM)) + closed = False + + async def _close(): + nonlocal closed + closed = True + + provider.close = _close + await pm.unregister_instance("rt", close=False) + assert closed is False + + async def test_close_failure_does_not_block_unregister(self, pm): + provider = pm.register_instance("rt", "vllm", dict(VLLM)) + + async def _boom(): + raise RuntimeError("vendor down") + + provider.close = _boom + assert await pm.unregister_instance("rt") is True + assert "rt" not in pm.get_available_providers() + + +# --------------------------------------------------------------------------- +# Manager built from a real ProviderManager +# --------------------------------------------------------------------------- + + +class TestManagerFromRealProviderManager: + def test_chat_only_providers_are_not_adapters(self): + """M1 adds the core only; no shipped provider implements the protocols yet.""" + pm = ProviderManager(_config({"vllm": dict(VLLM)})) + media = MediaManager.from_provider_manager(pm, lambda k, d=None: d) + assert media.adapter_names == [] + assert media.has_adapters() is False + + def test_a_registered_media_adapter_is_routable(self): + pm = ProviderManager(_config({"vllm": dict(VLLM)})) + media = MediaManager.from_provider_manager(pm, lambda k, d=None: d) + media.register_adapter("fake", FakeMediaProvider("fake")) + assert media.who_can(MediaCapability.IMAGE_GENERATE) == ["fake"] + + async def test_end_to_end_through_the_manager(self): + pm = ProviderManager(_config({"vllm": dict(VLLM)})) + media = MediaManager.from_provider_manager(pm, lambda k, d=None: d) + media.register_adapter("fake", FakeMediaProvider("fake", poll_count=0)) + result = await media.images.generate("a tabby", n=2) + assert len(result.artifacts) == 2 + job = await media.video.generate("dunes") + assert (await media.wait(job)).usage.seconds == 4.0 diff --git a/tests/media/test_media_models.py b/tests/media/test_media_models.py new file mode 100644 index 00000000..66d8a4f5 --- /dev/null +++ b/tests/media/test_media_models.py @@ -0,0 +1,228 @@ +# tests/media/test_media_models.py +"""Tests for the provider-agnostic media types. + +Covers MediaRef construction/validation, artifact expiry and materialization +signalling, usage/result projection, job lifecycle, and the capability tables. +""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta + +import pytest + +from llmcore.media.models import ( + AUDIO_CAPABILITIES, + CAPABILITY_KINDS, + IMAGE_CAPABILITIES, + TERMINAL_JOB_STATUSES, + VIDEO_CAPABILITIES, + MediaArtifact, + MediaCapability, + MediaJob, + MediaJobStatus, + MediaKind, + MediaProvenance, + MediaRef, + MediaResult, + MediaUsage, +) + + +class TestMediaRef: + def test_from_bytes(self): + r = MediaRef.from_bytes(b"abc", mime_type="image/png") + assert r.read_bytes() == b"abc" + assert r.is_remote is False + + def test_from_path_infers_mime_and_filename(self, tmp_path): + p = tmp_path / "pic.png" + p.write_bytes(b"data") + r = MediaRef.from_path(p) + assert r.mime_type == "image/png" + assert r.filename == "pic.png" + assert r.read_bytes() == b"data" + + def test_from_url_is_remote(self): + r = MediaRef.from_url("https://example.invalid/a.mp4") + assert r.is_remote is True + assert r.mime_type == "video/mp4" + + def test_remote_read_bytes_refuses(self): + r = MediaRef.from_url("https://example.invalid/a.png") + with pytest.raises(ValueError, match="provider adapter must fetch"): + r.read_bytes() + + def test_exactly_one_source_required(self): + with pytest.raises(ValueError, match="exactly one"): + MediaRef() + with pytest.raises(ValueError, match="exactly one"): + MediaRef(url="https://x", data=b"y") + + def test_as_data_uri(self): + r = MediaRef.from_bytes(b"abc", mime_type="image/png") + assert r.as_data_uri() == "data:image/png;base64,YWJj" + + def test_as_data_uri_without_mime_falls_back(self): + assert MediaRef.from_bytes(b"abc").as_data_uri().startswith( + "data:application/octet-stream;base64," + ) + + def test_from_artifact_prefers_bytes(self): + a = MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y.png", data=b"raw") + assert MediaRef.from_artifact(a).data == b"raw" + + def test_from_artifact_falls_back_to_uri(self): + a = MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y.png") + assert MediaRef.from_artifact(a).url == "https://x/y.png" + + def test_from_artifact_requires_content(self): + with pytest.raises(ValueError, match="neither data nor uri"): + MediaRef.from_artifact(MediaArtifact(kind=MediaKind.IMAGE)) + + +class TestMediaArtifact: + def test_expiry_flags(self): + past = MediaArtifact( + kind=MediaKind.IMAGE, uri="u", expires_at=datetime.now(UTC) - timedelta(minutes=1) + ) + future = MediaArtifact( + kind=MediaKind.IMAGE, uri="u", expires_at=datetime.now(UTC) + timedelta(hours=1) + ) + assert past.is_expired is True + assert future.is_expired is False + + def test_needs_materialization_requires_uri_and_expiry(self): + assert MediaArtifact( + kind=MediaKind.IMAGE, uri="u", expires_at=datetime.now(UTC) + ).needs_materialization is True + # no expiry stated -> we cannot know it dies + assert MediaArtifact(kind=MediaKind.IMAGE, uri="u").needs_materialization is False + # already has bytes + assert MediaArtifact( + kind=MediaKind.IMAGE, uri="u", data=b"x", expires_at=datetime.now(UTC) + ).needs_materialization is False + + def test_with_data_sets_checksum_and_preserves_fields(self): + a = MediaArtifact(kind=MediaKind.AUDIO, uri="u", mime_type="audio/mpeg", sample_rate_hz=44100) + b = a.with_data(b"hello") + assert b.data == b"hello" + assert b.checksum_sha256 == ( + "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824" + ) + assert b.sample_rate_hz == 44100 and b.mime_type == "audio/mpeg" + + def test_to_dict_summarizes_bytes_not_embeds_them(self): + d = MediaArtifact(kind=MediaKind.IMAGE, data=b"0" * 4096).to_dict() + assert d["data"] == "<4096 bytes>" + assert d["kind"] == "image" + + def test_to_dict_serializes_datetimes(self): + when = datetime(2026, 1, 2, 3, 4, 5, tzinfo=UTC) + assert MediaArtifact(kind=MediaKind.IMAGE, expires_at=when).to_dict()["expires_at"] == ( + when.isoformat() + ) + + def test_provenance_is_carried(self): + a = MediaArtifact( + kind=MediaKind.IMAGE, + provenance=MediaProvenance(watermarked=True, generator="fake-1"), + ) + assert a.provenance.watermarked is True + assert a.to_dict()["provenance"]["generator"] == "fake-1" + + +class TestMediaResult: + def _result(self, **kw): + return MediaResult( + capability=MediaCapability.IMAGE_GENERATE, provider="p", model="m", **kw + ) + + def test_artifact_shortcut(self): + a = MediaArtifact(kind=MediaKind.IMAGE) + assert self._result(artifacts=(a,)).artifact is a + + def test_artifact_shortcut_raises_when_empty(self): + with pytest.raises(IndexError, match="produced no artifacts"): + _ = self._result().artifact + + def test_text_joins_artifact_text(self): + r = self._result( + artifacts=( + MediaArtifact(kind=MediaKind.TEXT, text="one"), + MediaArtifact(kind=MediaKind.TEXT, text="two"), + ) + ) + assert r.text == "one\ntwo" + + def test_text_is_none_without_text(self): + assert self._result(artifacts=(MediaArtifact(kind=MediaKind.IMAGE),)).text is None + + def test_usage_keeps_native_units(self): + u = MediaUsage(provider="p", model="m", basis="per_second", seconds=8.0, images=None) + d = u.to_dict() + assert d["seconds"] == 8.0 and d["basis"] == "per_second" + # a unit the vendor did not report stays absent rather than synthesized + assert d["input_tokens"] is None + + +class TestMediaJob: + def _job(self, **kw): + return MediaJob( + capability=MediaCapability.VIDEO_GENERATE, provider="p", model="m", **kw + ) + + def test_default_status_and_id(self): + j = self._job() + assert j.status is MediaJobStatus.QUEUED + assert j.id.startswith("mj_") and not j.is_terminal + + @pytest.mark.parametrize("status", sorted(TERMINAL_JOB_STATUSES)) + def test_terminal_statuses(self, status): + assert self._job(status=status).is_terminal is True + + @pytest.mark.parametrize( + "status", [MediaJobStatus.QUEUED, MediaJobStatus.RUNNING] + ) + def test_non_terminal_statuses(self, status): + assert self._job(status=status).is_terminal is False + + def test_to_result_requires_success(self): + with pytest.raises(ValueError, match="not succeeded"): + self._job(status=MediaJobStatus.FAILED).to_result() + + def test_to_result_projects_artifacts_and_usage(self): + usage = MediaUsage(provider="p", model="m", seconds=4.0) + j = self._job( + status=MediaJobStatus.SUCCEEDED, + artifacts=[MediaArtifact(kind=MediaKind.VIDEO, uri="u")], + usage=usage, + ) + r = j.to_result() + assert r.capability is MediaCapability.VIDEO_GENERATE + assert r.artifacts[0].uri == "u" + assert r.usage is usage + + def test_touch_advances_updated_at(self): + j = self._job() + before = j.updated_at + j.touch() + assert j.updated_at >= before + + +class TestCapabilityTables: + def test_every_capability_has_a_kind(self): + missing = [c for c in MediaCapability if c not in CAPABILITY_KINDS] + assert missing == [] + + def test_router_groups_are_disjoint_and_complete(self): + groups = [AUDIO_CAPABILITIES, IMAGE_CAPABILITIES, VIDEO_CAPABILITIES] + union = set().union(*groups) + assert union == set(MediaCapability) + for i, a in enumerate(groups): + for b in groups[i + 1 :]: + assert not (a & b) + + def test_text_producing_capabilities_report_text(self): + for cap in (MediaCapability.ASR, MediaCapability.ASR_STREAM, MediaCapability.OCR): + assert CAPABILITY_KINDS[cap] is MediaKind.TEXT diff --git a/tests/media/test_media_subsystem.py b/tests/media/test_media_subsystem.py new file mode 100644 index 00000000..09919056 --- /dev/null +++ b/tests/media/test_media_subsystem.py @@ -0,0 +1,566 @@ +# tests/media/test_media_subsystem.py +"""Tests for the media manager, routers, job manager and artifact store. + +Everything runs against the in-package FakeMediaProvider, so the whole +subsystem is exercised with no network and no vendor account. +""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta +from unittest.mock import MagicMock + +import pytest + +from llmcore.exceptions import ( + MediaCapabilityError, + MediaError, + MediaJobError, + MediaJobTimeoutError, +) +from llmcore.media import ( + ArtifactStore, + JobPolicy, + MaterializePolicy, + MediaArtifact, + MediaCapability, + MediaExecution, + MediaJobStatus, + MediaKind, + MediaManager, + MediaRef, +) +from llmcore.media.protocols import CAPABILITY_PROTOCOLS, MediaCapableProvider +from llmcore.media.testing import FakeMediaProvider + +# A fast policy so timeout/backoff tests do not actually sleep for seconds. +FAST = JobPolicy(poll_initial_seconds=0.0, poll_max_seconds=0.0, job_timeout_seconds=5.0, jitter=0.0) + + +@pytest.fixture +def fake() -> FakeMediaProvider: + return FakeMediaProvider("fake") + + +@pytest.fixture +def manager(fake: FakeMediaProvider) -> MediaManager: + return MediaManager({"fake": fake}, job_policy=FAST) + + +# --------------------------------------------------------------------------- +# Protocol coverage +# --------------------------------------------------------------------------- + + +class TestProtocolCoverage: + def test_every_capability_maps_to_a_protocol(self): + missing = [c for c in MediaCapability if c not in CAPABILITY_PROTOCOLS] + assert missing == [] + + def test_fake_satisfies_every_mapped_protocol(self, fake): + unmet = [ + cap.value + for cap, proto in CAPABILITY_PROTOCOLS.items() + if not isinstance(fake, proto) + ] + assert unmet == [] + + def test_fake_is_media_capable(self, fake): + assert isinstance(fake, MediaCapableProvider) + + def test_execution_classes_are_declared(self, fake): + assert fake.media_execution(MediaCapability.VIDEO_GENERATE) is MediaExecution.ASYNC_JOB + assert fake.media_execution(MediaCapability.TTS_STREAM) is MediaExecution.STREAM + assert ( + fake.media_execution(MediaCapability.IMAGE_GENERATE) + is MediaExecution.REQUEST_RESPONSE + ) + + +# --------------------------------------------------------------------------- +# Discovery and resolution +# --------------------------------------------------------------------------- + + +class TestDiscovery: + def test_capabilities_lists_providers(self, manager): + caps = manager.capabilities() + assert caps[MediaCapability.IMAGE_GENERATE] == ["fake"] + assert len(caps) == len(list(MediaCapability)) + + def test_who_can(self, manager): + assert manager.who_can(MediaCapability.VIDEO_GENERATE) == ["fake"] + assert manager.who_can("image_generate") == ["fake"] + + def test_no_adapters_is_not_an_error(self): + m = MediaManager() + assert m.has_adapters() is False + assert m.capabilities() == {} + assert m.who_can(MediaCapability.TTS) == [] + + def test_register_and_unregister(self, manager): + manager.register_adapter("second", FakeMediaProvider("second")) + assert manager.adapter_names == ["fake", "second"] + manager.unregister_adapter("second") + assert manager.adapter_names == ["fake"] + + def test_declared_but_unimplemented_capability_is_dropped(self): + # Declares TTS without implementing the TTS protocol. + class Bare: + def get_name(self) -> str: + return "bare" + + def media_capabilities(self): + return frozenset({MediaCapability.TTS}) + + def media_execution(self, capability, model=None): + return MediaExecution.REQUEST_RESPONSE + + m = MediaManager({"bare": Bare()}) + assert m.who_can(MediaCapability.TTS) == [] + + def test_broken_adapter_does_not_break_discovery(self): + class Broken: + def get_name(self) -> str: + return "broken" + + def media_capabilities(self): + raise RuntimeError("boom") + + def media_execution(self, capability, model=None): + return MediaExecution.REQUEST_RESPONSE + + m = MediaManager({"broken": Broken(), "fake": FakeMediaProvider("fake")}) + assert m.who_can(MediaCapability.TTS) == ["fake"] + + +class TestResolution: + def test_explicit_provider(self, manager, fake): + assert manager.resolve(MediaCapability.TTS, provider="fake") is fake + + def test_explicit_unknown_provider_raises(self, manager): + with pytest.raises(MediaCapabilityError, match="not configured"): + manager.resolve(MediaCapability.TTS, provider="nope") + + def test_explicit_provider_lacking_capability_raises(self): + limited = FakeMediaProvider("limited", capabilities=[MediaCapability.TTS]) + m = MediaManager({"limited": limited}) + with pytest.raises(MediaCapabilityError, match="does not support"): + m.resolve(MediaCapability.VIDEO_GENERATE, provider="limited") + + def test_no_candidate_raises_with_hints(self): + m = MediaManager() + with pytest.raises(MediaCapabilityError) as exc: + m.resolve(MediaCapability.VIDEO_GENERATE) + # the built-in preference list becomes the actionable hint + assert "fal" in str(exc.value) or "gemini" in str(exc.value) + + def test_routing_preference_decides_between_providers(self): + a = FakeMediaProvider("alpha") + b = FakeMediaProvider("beta") + m = MediaManager( + {"alpha": a, "beta": b}, + routing={MediaCapability.IMAGE_GENERATE: ("beta", "alpha")}, + ) + assert m.who_can(MediaCapability.IMAGE_GENERATE) == ["beta", "alpha"] + assert m.resolve(MediaCapability.IMAGE_GENERATE) is b + + def test_unconfigured_names_in_routing_are_skipped(self): + m = MediaManager( + {"fake": FakeMediaProvider("fake")}, + routing={MediaCapability.TTS: ("elevenlabs", "fake")}, + ) + assert m.who_can(MediaCapability.TTS) == ["fake"] + + +class TestRoutingConfig: + def test_routing_read_from_config(self): + store = {"media.routing": {"image_generate": ["beta", "alpha"]}} + get = lambda k, d=None: store.get(k, d) + routing = MediaManager._routing_from_config(get) + assert routing[MediaCapability.IMAGE_GENERATE] == ("beta", "alpha") + + def test_unknown_capability_in_config_is_ignored(self): + get = lambda k, d=None: {"media.routing": {"teleport": ["x"]}}.get(k, d) + assert MediaManager._routing_from_config(get) == {} + + def test_scalar_value_is_accepted(self): + get = lambda k, d=None: {"media.routing": {"tts": "fake"}}.get(k, d) + assert MediaManager._routing_from_config(get)[MediaCapability.TTS] == ("fake",) + + def test_non_table_routing_is_ignored(self): + get = lambda k, d=None: {"media.routing": ["nope"]}.get(k, d) + assert MediaManager._routing_from_config(get) == {} + + +class TestFromProviderManager: + def _pm(self, providers: dict): + pm = MagicMock() + pm.get_available_providers.return_value = list(providers) + pm.get_provider.side_effect = lambda n: providers[n] + return pm + + def test_discovers_only_media_capable_providers(self): + chat_only = MagicMock(spec=["get_name"]) + m = MediaManager.from_provider_manager( + self._pm({"chat": chat_only, "fake": FakeMediaProvider("fake")}) + ) + assert m.adapter_names == ["fake"] + + def test_broken_provider_is_skipped(self): + pm = MagicMock() + pm.get_available_providers.return_value = ["bad", "fake"] + good = FakeMediaProvider("fake") + pm.get_provider.side_effect = lambda n: (_ for _ in ()).throw(RuntimeError()) if n == "bad" else good + m = MediaManager.from_provider_manager(pm) + assert m.adapter_names == ["fake"] + + def test_config_is_optional(self): + m = MediaManager.from_provider_manager(self._pm({})) + assert m.has_adapters() is False + assert m.artifacts.policy is MaterializePolicy.ON_EXPIRY + + def test_config_drives_artifact_policy(self, tmp_path): + store = {"media.artifact_materialize": "always", "media.artifact_path": str(tmp_path)} + m = MediaManager.from_provider_manager(self._pm({}), lambda k, d=None: store.get(k, d)) + assert m.artifacts.policy is MaterializePolicy.ALWAYS + assert m.artifacts.base_path == tmp_path + + def test_unknown_policy_falls_back(self): + store = {"media.artifact_materialize": "sometimes"} + m = MediaManager.from_provider_manager(self._pm({}), lambda k, d=None: store.get(k, d)) + assert m.artifacts.policy is MaterializePolicy.ON_EXPIRY + + +# --------------------------------------------------------------------------- +# Routers +# --------------------------------------------------------------------------- + + +class TestImageRouter: + async def test_generate_forwards_arguments(self, manager, fake): + r = await manager.images.generate("a tabby", n=3, size="512x512", seed=7) + assert r.capability is MediaCapability.IMAGE_GENERATE + assert len(r.artifacts) == 3 + method, kwargs = fake.calls[-1] + assert method == "generate_image_media" + assert kwargs["prompt"] == "a tabby" and kwargs["n"] == 3 + assert kwargs["size"] == "512x512" and kwargs["seed"] == 7 + + async def test_vendor_kwargs_pass_through(self, manager, fake): + await manager.images.generate("x", guidance_scale=8.5) + assert fake.calls[-1][1]["guidance_scale"] == 8.5 + + async def test_edit_and_upscale(self, manager): + ref = MediaRef.from_bytes(b"img", mime_type="image/png") + assert (await manager.images.edit("night", image=ref)).capability is ( + MediaCapability.IMAGE_EDIT + ) + assert (await manager.images.upscale(image=ref, scale=2)).capability is ( + MediaCapability.IMAGE_UPSCALE + ) + + async def test_ocr_returns_text(self, manager): + r = await manager.images.ocr(document=MediaRef.from_bytes(b"pdf")) + assert r.text == "fake ocr text" + + +class TestAudioRouter: + async def test_speak(self, manager): + r = await manager.audio.speak("hello", voice="rachel") + assert r.artifacts[0].kind is MediaKind.AUDIO + + async def test_stream_tts_yields_chunks(self, manager): + chunks = [c async for c in manager.audio.stream_tts("one two three")] + assert chunks == [b"one", b"two", b"three"] + + async def test_transcribe(self, manager): + r = await manager.audio.transcribe(audio=MediaRef.from_bytes(b"wav")) + assert r.text == "fake transcript" + + async def test_open_session(self, manager): + s = await manager.audio.open_transcription_session(model="nova") + assert s == {"session": "fake", "model": "nova"} + + async def test_music_and_sfx(self, manager): + assert (await manager.audio.music("lofi")).capability is MediaCapability.MUSIC + assert (await manager.audio.sfx("thunder")).capability is MediaCapability.SFX + + +class TestVideoRouter: + async def test_generate_returns_tracked_job(self, manager): + job = await manager.video.generate("dunes", duration_seconds=8, with_audio=True) + assert job.capability is MediaCapability.VIDEO_GENERATE + assert manager.jobs.get(job.id) is job + + async def test_edit_and_interpolate_return_jobs(self, manager): + ref = MediaRef.from_url("https://x/in.mp4") + assert (await manager.video.edit("brighter", video=ref)).capability is ( + MediaCapability.VIDEO_EDIT + ) + assert (await manager.video.interpolate(video=ref, target_fps=60)).capability is ( + MediaCapability.VIDEO_INTERPOLATE + ) + + +# --------------------------------------------------------------------------- +# Job lifecycle +# --------------------------------------------------------------------------- + + +class TestJobPolicy: + def test_backoff_grows_and_is_capped(self): + p = JobPolicy(poll_initial_seconds=1, poll_max_seconds=8, jitter=0.0) + assert [p.delay_for(i) for i in range(1, 6)] == [1, 2, 4, 8, 8] + + def test_jitter_stays_in_band(self): + p = JobPolicy(poll_initial_seconds=10, poll_max_seconds=10, jitter=0.5) + assert all(5.0 <= p.delay_for(1) <= 15.0 for _ in range(50)) + + def test_max_is_never_below_initial(self): + p = JobPolicy(poll_initial_seconds=30, poll_max_seconds=1) + assert p.poll_max_seconds == 30 + + def test_backoff_never_overflows(self): + """A long-running job can be polled thousands of times. + + Without an exponent cap, ``2 ** attempt`` stops converting to float and + the whole wait loop dies with OverflowError — which a multi-hour video + job would actually reach. + """ + p = JobPolicy(poll_initial_seconds=0.0, poll_max_seconds=0.0, jitter=0.0) + assert p.delay_for(10_000) == 0.0 + p2 = JobPolicy(poll_initial_seconds=2, poll_max_seconds=30, jitter=0.0) + assert p2.delay_for(5_000) == 30 + + def test_zero_initial_delay_is_allowed(self): + p = JobPolicy(poll_initial_seconds=0, poll_max_seconds=0, jitter=0.0) + assert p.delay_for(1) == 0.0 + + def test_from_config(self): + store = { + "media.jobs.poll_initial_seconds": 5, + "media.jobs.poll_max_seconds": 50, + "media.jobs.job_timeout_seconds": 60, + } + p = JobPolicy.from_config(lambda k, d=None: store.get(k, d)) + assert (p.poll_initial_seconds, p.poll_max_seconds, p.job_timeout_seconds) == (5, 50, 60) + + +class TestJobLifecycle: + async def test_wait_succeeds(self, manager): + job = await manager.video.generate("dunes") + finished = await manager.jobs.wait(job) + assert finished.succeeded and finished.artifacts + + async def test_manager_wait_returns_result(self, manager): + job = await manager.video.generate("dunes") + result = await manager.wait(job) + assert result.capability is MediaCapability.VIDEO_GENERATE + assert result.usage.seconds == 4.0 + + async def test_progress_is_reported_across_polls(self): + fake = FakeMediaProvider("fake", poll_count=4) + m = MediaManager({"fake": fake}, job_policy=FAST) + job = await m.video.generate("x") + job = await m.jobs.poll(job) + assert job.status is MediaJobStatus.RUNNING + assert 0 < job.progress < 1 + + async def test_poll_on_terminal_job_is_a_noop(self, manager, fake): + job = await manager.video.generate("x") + await manager.jobs.wait(job) + before = len(fake.calls) + assert (await manager.jobs.poll(job)) is job + assert len(fake.calls) == before + + async def test_failure_raises(self): + m = MediaManager({"fake": FakeMediaProvider("fake", fail_jobs=True)}, job_policy=FAST) + job = await m.video.generate("x") + with pytest.raises(MediaJobError, match="fake failure"): + await m.jobs.wait(job) + + async def test_failure_can_be_returned_instead_of_raised(self): + m = MediaManager({"fake": FakeMediaProvider("fake", fail_jobs=True)}, job_policy=FAST) + job = await m.video.generate("x") + finished = await m.jobs.wait(job, raise_on_failure=False) + assert finished.status is MediaJobStatus.FAILED + + async def test_timeout_raises_but_keeps_the_job_alive(self): + # never completes within the budget + m = MediaManager({"fake": FakeMediaProvider("fake", poll_count=10_000)}, job_policy=FAST) + job = await m.video.generate("x") + with pytest.raises(MediaJobTimeoutError, match="still live"): + await m.jobs.wait(job, timeout=0.05) + # the handle survives: an expensive generation is not discarded + assert m.jobs.get(job.id) is not None + assert not job.is_terminal + + async def test_immediate_success_without_polling(self): + m = MediaManager({"fake": FakeMediaProvider("fake", poll_count=0)}, job_policy=FAST) + job = await m.video.generate("x") + assert job.succeeded + assert (await m.jobs.wait(job)).succeeded + + async def test_cancel(self, manager): + job = await manager.video.generate("x") + assert (await manager.jobs.cancel(job)).status is MediaJobStatus.CANCELED + + async def test_cancel_terminal_is_noop(self, manager): + job = await manager.video.generate("x") + await manager.jobs.wait(job) + assert (await manager.jobs.cancel(job)).succeeded + + async def test_cancel_all_active(self, manager): + await manager.video.generate("a") + await manager.video.generate("b") + cancelled = await manager.jobs.cancel_all() + assert len(cancelled) == 2 + + async def test_cancel_all_swallows_failures(self, manager, fake): + job = await manager.video.generate("a") + + async def boom(_job): + raise RuntimeError("vendor down") + + fake.cancel_media_job = boom + assert await manager.jobs.cancel_all() == [] + assert not job.is_terminal + + async def test_registry_listing_and_forget(self, manager): + a = await manager.video.generate("a") + b = await manager.video.generate("b") + await manager.jobs.wait(a) + assert {j.id for j in manager.jobs.list()} == {a.id, b.id} + assert [j.id for j in manager.jobs.list(active_only=True)] == [b.id] + manager.jobs.forget(b.id) + assert manager.jobs.get(b.id) is None + + async def test_unknown_provider_cannot_be_polled(self, manager): + job = await manager.video.generate("x") + manager.unregister_adapter("fake") + with pytest.raises(MediaJobError, match="no longer configured"): + await manager.jobs.poll(job) + + async def test_non_polling_provider_is_rejected(self): + class NoPoll(FakeMediaProvider): + poll_media_job = None # type: ignore[assignment] + + adapter = FakeMediaProvider("fake") + m = MediaManager({"fake": adapter}, job_policy=FAST) + job = await m.video.generate("x") + m.register_adapter("fake", MagicMock(spec=["get_name", "media_capabilities", "media_execution"])) + with pytest.raises(MediaJobError, match="does not implement"): + await m.jobs.poll(job) + + async def test_close_leaves_active_jobs_running(self, manager): + job = await manager.video.generate("x") + await manager.close() + assert not job.is_terminal + + +# --------------------------------------------------------------------------- +# Artifact store +# --------------------------------------------------------------------------- + + +class TestArtifactStore: + def test_put_is_content_addressed_and_idempotent(self, tmp_path): + store = ArtifactStore(tmp_path) + cs1, p1 = store.put(b"hello", suffix=".txt") + cs2, p2 = store.put(b"hello", suffix=".txt") + assert cs1 == cs2 and p1 == p2 + assert store.has(cs1, suffix=".txt") + # sharded two levels deep + assert p1.relative_to(tmp_path).parts[:2] == (cs1[:2], cs1[2:4]) + + def test_put_leaves_no_partial_file(self, tmp_path): + store = ArtifactStore(tmp_path) + store.put(b"data") + assert list(tmp_path.rglob("*.part")) == [] + + @pytest.mark.parametrize( + ("policy", "with_expiry", "expected"), + [ + (MaterializePolicy.NEVER, True, False), + (MaterializePolicy.ALWAYS, False, True), + (MaterializePolicy.ON_EXPIRY, True, True), + (MaterializePolicy.ON_EXPIRY, False, False), + ], + ) + def test_should_materialize(self, tmp_path, policy, with_expiry, expected): + store = ArtifactStore(tmp_path, policy=policy) + a = MediaArtifact( + kind=MediaKind.IMAGE, + uri="https://x/y.png", + expires_at=datetime.now(UTC) + timedelta(hours=1) if with_expiry else None, + ) + assert store.should_materialize(a) is expected + + def test_inline_bytes_are_never_refetched(self, tmp_path): + store = ArtifactStore(tmp_path, policy=MaterializePolicy.ALWAYS) + assert store.should_materialize( + MediaArtifact(kind=MediaKind.IMAGE, uri="u", data=b"x") + ) is False + + async def test_materialize_rewrites_uri_and_keeps_source(self, tmp_path): + async def fetch(url: str) -> bytes: + return b"PNG" + url.encode() + + store = ArtifactStore(tmp_path, policy=MaterializePolicy.ALWAYS, fetcher=fetch) + a = MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y.png", mime_type="image/png") + m = await store.materialize(a) + assert m.uri.startswith("file://") and m.data == b"PNGhttps://x/y.png" + assert m.checksum_sha256 and m.expires_at is None + assert m.provider_metadata["source_uri"] == "https://x/y.png" + + async def test_materialize_is_a_noop_under_never(self, tmp_path): + store = ArtifactStore(tmp_path, policy=MaterializePolicy.NEVER, fetcher=None) + a = MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y.png") + assert (await store.materialize(a)) is a + + async def test_materialize_without_fetcher_raises(self, tmp_path): + store = ArtifactStore(tmp_path, policy=MaterializePolicy.ALWAYS) + with pytest.raises(MediaError, match="requires a fetcher"): + await store.materialize(MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y")) + + async def test_materialize_wraps_fetch_failures(self, tmp_path): + async def boom(_url): + raise RuntimeError("404") + + store = ArtifactStore(tmp_path, policy=MaterializePolicy.ALWAYS, fetcher=boom) + with pytest.raises(MediaError, match="Failed to materialize"): + await store.materialize(MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y")) + + async def test_download_uses_inline_bytes(self, tmp_path): + store = ArtifactStore(tmp_path) + dest = tmp_path / "out" / "a.png" + out = await store.download(MediaArtifact(kind=MediaKind.IMAGE, data=b"raw"), dest) + assert out.read_bytes() == b"raw" + + async def test_download_forces_fetch_even_under_never(self, tmp_path): + async def fetch(_url): + return b"fetched" + + store = ArtifactStore(tmp_path, policy=MaterializePolicy.NEVER, fetcher=fetch) + dest = tmp_path / "b.png" + await store.download(MediaArtifact(kind=MediaKind.IMAGE, uri="https://x/y"), dest) + assert dest.read_bytes() == b"fetched" + + def test_gc_without_keep_set_is_a_noop(self, tmp_path): + store = ArtifactStore(tmp_path) + cs, _ = store.put(b"keep") + assert store.gc() == 0 + assert store.has(cs) + + def test_gc_removes_unlisted(self, tmp_path): + store = ArtifactStore(tmp_path) + keep, _ = store.put(b"keep") + drop, _ = store.put(b"drop") + assert store.gc(keep_checksums={keep}) == 1 + assert store.has(keep) and not store.has(drop) + + def test_construction_does_not_touch_the_filesystem(self, tmp_path): + target = tmp_path / "not-created-yet" + ArtifactStore(target) + assert not target.exists() From 1f45a792ed6a7fc7a3d23ad9a108616c4a491bd9 Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Wed, 30 Sep 2026 02:15:05 -0300 Subject: [PATCH 10/11] feat(media): migrate Deepgram behind the audio protocols (M2) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase M2 of docs/MEDIA_SUBSYSTEM_SPEC.md. Deepgram becomes the first real media adapter, and the legacy multimodal result types gain a bridge to MediaArtifact. Deepgram is the reference migration on purpose: it is the only integration that already exercises batch STT, realtime WebSocket STT and a bidirectional voice agent, so it stress-tests the parts of the abstraction most likely to be wrong before any new vendor commits to them. Its twelve provider-specific methods were the evidence that the chat facade had nowhere to put realtime audio. DeepgramProvider now implements MediaCapableProvider, ASRProvider, TTSProvider, StreamingTTSProvider and StreamingASRProvider. It declares exactly the five capabilities it can serve and reports no async-job execution class, because it has none. The new methods delegate rather than duplicate: transcribe_media, synthesize_speech_media, stream_speech_media and open_transcription_session translate MediaRef in and MediaArtifact out, then call the existing implementations. One code path per operation, so the two surfaces cannot drift. Details worth noting: - A remote MediaRef is handed to Deepgram's own transcribe_url path rather than downloaded locally and re-uploaded. - `timestamps` maps onto Deepgram's `utterances`, which is what actually produces per-segment timings. - `voice` folds into `model` because Deepgram encodes the voice in the model id; an explicit model wins. - Options the caller omits are NOT forwarded, so they cannot override the provider's configured [providers.deepgram.*] defaults. All twelve provider-specific methods are untouched and still return the legacy types; the three existing Deepgram suites pass unchanged. Also adds the models_multimodal <-> MediaArtifact bridge (spec §4.3). SpeechResult, TranscriptionResult, OCRResult, GeneratedImage and ImageGenerationResult gain to_artifact()/to_artifacts(), and the two round-trippable ones gain from_artifact(). These types are public API returned by seven providers, so they are bridged rather than replaced. The conversions preserve what matters — audio format to MIME type, diarization segments and timings, revised_prompt, OCR page structure — and decode base64 image payloads to real bytes, since the media layer deals in bytes. Malformed base64 degrades to the URI path instead of raising. Verified live against the Deepgram API through the media routers: TTS produced 146 KB of WAV with character-based usage, that artifact was fed straight back in as an ASR input via MediaRef.from_artifact() and transcribed correctly, and streaming TTS yielded 39 chunks. Offline: 74 new tests; media + all three Deepgram suites 262 passed; full unit suite 5372 passed, 29 skipped. Co-Authored-By: Claude Opus 5 (1M context) --- CHANGELOG.md | 37 ++ docs/MEDIA_SUBSYSTEM_SPEC.md | 6 +- src/llmcore/models_multimodal.py | 174 ++++++++++ src/llmcore/providers/deepgram_provider.py | 208 +++++++++++ tests/media/test_deepgram_media_adapter.py | 384 +++++++++++++++++++++ tests/media/test_media_bridge.py | 172 +++++++++ 6 files changed, 978 insertions(+), 3 deletions(-) create mode 100644 tests/media/test_deepgram_media_adapter.py create mode 100644 tests/media/test_media_bridge.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 57091360..b304249a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,43 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## Unreleased +### Added — Deepgram behind the media protocols (M2) + +- **Deepgram is now a media adapter**, implementing `MediaCapableProvider`, + `ASRProvider`, `TTSProvider`, `StreamingTTSProvider` and + `StreamingASRProvider`. It declares exactly the five capabilities it can + serve (`asr`, `asr_stream`, `tts`, `tts_stream`, `voice_agent`) and no + execution class it does not have — Deepgram is request/response or live + stream only, never an async job. +- Deepgram was chosen as the **reference migration** because it is the only + integration that already exercises batch STT, realtime WebSocket STT *and* a + bidirectional voice agent, so it stress-tests the hard parts of the + abstraction before any new vendor lands. +- **The new methods delegate; they do not duplicate.** `transcribe_media`, + `synthesize_speech_media`, `stream_speech_media` and + `open_transcription_session` translate `MediaRef` in and `MediaArtifact` out, + then call the existing implementations — one code path per operation rather + than two that can drift. A remote `MediaRef` is handed to Deepgram's own + `transcribe_url` path rather than downloaded locally. +- **All twelve provider-specific methods are untouched** and still return the + legacy types; the three existing Deepgram test suites pass unchanged. + +### Added — `models_multimodal` ↔ `MediaArtifact` bridge + +- `SpeechResult`, `TranscriptionResult`, `OCRResult`, `GeneratedImage` and + `ImageGenerationResult` gained `to_artifact()` / `to_artifacts()`, and the two + round-trippable ones gained `from_artifact()`. These types are public API + returned by seven providers, so they are **bridged, not replaced** (spec §4.3). +- Data that must survive the conversion does: audio format → MIME type, + diarization segments and timings, `revised_prompt`, OCR page structure. Base64 + image payloads are decoded to real bytes, because the media layer deals in + bytes; malformed base64 degrades to the URI path instead of raising. + +Live-validated end to end: TTS through the router produced 146 KB of WAV, the +resulting artifact was fed straight back in as an ASR input via +`MediaRef.from_artifact()` and transcribed correctly, and streaming TTS yielded +39 chunks. 74 new tests; full unit suite 5372 passed. + ### Added — media subsystem core (M1) - **`llmcore.media`**, reached through `llm.media`: a sibling subsystem to chat diff --git a/docs/MEDIA_SUBSYSTEM_SPEC.md b/docs/MEDIA_SUBSYSTEM_SPEC.md index 40a7be9e..5c506b11 100644 --- a/docs/MEDIA_SUBSYSTEM_SPEC.md +++ b/docs/MEDIA_SUBSYSTEM_SPEC.md @@ -3,8 +3,8 @@ Generative image, audio and video as a first-class `llmcore` subsystem, plus the provider adapters that sit behind it. -- **Status:** **M1 implemented** (core types, protocols, routers, job manager, - artifact store, config, fake adapter). M2 onward not started — see §5. +- **Status:** **M1 + M2 implemented** (core subsystem; Deepgram migrated behind + the audio protocols). M3 onward not started — see §5. - **Written:** 2026-09-29 - **Primary input:** `/av/data/repos/docs/llmcore/researches/image-audio-video_providers_2026september.md` (the provider survey and priority matrix; this document is the llmcore-side design) @@ -334,7 +334,7 @@ Order follows the research doc's rollout, with llmcore-specific gates. | Phase | Scope | Gate | |---|---|---| | **M1** ✅ | `llmcore.media` core: types, protocols, routers, `MediaJobManager` (poll only), `ArtifactStore`, config section, fake adapter + tests. *Card schema blocks deferred to M2, where the first real adapter needs them.* | Landed 2026-09-30, 134 tests | -| **M2** | Refactor **Deepgram** behind the audio protocols; keep its public methods | Realtime + batch both pass through the abstraction | +| **M2** ✅ | Refactor **Deepgram** behind the audio protocols; keep its public methods. Added the `models_multimodal` ↔ `MediaArtifact` bridge (§4.3). | Landed 2026-09-30, 74 tests. Live: TTS → artifact → ASR round trip, plus streaming TTS | | **M3** | **OpenAI** media: images generate/edit, TTS, ASR, realtime audio | Lowest marginal cost — adapter already exists | | **M4** | **Google** media: Imagen/Nano-Banana images, **Veo** video (async job), native TTS | First true async-job provider; validates §2.8 | | **M5** | **fal** — queue/webhook lifecycle, URL inputs, video, SFX, FILM interpolation | The provider-neutrality test: if the abstraction bends here, fix the abstraction | diff --git a/src/llmcore/models_multimodal.py b/src/llmcore/models_multimodal.py index fe818097..b7e10c63 100644 --- a/src/llmcore/models_multimodal.py +++ b/src/llmcore/models_multimodal.py @@ -8,11 +8,33 @@ from __future__ import annotations +# Used by the media-subsystem bridge at the bottom of this module. +import base64 as _base64 +import hashlib as _hashlib from enum import Enum from typing import Any from pydantic import BaseModel, Field +_AUDIO_MIME_TYPES: dict[str, str] = { + "mp3": "audio/mpeg", + "wav": "audio/wav", + "linear16": "audio/wav", + "opus": "audio/opus", + "flac": "audio/flac", + "aac": "audio/aac", + "pcm": "audio/L16", + "mulaw": "audio/basic", + "alaw": "audio/basic", +} +_IMAGE_MIME_TYPES: dict[str, str] = { + "png": "image/png", + "jpeg": "image/jpeg", + "jpg": "image/jpeg", + "webp": "image/webp", + "gif": "image/gif", +} + # --------------------------------------------------------------------------- # Text-to-Speech (TTS) Result # --------------------------------------------------------------------------- @@ -440,3 +462,155 @@ class TextAnalysisResult(BaseModel): raw: dict[str, Any] = Field( default_factory=dict, description="Full provider response." ) + + +# ============================================================================= +# MEDIA SUBSYSTEM BRIDGE +# ============================================================================= +# +# The types above predate `llmcore.media` and are public API: seven providers +# return them today. Rather than replace them, each gains a conversion to and +# from `llmcore.media.MediaArtifact`, so the legacy provider methods and the +# media routers describe the same asset. See docs/MEDIA_SUBSYSTEM_SPEC.md §4.3. +# +# Imports are local to each method so this module keeps no import-time +# dependency on the media subsystem. + + +def _media_artifact(kind: str, **fields: Any) -> Any: + """Build a :class:`~llmcore.media.MediaArtifact` without a module-level import.""" + from .media.models import MediaArtifact, MediaKind + + return MediaArtifact(kind=MediaKind(kind), **fields) + + +def _speech_to_artifact(self: "SpeechResult") -> Any: + """Represent this speech result as a media artifact. + + Returns: + A :class:`~llmcore.media.MediaArtifact` carrying the audio bytes, its + checksum, and the voice/format in ``provider_metadata``. + """ + mime = _AUDIO_MIME_TYPES.get(self.format.lower(), f"audio/{self.format.lower()}") + return _media_artifact( + "audio", + data=self.audio_data, + mime_type=mime, + duration_seconds=self.duration_seconds, + checksum_sha256=_hashlib.sha256(self.audio_data).hexdigest() + if self.audio_data + else None, + provider_metadata={"voice": self.voice, "format": self.format, **self.metadata}, + ) + + +def _transcription_to_artifact(self: "TranscriptionResult") -> Any: + """Represent this transcript as a text media artifact. + + Segments are preserved in ``provider_metadata`` so diarization and timings + survive the round trip. + """ + return _media_artifact( + "text", + text=self.text, + mime_type="text/plain", + duration_seconds=self.duration_seconds, + provider_metadata={ + "language": self.language, + "segments": [seg.model_dump() for seg in self.segments], + **self.metadata, + }, + ) + + +def _ocr_to_artifact(self: "OCRResult") -> Any: + """Represent this OCR result as a text media artifact.""" + text = "\n".join( + str(page.get("markdown") or page.get("text") or "") for page in self.pages + ).strip() + return _media_artifact( + "text", + text=text, + mime_type="text/markdown", + provider_metadata={ + "pages": self.pages, + "pages_processed": self.pages_processed, + "document_annotation": self.document_annotation, + **self.metadata, + }, + ) + + +def _generated_image_to_artifact(self: "GeneratedImage") -> Any: + """Represent this generated image as a media artifact. + + ``GeneratedImage.data`` is base64 text (the OpenAI ``b64_json`` shape), so it + is decoded into real bytes here — the media layer deals in bytes. + """ + raw: bytes | None = None + if self.data: + try: + raw = _base64.b64decode(self.data) + except Exception: + raw = None + fmt = (self.format or "png").lower() + return _media_artifact( + "image", + data=raw, + uri=self.url, + mime_type=_IMAGE_MIME_TYPES.get(fmt, f"image/{fmt}"), + checksum_sha256=_hashlib.sha256(raw).hexdigest() if raw else None, + provider_metadata={"revised_prompt": self.revised_prompt, "format": fmt}, + ) + + +def _image_result_to_artifacts(self: "ImageGenerationResult") -> list[Any]: + """Represent every produced image as a media artifact.""" + return [img.to_artifact() for img in self.images] + + +def _speech_from_artifact(cls: type["SpeechResult"], artifact: Any, **overrides: Any) -> "SpeechResult": + """Build a :class:`SpeechResult` from a media artifact. + + Raises: + ValueError: If the artifact carries no inline audio bytes (a remote URI + must be materialized first — the legacy type has nowhere to put a URL). + """ + if artifact.data is None: + raise ValueError( + "SpeechResult requires inline audio bytes; materialize the artifact first." + ) + meta = dict(artifact.provider_metadata or {}) + return cls( + audio_data=artifact.data, + format=overrides.get("format") or meta.get("format") or "mp3", + model=overrides.get("model") or meta.get("model") or "unknown", + voice=overrides.get("voice") or meta.get("voice") or "unknown", + duration_seconds=artifact.duration_seconds, + metadata=meta, + ) + + +def _transcription_from_artifact( + cls: type["TranscriptionResult"], artifact: Any, **overrides: Any +) -> "TranscriptionResult": + """Build a :class:`TranscriptionResult` from a media artifact.""" + meta = dict(artifact.provider_metadata or {}) + segments = [TranscriptionSegment(**seg) for seg in meta.get("segments", []) or []] + return cls( + text=artifact.text or "", + language=overrides.get("language") or meta.get("language"), + duration_seconds=artifact.duration_seconds, + segments=segments, + model=overrides.get("model") or meta.get("model") or "unknown", + metadata=meta, + ) + + +SpeechResult.to_artifact = _speech_to_artifact # type: ignore[attr-defined] +SpeechResult.from_artifact = classmethod(_speech_from_artifact) # type: ignore[attr-defined] +TranscriptionResult.to_artifact = _transcription_to_artifact # type: ignore[attr-defined] +TranscriptionResult.from_artifact = classmethod(_transcription_from_artifact) # type: ignore[attr-defined] +OCRResult.to_artifact = _ocr_to_artifact # type: ignore[attr-defined] +GeneratedImage.to_artifact = _generated_image_to_artifact # type: ignore[attr-defined] +ImageGenerationResult.to_artifacts = _image_result_to_artifacts # type: ignore[attr-defined] diff --git a/src/llmcore/providers/deepgram_provider.py b/src/llmcore/providers/deepgram_provider.py index b1a0eb79..174a9c8a 100644 --- a/src/llmcore/providers/deepgram_provider.py +++ b/src/llmcore/providers/deepgram_provider.py @@ -88,6 +88,12 @@ if TYPE_CHECKING: # pragma: no cover - typing only from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator + from ..media.models import ( + MediaCapability, + MediaExecution, + MediaRef, + MediaResult, + ) from ..models import Tool logger = logging.getLogger(__name__) @@ -2159,6 +2165,208 @@ async def get_projects(self) -> dict[str, Any]: # Escape hatch # ------------------------------------------------------------------ + # ================================================================== + # Media subsystem adapter (llmcore.media protocols) + # ================================================================== + # + # Deepgram is the reference migration for the media subsystem (see + # docs/MEDIA_SUBSYSTEM_SPEC.md §4.4): it is the only integration that + # already exercises batch STT, realtime WebSocket STT and a bidirectional + # voice agent, so it validates the hard parts of the abstraction before any + # new vendor lands. + # + # The twelve provider-specific methods above are UNCHANGED and remain the + # full-power surface. These adapters are thin: they translate MediaRef in + # and MediaArtifact out, and delegate to the existing implementations. That + # keeps one code path per operation rather than two that can drift. + + def media_capabilities(self) -> "frozenset[MediaCapability]": + """Capabilities Deepgram can serve with the current configuration. + + Deepgram is speech-only by design: no image, video, music or OCR. + """ + from ..media.models import MediaCapability + + return frozenset( + { + MediaCapability.ASR, + MediaCapability.ASR_STREAM, + MediaCapability.TTS, + MediaCapability.TTS_STREAM, + MediaCapability.VOICE_AGENT, + } + ) + + def media_execution( + self, capability: "MediaCapability", model: str | None = None + ) -> "MediaExecution": + """Return how *capability* completes on Deepgram. + + Everything Deepgram does is either request/response or a live stream — + it has no asynchronous job surface, so ``ASYNC_JOB`` is never returned. + """ + from ..media.models import MediaCapability, MediaExecution + + streaming = { + MediaCapability.ASR_STREAM, + MediaCapability.TTS_STREAM, + MediaCapability.VOICE_AGENT, + } + return MediaExecution.STREAM if capability in streaming else MediaExecution.REQUEST_RESPONSE + + @staticmethod + def _media_ref_to_audio(audio: "MediaRef") -> tuple[bytes | str, dict[str, Any]]: + """Translate a :class:`MediaRef` into ``transcribe_audio`` arguments. + + Deepgram transcribes a remote URL natively, so a remote ref is passed + through as ``url=`` rather than downloaded locally. + + Returns: + ``(audio_data, extra_kwargs)`` for ``transcribe_audio``. + """ + if audio.is_remote: + # transcribe_url path: audio_data is unused but must be well-formed. + return b"", {"url": audio.url} + return audio.read_bytes(), {} + + async def transcribe_media( + self, + *, + audio: "MediaRef", + model: str | None = None, + language: str | None = None, + diarize: bool | None = None, + timestamps: bool | None = None, + **kwargs: Any, + ) -> "MediaResult": + """Transcribe audio, returning a normalized media result. + + Delegates to :meth:`transcribe_audio`; ``timestamps`` maps onto + Deepgram's ``utterances`` parameter, which is what produces per-segment + timings. + """ + from ..media.models import MediaCapability, MediaResult, MediaUsage + + payload, extra = self._media_ref_to_audio(audio) + if diarize is not None: + extra["diarize"] = diarize + if timestamps is not None: + extra["utterances"] = timestamps + + transcript = await self.transcribe_audio( + payload, model=model, language=language, **extra, **kwargs + ) + artifact = transcript.to_artifact() + return MediaResult( + capability=MediaCapability.ASR, + provider=self.get_name(), + model=transcript.model, + artifacts=(artifact,), + usage=MediaUsage( + provider=self.get_name(), + model=transcript.model, + basis="per_audio_minute", + audio_minutes=(transcript.duration_seconds / 60.0) + if transcript.duration_seconds + else None, + seconds=transcript.duration_seconds, + ), + raw=dict(transcript.metadata), + ) + + async def synthesize_speech_media( + self, + text: str, + *, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + speed: float | None = None, + **kwargs: Any, + ) -> "MediaResult": + """Synthesize speech, returning a normalized media result. + + In Deepgram the voice is encoded in the model id (``aura-2-thalia-en``); + :meth:`generate_speech` already reconciles ``voice``/``model``, so both + are forwarded unchanged. + """ + from ..media.models import MediaCapability, MediaResult, MediaUsage + + if sample_rate_hz is not None: + kwargs.setdefault("sample_rate", sample_rate_hz) + speech_kwargs: dict[str, Any] = {} + if voice is not None: + speech_kwargs["voice"] = voice + if audio_format is not None: + speech_kwargs["response_format"] = audio_format + if speed is not None: + speech_kwargs["speed"] = speed + + speech = await self.generate_speech(text, model=model, **speech_kwargs, **kwargs) + artifact = speech.to_artifact() + return MediaResult( + capability=MediaCapability.TTS, + provider=self.get_name(), + model=speech.model, + artifacts=(artifact,), + usage=MediaUsage( + provider=self.get_name(), + model=speech.model, + basis="per_character", + characters=len(text), + seconds=speech.duration_seconds, + ), + raw=dict(speech.metadata), + ) + + def stream_speech_media( + self, + text: str, + *, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> AsyncIterator[bytes]: + """Stream synthesized speech as it is produced. + + Returns the async iterator directly (not a coroutine), matching + :class:`~llmcore.media.protocols.StreamingTTSProvider`. + + Note ``voice`` is folded into ``model`` because Deepgram encodes the + voice in the model id; an explicit ``model`` wins. + """ + return self.stream_speech( + text, + model=model or voice, + response_format=audio_format, + sample_rate=sample_rate_hz, + **kwargs, + ) + + async def open_transcription_session( + self, + *, + model: str | None = None, + language: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> Any: + """Open a realtime, bidirectional transcription session. + + Returns the async context manager from + :meth:`open_transcription_socket`, so the caller keeps full duplex + control (``send_audio`` / ``finalize`` / ``keepalive`` interleaved with + iterating events) — realtime ASR is not expressible as an iterator. + """ + return self.open_transcription_socket( + model=model, language=language, sample_rate=sample_rate_hz, **kwargs + ) + + # ------------------------------------------------------------------ + @property def client(self) -> Any: """The underlying ``AsyncDeepgramClient`` (escape hatch for power users).""" diff --git a/tests/media/test_deepgram_media_adapter.py b/tests/media/test_deepgram_media_adapter.py new file mode 100644 index 00000000..95057608 --- /dev/null +++ b/tests/media/test_deepgram_media_adapter.py @@ -0,0 +1,384 @@ +# tests/media/test_deepgram_media_adapter.py +"""Deepgram as the reference media adapter (spec phase M2). + +Deepgram is the first migration precisely because it already exercises batch +STT, realtime WebSocket STT and a bidirectional voice agent — the hard parts of +the abstraction. These tests assert: + +* it satisfies every capability protocol it declares; +* the new protocol methods *delegate* to the existing implementations rather + than duplicating them, translating ``MediaRef`` in and ``MediaArtifact`` out; +* the twelve provider-specific methods are untouched (backward compatibility); +* ``MediaManager`` discovers and routes to it without naming it in code. + +The Deepgram SDK is mocked throughout — no network. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from llmcore.media import MediaCapability, MediaExecution, MediaManager, MediaRef +from llmcore.media.protocols import ( + CAPABILITY_PROTOCOLS, + ASRProvider, + MediaCapableProvider, + StreamingASRProvider, + StreamingTTSProvider, + TTSProvider, +) +from llmcore.models_multimodal import SpeechResult, TranscriptionResult, TranscriptionSegment +from llmcore.providers.deepgram_provider import deepgram_available + +# The Deepgram SDK is an optional extra (``pip install llmcore[deepgram]``). +# Constructing the provider touches real SDK symbols, so skip the module rather +# than fail when it is absent — mirroring tests/providers/test_deepgram_*.py. +# CI installs ``.[dev,all]``, so these do run there. +pytestmark = pytest.mark.skipif( + not deepgram_available, + reason="deepgram-sdk not installed (optional extra: pip install llmcore[deepgram])", +) + + +@pytest.fixture +def provider(): + """A DeepgramProvider with the SDK client mocked out.""" + with patch("llmcore.providers.deepgram_provider.AsyncDeepgramClient") as client_cls, patch( + "llmcore.providers.deepgram_provider.deepgram_available", True + ): + client_cls.return_value = MagicMock() + from llmcore.providers.deepgram_provider import DeepgramProvider + + return DeepgramProvider({"api_key": "test-key", "_instance_name": "deepgram"}) + + +def _transcript(**kw) -> TranscriptionResult: + return TranscriptionResult( + text=kw.pop("text", "hello there"), + language=kw.pop("language", "en"), + duration_seconds=kw.pop("duration_seconds", 120.0), + model=kw.pop("model", "nova-3"), + segments=kw.pop("segments", [TranscriptionSegment(text="hello", start=0, end=1)]), + metadata=kw.pop("metadata", {"request_id": "r1"}), + ) + + +def _speech(**kw) -> SpeechResult: + return SpeechResult( + audio_data=kw.pop("audio_data", b"ID3audio"), + format=kw.pop("format", "mp3"), + model=kw.pop("model", "aura-2-thalia-en"), + voice=kw.pop("voice", "thalia"), + duration_seconds=kw.pop("duration_seconds", 2.0), + metadata=kw.pop("metadata", {}), + ) + + +# --------------------------------------------------------------------------- +# Protocol conformance +# --------------------------------------------------------------------------- + + +class TestProtocolConformance: + def test_is_media_capable(self, provider): + assert isinstance(provider, MediaCapableProvider) + + @pytest.mark.parametrize( + "protocol", [ASRProvider, TTSProvider, StreamingTTSProvider, StreamingASRProvider] + ) + def test_implements_audio_protocols(self, provider, protocol): + assert isinstance(provider, protocol) + + def test_declares_only_speech_capabilities(self, provider): + assert provider.media_capabilities() == frozenset( + { + MediaCapability.ASR, + MediaCapability.ASR_STREAM, + MediaCapability.TTS, + MediaCapability.TTS_STREAM, + MediaCapability.VOICE_AGENT, + } + ) + + def test_every_declared_capability_is_backed(self, provider): + """A declaration the manager would have to drop is a bug here.""" + unmet = [ + cap.value + for cap in provider.media_capabilities() + if not isinstance(provider, CAPABILITY_PROTOCOLS[cap]) + ] + assert unmet == [] + + def test_declares_no_image_or_video(self, provider): + caps = provider.media_capabilities() + assert MediaCapability.IMAGE_GENERATE not in caps + assert MediaCapability.VIDEO_GENERATE not in caps + assert MediaCapability.OCR not in caps + + @pytest.mark.parametrize( + ("capability", "expected"), + [ + (MediaCapability.ASR, MediaExecution.REQUEST_RESPONSE), + (MediaCapability.TTS, MediaExecution.REQUEST_RESPONSE), + (MediaCapability.ASR_STREAM, MediaExecution.STREAM), + (MediaCapability.TTS_STREAM, MediaExecution.STREAM), + (MediaCapability.VOICE_AGENT, MediaExecution.STREAM), + ], + ) + def test_execution_classes(self, provider, capability, expected): + assert provider.media_execution(capability) is expected + + def test_has_no_async_job_surface(self, provider): + """Deepgram is request/response or live stream only.""" + assert all( + provider.media_execution(c) is not MediaExecution.ASYNC_JOB + for c in provider.media_capabilities() + ) + + +# --------------------------------------------------------------------------- +# transcribe_media +# --------------------------------------------------------------------------- + + +class TestTranscribeMedia: + async def test_delegates_with_inline_bytes(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + result = await provider.transcribe_media( + audio=MediaRef.from_bytes(b"WAVDATA", mime_type="audio/wav") + ) + payload, kwargs = provider.transcribe_audio.call_args[0], provider.transcribe_audio.call_args.kwargs + assert payload[0] == b"WAVDATA" + assert "url" not in kwargs + assert result.capability is MediaCapability.ASR + assert result.text == "hello there" + + async def test_remote_ref_uses_deepgram_url_path(self, provider): + """A remote URL is handed to Deepgram, not downloaded locally.""" + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + await provider.transcribe_media(audio=MediaRef.from_url("https://x/a.wav")) + assert provider.transcribe_audio.call_args.kwargs["url"] == "https://x/a.wav" + + async def test_reads_a_local_path(self, provider, tmp_path): + f = tmp_path / "a.wav" + f.write_bytes(b"FROMFILE") + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + await provider.transcribe_media(audio=MediaRef.from_path(f)) + assert provider.transcribe_audio.call_args[0][0] == b"FROMFILE" + + async def test_maps_diarize_and_timestamps(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + await provider.transcribe_media( + audio=MediaRef.from_bytes(b"x"), diarize=True, timestamps=True + ) + kwargs = provider.transcribe_audio.call_args.kwargs + assert kwargs["diarize"] is True + # timestamps maps onto Deepgram's utterances, which is what produces timings + assert kwargs["utterances"] is True + + async def test_forwards_language_and_vendor_kwargs(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + await provider.transcribe_media( + audio=MediaRef.from_bytes(b"x"), language="pt", smart_format=True + ) + kwargs = provider.transcribe_audio.call_args.kwargs + assert kwargs["language"] == "pt" and kwargs["smart_format"] is True + + async def test_result_carries_artifact_and_usage(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript(duration_seconds=120.0)) + result = await provider.transcribe_media(audio=MediaRef.from_bytes(b"x")) + artifact = result.artifact + assert artifact.text == "hello there" + assert artifact.provider_metadata["language"] == "en" + assert result.usage.audio_minutes == 2.0 # 120s + assert result.usage.basis == "per_audio_minute" + assert result.provider == "deepgram" + + async def test_missing_duration_leaves_minutes_unset(self, provider): + provider.transcribe_audio = AsyncMock( + return_value=_transcript(duration_seconds=None) + ) + result = await provider.transcribe_media(audio=MediaRef.from_bytes(b"x")) + assert result.usage.audio_minutes is None + + +# --------------------------------------------------------------------------- +# TTS +# --------------------------------------------------------------------------- + + +class TestSynthesizeSpeechMedia: + async def test_delegates_and_normalizes(self, provider): + provider.generate_speech = AsyncMock(return_value=_speech()) + result = await provider.synthesize_speech_media("hello world") + assert provider.generate_speech.call_args[0][0] == "hello world" + assert result.capability is MediaCapability.TTS + assert result.artifact.data == b"ID3audio" + assert result.artifact.mime_type == "audio/mpeg" + + async def test_maps_voice_format_and_speed(self, provider): + provider.generate_speech = AsyncMock(return_value=_speech()) + await provider.synthesize_speech_media( + "hi", voice="aura-2-luna-en", audio_format="linear16", speed=1.2 + ) + kwargs = provider.generate_speech.call_args.kwargs + assert kwargs["voice"] == "aura-2-luna-en" + assert kwargs["response_format"] == "linear16" + assert kwargs["speed"] == 1.2 + + async def test_sample_rate_is_forwarded_under_the_vendor_name(self, provider): + provider.generate_speech = AsyncMock(return_value=_speech()) + await provider.synthesize_speech_media("hi", sample_rate_hz=24000) + assert provider.generate_speech.call_args.kwargs["sample_rate"] == 24000 + + async def test_omitted_options_are_not_forced(self, provider): + """Unset options must not override the provider's configured defaults.""" + provider.generate_speech = AsyncMock(return_value=_speech()) + await provider.synthesize_speech_media("hi") + kwargs = provider.generate_speech.call_args.kwargs + for absent in ("voice", "response_format", "speed", "sample_rate"): + assert absent not in kwargs + + async def test_usage_counts_characters(self, provider): + provider.generate_speech = AsyncMock(return_value=_speech(duration_seconds=2.0)) + result = await provider.synthesize_speech_media("hello") + assert result.usage.basis == "per_character" + assert result.usage.characters == 5 + assert result.usage.seconds == 2.0 + + +class TestStreamSpeechMedia: + async def test_returns_an_iterator_not_a_coroutine(self, provider): + async def _chunks(*_a, **_k): + for piece in (b"a", b"b"): + yield piece + + provider.stream_speech = _chunks + stream = provider.stream_speech_media("hi") # no await + assert [c async for c in stream] == [b"a", b"b"] + + async def test_voice_folds_into_model(self, provider): + """Deepgram encodes the voice in the model id.""" + captured: dict[str, Any] = {} + + async def _chunks(text, **kwargs): + captured.update(kwargs) + yield b"" + + provider.stream_speech = _chunks + [c async for c in provider.stream_speech_media("hi", voice="aura-2-luna-en")] + assert captured["model"] == "aura-2-luna-en" + + async def test_explicit_model_beats_voice(self, provider): + captured: dict[str, Any] = {} + + async def _chunks(text, **kwargs): + captured.update(kwargs) + yield b"" + + provider.stream_speech = _chunks + [ + c + async for c in provider.stream_speech_media( + "hi", model="aura-2-thalia-en", voice="aura-2-luna-en" + ) + ] + assert captured["model"] == "aura-2-thalia-en" + + +class TestOpenTranscriptionSession: + async def test_delegates_to_the_socket(self, provider): + sentinel = object() + provider.open_transcription_socket = MagicMock(return_value=sentinel) + session = await provider.open_transcription_session( + model="nova-3", language="en", sample_rate_hz=16000 + ) + assert session is sentinel + kwargs = provider.open_transcription_socket.call_args.kwargs + assert kwargs["model"] == "nova-3" + assert kwargs["language"] == "en" + assert kwargs["sample_rate"] == 16000 + + +# --------------------------------------------------------------------------- +# Backward compatibility +# --------------------------------------------------------------------------- + + +class TestBackwardCompatibility: + LEGACY_METHODS = [ + "transcribe_audio", + "generate_speech", + "stream_speech", + "transcribe_stream", + "transcribe_stream_flux", + "open_transcription_socket", + "open_flux_socket", + "open_speech_socket", + "open_voice_agent", + "run_voice_agent", + "analyze_text", + "grant_token", + "get_projects", + ] + + @pytest.mark.parametrize("name", LEGACY_METHODS) + def test_legacy_method_still_present(self, provider, name): + assert callable(getattr(provider, name)) + + async def test_legacy_return_types_unchanged(self, provider): + """The legacy surface still returns the legacy types, not MediaResult.""" + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + assert isinstance(await provider.transcribe_audio(b"x"), TranscriptionResult) + provider.generate_speech = AsyncMock(return_value=_speech()) + assert isinstance(await provider.generate_speech("x"), SpeechResult) + + def test_chat_completion_is_still_refused(self, provider): + assert callable(provider.chat_completion) + + +# --------------------------------------------------------------------------- +# Routing through MediaManager +# --------------------------------------------------------------------------- + + +class TestRoutingThroughManager: + def test_manager_discovers_deepgram(self, provider): + pm = MagicMock() + pm.get_available_providers.return_value = ["deepgram"] + pm.get_provider.side_effect = lambda _n: provider + media = MediaManager.from_provider_manager(pm, lambda k, d=None: d) + assert media.adapter_names == ["deepgram"] + + def test_asr_routes_to_deepgram_by_default(self, provider): + """The built-in preference puts Deepgram first for production ASR.""" + media = MediaManager({"deepgram": provider}) + assert media.who_can(MediaCapability.ASR) == ["deepgram"] + assert media.resolve(MediaCapability.ASR) is provider + + def test_unsupported_capability_is_refused_with_hints(self, provider): + from llmcore.exceptions import MediaCapabilityError + + media = MediaManager({"deepgram": provider}) + with pytest.raises(MediaCapabilityError) as exc: + media.resolve(MediaCapability.VIDEO_GENERATE) + assert "gemini" in str(exc.value) or "fal" in str(exc.value) + + async def test_router_transcribes_through_the_manager(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + media = MediaManager({"deepgram": provider}) + result = await media.audio.transcribe(audio=MediaRef.from_bytes(b"x")) + assert result.text == "hello there" + assert result.provider == "deepgram" + + async def test_router_streams_tts_through_the_manager(self, provider): + async def _chunks(*_a, **_k): + for piece in (b"x", b"y"): + yield piece + + provider.stream_speech = _chunks + media = MediaManager({"deepgram": provider}) + assert [c async for c in media.audio.stream_tts("hi")] == [b"x", b"y"] diff --git a/tests/media/test_media_bridge.py b/tests/media/test_media_bridge.py new file mode 100644 index 00000000..61a949da --- /dev/null +++ b/tests/media/test_media_bridge.py @@ -0,0 +1,172 @@ +# tests/media/test_media_bridge.py +"""Tests for the legacy-result ↔ MediaArtifact bridge. + +``models_multimodal`` predates ``llmcore.media`` and is public API — seven +providers return those types today. Rather than replacing them, each converts +to and from :class:`~llmcore.media.MediaArtifact` so the legacy provider methods +and the media routers describe the same asset (spec §4.3). These tests pin the +round trips, including the data that must survive them. +""" + +from __future__ import annotations + +import base64 +import hashlib + +import pytest + +from llmcore.media.models import MediaArtifact, MediaKind +from llmcore.models_multimodal import ( + GeneratedImage, + ImageGenerationResult, + OCRResult, + SpeechResult, + TranscriptionResult, + TranscriptionSegment, +) + +AUDIO = b"ID3\x04audio-bytes" + + +class TestSpeechResultBridge: + def _speech(self, **kw) -> SpeechResult: + return SpeechResult( + audio_data=AUDIO, + format=kw.pop("format", "mp3"), + model="aura-2-thalia-en", + voice="thalia", + duration_seconds=1.5, + **kw, + ) + + def test_to_artifact_carries_bytes_and_checksum(self): + a = self._speech().to_artifact() + assert a.kind is MediaKind.AUDIO + assert a.data == AUDIO + assert a.checksum_sha256 == hashlib.sha256(AUDIO).hexdigest() + assert a.duration_seconds == 1.5 + + @pytest.mark.parametrize( + ("fmt", "mime"), + [("mp3", "audio/mpeg"), ("linear16", "audio/wav"), ("opus", "audio/opus"), + ("flac", "audio/flac"), ("mulaw", "audio/basic")], + ) + def test_format_maps_to_mime_type(self, fmt, mime): + assert self._speech(format=fmt).to_artifact().mime_type == mime + + def test_unknown_format_degrades_gracefully(self): + assert self._speech(format="weird").to_artifact().mime_type == "audio/weird" + + def test_voice_and_format_survive_in_metadata(self): + meta = self._speech().to_artifact().provider_metadata + assert meta["voice"] == "thalia" and meta["format"] == "mp3" + + def test_round_trip(self): + back = SpeechResult.from_artifact(self._speech().to_artifact()) + assert back.audio_data == AUDIO + assert back.voice == "thalia" and back.format == "mp3" + assert back.duration_seconds == 1.5 + + def test_from_artifact_requires_inline_bytes(self): + remote = MediaArtifact(kind=MediaKind.AUDIO, uri="https://x/a.mp3") + with pytest.raises(ValueError, match="materialize"): + SpeechResult.from_artifact(remote) + + def test_from_artifact_overrides_win(self): + a = self._speech().to_artifact() + assert SpeechResult.from_artifact(a, voice="other").voice == "other" + + +class TestTranscriptionResultBridge: + def _transcript(self) -> TranscriptionResult: + return TranscriptionResult( + text="hello there", + language="en", + duration_seconds=2.0, + model="nova-3", + segments=[ + TranscriptionSegment(text="hello", start=0.0, end=1.0, speaker="0"), + TranscriptionSegment(text="there", start=1.0, end=2.0, speaker="1"), + ], + ) + + def test_to_artifact_is_text_kind(self): + a = self._transcript().to_artifact() + assert a.kind is MediaKind.TEXT + assert a.text == "hello there" + assert a.mime_type == "text/plain" + assert a.duration_seconds == 2.0 + + def test_diarization_survives(self): + segs = self._transcript().to_artifact().provider_metadata["segments"] + assert [s["speaker"] for s in segs] == ["0", "1"] + + def test_round_trip_preserves_segments(self): + back = TranscriptionResult.from_artifact(self._transcript().to_artifact()) + assert back.text == "hello there" and back.language == "en" + assert len(back.segments) == 2 + assert back.segments[1].end == 2.0 and back.segments[1].speaker == "1" + + def test_round_trip_of_empty_transcript(self): + empty = TranscriptionResult(text="", model="nova-3") + back = TranscriptionResult.from_artifact(empty.to_artifact()) + assert back.text == "" and back.segments == [] + + +class TestImageBridge: + def test_base64_data_is_decoded_to_bytes(self): + img = GeneratedImage(data=base64.b64encode(b"PNGBYTES").decode(), format="png") + a = img.to_artifact() + assert a.data == b"PNGBYTES" + assert a.mime_type == "image/png" + assert a.checksum_sha256 == hashlib.sha256(b"PNGBYTES").hexdigest() + + def test_url_only_image_keeps_uri(self): + a = GeneratedImage(url="https://x/y.jpeg", format="jpeg").to_artifact() + assert a.uri == "https://x/y.jpeg" and a.data is None + assert a.mime_type == "image/jpeg" + + def test_bad_base64_falls_back_to_uri_path(self): + a = GeneratedImage(data="!!!not-base64!!!", url="https://x/y.png").to_artifact() + assert a.data is None and a.uri == "https://x/y.png" + + def test_revised_prompt_is_kept(self): + a = GeneratedImage(url="u", revised_prompt="a cat, oil painting").to_artifact() + assert a.provider_metadata["revised_prompt"] == "a cat, oil painting" + + def test_result_maps_every_image(self): + result = ImageGenerationResult( + images=[ + GeneratedImage(data=base64.b64encode(b"a").decode()), + GeneratedImage(url="https://x/b.png"), + ], + model="gpt-image", + ) + arts = result.to_artifacts() + assert len(arts) == 2 + assert arts[0].data == b"a" and arts[1].uri == "https://x/b.png" + + def test_empty_result(self): + assert ImageGenerationResult(images=[], model="m").to_artifacts() == [] + + +class TestOCRBridge: + def test_markdown_and_text_pages_are_joined(self): + a = OCRResult( + pages=[{"markdown": "# Title"}, {"text": "body"}], + model="mistral-ocr", + pages_processed=2, + ).to_artifact() + assert a.kind is MediaKind.TEXT + assert a.text == "# Title\nbody" + assert a.mime_type == "text/markdown" + assert a.provider_metadata["pages_processed"] == 2 + + def test_pages_are_preserved_verbatim(self): + pages = [{"markdown": "x", "images": [{"id": "img-1"}]}] + a = OCRResult(pages=pages, model="m", pages_processed=1).to_artifact() + assert a.provider_metadata["pages"] == pages + + def test_empty_pages(self): + a = OCRResult(pages=[], model="m", pages_processed=0).to_artifact() + assert a.text == "" From 60e9b9c10cda721f312c56439894919810dcf49c Mon Sep 17 00:00:00 2001 From: Araray Velho Date: Wed, 30 Sep 2026 02:50:25 -0300 Subject: [PATCH 11/11] =?UTF-8?q?feat(media):=20OpenAI=20media=20adapter?= =?UTF-8?q?=20=E2=80=94=20images,=20speech,=20embeddings=20(M3)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase M3 of docs/MEDIA_SUBSYSTEM_SPEC.md. OpenAI becomes a media adapter for image_generate, image_edit, tts, tts_stream and asr, and gains provider-level embeddings. Image generation, TTS and ASR delegate to the existing provider methods. Two things are new: image editing (POST /v1/images/edits, including mask support and remote-ref upload) and streaming TTS via the SDK's streaming response, so the first bytes arrive before synthesis completes. create_embeddings() closes a long-standing gap recorded in the support matrix: OpenAI embeddings were reachable only through the separate [embedding.openai] subsystem, so a caller holding a provider could not embed with it. Sora is deliberately absent. openai 3.1 deprecated the video APIs, so video_generate is not declared and a test asserts it stays that way — frontier video comes from Veo (M4) and fal (M5). Parameters with no OpenAI equivalent (seed, negative_prompt, sample_rate_hz) are dropped with a debug log rather than forwarded, where forwarding would 400. Supplying reference_images routes to the edit endpoint, which is how OpenAI expresses reference-conditioned generation. --- The subclassing hazard this phase surfaced --- DeepInfra, vLLM, Poe and OpenRouter all extend OpenAIProvider, so they inherit the media protocol METHODS without inheriting the endpoints behind them. Left alone, the router would confidently call /v1/images on a vLLM server. Each subclass now declares its own _MEDIA_CAPABILITIES (DeepInfra: image/TTS/ASR; the other three: none), and a test walks the subclass tree asserting every one declares explicitly — so a future subclass cannot silently inherit. MediaManager.from_provider_manager() also now skips providers that implement the protocols but declare no capabilities, so adapter_names keeps meaning "can actually do something" rather than listing providers that route nothing. The M1 test that asserted no shipped provider implements the protocols was updated to the new truth rather than the assertion being relaxed. --- A real bug live validation caught --- transcribe_audio() labelled every raw-bytes upload "audio.wav". OpenAI infers the container format from the upload filename, so mp3 bytes were rejected with "This model does not support the format you provided". This surfaced the first time a TTS artifact was fed straight back in as an ASR input — exactly the chaining the media subsystem makes natural, and something no unit test with a mocked client would have caught. transcribe_audio() gained an optional `filename` parameter defaulting to the previous "audio.wav", so existing callers are unaffected, and the media adapter derives the correct name from the MediaRef's mime type, filename or URL. Regression tested both at the helper and through the chaining path. Verified live: TTS produced 55 KB of mp3, the artifact round-tripped through ASR and transcribed correctly, streaming TTS yielded 4 chunks, and embeddings returned 256-dimension vectors. Offline: 42 new tests; media suite 254 passed; full unit suite 5418 passed, 29 skipped. Co-Authored-By: Claude Opus 5 (1M context) --- CHANGELOG.md | 44 ++ docs/MEDIA_SUBSYSTEM_SPEC.md | 6 +- src/llmcore/media/manager.py | 26 +- src/llmcore/providers/deepinfra_provider.py | 6 + src/llmcore/providers/openai_provider.py | 424 ++++++++++++++++- src/llmcore/providers/openrouter_provider.py | 5 + src/llmcore/providers/poe_provider.py | 6 + src/llmcore/providers/vllm_provider.py | 7 + tests/media/test_media_integration.py | 15 +- tests/media/test_openai_media_adapter.py | 461 +++++++++++++++++++ 10 files changed, 989 insertions(+), 11 deletions(-) create mode 100644 tests/media/test_openai_media_adapter.py diff --git a/CHANGELOG.md b/CHANGELOG.md index b304249a..81c41ac0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,50 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## Unreleased +### Added — OpenAI media adapter (M3) + +- **OpenAI is now a media adapter** for `image_generate`, `image_edit`, `tts`, + `tts_stream` and `asr`. Image generation, TTS and ASR delegate to the existing + provider methods; **image editing (`POST /v1/images/edits`) and streaming TTS + are new**. +- **`create_embeddings()` on the provider** — previously OpenAI embeddings were + reachable only through the separate `[embedding.openai]` subsystem, so a + caller holding a provider could not embed with it. Closes the gap recorded in + the support matrix. +- **Sora is deliberately not offered.** `openai` 3.1 deprecated the video APIs, + so `video_generate` is absent and a test asserts it stays that way. +- Parameters with no OpenAI equivalent (`seed`, `negative_prompt`, + `sample_rate_hz`) are dropped with a debug log rather than forwarded, where + forwarding would 400. Supplying `reference_images` routes to the edit + endpoint, which is how OpenAI expresses reference-conditioned generation. + +### Fixed — audio format was hard-coded on upload + +- **`OpenAIProvider.transcribe_audio()` labelled every raw-bytes upload + `audio.wav`.** OpenAI infers the container format from the upload filename, so + passing mp3 bytes was rejected with *"This model does not support the format + you provided"*. Found by feeding a TTS artifact straight back in as an ASR + input — the exact chaining the media subsystem makes natural. +- `transcribe_audio()` gained an optional `filename` parameter (defaulting to + the previous `"audio.wav"`, so existing callers are unaffected), and the media + adapter derives the right name from the `MediaRef`'s mime type, filename or + URL. + +### Changed — capability-less providers are no longer registered as adapters + +- Four providers subclass `OpenAIProvider` and therefore inherit its media + protocol *methods* — but not the endpoints behind them. Each now declares its + own `_MEDIA_CAPABILITIES` (DeepInfra: image/TTS/ASR; vLLM, Poe and OpenRouter: + none), and a test asserts **every** subclass declares explicitly, so a future + one cannot silently inherit and advertise endpoints that 404. +- `MediaManager.from_provider_manager()` now skips providers that implement the + protocols but declare no capabilities, so `adapter_names` keeps meaning "can + actually do something". + +Live-validated: TTS (55 KB mp3) → artifact → ASR round trip transcribed +correctly, streaming TTS, and 256-dimension embeddings. 42 new tests; full unit +suite 5410 passed. + ### Added — Deepgram behind the media protocols (M2) - **Deepgram is now a media adapter**, implementing `MediaCapableProvider`, diff --git a/docs/MEDIA_SUBSYSTEM_SPEC.md b/docs/MEDIA_SUBSYSTEM_SPEC.md index 5c506b11..3efe95e7 100644 --- a/docs/MEDIA_SUBSYSTEM_SPEC.md +++ b/docs/MEDIA_SUBSYSTEM_SPEC.md @@ -3,8 +3,8 @@ Generative image, audio and video as a first-class `llmcore` subsystem, plus the provider adapters that sit behind it. -- **Status:** **M1 + M2 implemented** (core subsystem; Deepgram migrated behind - the audio protocols). M3 onward not started — see §5. +- **Status:** **M1–M3 implemented** (core subsystem; Deepgram and OpenAI + migrated behind the protocols). M4 onward not started — see §5. - **Written:** 2026-09-29 - **Primary input:** `/av/data/repos/docs/llmcore/researches/image-audio-video_providers_2026september.md` (the provider survey and priority matrix; this document is the llmcore-side design) @@ -335,7 +335,7 @@ Order follows the research doc's rollout, with llmcore-specific gates. |---|---|---| | **M1** ✅ | `llmcore.media` core: types, protocols, routers, `MediaJobManager` (poll only), `ArtifactStore`, config section, fake adapter + tests. *Card schema blocks deferred to M2, where the first real adapter needs them.* | Landed 2026-09-30, 134 tests | | **M2** ✅ | Refactor **Deepgram** behind the audio protocols; keep its public methods. Added the `models_multimodal` ↔ `MediaArtifact` bridge (§4.3). | Landed 2026-09-30, 74 tests. Live: TTS → artifact → ASR round trip, plus streaming TTS | -| **M3** | **OpenAI** media: images generate/edit, TTS, ASR, realtime audio | Lowest marginal cost — adapter already exists | +| **M3** ✅ | **OpenAI** media: images generate/edit, TTS (+streaming), ASR, and provider-level embeddings. Realtime audio deferred to a later phase with Gemini Live. | Landed 2026-09-30, 42 tests. Live: TTS → artifact → ASR round trip, streaming TTS, embeddings | | **M4** | **Google** media: Imagen/Nano-Banana images, **Veo** video (async job), native TTS | First true async-job provider; validates §2.8 | | **M5** | **fal** — queue/webhook lifecycle, URL inputs, video, SFX, FILM interpolation | The provider-neutrality test: if the abstraction bends here, fix the abstraction | | **M6** | **ElevenLabs** — batch + realtime STT, TTS, SFX, music, voice design | Consent/provenance as first-class metadata | diff --git a/src/llmcore/media/manager.py b/src/llmcore/media/manager.py index 4f7c44c4..f3ec8ae7 100644 --- a/src/llmcore/media/manager.py +++ b/src/llmcore/media/manager.py @@ -121,8 +121,22 @@ def from_provider_manager( except Exception as e: logger.debug("Skipping provider '%s' during media discovery: %s", name, e) continue - if isinstance(provider, MediaCapableProvider): - adapters[name] = provider + if not isinstance(provider, MediaCapableProvider): + continue + # Implementing the protocols is not the same as serving anything. + # Four providers subclass OpenAIProvider and inherit its media + # methods without inheriting its endpoints, so they declare an + # empty capability set. Registering them would put a provider in + # `adapter_names` that can route nothing — so skip them here and + # keep the adapter list meaning "can actually do something". + if not _declares_any_capability(provider): + logger.debug( + "Provider '%s' implements the media protocols but declares no " + "capabilities; not registering it as an adapter.", + name, + ) + continue + adapters[name] = provider routing = cls._routing_from_config(get) @@ -358,6 +372,14 @@ async def close(self) -> None: ) +def _declares_any_capability(adapter: Any) -> bool: + """Whether *adapter* declares at least one media capability.""" + try: + return bool(adapter.media_capabilities()) + except Exception: # noqa: BLE001 - a broken adapter is simply not registered + return False + + def _name_of(adapter: Any) -> str: """Best-effort adapter name for log messages.""" try: diff --git a/src/llmcore/providers/deepinfra_provider.py b/src/llmcore/providers/deepinfra_provider.py index 96f4e20b..8167c375 100644 --- a/src/llmcore/providers/deepinfra_provider.py +++ b/src/llmcore/providers/deepinfra_provider.py @@ -170,6 +170,12 @@ class DeepInfraProvider(OpenAIProvider): ``Message.metadata["content_parts"]``. """ + #: DeepInfra overrides generate_speech / transcribe_audio / generate_image / + #: create_embeddings with its own native endpoints, so it serves those. It + #: has no image-edit endpoint and no streaming TTS, so neither is declared. + _MEDIA_CAPABILITIES: frozenset[Any] = frozenset({"image_generate", "tts", "asr"}) + + def __init__(self, config: dict[str, Any], log_raw_payloads: bool = False): """Initialise the DeepInfra provider. diff --git a/src/llmcore/providers/openai_provider.py b/src/llmcore/providers/openai_provider.py index d8a0d620..2777b7c7 100644 --- a/src/llmcore/providers/openai_provider.py +++ b/src/llmcore/providers/openai_provider.py @@ -17,7 +17,7 @@ import logging import os from collections.abc import AsyncGenerator -from typing import Any +from typing import TYPE_CHECKING, Any try: import openai @@ -78,6 +78,11 @@ from ..tokens import get_counter as _get_token_counter from .base import BaseProvider, ContextPayload +if TYPE_CHECKING: # pragma: no cover - typing only + from collections.abc import AsyncIterator, Sequence + + from ..media.models import MediaCapability, MediaExecution, MediaRef, MediaResult + logger = logging.getLogger(__name__) DEFAULT_OPENAI_TOKEN_LIMITS = { @@ -836,6 +841,7 @@ async def transcribe_audio( response_format: str = "json", temperature: float | None = None, timestamp_granularities: list[str] | None = None, + filename: str | None = None, **kwargs: Any, ) -> Any: """Transcribe audio to text using OpenAI Whisper / GPT-4o-transcribe. @@ -872,11 +878,17 @@ async def transcribe_audio( file_obj: Any = open(file_path, "rb") file_name = file_path.name else: - # Raw bytes — wrap in a tuple for the SDK + # Raw bytes — wrap in a tuple for the SDK. The *name* is what tells + # OpenAI the container format: sending mp3 bytes labelled + # "audio.wav" is rejected with + # "This model does not support the format you provided". import io + import mimetypes - file_obj = ("audio.wav", io.BytesIO(audio_data), "audio/wav") - file_name = "audio.wav" + name = filename or "audio.wav" + mime = mimetypes.guess_type(name)[0] or "audio/wav" + file_obj = (name, io.BytesIO(audio_data), mime) + file_name = name api_kwargs: dict[str, Any] = { "file": file_obj, @@ -972,6 +984,387 @@ async def transcribe_audio( # Multimodal: Image Generation # ------------------------------------------------------------------ + + # ================================================================== + # Media subsystem adapter (llmcore.media protocols) + # ================================================================== + # + # Phase M3 of docs/MEDIA_SUBSYSTEM_SPEC.md. The image/speech methods above + # are unchanged; these adapters translate MediaRef in and MediaArtifact out + # and delegate, so there is one code path per operation. + # + # NOTE FOR SUBCLASSES: DeepInfra, vLLM, Poe and OpenRouter all extend this + # class and therefore INHERIT these methods — but not the endpoints behind + # them. vLLM has no image generation; OpenRouter is a chat gateway. Each + # subclass must set ``_MEDIA_CAPABILITIES`` to what IT can actually serve, + # and a test asserts every subclass does so explicitly. Inheriting the + # declaration would make the router confidently call an endpoint that + # returns 404. + + #: Media capabilities this provider class serves. Subclasses MUST override. + #: Deliberately NOT inherited meaningfully: see the note above. + #: Sora video is excluded — openai 3.1 deprecated the video APIs. + _MEDIA_CAPABILITIES: frozenset[Any] = frozenset( + { + "image_generate", + "image_edit", + "tts", + "tts_stream", + "asr", + } + ) + + def media_capabilities(self) -> "frozenset[MediaCapability]": + """Capabilities this provider serves, declared per class.""" + from ..media.models import MediaCapability + + return frozenset(MediaCapability(c) for c in self._MEDIA_CAPABILITIES) + + def media_execution( + self, capability: "MediaCapability", model: str | None = None + ) -> "MediaExecution": + """Return how *capability* completes. + + Every OpenAI media surface is synchronous; only TTS has a stream. + """ + from ..media.models import MediaCapability, MediaExecution + + if capability is MediaCapability.TTS_STREAM: + return MediaExecution.STREAM + return MediaExecution.REQUEST_RESPONSE + + # --- image --- + + async def generate_image_media( + self, + prompt: str, + *, + model: str | None = None, + n: int = 1, + size: str | None = None, + seed: int | None = None, + negative_prompt: str | None = None, + reference_images: "Sequence[MediaRef] | None" = None, + **kwargs: Any, + ) -> "MediaResult": + """Generate images, returning a normalized media result. + + ``seed`` and ``negative_prompt`` have no OpenAI equivalent and are + ignored with a debug log rather than silently dropped. Supplying + ``reference_images`` routes to the *edit* endpoint, which is how + OpenAI expresses reference-conditioned generation. + """ + from ..media.models import MediaCapability, MediaResult, MediaUsage + + if reference_images: + return await self.edit_image_media( + prompt, image=reference_images[0], model=model, n=n, size=size, **kwargs + ) + for unsupported, value in (("seed", seed), ("negative_prompt", negative_prompt)): + if value is not None: + logger.debug("OpenAI image generation ignores '%s'.", unsupported) + + result = await self.generate_image(prompt, model=model, n=n, size=size, **kwargs) + return MediaResult( + capability=MediaCapability.IMAGE_GENERATE, + provider=self.get_name(), + model=result.model, + artifacts=tuple(result.to_artifacts()), + usage=MediaUsage( + provider=self.get_name(), + model=result.model, + basis="per_image", + images=len(result.images), + ), + raw=dict(result.metadata), + ) + + async def edit_image_media( + self, + prompt: str, + *, + image: "MediaRef", + mask: "MediaRef | None" = None, + model: str | None = None, + n: int = 1, + size: str | None = None, + **kwargs: Any, + ) -> "MediaResult": + """Edit an image via ``POST /v1/images/edits``. + + Args: + prompt: What the edit should produce. + image: The image to edit. A remote ref is fetched first, because the + edits endpoint takes an upload rather than a URL. + mask: Optional transparency mask restricting the edited region. + model: Image model; defaults to ``gpt-image-1``. + n: Number of variants. + size: Output dimensions. + **kwargs: Extra OpenAI parameters (``quality``, ``background``, + ``input_fidelity``, ``output_format`` …). + + Returns: + A :class:`~llmcore.media.MediaResult` with the edited images. + + Raises: + ProviderError: If the client is missing or the API call fails. + """ + from ..media.models import MediaCapability, MediaResult, MediaUsage + from ..models_multimodal import GeneratedImage, ImageGenerationResult + + if not self._client: + raise ProviderError(self.get_name(), "Client not initialized.") + + img_model = model or "gpt-image-1" + api_kwargs: dict[str, Any] = { + "prompt": prompt, + "model": img_model, + "n": n, + "image": await self._media_upload_tuple(image, "image.png"), + } + if mask is not None: + api_kwargs["mask"] = await self._media_upload_tuple(mask, "mask.png") + if size is not None: + api_kwargs["size"] = size + api_kwargs.update(kwargs) + + try: + response = await self._client.images.edit(**api_kwargs) + except OpenAIAPIStatusError as e: + raise ProviderError(self.get_name(), f"Image edit error ({e.status_code}): {e}") + except OpenAIError as e: + raise ProviderError(self.get_name(), f"Image edit error: {e}") + + images = [ + GeneratedImage( + data=getattr(item, "b64_json", None), + url=getattr(item, "url", None), + revised_prompt=getattr(item, "revised_prompt", None), + format=kwargs.get("output_format", "png"), + ) + for item in response.data + ] + result = ImageGenerationResult( + images=images, + model=img_model, + metadata={"created": getattr(response, "created", None)}, + ) + return MediaResult( + capability=MediaCapability.IMAGE_EDIT, + provider=self.get_name(), + model=img_model, + artifacts=tuple(result.to_artifacts()), + usage=MediaUsage( + provider=self.get_name(), + model=img_model, + basis="per_image", + images=len(images), + ), + raw=dict(result.metadata), + ) + + async def _media_upload_tuple(self, ref: "MediaRef", fallback_name: str) -> Any: + """Turn a :class:`MediaRef` into an OpenAI file-upload tuple. + + The images endpoints take an upload, not a URL, so a remote ref is + fetched here rather than handed through. + """ + if ref.is_remote: + from ..media.artifacts import default_fetcher + + data = await default_fetcher()(ref.url or "") + else: + data = ref.read_bytes() + return (ref.filename or fallback_name, data, ref.mime_type or "image/png") + + # --- audio --- + + async def synthesize_speech_media( + self, + text: str, + *, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + speed: float | None = None, + **kwargs: Any, + ) -> "MediaResult": + """Synthesize speech, returning a normalized media result.""" + from ..media.models import MediaCapability, MediaResult, MediaUsage + + if sample_rate_hz is not None: + logger.debug("OpenAI TTS has no sample-rate parameter; ignoring it.") + speech_kwargs: dict[str, Any] = {} + if voice is not None: + speech_kwargs["voice"] = voice + if audio_format is not None: + speech_kwargs["response_format"] = audio_format + if speed is not None: + speech_kwargs["speed"] = speed + + speech = await self.generate_speech(text, model=model, **speech_kwargs, **kwargs) + return MediaResult( + capability=MediaCapability.TTS, + provider=self.get_name(), + model=speech.model, + artifacts=(speech.to_artifact(),), + usage=MediaUsage( + provider=self.get_name(), + model=speech.model, + basis="per_character", + characters=len(text), + ), + raw=dict(speech.metadata), + ) + + def stream_speech_media( + self, + text: str, + *, + model: str | None = None, + voice: str | None = None, + audio_format: str | None = None, + sample_rate_hz: int | None = None, + **kwargs: Any, + ) -> "AsyncIterator[bytes]": + """Stream synthesized speech as it is produced. + + Uses the SDK's streaming response so the first bytes arrive before + synthesis completes. Returns the iterator directly, not a coroutine. + """ + client = self._client + provider_name = self.get_name() + + async def _stream() -> "AsyncIterator[bytes]": + if not client: + raise ProviderError(provider_name, "Client not initialized.") + api_kwargs: dict[str, Any] = { + "model": model or "gpt-4o-mini-tts", + "voice": voice or "alloy", + "input": text, + "response_format": audio_format or "mp3", + **kwargs, + } + try: + async with client.audio.speech.with_streaming_response.create( + **api_kwargs + ) as response: + async for chunk in response.iter_bytes(): + yield chunk + except OpenAIAPIStatusError as e: + raise ProviderError(provider_name, f"TTS stream error ({e.status_code}): {e}") + except OpenAIError as e: + raise ProviderError(provider_name, f"TTS stream error: {e}") + + return _stream() + + async def transcribe_media( + self, + *, + audio: "MediaRef", + model: str | None = None, + language: str | None = None, + diarize: bool | None = None, + timestamps: bool | None = None, + **kwargs: Any, + ) -> "MediaResult": + """Transcribe audio, returning a normalized media result. + + ``timestamps`` maps onto ``timestamp_granularities=["segment"]``. + ``diarize`` is only meaningful on the diarizing models, so it is passed + through rather than emulated. + """ + from ..media.models import MediaCapability, MediaResult, MediaUsage + + if audio.is_remote: + from ..media.artifacts import default_fetcher + + payload: bytes = await default_fetcher()(audio.url or "") + else: + payload = audio.read_bytes() + + extra: dict[str, Any] = dict(kwargs) + if timestamps: + extra.setdefault("timestamp_granularities", ["segment"]) + if diarize is not None: + extra.setdefault("diarize", diarize) + + transcript = await self.transcribe_audio( + payload, + model=model, + language=language, + filename=_audio_filename_for(audio), + **extra, + ) + return MediaResult( + capability=MediaCapability.ASR, + provider=self.get_name(), + model=transcript.model, + artifacts=(transcript.to_artifact(),), + usage=MediaUsage( + provider=self.get_name(), + model=transcript.model, + basis="per_audio_minute", + audio_minutes=(transcript.duration_seconds / 60.0) + if transcript.duration_seconds + else None, + seconds=transcript.duration_seconds, + ), + raw=dict(transcript.metadata), + ) + + # --- embeddings --- + + async def create_embeddings( + self, + input_texts: str | list[str], + *, + model: str | None = None, + dimensions: int | None = None, + encoding_format: str | None = None, + **kwargs: Any, + ) -> dict[str, Any]: + """Create text embeddings via ``POST /v1/embeddings``. + + Closes a long-standing gap: OpenAI embeddings were reachable only + through the separate ``[embedding.openai]`` subsystem, so callers + holding a provider could not embed with it. + + Args: + input_texts: A string or list of strings to embed. + model: Embedding model; defaults to ``text-embedding-3-small``. + dimensions: Output dimensionality, where the model supports it. + encoding_format: ``"float"`` or ``"base64"``. + **kwargs: Extra API parameters. + + Returns: + The raw API response dict with ``data``, ``model`` and ``usage``. + + Raises: + ProviderError: If the client is missing or the API call fails. + """ + if not self._client: + raise ProviderError(self.get_name(), "Client not initialized.") + + api_kwargs: dict[str, Any] = { + "model": model or "text-embedding-3-small", + "input": input_texts, + **kwargs, + } + if dimensions is not None: + api_kwargs["dimensions"] = dimensions + if encoding_format is not None: + api_kwargs["encoding_format"] = encoding_format + + try: + response = await self._client.embeddings.create(**api_kwargs) + except OpenAIAPIStatusError as e: + raise ProviderError(self.get_name(), f"Embeddings error ({e.status_code}): {e}") + except OpenAIError as e: + raise ProviderError(self.get_name(), f"Embeddings error: {e}") + return response.model_dump(exclude_none=True) + async def generate_image( self, prompt: str, @@ -1066,3 +1459,26 @@ async def generate_image( except Exception as e: logger.error(f"Unexpected image error: {e}", exc_info=True) raise ProviderError(self.get_name(), f"Image unexpected error: {e}") + + +def _audio_filename_for(ref: "MediaRef") -> str: + """Pick an upload filename whose extension names the real audio format. + + OpenAI infers the container from the filename, so an artifact chained in + from TTS (mp3 bytes) must not be uploaded as ``audio.wav``. + """ + import mimetypes + + if ref.filename: + return ref.filename + if ref.mime_type: + ext = mimetypes.guess_extension(ref.mime_type) + if ext: + return f"audio{ext}" + if ref.url: + from pathlib import Path as _Path + + suffix = _Path(ref.url.split("?", 1)[0]).suffix + if suffix: + return f"audio{suffix}" + return "audio.wav" diff --git a/src/llmcore/providers/openrouter_provider.py b/src/llmcore/providers/openrouter_provider.py index 061566eb..6ae7130c 100644 --- a/src/llmcore/providers/openrouter_provider.py +++ b/src/llmcore/providers/openrouter_provider.py @@ -81,6 +81,11 @@ class OpenRouterProvider(OpenAIProvider): timeout = 120 """ + #: OpenRouter is a chat-completions gateway; it exposes no image or audio + #: endpoints of its own. + _MEDIA_CAPABILITIES: frozenset[Any] = frozenset() + + # OpenRouter-specific state _app_url: str | None = None _app_title: str | None = None diff --git a/src/llmcore/providers/poe_provider.py b/src/llmcore/providers/poe_provider.py index b4c20cb3..938d87c4 100644 --- a/src/llmcore/providers/poe_provider.py +++ b/src/llmcore/providers/poe_provider.py @@ -132,6 +132,12 @@ class PoeProvider(OpenAIProvider): timeout = 120 """ + #: Poe's media bots are reached as chat *bots*, not through /v1/images or + #: /v1/audio, so none of the OpenAI media endpoints exist here. Surfacing + #: them through the media routers is tracked separately. + _MEDIA_CAPABILITIES: frozenset[Any] = frozenset() + + # Poe-specific state _backend: str = "openai" _poe_api_key: str = "" diff --git a/src/llmcore/providers/vllm_provider.py b/src/llmcore/providers/vllm_provider.py index a120abec..d87fbc55 100644 --- a/src/llmcore/providers/vllm_provider.py +++ b/src/llmcore/providers/vllm_provider.py @@ -211,6 +211,13 @@ class VLLMProvider(OpenAIProvider): # type: ignore[misc,valid-type] default_model = "meta-llama/Llama-3.3-70B-Instruct" """ + #: vLLM serves language models over an OpenAI-compatible surface. It has no + #: image, speech or audio endpoints, so it declares nothing — inheriting + #: OpenAI's declaration would make the router call endpoints that 404. + #: (Embeddings/rerank are a tracked gap; see PROVIDER_MODERNIZATION_PLAN §8.) + _MEDIA_CAPABILITIES: frozenset[Any] = frozenset() + + #: Per-model cache of ``max_model_len`` discovered from #: ``/v1/models``. Populated on :meth:`get_models_details` call; #: consulted by :meth:`get_max_context_length`. Process-local; diff --git a/tests/media/test_media_integration.py b/tests/media/test_media_integration.py index 9183d997..527f561b 100644 --- a/tests/media/test_media_integration.py +++ b/tests/media/test_media_integration.py @@ -16,6 +16,7 @@ import llmcore from llmcore.exceptions import ConfigError from llmcore.media import MediaCapability, MediaManager +from llmcore.media.protocols import MediaCapableProvider from llmcore.media.testing import FakeMediaProvider from llmcore.providers.manager import ProviderManager @@ -201,9 +202,19 @@ async def _boom(): class TestManagerFromRealProviderManager: - def test_chat_only_providers_are_not_adapters(self): - """M1 adds the core only; no shipped provider implements the protocols yet.""" + def test_capability_less_providers_are_not_adapters(self): + """Implementing the protocols is not the same as serving anything. + + vLLM subclasses OpenAIProvider, so since M3 it *inherits* the media + protocol methods — but not the endpoints behind them, so it declares an + empty capability set. It must not appear as an adapter that can route + nothing. + """ pm = ProviderManager(_config({"vllm": dict(VLLM)})) + provider = pm.get_provider("vllm") + assert isinstance(provider, MediaCapableProvider) # has the methods + assert provider.media_capabilities() == frozenset() # serves nothing + media = MediaManager.from_provider_manager(pm, lambda k, d=None: d) assert media.adapter_names == [] assert media.has_adapters() is False diff --git a/tests/media/test_openai_media_adapter.py b/tests/media/test_openai_media_adapter.py new file mode 100644 index 00000000..4d785ff6 --- /dev/null +++ b/tests/media/test_openai_media_adapter.py @@ -0,0 +1,461 @@ +# tests/media/test_openai_media_adapter.py +"""OpenAI as a media adapter (spec phase M3): images, speech, embeddings. + +Also guards the subclassing hazard this phase introduced. ``DeepInfraProvider``, +``VLLMProvider``, ``PoeProvider`` and ``OpenRouterProvider`` all extend +``OpenAIProvider``, so they inherit the media protocol *methods* — but not the +endpoints behind them. Each must declare what it can actually serve, or the +router will confidently call an endpoint that 404s. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from llmcore.exceptions import ProviderError +from llmcore.media import MediaCapability, MediaExecution, MediaManager, MediaRef +from llmcore.media.protocols import ( + CAPABILITY_PROTOCOLS, + ASRProvider, + ImageEditProvider, + ImageGenerationProvider, + MediaCapableProvider, + StreamingTTSProvider, + TTSProvider, +) +from llmcore.models_multimodal import ( + GeneratedImage, + ImageGenerationResult, + SpeechResult, + TranscriptionResult, +) + + +@pytest.fixture +def provider(): + """An OpenAIProvider with the SDK client mocked out.""" + with patch("llmcore.providers.openai_provider.AsyncOpenAI") as client_cls: + client_cls.return_value = MagicMock() + from llmcore.providers.openai_provider import OpenAIProvider + + return OpenAIProvider({"api_key": "sk-test", "_instance_name": "openai"}) + + +def _image_result(n: int = 1) -> ImageGenerationResult: + return ImageGenerationResult( + images=[GeneratedImage(url=f"https://x/{i}.png", format="png") for i in range(n)], + model="gpt-image-1", + metadata={"created": 1}, + ) + + +def _speech() -> SpeechResult: + return SpeechResult( + audio_data=b"MP3", format="mp3", model="gpt-4o-mini-tts", voice="alloy" + ) + + +def _transcript(duration: float | None = 60.0) -> TranscriptionResult: + return TranscriptionResult( + text="hello", language="en", duration_seconds=duration, model="whisper-1" + ) + + +# --------------------------------------------------------------------------- +# Conformance +# --------------------------------------------------------------------------- + + +class TestConformance: + def test_is_media_capable(self, provider): + assert isinstance(provider, MediaCapableProvider) + + @pytest.mark.parametrize( + "protocol", + [ImageGenerationProvider, ImageEditProvider, TTSProvider, StreamingTTSProvider, ASRProvider], + ) + def test_implements_protocols(self, provider, protocol): + assert isinstance(provider, protocol) + + def test_declared_capabilities(self, provider): + assert provider.media_capabilities() == frozenset( + { + MediaCapability.IMAGE_GENERATE, + MediaCapability.IMAGE_EDIT, + MediaCapability.TTS, + MediaCapability.TTS_STREAM, + MediaCapability.ASR, + } + ) + + def test_sora_video_is_not_declared(self, provider): + """openai 3.1 deprecated the Sora video APIs; we must not offer them.""" + caps = provider.media_capabilities() + assert MediaCapability.VIDEO_GENERATE not in caps + assert MediaCapability.VIDEO_EDIT not in caps + + def test_every_declared_capability_is_backed(self, provider): + unmet = [ + c.value + for c in provider.media_capabilities() + if not isinstance(provider, CAPABILITY_PROTOCOLS[c]) + ] + assert unmet == [] + + def test_execution_classes(self, provider): + assert provider.media_execution(MediaCapability.TTS_STREAM) is MediaExecution.STREAM + for cap in (MediaCapability.IMAGE_GENERATE, MediaCapability.TTS, MediaCapability.ASR): + assert provider.media_execution(cap) is MediaExecution.REQUEST_RESPONSE + + +# --------------------------------------------------------------------------- +# The subclassing guard +# --------------------------------------------------------------------------- + + +class TestSubclassCapabilityDeclaration: + """Every OpenAIProvider subclass must declare its own capabilities. + + Inheriting OpenAI's declaration would advertise /v1/images and /v1/audio on + providers that do not serve them. + """ + + @staticmethod + def _subclasses() -> list[type]: + from llmcore.providers.openai_provider import OpenAIProvider + + # Import every module that defines one, then walk the tree. + import llmcore.providers.deepinfra_provider # noqa: F401 + import llmcore.providers.openrouter_provider # noqa: F401 + import llmcore.providers.poe_provider # noqa: F401 + import llmcore.providers.vllm_provider # noqa: F401 + + found: list[type] = [] + stack = list(OpenAIProvider.__subclasses__()) + while stack: + cls = stack.pop() + found.append(cls) + stack.extend(cls.__subclasses__()) + return found + + def test_subclasses_were_found(self): + assert len(self._subclasses()) >= 4 + + def test_every_subclass_declares_its_own(self): + offenders = [ + cls.__name__ + for cls in self._subclasses() + if "_MEDIA_CAPABILITIES" not in cls.__dict__ + ] + assert offenders == [], ( + "These OpenAIProvider subclasses inherit OpenAI's media capability " + f"declaration instead of stating their own: {offenders}" + ) + + @pytest.mark.parametrize( + ("name", "expected"), + [ + ("DeepInfraProvider", {"image_generate", "tts", "asr"}), + ("VLLMProvider", set()), + ("PoeProvider", set()), + ("OpenRouterProvider", set()), + ], + ) + def test_subclass_declarations(self, name, expected): + cls = next(c for c in self._subclasses() if c.__name__ == name) + assert set(cls.__dict__["_MEDIA_CAPABILITIES"]) == expected + + def test_declarations_use_valid_capability_names(self): + for cls in self._subclasses(): + for raw in cls.__dict__.get("_MEDIA_CAPABILITIES", frozenset()): + MediaCapability(raw) # raises on a typo + + +# --------------------------------------------------------------------------- +# Images +# --------------------------------------------------------------------------- + + +class TestImageGeneration: + async def test_delegates_and_normalizes(self, provider): + provider.generate_image = AsyncMock(return_value=_image_result(2)) + result = await provider.generate_image_media("a tabby", n=2, size="1024x1024") + assert result.capability is MediaCapability.IMAGE_GENERATE + assert len(result.artifacts) == 2 + assert result.usage.basis == "per_image" and result.usage.images == 2 + kwargs = provider.generate_image.call_args.kwargs + assert kwargs["n"] == 2 and kwargs["size"] == "1024x1024" + + async def test_unsupported_params_are_not_forwarded(self, provider): + """seed / negative_prompt have no OpenAI equivalent.""" + provider.generate_image = AsyncMock(return_value=_image_result()) + await provider.generate_image_media("x", seed=7, negative_prompt="blurry") + kwargs = provider.generate_image.call_args.kwargs + assert "seed" not in kwargs and "negative_prompt" not in kwargs + + async def test_reference_images_route_to_edit(self, provider): + """OpenAI expresses reference-conditioned generation as an edit.""" + provider.edit_image_media = AsyncMock(return_value="edited") + ref = MediaRef.from_bytes(b"png", mime_type="image/png") + out = await provider.generate_image_media("x", reference_images=[ref]) + assert out == "edited" + assert provider.edit_image_media.call_args.kwargs["image"] is ref + + +class TestImageEdit: + async def test_calls_the_edits_endpoint(self, provider): + response = MagicMock() + response.data = [MagicMock(b64_json=None, url="https://x/e.png", revised_prompt=None)] + response.created = 1 + provider._client.images.edit = AsyncMock(return_value=response) + + result = await provider.edit_image_media( + "make it night", + image=MediaRef.from_bytes(b"PNGDATA", mime_type="image/png"), + size="512x512", + ) + assert result.capability is MediaCapability.IMAGE_EDIT + assert result.artifact.uri == "https://x/e.png" + kwargs = provider._client.images.edit.call_args.kwargs + assert kwargs["prompt"] == "make it night" + assert kwargs["size"] == "512x512" + # uploads, not URLs + assert kwargs["image"][1] == b"PNGDATA" + + async def test_mask_is_uploaded(self, provider): + response = MagicMock() + response.data = [] + provider._client.images.edit = AsyncMock(return_value=response) + await provider.edit_image_media( + "x", + image=MediaRef.from_bytes(b"IMG"), + mask=MediaRef.from_bytes(b"MASK"), + ) + assert provider._client.images.edit.call_args.kwargs["mask"][1] == b"MASK" + + async def test_remote_ref_is_fetched(self, provider): + """The edits endpoint takes an upload, so a URL must be materialized.""" + response = MagicMock() + response.data = [] + provider._client.images.edit = AsyncMock(return_value=response) + + async def fake_fetch(url): + return b"DOWNLOADED" + + with patch("llmcore.media.artifacts.default_fetcher", return_value=fake_fetch): + await provider.edit_image_media( + "x", image=MediaRef.from_url("https://x/in.png") + ) + assert provider._client.images.edit.call_args.kwargs["image"][1] == b"DOWNLOADED" + + async def test_api_error_maps_to_provider_error(self, provider): + from llmcore.providers.openai_provider import OpenAIError + + provider._client.images.edit = AsyncMock(side_effect=OpenAIError("nope")) + with pytest.raises(ProviderError, match="Image edit error"): + await provider.edit_image_media("x", image=MediaRef.from_bytes(b"i")) + + +# --------------------------------------------------------------------------- +# Audio +# --------------------------------------------------------------------------- + + +class TestSpeech: + async def test_synthesize_delegates(self, provider): + provider.generate_speech = AsyncMock(return_value=_speech()) + result = await provider.synthesize_speech_media("hello", voice="nova", speed=1.1) + assert result.capability is MediaCapability.TTS + assert result.artifact.data == b"MP3" + assert result.usage.characters == 5 + kwargs = provider.generate_speech.call_args.kwargs + assert kwargs["voice"] == "nova" and kwargs["speed"] == 1.1 + + async def test_sample_rate_is_ignored_not_forwarded(self, provider): + """OpenAI TTS has no sample-rate parameter; forwarding it would 400.""" + provider.generate_speech = AsyncMock(return_value=_speech()) + await provider.synthesize_speech_media("hi", sample_rate_hz=24000) + assert "sample_rate_hz" not in provider.generate_speech.call_args.kwargs + assert "sample_rate" not in provider.generate_speech.call_args.kwargs + + async def test_omitted_options_are_not_forced(self, provider): + provider.generate_speech = AsyncMock(return_value=_speech()) + await provider.synthesize_speech_media("hi") + kwargs = provider.generate_speech.call_args.kwargs + for absent in ("voice", "response_format", "speed"): + assert absent not in kwargs + + async def test_stream_uses_streaming_response(self, provider): + class _Resp: + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + return False + + async def iter_bytes(self): + for chunk in (b"a", b"b"): + yield chunk + + provider._client.audio.speech.with_streaming_response.create = MagicMock( + return_value=_Resp() + ) + stream = provider.stream_speech_media("hi", voice="nova") # no await + assert [c async for c in stream] == [b"a", b"b"] + kwargs = provider._client.audio.speech.with_streaming_response.create.call_args.kwargs + assert kwargs["voice"] == "nova" and kwargs["input"] == "hi" + + +class TestTranscription: + async def test_delegates_with_bytes(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + result = await provider.transcribe_media(audio=MediaRef.from_bytes(b"WAV")) + assert provider.transcribe_audio.call_args[0][0] == b"WAV" + assert result.text == "hello" + assert result.usage.audio_minutes == 1.0 + + async def test_timestamps_map_to_granularities(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + await provider.transcribe_media(audio=MediaRef.from_bytes(b"x"), timestamps=True) + assert provider.transcribe_audio.call_args.kwargs["timestamp_granularities"] == [ + "segment" + ] + + async def test_remote_audio_is_fetched(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + + async def fake_fetch(url): + return b"REMOTE" + + with patch("llmcore.media.artifacts.default_fetcher", return_value=fake_fetch): + await provider.transcribe_media(audio=MediaRef.from_url("https://x/a.wav")) + assert provider.transcribe_audio.call_args[0][0] == b"REMOTE" + + async def test_artifact_chaining_names_the_real_format(self, provider): + """Regression: chaining TTS output into ASR must not claim it is wav. + + OpenAI infers the container from the upload filename, so mp3 bytes + uploaded as ``audio.wav`` are rejected with + "This model does not support the format you provided" — which is + exactly what happened the first time a TTS artifact was fed back in. + """ + from llmcore.media.models import MediaArtifact, MediaKind + + provider.transcribe_audio = AsyncMock(return_value=_transcript()) + mp3 = MediaArtifact(kind=MediaKind.AUDIO, data=b"ID3", mime_type="audio/mpeg") + await provider.transcribe_media(audio=MediaRef.from_artifact(mp3)) + assert provider.transcribe_audio.call_args.kwargs["filename"] == "audio.mp3" + + @pytest.mark.parametrize( + ("ref", "expected"), + [ + (MediaRef.from_bytes(b"x", mime_type="audio/mpeg"), "audio.mp3"), + (MediaRef.from_bytes(b"x", mime_type="audio/wav"), "audio.wav"), + (MediaRef.from_bytes(b"x", filename="clip.flac"), "clip.flac"), + (MediaRef.from_url("https://x/a.ogg?t=1"), "audio.oga"), + (MediaRef.from_bytes(b"x"), "audio.wav"), + ], + ) + def test_filename_inference(self, ref, expected): + from llmcore.providers.openai_provider import _audio_filename_for + + assert _audio_filename_for(ref) == expected + + async def test_transcribe_audio_honours_filename(self, provider): + """The legacy method now labels raw bytes correctly too.""" + captured: dict[str, Any] = {} + + async def _create(**kwargs): + captured.update(kwargs) + return MagicMock(text="hi", language=None, duration=None, segments=[]) + + provider._client.audio.transcriptions.create = _create + await provider.transcribe_audio(b"ID3", filename="clip.mp3") + name, _stream, mime = captured["file"] + assert name == "clip.mp3" and mime == "audio/mpeg" + + async def test_transcribe_audio_default_filename_is_unchanged(self, provider): + """Existing callers keep the previous behaviour.""" + captured: dict[str, Any] = {} + + async def _create(**kwargs): + captured.update(kwargs) + return MagicMock(text="hi", language=None, duration=None, segments=[]) + + provider._client.audio.transcriptions.create = _create + await provider.transcribe_audio(b"RIFF") + assert captured["file"][0] == "audio.wav" + + async def test_missing_duration_leaves_minutes_unset(self, provider): + provider.transcribe_audio = AsyncMock(return_value=_transcript(duration=None)) + result = await provider.transcribe_media(audio=MediaRef.from_bytes(b"x")) + assert result.usage.audio_minutes is None + + +# --------------------------------------------------------------------------- +# Embeddings +# --------------------------------------------------------------------------- + + +class TestEmbeddings: + async def test_create_embeddings(self, provider): + response = MagicMock() + response.model_dump.return_value = {"data": [{"embedding": [0.1]}], "model": "m"} + provider._client.embeddings.create = AsyncMock(return_value=response) + out = await provider.create_embeddings(["hello"], dimensions=256) + assert out["data"][0]["embedding"] == [0.1] + kwargs = provider._client.embeddings.create.call_args.kwargs + assert kwargs["input"] == ["hello"] and kwargs["dimensions"] == 256 + + async def test_default_model(self, provider): + response = MagicMock() + response.model_dump.return_value = {} + provider._client.embeddings.create = AsyncMock(return_value=response) + await provider.create_embeddings("hi") + assert ( + provider._client.embeddings.create.call_args.kwargs["model"] + == "text-embedding-3-small" + ) + + async def test_error_maps(self, provider): + from llmcore.providers.openai_provider import OpenAIError + + provider._client.embeddings.create = AsyncMock(side_effect=OpenAIError("bad")) + with pytest.raises(ProviderError, match="Embeddings error"): + await provider.create_embeddings("hi") + + +# --------------------------------------------------------------------------- +# Routing +# --------------------------------------------------------------------------- + + +class TestRouting: + async def test_image_generate_routes_to_openai(self, provider): + media = MediaManager({"openai": provider}) + assert media.who_can(MediaCapability.IMAGE_GENERATE) == ["openai"] + provider.generate_image = AsyncMock(return_value=_image_result()) + result = await media.images.generate("a tabby") + assert result.provider == "openai" + + async def test_edit_routes_through_the_router(self, provider): + response = MagicMock() + response.data = [] + provider._client.images.edit = AsyncMock(return_value=response) + media = MediaManager({"openai": provider}) + out = await media.images.edit("night", image=MediaRef.from_bytes(b"i")) + assert out.capability is MediaCapability.IMAGE_EDIT + + def test_a_chat_only_subclass_is_not_an_adapter(self): + """An OpenRouter instance must not advertise image or audio.""" + with patch("llmcore.providers.openai_provider.AsyncOpenAI") as c: + c.return_value = MagicMock() + from llmcore.providers.openrouter_provider import OpenRouterProvider + + router_provider = OpenRouterProvider({"api_key": "k"}) + media = MediaManager({"openrouter": router_provider}) + assert media.who_can(MediaCapability.IMAGE_GENERATE) == [] + assert media.capabilities() == {}