feat(providers): Support DeepInfra as an LLM provider

This commit is contained in:
Georgi Atsev
2026-04-06 12:22:42 +03:00
committed by kshitij
parent ed8ce1f96c
commit fe002eb124
34 changed files with 2499 additions and 84 deletions
+236
View File
@@ -9,6 +9,7 @@ from __future__ import annotations
import json
import os
import re
import urllib.parse
import urllib.request
import urllib.error
@@ -516,6 +517,12 @@ _PROVIDER_MODELS: dict[str, list[str]] = {
"glm-4.7",
"MiniMax-M2.5",
],
# DeepInfra: empty by design. The live catalog at
# _fetch_deepinfra_models() (filtered by the ``chat`` surface tag) is
# the only source of truth. Hardcoding ids here would rot as models
# are deprecated upstream; the picker shows "no models" when the
# catalog is unreachable, which is honest.
"deepinfra": [],
# Curated HF model list — only agentic models that map to OpenRouter defaults.
"huggingface": [
"moonshotai/Kimi-K2.5",
@@ -1087,6 +1094,7 @@ CANONICAL_PROVIDERS: list[ProviderEntry] = [
ProviderEntry("bedrock", "AWS Bedrock", "AWS Bedrock (Claude, Nova, Llama, DeepSeek; IAM or API key)"),
ProviderEntry("azure-foundry", "Azure Foundry", "Azure Foundry (OpenAI-style or Anthropic-style endpoint, your Azure AI deployment)"),
ProviderEntry("qwen-oauth", "Qwen OAuth (Portal)", "Qwen OAuth (Reuses local Qwen CLI login)"),
ProviderEntry("deepinfra", "DeepInfra", "DeepInfra (100+ top models)"),
]
# Auto-extend CANONICAL_PROVIDERS with any provider registered in providers/
@@ -1297,6 +1305,7 @@ _PROVIDER_ALIASES = {
"lm_studio": "lmstudio",
"ollama": "custom", # bare "ollama" = local; use "ollama-cloud" for cloud
"ollama_cloud": "ollama-cloud",
"deep-infra": "deepinfra",
}
@@ -1593,6 +1602,8 @@ def get_pricing_for_provider(provider: str, *, force_refresh: bool = False) -> d
)
if normalized == "novita":
return _fetch_novita_pricing(force_refresh=force_refresh)
if normalized == "deepinfra":
return _fetch_deepinfra_pricing(force_refresh=force_refresh)
if normalized == "nous":
api_key, base_url = _resolve_nous_pricing_credentials()
if base_url:
@@ -2385,6 +2396,10 @@ def provider_model_ids(provider: Optional[str], *, force_refresh: bool = False)
merged_lower.add(m.lower())
return merged
return list(_PROVIDER_MODELS.get("anthropic", []))
if normalized == "deepinfra":
live = _fetch_deepinfra_models()
if live:
return live
if normalized == "ollama-cloud":
live = fetch_ollama_cloud_models(force_refresh=force_refresh)
if live:
@@ -3626,6 +3641,227 @@ def probe_api_models(
}
# Legacy filter — used when an item has no surface tag (rolling out
# 2026-05). Once every model returned by the catalog endpoint carries an
# explicit surface tag (``chat``/``embed``/``image-gen``/``tts``/``stt``)
# the regex path becomes unreachable and can be removed.
_DEEPINFRA_EXCLUDE_RE = re.compile(
r"(?i)(embed|rerank|whisper|stable-diffusion|flux|sdxl|"
r"tts|bark|speech|image-gen|clip|vit-|dpt-)",
)
# Surface tags announce *what kind of model* this is. When none of these
# are present on a catalog entry, the tags array only carries capability
# tags (``reasoning``, ``vision``, ``prompt_cache``, …) and we have to
# fall back to id-regex inference for the chat surface.
_DEEPINFRA_SURFACE_TAGS: frozenset[str] = frozenset({
"chat", "embed", "image-gen", "tts", "stt", "video-gen",
})
_DEEPINFRA_DEFAULT_BASE_URL = "https://api.deepinfra.com/v1/openai"
_DEEPINFRA_MODELS_QUERY = "filter=true&sort_by=hermes"
# Module-level cache for the full tagged catalog response, keyed by base URL.
# Each value is the parsed ``data`` list. Surface-specific filters read from
# this cache so a single network round-trip serves chat / image-gen / tts /
# stt callers across the whole process lifetime.
_deepinfra_catalog_cache: dict[str, list[dict]] = {}
# Negative cache: monotonic timestamp of the last failed fetch, keyed by base
# URL. Without this, an unreachable catalog (offline / DNS / firewall) makes
# every surface helper (chat picker, pricing, image/video/tts/stt defaults,
# vision) re-attempt a fresh blocking fetch that eats the full timeout each
# time — several sequential stalls in one user-visible operation. A short TTL
# lets connectivity recover without a process restart.
_deepinfra_catalog_neg_cache: dict[str, float] = {}
_DEEPINFRA_CATALOG_NEG_TTL = 60.0 # seconds
def _deepinfra_catalog_url() -> tuple[str, str]:
"""Return ``(cache_key, full_url)`` for the DeepInfra catalog endpoint."""
base = os.getenv("DEEPINFRA_BASE_URL", "").strip() or _DEEPINFRA_DEFAULT_BASE_URL
cache_key = base.rstrip("/")
return cache_key, f"{cache_key}/models?{_DEEPINFRA_MODELS_QUERY}"
def _fetch_deepinfra_catalog(
*,
timeout: float = 5.0,
force_refresh: bool = False,
) -> Optional[list[dict]]:
"""Fetch the raw DeepInfra catalog list with module-level caching.
The endpoint serves chat + embed + image-gen + tts + stt models in one
response. Authentication is optional but Bearer-attached when available
so user-scoped catalogs (private fine-tunes etc.) are visible.
"""
cache_key, url = _deepinfra_catalog_url()
if not force_refresh:
if cache_key in _deepinfra_catalog_cache:
return _deepinfra_catalog_cache[cache_key]
last_fail = _deepinfra_catalog_neg_cache.get(cache_key)
if last_fail is not None and (time.monotonic() - last_fail) < _DEEPINFRA_CATALOG_NEG_TTL:
return None
headers: dict[str, str] = {"User-Agent": _HERMES_USER_AGENT}
api_key = os.getenv("DEEPINFRA_API_KEY", "").strip()
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
req = urllib.request.Request(url, headers=headers)
try:
with urllib.request.urlopen(req, timeout=timeout) as resp:
payload = json.loads(resp.read().decode())
except Exception:
_deepinfra_catalog_neg_cache[cache_key] = time.monotonic()
return None
data = payload.get("data")
if not isinstance(data, list):
_deepinfra_catalog_neg_cache[cache_key] = time.monotonic()
return None
_deepinfra_catalog_cache[cache_key] = data
_deepinfra_catalog_neg_cache.pop(cache_key, None)
return data
def _fetch_deepinfra_models_by_tag(
tag: str,
*,
timeout: float = 5.0,
force_refresh: bool = False,
) -> Optional[list[dict]]:
"""Return DeepInfra models whose ``metadata.tags`` includes *tag*.
Each returned item is ``{"id": str, "metadata": dict}`` so callers can
inspect context length, pricing, default dimensions (image-gen),
pricing units (tts ``input_characters``, stt ``input_seconds``), etc.
For the chat surface, items without any ``tags`` field fall through
to the legacy name-regex exclusion so this keeps working while the
tag rollout (mid-2026) is still in flight.
Returns ``None`` on network failure.
"""
data = _fetch_deepinfra_catalog(timeout=timeout, force_refresh=force_refresh)
if data is None:
return None
matched: list[dict] = []
for item in data:
mid = item.get("id")
if not mid:
continue
# ``metadata is None`` means DeepInfra returns a stub without
# pricing/context — typically a model that's listed but not
# served. Skip those for every surface.
raw_metadata = item.get("metadata")
if raw_metadata is None:
continue
metadata = raw_metadata if isinstance(raw_metadata, dict) else {}
raw_tags = metadata.get("tags")
tags = raw_tags if isinstance(raw_tags, list) else []
has_surface_tag = any(t in _DEEPINFRA_SURFACE_TAGS for t in tags)
if has_surface_tag:
if tag in tags:
matched.append({"id": mid, "metadata": metadata})
continue
# Surface-tag rollout incomplete — fall back to id-regex inference.
# Only meaningful for the chat surface; embed/image-gen/tts/stt
# cannot be safely inferred from an id alone.
if tag == "chat" and not _DEEPINFRA_EXCLUDE_RE.search(mid):
matched.append({"id": mid, "metadata": metadata})
return matched
def _fetch_deepinfra_models(
timeout: float = 5.0,
*,
force_refresh: bool = False,
) -> Optional[list[str]]:
"""Return DeepInfra chat-model ids (tag-aware, regex fallback).
Thin wrapper over :func:`_fetch_deepinfra_models_by_tag` so historical
callers in :func:`provider_model_ids` keep their string-list contract.
Returns ``None`` on network failure, an empty list if the catalog
contains no chat-tagged ids (which would itself be surprising).
"""
items = _fetch_deepinfra_models_by_tag(
"chat", timeout=timeout, force_refresh=force_refresh
)
if items is None:
return None
return [item["id"] for item in items] or None
def deepinfra_model_ids(tag: str, *, force_refresh: bool = False) -> list[str]:
"""Return DeepInfra model ids carrying surface *tag* (``[]`` on failure).
Single source of truth for the per-surface model shims (TTS/STT/vision),
replacing the copy-pasted ``import _fetch_deepinfra_models_by_tag → fetch
→ [item["id"] …]`` wrapper each of them used to carry.
"""
items = _fetch_deepinfra_models_by_tag(tag, force_refresh=force_refresh)
return [item["id"] for item in items] if items else []
def deepinfra_base_url(section: Optional[dict] = None) -> str:
"""Resolve the DeepInfra OpenAI-compatible base URL, normalized.
Precedence: config-section ``base_url`` → ``DEEPINFRA_BASE_URL`` env →
default. Always stripped with any trailing slash removed. Single source
of truth for the base-URL chain the TTS/STT/image/video shims each used
to re-code (with subtly divergent normalization).
"""
candidate = section.get("base_url") if isinstance(section, dict) else None
value = candidate or os.getenv("DEEPINFRA_BASE_URL") or _DEEPINFRA_DEFAULT_BASE_URL
return str(value).strip().rstrip("/")
def _fetch_deepinfra_pricing(
timeout: float = 5.0,
*,
force_refresh: bool = False,
) -> dict[str, dict[str, str]]:
"""Return picker-shape pricing for DeepInfra chat models.
DeepInfra publishes ``input_tokens`` / ``output_tokens`` /
``cache_read_tokens`` in $/MTok; the picker expects per-token strings
under ``prompt`` / ``completion`` / ``input_cache_read`` (mirrors the
OpenRouter shape consumed by
:func:`format_model_pricing_table`). Cached via the catalog helper so
repeated picker renders are free.
"""
items = _fetch_deepinfra_models_by_tag(
"chat", timeout=timeout, force_refresh=force_refresh
)
if not items:
return {}
result: dict[str, dict[str, str]] = {}
for item in items:
metadata = item.get("metadata") or {}
pricing = metadata.get("pricing") if isinstance(metadata, dict) else None
if not isinstance(pricing, dict):
continue
entry: dict[str, str] = {}
inp = pricing.get("input_tokens")
out = pricing.get("output_tokens")
cache_read = pricing.get("cache_read_tokens")
if inp is not None:
entry["prompt"] = str(float(inp) / 1_000_000)
if out is not None:
entry["completion"] = str(float(out) / 1_000_000)
if cache_read is not None:
entry["input_cache_read"] = str(float(cache_read) / 1_000_000)
if entry:
result[item["id"]] = entry
return result
def fetch_api_models(
api_key: Optional[str],
base_url: Optional[str],