diff --git a/studio/backend/core/inference/external_provider.py b/studio/backend/core/inference/external_provider.py index cb4b391864..c628f54d0e 100644 --- a/studio/backend/core/inference/external_provider.py +++ b/studio/backend/core/inference/external_provider.py @@ -34,10 +34,17 @@ class ExternalProviderClient: def _auth_headers(self) -> dict[str, str]: """Build authentication headers. All supported providers use Bearer tokens.""" - return { + from core.inference.providers import get_provider_info + + headers = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } + # Merge any provider-specific extra headers (e.g. OpenRouter attribution headers) + provider_info = get_provider_info(self.provider_type) + if provider_info: + headers.update(provider_info.get("extra_headers", {})) + return headers async def stream_chat_completion( self, diff --git a/studio/backend/core/inference/providers.py b/studio/backend/core/inference/providers.py index 7c98c2dc90..9d049d1d64 100644 --- a/studio/backend/core/inference/providers.py +++ b/studio/backend/core/inference/providers.py @@ -119,6 +119,27 @@ PROVIDER_REGISTRY: dict[str, dict[str, Any]] = { "auth_prefix": "Bearer ", "notes": "Web-grounded responses with built-in search.", }, + "openrouter": { + "display_name": "OpenRouter", + "base_url": "https://openrouter.ai/api/v1", + "default_models": [ + "openai/gpt-4o", + "anthropic/claude-sonnet-4-5", + "google/gemini-2.5-flash", + "meta-llama/llama-4-maverick", + "mistralai/mistral-small-3.1-24b-instruct", + ], + "supports_streaming": True, + "supports_vision": True, + "supports_tool_calling": True, + "auth_header": "Authorization", + "auth_prefix": "Bearer ", + "extra_headers": { + "HTTP-Referer": "https://unsloth.ai", + "X-Title": "Unsloth Studio", + }, + "notes": "Unified gateway to 300+ models across all major providers. HTTP-Referer and X-Title headers sent for attribution.", + }, } diff --git a/studio/backend/tests/test_providers_api.py b/studio/backend/tests/test_providers_api.py index d044f24c33..3e00c0d66e 100644 --- a/studio/backend/tests/test_providers_api.py +++ b/studio/backend/tests/test_providers_api.py @@ -48,6 +48,7 @@ _PROVIDER_CONFIGS: dict[str, tuple[str, str]] = { "together": ("TOGETHER_API_KEY", "meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8"), "fireworks": ("FIREWORKS_API_KEY", "accounts/fireworks/models/llama4-maverick-instruct-basic"), "perplexity": ("PERPLEXITY_API_KEY", "sonar"), + "openrouter": ("OPENROUTER_API_KEY", "openai/gpt-4o-mini"), } PROVIDER_KEYS: dict[str, str] = { @@ -230,7 +231,7 @@ class TestPublicKey: class TestRegistry: - def test_registry_returns_7_providers(self, auth_headers: dict[str, str]): + def test_registry_returns_8_providers(self, auth_headers: dict[str, str]): """GET /api/providers/registry returns all 7 supported providers.""" resp = requests.get( _url("/api/providers/registry"), @@ -239,7 +240,7 @@ class TestRegistry: ) assert resp.status_code == 200, f"Registry failed: {resp.text}" providers = resp.json() - assert len(providers) == 7, f"Expected 7 providers, got {len(providers)}: {providers}" + assert len(providers) == 8, f"Expected 8 providers, got {len(providers)}: {providers}" print(f"\n {'Provider':<12} {'Base URL'}") print(f" {'-'*12} {'-'*45}") for p in providers: