mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-03 08:46:16 +00:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1df66f48b8 | ||
|
|
a9c5458660 |
+41
-5
@@ -1,5 +1,7 @@
|
|||||||
"""Model prioritization algorithm for selecting cheapest upstream providers."""
|
"""Model prioritization algorithm for selecting cheapest upstream providers."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from .core.logging import get_logger
|
from .core.logging import get_logger
|
||||||
@@ -157,15 +159,21 @@ def create_model_mappings(
|
|||||||
upstreams: list["BaseUpstreamProvider"],
|
upstreams: list["BaseUpstreamProvider"],
|
||||||
overrides_by_id: dict[str, tuple],
|
overrides_by_id: dict[str, tuple],
|
||||||
disabled_model_ids: set[str],
|
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.
|
"""Create optimal model mappings based on cost and provider preferences.
|
||||||
|
|
||||||
This is the main entry point for the algorithm. It processes all upstream providers
|
This is the main entry point for the algorithm. It processes all upstream providers
|
||||||
and creates three mappings based on cost optimization:
|
and creates three mappings based on cost optimization:
|
||||||
|
|
||||||
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
|
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)
|
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:
|
The algorithm:
|
||||||
- Processes non-OpenRouter providers first (they're typically cheaper)
|
- 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
|
disabled_model_ids: Set of model IDs that should be excluded
|
||||||
|
|
||||||
Returns:
|
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 .payment.models import _row_to_model
|
||||||
from .upstream.helpers import resolve_model_alias
|
from .upstream.helpers import resolve_model_alias
|
||||||
@@ -186,9 +194,10 @@ def create_model_mappings(
|
|||||||
model_instances: dict[str, "Model"] = {}
|
model_instances: dict[str, "Model"] = {}
|
||||||
provider_map: dict[str, "BaseUpstreamProvider"] = {}
|
provider_map: dict[str, "BaseUpstreamProvider"] = {}
|
||||||
unique_models: dict[str, "Model"] = {}
|
unique_models: dict[str, "Model"] = {}
|
||||||
|
provider_candidates_map: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||||
|
|
||||||
# Separate OpenRouter from other providers
|
# Separate OpenRouter from other providers
|
||||||
openrouter: "BaseUpstreamProvider" | None = None
|
openrouter: BaseUpstreamProvider | None = None
|
||||||
other_upstreams: list["BaseUpstreamProvider"] = []
|
other_upstreams: list["BaseUpstreamProvider"] = []
|
||||||
|
|
||||||
for upstream in upstreams:
|
for upstream in upstreams:
|
||||||
@@ -207,6 +216,13 @@ def create_model_mappings(
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""Set alias to model/provider if not set or if new model is preferred."""
|
"""Set alias to model/provider if not set or if new model is preferred."""
|
||||||
alias_lower = alias.lower()
|
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)
|
existing_model = model_instances.get(alias_lower)
|
||||||
if not existing_model:
|
if not existing_model:
|
||||||
# No existing mapping, set it
|
# No existing mapping, set it
|
||||||
@@ -276,6 +292,26 @@ def create_model_mappings(
|
|||||||
if openrouter:
|
if openrouter:
|
||||||
process_provider_models(openrouter, is_openrouter=True)
|
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
|
# Log provider distribution
|
||||||
provider_counts: dict[str, int] = {}
|
provider_counts: dict[str, int] = {}
|
||||||
for provider in provider_map.values():
|
for provider in provider_map.values():
|
||||||
@@ -287,4 +323,4 @@ def create_model_mappings(
|
|||||||
extra={"provider_distribution": provider_counts},
|
extra={"provider_distribution": provider_counts},
|
||||||
)
|
)
|
||||||
|
|
||||||
return model_instances, provider_map, unique_models
|
return model_instances, provider_map, unique_models, provider_candidates_map
|
||||||
|
|||||||
+193
-60
@@ -1,6 +1,7 @@
|
|||||||
import json
|
import json
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||||
from fastapi.responses import Response, StreamingResponse
|
from fastapi.responses import Response, StreamingResponse
|
||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
@@ -26,7 +27,7 @@ from .payment.helpers import (
|
|||||||
from .payment.models import Model
|
from .payment.models import Model
|
||||||
from .upstream import BaseUpstreamProvider
|
from .upstream import BaseUpstreamProvider
|
||||||
from .upstream.helpers import init_upstreams
|
from .upstream.helpers import init_upstreams
|
||||||
from .wallet import deserialize_token_from_string
|
from .wallet import deserialize_token_from_string, recieve_token
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
proxy_router = APIRouter()
|
proxy_router = APIRouter()
|
||||||
@@ -34,6 +35,9 @@ proxy_router = APIRouter()
|
|||||||
_upstreams: list[BaseUpstreamProvider] = []
|
_upstreams: list[BaseUpstreamProvider] = []
|
||||||
_model_instances: dict[str, Model] = {} # All aliases -> Model
|
_model_instances: dict[str, Model] = {} # All aliases -> Model
|
||||||
_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider
|
_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)
|
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
|
||||||
|
|
||||||
|
|
||||||
@@ -75,6 +79,21 @@ def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
|
|||||||
return _provider_map.get(model_id.lower())
|
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]:
|
def get_unique_models() -> list[Model]:
|
||||||
"""Get list of unique models (no duplicates from aliases)."""
|
"""Get list of unique models (no duplicates from aliases)."""
|
||||||
return list(_unique_models.values())
|
return list(_unique_models.values())
|
||||||
@@ -84,7 +103,7 @@ async def refresh_model_maps() -> None:
|
|||||||
"""Refresh global model and provider maps using the cost-based algorithm."""
|
"""Refresh global model and provider maps using the cost-based algorithm."""
|
||||||
from sqlalchemy.orm import selectinload
|
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:
|
async with create_session() as session:
|
||||||
# Fetch all providers with their models in a single logical operation
|
# Fetch all providers with their models in a single logical operation
|
||||||
@@ -104,10 +123,12 @@ async def refresh_model_maps() -> None:
|
|||||||
else:
|
else:
|
||||||
disabled_model_ids.add(model.id)
|
disabled_model_ids.add(model.id)
|
||||||
|
|
||||||
_model_instances, _provider_map, _unique_models = create_model_mappings(
|
_model_instances, _provider_map, _unique_models, _provider_candidates_map = (
|
||||||
upstreams=_upstreams,
|
create_model_mappings(
|
||||||
overrides_by_id=overrides_by_id,
|
upstreams=_upstreams,
|
||||||
disabled_model_ids=disabled_model_ids,
|
overrides_by_id=overrides_by_id,
|
||||||
|
disabled_model_ids=disabled_model_ids,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -172,8 +193,8 @@ async def proxy(
|
|||||||
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
||||||
)
|
)
|
||||||
|
|
||||||
upstream = get_provider_for_model(model_id)
|
upstreams = get_providers_for_model(model_id)
|
||||||
if not upstream:
|
if not upstreams:
|
||||||
return create_error_response(
|
return create_error_response(
|
||||||
"invalid_model",
|
"invalid_model",
|
||||||
f"No provider found for model '{model_id}'",
|
f"No provider found for model '{model_id}'",
|
||||||
@@ -190,14 +211,80 @@ async def proxy(
|
|||||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||||
|
|
||||||
if x_cashu := headers.get("x-cashu", None):
|
if x_cashu := headers.get("x-cashu", None):
|
||||||
if is_responses_api:
|
# Redeem token once before trying any providers
|
||||||
return await upstream.handle_x_cashu_responses(
|
amount, unit, mint = await recieve_token(x_cashu)
|
||||||
request, x_cashu, path, max_cost_for_model, model_obj
|
|
||||||
)
|
# Fallback for X-Cashu payments
|
||||||
else:
|
last_exception = None
|
||||||
return await upstream.handle_x_cashu(
|
for i, upstream in enumerate(upstreams):
|
||||||
request, x_cashu, path, max_cost_for_model, model_obj
|
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):
|
elif auth := headers.get("authorization", None):
|
||||||
key = await get_bearer_token_key(headers, path, session, auth)
|
key = await get_bearer_token_key(headers, path, session, auth)
|
||||||
@@ -212,56 +299,102 @@ async def proxy(
|
|||||||
)
|
)
|
||||||
|
|
||||||
logger.debug("Processing unauthenticated GET request", extra={"path": path})
|
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:
|
if request_body_dict:
|
||||||
await pay_for_request(key, max_cost_for_model, session)
|
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:
|
if response.status_code != 200:
|
||||||
response = await upstream.forward_responses_request(
|
# If it's a 429 (rate limit) or 503 (service unavailable), we might also want to fallback
|
||||||
request,
|
if response.status_code in (429, 503, 502) and i < len(upstreams) - 1:
|
||||||
path,
|
logger.warning(
|
||||||
headers,
|
f"Upstream provider {i + 1}/{len(upstreams)} returned {response.status_code}, trying fallback",
|
||||||
request_body,
|
extra={"model": model_id, "status": response.status_code},
|
||||||
key,
|
)
|
||||||
max_cost_for_model,
|
continue
|
||||||
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:
|
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
logger.warning(
|
||||||
logger.warning(
|
"Upstream request failed, revert payment",
|
||||||
"Upstream request failed, revert payment",
|
extra={
|
||||||
extra={
|
"status_code": response.status_code,
|
||||||
"status_code": response.status_code,
|
"path": path,
|
||||||
"path": path,
|
"key_hash": key.hashed_key[:8] + "...",
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_balance": key.balance,
|
||||||
"key_balance": key.balance,
|
"max_cost_for_model": max_cost_for_model,
|
||||||
"max_cost_for_model": max_cost_for_model,
|
"upstream_headers": response.headers
|
||||||
"upstream_headers": response.headers
|
if hasattr(response, "headers")
|
||||||
if hasattr(response, "headers")
|
else None,
|
||||||
else None,
|
},
|
||||||
},
|
)
|
||||||
)
|
return response
|
||||||
# Return the mapped error response generated earlier rather than masking with 502
|
|
||||||
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(
|
async def get_bearer_token_key(
|
||||||
@@ -357,7 +490,7 @@ def extract_model_from_responses_request(request_body_dict: dict[str, Any]) -> s
|
|||||||
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"No model found in Responses API request",
|
"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"
|
return "unknown"
|
||||||
|
|
||||||
|
|||||||
+39
-13
@@ -37,6 +37,8 @@ from ..wallet import recieve_token, send_token
|
|||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
DEFAULT_PROXY_TIMEOUT = 30.0
|
||||||
|
|
||||||
|
|
||||||
class TopupData(BaseModel):
|
class TopupData(BaseModel):
|
||||||
"""Universal top-up data schema for Lightning Network invoices."""
|
"""Universal top-up data schema for Lightning Network invoices."""
|
||||||
@@ -234,7 +236,11 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Handle model in input field (alternative format)
|
# 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
|
original_model = model_obj.id
|
||||||
transformed_model = self.transform_model_name(original_model)
|
transformed_model = self.transform_model_name(original_model)
|
||||||
data["input"]["model"] = transformed_model
|
data["input"]["model"] = transformed_model
|
||||||
@@ -686,6 +692,7 @@ class BaseUpstreamProvider:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
"Error processing non-streaming chat completion",
|
"Error processing non-streaming chat completion",
|
||||||
@@ -779,8 +786,13 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
# Track reasoning tokens for Responses API
|
# Track reasoning tokens for Responses API
|
||||||
if usage := obj.get("usage", {}):
|
if usage := obj.get("usage", {}):
|
||||||
if isinstance(usage, dict) and "reasoning_tokens" in usage:
|
if (
|
||||||
reasoning_tokens += usage.get("reasoning_tokens", 0)
|
isinstance(usage, dict)
|
||||||
|
and "reasoning_tokens" in usage
|
||||||
|
):
|
||||||
|
reasoning_tokens += usage.get(
|
||||||
|
"reasoning_tokens", 0
|
||||||
|
)
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
pass
|
pass
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -933,8 +945,8 @@ class BaseUpstreamProvider:
|
|||||||
"model": response_json.get("model", "unknown"),
|
"model": response_json.get("model", "unknown"),
|
||||||
"has_usage": "usage" in response_json,
|
"has_usage": "usage" in response_json,
|
||||||
"has_reasoning_tokens": "usage" in response_json
|
"has_reasoning_tokens": "usage" in response_json
|
||||||
and isinstance(response_json.get("usage"), dict)
|
and isinstance(response_json.get("usage"), dict)
|
||||||
and "reasoning_tokens" in response_json["usage"],
|
and "reasoning_tokens" in response_json["usage"],
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1047,7 +1059,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
client = httpx.AsyncClient(
|
client = httpx.AsyncClient(
|
||||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||||
timeout=None,
|
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1190,7 +1202,8 @@ class BaseUpstreamProvider:
|
|||||||
if isinstance(exc, httpx.ConnectError):
|
if isinstance(exc, httpx.ConnectError):
|
||||||
error_message = "Unable to connect to upstream service"
|
error_message = "Unable to connect to upstream service"
|
||||||
elif isinstance(exc, httpx.TimeoutException):
|
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):
|
elif isinstance(exc, httpx.NetworkError):
|
||||||
error_message = "Network error while connecting to upstream service"
|
error_message = "Network error while connecting to upstream service"
|
||||||
else:
|
else:
|
||||||
@@ -1273,7 +1286,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
client = httpx.AsyncClient(
|
client = httpx.AsyncClient(
|
||||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||||
timeout=None,
|
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1393,7 +1406,8 @@ class BaseUpstreamProvider:
|
|||||||
if isinstance(exc, httpx.ConnectError):
|
if isinstance(exc, httpx.ConnectError):
|
||||||
error_message = "Unable to connect to upstream service"
|
error_message = "Unable to connect to upstream service"
|
||||||
elif isinstance(exc, httpx.TimeoutException):
|
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):
|
elif isinstance(exc, httpx.NetworkError):
|
||||||
error_message = "Network error while connecting to upstream service"
|
error_message = "Network error while connecting to upstream service"
|
||||||
else:
|
else:
|
||||||
@@ -1456,7 +1470,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||||
timeout=None,
|
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||||
) as client:
|
) as client:
|
||||||
try:
|
try:
|
||||||
response = await client.send(
|
response = await client.send(
|
||||||
@@ -1487,6 +1501,9 @@ class BaseUpstreamProvider:
|
|||||||
status_code=response.status_code,
|
status_code=response.status_code,
|
||||||
headers=dict(response.headers),
|
headers=dict(response.headers),
|
||||||
)
|
)
|
||||||
|
except (httpx.TimeoutException, httpx.ConnectError):
|
||||||
|
# Re-raise to allow fallback handling in proxy layer
|
||||||
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
tb = traceback.format_exc()
|
tb = traceback.format_exc()
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -2019,7 +2036,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||||
timeout=None,
|
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||||
) as client:
|
) as client:
|
||||||
try:
|
try:
|
||||||
response = await client.send(
|
response = await client.send(
|
||||||
@@ -2113,6 +2130,9 @@ class BaseUpstreamProvider:
|
|||||||
headers=dict(response.headers),
|
headers=dict(response.headers),
|
||||||
background=background_tasks,
|
background=background_tasks,
|
||||||
)
|
)
|
||||||
|
except (httpx.TimeoutException, httpx.ConnectError):
|
||||||
|
# Re-raise to allow fallback handling in proxy layer
|
||||||
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
tb = traceback.format_exc()
|
tb = traceback.format_exc()
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -2280,7 +2300,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||||
timeout=None,
|
timeout=DEFAULT_PROXY_TIMEOUT,
|
||||||
) as client:
|
) as client:
|
||||||
try:
|
try:
|
||||||
response = await client.send(
|
response = await client.send(
|
||||||
@@ -2374,6 +2394,9 @@ class BaseUpstreamProvider:
|
|||||||
headers=dict(response.headers),
|
headers=dict(response.headers),
|
||||||
background=background_tasks,
|
background=background_tasks,
|
||||||
)
|
)
|
||||||
|
except (httpx.TimeoutException, httpx.ConnectError):
|
||||||
|
# Re-raise to allow fallback handling in proxy layer
|
||||||
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
tb = traceback.format_exc()
|
tb = traceback.format_exc()
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -2503,7 +2526,10 @@ class BaseUpstreamProvider:
|
|||||||
usage_data = data_json["usage"]
|
usage_data = data_json["usage"]
|
||||||
model = data_json.get("model")
|
model = data_json.get("model")
|
||||||
# Track reasoning tokens for Responses API
|
# 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)
|
reasoning_tokens = usage_data.get("reasoning_tokens", 0)
|
||||||
elif "model" in data_json and not model:
|
elif "model" in data_json and not model:
|
||||||
model = data_json["model"]
|
model = data_json["model"]
|
||||||
|
|||||||
Reference in New Issue
Block a user