mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-09 19:04:47 +00:00
128 lines
4.3 KiB
Python
128 lines
4.3 KiB
Python
import os
|
|
from typing import Any
|
|
from unittest.mock import AsyncMock, Mock, patch
|
|
|
|
# Set required env vars before importing
|
|
os.environ["UPSTREAM_BASE_URL"] = "http://test"
|
|
os.environ["UPSTREAM_API_KEY"] = "test"
|
|
|
|
from routstr.core.settings import settings # noqa: E402
|
|
from routstr.payment.helpers import get_max_cost_for_model # noqa: E402
|
|
|
|
|
|
async def test_get_max_cost_for_model_known() -> None:
|
|
from routstr.payment.models import Pricing
|
|
|
|
# Mock DB session behavior
|
|
mock_session = AsyncMock()
|
|
|
|
# Mock upstream provider rows
|
|
mock_provider_result = Mock()
|
|
mock_provider_result.all = Mock(return_value=[])
|
|
|
|
# Mock model row with proper JSON fields
|
|
row = Mock()
|
|
row.id = "gpt-4"
|
|
row.name = "GPT-4"
|
|
row.created = 1234567890
|
|
row.description = "Test model"
|
|
row.context_length = 8192
|
|
row.architecture = '{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "gpt", "instruct_type": null}'
|
|
row.pricing = '{"prompt": 0.0, "completion": 0.0, "request": 0.0, "image": 0.0, "web_search": 0.0, "internal_reasoning": 0.0, "max_cost": 0.0}'
|
|
row.per_request_limits = None
|
|
row.top_provider = None
|
|
row.enabled = True
|
|
row.upstream_provider_id = 1
|
|
|
|
# Mock the exec results to return model row when querying for override
|
|
def mock_exec(query: Any) -> Any:
|
|
result = Mock()
|
|
result.first = Mock(return_value=row)
|
|
result.all = Mock(return_value=[row])
|
|
return result
|
|
|
|
mock_session.exec = Mock(side_effect=mock_exec)
|
|
|
|
# Mock get for UpstreamProviderRow
|
|
mock_provider = Mock()
|
|
mock_provider.provider_fee = 1.01
|
|
mock_session.get = Mock(return_value=mock_provider)
|
|
|
|
# Mock the model with sats_pricing
|
|
mock_pricing = Pricing(
|
|
prompt=0.0,
|
|
completion=0.0,
|
|
request=0.0,
|
|
image=0.0,
|
|
web_search=0.0,
|
|
internal_reasoning=0.0,
|
|
max_cost=500.0,
|
|
)
|
|
mock_model = Mock()
|
|
mock_model.sats_pricing = mock_pricing
|
|
|
|
with patch.object(settings, "fixed_pricing", False):
|
|
with patch.object(settings, "tolerance_percentage", 0):
|
|
cost = await get_max_cost_for_model(
|
|
"gpt-4", session=mock_session, model_obj=mock_model
|
|
)
|
|
assert cost == 500000 # 500 sats * 1000 = msats
|
|
|
|
|
|
async def test_get_max_cost_for_model_unknown() -> None:
|
|
mock_session = AsyncMock()
|
|
|
|
# Mock the exec results to return no model override
|
|
async def async_mock_exec(query: Any) -> Any:
|
|
result = Mock()
|
|
result.first = Mock(return_value=None)
|
|
result.all = Mock(return_value=[])
|
|
return result
|
|
|
|
mock_session.exec = AsyncMock(side_effect=async_mock_exec)
|
|
mock_session.get = AsyncMock(return_value=None)
|
|
|
|
# Mock get_upstreams to return empty list
|
|
with patch("routstr.proxy.get_upstreams", return_value=[]):
|
|
with patch.object(settings, "fixed_cost_per_request", 100):
|
|
with patch.object(settings, "tolerance_percentage", 0):
|
|
cost = await get_max_cost_for_model(
|
|
"unknown-model", session=mock_session, model_obj=None
|
|
)
|
|
assert cost == 100000
|
|
|
|
|
|
async def test_get_max_cost_for_model_disabled() -> None:
|
|
mock_session = AsyncMock()
|
|
with patch.object(settings, "fixed_pricing", True):
|
|
with patch.object(settings, "fixed_cost_per_request", 200):
|
|
with patch.object(settings, "tolerance_percentage", 0):
|
|
cost = await get_max_cost_for_model("any-model", session=mock_session)
|
|
assert cost == 200000
|
|
|
|
|
|
async def test_get_max_cost_for_model_tolerance() -> None:
|
|
from routstr.payment.models import Pricing
|
|
|
|
mock_session = AsyncMock()
|
|
|
|
# Mock the model with sats_pricing
|
|
mock_pricing = Pricing(
|
|
prompt=0.0,
|
|
completion=0.0,
|
|
request=0.0,
|
|
image=0.0,
|
|
web_search=0.0,
|
|
internal_reasoning=0.0,
|
|
max_cost=500.0,
|
|
)
|
|
mock_model = Mock()
|
|
mock_model.sats_pricing = mock_pricing
|
|
|
|
with patch.object(settings, "fixed_pricing", False):
|
|
with patch.object(settings, "tolerance_percentage", 10):
|
|
cost = await get_max_cost_for_model(
|
|
"gpt-4", session=mock_session, model_obj=mock_model
|
|
)
|
|
assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000
|