fix mypy issues

This commit is contained in:
Shroominic
2025-06-13 17:06:08 +02:00
parent 47c8c34d5c
commit 324c406ed0
13 changed files with 147 additions and 122 deletions
+5 -3
View File
@@ -13,7 +13,7 @@ client = openai.OpenAI(
history: list = []
def chat():
def chat() -> None:
while True:
user_msg = {"role": "user", "content": input("\nYou: ")}
history.append(user_msg)
@@ -25,8 +25,10 @@ def chat():
stream=True,
):
if len(chunk.choices) > 0:
ai_msg["content"] += chunk.choices[0].delta.content
print(chunk.choices[0].delta.content, end="", flush=True)
content = chunk.choices[0].delta.content
if content is not None:
ai_msg["content"] += content
print(content, end="", flush=True)
print()
history.append(ai_msg)
+9 -5
View File
@@ -1,4 +1,4 @@
from typing import Annotated
from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
@@ -49,8 +49,9 @@ async def topup_wallet_endpoint(
cashu_token: str,
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
):
return await credit_balance(cashu_token, key, session)
) -> dict[str, int]:
amount_msats = await credit_balance(cashu_token, key, session)
return {"msats": amount_msats}
@wallet_router.post("/refund")
@@ -88,9 +89,12 @@ async def refund_wallet_endpoint(
@wallet_router.api_route(
"/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], include_in_schema=False
"/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"],
include_in_schema=False,
response_model=None,
)
async def wallet_catch_all(path: str):
async def wallet_catch_all(path: str) -> NoReturn:
raise HTTPException(
status_code=404, detail="Not found check /docs for available endpoints"
)
+1 -1
View File
@@ -155,7 +155,7 @@ async def dashboard(request: Request) -> str:
@admin_router.get("/", response_class=HTMLResponse)
async def admin(request: Request):
async def admin(request: Request) -> str:
admin_cookie = request.cookies.get("admin_password")
if admin_cookie and admin_cookie == os.getenv("ADMIN_PASSWORD"):
return await dashboard(request)
+2 -2
View File
@@ -25,12 +25,12 @@ async def delete_key_if_zero_balance(key: ApiKey, session: AsyncSession) -> None
await session.commit()
async def init_wallet():
async def init_wallet() -> None:
global WALLET
WALLET = await Wallet.create(nsec=NSEC, mint_urls=[MINT])
async def close_wallet():
async def close_wallet() -> None:
global WALLET
await WALLET.aclose()
+1 -1
View File
@@ -139,7 +139,7 @@ async def fetch_onion(provider: str) -> dict:
@providers_router.get("/")
async def get_providers(include_json: bool = False):
async def get_providers(include_json: bool = False) -> dict[str, list[dict | str]]:
npub = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s"
# Relays that support NIP-50 text search
+3 -2
View File
@@ -1,6 +1,7 @@
import asyncio
import os
from contextlib import asynccontextmanager
from typing import AsyncGenerator
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
@@ -17,7 +18,7 @@ __version__ = "0.0.1"
@asynccontextmanager
async def lifespan(_: FastAPI):
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
await init_db()
await init_wallet()
pricing_task = asyncio.create_task(update_sats_pricing())
@@ -51,7 +52,7 @@ app.add_middleware(
@app.get("/")
async def info():
async def info() -> dict:
return {
"name": app.title,
"description": app.description,
+4 -5
View File
@@ -1,6 +1,7 @@
import json
import os
import re
from typing import AsyncGenerator
import httpx
from fastapi import APIRouter, BackgroundTasks, Depends, Request
@@ -16,12 +17,10 @@ UPSTREAM_API_KEY = os.environ.get("UPSTREAM_API_KEY", "")
proxy_router = APIRouter()
@proxy_router.api_route(
"/{path:path}", methods=["GET", "POST"]
)
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session)
):
) -> Response | StreamingResponse:
auth = request.headers.get("Authorization", "")
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
refund_address = request.headers.get("Refund-LNURL", None)
@@ -146,7 +145,7 @@ async def proxy(
if is_streaming and response.status_code == 200:
# Process streaming response and extract cost from the last chunk
async def stream_with_cost():
async def stream_with_cost() -> AsyncGenerator[bytes, None]:
# Store all chunks to analyze
stored_chunks = []
+8 -8
View File
@@ -7,7 +7,7 @@ import pytest
import pytest_asyncio
from fastapi.testclient import TestClient
from httpx import ASGITransport, AsyncClient
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
@@ -61,7 +61,7 @@ with patch("sixty_nuts.Wallet") as mock_wallet_class:
@pytest.fixture(scope="session")
def event_loop():
def event_loop() -> Generator[asyncio.AbstractEventLoop, None, None]:
"""Create an instance of the default event loop for the test session."""
loop = asyncio.get_event_loop_policy().new_event_loop()
yield loop
@@ -69,7 +69,7 @@ def event_loop():
@pytest_asyncio.fixture(scope="function")
async def test_engine():
async def test_engine() -> AsyncGenerator[AsyncEngine, None]:
"""Create a test database engine - new for each test."""
engine = create_async_engine(
"sqlite+aiosqlite:///:memory:",
@@ -86,7 +86,7 @@ async def test_engine():
@pytest_asyncio.fixture
async def test_session(test_engine) -> AsyncGenerator[AsyncSession, None]:
async def test_session(test_engine: AsyncEngine) -> AsyncGenerator[AsyncSession, None]:
"""Create a test database session."""
from sqlmodel.ext.asyncio.session import AsyncSession as SqlModelAsyncSession
@@ -123,10 +123,10 @@ def test_client() -> Generator[TestClient, None, None]:
@pytest_asyncio.fixture
async def async_client(test_session) -> AsyncGenerator[AsyncClient, None]:
async def async_client(test_session: AsyncSession) -> AsyncGenerator[AsyncClient, None]:
"""Create an async test client with dependency overrides."""
async def override_get_session():
async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
yield test_session
app.dependency_overrides[get_session] = override_get_session
@@ -164,7 +164,7 @@ async def async_client(test_session) -> AsyncGenerator[AsyncClient, None]:
@pytest.fixture
def mock_models():
def mock_models() -> list[dict]:
"""Mock models data for testing."""
return [
{
@@ -199,7 +199,7 @@ def mock_models():
# Cleanup after all tests
@pytest.fixture(scope="session", autouse=True)
def cleanup():
def cleanup() -> Generator[None, None, None]:
yield
# Restore original environment carefully
current_keys = set(os.environ.keys())
+12 -12
View File
@@ -39,10 +39,11 @@ async def test_api_key(test_session: AsyncSession) -> ApiKey:
@pytest.mark.asyncio
async def test_account_info_with_valid_key(
async_client: AsyncClient, test_api_key: ApiKey
):
) -> None:
"""Test getting account info with a valid API key."""
response = await async_client.get(
"/v1/wallet/info", headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"}
"/v1/wallet/info",
headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"},
)
assert response.status_code == 200
@@ -53,7 +54,7 @@ async def test_account_info_with_valid_key(
@pytest.mark.asyncio
async def test_account_info_without_auth(async_client: AsyncClient):
async def test_account_info_without_auth(async_client: AsyncClient) -> None:
"""Test that account info requires authentication."""
response = await async_client.get("/v1/wallet/")
@@ -61,7 +62,7 @@ async def test_account_info_without_auth(async_client: AsyncClient):
@pytest.mark.asyncio
async def test_account_info_with_invalid_key(async_client: AsyncClient):
async def test_account_info_with_invalid_key(async_client: AsyncClient) -> None:
"""Test account info with an invalid API key."""
response = await async_client.get(
"/v1/wallet/info", headers={"Authorization": "Bearer invalid-key"}
@@ -73,7 +74,7 @@ async def test_account_info_with_invalid_key(async_client: AsyncClient):
@pytest.mark.asyncio
async def test_refund_balance_with_address(
async_client: AsyncClient, test_api_key: ApiKey, test_session: AsyncSession
):
) -> None:
"""Test refunding balance when refund address is set."""
# Need to patch the refund_balance at the module level to intercept the call
with patch("router.account.refund_balance", new_callable=AsyncMock) as mock_refund:
@@ -101,7 +102,7 @@ async def test_refund_balance_with_address(
@pytest.mark.asyncio
async def test_refund_balance_without_address(
async_client: AsyncClient, test_session: AsyncSession
):
) -> None:
"""Test refunding balance when no refund address is set."""
# Create key without refund address - with unique ID
unique_id = str(uuid.uuid4())[:8]
@@ -144,11 +145,11 @@ async def test_refund_balance_without_address(
@pytest.mark.asyncio
async def test_topup_balance_endpoint(
async_client: AsyncClient, test_api_key: ApiKey, test_session: AsyncSession
):
) -> None:
"""Test topping up balance with a cashu token."""
# Mock at the router.account module level to intercept the import
with patch("router.account.credit_balance", new_callable=AsyncMock) as mock_credit:
mock_credit.return_value = {"msats": 500000}
mock_credit.return_value = 500000 # Return integer msats value
response = await async_client.post(
"/v1/wallet/topup?cashu_token=cashuBqQSEQ...",
@@ -166,7 +167,7 @@ async def test_topup_balance_endpoint(
@pytest.mark.asyncio
async def test_topup_balance_requires_cashu_token(
async_client: AsyncClient, test_api_key: ApiKey
):
) -> None:
"""Test that topup endpoint requires a cashu token."""
response = await async_client.post(
"/v1/wallet/topup",
@@ -180,7 +181,7 @@ async def test_topup_balance_requires_cashu_token(
@pytest.mark.asyncio
async def test_account_with_cashu_token(
async_client: AsyncClient, test_session: AsyncSession
):
) -> None:
"""Test authentication with a cashu token creates a new account."""
cashu_token = "cashuBqQSEQ123456"
@@ -212,7 +213,7 @@ async def test_account_with_cashu_token(
@pytest.mark.asyncio
async def test_account_with_invalid_cashu_token(async_client: AsyncClient):
async def test_account_with_invalid_cashu_token(async_client: AsyncClient) -> None:
"""Test authentication with an invalid cashu token returns 401."""
with patch("router.auth.credit_balance", new_callable=AsyncMock) as mock_credit:
@@ -225,4 +226,3 @@ async def test_account_with_invalid_cashu_token(async_client: AsyncClient):
assert response.status_code == 401
error = response.json()
assert error["detail"]["error"]["code"] == "invalid_api_key"
+9 -8
View File
@@ -1,11 +1,12 @@
from unittest.mock import patch
import pytest
from fastapi.testclient import TestClient
from httpx import AsyncClient
@pytest.mark.asyncio
async def test_root_endpoint(async_client: AsyncClient):
async def test_root_endpoint(async_client: AsyncClient) -> None:
"""Test the root endpoint returns expected information."""
# Mock the environment variables for this specific test
env_vars = {
@@ -16,13 +17,13 @@ async def test_root_endpoint(async_client: AsyncClient):
"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
data = response.json()
# The app reads from env vars during import, so check what we actually get
assert "name" in data
assert "description" in data
@@ -35,16 +36,16 @@ async def test_root_endpoint(async_client: AsyncClient):
@pytest.mark.asyncio
async def test_cors_headers(async_client: AsyncClient):
async def test_cors_headers(async_client: AsyncClient) -> None:
"""Test that CORS headers are properly set."""
response = await async_client.options(
"/",
headers={
"Origin": "http://localhost:3000",
"Access-Control-Request-Method": "GET",
}
},
)
assert response.status_code == 200
# Check that CORS is working (might be * or specific origin)
assert "access-control-allow-origin" in response.headers
@@ -52,7 +53,7 @@ async def test_cors_headers(async_client: AsyncClient):
@pytest.mark.asyncio
async def test_startup_event_initializes_properly(test_client):
async def test_startup_event_initializes_properly(test_client: TestClient) -> None:
"""Test that the startup event runs without errors."""
# The test_client fixture already triggers the startup event
# This test ensures no exceptions are raised during startup
+71 -57
View File
@@ -1,6 +1,6 @@
import asyncio
from typing import Any
from unittest.mock import AsyncMock, patch
import pytest
from router.models import (
@@ -27,7 +27,7 @@ def sample_model() -> Model:
input_modalities=["text"],
output_modalities=["text"],
tokenizer="test_tokenizer",
instruct_type="chat"
instruct_type="chat",
),
pricing=Pricing(
prompt=0.01,
@@ -35,63 +35,71 @@ def sample_model() -> Model:
request=0.001,
image=0.0,
web_search=0.0,
internal_reasoning=0.0
internal_reasoning=0.0,
),
top_provider=TopProvider(
context_length=4096,
max_completion_tokens=2048,
is_moderated=False
)
context_length=4096, max_completion_tokens=2048, is_moderated=False
),
)
@pytest.mark.asyncio
async def test_update_sats_pricing_calculation(sample_model: Model):
async def test_update_sats_pricing_calculation(sample_model: Model) -> None:
"""Test that sats pricing is calculated correctly."""
# Mock the sats_usd_ask_price function
with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price:
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD
# Temporarily replace MODELS
original_models = MODELS[:]
MODELS.clear()
MODELS.append(sample_model)
# Run one iteration of the pricing update
sleep_called = asyncio.Event()
async def mock_sleep(duration):
async def mock_sleep(duration: float) -> None:
sleep_called.set()
raise asyncio.CancelledError()
with patch("asyncio.sleep", side_effect=mock_sleep):
try:
# Create and run the task
task = asyncio.create_task(update_sats_pricing())
# Wait for the first iteration to complete
await sleep_called.wait()
# Check that sats pricing was calculated
assert sample_model.sats_pricing is not None
# Verify calculations (prices in USD / sats_to_usd)
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
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 == pytest.approx(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
finally:
@@ -101,7 +109,7 @@ async def test_update_sats_pricing_calculation(sample_model: Model):
@pytest.mark.asyncio
async def test_update_sats_pricing_without_top_provider():
async def test_update_sats_pricing_without_top_provider() -> None:
"""Test sats pricing calculation for models without top_provider."""
model_without_top = Model(
id="test-model-no-top",
@@ -114,7 +122,7 @@ async def test_update_sats_pricing_without_top_provider():
input_modalities=["text"],
output_modalities=["text"],
tokenizer="test_tokenizer",
instruct_type=None
instruct_type=None,
),
pricing=Pricing(
prompt=0.01,
@@ -122,31 +130,33 @@ async def test_update_sats_pricing_without_top_provider():
request=0.001,
image=0.01,
web_search=0.005,
internal_reasoning=0.015
internal_reasoning=0.015,
),
top_provider=None
top_provider=None,
)
with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price:
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.return_value = 0.0001 # 1 sat = 0.0001 USD
original_models = MODELS[:]
MODELS.clear()
MODELS.append(model_without_top)
sleep_called = asyncio.Event()
async def mock_sleep(duration):
async def mock_sleep(duration: float) -> None:
sleep_called.set()
raise asyncio.CancelledError()
with patch("asyncio.sleep", side_effect=mock_sleep):
try:
task = asyncio.create_task(update_sats_pricing())
await sleep_called.wait()
assert model_without_top.sats_pricing is not None
# Verify the fallback max_cost calculation
p = model_without_top.sats_pricing.prompt * 1_000_000
c = model_without_top.sats_pricing.completion * 32_000
@@ -155,16 +165,18 @@ async def test_update_sats_pricing_without_top_provider():
w = model_without_top.sats_pricing.web_search * 1000
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 == pytest.approx(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
finally:
@@ -173,59 +185,61 @@ async def test_update_sats_pricing_without_top_provider():
@pytest.mark.asyncio
async def test_update_sats_pricing_handles_errors():
async def test_update_sats_pricing_handles_errors() -> None:
"""Test that update_sats_pricing handles errors gracefully."""
with patch("router.models.sats_usd_ask_price", new_callable=AsyncMock) as mock_price:
with patch(
"router.models.sats_usd_ask_price", new_callable=AsyncMock
) as mock_price:
mock_price.side_effect = Exception("API Error")
error_printed = False
original_print = print
def mock_print(*args, **kwargs):
def mock_print(*args: Any, **kwargs: Any) -> None:
nonlocal error_printed
message = " ".join(str(a) for a in args)
if "API Error" in message and "Error updating sats pricing" in message:
error_printed = True
original_print(*args, **kwargs)
with patch("builtins.print", side_effect=mock_print):
sleep_called = asyncio.Event()
async def mock_sleep(duration):
async def mock_sleep(duration: float) -> None:
sleep_called.set()
raise asyncio.CancelledError()
with patch("asyncio.sleep", side_effect=mock_sleep):
try:
task = asyncio.create_task(update_sats_pricing())
await sleep_called.wait()
# 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
def test_model_serialization(sample_model: Model):
def test_model_serialization(sample_model: Model) -> None:
"""Test that models can be serialized and deserialized correctly."""
model_dict = sample_model.dict()
# Verify all fields are present
assert model_dict["id"] == "test-model"
assert model_dict["name"] == "Test Model"
assert model_dict["pricing"]["prompt"] == 0.01
assert model_dict["architecture"]["modality"] == "text"
assert model_dict["top_provider"]["context_length"] == 4096
# Test deserialization
new_model = Model(**model_dict)
assert new_model.id == sample_model.id
assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt)
assert new_model.pricing.prompt == pytest.approx(sample_model.pricing.prompt)
+9 -8
View File
@@ -1,6 +1,7 @@
import json
import os
import uuid
from typing import AsyncGenerator
from unittest.mock import AsyncMock, patch
import pytest
@@ -28,7 +29,7 @@ async def api_key_with_balance(test_session: AsyncSession) -> ApiKey:
@pytest.mark.asyncio
async def test_proxy_requires_authentication(async_client: AsyncClient):
async def test_proxy_requires_authentication(async_client: AsyncClient) -> None:
"""Test that proxy endpoints require authentication."""
response = await async_client.post("/v1/chat/completions")
@@ -42,7 +43,7 @@ async def test_proxy_requires_authentication(async_client: AsyncClient):
@pytest.mark.asyncio
async def test_proxy_with_insufficient_balance(
async_client: AsyncClient, test_session: AsyncSession
):
) -> None:
"""Test proxy request with insufficient balance."""
# Create key with minimal balance
unique_id = str(uuid.uuid4())[:8]
@@ -71,7 +72,7 @@ async def test_proxy_with_insufficient_balance(
@pytest.mark.asyncio
async def test_proxy_invalid_json_body(
async_client: AsyncClient, api_key_with_balance: ApiKey
):
) -> None:
"""Test proxy request with invalid JSON body."""
response = await async_client.post(
"/v1/chat/completions",
@@ -91,7 +92,7 @@ async def test_proxy_invalid_json_body(
@pytest.mark.asyncio
async def test_proxy_successful_request_mock(
async_client: AsyncClient, api_key_with_balance: ApiKey, test_session: AsyncSession
):
) -> None:
"""Test successful proxy request with mocked upstream."""
mock_response_data = {
"id": "chatcmpl-123",
@@ -162,7 +163,7 @@ async def test_proxy_successful_request_mock(
@pytest.mark.asyncio
async def test_proxy_streaming_response(
async_client: AsyncClient, api_key_with_balance: ApiKey
):
) -> None:
"""Test proxy request with streaming response."""
# Mock SSE stream chunks
stream_chunks = [
@@ -172,7 +173,7 @@ async def test_proxy_streaming_response(
b"data: [DONE]\n\n",
]
async def mock_aiter_bytes():
async def mock_aiter_bytes() -> AsyncGenerator[bytes, None]:
for chunk in stream_chunks:
yield chunk
@@ -213,7 +214,7 @@ async def test_proxy_streaming_response(
@pytest.mark.asyncio
async def test_proxy_handles_upstream_errors(
async_client: AsyncClient, api_key_with_balance: ApiKey
):
) -> None:
"""Test proxy handles upstream connection errors gracefully."""
with patch("httpx.AsyncClient") as mock_client_class:
mock_client = AsyncMock()
@@ -245,7 +246,7 @@ async def test_proxy_handles_upstream_errors(
@pytest.mark.asyncio
async def test_proxy_with_model_based_pricing(
async_client: AsyncClient, test_session: AsyncSession
):
) -> None:
"""Test proxy with model-based pricing enabled."""
# Create API key with sufficient balance
unique_id = str(uuid.uuid4())[:8]
+13 -10
View File
@@ -8,11 +8,11 @@ from tests.conftest import TEST_ENV
@pytest.mark.asyncio
async def test_background_tasks_cancel_on_shutdown():
async def test_background_tasks_cancel_on_shutdown() -> None:
pricing_started = asyncio.Event()
pricing_cancelled = asyncio.Event()
async def fake_update():
async def fake_update() -> None:
pricing_started.set()
try:
await asyncio.Event().wait()
@@ -23,7 +23,7 @@ async def test_background_tasks_cancel_on_shutdown():
refund_started = asyncio.Event()
refund_cancelled = asyncio.Event()
async def fake_refund():
async def fake_refund() -> None:
refund_started.set()
try:
await asyncio.Event().wait()
@@ -31,7 +31,7 @@ async def test_background_tasks_cancel_on_shutdown():
refund_cancelled.set()
raise
with patch.dict('os.environ', TEST_ENV, clear=True):
with patch.dict("os.environ", TEST_ENV, clear=True):
mock_wallet = AsyncMock()
mock_wallet.__aenter__ = AsyncMock(return_value=mock_wallet)
mock_wallet.__aexit__ = AsyncMock(return_value=None)
@@ -40,13 +40,16 @@ async def test_background_tasks_cancel_on_shutdown():
mock_wallet.fetch_wallet_state = AsyncMock(return_value=mock_state)
mock_wallet.send_to_lnurl = AsyncMock(return_value=100)
mock_wallet.redeem = AsyncMock(return_value=1)
mock_wallet.send = AsyncMock(return_value='cashu:token123')
mock_wallet.send = AsyncMock(return_value="cashu:token123")
with patch('router.cashu.Wallet.create', AsyncMock(return_value=mock_wallet)), \
patch('router.cashu.WALLET', mock_wallet):
with patch('router.main.update_sats_pricing', new=fake_update), \
patch('router.main.check_for_refunds', new=fake_refund):
with (
patch("router.cashu.Wallet.create", AsyncMock(return_value=mock_wallet)),
patch("router.cashu.WALLET", mock_wallet),
):
with (
patch("router.main.update_sats_pricing", new=fake_update),
patch("router.main.check_for_refunds", new=fake_refund),
):
async with lifespan(app):
await pricing_started.wait()
await refund_started.wait()