mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix mypy issues
This commit is contained in:
+5
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user