mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-06 12:38:23 +00:00
Merge branch 'main' into bug-fixes
This commit is contained in:
+2
-2
@@ -3,9 +3,9 @@ __pycache__
|
||||
keys.db
|
||||
wallet.sqlite3
|
||||
|
||||
|
||||
# Development
|
||||
.notes
|
||||
.keys.db
|
||||
.wallet.sqlite3
|
||||
.models.json
|
||||
.models.json
|
||||
compose.override.yml
|
||||
|
||||
@@ -1,12 +1,4 @@
|
||||
# proxy
|
||||
|
||||
a reverse proxy that you can plug in front of any openai compatible api endpoint
|
||||
to handle api-key based payments using cashu tokens or bold12 lightning invoices
|
||||
|
||||
we also want to provice an internal dashboard
|
||||
|
||||
- general settings
|
||||
- configure payment methods
|
||||
- request evals
|
||||
- publish your listing to nostr
|
||||
- monitor traffic
|
||||
to handle payments using the cashu protocol (Bitcoin L3)
|
||||
|
||||
+3
-2
@@ -7,9 +7,10 @@ services:
|
||||
- .:/app
|
||||
env_file:
|
||||
- .env
|
||||
environment:
|
||||
- TOR_PROXY_URL=socks5://tor:9050
|
||||
ports:
|
||||
# Allows connectio via nginx
|
||||
- "8000:8000"
|
||||
- 8000:8000
|
||||
|
||||
tor:
|
||||
image: ghcr.io/hundehausen/tor-hidden-service:latest
|
||||
|
||||
+4845
-869
File diff suppressed because it is too large
Load Diff
+109
-12
@@ -1,9 +1,11 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import json
|
||||
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from .cashu import credit_balance, pay_out
|
||||
from .cashu import credit_balance, pay_out_with_new_session
|
||||
from .db import ApiKey, AsyncSession
|
||||
from .models import MODELS
|
||||
|
||||
@@ -26,7 +28,16 @@ async def validate_bearer_key(bearer_key: str, session: AsyncSession) -> ApiKey:
|
||||
Otherwise checks if the hash of the key exists.
|
||||
"""
|
||||
if not bearer_key:
|
||||
raise HTTPException(status_code=401, detail="api-key or cashu-token required")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "API key or Cashu token required",
|
||||
"type": "invalid_request_error",
|
||||
"code": "missing_api_key"
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
if bearer_key.startswith("sk-"):
|
||||
if exsisting_key := await session.get(ApiKey, bearer_key[3:]):
|
||||
@@ -44,14 +55,69 @@ async def validate_bearer_key(bearer_key: str, session: AsyncSession) -> ApiKey:
|
||||
except Exception as e:
|
||||
print(f"Redemption failed: {e}")
|
||||
raise HTTPException(
|
||||
status_code=401, detail=f"Invalid or expired cashu key: {e}"
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Invalid or expired Cashu key: {str(e)}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_api_key"
|
||||
}
|
||||
}
|
||||
)
|
||||
raise HTTPException(status_code=401, detail="Invalid API key")
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Invalid API key",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_api_key"
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def pay_for_request(key: ApiKey, session: AsyncSession) -> None:
|
||||
async def pay_for_request(key: ApiKey, session: AsyncSession, request: Request | None, request_body: bytes | None = None) -> None:
|
||||
if MODEL_BASED_PRICING and os.path.exists("models.json"):
|
||||
if request_body:
|
||||
body = json.loads(request_body)
|
||||
else:
|
||||
body = await request.json()
|
||||
if request_model := body.get("model"):
|
||||
if request_model not in [model.id for model in MODELS]:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Invalid model: {request_model}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "model_not_found"
|
||||
}
|
||||
}
|
||||
)
|
||||
model = next(model for model in MODELS if model.id == request_model)
|
||||
if key.balance < model.sats_pricing.max_cost * 1000:
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"This model requires a minimum balance of {model.sats_pricing.max_cost} sats",
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance"
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
if key.balance < COST_PER_REQUEST:
|
||||
raise HTTPException(status_code=402, detail=f"Insufficient balance: {COST_PER_REQUEST} mSats required. {key.balance} available.")
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Insufficient balance: {COST_PER_REQUEST} mSats required. {key.balance} available.",
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance"
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
# Charge the base cost for the request
|
||||
key.balance -= COST_PER_REQUEST
|
||||
@@ -77,19 +143,47 @@ async def adjust_payment_for_tokens(
|
||||
"total_msats": COST_PER_REQUEST,
|
||||
}
|
||||
|
||||
# Check if we have usage data
|
||||
if "usage" not in response_data or response_data["usage"] is None:
|
||||
print("No usage data in response, using base cost only")
|
||||
return cost_data
|
||||
|
||||
# Default to configured pricing
|
||||
MSATS_PER_1K_INPUT_TOKENS = COST_PER_1K_INPUT_TOKENS
|
||||
MSATS_PER_1K_OUTPUT_TOKENS = COST_PER_1K_OUTPUT_TOKENS
|
||||
|
||||
if MODEL_BASED_PRICING and os.path.exists("models.json"):
|
||||
response_model = response_data.get("model", "")
|
||||
if response_model not in [model.id for model in MODELS]:
|
||||
raise HTTPException(status_code=400, detail="Invalid model")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Invalid model in response: {response_model}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "model_not_found"
|
||||
}
|
||||
}
|
||||
)
|
||||
model = next(model for model in MODELS if model.id == response_model)
|
||||
if model.sats_pricing is None:
|
||||
raise HTTPException(status_code=400, detail="Model pricing not defined")
|
||||
# TODO: Rename, This is named very close to COST_PER_1K_OUTPUT_TOKENS
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Model pricing not defined",
|
||||
"type": "invalid_request_error",
|
||||
"code": "pricing_not_found"
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
MSATS_PER_1K_INPUT_TOKENS = model.sats_pricing.prompt * 1_000_000
|
||||
MSATS_PER_1K_OUTPUT_TOKENS = model.sats_pricing.completion * 1_000_000
|
||||
|
||||
if not (MSATS_PER_1K_OUTPUT_TOKENS and MSATS_PER_1K_INPUT_TOKENS):
|
||||
raise HTTPException(status_code=400, detail="Model pricing not defined")
|
||||
# If no token pricing is configured, just return base cost
|
||||
return cost_data
|
||||
|
||||
input_tokens = response_data.get("usage", {}).get("prompt_tokens", 0)
|
||||
output_tokens = response_data.get("usage", {}).get("completion_tokens", 0)
|
||||
@@ -117,6 +211,9 @@ async def adjust_payment_for_tokens(
|
||||
f"Warning: Insufficient balance for token-based pricing adjustment: {key.hashed_key[:10]}..."
|
||||
)
|
||||
# Still proceed but log the issue - we already provided the service
|
||||
# Add information about insufficient balance to cost data
|
||||
cost_data["warning"] = "Insufficient balance for full token-based pricing"
|
||||
cost_data["balance_shortage_msats"] = cost_difference - key.balance
|
||||
else:
|
||||
key.balance -= cost_difference
|
||||
key.total_spent += cost_difference
|
||||
@@ -131,6 +228,6 @@ async def adjust_payment_for_tokens(
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
await pay_out(session)
|
||||
asyncio.create_task(pay_out_with_new_session())
|
||||
|
||||
return cost_data
|
||||
|
||||
+48
-25
@@ -10,6 +10,7 @@ from .db import ApiKey, AsyncSession
|
||||
RECEIVE_LN_ADDRESS = os.environ["RECEIVE_LN_ADDRESS"]
|
||||
MINT = os.environ.get("MINT", "https://mint.minibits.cash/Bitcoin")
|
||||
MINIMUM_PAYOUT = int(os.environ.get("MINIMUM_PAYOUT", 100))
|
||||
DEV_LN_ADDRESS = "routstr@minibits.cash"
|
||||
DEVS_DONATION_RATE = 0.021 # 2.1%
|
||||
WALLET = None
|
||||
|
||||
@@ -85,42 +86,64 @@ async def _pay_invoice_with_cashu(
|
||||
return quote.amount
|
||||
|
||||
|
||||
async def pay_out_with_new_session() -> None:
|
||||
"""
|
||||
Wrapper for pay_out that creates its own database session.
|
||||
This prevents database connection conflicts when called as a background task.
|
||||
"""
|
||||
from .db import create_session
|
||||
|
||||
try:
|
||||
async with create_session() as session:
|
||||
await pay_out(session)
|
||||
except Exception as e:
|
||||
print(f"Error in pay_out_with_new_session: {e}")
|
||||
|
||||
|
||||
async def pay_out(session: AsyncSession) -> None:
|
||||
"""
|
||||
Calculates the pay-out amount based on the spent balance, profit, and donation rate.
|
||||
"""
|
||||
balance = (
|
||||
await session.exec(
|
||||
select(func.sum(col(ApiKey.balance))).where(ApiKey.balance > 0)
|
||||
)
|
||||
).one()
|
||||
if balance is None:
|
||||
raise ValueError("No balance to pay out.")
|
||||
user_balance = balance // 1000 # conversion to sats
|
||||
wallet = await _initialize_wallet()
|
||||
wallet_balance = wallet.available_balance
|
||||
try:
|
||||
balance = (
|
||||
await session.exec(
|
||||
select(func.sum(col(ApiKey.balance))).where(ApiKey.balance > 0)
|
||||
)
|
||||
).one()
|
||||
if balance is None or balance == 0:
|
||||
# No balance to pay out - this is OK, not an error
|
||||
return
|
||||
|
||||
user_balance_sats = balance // 1000 # Convert msats to sats
|
||||
wallet = await _initialize_wallet()
|
||||
wallet_balance_sats = wallet.available_balance # Already in sats
|
||||
|
||||
assert wallet_balance <= user_balance, "Something went deeply wrong."
|
||||
# Handle edge cases more gracefully
|
||||
if wallet_balance_sats < user_balance_sats:
|
||||
print(f"Warning: Wallet balance ({wallet_balance_sats} sats) is less than user balance ({user_balance_sats} sats). Skipping payout.")
|
||||
return
|
||||
|
||||
print(f"Wallet-balance: {wallet_balance}, User-balance: {user_balance}, Revenue: {wallet_balance - user_balance}, MinPayout:{MINIMUM_PAYOUT}", flush=True)
|
||||
# Why is that bad?
|
||||
#assert wallet_balance <= user_balance, f"Something went deeply wrong. Wallet-balance: {wallet_balance}, User-Balance: {user_balance}"
|
||||
if (revenue := wallet_balance - user_balance) <= MINIMUM_PAYOUT:
|
||||
return
|
||||
if (revenue := wallet_balance_sats - user_balance_sats) <= MINIMUM_PAYOUT:
|
||||
# Not enough revenue yet - this is OK
|
||||
return
|
||||
|
||||
devs_donation = int(revenue * DEVS_DONATION_RATE)
|
||||
owners_draw = revenue - devs_donation
|
||||
devs_donation = int(revenue * DEVS_DONATION_RATE)
|
||||
owners_draw = revenue - devs_donation
|
||||
|
||||
|
||||
await send_to_lnurl(wallet, RECEIVE_LN_ADDRESS, owners_draw * 1000) # conversion to msats for send_to_lnurl
|
||||
|
||||
if devs_donation > 0:
|
||||
# Send payouts
|
||||
print(f"Sending {owners_draw} sats to {RECEIVE_LN_ADDRESS}")
|
||||
await send_to_lnurl(wallet, RECEIVE_LN_ADDRESS, owners_draw * 1000) # Convert to msats
|
||||
print(f"Sending {devs_donation} sats to {DEV_LN_ADDRESS}")
|
||||
await send_to_lnurl(
|
||||
wallet,
|
||||
"npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash",
|
||||
devs_donation * 1000,
|
||||
DEV_LN_ADDRESS,
|
||||
devs_donation * 1000, # Convert to msats
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
# Log the error but don't crash - payouts can be retried later
|
||||
print(f"Error in pay_out: {e}")
|
||||
|
||||
|
||||
async def credit_balance(cashu_token: str, key: ApiKey, session: AsyncSession) -> int:
|
||||
token_obj: Token = deserialize_token_from_string(cashu_token)
|
||||
@@ -246,7 +269,7 @@ async def get_lnurl_data(lnurl: str) -> tuple[str, int, int]:
|
||||
elif lnurl.lower().startswith("lnurl"):
|
||||
try:
|
||||
# Optional import for environments where bech32 might not be present initially
|
||||
from bech32 import bech32_decode, convertbits
|
||||
from bech32 import bech32_decode, convertbits # type: ignore
|
||||
|
||||
hrp, data = bech32_decode(lnurl)
|
||||
if data is None:
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
from fastapi import APIRouter
|
||||
import asyncio
|
||||
import json
|
||||
import websockets
|
||||
import random
|
||||
import string
|
||||
import re
|
||||
import httpx
|
||||
import os
|
||||
|
||||
providers_router = APIRouter(prefix="/v1/providers")
|
||||
|
||||
|
||||
def generate_subscription_id() -> str:
|
||||
"""Generate a random subscription ID."""
|
||||
return "".join(random.choices(string.ascii_lowercase + string.digits, k=10))
|
||||
|
||||
|
||||
def extract_onion_urls(content: str) -> list[str]:
|
||||
"""Extract onion URLs from content."""
|
||||
pattern = r"http?://[a-zA-Z0-9\-._~]+\.onion"
|
||||
return re.findall(pattern, content)
|
||||
|
||||
|
||||
async def query_nostr_relay_with_search(
|
||||
search_term: str,
|
||||
relay_url: str,
|
||||
kinds: list[int] | None = None,
|
||||
limit: int = 1000,
|
||||
timeout: int = 30,
|
||||
) -> list[dict]:
|
||||
"""
|
||||
Query a Nostr relay and filter for events containing a search term.
|
||||
"""
|
||||
if kinds is None:
|
||||
kinds = [1]
|
||||
|
||||
events = []
|
||||
|
||||
# If searching for an npub mention, try tag-based search first
|
||||
if search_term.startswith("nostr:npub"):
|
||||
# Extract the npub and convert to hex
|
||||
npub = search_term.replace("nostr:", "")
|
||||
try:
|
||||
# Convert npub to hex (you might need to implement or import this)
|
||||
# For now, try tag-based search with the npub
|
||||
filter_obj = {
|
||||
"kinds": kinds,
|
||||
"limit": limit,
|
||||
"#p": [npub], # Posts that tag this pubkey
|
||||
}
|
||||
except:
|
||||
# If conversion fails, try regular search
|
||||
filter_obj = {
|
||||
"kinds": kinds,
|
||||
"limit": limit,
|
||||
}
|
||||
else:
|
||||
# Try relay's search functionality (NIP-50)
|
||||
filter_obj = {
|
||||
"kinds": kinds,
|
||||
"search": search_term,
|
||||
"limit": limit,
|
||||
}
|
||||
|
||||
sub_id = generate_subscription_id()
|
||||
req_message = json.dumps(["REQ", sub_id, filter_obj])
|
||||
|
||||
try:
|
||||
async with websockets.connect(relay_url, timeout=timeout) as websocket:
|
||||
print(f"Connected to relay, sending request with filter: {filter_obj}")
|
||||
await websocket.send(req_message)
|
||||
|
||||
while True:
|
||||
try:
|
||||
message = await asyncio.wait_for(websocket.recv(), timeout=5)
|
||||
data = json.loads(message)
|
||||
|
||||
if data[0] == "EVENT" and data[1] == sub_id:
|
||||
# For tag-based search, also check content
|
||||
if search_term.startswith("nostr:npub"):
|
||||
if search_term.lower() in data[2]["content"].lower():
|
||||
print(f"Found matching event: {data[2]['id']}")
|
||||
events.append(data[2])
|
||||
else:
|
||||
print(f"Found matching event: {data[2]['id']}")
|
||||
events.append(data[2])
|
||||
elif data[0] == "EOSE" and data[1] == sub_id:
|
||||
print("Received EOSE message")
|
||||
break
|
||||
elif data[0] == "NOTICE":
|
||||
print(f"Relay notice: {data[1]}")
|
||||
# If search not supported, could break and try different approach
|
||||
if "unrecognised filter item" in data[1] and "search" in str(
|
||||
filter_obj
|
||||
):
|
||||
print("Search not supported on this relay")
|
||||
break
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
print("Timeout waiting for message")
|
||||
break
|
||||
except json.JSONDecodeError:
|
||||
print("Failed to decode message as JSON")
|
||||
continue
|
||||
|
||||
await websocket.send(json.dumps(["CLOSE", sub_id]))
|
||||
|
||||
except Exception as e:
|
||||
print(f"Query failed: {e}")
|
||||
|
||||
print(f"Query complete. Found {len(events)} matching events")
|
||||
return events
|
||||
|
||||
|
||||
async def get_cache() -> list[dict]:
|
||||
return [] # TODO: Implement cache
|
||||
|
||||
|
||||
async def fetch_onion(provider: str) -> dict:
|
||||
"""Check if an onion service is healthy by making a GET request to its root."""
|
||||
try:
|
||||
# Get Tor proxy URL from environment variable, default to local Tor SOCKS5 proxy
|
||||
tor_proxy = os.getenv("TOR_PROXY_URL", "socks5://127.0.0.1:9050")
|
||||
|
||||
# Configure httpx to use Tor SOCKS5 proxy
|
||||
async with httpx.AsyncClient(
|
||||
proxies={"http://": tor_proxy, "https://": tor_proxy},
|
||||
timeout=httpx.Timeout(30.0),
|
||||
follow_redirects=True,
|
||||
) as client:
|
||||
response = await client.get(provider)
|
||||
# Consider 2xx and 3xx status codes as healthy
|
||||
return {"status_code": response.status_code, "json": response.json()}
|
||||
except Exception:
|
||||
# Any exception means the service is not healthy
|
||||
return {"status_code": 500, "json": {"error": "Failed to fetch onion"}}
|
||||
|
||||
|
||||
@providers_router.get("/")
|
||||
async def get_providers(include_json: bool = False):
|
||||
npub = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s"
|
||||
|
||||
# Relays that support NIP-50 text search
|
||||
search_relays = [
|
||||
"wss://relay.nostr.band", # Known to support search
|
||||
"wss://nostr.wine", # Known to support search
|
||||
"wss://relay.damus.io",
|
||||
"wss://nos.lol",
|
||||
]
|
||||
|
||||
# Search for the mention format that appears in posts
|
||||
search_term = f"nostr:{npub}"
|
||||
|
||||
all_events = []
|
||||
event_ids = set() # To avoid duplicates
|
||||
|
||||
# Try multiple relays
|
||||
for relay_url in search_relays:
|
||||
print(f"\nTrying relay: {relay_url}")
|
||||
try:
|
||||
events = await query_nostr_relay_with_search(
|
||||
search_term=search_term,
|
||||
relay_url=relay_url,
|
||||
kinds=[1], # Text notes
|
||||
limit=500,
|
||||
)
|
||||
|
||||
# Add unique events
|
||||
for event in events:
|
||||
if event["id"] not in event_ids:
|
||||
event_ids.add(event["id"])
|
||||
all_events.append(event)
|
||||
|
||||
print(f"Got {len(events)} events from {relay_url}")
|
||||
|
||||
# If we have enough events, we can stop
|
||||
if len(all_events) >= 100:
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
print(f"Failed to query {relay_url}: {e}")
|
||||
continue
|
||||
|
||||
print(f"Found {len(all_events)} total unique events mentioning routstr")
|
||||
|
||||
providers = []
|
||||
for event in all_events:
|
||||
onion_urls = extract_onion_urls(event["content"])
|
||||
providers.extend(onion_urls)
|
||||
|
||||
unique_providers = list(set(providers))
|
||||
|
||||
print(f"Found {len(unique_providers)} unique onion URLs")
|
||||
print(unique_providers)
|
||||
|
||||
healthy_providers: list[dict | str] = []
|
||||
for provider in unique_providers:
|
||||
response = await fetch_onion(provider)
|
||||
|
||||
if include_json:
|
||||
healthy_providers.append({provider: response["json"]})
|
||||
else:
|
||||
healthy_providers.append(provider)
|
||||
|
||||
return {"providers": healthy_providers}
|
||||
@@ -9,6 +9,7 @@ from .proxy import proxy_router
|
||||
from .account import account_router
|
||||
from .cashu import _initialize_wallet
|
||||
from .models import MODELS, update_sats_pricing
|
||||
from .discovery import providers_router
|
||||
|
||||
__version__ = "0.0.1"
|
||||
|
||||
@@ -45,6 +46,7 @@ async def info():
|
||||
|
||||
app.include_router(admin_router)
|
||||
app.include_router(account_router)
|
||||
app.include_router(providers_router)
|
||||
app.include_router(proxy_router)
|
||||
|
||||
|
||||
|
||||
+21
-2
@@ -20,7 +20,12 @@ class Pricing(BaseModel):
|
||||
image: float
|
||||
web_search: float
|
||||
internal_reasoning: float
|
||||
max_cost: float = 0.0 # in sats not msats
|
||||
|
||||
class TopProvider(BaseModel):
|
||||
context_length: int | None = None
|
||||
max_completion_tokens: int | None = None
|
||||
is_moderated: bool | None = None
|
||||
|
||||
class Model(BaseModel):
|
||||
id: str
|
||||
@@ -30,8 +35,9 @@ class Model(BaseModel):
|
||||
context_length: int
|
||||
architecture: Architecture
|
||||
pricing: Pricing
|
||||
sats_pricing: Pricing | None
|
||||
per_request_limits: dict | None
|
||||
sats_pricing: Pricing | None = None
|
||||
per_request_limits: dict | None = None
|
||||
top_provider: TopProvider | None = None
|
||||
|
||||
|
||||
MODELS: list[Model] = []
|
||||
@@ -48,6 +54,19 @@ async def update_sats_pricing() -> None:
|
||||
model.sats_pricing = Pricing(
|
||||
**{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
|
||||
)
|
||||
if model.top_provider:
|
||||
if model.top_provider.context_length and model.top_provider.max_completion_tokens:
|
||||
max_context_cost = model.top_provider.context_length * model.sats_pricing.prompt
|
||||
max_completion_cost = model.top_provider.max_completion_tokens * model.sats_pricing.completion
|
||||
model.sats_pricing.max_cost = max_context_cost + max_completion_cost
|
||||
else:
|
||||
p = model.sats_pricing.prompt * 1_000_000
|
||||
c = model.sats_pricing.completion * 32_000
|
||||
r = model.sats_pricing.request * 100_000
|
||||
i = model.sats_pricing.image * 100
|
||||
w = model.sats_pricing.web_search * 1000
|
||||
ir = model.sats_pricing.internal_reasoning * 100
|
||||
model.sats_pricing.max_cost = p + c + r + i + w + ir
|
||||
except Exception as e:
|
||||
print(e)
|
||||
await asyncio.sleep(10)
|
||||
|
||||
+110
-30
@@ -6,7 +6,7 @@ from fastapi.responses import Response, StreamingResponse
|
||||
import httpx
|
||||
import re
|
||||
|
||||
from router.cashu import pay_out
|
||||
from router.cashu import pay_out_with_new_session
|
||||
|
||||
from .auth import validate_bearer_key, pay_for_request, adjust_payment_for_tokens
|
||||
from .db import AsyncSession, get_session
|
||||
@@ -27,7 +27,41 @@ async def proxy(
|
||||
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
|
||||
|
||||
key = await validate_bearer_key(bearer_key, session)
|
||||
await pay_for_request(key, session)
|
||||
|
||||
# Pre-validate JSON for requests that require it
|
||||
request_body = None
|
||||
if request.method in ["POST", "PUT", "PATCH"] and path.endswith("chat/completions"):
|
||||
try:
|
||||
request_body = await request.body()
|
||||
# Try to parse JSON to validate it
|
||||
if request_body:
|
||||
json.loads(request_body)
|
||||
except json.JSONDecodeError as e:
|
||||
return Response(
|
||||
content=json.dumps({
|
||||
"error": {
|
||||
"message": f"Invalid JSON in request body: {str(e)}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_json"
|
||||
}
|
||||
}),
|
||||
status_code=400,
|
||||
media_type="application/json"
|
||||
)
|
||||
except Exception as e:
|
||||
return Response(
|
||||
content=json.dumps({
|
||||
"error": {
|
||||
"message": "Error reading request body",
|
||||
"type": "invalid_request_error",
|
||||
"code": "request_error"
|
||||
}
|
||||
}),
|
||||
status_code=400,
|
||||
media_type="application/json"
|
||||
)
|
||||
|
||||
await pay_for_request(key, session, request, request_body)
|
||||
|
||||
# Prepare headers, removing sensitive/problematic ones
|
||||
headers = dict(request.headers)
|
||||
@@ -45,19 +79,35 @@ async def proxy(
|
||||
path = path.replace("v1/", "")
|
||||
|
||||
url = f"{UPSTREAM_BASE_URL}/{path}"
|
||||
client = httpx.AsyncClient(transport=httpx.AsyncHTTPTransport(retries=1))
|
||||
client = httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=None # No timeout - requests can take as long as needed
|
||||
)
|
||||
|
||||
try:
|
||||
response = await client.send(
|
||||
client.build_request(
|
||||
request.method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=request.stream(),
|
||||
params=request.query_params,
|
||||
),
|
||||
stream=True,
|
||||
)
|
||||
# Use the pre-read body if available, otherwise stream
|
||||
if request_body is not None:
|
||||
response = await client.send(
|
||||
client.build_request(
|
||||
request.method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=request_body,
|
||||
params=request.query_params,
|
||||
),
|
||||
stream=True,
|
||||
)
|
||||
else:
|
||||
response = await client.send(
|
||||
client.build_request(
|
||||
request.method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=request.stream(),
|
||||
params=request.query_params,
|
||||
),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
# For chat completions, we need to handle token-based pricing
|
||||
if path.endswith("chat/completions"):
|
||||
@@ -114,17 +164,10 @@ async def proxy(
|
||||
usage_data_found = True
|
||||
break
|
||||
except json.JSONDecodeError:
|
||||
# Not valid JSON, skip
|
||||
continue
|
||||
|
||||
if usage_data_found:
|
||||
break
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing chunk for cost: {e}")
|
||||
|
||||
if not usage_data_found:
|
||||
print("No usage data found in any chunks")
|
||||
print(f"Error processing streaming response for cost: {e}")
|
||||
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -146,7 +189,6 @@ async def proxy(
|
||||
key, response_json, session
|
||||
)
|
||||
response_json["cost"] = cost_data
|
||||
asyncio.create_task(pay_out(session))
|
||||
return Response(
|
||||
content=json.dumps(response_json).encode(),
|
||||
status_code=response.status_code,
|
||||
@@ -165,8 +207,8 @@ async def proxy(
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
background_tasks.add_task(pay_out_with_new_session)
|
||||
|
||||
asyncio.create_task(pay_out(session))
|
||||
return StreamingResponse(
|
||||
response.aiter_bytes(),
|
||||
status_code=response.status_code,
|
||||
@@ -176,15 +218,53 @@ async def proxy(
|
||||
|
||||
except httpx.RequestError as exc:
|
||||
await client.aclose()
|
||||
print(f"Error forwarding request to upstream: {exc}")
|
||||
error_type = type(exc).__name__
|
||||
error_details = str(exc)
|
||||
print(
|
||||
f"Error forwarding request to upstream: {error_type}: {error_details}\n"
|
||||
f"Request details: method={request.method}, url={url}, headers={headers}, "
|
||||
f"path={path}, query_params={dict(request.query_params)}"
|
||||
)
|
||||
|
||||
# Provide more specific error messages based on the error type
|
||||
if isinstance(exc, httpx.ConnectError):
|
||||
error_message = "Unable to connect to upstream service"
|
||||
elif isinstance(exc, httpx.TimeoutException):
|
||||
error_message = "Upstream service request timed out"
|
||||
elif isinstance(exc, httpx.NetworkError):
|
||||
error_message = "Network error while connecting to upstream service"
|
||||
else:
|
||||
error_message = f"Error connecting to upstream service: {error_type}"
|
||||
|
||||
return Response(
|
||||
content=f"Error connecting to upstream service: {exc}",
|
||||
content=json.dumps({
|
||||
"error": {
|
||||
"message": error_message,
|
||||
"type": "upstream_error",
|
||||
"code": 502
|
||||
}
|
||||
}),
|
||||
status_code=502,
|
||||
media_type="application/json"
|
||||
)
|
||||
except Exception as exc:
|
||||
await client.aclose()
|
||||
print(f"Unexpected error: {exc}")
|
||||
return Response(
|
||||
content=f"Unexpected server error: {exc}",
|
||||
status_code=500,
|
||||
import traceback
|
||||
tb = traceback.format_exc()
|
||||
print(
|
||||
f"Unexpected error: {exc}\n"
|
||||
f"Request details: method={request.method}, url={url}, headers={headers}, "
|
||||
f"path={path}, query_params={dict(request.query_params)}\n"
|
||||
f"Traceback:\n{tb}"
|
||||
)
|
||||
return Response(
|
||||
content=json.dumps({
|
||||
"error": {
|
||||
"message": "An unexpected server error occurred",
|
||||
"type": "internal_error",
|
||||
"code": 500
|
||||
}
|
||||
}),
|
||||
status_code=500,
|
||||
media_type="application/json"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user