fix(providers): harden OAuth model discovery

This commit is contained in:
Xubin Ren
2026-08-29 21:22:20 +08:00
parent 1c6483147e
commit c02f013b17
8 changed files with 188 additions and 73 deletions
+14 -5
View File
@@ -314,7 +314,7 @@ def get_github_copilot_model_catalog(
account_key = _catalog_account_key(getattr(token, "account_id", None))
cache_key = (
f"{storage.get_token_path()}\0{account_key}\0"
f"{_resolve('NANOBOT_COPILOT_BASE_URL', DEFAULT_COPILOT_BASE_URL)}"
f"{_resolve('NANOBOT_COPILOT_BASE_URL', DEFAULT_COPILOT_BASE_URL)}\0{proxy or ''}"
)
return _GITHUB_COPILOT_MODEL_CATALOG.get(cache_key=cache_key, proxy=proxy)
@@ -389,10 +389,7 @@ def _parse_github_copilot_models(payload: Any) -> tuple[ProviderModelSpec, ...]:
or wire_id in seen
or row.get("model_picker_enabled") is not True
or policy.get("state") == "disabled"
or (
isinstance(endpoints, list)
and "/chat/completions" not in cast(list[object], endpoints)
)
or not _copilot_transport_supported(wire_id, endpoints)
):
continue
seen.add(wire_id)
@@ -419,6 +416,18 @@ def _parse_github_copilot_models(payload: Any) -> tuple[ProviderModelSpec, ...]:
return tuple(models)
def _copilot_transport_supported(wire_id: str, endpoints: object) -> bool:
if not isinstance(endpoints, list):
return True
supported = cast(list[object], endpoints)
if "/chat/completions" in supported:
return True
model = wire_id.lower()
return "/responses" in supported and any(
token in model for token in ("gpt-5", "o1", "o3", "o4")
)
def _oauth_fallback_models(provider_name: str) -> tuple[ProviderModelSpec, ...]:
spec = find_by_name(provider_name)
assert spec is not None
+61 -42
View File
@@ -73,49 +73,51 @@ class OAuthModelCatalog:
def get(self, *, cache_key: str, proxy: str | None = None) -> OAuthModelCatalogSnapshot:
"""Return a fresh catalog, sharing concurrent work and retaining a fallback."""
while True:
with self._condition:
with self._condition:
generation = self._generation
cached = self._cached_result(cache_key)
if cached is not None:
return cached
while cache_key in self._inflight:
self._condition.wait()
if generation != self._generation:
return self._stale_or_fallback(None, self._monotonic())
cached = self._cached_result(cache_key)
if cached is not None:
return cached
while cache_key in self._inflight:
self._condition.wait()
cached = self._cached_result(cache_key)
if cached is not None:
return cached
generation = self._generation
self._inflight.add(cache_key)
self._inflight.add(cache_key)
try:
models = tuple(self._fetch(proxy))
if not models:
raise ValueError("provider returned an empty model catalog")
except Exception as exc:
logger.warning("OAuth model catalog refresh failed: type={}", type(exc).__name__)
with self._condition:
invalidated = generation != self._generation
result = self._failure_result(cache_key) if not invalidated else None
else:
now = self._monotonic()
result = OAuthModelCatalogSnapshot(
models=models,
source="remote",
fetched_at=self._wall_clock(),
try:
models = tuple(self._fetch(proxy))
if not models:
raise ValueError("provider returned an empty model catalog")
except Exception as exc:
logger.warning("OAuth model catalog refresh failed: type={}", type(exc).__name__)
with self._condition:
result = (
self._stale_or_fallback(None, self._monotonic())
if generation != self._generation
else self._failure_result(cache_key)
)
with self._condition:
invalidated = generation != self._generation
if not invalidated:
self._store(cache_key, _CacheEntry(snapshot=result, stored_at=now))
self._failures.pop(cache_key, None)
finally:
with self._condition:
self._inflight.discard(cache_key)
self._condition.notify_all()
else:
now = self._monotonic()
result = OAuthModelCatalogSnapshot(
models=models,
source="remote",
fetched_at=self._wall_clock(),
)
with self._condition:
if generation != self._generation:
result = self._stale_or_fallback(None, now)
else:
self._store(cache_key, _CacheEntry(snapshot=result, stored_at=now))
self._failures.pop(cache_key, None)
finally:
with self._condition:
self._inflight.discard(cache_key)
self._condition.notify_all()
if invalidated:
continue
assert result is not None
return result
return result
def invalidate(self) -> None:
"""Drop cached work and prevent an older identity refresh from being stored."""
@@ -123,18 +125,23 @@ class OAuthModelCatalog:
self._generation += 1
self._entries.clear()
self._failures.clear()
self._condition.notify_all()
def _cached_result(self, cache_key: str) -> OAuthModelCatalogSnapshot | None:
now = self._monotonic()
entry = self._entries.get(cache_key)
if entry is not None and now - entry.stored_at < self._fresh_ttl_s:
return replace(entry.snapshot, source="cache")
if self._failures.get(cache_key, 0) > now:
failure_until = self._failures.get(cache_key)
if failure_until is not None and failure_until <= now:
self._failures.pop(cache_key, None)
elif failure_until is not None:
return self._stale_or_fallback(entry, now)
return None
def _failure_result(self, cache_key: str) -> OAuthModelCatalogSnapshot:
now = self._monotonic()
self._reserve(cache_key)
self._failures[cache_key] = now + self._failure_ttl_s
return self._stale_or_fallback(self._entries.get(cache_key), now)
@@ -157,12 +164,24 @@ class OAuthModelCatalog:
)
def _store(self, cache_key: str, entry: _CacheEntry) -> None:
if cache_key not in self._entries and len(self._entries) >= self._max_entries:
oldest = min(self._entries, key=lambda key: self._entries[key].stored_at)
self._entries.pop(oldest, None)
self._failures.pop(oldest, None)
self._reserve(cache_key)
self._entries[cache_key] = entry
def _reserve(self, cache_key: str) -> None:
known = set(self._entries) | set(self._failures)
if cache_key in known or len(known) < self._max_entries:
return
oldest = min(
known,
key=lambda key: (
self._entries[key].stored_at
if key in self._entries
else self._failures[key] - self._failure_ttl_s
),
)
self._entries.pop(oldest, None)
self._failures.pop(oldest, None)
def get_oauth_model_catalog(
provider_name: str,
+14 -1
View File
@@ -153,6 +153,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
backend="openai_compat",
is_direct=True,
),
# === Azure OpenAI (direct API calls with API version 2024-10-21) =====
ProviderSpec(
name="azure_openai",
@@ -315,6 +316,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
detect_by_base_keyword="siliconflow",
default_api_base="https://api.siliconflow.cn/v1",
),
# Novita AI: OpenAI-compatible gateway for hosted model APIs.
ProviderSpec(
name="novita",
@@ -326,6 +328,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
detect_by_base_keyword="novita",
default_api_base="https://api.novita.ai/openai",
),
# VolcEngine (火山引擎): OpenAI-compatible gateway, pay-per-use models
ProviderSpec(
name="volcengine",
@@ -339,6 +342,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
thinking_style="thinking_type",
supports_max_completion_tokens=True,
),
# VolcEngine Coding Plan (火山引擎 Coding Plan): same key as volcengine
ProviderSpec(
name="volcengine_coding_plan",
@@ -352,6 +356,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
thinking_style="thinking_type",
supports_max_completion_tokens=True,
),
# BytePlus: VolcEngine international, pay-per-use models
ProviderSpec(
name="byteplus",
@@ -365,6 +370,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
strip_model_prefix=True,
thinking_style="thinking_type",
),
# BytePlus Coding Plan: same key as byteplus
ProviderSpec(
name="byteplus_coding_plan",
@@ -377,6 +383,8 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
strip_model_prefix=True,
thinking_style="thinking_type",
),
# === Standard providers (matched by model-name keywords) ===============
# Anthropic: native Anthropic SDK
ProviderSpec(
@@ -492,6 +500,11 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
display_name="Github Copilot",
model_catalog="hybrid",
builtin_models=(
ProviderModelSpec(
id="github-copilot/gpt-5.4-mini",
label="GPT-5.4 Mini",
description="GitHub Copilot Responses model.",
),
ProviderModelSpec(
id="github-copilot/gpt-4.1",
label="GPT-4.1",
@@ -768,7 +781,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
env_key="QIANFAN_API_KEY",
display_name="Qianfan",
backend="openai_compat",
default_api_base="https://qianfan.baidubce.com/v2",
default_api_base="https://qianfan.baidubce.com/v2"
),
)
+3 -1
View File
@@ -207,6 +207,7 @@ class XAIGrokProvider(LLMProvider):
)
if on_stream_recover is not None:
await on_stream_recover()
headers = _build_headers(token.access, wire_model)
content, tool_calls, finish_reason, usage, reasoning_content = result
usage = _combine_usage(retry_usage, usage)
@@ -614,10 +615,11 @@ def _xai_error_response(exc: Exception) -> LLMResponse:
)
message = str(exc).strip() or "unexpected error"
retry_after = getattr(exc, "retry_after", None)
usage = getattr(exc, "usage", None)
return LLMResponse(
content=f"Error calling xAI ({type(exc).__name__}): {message}",
finish_reason="error",
usage=getattr(exc, "usage", None),
usage=usage if isinstance(usage, LLMUsage) else None,
retry_after=retry_after,
error_status_code=int(status_code) if status_code is not None else None,
error_kind=error_kind,