diff --git a/migrations/versions/a3b4c5d6e7f8_add_provider_fee_schedules.py b/migrations/versions/a3b4c5d6e7f8_add_provider_fee_schedules.py new file mode 100644 index 00000000..f0e3be5a --- /dev/null +++ b/migrations/versions/a3b4c5d6e7f8_add_provider_fee_schedules.py @@ -0,0 +1,40 @@ +"""add provider_fee_schedules to upstream_providers + +Revision ID: a3b4c5d6e7f8 +Revises: 614c0a740e68 +Create Date: 2026-04-12 00:00:00.000000 +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "a3b4c5d6e7f8" +down_revision = "b1c2d3e4f5a6" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + columns = [c["name"] for c in inspector.get_columns("upstream_providers")] + + if "provider_fee_schedules" not in columns: + op.add_column( + "upstream_providers", + sa.Column("provider_fee_schedules", sa.Text(), nullable=True), + ) + + if "provider_fee_default" not in columns: + op.add_column( + "upstream_providers", + sa.Column("provider_fee_default", sa.Float(), nullable=True), + ) + # Initialize default fee from current fee + op.execute("UPDATE upstream_providers SET provider_fee_default = provider_fee") + + +def downgrade() -> None: + op.drop_column("upstream_providers", "provider_fee_schedules") + op.drop_column("upstream_providers", "provider_fee_default") diff --git a/routstr/core/admin.py b/routstr/core/admin.py index d87120e2..45cb9be7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -9,7 +9,7 @@ from pydantic import BaseModel from sqlmodel import select from ..payment.models import _row_to_model, list_models -from ..proxy import refresh_model_maps, reinitialize_upstreams +from ..proxy import refresh_model_maps, reinitialize_upstreams, sync_provider_fees from ..wallet import ( fetch_all_balances, get_proofs_per_mint_and_unit, @@ -556,6 +556,7 @@ class UpstreamProviderCreate(BaseModel): api_version: str | None = None enabled: bool = True provider_fee: float = 1.01 + provider_fee_default: float | None = None provider_settings: dict | None = None @@ -566,29 +567,37 @@ class UpstreamProviderUpdate(BaseModel): api_version: str | None = None enabled: bool | None = None provider_fee: float | None = None + provider_fee_default: float | None = None provider_settings: dict | None = None +def _provider_to_dict( + p: UpstreamProviderRow, redact_key: bool = True +) -> dict[str, object]: + return { + "id": p.id, + "provider_type": p.provider_type, + "base_url": p.base_url, + "api_key": "[REDACTED]" if (redact_key and p.api_key) else (p.api_key or ""), + "api_version": p.api_version, + "enabled": p.enabled, + "provider_fee": p.provider_fee, + "provider_fee_default": p.provider_fee_default, + "provider_settings": json.loads(p.provider_settings) + if p.provider_settings + else None, + "provider_fee_schedules": json.loads(p.provider_fee_schedules) + if p.provider_fee_schedules + else [], + } + + @admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) async def get_upstream_providers() -> list[dict[str, object]]: async with create_session() as session: result = await session.exec(select(UpstreamProviderRow)) providers = result.all() - return [ - { - "id": p.id, - "provider_type": p.provider_type, - "base_url": p.base_url, - "api_key": "[REDACTED]" if p.api_key else "", - "api_version": p.api_version, - "enabled": p.enabled, - "provider_fee": p.provider_fee, - "provider_settings": json.loads(p.provider_settings) - if p.provider_settings - else None, - } - for p in providers - ] + return [_provider_to_dict(p) for p in providers] @admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) @@ -615,6 +624,9 @@ async def create_upstream_provider( api_version=payload.api_version, enabled=payload.enabled, provider_fee=payload.provider_fee, + provider_fee_default=payload.provider_fee_default + if payload.provider_fee_default is not None + else payload.provider_fee, provider_settings=json.dumps(payload.provider_settings) if payload.provider_settings else None, @@ -624,17 +636,7 @@ async def create_upstream_provider( await session.refresh(provider) await reinitialize_upstreams() - await refresh_model_maps() - return { - "id": provider.id, - "provider_type": provider.provider_type, - "base_url": provider.base_url, - "api_key": "[REDACTED]", - "api_version": provider.api_version, - "enabled": provider.enabled, - "provider_fee": provider.provider_fee, - "provider_settings": payload.provider_settings, - } + return _provider_to_dict(provider) @admin_router.get( @@ -645,18 +647,7 @@ async def get_upstream_provider(provider_id: int) -> dict[str, object]: provider = await session.get(UpstreamProviderRow, provider_id) if not provider: raise HTTPException(status_code=404, detail="Provider not found") - return { - "id": provider.id, - "provider_type": provider.provider_type, - "base_url": provider.base_url, - "api_key": "[REDACTED]" if provider.api_key else "", - "api_version": provider.api_version, - "enabled": provider.enabled, - "provider_fee": provider.provider_fee, - "provider_settings": json.loads(provider.provider_settings) - if provider.provider_settings - else None, - } + return _provider_to_dict(provider) @admin_router.patch( @@ -682,6 +673,8 @@ async def update_upstream_provider( provider.enabled = payload.enabled if payload.provider_fee is not None: provider.provider_fee = payload.provider_fee + if payload.provider_fee_default is not None: + provider.provider_fee_default = payload.provider_fee_default if payload.provider_settings is not None: provider.provider_settings = json.dumps(payload.provider_settings) @@ -690,19 +683,7 @@ async def update_upstream_provider( await session.refresh(provider) await reinitialize_upstreams() - await refresh_model_maps() - return { - "id": provider.id, - "provider_type": provider.provider_type, - "base_url": provider.base_url, - "api_key": "[REDACTED]", - "api_version": provider.api_version, - "enabled": provider.enabled, - "provider_fee": provider.provider_fee, - "provider_settings": json.loads(provider.provider_settings) - if provider.provider_settings - else None, - } + return _provider_to_dict(provider) @admin_router.delete( @@ -720,6 +701,78 @@ async def delete_upstream_provider(provider_id: int) -> dict[str, object]: return {"ok": True, "deleted_id": provider_id} +@admin_router.get( + "/api/upstream-providers/{provider_id}/fee-schedules", + dependencies=[Depends(require_admin_api)], +) +async def get_fee_schedules(provider_id: int) -> list[dict]: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + return ( + json.loads(provider.provider_fee_schedules) + if provider.provider_fee_schedules + else [] + ) + + +class FeeScheduleUpdate(BaseModel): + schedules: list[dict] + + +@admin_router.put( + "/api/upstream-providers/{provider_id}/fee-schedules", + dependencies=[Depends(require_admin_api)], +) +async def update_fee_schedules( + provider_id: int, payload: FeeScheduleUpdate +) -> list[dict]: + from ..payment.fee_schedule import FeeTimeRange, validate_no_overlaps + + try: + ranges = [FeeTimeRange(**s) for s in payload.schedules] + except Exception as e: + raise HTTPException(status_code=400, detail=f"Invalid schedule data: {e}") + + try: + validate_no_overlaps(ranges) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + + serialized = [r.dict() for r in ranges] + + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + provider.provider_fee_schedules = json.dumps(serialized) + session.add(provider) + await session.commit() + + await sync_provider_fees() + await refresh_model_maps() + return serialized + + +@admin_router.delete( + "/api/upstream-providers/{provider_id}/fee-schedules", + dependencies=[Depends(require_admin_api)], +) +async def delete_fee_schedules(provider_id: int) -> dict: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + provider.provider_fee_schedules = None + session.add(provider) + await session.commit() + + await sync_provider_fees() + await refresh_model_maps() + return {"ok": True} + + @admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)]) async def get_provider_types() -> list[dict[str, object]]: """Get metadata about available provider types including default URLs and whether they're fixed.""" diff --git a/routstr/core/db.py b/routstr/core/db.py index 4508a567..554eeecb 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -205,11 +205,17 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore ) enabled: bool = Field(default=True, description="Whether this provider is enabled") provider_fee: float = Field( - default=1.01, description="Provider fee multiplier (default 1%)" + default=1.01, description="Active fee multiplier (can be set by schedule)" + ) + provider_fee_default: float = Field( + default=1.01, description="Default fee multiplier (outside schedules)" ) provider_settings: str | None = Field( default=None, description="JSON string for provider-specific settings" ) + provider_fee_schedules: str | None = Field( + default=None, description="JSON array of fee time ranges (HH:MM UTC)" + ) models: list["ModelRow"] = Relationship( back_populates="upstream_provider", sa_relationship_kwargs={"cascade": "all, delete-orphan"}, diff --git a/routstr/payment/fee_schedule.py b/routstr/payment/fee_schedule.py new file mode 100644 index 00000000..a83739a8 --- /dev/null +++ b/routstr/payment/fee_schedule.py @@ -0,0 +1,113 @@ +"""Dynamic provider fee schedule logic. + +Supports time-based fee ranges (HH:MM UTC) with overlap validation and active fee resolution. +""" + +from __future__ import annotations + +import re +from datetime import datetime, timezone + +from pydantic.v1 import BaseModel, validator + +_HH_MM_RE = re.compile(r"^([01]\d|2[0-3]):([0-5]\d)$") + + +class FeeTimeRange(BaseModel): + start_time: str # HH:MM UTC + end_time: str # HH:MM UTC + provider_fee: float + + @validator("start_time", "end_time") + @classmethod + def validate_time_format(cls, v: str) -> str: + if not _HH_MM_RE.match(v): + raise ValueError(f"Time must be in HH:MM format (00:00–23:59), got: {v!r}") + return v + + @validator("provider_fee") + @classmethod + def validate_fee(cls, v: float) -> float: + if v <= 0: + raise ValueError(f"provider_fee must be > 0 (got {v})") + return v + + +def _to_minutes(t: str) -> int: + h, m = map(int, t.split(":")) + return h * 60 + m + + +def _range_intervals(r: FeeTimeRange) -> list[tuple[int, int]]: + """Return list of [start, end) minute intervals for this range. + + Handles midnight-crossing (e.g. 22:00–06:00 → [(1320,1440),(0,360)]). + start == end is treated as a full-day range. + """ + start = _to_minutes(r.start_time) + end = _to_minutes(r.end_time) + if start < end: + return [(start, end)] + if start > end: + return [(start, 1440), (0, end)] + # start == end → full day + return [(0, 1440)] + + +def _intervals_overlap(a: tuple[int, int], b: tuple[int, int]) -> bool: + return a[0] < b[1] and b[0] < a[1] + + +def ranges_overlap(a: FeeTimeRange, b: FeeTimeRange) -> bool: + """Return True if two fee time ranges overlap at any point in the day.""" + for ia in _range_intervals(a): + for ib in _range_intervals(b): + if _intervals_overlap(ia, ib): + return True + return False + + +def validate_no_overlaps(ranges: list[FeeTimeRange]) -> None: + """Raise ValueError if any two ranges in the list overlap.""" + for i in range(len(ranges)): + for j in range(i + 1, len(ranges)): + if ranges_overlap(ranges[i], ranges[j]): + raise ValueError( + f"Fee ranges overlap: [{ranges[i].start_time}–{ranges[i].end_time}]" + f" and [{ranges[j].start_time}–{ranges[j].end_time}]" + ) + + +def get_active_fee( + ranges: list[FeeTimeRange] | None, + default_fee: float, + *, + _now: datetime | None = None, +) -> float: + """Return the provider fee for the current UTC time. + + Falls back to *default_fee* when no range matches or *ranges* is empty/None. + The *_now* parameter is for testing only. + """ + if not ranges or not isinstance(ranges, list): + return default_fee + + now = _now if _now is not None else datetime.now(timezone.utc) + # Normalize to UTC + if now.tzinfo is not None: + now = now.astimezone(timezone.utc) + current = now.hour * 60 + now.minute + + for r in ranges: + start = _to_minutes(r.start_time) + end = _to_minutes(r.end_time) + if start < end: + if start <= current < end: + return r.provider_fee + elif start > end: # midnight-crossing + if current >= start or current < end: + return r.provider_fee + else: # full day (start == end) + return r.provider_fee + + return default_fee diff --git a/routstr/proxy.py b/routstr/proxy.py index ddab1c31..bfcce1d7 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -44,6 +44,7 @@ async def initialize_upstreams() -> None: global _upstreams _upstreams = await init_upstreams() logger.info(f"Initialized {len(_upstreams)} upstream providers") + await sync_provider_fees() await refresh_model_maps() @@ -55,6 +56,7 @@ async def reinitialize_upstreams() -> None: "Re-initialized upstream providers from admin action", extra={"provider_count": len(_upstreams)}, ) + await sync_provider_fees() await refresh_model_maps() @@ -100,6 +102,12 @@ async def refresh_model_maps() -> None: disabled_model_ids: set[str] = set() for provider in provider_rows: + # Match with instance in _upstreams to update its state from DB + for upstream in _upstreams: + if getattr(upstream, "db_id", None) == provider.id: + # This updates fee and merges DB models WITHOUT hitting network + await upstream.refresh_models_cache(skip_network=True) + if not provider.enabled: continue for model in provider.models: @@ -115,6 +123,39 @@ async def refresh_model_maps() -> None: ) +async def sync_provider_fees() -> None: + """Update active provider_fee in database based on schedules and defaults.""" + from .payment.fee_schedule import FeeTimeRange, get_active_fee + + async with create_session() as session: + result = await session.exec(select(UpstreamProviderRow)) + provider_rows = result.all() + + updated = False + for p in provider_rows: + schedules = None + if p.provider_fee_schedules: + try: + schedules = [ + FeeTimeRange(**s) for s in json.loads(p.provider_fee_schedules) + ] + except Exception: + pass + + active_fee = get_active_fee(schedules, p.provider_fee_default) + if p.provider_fee != active_fee: + logger.info( + f"Updating active fee for provider {p.id}: {p.provider_fee} -> {active_fee}", + extra={"provider_id": p.id, "active_fee": active_fee}, + ) + p.provider_fee = active_fee + session.add(p) + updated = True + + if updated: + await session.commit() + + async def refresh_model_maps_periodically() -> None: """Background task to refresh model maps every minute.""" import asyncio @@ -122,6 +163,7 @@ async def refresh_model_maps_periodically() -> None: while True: try: await asyncio.sleep(60) + await sync_provider_fees() await refresh_model_maps() except asyncio.CancelledError: break diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index a39a54de..b2f121b6 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -66,6 +66,7 @@ class BaseUpstreamProvider: api_key: str provider_fee: float = 1.05 _models_cache: list[Model] = [] + _raw_models_cache: list[Model] = [] _models_by_id: dict[str, Model] = {} def __init__(self, base_url: str, api_key: str, provider_fee: float = 1.01): @@ -80,6 +81,7 @@ class BaseUpstreamProvider: self.api_key = api_key self.provider_fee = provider_fee self._models_cache = [] + self._raw_models_cache = [] self._models_by_id = {} @classmethod @@ -677,7 +679,9 @@ class BaseUpstreamProvider: response_json["usage"]["cost_sats"] = ( cost_data.get("total_msats", 0) // 1000 ) - response_json["usage"]["remaining_balance_msats"] = remaining_balance_msats + response_json["usage"]["remaining_balance_msats"] = ( + remaining_balance_msats + ) # Keep detailed cost response_json["metadata"] = response_json.get("metadata", {}) @@ -753,7 +757,11 @@ class BaseUpstreamProvider: raise async def handle_streaming_responses_completion( - self, response: httpx.Response, key: ApiKey, max_cost_for_model: int, requested_model: str | None = None + self, + response: httpx.Response, + key: ApiKey, + max_cost_for_model: int, + requested_model: str | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -1000,7 +1008,9 @@ class BaseUpstreamProvider: response_json["usage"]["cost_sats"] = ( cost_data.get("total_msats", 0) // 1000 ) - response_json["usage"]["remaining_balance_msats"] = remaining_balance_msats + response_json["usage"]["remaining_balance_msats"] = ( + remaining_balance_msats + ) # Keep detailed cost response_json["metadata"] = response_json.get("metadata", {}) @@ -1309,7 +1319,9 @@ class BaseUpstreamProvider: path = self.normalize_request_path(path, model_obj) url = self.build_request_url(path, model_obj) - original_model_id = (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + original_model_id = ( + (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + ) transformed_body = self.prepare_request_body(request_body, model_obj) @@ -1469,7 +1481,10 @@ class BaseUpstreamProvider: background_tasks.add_task(response.aclose) background_tasks.add_task(client.aclose) result = await self.handle_streaming_chat_completion( - response, key, max_cost_for_model, background_tasks, + response, + key, + max_cost_for_model, + background_tasks, requested_model=original_model_id, ) result.background = background_tasks @@ -1479,7 +1494,10 @@ class BaseUpstreamProvider: if response.status_code == 200: try: return await self.handle_non_streaming_chat_completion( - response, key, session, max_cost_for_model, + response, + key, + session, + max_cost_for_model, requested_model=original_model_id, ) finally: @@ -1595,7 +1613,9 @@ class BaseUpstreamProvider: path = self.normalize_request_path(path, model_obj) url = self.build_request_url(path, model_obj) - original_model_id = (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + original_model_id = ( + (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + ) transformed_body = self.prepare_responses_request_body(request_body, model_obj) @@ -1683,7 +1703,9 @@ class BaseUpstreamProvider: if is_streaming and response.status_code == 200: result = await self.handle_streaming_responses_completion( - response, key, max_cost_for_model, + response, + key, + max_cost_for_model, requested_model=original_model_id, ) background_tasks = BackgroundTasks() @@ -1695,7 +1717,10 @@ class BaseUpstreamProvider: if response.status_code == 200: try: return await self.handle_non_streaming_responses_completion( - response, key, session, max_cost_for_model, + response, + key, + session, + max_cost_for_model, requested_model=original_model_id, ) finally: @@ -2122,7 +2147,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -2164,7 +2192,9 @@ class BaseUpstreamProvider: try: data_json = json.loads(line[6:]) if "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + data_json["usage"]["cost_sats"] = ( + cost_data.total_msats // 1000 + ) lines[i] = "data: " + json.dumps(data_json) except json.JSONDecodeError: pass @@ -2265,7 +2295,10 @@ class BaseUpstreamProvider: if refund_amount > 0: refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -2495,7 +2528,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - amount - 60, unit, mint, payment_token_hash, + amount - 60, + unit, + mint, + payment_token_hash, request_id=getattr(request.state, "request_id", None), ) @@ -2787,7 +2823,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - amount - 60, unit, mint, payment_token_hash, + amount - 60, + unit, + mint, + payment_token_hash, request_id=getattr(request.state, "request_id", None), ) @@ -3051,7 +3090,10 @@ class BaseUpstreamProvider: ) refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -3093,7 +3135,9 @@ class BaseUpstreamProvider: try: data_json = json.loads(line[6:]) if "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + data_json["usage"]["cost_sats"] = ( + cost_data.total_msats // 1000 + ) lines[i] = "data: " + json.dumps(data_json) except json.JSONDecodeError: pass @@ -3182,7 +3226,10 @@ class BaseUpstreamProvider: if refund_amount > 0: refund_token = await self.send_refund( - refund_amount, unit, mint, payment_token_hash, + refund_amount, + unit, + mint, + payment_token_hash, request_id=request_id, ) response_headers["X-Cashu"] = refund_token @@ -3521,8 +3568,28 @@ class BaseUpstreamProvider: None, ) - async def refresh_models_cache(self) -> None: - """Refresh the in-memory models cache from upstream API.""" + def apply_fee_to_cache(self) -> None: + """Apply current provider_fee to raw models and update active cache.""" + models_with_fees = [ + self._apply_provider_fee_to_model(m) for m in self._raw_models_cache + ] + + 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} + + async def refresh_models_cache(self, skip_network: bool = False) -> None: + """Refresh the in-memory models cache from upstream API and database. + + Args: + skip_network: If True, only refresh from database, skip hitting upstream API. + """ try: async with create_session() as session: stmt = select(UpstreamProviderRow).where( @@ -3536,6 +3603,9 @@ class BaseUpstreamProvider: if not provider or not provider.id: raise HTTPException(status_code=404, detail="Provider not found") + # Update fee from DB if it changed + self.provider_fee = provider.provider_fee + db_models = await list_models( session=session, upstream_id=provider.id, @@ -3543,34 +3613,37 @@ class BaseUpstreamProvider: apply_fees=False, ) db_model_ids: set[str] = {model.id for model in db_models} - models = await self.fetch_models() - model_ids = [model.id for model in models] - diff = set(db_model_ids) - set(model_ids) - for db_model_id in diff: - found_db_model = next( - ( - model_obj - for model_obj in db_models - if model_obj.id == db_model_id + if skip_network: + # Use existing raw models but filter/merge with DB models + # This avoids hitting the network + current_raw = {m.id: m for m in self._raw_models_cache} + # Keep only those still in current_raw (if we wanted to be strict) + # but actually we want to merge with db_models + models = [] + # Add all db_models (they take precedence as overrides) + models.extend(db_models) + # Add current raw models that are not in DB + for m_id, m in current_raw.items(): + if m_id not in db_model_ids: + models.append(m) + else: + models = await self.fetch_models() + model_ids = [model.id for model in models] + diff = set(db_model_ids) - set(model_ids) + + for db_model_id in diff: + found_db_model = next( + ( + model_obj + for model_obj in db_models + if model_obj.id == db_model_id + ) ) - ) - models.append(found_db_model) + models.append(found_db_model) - 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} + self._raw_models_cache = models + self.apply_fee_to_cache() except Exception as e: logger.error( diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index eff5d5bb..7ef27a09 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -65,9 +65,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - def get_request_base_url( - self, path: str, model_obj: Model | None = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" return f"{self.base_url.rstrip('/')}/v1" @@ -166,103 +164,3 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): }, ) 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.base_url}", - extra={"model_count": len(models)}, - ) - except Exception as e: - logger.error( - f"Failed to refresh models cache for {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/tests/integration/test_model_price_sync.py b/tests/integration/test_model_price_sync.py new file mode 100644 index 00000000..f73c4553 --- /dev/null +++ b/tests/integration/test_model_price_sync.py @@ -0,0 +1,190 @@ +"""Integration tests for model price updates when provider fee schedules change.""" + +import time +from typing import Any, Generator + +import pytest +from httpx import AsyncClient + +from routstr.core.admin import admin_sessions + +ADMIN_TOKEN = "test-admin-token" + + +def _auth_header() -> dict[str, str]: + return {"Authorization": f"Bearer {ADMIN_TOKEN}"} + + +@pytest.fixture(autouse=True) +def _inject_admin_session() -> Generator[None, None, None]: + admin_sessions[ADMIN_TOKEN] = int(time.time()) + 3600 + yield + admin_sessions.pop(ADMIN_TOKEN, None) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_price_updates_on_fee_schedule_change( + integration_client: AsyncClient, + patched_db_engine: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Patch fetch_models to return empty list to avoid network errors + # and allow DB models to be used + from routstr.upstream.base import BaseUpstreamProvider + + async def mock_fetch_models(self: BaseUpstreamProvider) -> list: + return [] + + monkeypatch.setattr(BaseUpstreamProvider, "fetch_models", mock_fetch_models) + + # 1. Create a provider + provider_resp = await integration_client.post( + "/admin/api/upstream-providers", + json={ + "provider_type": "custom", + "base_url": "https://api.example.com/v1", + "api_key": "test-key", + "enabled": True, + "provider_fee": 1.0, + }, + headers=_auth_header(), + ) + provider_id = provider_resp.json()["id"] + + # 2. Add a model to this provider + model_id = "test-model-price-update" + await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + json={ + "id": model_id, + "name": "Test Model", + "created": int(time.time()), + "description": "Test", + "context_length": 4096, + "architecture": { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "gpt2", + "instruct_type": "none", + }, + "pricing": {"prompt": 1.0, "completion": 2.0}, + "enabled": True, + }, + headers=_auth_header(), + ) + + # 3. Check initial price (should be prompt=1.0 * fee=1.0 = 1.0) + # We use /models endpoint + resp = await integration_client.get("/models") + models = resp.json()["data"] + target = next((m for m in models if m["id"] == model_id), None) + assert target is not None + assert target["pricing"]["prompt"] == 1.0 + + # 4. Update provider fee schedule to a very high value for the current time + # We'll use a range that covers the whole day to be safe + schedules = [ + {"start_time": "00:00", "end_time": "23:59", "provider_fee": 2.5}, + ] + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={"schedules": schedules}, + headers=_auth_header(), + ) + + # 5. Check price again - should be updated instantly + resp = await integration_client.get("/models") + models = resp.json()["data"] + target = next((m for m in models if m["id"] == model_id), None) + assert target is not None + # 1.0 * 2.5 = 2.5 + assert target["pricing"]["prompt"] == 2.5 + + # 6. Delete schedules + await integration_client.delete( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + headers=_auth_header(), + ) + + # 7. Should revert to default fee (1.0) + resp = await integration_client.get("/models") + models = resp.json()["data"] + target = next((m for m in models if m["id"] == model_id), None) + assert target is not None + assert target["pricing"]["prompt"] == 1.0 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_upstream_model_price_updates_on_fee_schedule_change( + integration_client: AsyncClient, + patched_db_engine: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.payment.models import Architecture, Model, Pricing + from routstr.upstream.base import BaseUpstreamProvider + + upstream_model_id = "upstream-model-only" + + # Mock fetch_models to return a model + async def mock_fetch_models(self: BaseUpstreamProvider) -> list[Model]: + return [ + Model( + id=upstream_model_id, + name="Upstream Model", + created=int(time.time()), + description="Test", + context_length=4096, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="gpt2", + instruct_type="none", + ), + pricing=Pricing(prompt=1.0, completion=2.0), + enabled=True, + ) + ] + + monkeypatch.setattr(BaseUpstreamProvider, "fetch_models", mock_fetch_models) + + # 1. Create a provider + provider_resp = await integration_client.post( + "/admin/api/upstream-providers", + json={ + "provider_type": "custom", + "base_url": "https://api.example.com/v1", + "api_key": "test-key-2", + "enabled": True, + "provider_fee": 1.0, + }, + headers=_auth_header(), + ) + provider_id = provider_resp.json()["id"] + + # 2. Check initial price (should be prompt=1.0 * fee=1.0 = 1.0) + resp = await integration_client.get("/models") + models = resp.json()["data"] + target = next((m for m in models if m["id"] == upstream_model_id), None) + assert target is not None + assert target["pricing"]["prompt"] == 1.0 + + # 3. Update provider fee schedule + schedules = [ + {"start_time": "00:00", "end_time": "23:59", "provider_fee": 3.0}, + ] + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={"schedules": schedules}, + headers=_auth_header(), + ) + + # 4. Check price again - I expect this to FAIL (still 1.0 instead of 3.0) + resp = await integration_client.get("/models") + models = resp.json()["data"] + target = next((m for m in models if m["id"] == upstream_model_id), None) + assert target is not None + assert target["pricing"]["prompt"] == 3.0 diff --git a/tests/integration/test_provider_fee_enforcement.py b/tests/integration/test_provider_fee_enforcement.py index ff93a231..fc6321ba 100644 --- a/tests/integration/test_provider_fee_enforcement.py +++ b/tests/integration/test_provider_fee_enforcement.py @@ -101,7 +101,7 @@ async def test_enforce_lowest_provider_fee_for_same_url( ) ] - async def refresh_models_cache(self) -> None: + async def refresh_models_cache(self, skip_network: bool = False) -> None: pass def prepare_headers(self, request_headers: dict[str, str]) -> dict[str, str]: diff --git a/tests/integration/test_provider_fee_schedules.py b/tests/integration/test_provider_fee_schedules.py new file mode 100644 index 00000000..4bd814cf --- /dev/null +++ b/tests/integration/test_provider_fee_schedules.py @@ -0,0 +1,402 @@ +"""Integration tests for provider fee schedule API endpoints.""" + +from typing import Any, Generator + +import pytest +from httpx import AsyncClient + +from routstr.core.admin import admin_sessions + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +ADMIN_TOKEN = "test-admin-token" + + +def _auth_header() -> dict[str, str]: + return {"Authorization": f"Bearer {ADMIN_TOKEN}"} + + +async def _create_provider(client: AsyncClient, *, fee: float = 1.02) -> int: + """Create a test provider and return its ID.""" + resp = await client.post( + "/admin/api/upstream-providers", + json={ + "provider_type": "custom", + "base_url": "https://api.example.com/v1", + "api_key": "test-key", + "enabled": True, + "provider_fee": fee, + }, + headers=_auth_header(), + ) + assert resp.status_code == 200, resp.text + return resp.json()["id"] + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture(autouse=True) +def _inject_admin_session() -> Generator[None, None, None]: + """Inject a valid admin session token for all tests.""" + import time + + admin_sessions[ADMIN_TOKEN] = int(time.time()) + 3600 + yield + admin_sessions.pop(ADMIN_TOKEN, None) + + +@pytest.fixture(autouse=True) +def _patch_reinitialize(monkeypatch: Any) -> None: + async def _noop(*args: Any, **kwargs: Any) -> None: + pass + + monkeypatch.setattr("routstr.core.admin.reinitialize_upstreams", _noop) + monkeypatch.setattr("routstr.core.admin.refresh_model_maps", _noop) + + +# --------------------------------------------------------------------------- +# GET fee schedules +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_get_fee_schedules_empty_for_new_provider( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + headers=_auth_header(), + ) + assert resp.status_code == 200 + assert resp.json() == [] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_get_fee_schedules_404_for_missing_provider( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + resp = await integration_client.get( + "/admin/api/upstream-providers/99999/fee-schedules", + headers=_auth_header(), + ) + assert resp.status_code == 404 + + +# --------------------------------------------------------------------------- +# PUT fee schedules +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_success( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + schedules = [ + {"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05}, + {"start_time": "18:00", "end_time": "08:00", "provider_fee": 1.02}, + ] + resp = await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={"schedules": schedules}, + headers=_auth_header(), + ) + assert resp.status_code == 200 + data = resp.json() + assert len(data) == 2 + assert data[0]["start_time"] == "08:00" + assert data[0]["provider_fee"] == 1.05 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_persisted( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + """Saved schedules are returned by a subsequent GET.""" + provider_id = await _create_provider(integration_client) + schedules = [{"start_time": "09:00", "end_time": "17:00", "provider_fee": 1.07}] + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={"schedules": schedules}, + headers=_auth_header(), + ) + get_resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + headers=_auth_header(), + ) + assert get_resp.status_code == 200 + assert get_resp.json()[0]["provider_fee"] == 1.07 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_replaces_existing( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + # Set initial schedule + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "08:00", "end_time": "12:00", "provider_fee": 1.03} + ] + }, + headers=_auth_header(), + ) + # Replace with different schedule + resp = await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "14:00", "end_time": "20:00", "provider_fee": 1.08} + ] + }, + headers=_auth_header(), + ) + assert resp.status_code == 200 + data = resp.json() + assert len(data) == 1 + assert data[0]["start_time"] == "14:00" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_overlap_rejected( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + schedules = [ + {"start_time": "08:00", "end_time": "14:00", "provider_fee": 1.05}, + {"start_time": "12:00", "end_time": "18:00", "provider_fee": 1.03}, + ] + resp = await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={"schedules": schedules}, + headers=_auth_header(), + ) + assert resp.status_code == 400 + assert "overlap" in resp.json()["detail"].lower() + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_invalid_time_format_rejected( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + resp = await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "8:00", "end_time": "18:00", "provider_fee": 1.05} + ] + }, + headers=_auth_header(), + ) + assert resp.status_code == 400 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_invalid_fee_rejected( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + resp = await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "08:00", "end_time": "18:00", "provider_fee": -0.5} + ] + }, + headers=_auth_header(), + ) + assert resp.status_code == 400 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_empty_clears_schedules( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + # Set a schedule + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05} + ] + }, + headers=_auth_header(), + ) + # Clear with empty list + resp = await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={"schedules": []}, + headers=_auth_header(), + ) + assert resp.status_code == 200 + assert resp.json() == [] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_put_fee_schedules_404_for_missing_provider( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + resp = await integration_client.put( + "/admin/api/upstream-providers/99999/fee-schedules", + json={"schedules": []}, + headers=_auth_header(), + ) + assert resp.status_code == 404 + + +# --------------------------------------------------------------------------- +# DELETE fee schedules +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_delete_fee_schedules( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + # Add schedules + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05} + ] + }, + headers=_auth_header(), + ) + # Delete + del_resp = await integration_client.delete( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + headers=_auth_header(), + ) + assert del_resp.status_code == 200 + assert del_resp.json()["ok"] is True + + # Verify schedules are gone + get_resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + headers=_auth_header(), + ) + assert get_resp.json() == [] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_delete_fee_schedules_404_for_missing_provider( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + resp = await integration_client.delete( + "/admin/api/upstream-providers/99999/fee-schedules", + headers=_auth_header(), + ) + assert resp.status_code == 404 + + +# --------------------------------------------------------------------------- +# Fee schedules appear in provider list and detail +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_fee_schedules_in_provider_list( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05} + ] + }, + headers=_auth_header(), + ) + list_resp = await integration_client.get( + "/admin/api/upstream-providers", headers=_auth_header() + ) + assert list_resp.status_code == 200 + providers = list_resp.json() + target = next((p for p in providers if p["id"] == provider_id), None) + assert target is not None + assert len(target["provider_fee_schedules"]) == 1 + assert target["provider_fee_schedules"][0]["provider_fee"] == 1.05 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_fee_schedules_in_provider_detail( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "10:00", "end_time": "22:00", "provider_fee": 1.06} + ] + }, + headers=_auth_header(), + ) + detail_resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}", headers=_auth_header() + ) + assert detail_resp.status_code == 200 + data = detail_resp.json() + assert len(data["provider_fee_schedules"]) == 1 + assert data["provider_fee_schedules"][0]["start_time"] == "10:00" + + +# --------------------------------------------------------------------------- +# Provider deletion clears fee schedules +# --------------------------------------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_provider_delete_clears_fee_schedules( + integration_client: AsyncClient, patched_db_engine: Any +) -> None: + provider_id = await _create_provider(integration_client) + await integration_client.put( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + json={ + "schedules": [ + {"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05} + ] + }, + headers=_auth_header(), + ) + # Delete provider + del_resp = await integration_client.delete( + f"/admin/api/upstream-providers/{provider_id}", headers=_auth_header() + ) + assert del_resp.status_code == 200 + + # Provider is gone → schedule endpoint returns 404 + get_resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/fee-schedules", + headers=_auth_header(), + ) + assert get_resp.status_code == 404 diff --git a/tests/unit/test_fee_schedule.py b/tests/unit/test_fee_schedule.py new file mode 100644 index 00000000..6f329e71 --- /dev/null +++ b/tests/unit/test_fee_schedule.py @@ -0,0 +1,256 @@ +"""Unit tests for routstr.payment.fee_schedule.""" + +from datetime import datetime, timedelta, timezone + +import pytest +from pydantic.v1 import ValidationError + +from routstr.payment.fee_schedule import ( + FeeTimeRange, + get_active_fee, + ranges_overlap, + validate_no_overlaps, +) + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _r(start: str, end: str, fee: float = 1.05) -> FeeTimeRange: + return FeeTimeRange(start_time=start, end_time=end, provider_fee=fee) + + +def _now(h: int, m: int = 0) -> datetime: + return datetime(2026, 1, 1, h, m, tzinfo=timezone.utc) + + +# --------------------------------------------------------------------------- +# FeeTimeRange validation +# --------------------------------------------------------------------------- + + +class TestFeeTimeRangeValidation: + def test_valid_range(self) -> None: + r = _r("08:00", "18:00", 1.05) + assert r.start_time == "08:00" + assert r.end_time == "18:00" + assert r.provider_fee == 1.05 + + def test_invalid_start_time_format(self) -> None: + with pytest.raises(ValidationError, match="HH:MM"): + _r("8:00", "18:00") + + def test_invalid_end_time_hour_out_of_range(self) -> None: + with pytest.raises(ValidationError): + _r("08:00", "24:00") + + def test_invalid_end_time_minute_out_of_range(self) -> None: + with pytest.raises(ValidationError): + _r("08:00", "18:60") + + def test_invalid_time_letters(self) -> None: + with pytest.raises(ValidationError): + _r("ab:cd", "18:00") + + def test_fee_must_be_positive(self) -> None: + with pytest.raises(ValidationError, match="provider_fee must be > 0"): + _r("08:00", "18:00", fee=0.0) + + def test_fee_negative_rejected(self) -> None: + with pytest.raises(ValidationError): + _r("08:00", "18:00", fee=-0.5) + + def test_fee_below_one_allowed(self) -> None: + r = _r("08:00", "18:00", fee=0.95) + assert r.provider_fee == 0.95 + + def test_boundary_times_valid(self) -> None: + r = _r("00:00", "23:59") + assert r.start_time == "00:00" + assert r.end_time == "23:59" + + +# --------------------------------------------------------------------------- +# ranges_overlap +# --------------------------------------------------------------------------- + + +class TestRangesOverlap: + def test_non_overlapping_ranges(self) -> None: + assert not ranges_overlap(_r("08:00", "12:00"), _r("12:00", "18:00")) + + def test_overlapping_ranges(self) -> None: + assert ranges_overlap(_r("08:00", "14:00"), _r("12:00", "18:00")) + + def test_one_contains_the_other(self) -> None: + assert ranges_overlap(_r("08:00", "20:00"), _r("10:00", "18:00")) + + def test_identical_ranges_overlap(self) -> None: + assert ranges_overlap(_r("08:00", "12:00"), _r("08:00", "12:00")) + + def test_adjacent_non_overlapping(self) -> None: + # end of first == start of second → no overlap (open interval [start, end)) + assert not ranges_overlap(_r("06:00", "12:00"), _r("12:00", "18:00")) + + def test_midnight_crossing_vs_day_range_overlap(self) -> None: + # 22:00–06:00 crosses midnight; 04:00–08:00 should overlap (both cover 04:00–06:00) + assert ranges_overlap(_r("22:00", "06:00"), _r("04:00", "08:00")) + + def test_midnight_crossing_vs_non_overlapping_day_range(self) -> None: + # 22:00–06:00 does NOT cover 10:00–18:00 + assert not ranges_overlap(_r("22:00", "06:00"), _r("10:00", "18:00")) + + def test_two_midnight_crossing_ranges_overlap(self) -> None: + assert ranges_overlap(_r("20:00", "04:00"), _r("22:00", "06:00")) + + def test_two_midnight_crossing_ranges_non_overlap(self) -> None: + # 21:00–23:00 and 23:00–21:00 (full day minus one hour): they do overlap + # Let's use a case that genuinely doesn't: 21:00–22:00 adjacent + # Actually for two midnight-crossing ranges it's hard to not overlap—let's test equal endpoints + assert not ranges_overlap(_r("22:00", "23:00"), _r("23:00", "01:00")) + + +# --------------------------------------------------------------------------- +# validate_no_overlaps +# --------------------------------------------------------------------------- + + +class TestValidateNoOverlaps: + def test_no_overlaps_passes(self) -> None: + validate_no_overlaps( + [_r("00:00", "08:00"), _r("08:00", "16:00"), _r("16:00", "23:59")] + ) + + def test_overlap_raises(self) -> None: + with pytest.raises(ValueError, match="overlap"): + validate_no_overlaps([_r("08:00", "14:00"), _r("12:00", "18:00")]) + + def test_single_range_passes(self) -> None: + validate_no_overlaps([_r("08:00", "18:00")]) + + def test_empty_list_passes(self) -> None: + validate_no_overlaps([]) + + def test_midnight_crossing_overlap_detected(self) -> None: + with pytest.raises(ValueError, match="overlap"): + validate_no_overlaps([_r("22:00", "06:00"), _r("04:00", "08:00")]) + + +# --------------------------------------------------------------------------- +# get_active_fee +# --------------------------------------------------------------------------- + + +class TestGetActiveFee: + def test_returns_default_when_no_ranges(self) -> None: + assert get_active_fee(None, 1.01) == 1.01 + + def test_returns_default_for_empty_list(self) -> None: + assert get_active_fee([], 1.01) == 1.01 + + def test_returns_matching_fee(self) -> None: + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=_now(12)) == 1.05 + + def test_returns_default_when_no_match(self) -> None: + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=_now(20)) == 1.01 + + def test_boundary_start_inclusive(self) -> None: + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=_now(8, 0)) == 1.05 + + def test_boundary_end_exclusive(self) -> None: + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=_now(18, 0)) == 1.01 + + def test_midnight_crossing_before_midnight(self) -> None: + ranges = [_r("22:00", "06:00", fee=1.03)] + assert get_active_fee(ranges, 1.01, _now=_now(23)) == 1.03 + + def test_midnight_crossing_after_midnight(self) -> None: + ranges = [_r("22:00", "06:00", fee=1.03)] + assert get_active_fee(ranges, 1.01, _now=_now(3)) == 1.03 + + def test_midnight_crossing_outside_range(self) -> None: + ranges = [_r("22:00", "06:00", fee=1.03)] + assert get_active_fee(ranges, 1.01, _now=_now(12)) == 1.01 + + def test_multiple_ranges_correct_match(self) -> None: + ranges = [ + _r("00:00", "08:00", fee=1.02), + _r("08:00", "16:00", fee=1.05), + _r("16:00", "23:59", fee=1.03), + ] + assert get_active_fee(ranges, 1.01, _now=_now(10)) == 1.05 + assert get_active_fee(ranges, 1.01, _now=_now(2)) == 1.02 + assert get_active_fee(ranges, 1.01, _now=_now(20)) == 1.03 + + def test_first_matching_range_wins(self) -> None: + # When multiple ranges could match (should not happen if validated), + # the first one wins. + ranges = [_r("08:00", "20:00", fee=1.05), _r("10:00", "12:00", fee=1.02)] + assert get_active_fee(ranges, 1.01, _now=_now(11)) == 1.05 + + +# --------------------------------------------------------------------------- +# Timezone-aware inputs (CEST / CET) +# --------------------------------------------------------------------------- + + +class TestGetActiveFeeTimezones: + """Verify that tz-aware datetimes are normalised to UTC before matching.""" + + # CEST = UTC+2 (Central European Summer Time, used ~late March – late Oct) + CEST = timezone(timedelta(hours=2)) + # CET = UTC+1 (Central European Time, used the rest of the year) + CET = timezone(timedelta(hours=1)) + + def test_cest_datetime_normalised_to_utc_matches(self) -> None: + # 10:00 CEST == 08:00 UTC — schedule 08:00–18:00 should match + now_cest = datetime(2026, 7, 1, 10, 0, tzinfo=self.CEST) + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.05 + + def test_cest_datetime_normalised_to_utc_no_match(self) -> None: + # 06:00 CEST == 04:00 UTC — schedule 08:00–18:00 should NOT match + now_cest = datetime(2026, 7, 1, 6, 0, tzinfo=self.CEST) + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.01 + + def test_cet_datetime_normalised_to_utc_matches(self) -> None: + # 09:00 CET == 08:00 UTC — schedule 08:00–18:00 should match + now_cet = datetime(2026, 1, 15, 9, 0, tzinfo=self.CET) + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=now_cet) == 1.05 + + def test_cet_datetime_before_utc_range(self) -> None: + # 08:30 CET == 07:30 UTC — schedule 08:00–18:00 should NOT match + now_cet = datetime(2026, 1, 15, 8, 30, tzinfo=self.CET) + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=now_cet) == 1.01 + + def test_cest_midnight_crossing_before_midnight(self) -> None: + # 00:30 CEST == 22:30 UTC — schedule 22:00–06:00 UTC should match + now_cest = datetime(2026, 7, 2, 0, 30, tzinfo=self.CEST) + ranges = [_r("22:00", "06:00", fee=1.03)] + assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.03 + + def test_cest_midnight_crossing_after_midnight(self) -> None: + # 05:00 CEST == 03:00 UTC — schedule 22:00–06:00 UTC should match + now_cest = datetime(2026, 7, 2, 5, 0, tzinfo=self.CEST) + ranges = [_r("22:00", "06:00", fee=1.03)] + assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.03 + + def test_cest_midnight_crossing_outside_range(self) -> None: + # 14:00 CEST == 12:00 UTC — schedule 22:00–06:00 UTC should NOT match + now_cest = datetime(2026, 7, 2, 14, 0, tzinfo=self.CEST) + ranges = [_r("22:00", "06:00", fee=1.03)] + assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.01 + + def test_naive_utc_datetime_still_works(self) -> None: + # Naive datetimes are treated as UTC (defensive fallback path) + now_naive = datetime(2026, 1, 1, 12, 0) # no tzinfo + ranges = [_r("08:00", "18:00", fee=1.05)] + assert get_active_fee(ranges, 1.01, _now=now_naive) == 1.05 diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index 4a7046a8..746b3f71 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -14,10 +14,11 @@ import { } from '@/lib/api/services/admin'; import { AddProviderModelDialog } from '@/components/add-provider-model-dialog'; import { BatchOverrideDialog } from '@/components/batch-override-dialog'; +import { ProviderFeeScheduleModal } from '@/components/provider-fee-schedule-modal'; import { ProviderCard } from '@/components/provider-card'; import { ProviderFormDialogContent } from '@/components/provider-form-dialog-content'; import { Skeleton } from '@/components/ui/skeleton'; -import { AlertCircle, Plus, Server } from 'lucide-react'; +import { AlertCircle, Clock, Plus, Server } from 'lucide-react'; import { Alert, AlertDescription } from '@/components/ui/alert'; import { Dialog, DialogTrigger } from '@/components/ui/dialog'; import { @@ -70,6 +71,10 @@ export default function ProvidersPage() { const [batchOverrideProviderId, setBatchOverrideProviderId] = useState< number | null >(null); + const [feeScheduleState, setFeeScheduleState] = useState<{ + open: boolean; + initialIds: number[]; + }>({ open: false, initialIds: [] }); const [providerDeleteTarget, setProviderDeleteTarget] = useState(null); const [modelDeleteTarget, setModelDeleteTarget] = useState<{ @@ -242,6 +247,7 @@ export default function ProvidersPage() { api_version: provider.api_version || null, enabled: provider.enabled, provider_fee: provider.provider_fee, + provider_fee_default: provider.provider_fee_default, provider_settings: provider.provider_settings || {}, }); setIsEditDialogOpen(true); @@ -254,7 +260,7 @@ export default function ProvidersPage() { base_url: formData.base_url, api_version: formData.api_version, enabled: formData.enabled, - provider_fee: formData.provider_fee, + provider_fee_default: formData.provider_fee_default, provider_settings: formData.provider_settings, }; if (formData.api_key) { @@ -342,6 +348,13 @@ export default function ProvidersPage() { setBatchOverrideProviderId(providerId); }; + const handleManageFeeSchedules = (providerId?: number) => { + setFeeScheduleState({ + open: true, + initialIds: providerId !== undefined ? [providerId] : [], + }); + }; + const availableMints = (globalSettings?.cashu_mints as string[]) || []; return ( @@ -352,12 +365,22 @@ export default function ProvidersPage() { title='Upstream Providers' description='Manage your AI provider connections and credentials.' actions={ - - - + + + + } /> handleEdit(provider)} onDeleteProvider={() => setProviderDeleteTarget(provider)} onBatchOverride={() => handleBatchOverride(provider.id)} + onManageFeeSchedules={() => + handleManageFeeSchedules(provider.id) + } onAddModel={() => handleAddModel(provider.id)} onEditModel={(model) => handleEditModel(provider.id, model)} onDeleteModel={(modelId) => @@ -562,6 +588,16 @@ export default function ProvidersPage() { }} /> )} + + setFeeScheduleState({ open: false, initialIds: [] })} + onSuccess={() => { + queryClient.invalidateQueries({ queryKey: ['upstream-providers'] }); + }} + /> ); } diff --git a/ui/components/provider-card.tsx b/ui/components/provider-card.tsx index b48dc64b..824d234a 100644 --- a/ui/components/provider-card.tsx +++ b/ui/components/provider-card.tsx @@ -20,6 +20,7 @@ import { Trash2, Key, RotateCcw, + Clock, } from 'lucide-react'; import { ProviderBalance } from '@/components/provider-balance'; import { ProviderModelsPanel } from '@/components/provider-models-panel'; @@ -54,6 +55,7 @@ interface ProviderCardProps { onDeleteModel: (modelId: string) => void; onOverrideModel: (model: AdminModel) => void; onUpdateApiKey: (newKey: string) => void; + onManageFeeSchedules: () => void; availableMints: string[]; } @@ -74,6 +76,7 @@ export function ProviderCard({ onDeleteModel, onOverrideModel, onUpdateApiKey, + onManageFeeSchedules, }: ProviderCardProps) { const queryClient = useQueryClient(); const [isKeyModalOpen, setIsKeyModalOpen] = useState(false); @@ -113,6 +116,14 @@ export function ProviderCard({ > {provider.enabled ? 'Enabled' : 'Disabled'} + + Fee: {provider.provider_fee}x + {provider.provider_fee !== provider.provider_fee_default && ( + + (default: {provider.provider_fee_default}x) + + )} + {provider.base_url} @@ -189,6 +200,22 @@ export function ProviderCard({ )} + + + +
+ {providers.map((p) => ( +
+ + +
+ ))} +
+ {noneSelected && ( +

+ Select at least one provider. +

+ )} + + + {/* Fee range editor */} +
+
+ +
+ setEnforceOverride(!!checked)} + /> + +
+
+ +
+ {rows.length === 0 && ( +

+ No ranges configured — saving with no ranges will clear + schedules. +

+ )} + + {rows.map((row) => { + const isOverlap = overlapping.has(row._id); + const badTime = + (row.start_time && !isValidTime(row.start_time)) || + (row.end_time && !isValidTime(row.end_time)); + const badFee = row.provider_fee <= 1.0; + const hasError = isOverlap || badTime || badFee; + + return ( +
+
+ + + updateRow( + row._id, + 'start_time', + normalizeTime(e.target.value) + ) + } + className='border-input bg-background ring-offset-background focus-visible:ring-ring flex h-10 w-full rounded-md border px-3 py-2 font-mono text-sm focus-visible:ring-2 focus-visible:ring-offset-2 focus-visible:outline-none disabled:cursor-not-allowed disabled:opacity-50' + /> +
+
+ + + updateRow( + row._id, + 'end_time', + normalizeTime(e.target.value) + ) + } + className='border-input bg-background ring-offset-background focus-visible:ring-ring flex h-10 w-full rounded-md border px-3 py-2 font-mono text-sm focus-visible:ring-2 focus-visible:ring-offset-2 focus-visible:outline-none disabled:cursor-not-allowed disabled:opacity-50' + /> +
+
+ + + updateRow( + row._id, + 'provider_fee', + parseFloat(e.target.value) || 0 + ) + } + /> +
+ +
+ ); + })} +
+ + {overlapping.size > 0 && ( +

+ Some ranges overlap — fix them before saving. +

+ )} +
+ +
+ +
+ + + + + + + + + ); +} diff --git a/ui/components/provider-form-fields.tsx b/ui/components/provider-form-fields.tsx index 46a7d2ef..ae95e394 100644 --- a/ui/components/provider-form-fields.tsx +++ b/ui/components/provider-form-fields.tsx @@ -202,26 +202,34 @@ export function ProviderFormFields({
- setFormData((prev) => ({ - ...prev, - provider_fee: e.target.value - ? parseFloat(e.target.value) - : undefined, - })) + value={ + (mode === 'edit' + ? formData.provider_fee_default + : formData.provider_fee) || '' } + onChange={(e) => { + const val = e.target.value ? parseFloat(e.target.value) : undefined; + setFormData((prev) => + mode === 'edit' + ? { ...prev, provider_fee_default: val } + : { ...prev, provider_fee: val } + ); + }} placeholder={providerFeePlaceholder} />

- 1.01 means +1% e.g. currency exchange, card fees, etc. + {mode === 'edit' + ? 'This is the default fee when no schedule is active. Updates will not affect currently active scheduled fees.' + : '1.01 means +1% e.g. currency exchange, card fees, etc.'}

diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index 22889264..e69cdd1d 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -12,6 +12,12 @@ export const ProviderTypeSchema = z.object({ can_show_balance: z.boolean(), }); +export const FeeTimeRangeSchema = z.object({ + start_time: z.string(), + end_time: z.string(), + provider_fee: z.number(), +}); + export const UpstreamProviderSchema = z.object({ id: z.number(), provider_type: z.string(), @@ -20,7 +26,9 @@ export const UpstreamProviderSchema = z.object({ api_version: z.string().nullable().optional(), enabled: z.boolean(), provider_fee: z.number().optional(), + provider_fee_default: z.number().optional(), provider_settings: z.record(z.string(), z.any()).nullable().optional(), + provider_fee_schedules: z.array(FeeTimeRangeSchema).optional().default([]), }); export const CreateUpstreamProviderSchema = z.object({ @@ -30,6 +38,7 @@ export const CreateUpstreamProviderSchema = z.object({ api_version: z.string().nullable().optional(), enabled: z.boolean().default(true), provider_fee: z.number().optional(), + provider_fee_default: z.number().optional(), provider_settings: z.record(z.string(), z.any()).nullable().optional(), }); @@ -40,6 +49,7 @@ export const UpdateUpstreamProviderSchema = z.object({ api_version: z.string().nullable().optional(), enabled: z.boolean().optional(), provider_fee: z.number().optional(), + provider_fee_default: z.number().optional(), provider_settings: z.record(z.string(), z.any()).nullable().optional(), }); @@ -97,6 +107,7 @@ export type CreateUpstreamProvider = z.infer< export type UpdateUpstreamProvider = z.infer< typeof UpdateUpstreamProviderSchema >; +export type FeeTimeRange = z.infer; export type AdminModel = z.infer; export type AdminModelPricing = z.infer; export type AdminModelArchitecture = z.infer< @@ -308,6 +319,30 @@ export class AdminService { ); } + static async getFeeSchedules(providerId: number): Promise { + return await apiClient.get( + `/admin/api/upstream-providers/${providerId}/fee-schedules` + ); + } + + static async updateFeeSchedules( + providerId: number, + schedules: FeeTimeRange[] + ): Promise { + return await apiClient.put( + `/admin/api/upstream-providers/${providerId}/fee-schedules`, + { schedules } + ); + } + + static async deleteFeeSchedules( + providerId: number + ): Promise<{ ok: boolean }> { + return await apiClient.delete<{ ok: boolean }>( + `/admin/api/upstream-providers/${providerId}/fee-schedules` + ); + } + static async getProviderModels(providerId: number): Promise { const data = await apiClient.get( `/admin/api/upstream-providers/${providerId}/models`