Source code for axio_transport_openai.openrouter

"""OpenRouter CompletionTransport - inherits from OpenAI-compatible transport."""

from __future__ import annotations

import logging
import os
from dataclasses import dataclass, field
from typing import Any, Literal

from axio.exceptions import StreamError
from axio.models import Capability, ModelSpec

from axio_transport_openai import OpenAITransport, ThinkingMixin

logger = logging.getLogger(__name__)


[docs] @dataclass(slots=True) class OpenRouterTransport(ThinkingMixin, OpenAITransport): # These point at servers that implement /v1/chat/completions and not /v1/responses. api: Literal["responses", "chat"] = field(default="chat", kw_only=True) name: str = "OpenRouter" api_key: str = field(default_factory=lambda: os.environ.get("OPENROUTER_API_KEY", "")) base_url: str = "https://openrouter.ai/api/v1" model: ModelSpec = ModelSpec(id="google/gemini-2.5-pro-preview") thinking: bool = False
[docs] async def fetch_models(self) -> None: """Fetch available models from OpenRouter ``/v1/models``.""" assert self.session is not None, "session is required for fetch_models" url = f"{self.base_url}/models" headers = {"Authorization": f"Bearer {self.api_key}"} async with self.session.get(url, headers=headers) as resp: if resp.status != 200: body = await resp.text() raise StreamError(f"OpenRouter API error {resp.status}: {body}") payload: dict[str, Any] = await resp.json() self.models.clear() for entry in payload.get("data", []): m = self._parse_model(entry) self.models[m.id] = m logger.info("Loaded %d models from %s", len(self.models), url)
@staticmethod def _parse_model(entry: dict[str, Any]) -> ModelSpec: caps: set[Capability] = set() params: list[str] = entry.get("supported_parameters", []) if "tools" in params: caps.add(Capability.tool_use) arch: dict[str, Any] = entry.get("architecture", {}) input_modalities: list[str] = arch.get("input_modalities", []) output_modalities: list[str] = arch.get("output_modalities", []) if "image" in input_modalities: caps.add(Capability.vision) if "embedding" in output_modalities: caps.add(Capability.embedding) top: dict[str, Any] = entry.get("top_provider", {}) context_window = int(entry.get("context_length") or top.get("context_length") or 128_000) max_output_tokens = int(top.get("max_completion_tokens") or 8_000) pricing: dict[str, Any] = entry.get("pricing", {}) return ModelSpec( id=entry["id"], context_window=context_window, max_output_tokens=max_output_tokens, capabilities=frozenset(caps), input_cost=float(pricing.get("prompt", 0)) * 1_000_000, output_cost=float(pricing.get("completion", 0)) * 1_000_000, )