mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-09 02:54:37 +00:00
741 lines
26 KiB
Python
741 lines
26 KiB
Python
import asyncio
|
|
import json
|
|
import os
|
|
from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
import pytest_asyncio
|
|
from fastapi import FastAPI
|
|
from httpx import AsyncClient
|
|
from sqlalchemy.ext.asyncio import create_async_engine
|
|
from sqlmodel import select
|
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
|
|
|
from routstr.core.logging import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
# Configure test environment based on whether we're using local services or not
|
|
use_local_services = os.environ.get("USE_LOCAL_SERVICES", "0") == "1"
|
|
|
|
if use_local_services:
|
|
# Docker mode: Use Docker services for more realistic testing
|
|
logger.info("🐳 Using Docker services for integration tests")
|
|
test_env = {
|
|
"DATABASE_URL": "sqlite+aiosqlite:///:memory:",
|
|
"UPSTREAM_BASE_URL": "http://localhost:3000", # Mock OpenAI service
|
|
"UPSTREAM_API_KEY": "test-upstream-key",
|
|
"CASHU_MINTS": "http://mint:3338", # Docker service name for routstr validation
|
|
"MINT": "http://mint:3338",
|
|
"MINT_URL": "http://mint:3338",
|
|
"NOSTR_RELAY_URL": "ws://localhost:8088",
|
|
"RECEIVE_LN_ADDRESS": "test@routstr.com",
|
|
"REFUND_PROCESSING_INTERVAL": "3600",
|
|
"NSEC": "nsec1testkey1234567890abcdef",
|
|
"FIXED_COST_PER_REQUEST": "10",
|
|
"FIXED_PRICING": "false",
|
|
"MINIMUM_PAYOUT": "1000",
|
|
"PAYOUT_INTERVAL": "86400",
|
|
"NAME": "TestRoutstrNode",
|
|
"DESCRIPTION": "Test Node for Integration Tests",
|
|
"NPUB": "npub1test",
|
|
"HTTP_URL": "http://localhost:8000",
|
|
"ONION_URL": "http://test.onion",
|
|
"CORS_ORIGINS": "*",
|
|
}
|
|
else:
|
|
# Mock mode: Use in-memory mocks for fast testing
|
|
logger.info("🎭 Using mocked services for integration tests")
|
|
test_env = {
|
|
"DATABASE_URL": "sqlite+aiosqlite:///:memory:",
|
|
"UPSTREAM_BASE_URL": "https://api.openai.com/v1",
|
|
"UPSTREAM_API_KEY": "test-upstream-key",
|
|
"CASHU_MINTS": "http://localhost:3338",
|
|
"RECEIVE_LN_ADDRESS": "test@routstr.com",
|
|
"REFUND_PROCESSING_INTERVAL": "3600",
|
|
"NSEC": "nsec1testkey1234567890abcdef",
|
|
"FIXED_COST_PER_REQUEST": "10",
|
|
"FIXED_PRICING": "false",
|
|
"MINIMUM_PAYOUT": "1000",
|
|
"PAYOUT_INTERVAL": "86400",
|
|
}
|
|
|
|
# Set test environment variables before importing the app
|
|
os.environ.update(test_env)
|
|
os.environ.pop("ADMIN_PASSWORD", None)
|
|
|
|
|
|
from routstr.core.db import ApiKey, get_session # noqa: E402
|
|
from routstr.core.main import app, lifespan # noqa: E402
|
|
|
|
|
|
@pytest.fixture(scope="session")
|
|
def test_mode() -> str:
|
|
"""Returns current test mode for clarity"""
|
|
if os.environ.get("USE_LOCAL_SERVICES") == "1":
|
|
print("\n🐳 Running with Docker services (realistic mode)")
|
|
return "docker"
|
|
else:
|
|
print("\n🎭 Running with mocked services (fast mode)")
|
|
return "mock"
|
|
|
|
|
|
class TestmintWallet:
|
|
"""Test wallet that simulates Cashu mint interactions for testing"""
|
|
|
|
def __init__(
|
|
self, mint_url: Optional[str] = None, nsec: Optional[str] = None
|
|
) -> None:
|
|
# Use the configured CASHU_MINTS URL, fallback to MINT, or default
|
|
configured_mint_url = (
|
|
mint_url
|
|
or os.environ.get("CASHU_MINTS", "").split(",")[0].strip()
|
|
or os.environ.get("MINT", "http://localhost:3338")
|
|
)
|
|
|
|
# For local services, use localhost for connection but mint service name for token creation
|
|
if os.environ.get("USE_LOCAL_SERVICES") == "1":
|
|
self.connection_url = configured_mint_url.replace(
|
|
"http://mint:", "http://localhost:"
|
|
)
|
|
self.mint_url = configured_mint_url # Keep Docker service name for tokens
|
|
else:
|
|
self.connection_url = configured_mint_url
|
|
self.mint_url = configured_mint_url
|
|
# Use a valid test nsec for testing (this is a well-known test key)
|
|
self.nsec = (
|
|
nsec or "nsec1vl029mgpspedva04g90vltkh6fvh240zqtv9k0t9af8935ke9laqsnlfe5"
|
|
)
|
|
self.wallet = None
|
|
self.tokens: List[Dict[str, Any]] = []
|
|
self.spent_tokens: List[str] = []
|
|
self.refund_history: List[Dict[str, Any]] = []
|
|
|
|
async def init(self) -> None:
|
|
"""Initialize the sixty_nuts wallet"""
|
|
# In mock mode, we don't actually create a real wallet
|
|
# This is just a placeholder for the mock implementation
|
|
self.wallet = None
|
|
|
|
async def mint_tokens(self, amount: int) -> str:
|
|
"""Create a test token for the testmint"""
|
|
logger.info(
|
|
f"Creating test token for {amount} sats from testmint {self.mint_url}"
|
|
)
|
|
|
|
# For integration tests, use fallback tokens to avoid external dependencies
|
|
return await self._create_fallback_token(amount)
|
|
|
|
async def _create_real_token(self, amount: int) -> str:
|
|
"""Create real tokens using the testmint"""
|
|
import tempfile
|
|
|
|
from cashu.wallet.wallet import Wallet
|
|
|
|
logger.info(
|
|
f"Creating real token for {amount} sats from testmint {self.connection_url}"
|
|
)
|
|
|
|
try:
|
|
# Create a temporary wallet to mint real tokens
|
|
with tempfile.TemporaryDirectory() as temp_dir:
|
|
wallet_db_path = os.path.join(temp_dir, "test_wallet.db")
|
|
|
|
wallet = await Wallet.with_db(
|
|
self.connection_url, # Connect via localhost
|
|
db=f"sqlite+aiosqlite:///{wallet_db_path}",
|
|
load_all_keysets=True,
|
|
unit="sat",
|
|
)
|
|
|
|
# Load mint information
|
|
await wallet.load_mint()
|
|
|
|
# Request a mint quote
|
|
quote_response = await wallet.mint_quote(amount=amount, unit="sat")
|
|
quote = quote_response.quote
|
|
|
|
# Mint tokens (simulate payment by directly calling mint endpoint)
|
|
mint_response = await wallet.mint(amount=amount, hash=quote)
|
|
token = mint_response.token
|
|
|
|
# Replace connection URL with Docker service name for routstr validation
|
|
if self.connection_url != self.mint_url:
|
|
token = token.replace(self.connection_url, self.mint_url)
|
|
|
|
logger.info(f"Successfully minted real token for {amount} sats")
|
|
return token
|
|
except Exception as e:
|
|
logger.error(f"Failed to mint real token: {e}")
|
|
raise
|
|
|
|
async def _create_fallback_token(self, amount: int) -> str:
|
|
"""Fallback method to create a basic test token"""
|
|
import base64
|
|
import json
|
|
import random
|
|
import time
|
|
|
|
unique_id = int(time.time() * 1000000) + random.randint(1000, 9999)
|
|
token_data = {
|
|
"token": [
|
|
{
|
|
"mint": self.mint_url,
|
|
"proofs": [
|
|
{
|
|
"id": f"009a1f293253e41e{unique_id % 100000000:08d}",
|
|
"amount": amount,
|
|
"secret": f"test-secret-{amount}-{unique_id}",
|
|
"C": "02194603ffa36356f4a56b7df9371fc3192472351453ec7398b8da8117e7c3e104",
|
|
}
|
|
],
|
|
}
|
|
],
|
|
"unit": "sat",
|
|
"memo": f"Test token {amount} sats",
|
|
}
|
|
|
|
token_json = json.dumps(token_data)
|
|
token_base64 = base64.urlsafe_b64encode(token_json.encode()).decode()
|
|
return f"cashuA{token_base64}"
|
|
|
|
async def redeem_token(self, token: str) -> Tuple[int, str, str]:
|
|
"""Redeem a Cashu token - compatible with wallet.recieve_token"""
|
|
if not self.wallet:
|
|
await self.init()
|
|
|
|
# For testing, simulate the redemption
|
|
import base64
|
|
|
|
if not token.startswith("cashuA"):
|
|
raise ValueError("Invalid token format")
|
|
|
|
try:
|
|
token_base64 = token[6:] # Remove "cashuA" prefix
|
|
# Add padding if necessary
|
|
padding = (4 - len(token_base64) % 4) % 4
|
|
token_base64 += "=" * padding
|
|
token_json = base64.urlsafe_b64decode(token_base64).decode()
|
|
token_data = json.loads(token_json)
|
|
|
|
total_amount = 0
|
|
mint_url = self.mint_url
|
|
unit = token_data.get("unit", "sat")
|
|
|
|
for mint_tokens in token_data["token"]:
|
|
mint_url = mint_tokens.get("mint", self.mint_url)
|
|
for proof in mint_tokens["proofs"]:
|
|
# Check if token was already spent
|
|
if proof["id"] in self.spent_tokens:
|
|
raise ValueError("Token already spent")
|
|
|
|
self.spent_tokens.append(proof["id"])
|
|
total_amount += proof["amount"]
|
|
|
|
return total_amount, unit, mint_url
|
|
|
|
except Exception as e:
|
|
raise ValueError(f"Failed to decode token: {str(e)}")
|
|
|
|
async def redeem_token_simple(self, token: str) -> Tuple[int, str]:
|
|
"""Redeem a Cashu token - simple version for credit_balance"""
|
|
amount, unit, mint_url = await self.redeem_token(token)
|
|
return amount, "test_metadata"
|
|
|
|
async def send(self, amount: int) -> str:
|
|
"""Create a token to send (for refunds)"""
|
|
if not self.wallet:
|
|
await self.init()
|
|
|
|
# For testing, create a refund token
|
|
return await self.mint_tokens(amount)
|
|
|
|
async def send_token(
|
|
self, amount: int, unit: str, mint_url: Optional[str] = None
|
|
) -> str:
|
|
"""Send token with compatible signature for mocking routstr.wallet.send_token"""
|
|
return await self.send(amount)
|
|
|
|
async def send_to_lnurl(self, lnurl: str, amount: int) -> int:
|
|
"""Send to lightning address - simulated for testing"""
|
|
if not self.wallet:
|
|
await self.init()
|
|
|
|
self.refund_history.append(
|
|
{
|
|
"amount": amount,
|
|
"ln_address": lnurl,
|
|
"timestamp": asyncio.get_event_loop().time(),
|
|
}
|
|
)
|
|
return amount
|
|
|
|
async def get_balance(self) -> int:
|
|
"""Get wallet balance"""
|
|
if not self.wallet:
|
|
await self.init()
|
|
|
|
# For testing, return a simulated balance
|
|
return 100000 # 100k sats
|
|
|
|
async def credit_balance(
|
|
self, cashu_token: str, key: ApiKey, session: AsyncSession
|
|
) -> int:
|
|
"""Credit balance to API key - test implementation"""
|
|
try:
|
|
logger.info(
|
|
f"TestmintWallet.credit_balance called with token: {cashu_token[:20]}..."
|
|
)
|
|
|
|
# Redeem the token to get amount
|
|
amount, _ = await self.redeem_token_simple(cashu_token)
|
|
logger.info(f"TestmintWallet.credit_balance redeemed amount: {amount}")
|
|
|
|
# For testing, convert to msat if needed
|
|
amount_msat = amount * 1000 # Assume tokens are in sats
|
|
logger.info(f"TestmintWallet.credit_balance amount in msat: {amount_msat}")
|
|
|
|
# Credit the balance using atomic database update to prevent race conditions
|
|
from sqlmodel import col, update
|
|
|
|
# Use atomic update to avoid lost update problem in concurrent scenarios
|
|
stmt = (
|
|
update(ApiKey)
|
|
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
.values(balance=ApiKey.balance + amount_msat)
|
|
)
|
|
await session.execute(stmt)
|
|
await session.commit()
|
|
|
|
# Refresh the key object to get the updated balance
|
|
await session.refresh(key)
|
|
|
|
logger.info(
|
|
f"TestmintWallet.credit_balance successfully credited {amount_msat} msat"
|
|
)
|
|
|
|
return amount_msat
|
|
except Exception as e:
|
|
logger.error(f"TestmintWallet.credit_balance failed: {e}")
|
|
import traceback
|
|
|
|
logger.error(
|
|
f"TestmintWallet.credit_balance full traceback: {traceback.format_exc()}"
|
|
)
|
|
raise ValueError(f"Failed to redeem token: {str(e)}")
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def testmint_wallet() -> TestmintWallet:
|
|
"""Fixture for testmint wallet instance"""
|
|
# Check if we should use real mint
|
|
mint_url = os.environ.get(
|
|
"MINT_URL", os.environ.get("MINT", "http://localhost:3338")
|
|
)
|
|
|
|
wallet = TestmintWallet(mint_url=mint_url)
|
|
await wallet.init()
|
|
return wallet
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def test_database_url(tmp_path: Any) -> str:
|
|
"""Create a temporary SQLite database file for integration tests"""
|
|
db_file = tmp_path / "test_integration.db"
|
|
return f"sqlite+aiosqlite:///{db_file}"
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None]:
|
|
"""Create an async engine for integration tests"""
|
|
engine = create_async_engine(
|
|
test_database_url,
|
|
echo=False,
|
|
future=True,
|
|
pool_pre_ping=True,
|
|
pool_size=5,
|
|
max_overflow=10,
|
|
)
|
|
|
|
# Initialize database schema
|
|
# Create tables using the engine directly since init_db uses the global engine
|
|
async with engine.begin() as conn:
|
|
from sqlmodel import SQLModel
|
|
|
|
await conn.run_sync(SQLModel.metadata.create_all)
|
|
|
|
yield engine
|
|
|
|
# Cleanup
|
|
await engine.dispose()
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def integration_session(
|
|
integration_engine: Any,
|
|
) -> AsyncGenerator[AsyncSession, None]:
|
|
"""Create a database session for integration tests"""
|
|
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
|
yield session
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def patched_db_engine(integration_engine: Any) -> AsyncGenerator[None, None]:
|
|
"""Patch the global db engine so create_session() uses the test engine."""
|
|
with patch("routstr.core.db.engine", integration_engine):
|
|
yield
|
|
|
|
|
|
class DatabaseSnapshot:
|
|
"""Utility to capture and compare database states"""
|
|
|
|
def __init__(self, session: AsyncSession) -> None:
|
|
self.session = session
|
|
self.snapshot: Optional[Dict[str, List[Dict]]] = None
|
|
|
|
async def capture(self) -> Dict[str, List[Dict]]:
|
|
"""Capture current database state"""
|
|
# Get all API keys with their data
|
|
result = await self.session.execute(select(ApiKey))
|
|
api_keys = result.scalars().all()
|
|
|
|
snapshot = {
|
|
"api_keys": [
|
|
{
|
|
"hashed_key": key.hashed_key,
|
|
"balance": key.balance,
|
|
"total_spent": key.total_spent,
|
|
"total_requests": key.total_requests,
|
|
"refund_address": key.refund_address,
|
|
"key_expiry_time": key.key_expiry_time,
|
|
}
|
|
for key in api_keys
|
|
]
|
|
}
|
|
|
|
self.snapshot = snapshot
|
|
return snapshot
|
|
|
|
async def diff(
|
|
self, new_snapshot: Optional[Dict[str, List[Dict]]] = None
|
|
) -> Dict[str, Any]:
|
|
"""Calculate differences between snapshots"""
|
|
if new_snapshot is None:
|
|
new_snapshot = await self.capture()
|
|
|
|
if self.snapshot is None:
|
|
raise ValueError("No initial snapshot to compare against")
|
|
|
|
diff: Dict[str, Dict[str, List[Any]]] = {
|
|
"api_keys": {"added": [], "removed": [], "modified": []}
|
|
}
|
|
|
|
# Create lookup maps
|
|
old_keys = {k["hashed_key"]: k for k in self.snapshot["api_keys"]}
|
|
new_keys = {k["hashed_key"]: k for k in new_snapshot["api_keys"]}
|
|
|
|
# Find added keys
|
|
for key_id in new_keys:
|
|
if key_id not in old_keys:
|
|
diff["api_keys"]["added"].append(new_keys[key_id])
|
|
|
|
# Find removed keys
|
|
for key_id in old_keys:
|
|
if key_id not in new_keys:
|
|
diff["api_keys"]["removed"].append(old_keys[key_id])
|
|
|
|
# Find modified keys
|
|
for key_id in old_keys:
|
|
if key_id in new_keys:
|
|
old = old_keys[key_id]
|
|
new = new_keys[key_id]
|
|
changes = {}
|
|
|
|
for field in [
|
|
"balance",
|
|
"total_spent",
|
|
"total_requests",
|
|
"refund_address",
|
|
"key_expiry_time",
|
|
]:
|
|
if old[field] != new[field]:
|
|
changes[field] = {
|
|
"old": old[field],
|
|
"new": new[field],
|
|
"delta": new[field] - old[field]
|
|
if isinstance(new[field], (int, float))
|
|
else None,
|
|
}
|
|
|
|
if changes:
|
|
diff["api_keys"]["modified"].append(
|
|
{"hashed_key": key_id, "changes": changes}
|
|
)
|
|
|
|
return diff
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def db_snapshot(integration_session: AsyncSession) -> DatabaseSnapshot:
|
|
"""Database snapshot utility for tracking state changes"""
|
|
return DatabaseSnapshot(integration_session)
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def integration_app(
|
|
integration_engine: Any,
|
|
integration_session: AsyncSession,
|
|
testmint_wallet: TestmintWallet,
|
|
test_database_url: str,
|
|
) -> AsyncGenerator[FastAPI, None]:
|
|
"""Create FastAPI app instance for integration tests"""
|
|
|
|
# Override environment with test database URL
|
|
os.environ["DATABASE_URL"] = test_database_url
|
|
|
|
# Create a new app instance with our lifespan
|
|
test_app = FastAPI(lifespan=lifespan)
|
|
|
|
# Copy all routes from the main app
|
|
test_app.router = app.router
|
|
|
|
# Override the get_session dependency
|
|
async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
|
|
yield integration_session
|
|
|
|
test_app.dependency_overrides[get_session] = override_get_session
|
|
|
|
# Check if we should use real mint
|
|
use_real_mint = os.environ.get("USE_REAL_MINT", "false").lower() == "true"
|
|
|
|
if use_real_mint:
|
|
# Use real mint - no wallet patches needed
|
|
with patch("routstr.core.db.engine", integration_engine):
|
|
yield test_app
|
|
else:
|
|
# Use testmint with wallet patches for all integration tests
|
|
mint_url = os.environ.get("CASHU_MINTS", "http://localhost:3338")
|
|
from routstr.core.settings import settings as _settings
|
|
|
|
# Passthrough discounted max cost to avoid dependence on MODELS in tests
|
|
async def _passthrough_discount(
|
|
max_cost_for_model: int,
|
|
body: dict,
|
|
model_obj: Any = None,
|
|
) -> int:
|
|
return max_cost_for_model
|
|
|
|
with (
|
|
patch("routstr.core.db.engine", integration_engine),
|
|
patch.object(_settings, "cashu_mints", [mint_url]),
|
|
patch("routstr.wallet.credit_balance", testmint_wallet.credit_balance),
|
|
patch("routstr.wallet.send_token", testmint_wallet.send_token),
|
|
patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl),
|
|
patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token),
|
|
patch("routstr.wallet.get_balance", testmint_wallet.get_balance),
|
|
patch("routstr.balance.send_token", testmint_wallet.send_token),
|
|
patch("routstr.balance.send_to_lnurl", testmint_wallet.send_to_lnurl),
|
|
patch("websockets.connect") as mock_websockets,
|
|
patch("routstr.payment.price.btc_usd_price", return_value=50000.0),
|
|
patch("routstr.payment.price.sats_usd_price", return_value=0.0005),
|
|
patch(
|
|
"routstr.payment.helpers.calculate_discounted_max_cost",
|
|
side_effect=_passthrough_discount,
|
|
),
|
|
):
|
|
# Configure the WebSocket mock for discovery service - fast failure for performance tests
|
|
async def mock_websocket_connect(*args: Any, **kwargs: Any) -> None:
|
|
raise ConnectionError("Mock connection failed")
|
|
|
|
mock_websockets.side_effect = mock_websocket_connect
|
|
|
|
yield test_app
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def integration_client(
|
|
integration_app: FastAPI,
|
|
integration_engine: Any, # Ensure engine is created first
|
|
) -> AsyncGenerator[AsyncClient, None]:
|
|
"""Create an async HTTP client for integration tests"""
|
|
from httpx import ASGITransport
|
|
|
|
async with AsyncClient(
|
|
transport=ASGITransport(app=integration_app), # type: ignore
|
|
base_url="http://test",
|
|
timeout=30.0,
|
|
) as client:
|
|
yield client
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def authenticated_client(
|
|
integration_client: AsyncClient,
|
|
testmint_wallet: TestmintWallet,
|
|
integration_session: AsyncSession,
|
|
) -> AsyncClient:
|
|
"""Create an authenticated client with a persistent API key"""
|
|
# Generate a cashu token
|
|
test_token = await testmint_wallet.mint_tokens(10000) # 10k sats
|
|
|
|
# Use the cashu token as Bearer auth to create an API key
|
|
integration_client.headers["Authorization"] = f"Bearer {test_token}"
|
|
|
|
# Make a request to create the API key (first use of cashu token creates the key)
|
|
response = await integration_client.get("/v1/wallet/info")
|
|
assert response.status_code == 200
|
|
wallet_info = response.json()
|
|
api_key = wallet_info["api_key"]
|
|
|
|
# Now switch to using the persistent API key
|
|
integration_client.headers["Authorization"] = f"Bearer {api_key}"
|
|
|
|
# Store the API key and balance for tests that need it
|
|
integration_client._test_api_key = api_key # type: ignore
|
|
integration_client._test_balance = wallet_info["balance"] # type: ignore
|
|
|
|
return integration_client
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def create_api_key() -> Callable:
|
|
"""Helper to create new API keys for testing"""
|
|
|
|
async def _create_key(
|
|
client: AsyncClient,
|
|
wallet: TestmintWallet,
|
|
amount: int = 1000,
|
|
refund_address: Optional[str] = None,
|
|
key_expiry_time: Optional[int] = None,
|
|
) -> Tuple[str, int]:
|
|
"""Create a new API key and return (api_key, balance)"""
|
|
# Generate cashu token
|
|
token = await wallet.mint_tokens(amount)
|
|
|
|
# Create headers
|
|
headers = {"Authorization": f"Bearer {token}"}
|
|
if refund_address:
|
|
headers["Refund-LNURL"] = refund_address
|
|
if key_expiry_time:
|
|
headers["Key-Expiry-Time"] = str(key_expiry_time)
|
|
|
|
# Use the token to create API key
|
|
response = await client.get("/v1/wallet/info", headers=headers)
|
|
assert response.status_code == 200
|
|
|
|
wallet_info = response.json()
|
|
return wallet_info["api_key"], wallet_info["balance"]
|
|
|
|
return _create_key
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_upstream_server() -> Any:
|
|
"""Mock upstream API server responses"""
|
|
responses: Dict[str, Any] = {}
|
|
|
|
class MockResponse:
|
|
def __init__(
|
|
self,
|
|
status_code: int,
|
|
json_data: Any = None,
|
|
text_data: Optional[str] = None,
|
|
) -> None:
|
|
self.status_code = status_code
|
|
self._json_data = json_data
|
|
self._text_data = text_data
|
|
self.headers = {"content-type": "application/json"}
|
|
|
|
def json(self) -> Any:
|
|
return self._json_data
|
|
|
|
@property
|
|
def text(self) -> str:
|
|
return self._text_data or ""
|
|
|
|
async def aiter_bytes(
|
|
self, chunk_size: Optional[int] = None
|
|
) -> AsyncGenerator[bytes, None]:
|
|
"""Async iterator for streaming responses"""
|
|
if self._text_data:
|
|
yield self._text_data.encode()
|
|
|
|
def add_response(method: str, path: str, response: MockResponse) -> None:
|
|
"""Add a mock response for a specific method and path"""
|
|
responses[f"{method}:{path}"] = response
|
|
|
|
def get_response(method: str, path: str) -> MockResponse:
|
|
"""Get mock response for a request"""
|
|
key = f"{method}:{path}"
|
|
if key in responses:
|
|
return responses[key]
|
|
# Default 404 response
|
|
return MockResponse(404, {"error": "Not found"})
|
|
|
|
mock_server = MagicMock()
|
|
mock_server.add_response = add_response
|
|
mock_server.get_response = get_response
|
|
mock_server.responses = responses
|
|
|
|
return mock_server
|
|
|
|
|
|
@pytest_asyncio.fixture
|
|
async def background_tasks_controller() -> AsyncGenerator[Any, None]:
|
|
"""Control background tasks during tests"""
|
|
tasks: List[asyncio.Task] = []
|
|
|
|
class TaskController:
|
|
def __init__(self) -> None:
|
|
self.paused = False
|
|
self.cancelled = False
|
|
|
|
async def pause(self) -> None:
|
|
"""Pause all background tasks"""
|
|
self.paused = True
|
|
|
|
async def resume(self) -> None:
|
|
"""Resume all background tasks"""
|
|
self.paused = False
|
|
|
|
async def cancel_all(self) -> None:
|
|
"""Cancel all background tasks"""
|
|
self.cancelled = True
|
|
for task in tasks:
|
|
task.cancel()
|
|
|
|
controller = TaskController()
|
|
|
|
# Patch background task functions to respect controller
|
|
original_update_pricing: Optional[Callable] = None
|
|
original_periodic_payout: Optional[Callable] = None
|
|
|
|
try:
|
|
from routstr.payment.models import update_sats_pricing
|
|
from routstr.wallet import periodic_payout
|
|
|
|
async def controlled_update_pricing() -> None:
|
|
while not controller.cancelled:
|
|
if not controller.paused and original_update_pricing:
|
|
await original_update_pricing()
|
|
await asyncio.sleep(1)
|
|
|
|
async def controlled_periodic_payout() -> None:
|
|
while not controller.cancelled:
|
|
if not controller.paused and original_periodic_payout:
|
|
await original_periodic_payout()
|
|
await asyncio.sleep(1)
|
|
|
|
# Store originals and patch
|
|
original_update_pricing = update_sats_pricing
|
|
original_periodic_payout = periodic_payout
|
|
|
|
except ImportError:
|
|
pass
|
|
|
|
yield controller
|
|
|
|
# Cleanup
|
|
controller.cancelled = True
|