Merge branch 'main' into x-cash-refac

This commit is contained in:
Shroominic
2025-07-01 13:44:13 -03:00
12 changed files with 121 additions and 315 deletions
+3
View File
@@ -14,3 +14,6 @@ compose.override.yml
# Coverage
.coverage
# deployment
*.log
+2 -2
View File
@@ -5,7 +5,7 @@ from fastapi import APIRouter, Request
from fastapi.responses import HTMLResponse
from sqlmodel import select
from .cashu import WALLET
from .cashu import wallet
from .db import ApiKey, create_session
admin_router = APIRouter(prefix="/admin")
@@ -112,7 +112,7 @@ async def dashboard(request: Request) -> str:
# avoid rounding issues.
total_user_balance = sum(key.balance for key in api_keys) // 1000
# Fetch balance from cashu
current_balance = (await WALLET.fetch_wallet_state()).balance
current_balance = await wallet().get_balance()
owner_balance = current_balance - total_user_balance
return f"""<!DOCTYPE html>
+4 -7
View File
@@ -190,20 +190,17 @@ async def refund_balance(amount_msats: int, key: ApiKey, session: AsyncSession)
if key.refund_address is None:
raise ValueError("Refund address not set.")
assert WALLET is not None, "Wallet not initialized"
return await WALLET.send_to_lnurl(key.refund_address, amount=amount_sats)
return await wallet().send_to_lnurl(key.refund_address, amount=amount_sats)
async def x_cashu_refund(key: ApiKey, session: AsyncSession) -> str:
assert WALLET is not None, "Wallet not initialized"
refund_token = await WALLET.send(key.balance)
refund_token = await wallet().send(key.balance)
await session.delete(key)
await session.commit()
return refund_token
async def redeem(cashu_token: str, lnurl: str) -> int:
assert WALLET is not None, "Wallet not initialized"
amount_sats, _ = await WALLET.redeem(cashu_token)
await WALLET.send_to_lnurl(lnurl, amount=amount_sats)
amount_sats, _ = await wallet().redeem(cashu_token)
await wallet().send_to_lnurl(lnurl, amount=amount_sats)
return amount_sats
+2 -3
View File
@@ -11,11 +11,10 @@ from .admin import admin_router
from .cashu import check_for_refunds, init_wallet, periodic_payout
from .db import init_db
from .discovery import providers_router
from .models import MODELS, update_sats_pricing
from .models import MODELS, models_router, update_sats_pricing
from .proxy import proxy_router
from .route.models import models_router
__version__ = "0.0.1"
__version__ = "0.1.0"
@asynccontextmanager
+8
View File
@@ -3,10 +3,13 @@ import json
import os
from pathlib import Path
from fastapi import APIRouter
from pydantic.v1 import BaseModel
from .price import sats_usd_ask_price
models_router = APIRouter()
class Architecture(BaseModel):
modality: str
@@ -139,3 +142,8 @@ async def update_sats_pricing() -> None:
await asyncio.sleep(10)
except asyncio.CancelledError:
break
@models_router.get("/v1/models")
async def models() -> dict:
return {"data": MODELS}
View File
-62
View File
@@ -1,62 +0,0 @@
from typing import List, Optional
from fastapi import APIRouter
from pydantic import BaseModel
from router.models import MODELS, Model
models_router = APIRouter(prefix="/proxy")
class ProxyModelFromApi(BaseModel):
name: str
input_cost: Optional[float] = None
output_cost: Optional[float] = None
min_cash_per_request: Optional[float] = None
min_cost_per_request: Optional[float] = None
provider: Optional[str] = None
soft_deleted: Optional[bool] = None
model_type: Optional[str] = None
description: Optional[str] = None
context_length: Optional[int] = None
is_free: Optional[bool] = None
def convert_model_to_proxy_format(model: Model) -> ProxyModelFromApi:
input_cost = None
output_cost = None
min_cash_per_request = None
min_cost_per_request = None
is_free = None
if model.sats_pricing:
input_cost = model.sats_pricing.prompt * 1000 * 1_000_000
output_cost = model.sats_pricing.completion * 1000 * 1_000_000
min_cash_per_request = (
model.sats_pricing.request * 1000 if model.sats_pricing.request else 0
)
min_cost_per_request = model.sats_pricing.max_cost * 1000
is_free = (
model.sats_pricing.prompt == 0
and model.sats_pricing.completion == 0
and model.sats_pricing.request == 0
)
return ProxyModelFromApi(
name=model.id,
input_cost=input_cost,
output_cost=output_cost,
min_cash_per_request=min_cash_per_request,
min_cost_per_request=min_cost_per_request,
provider=None,
model_type=model.architecture.modality,
description=model.description,
context_length=model.context_length,
is_free=is_free,
)
@models_router.get("/models")
async def get_models() -> List[ProxyModelFromApi]:
return [convert_model_to_proxy_format(model) for model in MODELS]
+92
View File
@@ -0,0 +1,92 @@
#!/bin/bash
# Auto-update script for Docker Compose projects
# This script checks for new commits and updates Docker Compose when changes are detected
# Configuration - can be overridden by environment variables with sensible defaults
REPO_DIR="${REPO_DIR:-/home/ubuntu/proxy}"
LOG_FILE="${LOG_FILE:-/home/ubuntu/proxy/update.log}"
LOCK_FILE="${LOCK_FILE:-/tmp/proxy_update.lock}"
MAX_LOG_LINES="${MAX_LOG_LINES:-10000}" # Maximum number of log lines to keep
# Function to log messages with timestamp
log_message() {
echo "$(date '+%Y-%m-%d %H:%M:%S') - $1" >> "$LOG_FILE"
}
# Function to prune log file to keep only recent entries
prune_logs() {
if [ -f "$LOG_FILE" ]; then
local line_count=$(wc -l < "$LOG_FILE" 2>/dev/null || echo 0)
if [ "$line_count" -gt "$MAX_LOG_LINES" ]; then
# Create a temporary file with only the last MAX_LOG_LINES lines
tail -n "$MAX_LOG_LINES" "$LOG_FILE" > "${LOG_FILE}.tmp"
mv "${LOG_FILE}.tmp" "$LOG_FILE"
log_message "Log file pruned. Kept last $MAX_LOG_LINES lines."
fi
fi
}
# Function to cleanup on exit
cleanup() {
rm -f "$LOCK_FILE"
}
# Set trap to cleanup on exit
trap cleanup EXIT
# Check if another instance is running
if [ -f "$LOCK_FILE" ]; then
log_message "Another update process is already running. Exiting."
exit 1
fi
# Create lock file
touch "$LOCK_FILE"
# Change to repository directory
cd "$REPO_DIR" || {
log_message "ERROR: Cannot change to repository directory $REPO_DIR"
exit 1
}
# Fetch latest changes from remote
log_message "Fetching latest changes from remote..."
git fetch origin 2>/dev/null
# Check if local branch is behind remote
LOCAL_HASH=$(git rev-parse HEAD)
REMOTE_HASH=$(git rev-parse origin/$(git branch --show-current))
if [ "$LOCAL_HASH" != "$REMOTE_HASH" ]; then
log_message "New commits detected. Current: $LOCAL_HASH, Remote: $REMOTE_HASH"
# Pull latest changes
log_message "Pulling latest changes..."
if git pull origin $(git branch --show-current) 2>&1 | tee -a "$LOG_FILE"; then
log_message "Successfully pulled latest changes"
# Stop current containers
log_message "Stopping current containers..."
sudo docker compose down 2>&1 | tee -a "$LOG_FILE"
# Build and start updated containers
log_message "Building and starting updated containers..."
if sudo docker compose up -d --build 2>&1 | tee -a "$LOG_FILE"; then
log_message "Successfully updated and restarted containers"
else
log_message "ERROR: Failed to start containers"
exit 1
fi
else
log_message "ERROR: Failed to pull changes"
exit 1
fi
else
log_message "No new commits found. Repository is up to date."
fi
log_message "Update check completed successfully"
# Prune logs after each run to prevent unlimited growth
prune_logs
+3
View File
@@ -0,0 +1,3 @@
REPO_DIR=/home/user/proxy
LOG_FILE=/home/user/proxy/update.log
* * * * * /home/user/proxy/scripts/auto_update.sh >/dev/null 2>&1
+7 -11
View File
@@ -1,8 +1,6 @@
import asyncio
import json
from typing import TypedDict
import httpx
from urllib.request import urlopen
class ModelArchitecture(TypedDict):
@@ -40,12 +38,10 @@ class Model(TypedDict):
per_request_limits: dict | None
async def fetch_openrouter_models() -> list[Model]:
def fetch_openrouter_models() -> list[Model]:
"""Fetches model information from OpenRouter API."""
async with httpx.AsyncClient() as client:
response = await client.get("https://openrouter.ai/api/v1/models")
response.raise_for_status()
data = response.json()
with urlopen("https://openrouter.ai/api/v1/models") as response:
data = json.loads(response.read().decode("utf-8"))
models_data: list[Model] = []
for model in data.get("data", []):
@@ -64,8 +60,8 @@ async def fetch_openrouter_models() -> list[Model]:
return models_data
async def main() -> None:
models = await fetch_openrouter_models()
def main() -> None:
models = fetch_openrouter_models()
# Print the first model data in a nicely indented JSON format
print(json.dumps(models[0], indent=4))
@@ -75,4 +71,4 @@ async def main() -> None:
if __name__ == "__main__":
asyncio.run(main())
main()
-228
View File
@@ -1,228 +0,0 @@
import hashlib
import uuid
from unittest.mock import AsyncMock, patch
import pytest
import pytest_asyncio
from httpx import AsyncClient
from router.db import ApiKey, AsyncSession
def hash_api_key(api_key: str) -> str:
"""Hash an API key for storage."""
return hashlib.sha256(api_key.encode()).hexdigest()
@pytest_asyncio.fixture
async def test_api_key(test_session: AsyncSession) -> ApiKey:
"""Create a test API key in the database."""
# Use unique key for each test
unique_id = str(uuid.uuid4())[:8]
api_key = f"test-api-key-{unique_id}"
key = ApiKey(
hashed_key=api_key,
balance=1000000, # 1000 sats in msats
refund_address="test@lightning.address",
total_spent=0,
total_requests=0,
)
test_session.add(key)
await test_session.commit()
await test_session.refresh(key)
return key
@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}"},
)
assert response.status_code == 200
data = response.json()
assert data["api_key"] == f"sk-{test_api_key.hashed_key}"
assert data["balance"] == 1000000
@pytest.mark.asyncio
async def test_account_info_without_auth(async_client: AsyncClient) -> None:
"""Test that account info requires authentication."""
response = await async_client.get("/v1/wallet/")
assert response.status_code == 422 # Missing required header
@pytest.mark.asyncio
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"}
)
assert response.status_code == 401
@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:
mock_refund.return_value = 1000000
response = await async_client.post(
"/v1/wallet/refund",
headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"},
)
assert response.status_code == 200
data = response.json()
assert data["recipient"] == "test@lightning.address"
assert data["msats"] == 1000000
# Verify the API key was deleted after refund
deleted_key = await test_session.get(ApiKey, test_api_key.hashed_key)
assert deleted_key is None
# Verify refund_balance was called
mock_refund.assert_called_once()
@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]
api_key = f"test-key-no-refund-{unique_id}"
key = ApiKey(
hashed_key=api_key,
balance=500000,
refund_address=None,
total_spent=0,
total_requests=0,
)
test_session.add(key)
await test_session.commit()
# Mock the WALLET instance at the router.account module level
with patch("router.account.WALLET") as mock_wallet:
mock_wallet.send = AsyncMock(return_value="cashuBqQSEQ...")
response = await async_client.post(
"/v1/wallet/refund", headers={"Authorization": f"Bearer sk-{api_key}"}
)
assert response.status_code == 200
data = response.json()
assert data["recipient"] is None
assert data["msats"] == 500000
assert data["token"] == "cashuBqQSEQ..."
# Verify wallet.send was called with the correct amount (msats converted to sats)
mock_wallet.send.assert_called_once_with(500)
# Verify the API key was deleted after refund
deleted = await test_session.get(ApiKey, api_key)
assert deleted is None
@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 = 500000 # Return integer msats value
response = await async_client.post(
"/v1/wallet/topup?cashu_token=cashuBqQSEQ...",
headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"},
)
assert response.status_code == 200
data = response.json()
assert data == {"msats": 500000}
# Verify credit_balance was called
mock_credit.assert_called_once()
@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",
headers={"Authorization": f"Bearer sk-{test_api_key.hashed_key}"},
json={},
)
assert response.status_code == 422 # Missing required field
@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"
async def mock_credit_balance(
token: str, key: ApiKey, session: AsyncSession
) -> int:
"""Mock credit_balance function that simulates adding balance and committing."""
amount = 5000000 # 5000 sats in msats
key.balance += amount
session.add(key)
await session.commit()
return amount
with patch(
"router.cashu.credit_balance",
new_callable=AsyncMock,
side_effect=mock_credit_balance,
):
response = await async_client.get(
"/v1/wallet/info", headers={"Authorization": f"Bearer {cashu_token}"}
)
assert response.status_code == 200
data = response.json()
# Check that a new key was created with the hashed token
assert data["api_key"].startswith("sk-")
assert data["balance"] >= 0 # Balance should be set after credit_balance
@pytest.mark.asyncio
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:
mock_credit.return_value = 0
response = await async_client.get(
"/v1/wallet/info", headers={"Authorization": "Bearer cashuInvalid"}
)
assert response.status_code == 401
error = response.json()
assert error["detail"]["error"]["code"] == "invalid_api_key"
-2
View File
@@ -27,12 +27,10 @@ async def test_root_endpoint(async_client: AsyncClient) -> None:
# The app reads from env vars during import, so check what we actually get
assert "name" in data
assert "description" in data
assert data["version"] == "0.0.1"
assert "npub" in data
assert "mint" in data
assert "http_url" in data
assert "onion_url" in data
assert "models" in data
@pytest.mark.asyncio