mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-22 12:22:20 +00:00
Compare commits
1 Commits
feat/vulne
...
fix-model-
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4b33d8304c |
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user