From f9eaf48f451fb39f500dbb1a8678a22923e76a59 Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Thu, 23 Oct 2025 21:03:10 +0200 Subject: [PATCH] add ollama upstream --- routstr/upstream.py | 23 ++++ routstr/upstreams/__init__.py | 3 + routstr/upstreams/ollama.py | 243 ++++++++++++++++++++++++++++++++++ ui/app/providers/page.tsx | 17 ++- 4 files changed, 279 insertions(+), 7 deletions(-) create mode 100644 routstr/upstreams/__init__.py create mode 100644 routstr/upstreams/ollama.py diff --git a/routstr/upstream.py b/routstr/upstream.py index 0ecff033..adeb76cb 100644 --- a/routstr/upstream.py +++ b/routstr/upstream.py @@ -21,6 +21,7 @@ from .core import get_logger from .core.db import ApiKey, AsyncSession, ModelRow, UpstreamProviderRow, create_session from .payment.helpers import create_error_response from .payment.models import Model, async_fetch_openrouter_models +from .upstreams import OllamaUpstreamProvider logger = get_logger(__name__) @@ -323,6 +324,24 @@ async def _seed_providers_from_settings( ) seeded_base_urls.add(base_url) + ollama_base_url = os.environ.get("OLLAMA_BASE_URL") + if ollama_base_url: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == ollama_base_url + ) + ) + if not result.first(): + providers_to_add.append( + UpstreamProviderRow( + provider_type="ollama", + base_url=ollama_base_url, + api_key=os.environ.get("OLLAMA_API_KEY", ""), + enabled=True, + ) + ) + seeded_base_urls.add(ollama_base_url) + if settings.chat_completions_api_version and settings.upstream_base_url: base_url = settings.upstream_base_url if base_url not in seeded_base_urls: @@ -433,6 +452,10 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider return OpenRouterUpstreamProvider( provider_row.api_key, provider_row.provider_fee ) + elif provider_row.provider_type == "ollama": + return OllamaUpstreamProvider( + provider_row.base_url, provider_row.api_key, provider_row.provider_fee + ) elif provider_row.provider_type == "custom": return UpstreamProvider( provider_row.base_url, provider_row.api_key, provider_row.provider_fee diff --git a/routstr/upstreams/__init__.py b/routstr/upstreams/__init__.py new file mode 100644 index 00000000..742e0836 --- /dev/null +++ b/routstr/upstreams/__init__.py @@ -0,0 +1,3 @@ +from .ollama import OllamaUpstreamProvider + +__all__ = ["OllamaUpstreamProvider"] diff --git a/routstr/upstreams/ollama.py b/routstr/upstreams/ollama.py new file mode 100644 index 00000000..300efda4 --- /dev/null +++ b/routstr/upstreams/ollama.py @@ -0,0 +1,243 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import httpx + +if TYPE_CHECKING: + from ..payment.models import Model + +from ..core.logging import get_logger + +logger = get_logger(__name__) + + +class OllamaUpstreamProvider: + """Upstream provider specifically configured for Ollama API.""" + + base_url: str + api_key: str + upstream_name: str = "ollama" + provider_fee: float = 1.01 + _models_cache: list[Model] = [] + _models_by_id: dict[str, Model] = {} + + def __init__( + self, + base_url: str = "http://localhost:11434", + api_key: str = "", + provider_fee: float = 1.01, + ): + """Initialize Ollama provider. + + Args: + base_url: Ollama API base URL (default http://localhost:11434) + api_key: Optional API key (Ollama typically doesn't require one) + provider_fee: Provider fee multiplier (default 1.01 for 1% fee) + """ + self.upstream_name = "ollama" + self.base_url = base_url + self.api_key = api_key + self.provider_fee = provider_fee + self._models_cache = [] + self._models_by_id = {} + + def transform_model_name(self, model_id: str) -> str: + """Strip 'ollama/' prefix for Ollama API compatibility.""" + return model_id.removeprefix("ollama/") + + async def fetch_models(self) -> list[Model]: + """Fetch models from Ollama API using /api/tags endpoint.""" + from ..payment.models import Architecture, Model, Pricing, TopProvider + + try: + async with httpx.AsyncClient(timeout=30.0) as client: + response = await client.get(f"{self.base_url}/api/tags") + response.raise_for_status() + data = response.json() + + models_list = [] + for model_data in data.get("models", []): + model_name = model_data.get("name", "") + if not model_name: + continue + + details = model_data.get("details", {}) + parameter_size = details.get("parameter_size", "") + + context_length = 4096 + if ( + "70b" in parameter_size.lower() + or "72b" in parameter_size.lower() + ): + context_length = 8192 + elif "13b" in parameter_size.lower(): + context_length = 4096 + elif "7b" in parameter_size.lower(): + context_length = 4096 + elif "3b" in parameter_size.lower(): + context_length = 2048 + elif "1b" in parameter_size.lower(): + context_length = 2048 + + model_family = details.get("family", "unknown") + model_format = details.get("format", "unknown") + + description = f"Ollama {model_family} model" + if parameter_size: + description += f" ({parameter_size})" + + models_list.append( + Model( + id=model_name, + name=model_name, + created=0, + description=description, + context_length=context_length, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer=model_format, + instruct_type=None, + ), + pricing=Pricing( + prompt=0.000003, + completion=0.000003, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + max_prompt_cost=0.001, + max_completion_cost=0.001, + max_cost=0.001, + ), + sats_pricing=None, + per_request_limits=None, + top_provider=TopProvider( + context_length=context_length, + max_completion_tokens=context_length // 2, + is_moderated=False, + ), + enabled=True, + upstream_provider_id=None, + canonical_slug=None, + ) + ) + + logger.info( + f"Fetched {len(models_list)} models from Ollama", + extra={"model_count": len(models_list), "base_url": self.base_url}, + ) + return models_list + + except Exception as e: + logger.error( + f"Failed to fetch models from Ollama API: {e}", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "base_url": self.base_url, + }, + ) + return [] + + async def refresh_models_cache(self) -> None: + """Refresh the in-memory models cache from upstream API.""" + try: + from ..payment.models import _update_model_sats_pricing + from ..payment.price import sats_usd_price + + models = await self.fetch_models() + models_with_fees = [self._apply_provider_fee_to_model(m) for m in models] + + try: + sats_to_usd = sats_usd_price() + self._models_cache = [ + _update_model_sats_pricing(m, sats_to_usd) for m in models_with_fees + ] + except Exception: + self._models_cache = models_with_fees + + self._models_by_id = {m.id: m for m in self._models_cache} + logger.info( + f"Refreshed models cache for {self.upstream_name or self.base_url}", + extra={"model_count": len(models)}, + ) + except Exception as e: + logger.error( + f"Failed to refresh models cache for {self.upstream_name or self.base_url}", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + def get_cached_models(self) -> list[Model]: + """Get cached models for this provider. + + Returns: + List of cached Model objects + """ + return self._models_cache + + def get_cached_model_by_id(self, model_id: str) -> Model | None: + """Get a specific cached model by ID. + + Args: + model_id: Model identifier + + Returns: + Model object or None if not found + """ + return self._models_by_id.get(model_id) + + def _apply_provider_fee_to_model(self, model: Model) -> Model: + """Apply provider fee to model's USD pricing and calculate max costs. + + Args: + model: Model object to update + + Returns: + Model with provider fee applied to pricing and max costs calculated + """ + from ..payment.models import Model, Pricing, _calculate_usd_max_costs + + adjusted_pricing = Pricing.parse_obj( + {k: v * self.provider_fee for k, v in model.pricing.dict().items()} + ) + + temp_model = Model( + id=model.id, + name=model.name, + created=model.created, + description=model.description, + context_length=model.context_length, + architecture=model.architecture, + pricing=adjusted_pricing, + sats_pricing=None, + per_request_limits=model.per_request_limits, + top_provider=model.top_provider, + enabled=model.enabled, + upstream_provider_id=model.upstream_provider_id, + canonical_slug=model.canonical_slug, + ) + + ( + adjusted_pricing.max_prompt_cost, + adjusted_pricing.max_completion_cost, + adjusted_pricing.max_cost, + ) = _calculate_usd_max_costs(temp_model) + + return Model( + id=model.id, + name=model.name, + created=model.created, + description=model.description, + context_length=model.context_length, + architecture=model.architecture, + pricing=adjusted_pricing, + sats_pricing=model.sats_pricing, + per_request_limits=model.per_request_limits, + top_provider=model.top_provider, + enabled=model.enabled, + upstream_provider_id=model.upstream_provider_id, + canonical_slug=model.canonical_slug, + ) diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index 64339515..20be393e 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -187,6 +187,7 @@ export default function ProvidersPage() { openai: 'https://api.openai.com/v1', anthropic: 'https://api.anthropic.com/v1', azure: '', + ollama: 'http://localhost:11434', generic: '', }; return defaults[type] || ''; @@ -261,6 +262,7 @@ export default function ProvidersPage() { OpenAI Anthropic Azure OpenAI + Ollama Generic @@ -475,20 +477,20 @@ export default function ProvidersPage() { ) : providerModels && viewingModels === provider.id ? ( - + - - Database Models - - {providerModels.db_models.length} - - Remote Models {providerModels.remote_models.length} + + Database Models + + {providerModels.db_models.length} + + OpenAI Anthropic Azure OpenAI + Ollama Generic