refactor: Antigravity/Codex 服务重构为插件化适配器架构

- 将 Antigravity 和 Codex 从独立模块迁移至 src/services/provider/adapters/ 插件体系
- 新增 provider_types 和 oauth_token 模块,移除 maintenance_scheduler 中的 OAuth 定时刷新
- 增强 admin API:扩展 keys 和 provider_query 端点,新增 dashboard 路由
- 大幅增强 ProviderDetailDrawer 组件,新增 AntigravityQuotaDialog
- 改进 handler 基类(chat/cli)和错误分类器
- 优化 fetch_scheduler 和 upstream_fetcher
- 前端 UI 组件清理和优化
- 更新测试以匹配新模块结构
This commit is contained in:
fawney19
2026-02-06 16:37:06 +08:00
parent e01dfee41d
commit b88fb6273b
112 changed files with 5193 additions and 1992 deletions

View File

@@ -643,6 +643,10 @@ class ChatAdapterBase(ApiAdapter):
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
# Provider 上下文Chat 适配器忽略这些参数,仅保持签名兼容)
auth_type: str | None = None, # noqa: ARG003
provider_type: str | None = None, # noqa: ARG003
decrypted_auth_config: dict[str, Any] | None = None, # noqa: ARG003
) -> dict[str, Any]:
"""
测试模型连接性(非流式)

View File

@@ -920,13 +920,48 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
# OAuth token may be revoked/expired earlier than expires_at indicates.
# Best-effort: force refresh once on 401 and retry a single time.
if (
resp.status_code == 401
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
):
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
if refreshed_auth:
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
ctx.provider_request_headers = provider_headers
# retry once
resp = await http_client.post(
url,
json=provider_payload,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout_sync),
)
ctx.status_code = resp.status_code
ctx.response_headers = dict(resp.headers)
if envelope:
envelope.on_http_status(
base_url=ctx.selected_base_url, status_code=ctx.status_code
)
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e2:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
e2.upstream_response = error_body # type: ignore[attr-defined]
raise
else:
error_body = ""
e.upstream_response = error_body # type: ignore[attr-defined]
raise
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
e.upstream_response = error_body # type: ignore[attr-defined]
raise
# Safe JSON parsing.
try:
@@ -1087,86 +1122,115 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
max_prefetch_lines=config.stream_prefetch_lines,
)
try:
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
if is_disconnected is not None:
await wait_for_with_disconnect_detection(
_connect_and_prefetch(),
timeout=request_timeout,
is_disconnected=is_disconnected,
request_id=self.request_id,
)
else:
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
except ClientDisconnectedException:
# 客户端断开连接,清理资源
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
ctx.status_code = 499
ctx.error_message = "client_disconnected_during_prefetch"
raise
except TimeoutError:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
logger.warning(
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
)
raise ProviderTimeoutException(
provider_name=str(provider.name),
timeout=int(request_timeout),
)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
# 连接/读写超时:清理可能已建立的连接上下文
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
if ctx.selected_base_url:
logger.warning(
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
)
raise
except httpx.HTTPStatusError as e:
error_text = await self._extract_error_text(e)
logger.error(f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}")
await http_client.aclose()
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
e.upstream_response = error_text # type: ignore[attr-defined]
raise
except EmbeddedErrorException:
for attempt in range(2):
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
raise
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
if is_disconnected is not None:
await wait_for_with_disconnect_detection(
_connect_and_prefetch(),
timeout=request_timeout,
is_disconnected=is_disconnected,
request_id=self.request_id,
)
else:
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
break
except Exception:
await http_client.aclose()
raise
except ClientDisconnectedException:
# 客户端断开连接,清理资源
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
ctx.status_code = 499
ctx.error_message = "client_disconnected_during_prefetch"
raise
except TimeoutError:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
logger.warning(
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
)
raise ProviderTimeoutException(
provider_name=str(provider.name),
timeout=int(request_timeout),
)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
# 连接/读写超时:清理可能已建立的连接上下文
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
if ctx.selected_base_url:
logger.warning(
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
)
raise
except httpx.HTTPStatusError as e:
status = int(getattr(e.response, "status_code", 0) or 0)
if (
attempt == 0
and status == 401
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
):
# OAuth token may be revoked/expired earlier than expires_at indicates.
# Best-effort: force refresh once on 401 and retry a single time.
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
if refreshed_auth:
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
ctx.provider_request_headers = provider_headers
# Reset state for the next attempt.
byte_iterator = None
prefetched_chunks = None
response_ctx = None
continue
error_text = await self._extract_error_text(e)
logger.error(
f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}"
)
await http_client.aclose()
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
e.upstream_response = error_text # type: ignore[attr-defined]
raise
except EmbeddedErrorException:
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
raise
except Exception:
await http_client.aclose()
raise
# 类型断言:成功执行后这些变量不会为 None
assert byte_iterator is not None

View File

@@ -609,6 +609,10 @@ class CliAdapterBase(ApiAdapter):
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
# Provider 上下文(用于 OAuth 认证和 Antigravity 等特殊路由)
auth_type: str | None = None,
provider_type: str | None = None,
decrypted_auth_config: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
@@ -634,6 +638,10 @@ class CliAdapterBase(ApiAdapter):
provider_id: 提供商ID
api_key_id: API密钥ID
model_name: 模型名称
auth_type: Key 认证类型("api_key"/"oauth"/"vertex_ai"
OAuth 类型自动使用 Authorization: Bearer 替代端点默认认证头
provider_type: 提供商类型(用于 Antigravity v1internal 等特殊路由)
decrypted_auth_config: 解密后的 OAuth 配置Antigravity 需要 project_id
Returns:
测试响应数据
@@ -641,37 +649,84 @@ class CliAdapterBase(ApiAdapter):
from src.api.handlers.base.endpoint_checker import run_endpoint_check
from src.api.handlers.base.request_builder import apply_body_rules
from src.core.api_format.headers import HeaderBuilder
from src.core.provider_types import ProviderType
# 构建请求组件
url = cls.build_endpoint_url(base_url, request_data, model_name)
is_antigravity = provider_type == ProviderType.ANTIGRAVITY
is_oauth = auth_type == "oauth"
# 合并 CLI 额外头部到 extra_headers
# ---- URL ----
if is_antigravity:
# Antigravity 走 v1internal 端点,模型名在请求体 envelope 中,不在 URL 路径里
from src.services.provider.adapters.antigravity.constants import (
V1INTERNAL_PATH_TEMPLATE,
)
from src.services.provider.adapters.antigravity.constants import (
get_http_user_agent as _get_antigravity_ua,
)
from src.services.provider.adapters.antigravity.envelope import wrap_v1internal_request
from src.services.provider.adapters.antigravity.url_availability import url_availability
ordered_urls = url_availability.get_ordered_urls(prefer_daily=True)
effective_base_url = ordered_urls[0] if ordered_urls else base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(effective_base_url).rstrip('/')}{path}"
else:
url = cls.build_endpoint_url(base_url, request_data, model_name)
# ---- Headers ----
cli_extra = cls.get_cli_extra_headers(base_url=base_url)
merged_extra = dict(extra_headers) if extra_headers else {}
merged_extra.update(cli_extra)
# 使用统一的头部构建函数
# Antigravity 需要特定的 User-Agent
if is_antigravity:
merged_extra["User-Agent"] = _get_antigravity_ua()
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
# OAuth 统一处理:替换端点默认认证头为 Authorization: Bearer
# (与 get_provider_auth 返回的 ProviderAuthInfo 行为一致)
if is_oauth:
from src.core.api_format import get_auth_config_for_endpoint
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
if default_auth_header.lower() != "authorization":
headers.pop(default_auth_header, None)
headers["Authorization"] = f"Bearer {api_key}"
# ---- Body ----
body = cls.build_request_body(request_data, base_url=base_url)
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
if body_rules:
body = apply_body_rules(body, body_rules)
# 应用请求头规则(在请求头构建后应用)
if header_rules:
# 获取认证头名称,防止被规则覆盖
from src.core.api_format import get_auth_config_for_endpoint
# Antigravity用 v1internal envelope 包装请求体
if is_antigravity:
project_id = (decrypted_auth_config or {}).get("project_id", "")
effective_model = model_name or request_data.get("model", "")
body = wrap_v1internal_request(
body,
project_id=project_id,
model=effective_model,
)
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
protected_keys = {auth_header.lower(), "content-type"}
# ---- Header Rules ----
if header_rules:
from src.core.api_format import get_auth_config_for_endpoint as _get_auth_cfg
# 保护实际使用的认证头,而非端点默认的
if is_oauth:
protected_keys = {"authorization", "content-type"}
else:
ep_auth_header, _ = _get_auth_cfg(cls.FORMAT_ID)
protected_keys = {ep_auth_header.lower(), "content-type"}
header_builder = HeaderBuilder()
header_builder.add_many(headers)
header_builder.apply_rules(header_rules, protected_keys)
headers = header_builder.build()
# 获取有效的模型名称
# ---- Execute ----
effective_model_name = model_name or request_data.get("model")
return await run_endpoint_check(
@@ -680,7 +735,6 @@ class CliAdapterBase(ApiAdapter):
headers=headers,
json_body=body,
api_format=cls.FORMAT_ID,
# 用量计算参数(现在强制记录)
db=db,
user=user,
provider_name=provider_name,

View File

@@ -892,13 +892,48 @@ class CliMessageHandlerBase(BaseMessageHandler):
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
# OAuth token may be revoked/expired earlier than expires_at indicates.
# Best-effort: force refresh once on 401 and retry a single time.
if (
resp.status_code == 401
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
):
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
if refreshed_auth:
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
ctx.provider_request_headers = provider_headers
# retry once
resp = await http_client.post(
url,
json=provider_payload,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout_sync),
)
ctx.status_code = resp.status_code
ctx.response_headers = dict(resp.headers)
if envelope:
envelope.on_http_status(
base_url=ctx.selected_base_url, status_code=ctx.status_code
)
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e2:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
e2.upstream_response = error_body # type: ignore[attr-defined]
raise
else:
error_body = ""
e.upstream_response = error_body # type: ignore[attr-defined]
raise
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
e.upstream_response = error_body # type: ignore[attr-defined]
raise
# Safe JSON parsing.
try:
@@ -1055,85 +1090,112 @@ class CliMessageHandlerBase(BaseMessageHandler):
byte_iterator, provider, endpoint, ctx
)
try:
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
if http_request is not None:
await wait_for_with_disconnect_detection(
_connect_and_prefetch(),
timeout=request_timeout,
is_disconnected=http_request.is_disconnected,
request_id=self.request_id,
)
else:
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
except TimeoutError as e:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
await http_client.aclose()
logger.warning(
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
)
raise ProviderTimeoutException(
provider_name=str(provider.name),
timeout=int(request_timeout),
)
except ClientDisconnectedException:
# 客户端断开连接,清理资源
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
ctx.status_code = 499
ctx.error_message = "client_disconnected_during_prefetch"
raise
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
if ctx.selected_base_url:
logger.warning(
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
)
await http_client.aclose()
raise
except httpx.HTTPStatusError as e:
error_text = await self._extract_error_text(e)
logger.error(
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
)
await http_client.aclose()
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
e.upstream_response = error_text # type: ignore[attr-defined]
raise
except EmbeddedErrorException:
# 嵌套错误需要触发重试,关闭连接后重新抛出
for attempt in range(2):
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
raise
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
if http_request is not None:
await wait_for_with_disconnect_detection(
_connect_and_prefetch(),
timeout=request_timeout,
is_disconnected=http_request.is_disconnected,
request_id=self.request_id,
)
else:
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
break
except Exception:
await http_client.aclose()
raise
except TimeoutError as e:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
await http_client.aclose()
logger.warning(
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
)
raise ProviderTimeoutException(
provider_name=str(provider.name),
timeout=int(request_timeout),
)
except ClientDisconnectedException:
# 客户端断开连接,清理资源
if response_ctx is not None:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
ctx.status_code = 499
ctx.error_message = "client_disconnected_during_prefetch"
raise
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
if ctx.selected_base_url:
logger.warning(
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
)
await http_client.aclose()
raise
except httpx.HTTPStatusError as e:
status = int(getattr(e.response, "status_code", 0) or 0)
if (
attempt == 0
and status == 401
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
):
# OAuth token may be revoked/expired earlier than expires_at indicates.
# Best-effort: force refresh once on 401 and retry a single time.
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
if refreshed_auth:
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
ctx.provider_request_headers = provider_headers
# Reset state for the next attempt.
byte_iterator = None
prefetched_chunks = None
response_ctx = None
continue
error_text = await self._extract_error_text(e)
logger.error(
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
)
await http_client.aclose()
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
e.upstream_response = error_text # type: ignore[attr-defined]
raise
except EmbeddedErrorException:
# 嵌套错误需要触发重试,关闭连接后重新抛出
try:
if response_ctx is not None:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass
await http_client.aclose()
raise
except Exception:
await http_client.aclose()
raise
# 类型断言:成功执行后这些变量不会为 None
assert byte_iterator is not None

View File

@@ -599,6 +599,8 @@ def build_passthrough_request(
async def get_provider_auth(
endpoint: "ProviderEndpoint",
key: "ProviderAPIKey",
*,
force_refresh: bool = False,
) -> ProviderAuthInfo | None:
"""
获取 Provider 的认证信息
@@ -638,7 +640,7 @@ async def get_provider_auth(
refresh_token = token_meta.get("refresh_token")
provider_type = str(token_meta.get("provider_type") or "")
# 120s skew
# 120s skew (or force refresh when upstream returns 401)
should_refresh = False
try:
if expires_at is not None:
@@ -646,6 +648,9 @@ async def get_provider_auth(
except Exception:
should_refresh = False
if force_refresh:
should_refresh = True
if should_refresh and refresh_token and provider_type:
try:
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS

View File

@@ -578,6 +578,15 @@ class StreamProcessor:
if not isinstance(data_obj, dict):
return []
# Provider envelope: unwrap v1internal wrapper before conversion
# (e.g. Antigravity {"response": {...}, "traceId": "..."} → inner response)
if envelope and isinstance(data_obj, dict):
data_obj = envelope.unwrap_response(data_obj)
envelope.postprocess_unwrapped_response(
model=str(ctx.model or ""),
data=data_obj,
)
try:
converted_events = registry.convert_stream_chunk(
data_obj,

View File

@@ -160,7 +160,11 @@ class ClaudeChatAdapter(ChatAdapterBase):
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表"""
"""查询 Claude API 支持的模型列表
Anthropic 的 /v1/models 是分页接口has_more/first_id/last_id
默认只返回一页。这里做 best-effort 的全量拉取,确保管理端能展示完整模型列表。
"""
headers = cls.build_headers_with_extra(api_key, extra_headers)
# 构建 /v1/models URL
@@ -171,24 +175,60 @@ class ClaudeChatAdapter(ChatAdapterBase):
models_url = f"{base_url}/v1/models"
try:
response = await client.get(models_url, headers=headers)
logger.debug(f"Claude models request to {models_url}: status={response.status_code}")
if response.status_code == 200:
all_models: list[dict] = []
seen_ids: set[str] = set()
after_id: str | None = None
limit = 100 # Anthropic 支持 limit尽量减少分页次数
max_pages = 20 # safety guard
for _ in range(max_pages):
params: dict[str, Any] = {"limit": limit}
if after_id:
params["after_id"] = after_id
response = await client.get(models_url, headers=headers, params=params)
logger.debug(
f"Claude models request to {models_url}: status={response.status_code}, after_id={after_id}"
)
if response.status_code != 200:
error_body = response.text[:500] if response.text else "(empty)"
error_msg = f"HTTP {response.status_code}: {error_body}"
logger.warning(f"Claude models request to {models_url} failed: {error_msg}")
return [], error_msg
data = response.json()
models = []
if "data" in data:
models = data["data"]
page_models: list[dict] = []
if isinstance(data, dict) and isinstance(data.get("data"), list):
page_models = [m for m in data["data"] if isinstance(m, dict)]
elif isinstance(data, list):
models = data
# 为每个模型添加 api_format 字段
for m in models:
page_models = [m for m in data if isinstance(m, dict)]
for m in page_models:
mid = m.get("id")
if isinstance(mid, str) and mid and mid in seen_ids:
continue
if isinstance(mid, str) and mid:
seen_ids.add(mid)
m["api_format"] = cls.FORMAT_ID
return models, None
else:
error_body = response.text[:500] if response.text else "(empty)"
error_msg = f"HTTP {response.status_code}: {error_body}"
logger.warning(f"Claude models request to {models_url} failed: {error_msg}")
return [], error_msg
all_models.append(m)
# Pagination (Anthropic list response shape)
if not isinstance(data, dict):
break
has_more = bool(data.get("has_more"))
last_id = data.get("last_id")
if not has_more:
break
if not isinstance(last_id, str) or not last_id:
break
if after_id == last_id:
# Prevent infinite loops on unexpected upstream behavior.
break
after_id = last_id
return all_models, None
except Exception as e:
error_msg = f"Request error: {str(e)}"
logger.warning(f"Failed to fetch Claude models from {models_url}: {e}")

View File

@@ -137,7 +137,10 @@ class GeminiChatAdapter(ChatAdapterBase):
contents = getattr(request_obj, "contents", []) or []
for content in contents:
role = getattr(content, "role", None) or content.get("role", "unknown")
if isinstance(content, dict):
role = content.get("role", "unknown")
else:
role = getattr(content, "role", None) or "unknown"
role_counts[role] = role_counts.get(role, 0) + 1
generation_config = getattr(request_obj, "generation_config", None) or {}
@@ -262,6 +265,10 @@ class GeminiChatAdapter(ChatAdapterBase):
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
# Provider 上下文Gemini Chat 适配器忽略,仅保持签名兼容)
auth_type: str | None = None, # noqa: ARG003
provider_type: str | None = None, # noqa: ARG003
decrypted_auth_config: dict[str, Any] | None = None, # noqa: ARG003
) -> dict[str, Any]:
"""测试 Gemini API 模型连接性(非流式)"""
from src.api.handlers.base.endpoint_checker import run_endpoint_check
@@ -292,6 +299,7 @@ class GeminiChatAdapter(ChatAdapterBase):
if header_rules:
# 获取认证头名称,防止被规则覆盖
from src.core.api_format import get_auth_config_for_endpoint
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
protected_keys = {auth_header.lower(), "content-type"}

View File

@@ -14,7 +14,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import get_provider_auth
from src.api.handlers.base.request_builder import apply_body_rules, get_provider_auth
from src.api.handlers.base.video_handler_base import (
VideoHandlerBase,
normalize_gemini_operation_id,
@@ -145,6 +145,9 @@ class GeminiVeoHandler(VideoHandlerBase):
format_conversion_info["provider_format"] = provider_format
format_conversion_info["converted"] = needs_conversion
# 应用端点的请求体规则
endpoint_body_rules = getattr(endpoint, "body_rules", None)
if needs_conversion and provider_format.upper().startswith("OPENAI:"):
# Gemini -> OpenAI 格式转换
converted_body = format_conversion_registry.convert_video_request(
@@ -156,6 +159,9 @@ class GeminiVeoHandler(VideoHandlerBase):
if "seconds" in converted_body and converted_body["seconds"] is not None:
converted_body["seconds"] = str(converted_body["seconds"])
if endpoint_body_rules:
converted_body = apply_body_rules(converted_body, endpoint_body_rules)
# 构建 OpenAI 风格的 URL
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
@@ -168,12 +174,18 @@ class GeminiVeoHandler(VideoHandlerBase):
return await client.post(upstream_url, headers=headers, json=converted_body)
else:
# 原始 Gemini 格式
request_body = (
original_request_body.copy() if endpoint_body_rules else original_request_body
)
if endpoint_body_rules:
request_body = apply_body_rules(request_body, endpoint_body_rules)
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(
original_headers, upstream_key, endpoint, auth_info
)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=original_request_body)
return await client.post(upstream_url, headers=headers, json=request_body)
def _extract_task_id(payload: dict[str, Any]) -> str | None:
# 根据响应格式提取 task ID
@@ -514,6 +526,7 @@ class GeminiVeoHandler(VideoHandlerBase):
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
if auth_info:
# 覆盖为 OAuth2 BearerVertex AI
@@ -557,6 +570,7 @@ class GeminiVeoHandler(VideoHandlerBase):
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
def _create_task_record(

View File

@@ -114,7 +114,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
Returns:
包含 input_tokens, output_tokens, cached_tokens 的字典
"""
if str(provider_type or "").lower() == "antigravity":
from src.core.provider_types import ProviderType
if str(provider_type or "").lower() == ProviderType.ANTIGRAVITY:
return self._extract_antigravity_usage(event)
from src.api.handlers.gemini.stream_parser import GeminiStreamParser

View File

@@ -16,7 +16,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import get_provider_auth
from src.api.handlers.base.request_builder import apply_body_rules, get_provider_auth
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
from src.clients.http_client import HTTPClientPool
from src.config.settings import config
@@ -139,6 +139,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
if "seconds" in request_body and request_body["seconds"] is not None:
request_body["seconds"] = str(request_body["seconds"])
# 应用端点的请求体规则
endpoint_body_rules = getattr(endpoint, "body_rules", None)
if needs_conversion and provider_format.upper().startswith("GEMINI:"):
# OpenAI -> Gemini 格式转换
converted_body = format_conversion_registry.convert_video_request(
@@ -150,6 +153,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
if "model" not in converted_body:
converted_body["model"] = internal_request.model
if endpoint_body_rules:
converted_body = apply_body_rules(converted_body, endpoint_body_rules)
# 构建 Gemini 风格的 URL
upstream_url = self._build_gemini_upstream_url(
endpoint.base_url, internal_request.model
@@ -165,6 +171,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
return await client.post(upstream_url, headers=headers, json=converted_body)
else:
# 原始 OpenAI 格式
if endpoint_body_rules:
request_body = apply_body_rules(request_body, endpoint_body_rules)
upstream_url = self._build_upstream_url(endpoint.base_url)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
@@ -500,6 +509,11 @@ class OpenAIVideoHandler(VideoHandlerBase):
if "seconds" in request_body and request_body["seconds"] is not None:
request_body["seconds"] = str(request_body["seconds"])
# 应用端点的请求体规则
endpoint_body_rules = getattr(endpoint, "body_rules", None)
if endpoint_body_rules:
request_body = apply_body_rules(request_body, endpoint_body_rules)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=request_body)
@@ -785,6 +799,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
# ------------------------------------------------------------------
@@ -816,6 +831,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
)
if auth_info:
# 覆盖为 OAuth2 BearerVertex AI

View File

@@ -126,7 +126,7 @@ class OpenAICliAdapter(CliAdapterBase):
# 仅 Codex 端点添加特定头部
if base_url and is_codex_url(base_url):
# 与运行时路径保持一致:使用 Codex envelope 的 best-effort headers。
from src.services.codex.envelope import codex_oauth_envelope
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
headers.update(codex_oauth_envelope.extra_headers() or {})