mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -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}")
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ Chat Adapter 通用基类
|
||||
- Handler 创建和调用
|
||||
|
||||
公共逻辑(异常处理、计费、头部构建等)继承自 HandlerAdapterBase。
|
||||
计费策略、模型抓取与 provider 格式能力由 `core.api_format` 注册表统一提供。
|
||||
|
||||
子类只需提供:
|
||||
- FORMAT_ID: API 格式标识
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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 均为 0,total 为合计值。
|
||||
"""
|
||||
# 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]:
|
||||
"""
|
||||
构建 SSE(text/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")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user