mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
ai fix tests
This commit is contained in:
+1
-1
@@ -97,7 +97,7 @@ def test_client() -> TestClient:
|
||||
|
||||
with patch("router.models.update_sats_pricing") as mock_update:
|
||||
mock_update.return_value = None
|
||||
return TestClient(app)
|
||||
yield TestClient(app)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
|
||||
@@ -188,7 +188,6 @@ async def test_account_with_cashu_token(
|
||||
):
|
||||
"""Test authentication with a cashu token creates a new account."""
|
||||
cashu_token = "cashuBqQSEQ123456"
|
||||
hashed = hash_api_key(cashu_token)
|
||||
|
||||
with patch("router.cashu.credit_balance", new_callable=AsyncMock) as mock_credit:
|
||||
# Mock successful token redemption
|
||||
|
||||
+10
-14
@@ -7,20 +7,16 @@ from unittest.mock import patch
|
||||
async def test_root_endpoint(async_client: AsyncClient):
|
||||
"""Test the root endpoint returns expected information."""
|
||||
# Mock the environment variables for this specific test
|
||||
with patch("os.environ.get") as mock_env_get:
|
||||
def env_side_effect(key, default=None):
|
||||
env_map = {
|
||||
"NAME": "TestRoutstrNode",
|
||||
"DESCRIPTION": "Test Node",
|
||||
"NPUB": "npub1test",
|
||||
"MINT": "https://test.mint.com",
|
||||
"HTTP_URL": "http://test.example.com",
|
||||
"ONION_URL": "http://test.onion",
|
||||
}
|
||||
return env_map.get(key, default)
|
||||
|
||||
mock_env_get.side_effect = env_side_effect
|
||||
|
||||
env_vars = {
|
||||
"NAME": "TestRoutstrNode",
|
||||
"DESCRIPTION": "Test Node",
|
||||
"NPUB": "npub1test",
|
||||
"MINT": "https://test.mint.com",
|
||||
"HTTP_URL": "http://test.example.com",
|
||||
"ONION_URL": "http://test.onion",
|
||||
}
|
||||
|
||||
with patch.dict("os.environ", env_vars, clear=False):
|
||||
response = await async_client.get("/")
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
+30
-9
@@ -67,14 +67,21 @@ async def test_update_sats_pricing_calculation(sample_model: Model):
|
||||
assert sample_model.sats_pricing is not None
|
||||
|
||||
# Verify calculations (prices in USD / sats_to_usd)
|
||||
assert sample_model.sats_pricing.prompt == 0.01 / 0.0001 # 100 sats
|
||||
assert sample_model.sats_pricing.completion == 0.02 / 0.0001 # 200 sats
|
||||
assert sample_model.sats_pricing.request == 0.001 / 0.0001 # 10 sats
|
||||
assert sample_model.sats_pricing.prompt == pytest.approx(0.01 / 0.0001) # 100 sats
|
||||
assert sample_model.sats_pricing.completion == pytest.approx(0.02 / 0.0001) # 200 sats
|
||||
assert sample_model.sats_pricing.request == pytest.approx(0.001 / 0.0001) # 10 sats
|
||||
|
||||
# Verify max_cost calculation for model with top_provider
|
||||
expected_max_context = 4096 * sample_model.sats_pricing.prompt
|
||||
expected_max_completion = 2048 * sample_model.sats_pricing.completion
|
||||
assert sample_model.sats_pricing.max_cost == expected_max_context + expected_max_completion
|
||||
assert sample_model.sats_pricing.max_cost == pytest.approx(expected_max_context + expected_max_completion)
|
||||
|
||||
# Cancel and await the task
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
@@ -140,7 +147,14 @@ async def test_update_sats_pricing_without_top_provider():
|
||||
ir = model_without_top.sats_pricing.internal_reasoning * 100
|
||||
expected_max = p + c + r + i + w + ir
|
||||
|
||||
assert model_without_top.sats_pricing.max_cost == expected_max
|
||||
assert model_without_top.sats_pricing.max_cost == pytest.approx(expected_max)
|
||||
|
||||
# Cancel and await the task
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
@@ -158,11 +172,11 @@ async def test_update_sats_pricing_handles_errors():
|
||||
error_printed = False
|
||||
original_print = print
|
||||
|
||||
def mock_print(msg):
|
||||
def mock_print(*args, **kwargs):
|
||||
nonlocal error_printed
|
||||
if isinstance(msg, Exception) and str(msg) == "API Error":
|
||||
if args and isinstance(args[0], Exception) and str(args[0]) == "API Error":
|
||||
error_printed = True
|
||||
original_print(msg)
|
||||
original_print(*args, **kwargs)
|
||||
|
||||
with patch("builtins.print", side_effect=mock_print):
|
||||
sleep_called = asyncio.Event()
|
||||
@@ -179,6 +193,13 @@ async def test_update_sats_pricing_handles_errors():
|
||||
# Verify error was printed
|
||||
assert error_printed
|
||||
|
||||
# Cancel and await the task
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
@@ -197,4 +218,4 @@ def test_model_serialization(sample_model: Model):
|
||||
# Test deserialization
|
||||
new_model = Model(**model_dict)
|
||||
assert new_model.id == sample_model.id
|
||||
assert new_model.pricing.prompt == sample_model.pricing.prompt
|
||||
assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt)
|
||||
+19
-2
@@ -113,6 +113,10 @@ async def test_proxy_successful_request_mock(
|
||||
mock_client = AsyncMock()
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
# Add async context manager methods
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock()
|
||||
|
||||
# Create a mock response
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
@@ -175,10 +179,14 @@ async def test_proxy_streaming_response(
|
||||
mock_client = AsyncMock()
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
# Add async context manager methods
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock()
|
||||
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "text/event-stream"}
|
||||
mock_response.aiter_bytes = mock_aiter_bytes
|
||||
mock_response.aiter_bytes = lambda: mock_aiter_bytes()
|
||||
mock_response.aclose = AsyncMock()
|
||||
|
||||
mock_client.send = AsyncMock(return_value=mock_response)
|
||||
@@ -213,6 +221,10 @@ async def test_proxy_handles_upstream_errors(
|
||||
mock_client = AsyncMock()
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
# Add async context manager methods
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock()
|
||||
|
||||
# Simulate connection error
|
||||
mock_client.send.side_effect = Exception("Connection refused")
|
||||
mock_client.build_request = AsyncMock()
|
||||
@@ -252,7 +264,8 @@ async def test_proxy_with_model_based_pricing(
|
||||
test_session.add(key)
|
||||
await test_session.commit()
|
||||
|
||||
with patch.dict(os.environ, {"MODEL_BASED_PRICING": "true"}):
|
||||
# Patch the MODEL_BASED_PRICING constant directly
|
||||
with patch("router.auth.MODEL_BASED_PRICING", True):
|
||||
with patch("os.path.exists", return_value=True):
|
||||
# Mock a model with pricing
|
||||
from router.models import MODELS, Model, Pricing, Architecture, TopProvider
|
||||
@@ -304,6 +317,10 @@ async def test_proxy_with_model_based_pricing(
|
||||
mock_client = AsyncMock()
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
# Add async context manager methods
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock()
|
||||
|
||||
# Create a mock response
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
|
||||
Reference in New Issue
Block a user