diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f5bfb9ff..54c360f6 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -59,14 +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]: 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. - # 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 ea1b7fb6..81c41ac0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,319 @@ 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 — 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`, + `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 + 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 + `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 + `>=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 + 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..961ffe5f 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,11 @@ 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) +- [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/CONFIG_REFERENCE.md b/docs/CONFIG_REFERENCE.md index 1ed7e750..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. @@ -34,6 +64,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 +740,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/MEDIA_SUBSYSTEM_SPEC.md b/docs/MEDIA_SUBSYSTEM_SPEC.md new file mode 100644 index 00000000..3efe95e7 --- /dev/null +++ b/docs/MEDIA_SUBSYSTEM_SPEC.md @@ -0,0 +1,444 @@ +# Media Subsystem — Design & Specification + +Generative image, audio and video as a first-class `llmcore` subsystem, plus the +provider adapters that sit behind it. + +- **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) +- **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`, 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 (+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 | +| **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 new file mode 100644 index 00000000..a2f8fc6e --- /dev/null +++ b/docs/PROVIDER_MODERNIZATION_PLAN.md @@ -0,0 +1,349 @@ +# 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? + +--- + +## 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: + +```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..2d219728 --- /dev/null +++ b/docs/PROVIDER_SUPPORT_MATRIX.md @@ -0,0 +1,332 @@ +# 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 +- **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) + +--- + +## 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 | `>=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 | 🟡 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 (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 | +| 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 | +| 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** | ✅ 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 +`httpx` transitively. Under `openai>=3` they fail at import: + +| Provider | Imports `httpx` | Own extra | +|---|:---:|:---:| +| `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.) + +--- + +## 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 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 | +| 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 | ✅ 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) | + +--- + +## 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. 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) 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()) diff --git a/pyproject.toml b/pyproject.toml index 91d74c50..530fca76 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,23 +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] +# 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 @@ -179,7 +217,14 @@ all = [ "llmcore[deepinfra]", "llmcore[deepgram]", "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/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 64eeab86..8919c713 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] @@ -1117,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..f3ec8ae7 --- /dev/null +++ b/src/llmcore/media/manager.py @@ -0,0 +1,409 @@ +# 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 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) + + 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 _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: + 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/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 + } + } +} 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/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/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/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..f9727aa4 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. @@ -137,6 +145,7 @@ class ProviderManager: """ _providers: dict[str, BaseProvider] + _ephemeral_instances: set[str] _config: ConfyConfig _default_provider_name: str _event_logger: Any | None @@ -167,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 @@ -707,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/src/llmcore/providers/openai_provider.py b/src/llmcore/providers/openai_provider.py index e79a376d..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 = { @@ -554,10 +559,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. @@ -837,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. @@ -873,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, @@ -973,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, @@ -1067,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/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/media/__init__.py b/tests/media/__init__.py new file mode 100644 index 00000000..e69de29b 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 == "" diff --git a/tests/media/test_media_integration.py b/tests/media/test_media_integration.py new file mode 100644 index 00000000..527f561b --- /dev/null +++ b/tests/media/test_media_integration.py @@ -0,0 +1,235 @@ +# 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.protocols import MediaCapableProvider +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_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 + + 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() 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() == {} 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) 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/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: 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"] 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",