ai fix tests

This commit is contained in:
shroominic
2025-05-28 11:57:01 +00:00
parent 5be70738db
commit 6c0ca7f697
5 changed files with 60 additions and 27 deletions
+1 -1
View File
@@ -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
-1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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