mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-22 12:22:20 +00:00
Compare commits
5 Commits
feature/fa
...
283-timeou
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1df66f48b8 | ||
|
|
a9c5458660 | ||
|
|
5a4ba60072 | ||
|
|
d41c214d9e | ||
|
|
ec0fcfb48b |
@@ -1,5 +1,7 @@
|
||||
"""Model prioritization algorithm for selecting cheapest upstream providers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from .core.logging import get_logger
|
||||
@@ -157,15 +159,21 @@ def create_model_mappings(
|
||||
upstreams: list["BaseUpstreamProvider"],
|
||||
overrides_by_id: dict[str, tuple],
|
||||
disabled_model_ids: set[str],
|
||||
) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]:
|
||||
) -> tuple[
|
||||
dict[str, "Model"],
|
||||
dict[str, "BaseUpstreamProvider"],
|
||||
dict[str, "Model"],
|
||||
dict[str, list["BaseUpstreamProvider"]],
|
||||
]:
|
||||
"""Create optimal model mappings based on cost and provider preferences.
|
||||
|
||||
This is the main entry point for the algorithm. It processes all upstream providers
|
||||
and creates three mappings based on cost optimization:
|
||||
|
||||
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
|
||||
2. provider_map: alias -> UpstreamProvider (which provider to use for each alias)
|
||||
2. provider_map: alias -> UpstreamProvider (the BEST provider to use for each alias)
|
||||
3. unique_models: base_id -> Model (unique models without provider prefixes)
|
||||
4. provider_candidates_map: alias -> list[UpstreamProvider] (all providers offering the model, sorted by preference)
|
||||
|
||||
The algorithm:
|
||||
- Processes non-OpenRouter providers first (they're typically cheaper)
|
||||
@@ -178,7 +186,7 @@ def create_model_mappings(
|
||||
disabled_model_ids: Set of model IDs that should be excluded
|
||||
|
||||
Returns:
|
||||
Tuple of (model_instances, provider_map, unique_models)
|
||||
Tuple of (model_instances, provider_map, unique_models, provider_candidates_map)
|
||||
"""
|
||||
from .payment.models import _row_to_model
|
||||
from .upstream.helpers import resolve_model_alias
|
||||
@@ -186,9 +194,10 @@ def create_model_mappings(
|
||||
model_instances: dict[str, "Model"] = {}
|
||||
provider_map: dict[str, "BaseUpstreamProvider"] = {}
|
||||
unique_models: dict[str, "Model"] = {}
|
||||
provider_candidates_map: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||
|
||||
# Separate OpenRouter from other providers
|
||||
openrouter: "BaseUpstreamProvider" | None = None
|
||||
openrouter: BaseUpstreamProvider | None = None
|
||||
other_upstreams: list["BaseUpstreamProvider"] = []
|
||||
|
||||
for upstream in upstreams:
|
||||
@@ -207,6 +216,13 @@ def create_model_mappings(
|
||||
) -> None:
|
||||
"""Set alias to model/provider if not set or if new model is preferred."""
|
||||
alias_lower = alias.lower()
|
||||
|
||||
# Add to candidates list, to be used later as fallback
|
||||
if alias_lower not in provider_candidates_map:
|
||||
provider_candidates_map[alias_lower] = [provider]
|
||||
else:
|
||||
provider_candidates_map[alias_lower].append(provider)
|
||||
|
||||
existing_model = model_instances.get(alias_lower)
|
||||
if not existing_model:
|
||||
# No existing mapping, set it
|
||||
@@ -276,6 +292,26 @@ def create_model_mappings(
|
||||
if openrouter:
|
||||
process_provider_models(openrouter, is_openrouter=True)
|
||||
|
||||
# Sort and filter provider candidates for each alias using provider_map as reference
|
||||
# We only keep entries that have more than one provider.
|
||||
final_candidates_map: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||
for alias_lower, best_provider in provider_map.items():
|
||||
candidates = provider_candidates_map.get(alias_lower, [])
|
||||
# Remove duplicates
|
||||
unique_candidates = []
|
||||
seen = set()
|
||||
for c in candidates:
|
||||
if c not in seen:
|
||||
unique_candidates.append(c)
|
||||
seen.add(c)
|
||||
|
||||
if len(unique_candidates) > 1:
|
||||
# Keep the best one at the front, others follow.
|
||||
if best_provider in unique_candidates:
|
||||
unique_candidates.remove(best_provider)
|
||||
unique_candidates.insert(0, best_provider)
|
||||
final_candidates_map[alias_lower] = unique_candidates
|
||||
|
||||
# Log provider distribution
|
||||
provider_counts: dict[str, int] = {}
|
||||
for provider in provider_map.values():
|
||||
@@ -287,4 +323,4 @@ def create_model_mappings(
|
||||
extra={"provider_distribution": provider_counts},
|
||||
)
|
||||
|
||||
return model_instances, provider_map, unique_models
|
||||
return model_instances, provider_map, unique_models, provider_candidates_map
|
||||
|
||||
@@ -6,7 +6,7 @@ import os
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from pydantic.v1 import BaseModel, BaseSettings, Field
|
||||
from pydantic.v1 import BaseModel, BaseSettings, Field, validator
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
|
||||
@@ -37,6 +37,13 @@ class Settings(BaseSettings):
|
||||
|
||||
# Cashu
|
||||
cashu_mints: list[str] = Field(default_factory=list, env="CASHU_MINTS")
|
||||
|
||||
@validator("cashu_mints", pre=True, each_item=True)
|
||||
def normalize_mint_url(cls, v: str) -> str:
|
||||
if isinstance(v, str):
|
||||
return v.rstrip("/")
|
||||
return v
|
||||
|
||||
receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS")
|
||||
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
|
||||
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
|
||||
|
||||
271
routstr/proxy.py
271
routstr/proxy.py
@@ -1,6 +1,7 @@
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from sqlmodel import select
|
||||
@@ -16,6 +17,7 @@ from .core.db import (
|
||||
create_session,
|
||||
get_session,
|
||||
)
|
||||
from .core.settings import settings
|
||||
from .payment.helpers import (
|
||||
calculate_discounted_max_cost,
|
||||
check_token_balance,
|
||||
@@ -25,6 +27,7 @@ from .payment.helpers import (
|
||||
from .payment.models import Model
|
||||
from .upstream import BaseUpstreamProvider
|
||||
from .upstream.helpers import init_upstreams
|
||||
from .wallet import deserialize_token_from_string, recieve_token
|
||||
|
||||
logger = get_logger(__name__)
|
||||
proxy_router = APIRouter()
|
||||
@@ -32,6 +35,9 @@ proxy_router = APIRouter()
|
||||
_upstreams: list[BaseUpstreamProvider] = []
|
||||
_model_instances: dict[str, Model] = {} # All aliases -> Model
|
||||
_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider
|
||||
_provider_candidates_map: dict[
|
||||
str, list[BaseUpstreamProvider]
|
||||
] = {} # All aliases -> [Providers]
|
||||
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
|
||||
|
||||
|
||||
@@ -73,6 +79,21 @@ def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
|
||||
return _provider_map.get(model_id.lower())
|
||||
|
||||
|
||||
def get_providers_for_model(model_id: str) -> list[BaseUpstreamProvider]:
|
||||
"""Get list of prioritized UpstreamProviders for model ID from global cache.
|
||||
|
||||
If multiple providers are available, returns the sorted list.
|
||||
Otherwise, returns a list containing only the single best provider.
|
||||
"""
|
||||
candidates = _provider_candidates_map.get(model_id.lower(), [])
|
||||
if candidates:
|
||||
return candidates
|
||||
|
||||
# Fallback to the single best provider if no multi-provider candidates exist
|
||||
best = get_provider_for_model(model_id)
|
||||
return [best] if best else []
|
||||
|
||||
|
||||
def get_unique_models() -> list[Model]:
|
||||
"""Get list of unique models (no duplicates from aliases)."""
|
||||
return list(_unique_models.values())
|
||||
@@ -82,7 +103,7 @@ async def refresh_model_maps() -> None:
|
||||
"""Refresh global model and provider maps using the cost-based algorithm."""
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
global _model_instances, _provider_map, _unique_models
|
||||
global _model_instances, _provider_map, _unique_models, _provider_candidates_map
|
||||
|
||||
async with create_session() as session:
|
||||
# Fetch all providers with their models in a single logical operation
|
||||
@@ -102,10 +123,12 @@ async def refresh_model_maps() -> None:
|
||||
else:
|
||||
disabled_model_ids.add(model.id)
|
||||
|
||||
_model_instances, _provider_map, _unique_models = create_model_mappings(
|
||||
upstreams=_upstreams,
|
||||
overrides_by_id=overrides_by_id,
|
||||
disabled_model_ids=disabled_model_ids,
|
||||
_model_instances, _provider_map, _unique_models, _provider_candidates_map = (
|
||||
create_model_mappings(
|
||||
upstreams=_upstreams,
|
||||
overrides_by_id=overrides_by_id,
|
||||
disabled_model_ids=disabled_model_ids,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -146,14 +169,32 @@ async def proxy(
|
||||
else:
|
||||
model_id = request_body_dict.get("model", "unknown")
|
||||
|
||||
if "https://testnut.cashu.space" in settings.cashu_mints:
|
||||
try:
|
||||
token_str = None
|
||||
if x_cashu_header := headers.get("x-cashu"):
|
||||
token_str = x_cashu_header
|
||||
elif auth_header := headers.get("authorization"):
|
||||
parts = auth_header.split(" ")
|
||||
if len(parts) > 1 and not parts[1].startswith("sk-"):
|
||||
token_str = parts[1]
|
||||
|
||||
if token_str:
|
||||
token_obj = deserialize_token_from_string(token_str)
|
||||
if token_obj.mint == "https://testnut.cashu.space":
|
||||
model_id = "mock/gpt-420-mock"
|
||||
request_body_dict["model"] = model_id
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
model_obj = get_model_instance(model_id)
|
||||
if not model_obj:
|
||||
return create_error_response(
|
||||
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
||||
)
|
||||
|
||||
upstream = get_provider_for_model(model_id)
|
||||
if not upstream:
|
||||
upstreams = get_providers_for_model(model_id)
|
||||
if not upstreams:
|
||||
return create_error_response(
|
||||
"invalid_model",
|
||||
f"No provider found for model '{model_id}'",
|
||||
@@ -170,14 +211,80 @@ async def proxy(
|
||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||
|
||||
if x_cashu := headers.get("x-cashu", None):
|
||||
if is_responses_api:
|
||||
return await upstream.handle_x_cashu_responses(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
else:
|
||||
return await upstream.handle_x_cashu(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
# Redeem token once before trying any providers
|
||||
amount, unit, mint = await recieve_token(x_cashu)
|
||||
|
||||
# Fallback for X-Cashu payments
|
||||
last_exception = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
try:
|
||||
# Prepare headers for this specific upstream
|
||||
upstream_headers = upstream.prepare_headers(dict(request.headers))
|
||||
|
||||
if is_responses_api:
|
||||
return await upstream.forward_x_cashu_responses_request(
|
||||
request,
|
||||
path,
|
||||
upstream_headers,
|
||||
amount,
|
||||
unit,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
mint,
|
||||
)
|
||||
else:
|
||||
return await upstream.forward_x_cashu_request(
|
||||
request,
|
||||
path,
|
||||
upstream_headers,
|
||||
amount,
|
||||
unit,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
mint,
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.ConnectError) as e:
|
||||
logger.warning(
|
||||
f"Upstream provider {i + 1}/{len(upstreams)} ({upstream.provider_type}) timed out, trying fallback",
|
||||
extra={
|
||||
"model": model_id,
|
||||
"error": str(e),
|
||||
"attempt": i + 1,
|
||||
},
|
||||
)
|
||||
last_exception = e
|
||||
continue
|
||||
|
||||
# If we get here, all providers failed
|
||||
# Since the token was already redeemed, we must issue a refund
|
||||
logger.error(
|
||||
"All providers failed for X-Cashu request, issuing emergency refund",
|
||||
extra={"amount": amount, "unit": unit, "mint": mint},
|
||||
)
|
||||
|
||||
# Try to use the first provider's refund mechanism
|
||||
refund_token = await upstreams[0].send_refund(amount - 60, unit, mint)
|
||||
|
||||
error_message = "All upstream providers timed out"
|
||||
if isinstance(last_exception, httpx.ConnectError):
|
||||
error_message = "Unable to connect to any upstream service"
|
||||
|
||||
error_response = Response(
|
||||
content=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": error_message,
|
||||
"type": "upstream_error",
|
||||
"code": 504,
|
||||
"refund_token": refund_token,
|
||||
}
|
||||
}
|
||||
),
|
||||
status_code=504,
|
||||
media_type="application/json",
|
||||
)
|
||||
error_response.headers["X-Cashu"] = refund_token
|
||||
return error_response
|
||||
|
||||
elif auth := headers.get("authorization", None):
|
||||
key = await get_bearer_token_key(headers, path, session, auth)
|
||||
@@ -192,56 +299,102 @@ async def proxy(
|
||||
)
|
||||
|
||||
logger.debug("Processing unauthenticated GET request", extra={"path": path})
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
return await upstream.forward_get_request(request, path, headers)
|
||||
|
||||
# Try fallback for GET requests too
|
||||
last_exception = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
try:
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
return await upstream.forward_get_request(request, path, headers)
|
||||
except (httpx.TimeoutException, httpx.ConnectError) as e:
|
||||
logger.warning(
|
||||
f"Upstream GET provider {i + 1}/{len(upstreams)} ({upstream.provider_type}) timed out, trying fallback",
|
||||
extra={"path": path, "error": str(e)},
|
||||
)
|
||||
last_exception = e
|
||||
continue
|
||||
|
||||
error_message = "Upstream service request timed out"
|
||||
if isinstance(last_exception, httpx.ConnectError):
|
||||
error_message = "Unable to connect to upstream service"
|
||||
|
||||
return create_error_response(
|
||||
"upstream_error", error_message, 502, request=request
|
||||
)
|
||||
|
||||
if request_body_dict:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
# Fallback for API Key payments
|
||||
last_exception = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
try:
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
if is_responses_api:
|
||||
response = await upstream.forward_responses_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
else:
|
||||
response = await upstream.forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
if is_responses_api:
|
||||
response = await upstream.forward_responses_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
else:
|
||||
response = await upstream.forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
if response.status_code != 200:
|
||||
# If it's a 429 (rate limit) or 503 (service unavailable), we might also want to fallback
|
||||
if response.status_code in (429, 503, 502) and i < len(upstreams) - 1:
|
||||
logger.warning(
|
||||
f"Upstream provider {i + 1}/{len(upstreams)} returned {response.status_code}, trying fallback",
|
||||
extra={"model": model_id, "status": response.status_code},
|
||||
)
|
||||
continue
|
||||
|
||||
if response.status_code != 200:
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
logger.warning(
|
||||
"Upstream request failed, revert payment",
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"key_balance": key.balance,
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
"upstream_headers": response.headers
|
||||
if hasattr(response, "headers")
|
||||
else None,
|
||||
},
|
||||
)
|
||||
# Return the mapped error response generated earlier rather than masking with 502
|
||||
return response
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
logger.warning(
|
||||
"Upstream request failed, revert payment",
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"key_balance": key.balance,
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
"upstream_headers": response.headers
|
||||
if hasattr(response, "headers")
|
||||
else None,
|
||||
},
|
||||
)
|
||||
return response
|
||||
|
||||
return response
|
||||
return response
|
||||
|
||||
except (httpx.TimeoutException, httpx.ConnectError) as e:
|
||||
logger.warning(
|
||||
f"Upstream provider {i + 1}/{len(upstreams)} ({upstream.provider_type}) timed out, trying fallback",
|
||||
extra={"model": model_id, "error": str(e)},
|
||||
)
|
||||
last_exception = e
|
||||
continue
|
||||
|
||||
# All providers failed with timeout/connect error
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
error_message = "Upstream service request timed out"
|
||||
if isinstance(last_exception, httpx.ConnectError):
|
||||
error_message = "Unable to connect to upstream service"
|
||||
|
||||
return create_error_response("upstream_error", error_message, 502, request=request)
|
||||
|
||||
|
||||
async def get_bearer_token_key(
|
||||
@@ -337,7 +490,7 @@ def extract_model_from_responses_request(request_body_dict: dict[str, Any]) -> s
|
||||
|
||||
logger.warning(
|
||||
"No model found in Responses API request",
|
||||
extra={"body_keys": list(request_body_dict.keys())}
|
||||
extra={"body_keys": list(request_body_dict.keys())},
|
||||
)
|
||||
return "unknown"
|
||||
|
||||
|
||||
@@ -37,6 +37,8 @@ from ..wallet import recieve_token, send_token
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
DEFAULT_PROXY_TIMEOUT = 30.0
|
||||
|
||||
|
||||
class TopupData(BaseModel):
|
||||
"""Universal top-up data schema for Lightning Network invoices."""
|
||||
@@ -234,7 +236,11 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
|
||||
# Handle model in input field (alternative format)
|
||||
if "input" in data and isinstance(data["input"], dict) and "model" in data["input"]:
|
||||
if (
|
||||
"input" in data
|
||||
and isinstance(data["input"], dict)
|
||||
and "model" in data["input"]
|
||||
):
|
||||
original_model = model_obj.id
|
||||
transformed_model = self.transform_model_name(original_model)
|
||||
data["input"]["model"] = transformed_model
|
||||
@@ -686,6 +692,7 @@ class BaseUpstreamProvider:
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error processing non-streaming chat completion",
|
||||
@@ -779,8 +786,13 @@ class BaseUpstreamProvider:
|
||||
|
||||
# Track reasoning tokens for Responses API
|
||||
if usage := obj.get("usage", {}):
|
||||
if isinstance(usage, dict) and "reasoning_tokens" in usage:
|
||||
reasoning_tokens += usage.get("reasoning_tokens", 0)
|
||||
if (
|
||||
isinstance(usage, dict)
|
||||
and "reasoning_tokens" in usage
|
||||
):
|
||||
reasoning_tokens += usage.get(
|
||||
"reasoning_tokens", 0
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
except Exception:
|
||||
@@ -933,8 +945,8 @@ class BaseUpstreamProvider:
|
||||
"model": response_json.get("model", "unknown"),
|
||||
"has_usage": "usage" in response_json,
|
||||
"has_reasoning_tokens": "usage" in response_json
|
||||
and isinstance(response_json.get("usage"), dict)
|
||||
and "reasoning_tokens" in response_json["usage"],
|
||||
and isinstance(response_json.get("usage"), dict)
|
||||
and "reasoning_tokens" in response_json["usage"],
|
||||
},
|
||||
)
|
||||
|
||||
@@ -1047,7 +1059,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
client = httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=None,
|
||||
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1190,7 +1202,8 @@ class BaseUpstreamProvider:
|
||||
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"
|
||||
# Re-raise timeout exception to allow fallback handling in proxy layer
|
||||
raise
|
||||
elif isinstance(exc, httpx.NetworkError):
|
||||
error_message = "Network error while connecting to upstream service"
|
||||
else:
|
||||
@@ -1273,7 +1286,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
client = httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=None,
|
||||
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1393,7 +1406,8 @@ class BaseUpstreamProvider:
|
||||
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"
|
||||
# Re-raise timeout exception to allow fallback handling in proxy layer
|
||||
raise
|
||||
elif isinstance(exc, httpx.NetworkError):
|
||||
error_message = "Network error while connecting to upstream service"
|
||||
else:
|
||||
@@ -1456,7 +1470,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=None,
|
||||
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||
) as client:
|
||||
try:
|
||||
response = await client.send(
|
||||
@@ -1487,6 +1501,9 @@ class BaseUpstreamProvider:
|
||||
status_code=response.status_code,
|
||||
headers=dict(response.headers),
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.ConnectError):
|
||||
# Re-raise to allow fallback handling in proxy layer
|
||||
raise
|
||||
except Exception as exc:
|
||||
tb = traceback.format_exc()
|
||||
logger.error(
|
||||
@@ -2019,7 +2036,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=None,
|
||||
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||
) as client:
|
||||
try:
|
||||
response = await client.send(
|
||||
@@ -2113,6 +2130,9 @@ class BaseUpstreamProvider:
|
||||
headers=dict(response.headers),
|
||||
background=background_tasks,
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.ConnectError):
|
||||
# Re-raise to allow fallback handling in proxy layer
|
||||
raise
|
||||
except Exception as exc:
|
||||
tb = traceback.format_exc()
|
||||
logger.error(
|
||||
@@ -2280,7 +2300,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=None,
|
||||
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||
) as client:
|
||||
try:
|
||||
response = await client.send(
|
||||
@@ -2374,6 +2394,9 @@ class BaseUpstreamProvider:
|
||||
headers=dict(response.headers),
|
||||
background=background_tasks,
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.ConnectError):
|
||||
# Re-raise to allow fallback handling in proxy layer
|
||||
raise
|
||||
except Exception as exc:
|
||||
tb = traceback.format_exc()
|
||||
logger.error(
|
||||
@@ -2503,7 +2526,10 @@ class BaseUpstreamProvider:
|
||||
usage_data = data_json["usage"]
|
||||
model = data_json.get("model")
|
||||
# Track reasoning tokens for Responses API
|
||||
if isinstance(usage_data, dict) and "reasoning_tokens" in usage_data:
|
||||
if (
|
||||
isinstance(usage_data, dict)
|
||||
and "reasoning_tokens" in usage_data
|
||||
):
|
||||
reasoning_tokens = usage_data.get("reasoning_tokens", 0)
|
||||
elif "model" in data_json and not model:
|
||||
model = data_json["model"]
|
||||
|
||||
265
routstr/upstream/fake.py
Normal file
265
routstr/upstream/fake.py
Normal file
@@ -0,0 +1,265 @@
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
from typing import AsyncIterator
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
|
||||
from ..core.db import ApiKey, AsyncSession
|
||||
from ..payment.models import Architecture, Model, Pricing
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
|
||||
class MockUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Fack Mock Upstream provider specifically for Testing."""
|
||||
|
||||
provider_type = "mock"
|
||||
|
||||
async def forward_request(
|
||||
self,
|
||||
request: Request,
|
||||
path: str,
|
||||
headers: dict,
|
||||
request_body: bytes | None,
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
session: AsyncSession,
|
||||
model_obj: Model,
|
||||
) -> Response | StreamingResponse:
|
||||
if path.endswith("chat/completions"):
|
||||
is_streaming = False
|
||||
if request_body:
|
||||
request_data = json.loads(request_body)
|
||||
is_streaming = request_data.get("stream", False)
|
||||
|
||||
if is_streaming:
|
||||
|
||||
async def fake_streaming_response(
|
||||
chunk_size: int | None = None,
|
||||
) -> AsyncIterator[bytes]:
|
||||
suffix = random.randint(1000, 9999)
|
||||
req_id = f"gen-mock-stream-{suffix}"
|
||||
created = 1766138895
|
||||
model = "mock/gpt-420-mock"
|
||||
|
||||
def make_chunk(
|
||||
delta: dict,
|
||||
finish_reason: str | None = None,
|
||||
usage: dict | None = None,
|
||||
) -> bytes:
|
||||
chunk = {
|
||||
"id": req_id,
|
||||
"provider": "MockProvider",
|
||||
"model": model,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": delta,
|
||||
"finish_reason": finish_reason,
|
||||
"native_finish_reason": "completed"
|
||||
if finish_reason
|
||||
else None,
|
||||
"logprobs": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
if usage:
|
||||
chunk["usage"] = usage
|
||||
return f"data: {json.dumps(chunk)}\n\n".encode()
|
||||
|
||||
# 1. Initial chunk
|
||||
yield make_chunk({"role": "assistant", "content": ""})
|
||||
await asyncio.sleep(0.02)
|
||||
|
||||
# 2. Reasoning chunks
|
||||
reasoning_tokens = ["Mock", " reason", "ing", "..."]
|
||||
for token in reasoning_tokens:
|
||||
delta = {
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"reasoning": token,
|
||||
"reasoning_details": [
|
||||
{
|
||||
"type": "reasoning.summary",
|
||||
"summary": token,
|
||||
"format": "openai-responses-v1",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
yield make_chunk(delta)
|
||||
await asyncio.sleep(0.03)
|
||||
|
||||
# 3. Content chunks
|
||||
content_tokens = ["This", " is", " a", " mock", " stream", "."]
|
||||
for token in content_tokens:
|
||||
yield make_chunk({"role": "assistant", "content": token})
|
||||
await asyncio.sleep(0.03)
|
||||
|
||||
# 4. Finish chunk
|
||||
yield make_chunk(
|
||||
{"role": "assistant", "content": ""}, finish_reason="stop"
|
||||
)
|
||||
|
||||
# 5. Usage chunk
|
||||
usage_data = {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 30,
|
||||
"cost": 0.001,
|
||||
"is_byok": False,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"audio_tokens": 0,
|
||||
"video_tokens": 0,
|
||||
},
|
||||
"cost_details": {
|
||||
"upstream_inference_cost": None,
|
||||
"upstream_inference_prompt_cost": 0,
|
||||
"upstream_inference_completions_cost": 0.001,
|
||||
},
|
||||
"completion_tokens_details": {
|
||||
"reasoning_tokens": 10,
|
||||
"image_tokens": 0,
|
||||
},
|
||||
}
|
||||
|
||||
usage_chunk = {
|
||||
"id": req_id,
|
||||
"provider": "MockProvider",
|
||||
"model": model,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant", "content": ""},
|
||||
"finish_reason": None,
|
||||
"native_finish_reason": None,
|
||||
"logprobs": None,
|
||||
}
|
||||
],
|
||||
"usage": usage_data,
|
||||
}
|
||||
yield f"data: {json.dumps(usage_chunk)}\n\n".encode()
|
||||
|
||||
# 6. DONE
|
||||
yield b"data: [DONE]\n\n"
|
||||
|
||||
# 7. Cost
|
||||
cost_chunk = {
|
||||
"cost": {
|
||||
"base_msats": 0,
|
||||
"input_msats": 2,
|
||||
"output_msats": 10,
|
||||
"total_msats": 12,
|
||||
}
|
||||
}
|
||||
yield f"data: {json.dumps(cost_chunk)}\n\n".encode()
|
||||
|
||||
return StreamingResponse(
|
||||
fake_streaming_response(),
|
||||
200,
|
||||
)
|
||||
|
||||
else:
|
||||
suffix = random.randint(1000, 9999)
|
||||
content_dict = {
|
||||
"id": f"gen-mock-{suffix}",
|
||||
"provider": "MockProvider",
|
||||
"model": "mock/gpt-5-mini",
|
||||
"object": "chat.completion",
|
||||
"created": 1766138655,
|
||||
"choices": [
|
||||
{
|
||||
"logprobs": None,
|
||||
"finish_reason": "length",
|
||||
"native_finish_reason": "max_output_tokens",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": f"Mock Content {suffix}",
|
||||
"refusal": None,
|
||||
"reasoning": f"Mock Reasoning {suffix}",
|
||||
"reasoning_details": [
|
||||
{
|
||||
"format": "openai-responses-v1",
|
||||
"index": 0,
|
||||
"type": "reasoning.summary",
|
||||
"summary": f"Mock Summary {suffix}",
|
||||
},
|
||||
{
|
||||
"id": f"rs_mock_{suffix}",
|
||||
"format": "openai-responses-v1",
|
||||
"index": 0,
|
||||
"type": "reasoning.encrypted",
|
||||
"data": "mock_encrypted_data",
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 10,
|
||||
"total_tokens": 20,
|
||||
"cost": 0,
|
||||
"is_byok": False,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"audio_tokens": 0,
|
||||
"video_tokens": 0,
|
||||
},
|
||||
"cost_details": {
|
||||
"upstream_inference_cost": None,
|
||||
"upstream_inference_prompt_cost": 0,
|
||||
"upstream_inference_completions_cost": 0,
|
||||
},
|
||||
"completion_tokens_details": {
|
||||
"reasoning_tokens": 5,
|
||||
"image_tokens": 0,
|
||||
},
|
||||
},
|
||||
"cost": {
|
||||
"base_msats": 0,
|
||||
"input_msats": 0,
|
||||
"output_msats": 0,
|
||||
"total_msats": 0,
|
||||
},
|
||||
}
|
||||
return Response(json.dumps(content_dict).encode(), 200)
|
||||
|
||||
elif path.endswith("embeddings"):
|
||||
raise NotImplementedError
|
||||
elif path.endswith("responses"):
|
||||
raise NotImplementedError
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
return [
|
||||
Model(
|
||||
id="mock/gpt-420-mock",
|
||||
name="mock/gpt-420-mock",
|
||||
created=0,
|
||||
description="mock model for testing",
|
||||
context_length=8192,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(prompt=0.01, completion=0.01),
|
||||
),
|
||||
]
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
return "fake-model"
|
||||
|
||||
async def get_balance(self) -> float | None:
|
||||
return 420.69
|
||||
@@ -218,6 +218,14 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
|
||||
results = await asyncio.gather(*tasks)
|
||||
upstreams = [p for p in results if p is not None]
|
||||
|
||||
if "https://testnut.cashu.space" in settings.cashu_mints:
|
||||
from .fake import MockUpstreamProvider
|
||||
|
||||
mock_provider = MockUpstreamProvider("mock", "mock")
|
||||
await mock_provider.refresh_models_cache()
|
||||
upstreams.append(mock_provider)
|
||||
logger.info("Initialized MockUpstreamProvider for testnut mint")
|
||||
|
||||
return upstreams
|
||||
|
||||
|
||||
|
||||
@@ -313,6 +313,8 @@ async def periodic_payout() -> None:
|
||||
try:
|
||||
async with db.create_session() as session:
|
||||
for mint_url in settings.cashu_mints:
|
||||
if mint_url == "https://testnut.cashu.space":
|
||||
continue
|
||||
for unit in ["sat", "msat"]:
|
||||
wallet = await get_wallet(mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
|
||||
Reference in New Issue
Block a user