refactor: 将 adapter 层的计费/模型抓取/行为变体能力下沉到 core.api_format 注册表

- 新增 core/api_format/capabilities.py,统一注册计费模板、模型抓取、
  total_input_context 计算和 provider behavior variant
- 新增 core/usage_tokens.py,抽取 cache token 解析逻辑到 core 层
- handler adapter 移除各自的 compute_total_input_context / fetch_models /
  BILLING_TEMPLATE 覆盖,改为委托 core 注册表解析
- provider/behavior.py 改为薄封装,底层委托 core registry
- 新增 tests/test_architecture_import_rules.py 架构导入约束测试
- 新增 tests/services/api_format/test_capabilities.py 能力注册表测试

Closes #207

Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-03-09 18:26:53 +08:00
parent 4999a1a0a8
commit 0258d01ee6
43 changed files with 1313 additions and 630 deletions

View File

@@ -19,6 +19,8 @@ from pydantic import BaseModel, Field
from sqlalchemy import update
from sqlalchemy.orm import Session, joinedload, make_transient
from src.api.handlers.base.chat_adapter_base import get_adapter_class
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class
from src.config.constants import TimeoutDefaults
from src.core.api_format import get_extra_headers_from_endpoint
from src.core.cache_service import CacheService
@@ -40,7 +42,6 @@ from src.services.model.upstream_fetcher import (
UpstreamModelsFetcherRegistry,
build_format_to_config,
fetch_models_for_key,
get_adapter_for_format,
)
from src.services.provider.oauth_token import resolve_oauth_access_token
from src.services.proxy_node.resolver import resolve_effective_proxy
@@ -77,6 +78,11 @@ async def _set_provider_upstream_models_cache(provider_id: str, models: list[dic
_ANTIGRAVITY_TIER_PRIORITY: dict[str, int] = {"ultra": 3, "pro": 2, "free": 1}
def _get_adapter_for_format(api_format: str) -> Any:
"""按 api_format 获取 Chat/CLI adapter 类。"""
return get_adapter_class(api_format) or get_cli_adapter_class(api_format)
def _antigravity_sort_keys(api_keys: list[Any]) -> list[Any]:
"""按 tier/可用性对 Antigravity Key 降序排列。
@@ -828,7 +834,7 @@ async def test_model(
try:
# 获取对应的 Adapter 类
adapter_class = get_adapter_for_format(endpoint.api_format)
adapter_class = _get_adapter_for_format(endpoint.api_format)
if not adapter_class:
return {
"success": False,
@@ -1497,7 +1503,7 @@ async def _execute_test_check(
if account_id:
extra_headers["chatgpt-account-id"] = str(account_id)
adapter_class = get_adapter_for_format(endpoint.api_format)
adapter_class = _get_adapter_for_format(endpoint.api_format)
if not adapter_class:
raise ValueError(f"Unknown API format: {endpoint.api_format}")

View File

@@ -7,6 +7,7 @@ Chat Adapter 通用基类
- Handler 创建和调用
公共逻辑(异常处理、计费、头部构建等)继承自 HandlerAdapterBase。
计费策略、模型抓取与 provider 格式能力由 `core.api_format` 注册表统一提供。
子类只需提供:
- FORMAT_ID: API 格式标识

View File

@@ -7,12 +7,12 @@ CLI Adapter 通用基类
- Handler 创建和调用
公共逻辑(异常处理、计费、头部构建等)继承自 HandlerAdapterBase。
计费策略、模型抓取与 provider 格式能力由 `core.api_format` 注册表统一提供。
子类只需提供:
- FORMAT_ID: API 格式标识
- HANDLER_CLASS: 对应的 MessageHandler 类
- 可选覆盖 _extract_message_count() 自定义消息计数逻辑
- 可选覆盖 compute_total_input_context() 自定义总输入上下文计算
"""
from __future__ import annotations

View File

@@ -429,7 +429,7 @@ def _extract_tokens_from_response(
# 尝试提取cache creation tokens
try:
from src.api.handlers.base.utils import extract_cache_creation_tokens
from src.core.usage_tokens import extract_cache_creation_tokens
cache_creation_input_tokens = extract_cache_creation_tokens(usage_info)
except Exception as e:
@@ -460,7 +460,7 @@ def _extract_tokens_from_response(
output_tokens = usage_info.get("output_tokens", 0)
cache_read_input_tokens = usage_info.get("cache_read_input_tokens", 0)
try:
from src.api.handlers.base.utils import extract_cache_creation_tokens
from src.core.usage_tokens import extract_cache_creation_tokens
cache_creation_input_tokens = extract_cache_creation_tokens(usage_info)
except Exception as e:

View File

@@ -4,11 +4,11 @@ Handler Adapter 公共基类
从 ChatAdapterBase 和 CliAdapterBase 提取的共享逻辑:
- API 格式与头部处理
- 异常处理和错误响应
- 计费策略
- 模型列表查询和端点测试
- 通过 `core.api_format` 注册表解析计费模板与抓模能力
- 端点测试辅助
- 路径参数合并
子类ChatAdapterBase / CliAdapterBase只需关注各自的 handle() 流程差异。
子类ChatAdapterBase / CliAdapterBase只需关注各自的 `handle()` 流程差异。
"""
from __future__ import annotations
@@ -27,9 +27,12 @@ from src.core.api_format import (
EndpointKind,
build_adapter_base_headers_for_endpoint,
build_adapter_headers_for_endpoint,
compute_total_input_context_for_api_format,
fetch_models_for_api_format,
get_adapter_protected_keys_for_endpoint,
get_auth_handler,
get_default_auth_method_for_endpoint,
resolve_billing_template_for_api_format,
resolve_header_name_case,
)
from src.core.exceptions import (
@@ -50,10 +53,10 @@ class HandlerAdapterBase(ApiAdapter):
封装两者共享的逻辑:
- API 格式与头部处理
- 异常处理和错误响应
- 计费策略
- 模型列表查询和端点测试
- 通过 `core.api_format` 注册表解析计费模板与模型抓取能力
- 端点测试辅助
子类ChatAdapterBase / CliAdapterBase只需实现 handle() 和格式特有的方法。
子类ChatAdapterBase / CliAdapterBase只需实现 `handle()` 和格式特有的方法。
"""
# 子类必须覆盖
@@ -63,7 +66,7 @@ class HandlerAdapterBase(ApiAdapter):
API_FAMILY: ClassVar[ApiFamily | None] = None
ENDPOINT_KIND: ClassVar[EndpointKind] = EndpointKind.CHAT
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini"
# 兼容性回退:若 api_format 注册表未声明计费模板,则使用该默认值。
BILLING_TEMPLATE: str = "claude"
def __init__(self, allowed_api_formats: list[str] | None = None):
@@ -242,18 +245,9 @@ class HandlerAdapterBase(ApiAdapter):
)
# =========================================================================
# 计费策略
# 计费能力委托
# =========================================================================
def compute_total_input_context(
self,
input_tokens: int,
cache_read_input_tokens: int,
cache_creation_input_tokens: int = 0,
) -> int:
"""计算总输入上下文(用于阶梯计费判定)- 子类可覆盖"""
return input_tokens + cache_read_input_tokens
def compute_cost(
self,
input_tokens: int,
@@ -269,8 +263,11 @@ class HandlerAdapterBase(ApiAdapter):
cache_ttl_minutes: int | None = None,
) -> dict[str, Any]:
"""计算请求成本"""
total_input_context = self.compute_total_input_context(
input_tokens, cache_read_input_tokens, cache_creation_input_tokens
total_input_context = compute_total_input_context_for_api_format(
self.FORMAT_ID, input_tokens, cache_read_input_tokens, cache_creation_input_tokens
)
billing_template = (
resolve_billing_template_for_api_format(self.FORMAT_ID) or self.BILLING_TEMPLATE
)
return _calculate_request_cost(
@@ -286,11 +283,11 @@ class HandlerAdapterBase(ApiAdapter):
tiered_pricing=tiered_pricing,
cache_ttl_minutes=cache_ttl_minutes,
total_input_context=total_input_context,
billing_template=self.BILLING_TEMPLATE,
billing_template=billing_template,
)
# =========================================================================
# 模型列表查询与端点测试
# 模型抓取委托与端点测试
# =========================================================================
@classmethod
@@ -301,8 +298,14 @@ class HandlerAdapterBase(ApiAdapter):
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询上游 API 支持的模型列表 - 子类应覆盖"""
return [], f"{cls.FORMAT_ID} adapter does not implement fetch_models"
"""查询上游 API 支持的模型列表"""
return await fetch_models_for_api_format(
client,
api_format=cls.FORMAT_ID,
base_url=base_url,
api_key=api_key,
extra_headers=extra_headers,
)
@classmethod
def build_request_body(

View File

@@ -14,10 +14,10 @@ from src.api.handlers.base.response_parser import (
ResponseParser,
StreamStats,
)
from src.api.handlers.base.utils import extract_cache_creation_tokens
# is_cli_format 权威定义在 core 层
from src.core.api_format import is_cli_format
from src.core.usage_tokens import extract_cache_creation_tokens
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:

View File

@@ -36,97 +36,6 @@ def get_format_converter_registry() -> FormatConversionRegistry:
return format_conversion_registry
def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
"""
提取缓存创建 tokens兼容三种格式
根据 Anthropic API 文档,支持三种格式(按优先级):
1. **嵌套格式(优先级最高)**
usage.cache_creation.ephemeral_5m_input_tokens
usage.cache_creation.ephemeral_1h_input_tokens
2. **扁平新格式(优先级第二)**
usage.claude_cache_creation_5_m_tokens
usage.claude_cache_creation_1_h_tokens
3. **旧格式(优先级第三)**
usage.cache_creation_input_tokens
说明:
- 只要检测到新格式字段(嵌套/扁平),即视为权威来源:哪怕值为 0 也不回退到旧字段。
- 仅当新格式字段完全不存在时,才回退到旧字段。
- 扁平格式和嵌套格式互斥,按顺序检查。
Args:
usage: API 响应中的 usage 字典
Returns:
缓存创建 tokens 总数
"""
# 1. 检查嵌套格式(最新格式)
cache_creation = usage.get("cache_creation")
has_nested_format = isinstance(cache_creation, dict) and (
"ephemeral_5m_input_tokens" in cache_creation
or "ephemeral_1h_input_tokens" in cache_creation
)
if has_nested_format:
cache_5m = int(cache_creation.get("ephemeral_5m_input_tokens", 0))
cache_1h = int(cache_creation.get("ephemeral_1h_input_tokens", 0))
total = cache_5m + cache_1h
logger.debug(f"Using nested cache_creation: 5m={cache_5m}, 1h={cache_1h}, total={total}")
return total
# 2. 检查扁平新格式
has_flat_format = (
"claude_cache_creation_5_m_tokens" in usage or "claude_cache_creation_1_h_tokens" in usage
)
if has_flat_format:
cache_5m = int(usage.get("claude_cache_creation_5_m_tokens", 0))
cache_1h = int(usage.get("claude_cache_creation_1_h_tokens", 0))
total = cache_5m + cache_1h
logger.debug(f"Using flat new format: 5m={cache_5m}, 1h={cache_1h}, total={total}")
return total
# 3. 回退到旧格式
old_format = int(usage.get("cache_creation_input_tokens", 0))
if old_format > 0:
logger.debug(f"Using old format: cache_creation_input_tokens={old_format}")
return old_format
def extract_cache_creation_tokens_detail(usage: dict[str, Any]) -> tuple[int, int, int]:
"""
提取缓存创建 tokens 细分(区分 5m 和 1h
返回 (total, tokens_5m, tokens_1h) 三元组。
当无法区分时tokens_5m 和 tokens_1h 均为 0total 为合计值。
"""
# 1. 嵌套格式
cache_creation = usage.get("cache_creation")
if isinstance(cache_creation, dict) and (
"ephemeral_5m_input_tokens" in cache_creation
or "ephemeral_1h_input_tokens" in cache_creation
):
t5m = int(cache_creation.get("ephemeral_5m_input_tokens", 0))
t1h = int(cache_creation.get("ephemeral_1h_input_tokens", 0))
return t5m + t1h, t5m, t1h
# 2. 扁平新格式
if "claude_cache_creation_5_m_tokens" in usage or "claude_cache_creation_1_h_tokens" in usage:
t5m = int(usage.get("claude_cache_creation_5_m_tokens", 0))
t1h = int(usage.get("claude_cache_creation_1_h_tokens", 0))
return t5m + t1h, t5m, t1h
# 3. 旧格式:无法区分
old = int(usage.get("cache_creation_input_tokens", 0))
return old, 0, 0
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
"""
构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。
@@ -205,9 +114,7 @@ def build_json_response_for_client(
}
cleaned_headers["Content-Encoding"] = "gzip"
existing_vary = next(
(v for k, v in response_headers.items() if k.lower() == "vary"), ""
)
existing_vary = next((v for k, v in response_headers.items() if k.lower() == "vary"), "")
vary_values = [part.strip() for part in str(existing_vary).split(",") if part.strip()]
if not any(part.lower() == "accept-encoding" for part in vary_values):
vary_values.append("Accept-Encoding")

View File

@@ -8,7 +8,6 @@ from __future__ import annotations
from typing import Any
import httpx
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse
@@ -111,7 +110,6 @@ class ClaudeChatAdapter(ChatAdapterBase):
FORMAT_ID = "claude:chat"
API_FAMILY = ApiFamily.CLAUDE
BILLING_TEMPLATE = "claude" # 使用 Claude 计费模板
name = "claude.chat"
@property
@@ -133,23 +131,6 @@ class ClaudeChatAdapter(ChatAdapterBase):
"""检测 Claude 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers, request_body)
# =========================================================================
# Claude 特定的计费逻辑
# =========================================================================
def compute_total_input_context(
self,
input_tokens: int,
cache_read_input_tokens: int,
cache_creation_input_tokens: int = 0,
) -> int:
"""
计算 Claude 的总输入上下文(用于阶梯计费判定)
Claude 的总输入 = input_tokens + cache_creation_input_tokens + cache_read_input_tokens
"""
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
@@ -203,102 +184,6 @@ class ClaudeChatAdapter(ChatAdapterBase):
"thinking_enabled": bool(request_obj.thinking),
}
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表(兼容 x-api-key 和 Bearer 认证)"""
headers = cls.build_headers_with_extra(api_key, extra_headers)
# 兼容第三方提供商:同时发送 Authorization: Bearer 认证头
# 官方 Claude API 使用 x-api-key第三方代理可能使用 Bearer Token
if "authorization" not in {k.lower() for k in headers}:
headers["Authorization"] = f"Bearer {api_key}"
return await cls._fetch_models_paginated(client, base_url, headers, cls.FORMAT_ID)
@staticmethod
async def _fetch_models_paginated(
client: httpx.AsyncClient,
base_url: str,
headers: dict[str, str],
format_id: str,
) -> tuple[list, str | None]:
"""Claude 模型列表分页获取核心逻辑
Anthropic 的 /v1/models 是分页接口has_more/first_id/last_id
默认只返回一页。这里做 best-effort 的全量拉取,确保管理端能展示完整模型列表。
"""
# 构建 /v1/models URL
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
models_url = f"{base_url}/models"
else:
models_url = f"{base_url}/v1/models"
try:
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()
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):
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"] = format_id
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}")
return [], error_msg
@classmethod
def build_endpoint_url(
cls,
base_url: str,

View File

@@ -8,8 +8,8 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.api.handlers.base.utils import extract_cache_creation_tokens_detail
from src.core.api_format import ApiFamily, EndpointKind
from src.core.usage_tokens import extract_cache_creation_tokens_detail
class ClaudeChatHandler(ChatHandlerBase):

View File

@@ -7,7 +7,7 @@ Claude SSE 流解析器
import json
from typing import Any
from src.api.handlers.base.utils import extract_cache_creation_tokens
from src.core.usage_tokens import extract_cache_creation_tokens
class ClaudeStreamParser:

View File

@@ -8,11 +8,9 @@ from __future__ import annotations
from typing import Any
import httpx
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.claude.adapter import ClaudeCapabilityDetector, ClaudeChatAdapter
from src.api.handlers.claude.adapter import ClaudeCapabilityDetector
from src.config.settings import config
from src.core.api_format import ApiFamily
@@ -27,7 +25,6 @@ class ClaudeCliAdapter(CliAdapterBase):
FORMAT_ID = "claude:cli"
API_FAMILY = ApiFamily.CLAUDE
BILLING_TEMPLATE = "claude" # 使用 Claude 计费模板
name = "claude.cli"
@property
@@ -48,23 +45,6 @@ class ClaudeCliAdapter(CliAdapterBase):
"""检测 Claude CLI 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers, request_body)
# =========================================================================
# Claude CLI 特定的计费逻辑
# =========================================================================
def compute_total_input_context(
self,
input_tokens: int,
cache_read_input_tokens: int,
cache_creation_input_tokens: int = 0,
) -> int:
"""
计算 Claude CLI 的总输入上下文(用于阶梯计费判定)
Claude 的总输入 = input_tokens + cache_creation_input_tokens + cache_read_input_tokens
"""
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Claude CLI 使用 messages 字段"""
messages = payload.get("messages", [])
@@ -98,28 +78,6 @@ class ClaudeCliAdapter(CliAdapterBase):
"system_present": bool(payload.get("system")),
}
# =========================================================================
# 模型列表查询
# =========================================================================
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表(使用 CLI Bearer 认证)"""
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
if extra_headers:
cli_headers.update(extra_headers)
# 使用 CLI adapter 自己的认证头Authorization: Bearer而非 Chat 的 x-api-key
headers = cls.build_headers_with_extra(api_key, cli_headers)
return await ClaudeChatAdapter._fetch_models_paginated(
client, base_url, headers, cls.FORMAT_ID
)
@classmethod
def build_endpoint_url(
cls,

View File

@@ -10,8 +10,8 @@ from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
StreamContext,
)
from src.api.handlers.base.utils import extract_cache_creation_tokens_detail
from src.core.api_format import ApiFamily, EndpointKind
from src.core.usage_tokens import extract_cache_creation_tokens_detail
class ClaudeCliMessageHandler(CliMessageHandlerBase):

View File

@@ -16,11 +16,9 @@ from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_ad
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, get_auth_handler, resolve_header_name_case
from src.core.api_format.enums import AuthMethod
from src.core.api_format.headers import BROWSER_FINGERPRINT_HEADERS
from src.core.logger import logger
from src.models.gemini import GeminiRequest
from src.services.gemini_files_mapping import extract_file_names_from_request
from src.services.provider.transport import redact_url_for_log
class GeminiCapabilityDetector:
@@ -54,7 +52,6 @@ class GeminiChatAdapter(ChatAdapterBase):
FORMAT_ID = "gemini:chat"
API_FAMILY = ApiFamily.GEMINI
BILLING_TEMPLATE = "gemini" # 使用 Gemini 计费模板
name = "gemini.chat"
@property
@@ -203,61 +200,6 @@ class GeminiChatAdapter(ChatAdapterBase):
},
)
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表"""
# Gemini 使用 URL 参数传递 key不需要 headers 中的认证
base_url_clean = base_url.rstrip("/")
if base_url_clean.endswith("/v1beta"):
models_url = f"{base_url_clean}/models?key={api_key}"
else:
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
headers: dict[str, str] = {**BROWSER_FINGERPRINT_HEADERS}
if extra_headers:
headers.update(extra_headers)
try:
response = await client.get(models_url, headers=headers)
logger.debug(
f"Gemini models request to {redact_url_for_log(models_url)}: status={response.status_code}"
)
if response.status_code == 200:
data = response.json()
if "models" in data:
# 转换为统一格式
return [
{
"id": m.get("name", "").replace("models/", ""),
"owned_by": "google",
"display_name": m.get("displayName", ""),
"api_format": cls.FORMAT_ID,
}
for m in data["models"]
], None
return [], None
else:
error_body = response.text[:500] if response.text else "(empty)"
error_msg = f"HTTP {response.status_code}: {error_body}"
logger.warning(
f"Gemini models request to {redact_url_for_log(models_url)} failed: {error_msg}"
)
return [], error_msg
except Exception as e:
# 异常信息可能包含带 key 参数的 URL需要脱敏
sanitized_error = redact_url_for_log(str(e))
error_msg = f"Request error: {sanitized_error}"
logger.warning(
f"Failed to fetch Gemini models from {redact_url_for_log(models_url)}: {sanitized_error}"
)
return [], error_msg
@classmethod
def build_endpoint_url(
cls,

View File

@@ -8,12 +8,11 @@ from __future__ import annotations
from typing import Any
import httpx
from fastapi import Request
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.gemini.adapter import GeminiCapabilityDetector, GeminiChatAdapter
from src.api.handlers.gemini.adapter import GeminiCapabilityDetector
from src.config.settings import config
from src.core.api_format import ApiFamily, get_auth_handler
from src.core.api_format.enums import AuthMethod
@@ -29,7 +28,6 @@ class GeminiCliAdapter(CliAdapterBase):
FORMAT_ID = "gemini:cli"
API_FAMILY = ApiFamily.GEMINI
BILLING_TEMPLATE = "gemini" # 使用 Gemini 计费模板
name = "gemini.cli"
@property
@@ -119,29 +117,6 @@ class GeminiCliAdapter(CliAdapterBase):
"safety_settings_count": len(payload.get("safety_settings") or []),
}
# =========================================================================
# 模型列表查询
# =========================================================================
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent"""
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
if extra_headers:
cli_headers.update(extra_headers)
models, error = await GeminiChatAdapter.fetch_models(client, base_url, api_key, cli_headers)
# 更新 api_format 为 CLI 格式
for m in models:
m["api_format"] = cls.FORMAT_ID
return models, error
@classmethod
def build_endpoint_url(
cls,

View File

@@ -8,7 +8,6 @@ from __future__ import annotations
from typing import Any
import httpx
from fastapi.responses import JSONResponse
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
@@ -28,7 +27,6 @@ class OpenAIChatAdapter(ChatAdapterBase):
FORMAT_ID = "openai:chat"
API_FAMILY = ApiFamily.OPENAI
BILLING_TEMPLATE = "openai" # 使用 OpenAI 计费模板
name = "openai.chat"
@property
@@ -105,48 +103,6 @@ class OpenAIChatAdapter(ChatAdapterBase):
},
)
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表"""
headers = cls.build_headers_with_extra(api_key, extra_headers)
# 构建 /v1/models URL
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
models_url = f"{base_url}/models"
else:
models_url = f"{base_url}/v1/models"
try:
response = await client.get(models_url, headers=headers)
logger.debug(f"OpenAI models request to {models_url}: status={response.status_code}")
if response.status_code == 200:
data = response.json()
models = []
if "data" in data:
models = data["data"]
elif isinstance(data, list):
models = data
# 为每个模型添加 api_format 字段
for m in models:
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"OpenAI models request to {models_url} failed: {error_msg}")
return [], error_msg
except Exception as e:
error_msg = f"Request error: {str(e)}"
logger.warning(f"Failed to fetch models from {models_url}: {e}")
return [], error_msg
@classmethod
def build_endpoint_url(
cls,

View File

@@ -8,12 +8,9 @@ from __future__ import annotations
from typing import Any
import httpx
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.openai.adapter import OpenAIChatAdapter
from src.config.settings import config
from src.core.api_format import ApiFamily, EndpointKind
from src.core.provider_types import ProviderType
@@ -30,7 +27,6 @@ class OpenAICliAdapter(CliAdapterBase):
FORMAT_ID = "openai:cli"
API_FAMILY = ApiFamily.OPENAI
BILLING_TEMPLATE = "openai" # 使用 OpenAI 计费模板
name = "openai.cli"
@property
@@ -67,29 +63,6 @@ class OpenAICliAdapter(CliAdapterBase):
set_codex_request_context(CodexRequestContext(is_compact=True))
return await super().handle(context)
# =========================================================================
# 模型列表查询
# =========================================================================
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent"""
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
if extra_headers:
cli_headers.update(extra_headers)
models, error = await OpenAIChatAdapter.fetch_models(client, base_url, api_key, cli_headers)
# 更新 api_format 为 CLI 格式
for m in models:
m["api_format"] = cls.FORMAT_ID
return models, error
@classmethod
def build_endpoint_url(
cls,