mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix groq +xai model fetching
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import traceback
|
||||
@@ -1698,7 +1699,95 @@ class BaseUpstreamProvider:
|
||||
List of Model objects with pricing
|
||||
"""
|
||||
logger.debug(f"Fetching models for {self.provider_type or self.base_url}")
|
||||
return []
|
||||
|
||||
try:
|
||||
or_models, provider_models_response = await asyncio.gather(
|
||||
self._fetch_openrouter_models(),
|
||||
self._fetch_provider_models(),
|
||||
)
|
||||
|
||||
provider_model_ids = self._parse_model_ids(provider_models_response)
|
||||
|
||||
found_models = []
|
||||
not_found_models = []
|
||||
|
||||
for model_id in provider_model_ids:
|
||||
or_model = self._match_model(model_id, or_models)
|
||||
if or_model:
|
||||
try:
|
||||
model = Model(**or_model) # type: ignore
|
||||
found_models.append(model)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to parse model {model_id}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
else:
|
||||
not_found_models.append(model_id)
|
||||
|
||||
logger.info(
|
||||
"Fetched models for provider",
|
||||
extra={
|
||||
"provider": self.provider_type or self.base_url,
|
||||
"found_count": len(found_models),
|
||||
"not_found_count": len(not_found_models),
|
||||
},
|
||||
)
|
||||
|
||||
if not_found_models:
|
||||
logger.debug(
|
||||
"Models not found in OpenRouter",
|
||||
extra={"not_found_models": not_found_models},
|
||||
)
|
||||
|
||||
return found_models
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error fetching models for {self.provider_type or self.base_url}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
return []
|
||||
|
||||
async def _fetch_openrouter_models(self) -> list[dict]:
|
||||
"""Fetch models from OpenRouter API."""
|
||||
url = "https://openrouter.ai/api/v1/models"
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url)
|
||||
response.raise_for_status()
|
||||
models = response.json()
|
||||
return [
|
||||
model
|
||||
for model in models.get("data", [])
|
||||
if ":free" not in model.get("id", "").lower()
|
||||
]
|
||||
|
||||
async def _fetch_provider_models(self) -> dict:
|
||||
"""Fetch models from provider's API."""
|
||||
url = f"{self.base_url.rstrip('/')}/models"
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else None
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
def _parse_model_ids(self, response: dict) -> list[str]:
|
||||
"""Parse model IDs from provider response."""
|
||||
return [model.get("id") for model in response.get("data", []) if "id" in model]
|
||||
|
||||
def _match_model(self, model_id: str, or_models: list[dict]) -> dict | None:
|
||||
"""Match provider model ID with OpenRouter model."""
|
||||
return next(
|
||||
(
|
||||
model
|
||||
for model in or_models
|
||||
if (model.get("id") == model_id)
|
||||
or (model.get("id", "").split("/")[-1] == model_id)
|
||||
or (model.get("canonical_slug") == model_id)
|
||||
or (model.get("canonical_slug", "").split("/")[-1] == model_id)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
async def refresh_models_cache(self) -> None:
|
||||
"""Refresh the in-memory models cache from upstream API."""
|
||||
|
||||
@@ -39,8 +39,3 @@ class GroqUpstreamProvider(BaseUpstreamProvider):
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'groq/' prefix for Groq API compatibility."""
|
||||
return model_id.removeprefix("groq/")
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch Groq models from OpenRouter API filtered by groq source."""
|
||||
models_data = await async_fetch_openrouter_models(source_filter="groq")
|
||||
return [Model(**model) for model in models_data] # type: ignore
|
||||
|
||||
@@ -10,7 +10,7 @@ if TYPE_CHECKING:
|
||||
class XAIUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider specifically configured for XAI API."""
|
||||
|
||||
provider_type = "xai"
|
||||
provider_type = "x-ai"
|
||||
default_base_url = "https://api.x.ai/v1"
|
||||
platform_url = "https://console.x.ai/"
|
||||
|
||||
@@ -38,9 +38,9 @@ class XAIUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Strip 'xai/' prefix for XAI API compatibility."""
|
||||
return model_id.removeprefix("xai/")
|
||||
return model_id.removeprefix("x-ai/")
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch XAI models from OpenRouter API filtered by xai source."""
|
||||
models_data = await async_fetch_openrouter_models(source_filter="xai")
|
||||
models_data = await async_fetch_openrouter_models(source_filter="x-ai")
|
||||
return [Model(**model) for model in models_data] # type: ignore
|
||||
|
||||
Reference in New Issue
Block a user