Compare commits

...

2 Commits

Author SHA1 Message Date
9qeklajc
4b33d8304c fix forwarding correct model ids 2026-03-13 22:42:41 +01:00
9qeklajc
9006709f8d Merge pull request #404 from Routstr/add-missing-refund-button
add missing refund button
2026-03-13 22:39:39 +01:00
3 changed files with 24 additions and 16 deletions

View File

@@ -89,7 +89,7 @@ def create_model_mappings(
overrides_by_id: dict[str, tuple],
disabled_model_ids: set[str],
) -> tuple[
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
dict[str, "Model"], dict[str, list[tuple["BaseUpstreamProvider", "Model"]]], dict[str, "Model"]
]:
"""Create optimal model mappings based on cost and provider preferences.
@@ -298,7 +298,7 @@ def create_model_mappings(
# Sort candidates and build final maps
model_instances: dict[str, "Model"] = {}
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
provider_map: dict[str, list[tuple["BaseUpstreamProvider", "Model"]]] = {}
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is."""
@@ -326,13 +326,13 @@ def create_model_mappings(
best_model, best_provider = items[0]
model_instances[alias] = best_model
provider_map[alias] = [p for _, p in items]
provider_map[alias] = [(p, m) for m, p in items]
# Log provider distribution (using top provider for stats)
provider_counts: dict[str, int] = {}
for providers in provider_map.values():
if providers:
provider = providers[0]
provider, _model = providers[0]
provider_name = getattr(provider, "upstream_name", "unknown")
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1

View File

@@ -34,8 +34,8 @@ proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[
str, list[BaseUpstreamProvider]
] = {} # All aliases -> List[Provider]
str, list[tuple[BaseUpstreamProvider, Model]]
] = {} # All aliases -> List[(Provider, Model)]
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
@@ -72,8 +72,10 @@ def get_model_instance(model_id: str) -> Model | None:
return _model_instances.get(model_id.lower())
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
"""Get UpstreamProvider list for model ID from global cache."""
def get_provider_for_model(
model_id: str,
) -> list[tuple[BaseUpstreamProvider, Model]] | None:
"""Get UpstreamProvider list (with their Model) for model ID from global cache."""
return _provider_map.get(model_id.lower())
@@ -180,15 +182,15 @@ async def proxy(
if x_cashu := headers.get("x-cashu", None):
last_error = None
for i, upstream in enumerate(upstreams):
for i, (upstream, upstream_model) in enumerate(upstreams):
try:
if is_responses_api:
return await upstream.handle_x_cashu_responses(
request, x_cashu, path, max_cost_for_model, model_obj
request, x_cashu, path, max_cost_for_model, upstream_model
)
else:
return await upstream.handle_x_cashu(
request, x_cashu, path, max_cost_for_model, model_obj
request, x_cashu, path, max_cost_for_model, upstream_model
)
except UpstreamError as e:
logger.warning(
@@ -222,7 +224,7 @@ async def proxy(
logger.debug("Processing unauthenticated GET request", extra={"path": path})
last_error_response = None
for i, upstream in enumerate(upstreams):
for i, (upstream, _upstream_model) in enumerate(upstreams):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
@@ -269,7 +271,11 @@ async def proxy(
if request_body_dict:
await pay_for_request(key, max_cost_for_model, session)
for i, upstream in enumerate(upstreams):
for i, (upstream, upstream_model) in enumerate(upstreams):
logger.info(f"Selected upstream provider: {upstream.provider_type}")
logger.info(
f"Forwarding request to {upstream.provider_type} for model {upstream_model.id}"
)
headers = upstream.prepare_headers(dict(request.headers))
try:
@@ -283,7 +289,7 @@ async def proxy(
key,
max_cost_for_model,
session,
model_obj,
upstream_model,
)
else:
response = await upstream.forward_request(
@@ -294,7 +300,7 @@ async def proxy(
key,
max_cost_for_model,
session,
model_obj,
upstream_model,
)
except UpstreamError:
# Let the outer UpstreamError handler manage retry/revert

View File

@@ -63,7 +63,9 @@ def resolve_model_alias(
aliases.append(canonical_base)
if alias_ids:
aliases.extend(alias_ids)
for aid in alias_ids:
if aid not in aliases:
aliases.append(aid)
return aliases