mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-22 12:22:20 +00:00
Compare commits
16 Commits
rollback-t
...
dynamic-pr
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
58b468018a | ||
|
|
2b79831c61 | ||
|
|
724de338f2 | ||
|
|
c98cc30fd4 | ||
|
|
4691294b24 | ||
|
|
111060c004 | ||
|
|
f2b700a2e6 | ||
|
|
9008b00fc9 | ||
|
|
d7467dbad2 | ||
|
|
1190a09acb | ||
|
|
29f52116c3 | ||
|
|
85532f32a1 | ||
|
|
ea8257c44b | ||
|
|
d89e740ff8 | ||
|
|
af4be5bfec | ||
|
|
df2a925577 |
@@ -0,0 +1,57 @@
|
||||
"""add provider_fee_schedules and provider_fee_default to upstream_providers
|
||||
|
||||
Revision ID: 6d2fa295fa43
|
||||
Revises: cli_tokens_001
|
||||
Create Date: 2026-04-28 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "6d2fa295fa43"
|
||||
down_revision = "cli_tokens_001"
|
||||
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_default" not in columns:
|
||||
op.add_column(
|
||||
"upstream_providers",
|
||||
sa.Column(
|
||||
"provider_fee_default",
|
||||
sa.Float(),
|
||||
nullable=False,
|
||||
server_default="1.01",
|
||||
),
|
||||
)
|
||||
# Preserve any custom per-provider fees by copying from provider_fee.
|
||||
op.execute(
|
||||
"UPDATE upstream_providers "
|
||||
"SET provider_fee_default = provider_fee "
|
||||
"WHERE provider_fee IS NOT NULL"
|
||||
)
|
||||
|
||||
if "provider_fee_schedules" not in columns:
|
||||
op.add_column(
|
||||
"upstream_providers",
|
||||
sa.Column("provider_fee_schedules", sa.Text(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
|
||||
|
||||
if "provider_fee_schedules" in columns:
|
||||
op.drop_column("upstream_providers", "provider_fee_schedules")
|
||||
if "provider_fee_default" in columns:
|
||||
op.drop_column("upstream_providers", "provider_fee_default")
|
||||
@@ -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,
|
||||
@@ -639,6 +639,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
|
||||
|
||||
|
||||
@@ -649,29 +650,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)])
|
||||
@@ -698,6 +707,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,
|
||||
@@ -707,17 +719,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(
|
||||
@@ -728,18 +730,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(
|
||||
@@ -765,6 +756,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)
|
||||
|
||||
@@ -773,19 +766,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(
|
||||
@@ -803,6 +784,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."""
|
||||
|
||||
@@ -220,11 +220,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"},
|
||||
|
||||
113
routstr/payment/fee_schedule.py
Normal file
113
routstr/payment/fee_schedule.py
Normal file
@@ -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
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -118,6 +120,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:
|
||||
@@ -133,6 +141,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
|
||||
@@ -140,6 +181,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
|
||||
|
||||
@@ -67,6 +67,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):
|
||||
@@ -81,6 +82,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
|
||||
@@ -3782,8 +3784,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(
|
||||
@@ -3797,6 +3819,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,
|
||||
@@ -3804,34 +3829,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(
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
190
tests/integration/test_model_price_sync.py
Normal file
190
tests/integration/test_model_price_sync.py
Normal file
@@ -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
|
||||
@@ -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]:
|
||||
|
||||
402
tests/integration/test_provider_fee_schedules.py
Normal file
402
tests/integration/test_provider_fee_schedules.py
Normal file
@@ -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
|
||||
256
tests/unit/test_fee_schedule.py
Normal file
256
tests/unit/test_fee_schedule.py
Normal file
@@ -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
|
||||
@@ -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<UpstreamProvider | null>(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={
|
||||
<DialogTrigger asChild>
|
||||
<Button>
|
||||
<Plus className='h-4 w-4' />
|
||||
Add Provider
|
||||
<div className='flex gap-2'>
|
||||
<Button
|
||||
variant='outline'
|
||||
onClick={() => handleManageFeeSchedules()}
|
||||
disabled={providers.length === 0}
|
||||
>
|
||||
<Clock className='h-4 w-4' />
|
||||
Fee Schedules
|
||||
</Button>
|
||||
</DialogTrigger>
|
||||
<DialogTrigger asChild>
|
||||
<Button>
|
||||
<Plus className='h-4 w-4' />
|
||||
Add Provider
|
||||
</Button>
|
||||
</DialogTrigger>
|
||||
</div>
|
||||
}
|
||||
/>
|
||||
<ProviderFormDialogContent
|
||||
@@ -437,6 +460,9 @@ export default function ProvidersPage() {
|
||||
onEditProvider={() => 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() {
|
||||
}}
|
||||
/>
|
||||
)}
|
||||
|
||||
<ProviderFeeScheduleModal
|
||||
providers={providers}
|
||||
initialSelectedIds={feeScheduleState.initialIds}
|
||||
isOpen={feeScheduleState.open}
|
||||
onClose={() => setFeeScheduleState({ open: false, initialIds: [] })}
|
||||
onSuccess={() => {
|
||||
queryClient.invalidateQueries({ queryKey: ['upstream-providers'] });
|
||||
}}
|
||||
/>
|
||||
</AppPageShell>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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'}
|
||||
</Badge>
|
||||
<Badge variant='outline' className='w-fit'>
|
||||
Fee: {provider.provider_fee}x
|
||||
{provider.provider_fee !== provider.provider_fee_default && (
|
||||
<span className='text-muted-foreground ml-1 font-normal'>
|
||||
(default: {provider.provider_fee_default}x)
|
||||
</span>
|
||||
)}
|
||||
</Badge>
|
||||
</div>
|
||||
<CardDescription className='break-all'>
|
||||
{provider.base_url}
|
||||
@@ -189,6 +200,22 @@ export function ProviderCard({
|
||||
)}
|
||||
</Button>
|
||||
|
||||
<Button
|
||||
variant='outline'
|
||||
size='sm'
|
||||
onClick={onManageFeeSchedules}
|
||||
className='justify-center gap-1.5'
|
||||
title='Manage fee schedules'
|
||||
>
|
||||
<Clock className='h-4 w-4' />
|
||||
<span>Fees</span>
|
||||
{(provider.provider_fee_schedules?.length ?? 0) > 0 && (
|
||||
<Badge variant='secondary' className='ml-0.5 h-4 px-1 text-xs'>
|
||||
{provider.provider_fee_schedules!.length}
|
||||
</Badge>
|
||||
)}
|
||||
</Button>
|
||||
|
||||
<Button
|
||||
variant='outline'
|
||||
size='sm'
|
||||
|
||||
500
ui/components/provider-fee-schedule-modal.tsx
Normal file
500
ui/components/provider-fee-schedule-modal.tsx
Normal file
@@ -0,0 +1,500 @@
|
||||
'use client';
|
||||
|
||||
import { useEffect, useState } from 'react';
|
||||
import { useQueryClient } from '@tanstack/react-query';
|
||||
import { toast } from 'sonner';
|
||||
import { Plus, Trash2 } from 'lucide-react';
|
||||
import {
|
||||
Dialog,
|
||||
DialogContent,
|
||||
DialogDescription,
|
||||
DialogFooter,
|
||||
DialogHeader,
|
||||
DialogTitle,
|
||||
} from '@/components/ui/dialog';
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { Input } from '@/components/ui/input';
|
||||
import { Label } from '@/components/ui/label';
|
||||
import { Badge } from '@/components/ui/badge';
|
||||
import { Checkbox } from '@/components/ui/checkbox';
|
||||
import {
|
||||
AdminService,
|
||||
FeeTimeRange,
|
||||
UpstreamProvider,
|
||||
} from '@/lib/api/services/admin';
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Types
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export interface ProviderFeeScheduleModalProps {
|
||||
providers: UpstreamProvider[];
|
||||
/** Pre-selected provider IDs (e.g. clicked from a card). Empty = all selected. */
|
||||
initialSelectedIds?: number[];
|
||||
isOpen: boolean;
|
||||
onClose: () => void;
|
||||
onSuccess: () => void;
|
||||
}
|
||||
|
||||
interface RangeRow extends FeeTimeRange {
|
||||
_id: number;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Overlap helpers (mirrored from backend logic)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
function _toMinutes(t: string): number {
|
||||
const [h, m] = t.split(':').map(Number);
|
||||
return h * 60 + m;
|
||||
}
|
||||
|
||||
function _rangeIntervals(start: string, end: string): Array<[number, number]> {
|
||||
const s = _toMinutes(start);
|
||||
const e = _toMinutes(end);
|
||||
if (s < e) return [[s, e]];
|
||||
if (s > e)
|
||||
return [
|
||||
[s, 1440],
|
||||
[0, e],
|
||||
];
|
||||
return [[0, 1440]];
|
||||
}
|
||||
|
||||
function _intervalsOverlap(a: [number, number], b: [number, number]): boolean {
|
||||
return a[0] < b[1] && b[0] < a[1];
|
||||
}
|
||||
|
||||
function findOverlappingIds(rows: RangeRow[]): Set<number> {
|
||||
const overlapping = new Set<number>();
|
||||
for (let i = 0; i < rows.length; i++) {
|
||||
for (let j = i + 1; j < rows.length; j++) {
|
||||
const a = rows[i];
|
||||
const b = rows[j];
|
||||
if (!a.start_time || !a.end_time || !b.start_time || !b.end_time)
|
||||
continue;
|
||||
for (const ia of _rangeIntervals(a.start_time, a.end_time)) {
|
||||
for (const ib of _rangeIntervals(b.start_time, b.end_time)) {
|
||||
if (_intervalsOverlap(ia, ib)) {
|
||||
overlapping.add(a._id);
|
||||
overlapping.add(b._id);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return overlapping;
|
||||
}
|
||||
|
||||
function isValidTime(t: string): boolean {
|
||||
return /^([01]\d|2[0-3]):([0-5]\d)$/.test(t);
|
||||
}
|
||||
|
||||
function utcTimeNow(): string {
|
||||
const now = new Date();
|
||||
return now.toUTCString().slice(17, 22);
|
||||
}
|
||||
|
||||
// Browsers may return "HH:MM:SS" from time inputs — strip seconds.
|
||||
function normalizeTime(v: string): string {
|
||||
return v.slice(0, 5);
|
||||
}
|
||||
|
||||
let _nextId = 1;
|
||||
|
||||
function makeRow(partial: Partial<FeeTimeRange> = {}): RangeRow {
|
||||
return {
|
||||
_id: _nextId++,
|
||||
start_time: partial.start_time ?? '',
|
||||
end_time: partial.end_time ?? '',
|
||||
provider_fee: partial.provider_fee ?? 1.05,
|
||||
};
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Sub-component: read-only range list under a provider
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
function ProviderRangePreview({ schedules }: { schedules: FeeTimeRange[] }) {
|
||||
if (schedules.length === 0) {
|
||||
return (
|
||||
<p className='text-muted-foreground pl-7 text-xs'>
|
||||
No scheduled ranges — default fee always applies.
|
||||
</p>
|
||||
);
|
||||
}
|
||||
return (
|
||||
<ul className='space-y-0.5 pl-7'>
|
||||
{schedules.map((s, i) => (
|
||||
<li key={i} className='flex items-center gap-2 text-xs'>
|
||||
<span className='text-muted-foreground font-mono'>
|
||||
{s.start_time} → {s.end_time} UTC
|
||||
</span>
|
||||
<Badge variant='outline' className='py-0 font-mono text-xs'>
|
||||
×{s.provider_fee.toFixed(3)}
|
||||
</Badge>
|
||||
</li>
|
||||
))}
|
||||
</ul>
|
||||
);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Modal
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export function ProviderFeeScheduleModal({
|
||||
providers,
|
||||
initialSelectedIds,
|
||||
isOpen,
|
||||
onClose,
|
||||
onSuccess,
|
||||
}: ProviderFeeScheduleModalProps) {
|
||||
const queryClient = useQueryClient();
|
||||
|
||||
const [selectedIds, setSelectedIds] = useState<Set<number>>(new Set());
|
||||
const [rows, setRows] = useState<RangeRow[]>([]);
|
||||
const [enforceOverride, setEnforceOverride] = useState(false);
|
||||
const [saving, setSaving] = useState(false);
|
||||
const [clearing, setClearing] = useState(false);
|
||||
|
||||
// Reset state when modal opens. If a single provider is pre-selected,
|
||||
// pre-populate the editor with its existing schedule so the user can edit it.
|
||||
useEffect(() => {
|
||||
if (!isOpen) return;
|
||||
|
||||
const ids =
|
||||
initialSelectedIds && initialSelectedIds.length > 0
|
||||
? new Set(initialSelectedIds)
|
||||
: new Set(providers.map((p) => p.id));
|
||||
|
||||
setSelectedIds(ids);
|
||||
setEnforceOverride(false);
|
||||
|
||||
if (initialSelectedIds && initialSelectedIds.length === 1) {
|
||||
const provider = providers.find((p) => p.id === initialSelectedIds[0]);
|
||||
const existing = provider?.provider_fee_schedules ?? [];
|
||||
setRows(existing.length > 0 ? existing.map((s) => makeRow(s)) : []);
|
||||
setEnforceOverride(true);
|
||||
} else {
|
||||
setRows([]);
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [isOpen]);
|
||||
|
||||
const overlapping = findOverlappingIds(rows);
|
||||
const allSelected =
|
||||
providers.length > 0 && selectedIds.size === providers.length;
|
||||
const noneSelected = selectedIds.size === 0;
|
||||
|
||||
const toggleProvider = (id: number) => {
|
||||
setSelectedIds((prev) => {
|
||||
const next = new Set(prev);
|
||||
next.has(id) ? next.delete(id) : next.add(id);
|
||||
return next;
|
||||
});
|
||||
};
|
||||
|
||||
const toggleAll = () => {
|
||||
setSelectedIds(
|
||||
allSelected ? new Set() : new Set(providers.map((p) => p.id))
|
||||
);
|
||||
};
|
||||
|
||||
const addRow = () => setRows((prev) => [...prev, makeRow()]);
|
||||
|
||||
const removeRow = (id: number) =>
|
||||
setRows((prev) => prev.filter((r) => r._id !== id));
|
||||
|
||||
const updateRow = (
|
||||
id: number,
|
||||
field: keyof FeeTimeRange,
|
||||
value: string | number
|
||||
) =>
|
||||
setRows((prev) =>
|
||||
prev.map((r) => (r._id === id ? { ...r, [field]: value } : r))
|
||||
);
|
||||
|
||||
const hasValidationErrors =
|
||||
noneSelected ||
|
||||
rows.some(
|
||||
(r) =>
|
||||
!isValidTime(r.start_time) ||
|
||||
!isValidTime(r.end_time) ||
|
||||
r.provider_fee <= 0
|
||||
) ||
|
||||
overlapping.size > 0;
|
||||
|
||||
const handleSave = async () => {
|
||||
if (hasValidationErrors) return;
|
||||
const newSchedules: FeeTimeRange[] = rows.map(
|
||||
({ start_time, end_time, provider_fee }) => ({
|
||||
start_time,
|
||||
end_time,
|
||||
provider_fee,
|
||||
})
|
||||
);
|
||||
setSaving(true);
|
||||
try {
|
||||
await Promise.all(
|
||||
[...selectedIds].map((id) => {
|
||||
const provider = providers.find((p) => p.id === id);
|
||||
const existing = provider?.provider_fee_schedules ?? [];
|
||||
|
||||
let finalSchedules: FeeTimeRange[];
|
||||
if (enforceOverride) {
|
||||
finalSchedules = newSchedules;
|
||||
} else {
|
||||
// Only override ranges that overlap with ANY of the new ranges.
|
||||
// Keep existing non-overlapping ranges.
|
||||
const keptExisting = existing.filter((ex) => {
|
||||
return !newSchedules.some((nw) =>
|
||||
_rangeIntervals(ex.start_time, ex.end_time).some((ia) =>
|
||||
_rangeIntervals(nw.start_time, nw.end_time).some((ib) =>
|
||||
_intervalsOverlap(ia, ib)
|
||||
)
|
||||
)
|
||||
);
|
||||
});
|
||||
finalSchedules = [...keptExisting, ...newSchedules];
|
||||
}
|
||||
|
||||
return AdminService.updateFeeSchedules(id, finalSchedules);
|
||||
})
|
||||
);
|
||||
queryClient.invalidateQueries({ queryKey: ['upstream-providers'] });
|
||||
toast.success(
|
||||
`Fee schedules saved for ${selectedIds.size} provider${selectedIds.size > 1 ? 's' : ''}`
|
||||
);
|
||||
onSuccess();
|
||||
onClose();
|
||||
} catch (err) {
|
||||
toast.error(
|
||||
`Failed to save: ${err instanceof Error ? err.message : 'Unknown error'}`
|
||||
);
|
||||
} finally {
|
||||
setSaving(false);
|
||||
}
|
||||
};
|
||||
|
||||
const handleClearAll = async () => {
|
||||
if (noneSelected) return;
|
||||
setClearing(true);
|
||||
try {
|
||||
await Promise.all(
|
||||
[...selectedIds].map((id) => AdminService.deleteFeeSchedules(id))
|
||||
);
|
||||
queryClient.invalidateQueries({ queryKey: ['upstream-providers'] });
|
||||
setRows([]);
|
||||
toast.success(
|
||||
`Fee schedules cleared for ${selectedIds.size} provider${selectedIds.size > 1 ? 's' : ''}`
|
||||
);
|
||||
onSuccess();
|
||||
} catch (err) {
|
||||
toast.error(
|
||||
`Failed to clear: ${err instanceof Error ? err.message : 'Unknown error'}`
|
||||
);
|
||||
} finally {
|
||||
setClearing(false);
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<Dialog open={isOpen} onOpenChange={(open) => !open && onClose()}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[660px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle>Fee Schedules</DialogTitle>
|
||||
<DialogDescription>
|
||||
Select providers and configure time-based fee ranges (UTC). Outside
|
||||
scheduled ranges each provider's default fee applies. Current
|
||||
UTC time:{' '}
|
||||
<Badge variant='outline' className='font-mono'>
|
||||
{utcTimeNow()}
|
||||
</Badge>
|
||||
</DialogDescription>
|
||||
</DialogHeader>
|
||||
|
||||
{/* Provider selection with existing-range read view */}
|
||||
<div className='space-y-2'>
|
||||
<div className='flex items-center justify-between'>
|
||||
<Label className='text-sm font-medium'>Apply to providers</Label>
|
||||
<button
|
||||
onClick={toggleAll}
|
||||
className='text-muted-foreground hover:text-foreground text-xs underline-offset-2 hover:underline'
|
||||
>
|
||||
{allSelected ? 'Deselect all' : 'Select all'}
|
||||
</button>
|
||||
</div>
|
||||
<div className='divide-y rounded-md border'>
|
||||
{providers.map((p) => (
|
||||
<div key={p.id} className='space-y-1.5 px-3 py-2'>
|
||||
<label className='hover:bg-muted/50 flex cursor-pointer items-center gap-3 rounded'>
|
||||
<Checkbox
|
||||
checked={selectedIds.has(p.id)}
|
||||
onCheckedChange={() => toggleProvider(p.id)}
|
||||
/>
|
||||
<span className='flex-1 text-sm font-medium'>
|
||||
{p.provider_type}
|
||||
</span>
|
||||
<span className='text-muted-foreground truncate text-xs'>
|
||||
{p.base_url}
|
||||
</span>
|
||||
</label>
|
||||
<ProviderRangePreview
|
||||
schedules={p.provider_fee_schedules ?? []}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
{noneSelected && (
|
||||
<p className='text-destructive text-xs'>
|
||||
Select at least one provider.
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Fee range editor */}
|
||||
<div className='space-y-4'>
|
||||
<div className='flex items-center justify-between'>
|
||||
<Label className='text-sm font-medium'>
|
||||
New schedule{' '}
|
||||
{!enforceOverride && (
|
||||
<span className='text-muted-foreground font-normal'>
|
||||
(merges with existing ranges, overriding only overlaps)
|
||||
</span>
|
||||
)}
|
||||
</Label>
|
||||
<div className='flex items-center space-x-2'>
|
||||
<Checkbox
|
||||
id='enforce-override'
|
||||
checked={enforceOverride}
|
||||
onCheckedChange={(checked) => setEnforceOverride(!!checked)}
|
||||
/>
|
||||
<label
|
||||
htmlFor='enforce-override'
|
||||
className='text-xs leading-none font-medium peer-disabled:cursor-not-allowed peer-disabled:opacity-70'
|
||||
>
|
||||
Enforce overriding everything
|
||||
</label>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className='space-y-2'>
|
||||
{rows.length === 0 && (
|
||||
<p className='text-muted-foreground rounded-md border border-dashed p-4 text-center text-sm'>
|
||||
No ranges configured — saving with no ranges will clear
|
||||
schedules.
|
||||
</p>
|
||||
)}
|
||||
|
||||
{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 (
|
||||
<div
|
||||
key={row._id}
|
||||
className={`flex flex-col gap-2 rounded-md border p-3 sm:flex-row sm:items-end ${
|
||||
hasError ? 'border-destructive bg-destructive/5' : ''
|
||||
}`}
|
||||
>
|
||||
<div className='flex flex-1 flex-col gap-1'>
|
||||
<Label className='text-xs'>Start (UTC)</Label>
|
||||
<input
|
||||
type='time'
|
||||
value={row.start_time}
|
||||
onChange={(e) =>
|
||||
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'
|
||||
/>
|
||||
</div>
|
||||
<div className='flex flex-1 flex-col gap-1'>
|
||||
<Label className='text-xs'>End (UTC)</Label>
|
||||
<input
|
||||
type='time'
|
||||
value={row.end_time}
|
||||
onChange={(e) =>
|
||||
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'
|
||||
/>
|
||||
</div>
|
||||
<div className='flex flex-1 flex-col gap-1'>
|
||||
<Label className='text-xs'>Fee multiplier</Label>
|
||||
<Input
|
||||
type='number'
|
||||
step='0.001'
|
||||
min='0.001'
|
||||
placeholder='1.05'
|
||||
value={row.provider_fee}
|
||||
onChange={(e) =>
|
||||
updateRow(
|
||||
row._id,
|
||||
'provider_fee',
|
||||
parseFloat(e.target.value) || 0
|
||||
)
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
<Button
|
||||
variant='ghost'
|
||||
size='icon'
|
||||
className='text-destructive hover:text-destructive shrink-0'
|
||||
onClick={() => removeRow(row._id)}
|
||||
>
|
||||
<Trash2 className='h-4 w-4' />
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
|
||||
{overlapping.size > 0 && (
|
||||
<p className='text-destructive text-xs'>
|
||||
Some ranges overlap — fix them before saving.
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<Button variant='outline' size='sm' onClick={addRow}>
|
||||
<Plus className='mr-1.5 h-4 w-4' />
|
||||
Add Range
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
<DialogFooter className='gap-2'>
|
||||
<Button
|
||||
variant='ghost'
|
||||
onClick={handleClearAll}
|
||||
disabled={clearing || noneSelected}
|
||||
className='text-destructive hover:text-destructive mr-auto'
|
||||
>
|
||||
{clearing ? 'Clearing…' : 'Clear Selected'}
|
||||
</Button>
|
||||
<Button variant='outline' onClick={onClose}>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button onClick={handleSave} disabled={saving || hasValidationErrors}>
|
||||
{saving
|
||||
? 'Saving…'
|
||||
: `Save to ${selectedIds.size} provider${selectedIds.size !== 1 ? 's' : ''}`}
|
||||
</Button>
|
||||
</DialogFooter>
|
||||
</DialogContent>
|
||||
</Dialog>
|
||||
);
|
||||
}
|
||||
@@ -202,26 +202,34 @@ export function ProviderFormFields({
|
||||
|
||||
<div className='grid gap-2'>
|
||||
<Label htmlFor={`${idPrefix}provider_fee`}>
|
||||
Provider Fee (Multiplier)
|
||||
{mode === 'edit'
|
||||
? 'Default Provider Fee (Multiplier)'
|
||||
: 'Provider Fee (Multiplier)'}
|
||||
</Label>
|
||||
<Input
|
||||
id={`${idPrefix}provider_fee`}
|
||||
type='number'
|
||||
step='0.001'
|
||||
min='1.0'
|
||||
value={formData.provider_fee || ''}
|
||||
onChange={(e) =>
|
||||
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}
|
||||
/>
|
||||
<p className='text-muted-foreground text-xs'>
|
||||
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.'}
|
||||
</p>
|
||||
</div>
|
||||
|
||||
|
||||
@@ -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<typeof FeeTimeRangeSchema>;
|
||||
export type AdminModel = z.infer<typeof AdminModelSchema>;
|
||||
export type AdminModelPricing = z.infer<typeof AdminModelPricingSchema>;
|
||||
export type AdminModelArchitecture = z.infer<
|
||||
@@ -308,6 +319,30 @@ export class AdminService {
|
||||
);
|
||||
}
|
||||
|
||||
static async getFeeSchedules(providerId: number): Promise<FeeTimeRange[]> {
|
||||
return await apiClient.get<FeeTimeRange[]>(
|
||||
`/admin/api/upstream-providers/${providerId}/fee-schedules`
|
||||
);
|
||||
}
|
||||
|
||||
static async updateFeeSchedules(
|
||||
providerId: number,
|
||||
schedules: FeeTimeRange[]
|
||||
): Promise<FeeTimeRange[]> {
|
||||
return await apiClient.put<FeeTimeRange[]>(
|
||||
`/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<ProviderModels> {
|
||||
const data = await apiClient.get<ProviderModels>(
|
||||
`/admin/api/upstream-providers/${providerId}/models`
|
||||
|
||||
Reference in New Issue
Block a user