From 6c0ca7f697b94369b26da82f62e19f69f15b0ecc Mon Sep 17 00:00:00 2001 From: shroominic Date: Wed, 28 May 2025 11:57:01 +0000 Subject: [PATCH] ai fix tests --- tests/conftest.py | 2 +- tests/test_account.py | 1 - tests/test_main.py | 24 ++++++++++-------------- tests/test_models.py | 39 ++++++++++++++++++++++++++++++--------- tests/test_proxy.py | 21 +++++++++++++++++++-- 5 files changed, 60 insertions(+), 27 deletions(-) diff --git a/tests/conftest.py b/tests/conftest.py index ed98bc35..c5ca43ff 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 diff --git a/tests/test_account.py b/tests/test_account.py index bdcb8f94..9105db3f 100644 --- a/tests/test_account.py +++ b/tests/test_account.py @@ -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 diff --git a/tests/test_main.py b/tests/test_main.py index 0a66632a..cee18b37 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -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 diff --git a/tests/test_models.py b/tests/test_models.py index d28a2f2f..c4a8494e 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -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 \ No newline at end of file + assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt) \ No newline at end of file diff --git a/tests/test_proxy.py b/tests/test_proxy.py index 603eb522..7386b5c9 100644 --- a/tests/test_proxy.py +++ b/tests/test_proxy.py @@ -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