feat(provider): 增加 Claude Code 适配器、高级配置能力与 OAuth 账号类型统一解析

- 新增 Claude Code provider adapter (context, envelope, plugin, constants)
- 扩展 provider admin 路由,支持 Claude Code 高级配置 (CRUD)
- 统一 OAuth 账号类型解析逻辑,前后端对齐
- 重构 BatchAssignModelsDialog / ModelMappingDialog,简化组件逻辑
- handler 基类增强: request_builder 支持 Claude Code 信封格式
- CLI stream/sync mixin 适配 Claude Code 流式与同步模式
- 扩展 candidate builder / failover / scheduler 对 Claude Code 的支持
- 前端增加请求时间线可视化 (HorizontalRequestTimeline)
- 补充 Claude Code envelope / runtime controls / distributed sessions 等测试

Closes #183
Closes #185

Co-Authored-By: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-02-27 13:54:46 +08:00
parent f2f2a2dbc4
commit 579b5e4623
55 changed files with 3106 additions and 652 deletions

View File

@@ -9,6 +9,7 @@ from .endpoints import router as endpoints_router
from .models import router as models_router
from .modules import router as modules_router
from .monitoring import router as monitoring_router
from .pool import router as pool_router
from .provider_oauth import router as provider_oauth_router
from .provider_ops import router as provider_ops_router
from .provider_query import router as provider_query_router
@@ -38,6 +39,7 @@ router.include_router(security_router)
router.include_router(stats_router)
router.include_router(provider_query_router)
router.include_router(modules_router)
router.include_router(pool_router)
router.include_router(provider_ops_router)
router.include_router(video_tasks_router)

View File

@@ -357,6 +357,43 @@ def _build_kiro_key_name(
return f"{base} ({method})"
def _normalize_codex_plan_group(plan_type: Any) -> str | None:
"""将 Codex plan_type 归一化到判重分组。
分组规则:
- free
- team/plus/enterprise同组
"""
if not isinstance(plan_type, str):
return None
normalized = plan_type.strip().lower()
if not normalized:
return None
if normalized == "free":
return "free"
if normalized in {"team", "plus", "enterprise"}:
return "team_plus_enterprise"
return None
def _is_codex_cross_plan_group_non_duplicate(
*,
new_provider_type: Any,
existing_provider_type: Any,
new_plan_type: Any,
existing_plan_type: Any,
) -> bool:
"""Codex 账号在 free 与 Team/Plus/Enterprise 之间不判重。"""
new_pt = str(new_provider_type or "").strip().lower()
existing_pt = str(existing_provider_type or "").strip().lower()
if new_pt != ProviderType.CODEX.value and existing_pt != ProviderType.CODEX.value:
return False
new_group = _normalize_codex_plan_group(new_plan_type)
existing_group = _normalize_codex_plan_group(existing_plan_type)
return bool(new_group and existing_group and new_group != existing_group)
def _check_duplicate_oauth_account(
db: Session,
provider_id: str,
@@ -368,6 +405,7 @@ def _check_duplicate_oauth_account(
通过以下字段判断重复:
- user_id: Codex 等使用用户级别 ID同 team 下不同成员共享 account_id 但 user_id 不同)
对 Codex 额外按账号类型分组free 与 Team/Plus/Enterprise 互不判重
- email + auth_method: Kiro 使用 email + auth_method 组合判断
(同一邮箱可能通过 Social 和 IdC 两种方式登录,视为不同账号)
- email: 其他 OAuth Provider 使用邮箱判断
@@ -383,6 +421,7 @@ def _check_duplicate_oauth_account(
new_user_id = auth_config.get("user_id")
new_auth_method = auth_config.get("auth_method") # Kiro: social / idc
new_provider_type = auth_config.get("provider_type")
new_plan_type = auth_config.get("plan_type")
# 如果没有可用于识别的字段,跳过检查
if not new_email and not new_user_id:
@@ -409,12 +448,19 @@ def _check_duplicate_oauth_account(
existing_user_id = decrypted_config.get("user_id")
existing_auth_method = decrypted_config.get("auth_method")
existing_provider_type = decrypted_config.get("provider_type")
existing_plan_type = decrypted_config.get("plan_type")
is_duplicate = False
# user_id 相同即重复Codex 等,同一 team 下不同成员共享 account_id 但 user_id 不同)
if new_user_id and existing_user_id and new_user_id == existing_user_id:
is_duplicate = True
if not _is_codex_cross_plan_group_non_duplicate(
new_provider_type=new_provider_type,
existing_provider_type=existing_provider_type,
new_plan_type=new_plan_type,
existing_plan_type=existing_plan_type,
):
is_duplicate = True
# email 判断
if not is_duplicate and new_email and existing_email and new_email == existing_email:
@@ -428,7 +474,13 @@ def _check_duplicate_oauth_account(
):
is_duplicate = True
else:
is_duplicate = True
if not _is_codex_cross_plan_group_non_duplicate(
new_provider_type=new_provider_type,
existing_provider_type=existing_provider_type,
new_plan_type=new_plan_type,
existing_plan_type=existing_plan_type,
):
is_duplicate = True
if is_duplicate:
# 失效账号允许覆盖

View File

@@ -37,6 +37,84 @@ MAPPING_PREVIEW_MAX_MODELS = 500
MAPPING_PREVIEW_TIMEOUT_SECONDS = 10.0
def _should_enable_format_conversion_by_default(provider_type: str | None) -> bool:
"""固定类型 Provider 默认是否开启格式转换。"""
pt = (provider_type or "custom").strip().lower()
envelope_provider_types = {
ProviderType.ANTIGRAVITY.value,
ProviderType.CLAUDE_CODE.value,
ProviderType.CODEX.value,
ProviderType.KIRO.value,
}
return pt in envelope_provider_types
def _normalize_provider_type(provider_type: str | None) -> str:
return (provider_type or "custom").strip().lower()
def _merge_pool_advanced_config(
*,
provider_config: dict[str, Any] | None,
pool_advanced: dict[str, Any] | None,
pool_advanced_in_payload: bool,
) -> tuple[dict[str, Any] | None, bool]:
"""合并 pool_advanced 到 provider.config任何 provider_type 均可使用)。"""
merged_config = dict(provider_config or {})
config_changed = False
if not pool_advanced_in_payload:
return merged_config or None, config_changed
if pool_advanced is None:
if "pool_advanced" in merged_config:
merged_config.pop("pool_advanced", None)
config_changed = True
else:
next_value = dict(pool_advanced)
if merged_config.get("pool_advanced") != next_value:
merged_config["pool_advanced"] = next_value
config_changed = True
return merged_config or None, config_changed
def _merge_claude_code_advanced_config(
*,
provider_type: str | None,
provider_config: dict[str, Any] | None,
claude_code_advanced: dict[str, Any] | None,
claude_advanced_in_payload: bool,
) -> tuple[dict[str, Any] | None, bool]:
"""合并并规范 claude_code_advanced确保仅在 claude_code 下保留。"""
normalized_provider_type = _normalize_provider_type(provider_type)
merged_config = dict(provider_config or {})
config_changed = False
if normalized_provider_type != ProviderType.CLAUDE_CODE.value:
if claude_advanced_in_payload and claude_code_advanced is not None:
raise InvalidRequestException("claude_code_advanced 仅适用于 provider_type=claude_code")
if "claude_code_advanced" in merged_config:
merged_config.pop("claude_code_advanced", None)
config_changed = True
return merged_config or None, config_changed
if not claude_advanced_in_payload:
return merged_config or None, config_changed
if claude_code_advanced is None:
if "claude_code_advanced" in merged_config:
merged_config.pop("claude_code_advanced", None)
config_changed = True
else:
next_value = dict(claude_code_advanced)
if merged_config.get("claude_code_advanced") != next_value:
merged_config["claude_code_advanced"] = next_value
config_changed = True
return merged_config or None, config_changed
# ========== Response Models ==========
@@ -161,7 +239,7 @@ async def create_provider(request: Request, db: Session = Depends(get_db)) -> An
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.put("/{provider_id}")
@router.patch("/{provider_id}")
async def update_provider(
provider_id: str, request: Request, db: Session = Depends(get_db)
) -> None:
@@ -290,15 +368,29 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
else ProviderBillingType.PAY_AS_YOU_GO
)
# 有 envelope 包装的 Provider 类型(如 Antigravity、Codex需要格式转换来正确
# 解包上游响应,创建时默认开启 enable_format_conversion。
pt = (validated_data.provider_type or "custom").strip()
envelope_provider_types = {
ProviderType.ANTIGRAVITY,
ProviderType.CODEX,
ProviderType.KIRO,
}
default_enable_format_conversion = pt in envelope_provider_types
# 有 envelope 包装的 Provider 类型(如 ClaudeCode、Antigravity、Codex需要
# 格式转换来正确解包上游响应,创建时默认开启 enable_format_conversion。
pt = _normalize_provider_type(validated_data.provider_type)
default_enable_format_conversion = _should_enable_format_conversion_by_default(pt)
provider_config, _ = _merge_claude_code_advanced_config(
provider_type=pt,
provider_config=validated_data.config,
claude_code_advanced=(
validated_data.claude_code_advanced.model_dump(exclude_none=True)
if validated_data.claude_code_advanced is not None
else None
),
claude_advanced_in_payload=validated_data.claude_code_advanced is not None,
)
provider_config, _pool_changed = _merge_pool_advanced_config(
provider_config=provider_config,
pool_advanced=(
validated_data.pool_advanced.model_dump(exclude_none=True)
if validated_data.pool_advanced is not None
else None
),
pool_advanced_in_payload=validated_data.pool_advanced is not None,
)
# 创建 Provider 对象
provider = Provider(
@@ -319,7 +411,7 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
# 超时配置
stream_first_byte_timeout=validated_data.stream_first_byte_timeout,
request_timeout=validated_data.request_timeout,
config=validated_data.config,
config=provider_config or None,
# 有 envelope 的反代类型默认开启格式转换
enable_format_conversion=default_enable_format_conversion,
)
@@ -411,6 +503,45 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
try:
# 更新字段(只更新非 None 的字段)
update_data = validated_data.model_dump(exclude_unset=True)
config_in_payload = "config" in update_data
claude_advanced_in_payload = "claude_code_advanced" in update_data
pool_advanced_in_payload = "pool_advanced" in update_data
provider_config = (
dict(update_data.pop("config") or {})
if config_in_payload
else dict(provider.config or {})
)
claude_advanced = (
update_data.pop("claude_code_advanced") if claude_advanced_in_payload else None
)
pool_advanced = update_data.pop("pool_advanced") if pool_advanced_in_payload else None
target_provider_type = (
update_data.get("provider_type")
or getattr(provider, "provider_type", None)
or "custom"
)
provider_config, config_changed_by_claude = _merge_claude_code_advanced_config(
provider_type=target_provider_type,
provider_config=provider_config,
claude_code_advanced=claude_advanced,
claude_advanced_in_payload=claude_advanced_in_payload,
)
provider_config, config_changed_by_pool = _merge_pool_advanced_config(
provider_config=provider_config,
pool_advanced=pool_advanced,
pool_advanced_in_payload=pool_advanced_in_payload,
)
config_touched = (
config_in_payload
or claude_advanced_in_payload
or config_changed_by_claude
or pool_advanced_in_payload
or config_changed_by_pool
)
if config_touched:
update_data["config"] = provider_config
for field, value in update_data.items():
if field == "billing_type" and value is not None:
@@ -719,3 +850,165 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
truncated_keys=truncated_keys,
truncated_models=truncated_models,
)
# ========== Claude Code Pool Management ==========
class PoolKeyStatus(BaseModel):
"""Single key's pool status."""
key_id: str
key_name: str
is_active: bool
cooldown_reason: str | None = None
cooldown_ttl_seconds: int | None = None
cost_window_usage: int = 0
cost_limit: int | None = None
sticky_sessions: int = 0
lru_score: float | None = None
model_config = ConfigDict(from_attributes=True)
class PoolStatusResponse(BaseModel):
"""Pool status for a Provider with pool config."""
provider_id: str
provider_name: str
pool_enabled: bool = False
total_keys: int = 0
total_sticky_sessions: int = 0
keys: list[PoolKeyStatus] = Field(default_factory=list)
model_config = ConfigDict(from_attributes=True)
@router.get("/{provider_id}/pool-status", response_model=PoolStatusResponse)
async def get_pool_status(
request: Request,
provider_id: str,
db: Session = Depends(get_db),
) -> PoolStatusResponse:
"""获取 Provider 的号池状态。"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException("提供商不存在", "provider")
from src.services.provider.pool import redis_ops as pool_redis
from src.services.provider.pool.config import parse_pool_config
pcfg = parse_pool_config(provider.config)
if pcfg is None:
return PoolStatusResponse(
provider_id=provider.id,
provider_name=provider.name,
pool_enabled=False,
)
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == provider_id).all()
key_ids = [str(k.id) for k in keys]
pid = str(provider.id)
import asyncio
# Batch fetch pool state (parallel)
lru_coro = (
pool_redis.get_lru_scores(pid, key_ids) if pcfg.lru_enabled else asyncio.sleep(0, result={})
)
cooldowns, cooldown_ttls, lru_scores, cost_totals, total_sticky = await asyncio.gather(
pool_redis.batch_get_cooldowns(pid, key_ids),
pool_redis.batch_get_cooldown_ttls(pid, key_ids),
lru_coro,
pool_redis.batch_get_cost_totals(pid, key_ids, pcfg.cost_window_seconds),
pool_redis.get_sticky_session_count(pid),
)
# Sticky count per key requires SCAN+MGET; batch with gather.
sticky_counts: dict[str, int] = {}
if key_ids:
counts = await asyncio.gather(
*(pool_redis.get_key_sticky_count(pid, kid) for kid in key_ids)
)
sticky_counts = dict(zip(key_ids, counts))
key_statuses: list[PoolKeyStatus] = []
for k in keys:
kid = str(k.id)
cd_reason = cooldowns.get(kid)
key_statuses.append(
PoolKeyStatus(
key_id=kid,
key_name=k.name or "",
is_active=bool(k.is_active),
cooldown_reason=cd_reason,
cooldown_ttl_seconds=cooldown_ttls.get(kid) if cd_reason else None,
cost_window_usage=cost_totals.get(kid, 0),
cost_limit=pcfg.cost_limit_per_key_tokens,
sticky_sessions=sticky_counts.get(kid, 0),
lru_score=lru_scores.get(kid),
)
)
return PoolStatusResponse(
provider_id=provider.id,
provider_name=provider.name,
pool_enabled=True,
total_keys=len(keys),
total_sticky_sessions=total_sticky,
keys=key_statuses,
)
@router.post("/{provider_id}/pool/clear-cooldown/{key_id}")
async def clear_pool_cooldown(
request: Request,
provider_id: str,
key_id: str,
db: Session = Depends(get_db),
) -> dict[str, str]:
"""手动清除指定 Key 的号池冷却状态。"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException("提供商不存在", "provider")
key = (
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.id == key_id, ProviderAPIKey.provider_id == provider_id)
.first()
)
if not key:
raise NotFoundException("密钥不存在", "key")
from src.services.provider.pool import redis_ops as pool_redis
await pool_redis.clear_cooldown(str(provider.id), str(key.id))
return {"message": f"已清除 Key {key.name or key_id} 的冷却状态"}
@router.post("/{provider_id}/pool/reset-cost/{key_id}")
async def reset_pool_cost(
request: Request,
provider_id: str,
key_id: str,
db: Session = Depends(get_db),
) -> dict[str, str]:
"""重置指定 Key 的号池成本窗口。"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException("提供商不存在", "provider")
key = (
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.id == key_id, ProviderAPIKey.provider_id == provider_id)
.first()
)
if not key:
raise NotFoundException("密钥不存在", "key")
from src.services.provider.pool import redis_ops as pool_redis
await pool_redis.clear_cost(str(provider.id), str(key.id))
return {"message": f"已重置 Key {key.name or key_id} 的成本窗口"}

View File

@@ -15,9 +15,10 @@ from src.api.base.context import ApiRequestContext
from src.api.base.models_service import invalidate_models_list_cache
from src.api.base.pipeline import ApiRequestPipeline
from src.core.enums import ProviderBillingType
from src.core.exceptions import NotFoundException
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.database import get_db
from src.models.admin_requests import ClaudeCodeAdvancedConfig, PoolAdvancedConfig
from src.models.database import (
Model,
Provider,
@@ -213,6 +214,74 @@ async def update_provider_settings(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
def _extract_pool_advanced_from_config(
provider_config: dict[str, Any] | None,
*,
provider_id: str,
) -> PoolAdvancedConfig | None:
"""从 Provider.config 中安全提取通用号池配置。
优先查找 ``pool_advanced``,回退查找 ``claude_code_advanced`` 中的号池字段。
"""
cfg = provider_config or {}
raw = cfg.get("pool_advanced")
if raw is None:
return None
if isinstance(raw, PoolAdvancedConfig):
return raw
if not isinstance(raw, dict):
logger.warning(
"Provider {} 的 pool_advanced 类型无效: {},已忽略",
provider_id,
type(raw).__name__,
)
return None
try:
return PoolAdvancedConfig.model_validate(raw)
except Exception as exc:
logger.warning(
"Provider {} 的 pool_advanced 配置无效,已忽略: {}",
provider_id,
str(exc),
)
return None
def _extract_claude_code_advanced_from_config(
provider_config: dict[str, Any] | None,
*,
provider_id: str,
) -> ClaudeCodeAdvancedConfig | None:
"""从 Provider.config 中安全提取 Claude Code 高级配置。"""
raw_config = (provider_config or {}).get("claude_code_advanced")
if raw_config is None:
return None
if isinstance(raw_config, ClaudeCodeAdvancedConfig):
return raw_config
if not isinstance(raw_config, dict):
logger.warning(
"Provider {} 的 claude_code_advanced 类型无效: {},已忽略",
provider_id,
type(raw_config).__name__,
)
return None
try:
return ClaudeCodeAdvancedConfig.model_validate(raw_config)
except Exception as exc:
logger.warning(
"Provider {} 的 claude_code_advanced 配置无效,已忽略: {}",
provider_id,
str(exc),
)
return None
def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndpointsSummary:
endpoints = db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
@@ -311,12 +380,29 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
for e in endpoints
]
provider_config_raw = provider.config
provider_config = provider_config_raw if isinstance(provider_config_raw, dict) else {}
if provider_config_raw is not None and not isinstance(provider_config_raw, dict):
logger.warning(
"Provider {} 的 config 类型无效: {},按空配置处理",
provider.id,
type(provider_config_raw).__name__,
)
# 检查是否配置了 Provider Ops余额监控等
provider_ops_config = (provider.config or {}).get("provider_ops")
provider_ops_config = provider_config.get("provider_ops")
ops_configured = bool(provider_ops_config)
ops_architecture_id = (
provider_ops_config.get("architecture_id") if provider_ops_config else None
)
claude_code_advanced = _extract_claude_code_advanced_from_config(
provider_config,
provider_id=str(provider.id),
)
pool_advanced = _extract_pool_advanced_from_config(
provider_config,
provider_id=str(provider.id),
)
return ProviderWithEndpointsSummary(
id=provider.id,
@@ -338,6 +424,8 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
proxy=provider.proxy,
stream_first_byte_timeout=provider.stream_first_byte_timeout,
request_timeout=provider.request_timeout,
claude_code_advanced=claude_code_advanced,
pool_advanced=pool_advanced,
total_endpoints=total_endpoints,
active_endpoints=active_endpoints,
total_keys=total_keys,
@@ -496,6 +584,30 @@ class AdminUpdateProviderSettingsAdapter(AdminApiAdapter):
raise NotFoundException("Provider not found", "provider")
update_dict = self.update_data.model_dump(exclude_unset=True)
if "claude_code_advanced" in update_dict:
claude_advanced = update_dict.pop("claude_code_advanced")
provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
if claude_advanced is not None and provider_type != "claude_code":
raise InvalidRequestException(
"claude_code_advanced 仅适用于 provider_type=claude_code"
)
provider_config = dict(provider.config or {})
if claude_advanced is None:
provider_config.pop("claude_code_advanced", None)
else:
provider_config["claude_code_advanced"] = dict(claude_advanced)
update_dict["config"] = provider_config or None
if "pool_advanced" in update_dict:
pool_advanced = update_dict.pop("pool_advanced")
provider_config = dict(update_dict.get("config") or provider.config or {})
if pool_advanced is None:
provider_config.pop("pool_advanced", None)
else:
provider_config["pool_advanced"] = dict(pool_advanced)
update_dict["config"] = provider_config or None
if "billing_type" in update_dict and update_dict["billing_type"] is not None:
update_dict["billing_type"] = ProviderBillingType(update_dict["billing_type"])

View File

@@ -102,6 +102,7 @@ class ProviderRequestResult:
provider_api_format: str = ""
client_api_format: str = ""
auth_info: Any = None
tls_profile: str | None = None
class ChatHandlerBase(BaseMessageHandler, ABC):
@@ -579,6 +580,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
ctx.provider_id = provider_id
ctx.endpoint_id = endpoint_id
ctx.key_id = key_id
if getattr(exec_result, "pool_summary", None):
ctx.pool_summary = exec_result.pool_summary
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_body_ref.get("_rectified", False)
@@ -714,6 +717,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
policy=upstream_policy,
)
# Envelope lifecycle: prepare_context (pre-wrap hook).
envelope_tls_profile: str | None = None
if envelope and hasattr(envelope, "prepare_context"):
envelope_tls_profile = envelope.prepare_context(
provider_config=getattr(provider, "config", None),
key_id=str(getattr(key, "id", "") or ""),
is_stream=upstream_is_stream,
provider_id=str(getattr(provider, "id", "") or ""),
)
# 跨格式:先做请求体转换(失败触发 failover
registry = get_format_converter_registry()
if needs_conversion:
@@ -776,6 +789,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
url_model=url_model,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# Envelope lifecycle: post_wrap_request (post-wrap hook).
if hasattr(envelope, "post_wrap_request"):
await envelope.post_wrap_request(request_body)
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
extra_headers: dict[str, str] = {}
@@ -793,6 +809,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider_api_format=provider_api_format,
client_api_format=client_api_format,
auth_info=auth_info,
tls_profile=envelope_tls_profile,
)
async def _execute_stream_request(
@@ -853,6 +870,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
envelope = prep.envelope
upstream_is_stream = prep.upstream_is_stream
auth_info = prep.auth_info
tls_profile = prep.tls_profile
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
provider_payload, provider_headers = self._request_builder.build(
@@ -863,6 +881,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
is_stream=upstream_is_stream,
extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
@@ -908,7 +927,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
request_timeout_sync = provider.request_timeout or config.http_request_timeout
delegate_cfg = resolve_delegate_config(effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=effective_proxy
delegate_cfg,
proxy_config=effective_proxy,
tls_profile=tls_profile,
)
try:
@@ -1104,7 +1125,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
delegate_cfg = resolve_delegate_config(effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=effective_proxy
delegate_cfg,
proxy_config=effective_proxy,
tls_profile=tls_profile,
)
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)

View File

@@ -193,6 +193,8 @@ class ChatSyncExecutor:
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
if getattr(exec_result, "pool_summary", None):
request_metadata["pool_summary"] = exec_result.pool_summary
total_cost = await handler.telemetry.record_success( # noqa: F841
provider=ctx.provider_name,
model=model,
@@ -425,6 +427,7 @@ class ChatSyncExecutor:
envelope = prep.envelope
upstream_is_stream = prep.upstream_is_stream
auth_info = prep.auth_info
tls_profile = prep.tls_profile
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
provider_payload, provider_hdrs = handler._request_builder.build(
@@ -435,6 +438,7 @@ class ChatSyncExecutor:
is_stream=upstream_is_stream,
extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
@@ -496,7 +500,9 @@ class ChatSyncExecutor:
delegate_cfg = resolve_delegate_config(_effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=_effective_proxy
delegate_cfg,
proxy_config=_effective_proxy,
tls_profile=tls_profile,
)
# 注意:不使用 async with因为复用的客户端不应该被关闭

View File

@@ -342,6 +342,8 @@ class CliPrefetchMixin:
last_data_time = time.time()
buffer = b""
output_state = {"first_yield": True, "streaming_updated": False}
_sample_lines: list[str] = [] # 采集前几行原始内容,用于空流诊断
_MAX_SAMPLE_LINES = 5
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
@@ -409,6 +411,8 @@ class CliPrefetchMixin:
continue
ctx.chunk_count += 1
if len(_sample_lines) < _MAX_SAMPLE_LINES:
_sample_lines.append(normalized_line[:200])
# 格式转换或直接透传
if needs_conversion:
@@ -468,6 +472,8 @@ class CliPrefetchMixin:
continue
ctx.chunk_count += 1
if len(_sample_lines) < _MAX_SAMPLE_LINES:
_sample_lines.append(normalized_line[:200])
# 空流检测:超过阈值且无数据,发送错误事件并结束
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
@@ -524,9 +530,10 @@ class CliPrefetchMixin:
# 检查是否收到数据
if ctx.data_count == 0:
# 空流通常意味着配置错误(如 base_url 指向了网页而非 API
sample_info = f", 前几行内容: {_sample_lines!r}" if _sample_lines else ""
logger.error(
f"Provider '{ctx.provider_name}' 返回空流式响应 (收到 {ctx.chunk_count} 个非数据行), "
f"可能是 endpoint base_url 配置错误"
f"可能是 endpoint base_url 配置错误{sample_info}"
)
# 设置错误状态用于后续记录
ctx.status_code = 503

View File

@@ -188,6 +188,8 @@ class CliStreamMixin:
ctx.endpoint_id = endpoint_id
if not ctx.key_id:
ctx.key_id = key_id
if getattr(exec_result, "pool_summary", None):
ctx.pool_summary = exec_result.pool_summary
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_body_ref.get("_rectified", False)
@@ -319,6 +321,15 @@ class CliStreamMixin:
client_is_stream=True,
policy=upstream_policy,
)
# Envelope lifecycle: prepare_context (pre-wrap hook).
envelope_tls_profile: str | None = None
if envelope and hasattr(envelope, "prepare_context"):
envelope_tls_profile = envelope.prepare_context(
provider_config=getattr(provider, "config", None),
key_id=str(getattr(key, "id", "") or ""),
is_stream=upstream_is_stream,
provider_id=str(getattr(provider, "id", "") or ""),
)
# 跨格式:先做请求体转换(失败触发 failover
if needs_conversion and provider_api_format:
@@ -374,6 +385,9 @@ class CliStreamMixin:
url_model=url_model,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# Envelope lifecycle: post_wrap_request (post-wrap hook).
if hasattr(envelope, "post_wrap_request"):
await envelope.post_wrap_request(request_body)
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
extra_headers: dict[str, str] = {}
@@ -391,6 +405,7 @@ class CliStreamMixin:
is_stream=upstream_is_stream,
extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
@@ -429,7 +444,9 @@ class CliStreamMixin:
request_timeout_sync = provider.request_timeout or config.http_request_timeout
delegate_cfg = resolve_delegate_config(effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=effective_proxy
delegate_cfg,
proxy_config=effective_proxy,
tls_profile=envelope_tls_profile,
)
try:
@@ -623,7 +640,9 @@ class CliStreamMixin:
delegate_cfg = resolve_delegate_config(effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=effective_proxy
delegate_cfg,
proxy_config=effective_proxy,
tls_profile=envelope_tls_profile,
)
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
@@ -798,6 +817,8 @@ class CliStreamMixin:
last_data_time = time.time()
buffer = b""
output_state = {"first_yield": True, "streaming_updated": False}
_sample_lines: list[str] = [] # 采集前几行原始内容,用于空流诊断
_MAX_SAMPLE_LINES = 5
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
@@ -864,6 +885,8 @@ class CliStreamMixin:
continue
ctx.chunk_count += 1
if len(_sample_lines) < _MAX_SAMPLE_LINES:
_sample_lines.append(normalized_line[:200])
# 空流检测:超过阈值且无数据,发送错误事件并结束
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
@@ -922,7 +945,12 @@ class CliStreamMixin:
# 检查是否收到数据
if ctx.data_count == 0:
logger.warning("Provider '{}' 返回空流式响应", ctx.provider_name)
sample_info = f", 前几行内容: {_sample_lines!r}" if _sample_lines else ""
logger.warning(
"Provider '{}' 返回空流式响应{}",
ctx.provider_name,
sample_info,
)
ctx.status_code = 503
ctx.error_message = "上游服务返回了空的流式响应"
ctx.upstream_response = (

View File

@@ -147,6 +147,15 @@ class CliSyncMixin:
client_is_stream=False,
policy=upstream_policy,
)
# Envelope lifecycle: prepare_context (pre-wrap hook).
envelope_tls_profile: str | None = None
if envelope and hasattr(envelope, "prepare_context"):
envelope_tls_profile = envelope.prepare_context(
provider_config=getattr(provider, "config", None),
key_id=str(getattr(key, "id", "") or ""),
is_stream=upstream_is_stream,
provider_id=str(getattr(provider, "id", "") or ""),
)
# 跨格式:先做请求体转换(失败触发 failover
if needs_conversion and provider_api_format:
@@ -202,6 +211,9 @@ class CliSyncMixin:
url_model=url_model,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# Envelope lifecycle: post_wrap_request (post-wrap hook).
if hasattr(envelope, "post_wrap_request"):
await envelope.post_wrap_request(request_body)
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
extra_headers: dict[str, str] = {}
@@ -219,6 +231,7 @@ class CliSyncMixin:
is_stream=upstream_is_stream,
extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
@@ -274,7 +287,9 @@ class CliSyncMixin:
delegate_cfg = resolve_delegate_config(_effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=_effective_proxy
delegate_cfg,
proxy_config=_effective_proxy,
tls_profile=envelope_tls_profile,
)
# 注意:不使用 async with因为复用的客户端不应该被关闭
@@ -521,6 +536,8 @@ class CliSyncMixin:
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
if getattr(exec_result, "pool_summary", None):
request_metadata["pool_summary"] = exec_result.pool_summary
total_cost = await self.telemetry.record_success(
provider=provider_name,
model=model,

View File

@@ -26,9 +26,9 @@ from src.core.api_format import (
make_signature_key,
)
from src.core.crypto import crypto_service
from src.core.provider_auth_types import ProviderAuthInfo
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
from src.services.provider.auth import get_provider_auth
from src.services.provider.auth import get_provider_auth # noqa: F401
from src.services.provider.envelope import ProviderEnvelope
# ==============================================================================
# 统一的头部配置常量
@@ -1036,6 +1036,7 @@ class RequestBuilder(ABC):
*,
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
envelope: ProviderEnvelope | None = None,
) -> dict[str, str]:
"""构建请求头"""
pass
@@ -1051,6 +1052,7 @@ class RequestBuilder(ABC):
is_stream: bool = False,
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
envelope: ProviderEnvelope | None = None,
) -> tuple[dict[str, Any], dict[str, str]]:
"""
构建完整的请求(请求体 + 请求头)
@@ -1085,6 +1087,7 @@ class RequestBuilder(ABC):
key,
extra_headers=extra_headers,
pre_computed_auth=pre_computed_auth,
envelope=envelope,
)
return payload, headers
@@ -1114,6 +1117,70 @@ class PassthroughRequestBuilder(RequestBuilder):
"""
return dict(original_body)
@staticmethod
def _merge_comma_header_values(primary: str, secondary: str) -> str:
"""合并逗号分隔 header 值并去重,保持 primary 在前。"""
seen: set[str] = set()
merged: list[str] = []
def _append(raw: str) -> None:
for token in str(raw or "").split(","):
token = token.strip()
if not token or token in seen:
continue
seen.add(token)
merged.append(token)
_append(primary)
_append(secondary)
return ",".join(merged)
@classmethod
def _drop_beta_token(cls, value: str, token: str) -> str:
"""从逗号分隔 header 中移除指定 token。"""
if not value or token not in value:
return value
return cls._merge_comma_header_values(
",".join(p.strip() for p in str(value).split(",") if p.strip() and p.strip() != token),
"",
)
@classmethod
def _merge_extra_headers_with_original(
cls,
original_headers: dict[str, str],
extra_headers: dict[str, str] | None,
*,
envelope: ProviderEnvelope | None = None,
) -> dict[str, str] | None:
"""合并 extra_headers 与原始头部中的特定字段。"""
if not extra_headers:
return None
merged_extra = dict(extra_headers)
beta_extra_key = next((k for k in merged_extra if k.lower() == "anthropic-beta"), None)
if beta_extra_key is None:
return merged_extra
incoming_beta = next(
(v for k, v in original_headers.items() if k.lower() == "anthropic-beta"),
"",
)
merged_beta = str(merged_extra.get(beta_extra_key) or "")
if incoming_beta:
merged_beta = cls._merge_comma_header_values(
merged_beta,
str(incoming_beta),
)
# 由 envelope 声明需要排除的 beta token如 Claude Code OAuth 的 context-1m
if envelope and hasattr(envelope, "excluded_beta_tokens"):
for token in envelope.excluded_beta_tokens():
merged_beta = cls._drop_beta_token(merged_beta, token)
merged_extra[beta_extra_key] = merged_beta
return merged_extra
def build_headers(
self,
original_headers: dict[str, str],
@@ -1122,6 +1189,7 @@ class PassthroughRequestBuilder(RequestBuilder):
*,
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
envelope: ProviderEnvelope | None = None,
) -> dict[str, str]:
"""
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
@@ -1177,8 +1245,13 @@ class PassthroughRequestBuilder(RequestBuilder):
builder.apply_rules(header_rules, protected_keys)
# 4. 添加额外头部
if extra_headers:
builder.add_many(extra_headers)
effective_extra_headers = self._merge_extra_headers_with_original(
original_headers,
extra_headers,
envelope=envelope,
)
if effective_extra_headers:
builder.add_many(effective_extra_headers)
# 5. 设置认证头(最高优先级,上游始终使用 header 认证)
builder.add(auth_header, auth_value)

View File

@@ -135,6 +135,9 @@ class StreamContext:
# 代理信息(用于 usage 记录和日志,含 ttfb_ms
proxy_info: dict[str, Any] | None = None
# 号池调度摘要(来自 ExecutionResult.pool_summary
pool_summary: dict[str, Any] | None = None
# 流式格式转换状态(跨 chunk 追踪)
stream_conversion_state: StreamState | None = None
stream_conversion_event_count: int = 0 # 流式转换成功的 event 计数

View File

@@ -208,6 +208,8 @@ class StreamTelemetryRecorder:
metadata["perf"] = ctx.perf_metrics
if ctx.proxy_info:
metadata["proxy"] = ctx.proxy_info
if ctx.pool_summary:
metadata["pool_summary"] = ctx.pool_summary
await writer.record_success(
provider=ctx.provider_name or "unknown",

View File

@@ -25,7 +25,7 @@ from src.services.proxy_node.resolver import (
get_system_proxy_config,
make_proxy_param,
)
from src.utils.ssl_utils import get_ssl_context
from src.utils.ssl_utils import get_ssl_context, get_ssl_context_for_profile
# 模块级锁,避免类属性延迟初始化的竞态条件
_proxy_clients_lock = asyncio.Lock()
@@ -186,6 +186,7 @@ class HTTPClientPool:
async def get_proxy_client(
cls,
proxy_config: dict[str, Any] | None = None,
tls_profile: str | None = None,
) -> httpx.AsyncClient:
"""
获取代理客户端(带缓存复用)
@@ -212,6 +213,9 @@ class HTTPClientPool:
return await cls._get_tunnel_client(delegate_cfg["node_id"])
cache_key = compute_proxy_cache_key(proxy_config)
tls_profile_key = str(tls_profile or "").strip().lower()
if tls_profile_key:
cache_key = f"{cache_key}::tls:{tls_profile_key}"
# 无代理时返回默认客户端
if cache_key == "__no_proxy__":
@@ -229,6 +233,10 @@ class HTTPClientPool:
else:
# 更新最后使用时间
cls._proxy_clients[cache_key] = (client, time.time())
if tls_profile_key:
logger.debug(
"复用代理客户端 TLS profile={} key={}", tls_profile_key, cache_key
)
return client
# 淘汰旧客户端(如果超过上限)
@@ -237,7 +245,7 @@ class HTTPClientPool:
# 创建新客户端(使用默认超时,请求时可覆盖)
client_config: dict[str, Any] = {
"http2": False,
"verify": get_ssl_context(),
"verify": get_ssl_context_for_profile(tls_profile),
"follow_redirects": True,
"limits": httpx.Limits(
max_connections=config.http_max_connections,
@@ -269,6 +277,8 @@ class HTTPClientPool:
logger.debug(
"创建代理客户端(缓存): {}, 缓存数量: {}", proxy_label, len(cls._proxy_clients)
)
if tls_profile_key:
logger.debug("创建代理客户端 TLS profile={} key={}", tls_profile_key, cache_key)
return client
@@ -342,6 +352,7 @@ class HTTPClientPool:
cls,
proxy_config: dict[str, Any] | None = None,
timeout: httpx.Timeout | None = None,
tls_profile: str | None = None,
**kwargs: Any,
) -> httpx.AsyncClient:
"""
@@ -359,7 +370,7 @@ class HTTPClientPool:
"""
client_config: dict[str, Any] = {
"http2": False,
"verify": get_ssl_context(),
"verify": get_ssl_context_for_profile(tls_profile),
"follow_redirects": True,
}
@@ -392,6 +403,7 @@ class HTTPClientPool:
cls,
delegate_cfg: dict[str, Any] | None,
proxy_config: dict[str, Any] | None = None,
tls_profile: str | None = None,
) -> httpx.AsyncClient:
"""
获取可复用的上游请求客户端(自动选择 tunnel/代理模式)
@@ -401,7 +413,7 @@ class HTTPClientPool:
"""
if delegate_cfg and delegate_cfg.get("tunnel"):
return await cls._get_tunnel_client(delegate_cfg["node_id"])
return await cls.get_proxy_client(proxy_config=proxy_config)
return await cls.get_proxy_client(proxy_config=proxy_config, tls_profile=tls_profile)
@classmethod
async def create_upstream_stream_client(
@@ -409,6 +421,7 @@ class HTTPClientPool:
delegate_cfg: dict[str, Any] | None,
proxy_config: dict[str, Any] | None = None,
timeout: httpx.Timeout | None = None,
tls_profile: str | None = None,
) -> httpx.AsyncClient:
"""
创建上游流式请求客户端(自动选择 tunnel/代理模式)
@@ -417,7 +430,11 @@ class HTTPClientPool:
"""
if delegate_cfg and delegate_cfg.get("tunnel"):
return await cls._get_tunnel_client(delegate_cfg["node_id"], timeout=timeout)
return cls.create_client_with_proxy(proxy_config=proxy_config, timeout=timeout)
return cls.create_client_with_proxy(
proxy_config=proxy_config,
timeout=timeout,
tls_profile=tls_profile,
)
@classmethod
async def _get_tunnel_client(

View File

@@ -520,6 +520,7 @@ class ClaudeNormalizer(FormatNormalizer):
message_raw = chunk.get("message")
message: dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
msg_id = str(message.get("id") or "")
usage_info = self._claude_usage_to_internal(message.get("usage"))
# 保留初始化时设置的 model客户端请求的模型仅在空时用上游值
model = state.model or str(message.get("model") or "")
state.message_id = msg_id or state.message_id
@@ -527,7 +528,9 @@ class ClaudeNormalizer(FormatNormalizer):
state.model = model
ss["message_started"] = True
ss.setdefault("block_index_to_tool_id", {})
events.append(MessageStartEvent(message_id=msg_id, model=model))
if usage_info is not None:
ss["usage"] = message.get("usage")
events.append(MessageStartEvent(message_id=msg_id, model=model, usage=usage_info))
return events
if event_type == "content_block_start":

View File

@@ -80,6 +80,87 @@ class ProxyConfig(BaseModel):
return self
class PoolAdvancedConfig(BaseModel):
"""通用号池配置(适用于所有 Provider 类型)。"""
sticky_session_ttl_seconds: int | None = Field(
None,
ge=60,
le=86400,
description="粘性会话 TTL同一对话始终路由到同一 Key。None = 禁用",
)
load_threshold_percent: int | None = Field(
None,
ge=10,
le=100,
description="负载率阈值(%),超过时该 Key 被降权。默认 80",
)
lru_enabled: bool = Field(True, description="LRU 调度(优先选择最久未用的 Key")
cost_window_seconds: int | None = Field(
None,
ge=3600,
le=86400,
description="滚动成本窗口(秒)。默认 180005 小时)",
)
cost_limit_per_key_tokens: int | None = Field(
None, ge=0, description="每个 Key 在窗口内的最大 token 用量。None = 不限"
)
cost_soft_threshold_percent: int | None = Field(
None,
ge=0,
le=100,
description="成本软阈值(%),超过时优先选用其他 Key。默认 80",
)
rate_limit_cooldown_seconds: int | None = Field(
None, ge=10, le=3600, description="429 冷却时间(秒)。默认 300"
)
overload_cooldown_seconds: int | None = Field(
None, ge=5, le=600, description="529 冷却时间(秒)。默认 30"
)
proactive_refresh_seconds: int | None = Field(
None,
ge=60,
le=600,
description="OAuth Token 提前刷新秒数。默认 1803 分钟)",
)
health_policy_enabled: bool = Field(
True, description="启用号池健康策略(按上游错误码自动冷却/禁用 Key"
)
unschedulable_rules: list[dict] | None = Field(
None,
description="关键词临时不可调度规则: [{'keyword': '...', 'duration_minutes': 5}]",
)
class ClaudeCodeAdvancedConfig(BaseModel):
"""Claude Code 特有配置。"""
max_sessions: int | None = Field(
None, ge=1, le=1000, description="最大活跃会话数(为空表示不限制)"
)
session_idle_timeout_minutes: int | None = Field(
None, ge=1, le=1440, description="会话空闲超时(分钟)"
)
enable_tls_fingerprint: bool = Field(
False, description="是否启用 TLS 指纹模拟(模拟 Node.js/Claude Code 客户端)"
)
session_id_masking_enabled: bool = Field(
False, description="是否启用会话 ID 伪装(固定 metadata.user_id 中 session 片段)"
)
@model_validator(mode="after")
def normalize_session_control(self) -> "ClaudeCodeAdvancedConfig":
# 未启用会话限制时,不保留超时配置,避免产生误导。
if self.max_sessions is None:
self.session_idle_timeout_minutes = None
return self
# 启用会话限制但未设置超时时,回落到 5 分钟默认值。
if self.session_idle_timeout_minutes is None:
self.session_idle_timeout_minutes = 5
return self
class CreateProviderRequest(BaseModel):
"""创建 Provider 请求"""
@@ -148,6 +229,12 @@ class CreateProviderRequest(BaseModel):
request_timeout: float | None = Field(
None, ge=1, le=600, description="非流式请求整体超时(秒)"
)
pool_advanced: PoolAdvancedConfig | None = Field(
None, description="号池高级配置(适用于所有 Provider 类型)"
)
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
None, description="Claude Code 特有配置"
)
config: dict[str, Any] | None = Field(None, description="其他配置")
@field_validator("provider_type")
@@ -214,6 +301,13 @@ class CreateProviderRequest(BaseModel):
valid_types = [t.value for t in ProviderBillingType]
raise ValueError(f"无效的计费类型,有效值为: {', '.join(valid_types)}")
@model_validator(mode="after")
def validate_claude_code_advanced_scope(self) -> "CreateProviderRequest":
provider_type = (self.provider_type or "custom").strip()
if self.claude_code_advanced is not None and provider_type != "claude_code":
raise ValueError("claude_code_advanced 仅适用于 provider_type=claude_code")
return self
class UpdateProviderRequest(BaseModel):
"""更新 Provider 请求"""
@@ -244,6 +338,12 @@ class UpdateProviderRequest(BaseModel):
request_timeout: float | None = Field(
None, ge=1, le=600, description="非流式请求整体超时(秒)"
)
pool_advanced: PoolAdvancedConfig | None = Field(
None, description="号池高级配置(适用于所有 Provider 类型)"
)
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
None, description="Claude Code 特有配置"
)
config: dict[str, Any] | None = None
# 复用相同的验证器
@@ -258,6 +358,15 @@ class UpdateProviderRequest(BaseModel):
CreateProviderRequest.validate_provider_type.__func__
)
@model_validator(mode="after")
def validate_claude_code_advanced_scope(self) -> "UpdateProviderRequest":
# 更新场景下 provider_type 可能不在 payload 中,最终校验由路由层结合数据库值完成。
if self.claude_code_advanced is not None and self.provider_type is not None:
provider_type = (self.provider_type or "custom").strip()
if provider_type != "claude_code":
raise ValueError("claude_code_advanced 仅适用于 provider_type=claude_code")
return self
class CreateEndpointRequest(BaseModel):
"""创建 Endpoint 请求"""

View File

@@ -10,7 +10,7 @@ from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator
from src.models.admin_requests import ProxyConfig
from src.models.admin_requests import ClaudeCodeAdvancedConfig, PoolAdvancedConfig, ProxyConfig
# ========== Header Rule 类型定义 ==========
# 请求头规则支持三种操作:
@@ -934,6 +934,10 @@ class ProviderUpdateRequest(BaseModel):
request_timeout: float | None = Field(
None, ge=1, le=600, description="非流式请求整体超时(秒)"
)
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
None, description="Claude Code 高级配置"
)
pool_advanced: PoolAdvancedConfig | None = Field(None, description="通用号池配置")
class ProviderWithEndpointsSummary(BaseModel):
@@ -974,6 +978,10 @@ class ProviderWithEndpointsSummary(BaseModel):
default=None, description="流式请求首字节超时(秒)"
)
request_timeout: float | None = Field(default=None, description="非流式请求整体超时(秒)")
claude_code_advanced: ClaudeCodeAdvancedConfig | None = Field(
default=None, description="Claude Code 高级配置"
)
pool_advanced: PoolAdvancedConfig | None = Field(default=None, description="通用号池配置")
# Endpoint 统计
total_endpoints: int = Field(default=0, description="总 Endpoint 数量")

View File

@@ -513,6 +513,10 @@ class FailoverEngine:
api_key_id: str | None,
) -> str:
# Create "available" record, then caller will mark pending.
extra: dict = {}
pool_extra = getattr(candidate, "_pool_extra_data", None)
if pool_extra:
extra.update(pool_extra)
row = RequestCandidateService.create_candidate(
db=self.db,
request_id=request_id,
@@ -525,7 +529,7 @@ class FailoverEngine:
key_id=str(candidate.key.id),
status="available",
is_cached=bool(getattr(candidate, "is_cached", False)),
extra_data={},
extra_data=extra,
)
return str(row.id)
@@ -539,6 +543,10 @@ class FailoverEngine:
api_key_id: str | None,
skip_reason: str | None,
) -> str:
extra: dict = {}
pool_extra = getattr(candidate, "_pool_extra_data", None)
if pool_extra:
extra.update(pool_extra)
row = RequestCandidateService.create_candidate(
db=self.db,
request_id=request_id,
@@ -552,7 +560,7 @@ class FailoverEngine:
status="skipped",
skip_reason=skip_reason,
is_cached=bool(getattr(candidate, "is_cached", False)),
extra_data={},
extra_data=extra,
)
# ensure visible for subsequent recorder reads
if self.db.in_transaction():

View File

@@ -52,6 +52,7 @@ class CandidateResolver:
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | None = None,
request_body: dict | None = None,
) -> tuple[list[ProviderCandidate], str]:
"""
获取所有可用候选
@@ -96,6 +97,7 @@ class CandidateResolver:
provider_limit=provider_batch_size,
is_stream=is_stream,
capability_requirements=capability_requirements,
request_body=request_body,
)
)

View File

@@ -0,0 +1,3 @@
"""Claude Code provider adapter."""
__all__ = []

View File

@@ -0,0 +1,52 @@
"""Claude Code adapter constants."""
from __future__ import annotations
CLAUDE_MESSAGES_PATH = "/v1/messages"
DEFAULT_ANTHROPIC_VERSION = "2023-06-01"
DEFAULT_ACCEPT = "application/json"
STREAM_HELPER_METHOD = "stream"
SESSION_ID_MASKING_TTL_SECONDS = 15 * 60
# 仅代表“启用 Claude Code TLS 配置”的客户端 profile 标识best-effort
TLS_PROFILE_CLAUDE_CODE = "claude_code_nodejs"
# Claude Code OAuth required betas.
BETA_CLAUDE_CODE = "claude-code-20250219"
BETA_OAUTH = "oauth-2025-04-20"
BETA_INTERLEAVED_THINKING = "interleaved-thinking-2025-05-14"
BETA_CONTEXT_1M = "context-1m-2025-08-07"
CLAUDE_CODE_REQUIRED_BETA_TOKENS: tuple[str, ...] = (
BETA_CLAUDE_CODE,
BETA_OAUTH,
BETA_INTERLEAVED_THINKING,
)
# Mimic headers observed from Claude Code traffic.
CLAUDE_CODE_DEFAULT_HEADERS: dict[str, str] = {
"X-Stainless-Lang": "js",
"X-Stainless-Package-Version": "0.70.0",
"X-Stainless-OS": "Linux",
"X-Stainless-Arch": "arm64",
"X-Stainless-Runtime": "node",
"X-Stainless-Runtime-Version": "v24.13.0",
"X-Stainless-Retry-Count": "0",
"X-Stainless-Timeout": "600",
"X-App": "cli",
"Anthropic-Dangerous-Direct-Browser-Access": "true",
}
__all__ = [
"BETA_CLAUDE_CODE",
"BETA_CONTEXT_1M",
"BETA_INTERLEAVED_THINKING",
"BETA_OAUTH",
"CLAUDE_CODE_DEFAULT_HEADERS",
"CLAUDE_CODE_REQUIRED_BETA_TOKENS",
"CLAUDE_MESSAGES_PATH",
"DEFAULT_ACCEPT",
"DEFAULT_ANTHROPIC_VERSION",
"SESSION_ID_MASKING_TTL_SECONDS",
"STREAM_HELPER_METHOD",
"TLS_PROFILE_CLAUDE_CODE",
]

View File

@@ -0,0 +1,141 @@
"""Claude Code request context using contextvars."""
from __future__ import annotations
import contextvars
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from src.core.logger import logger
from src.models.admin_requests import ClaudeCodeAdvancedConfig
from src.services.provider.adapters.claude_code.constants import TLS_PROFILE_CLAUDE_CODE
if TYPE_CHECKING:
from src.services.provider.pool.config import PoolConfig
@dataclass(frozen=True, slots=True)
class ClaudeCodeRequestContext:
is_stream: bool = False
# 使用 Key 级别作用域,确保会话限制按 OAuth 账号隔离。
scope_key: str | None = None
key_id: str | None = None
max_sessions: int | None = None
session_idle_timeout_minutes: int = 5
enable_tls_fingerprint: bool = False
session_id_masking_enabled: bool = False
# Account Pool fields
provider_id: str | None = None
pool_config: PoolConfig | None = None
session_uuid: str | None = None
_claude_code_request_context: contextvars.ContextVar[ClaudeCodeRequestContext | None] = (
contextvars.ContextVar(
"claude_code_request_context",
default=None,
)
)
def set_claude_code_request_context(ctx: ClaudeCodeRequestContext | None) -> None:
_claude_code_request_context.set(ctx)
def get_claude_code_request_context() -> ClaudeCodeRequestContext | None:
return _claude_code_request_context.get()
def build_claude_code_request_context(
*,
provider_config: Any,
key_id: str | None,
is_stream: bool,
provider_id: str | None = None,
) -> ClaudeCodeRequestContext:
"""根据 Provider.config 构建 Claude Code 请求上下文。"""
from src.services.provider.pool.config import parse_pool_config
normalized_key_id = str(key_id or "").strip() or None
advanced_config: ClaudeCodeAdvancedConfig | None = None
provider_config_dict = provider_config if isinstance(provider_config, dict) else {}
raw_advanced = provider_config_dict.get("claude_code_advanced")
if raw_advanced is not None:
try:
if isinstance(raw_advanced, ClaudeCodeAdvancedConfig):
advanced_config = raw_advanced
elif isinstance(raw_advanced, dict):
advanced_config = ClaudeCodeAdvancedConfig.model_validate(raw_advanced)
else:
logger.warning(
"Claude Code advanced config 类型无效: {},已忽略",
type(raw_advanced).__name__,
)
except Exception as exc:
logger.warning("Claude Code advanced config 解析失败,已忽略: {}", str(exc))
max_sessions = advanced_config.max_sessions if advanced_config else None
idle_timeout_minutes = (
advanced_config.session_idle_timeout_minutes
if advanced_config and advanced_config.session_idle_timeout_minutes is not None
else 5
)
enable_tls_fingerprint = (
bool(advanced_config.enable_tls_fingerprint) if advanced_config else False
)
session_id_masking_enabled = (
bool(advanced_config.session_id_masking_enabled) if advanced_config else False
)
# Parse pool config (None = non-pool provider, keep as None for semantic consistency)
pool_cfg = parse_pool_config(provider_config_dict)
return ClaudeCodeRequestContext(
is_stream=bool(is_stream),
scope_key=f"key:{normalized_key_id}" if normalized_key_id else None,
key_id=normalized_key_id,
max_sessions=max_sessions,
session_idle_timeout_minutes=idle_timeout_minutes,
enable_tls_fingerprint=enable_tls_fingerprint,
session_id_masking_enabled=session_id_masking_enabled,
provider_id=str(provider_id or "").strip() or None,
pool_config=pool_cfg,
)
def resolve_claude_code_tls_profile(
ctx: ClaudeCodeRequestContext | None,
) -> str | None:
if ctx and ctx.enable_tls_fingerprint:
return TLS_PROFILE_CLAUDE_CODE
return None
def build_and_set_claude_code_request_context(
*,
provider_config: Any,
key_id: str | None,
is_stream: bool,
provider_id: str | None = None,
) -> tuple[ClaudeCodeRequestContext, str | None]:
"""构建并写入 Claude Code 上下文,同时返回对应 TLS profile。"""
ctx = build_claude_code_request_context(
provider_config=provider_config,
key_id=key_id,
is_stream=is_stream,
provider_id=provider_id,
)
set_claude_code_request_context(ctx)
return ctx, resolve_claude_code_tls_profile(ctx)
__all__ = [
"build_and_set_claude_code_request_context",
"build_claude_code_request_context",
"ClaudeCodeRequestContext",
"get_claude_code_request_context",
"resolve_claude_code_tls_profile",
"set_claude_code_request_context",
]

View File

@@ -0,0 +1,526 @@
"""Claude Code upstream envelope hooks."""
from __future__ import annotations
import threading
import time
import uuid
from dataclasses import replace
from typing import Any
from src.clients.redis_client import get_redis_client, get_redis_client_sync
from src.config.settings import config
from src.core.exceptions import ConcurrencyLimitError
from src.core.logger import logger
from src.services.provider.adapters.claude_code.constants import (
BETA_CONTEXT_1M,
CLAUDE_CODE_DEFAULT_HEADERS,
CLAUDE_CODE_REQUIRED_BETA_TOKENS,
DEFAULT_ACCEPT,
DEFAULT_ANTHROPIC_VERSION,
SESSION_ID_MASKING_TTL_SECONDS,
STREAM_HELPER_METHOD,
)
from src.services.provider.adapters.claude_code.context import (
ClaudeCodeRequestContext,
get_claude_code_request_context,
set_claude_code_request_context,
)
_SESSION_MARKER = "_session_"
_DUMMY_THINKING_SIGNATURE = "skip_thought_signature_validator"
_session_runtime_lock = threading.Lock()
# key: scope_key -> {session_id -> last_seen_monotonic}
_active_sessions: dict[str, dict[str, float]] = {}
# key: scope_key -> (masked_session_uuid, expire_at_monotonic)
_masked_sessions: dict[str, tuple[str, float]] = {}
_REDIS_SESSION_KEY_PREFIX = "claude_code:sessions"
_REDIS_SESSION_RESERVE_LUA = """
local key = KEYS[1]
local sid = ARGV[1]
local now = tonumber(ARGV[2])
local expire_before = tonumber(ARGV[3])
local max_sessions = tonumber(ARGV[4])
local ttl_seconds = tonumber(ARGV[5])
redis.call("ZREMRANGEBYSCORE", key, "-inf", expire_before)
local exists = redis.call("ZSCORE", key, sid)
if exists then
redis.call("ZADD", key, now, sid)
redis.call("EXPIRE", key, ttl_seconds)
return {1, redis.call("ZCARD", key)}
end
local active = redis.call("ZCARD", key)
if active >= max_sessions then
return {0, active}
end
redis.call("ZADD", key, now, sid)
redis.call("EXPIRE", key, ttl_seconds)
return {1, active + 1}
"""
def merge_anthropic_beta_tokens(
incoming: str | None,
*,
required: tuple[str, ...] = CLAUDE_CODE_REQUIRED_BETA_TOKENS,
) -> str:
"""Merge required beta tokens and incoming anthropic-beta with deduplication."""
seen: set[str] = set()
merged: list[str] = []
def _append(token: str) -> None:
token = token.strip()
if not token or token in seen:
return
seen.add(token)
merged.append(token)
for token in required:
_append(token)
for token in str(incoming or "").split(","):
_append(token)
return ",".join(merged)
def _parse_stream_flag(raw_stream: Any) -> bool:
if isinstance(raw_stream, bool):
return raw_stream
return str(raw_stream).strip().lower() in {"1", "true", "yes", "on"}
def _get_metadata_user_id(request_body: dict[str, Any]) -> str | None:
metadata = request_body.get("metadata")
if not isinstance(metadata, dict):
return None
user_id = metadata.get("user_id")
if not isinstance(user_id, str):
return None
text = user_id.strip()
return text or None
def _set_metadata_user_id(request_body: dict[str, Any], user_id: str) -> None:
metadata = request_body.get("metadata")
if not isinstance(metadata, dict):
metadata = {}
request_body["metadata"] = metadata
metadata["user_id"] = user_id
def _extract_session_id(user_id: str) -> str | None:
idx = user_id.rfind(_SESSION_MARKER)
if idx == -1:
return None
session_id = user_id[idx + len(_SESSION_MARKER) :].strip()
return session_id or None
def _is_thinking_enabled(request_body: dict[str, Any]) -> bool:
thinking = request_body.get("thinking")
if not isinstance(thinking, dict):
return False
thinking_type = str(thinking.get("type") or "").strip().lower()
return thinking_type in {"enabled", "adaptive"}
def _sanitize_thinking_blocks(request_body: dict[str, Any]) -> None:
"""过滤可能导致 Claude Code 400 的无效 thinking 块。"""
messages = request_body.get("messages")
if not isinstance(messages, list) or not messages:
return
thinking_enabled = _is_thinking_enabled(request_body)
filtered_messages = 0
filtered_blocks = 0
for message in messages:
if not isinstance(message, dict):
continue
role = str(message.get("role") or "")
content = message.get("content")
if not isinstance(content, list):
continue
new_content: list[Any] = []
modified = False
for block in content:
if not isinstance(block, dict):
new_content.append(block)
continue
block_type = str(block.get("type") or "")
if block_type in {"thinking", "redacted_thinking"}:
keep = False
# 仅保留 assistant 且带真实 signature 的 thinking 块。
if thinking_enabled and role == "assistant":
signature = str(block.get("signature") or "").strip()
keep = bool(signature and signature != _DUMMY_THINKING_SIGNATURE)
if keep:
new_content.append(block)
else:
modified = True
filtered_blocks += 1
continue
# 兼容无 type 但带 thinking 字段的历史块,直接移除。
if not block_type and "thinking" in block:
modified = True
filtered_blocks += 1
continue
new_content.append(block)
if modified:
message["content"] = new_content
filtered_messages += 1
if filtered_blocks:
logger.info(
"Claude Code thinking 预过滤: messages={}, blocks={}, thinking_enabled={}",
filtered_messages,
filtered_blocks,
thinking_enabled,
)
def _get_or_create_masked_session(scope_key: str) -> str:
now = time.monotonic()
with _session_runtime_lock:
existing = _masked_sessions.get(scope_key)
if existing and existing[1] > now:
masked_session_id = existing[0]
else:
masked_session_id = str(uuid.uuid4())
_masked_sessions[scope_key] = (
masked_session_id,
now + SESSION_ID_MASKING_TTL_SECONDS,
)
return masked_session_id
def _apply_session_id_masking(request_body: dict[str, Any], *, scope_key: str) -> None:
user_id = _get_metadata_user_id(request_body)
if not user_id:
return
idx = user_id.rfind(_SESSION_MARKER)
if idx == -1:
return
masked_session_id = _get_or_create_masked_session(scope_key)
_set_metadata_user_id(
request_body,
user_id[: idx + len(_SESSION_MARKER)] + masked_session_id,
)
def _register_or_reject_session(
*,
scope_key: str,
session_id: str,
max_sessions: int,
idle_timeout_minutes: int,
) -> tuple[bool, int]:
now = time.monotonic()
idle_seconds = max(60, int(idle_timeout_minutes * 60))
with _session_runtime_lock:
bucket = _active_sessions.setdefault(scope_key, {})
# 先清理过期会话,避免误判占用。
expired = [sid for sid, last_seen in bucket.items() if now - last_seen > idle_seconds]
for sid in expired:
bucket.pop(sid, None)
if session_id in bucket:
bucket[session_id] = now
return True, len(bucket)
if len(bucket) >= max_sessions:
return False, len(bucket)
bucket[session_id] = now
return True, len(bucket)
def _build_session_limit_error(
*,
max_sessions: int,
active_count: int,
key_id: str | None,
) -> ConcurrencyLimitError:
return ConcurrencyLimitError(
message=(f"Claude Code 活跃会话数已达上限({max_sessions})。当前活跃会话: {active_count}"),
key_id=key_id,
)
def _redis_session_key(scope_key: str) -> str:
return f"{_REDIS_SESSION_KEY_PREFIX}:{scope_key}"
def _parse_redis_session_result(raw: Any) -> tuple[bool, int] | None:
if not isinstance(raw, (list, tuple)) or len(raw) < 2:
return None
try:
allowed = int(raw[0]) == 1
active_count = int(raw[1])
except Exception:
return None
return allowed, active_count
def _enforce_session_controls(
request_body: dict[str, Any],
ctx: ClaudeCodeRequestContext,
*,
enforce_max_sessions: bool = True,
) -> None:
"""同步执行会话限制 + masking。
当 ``enforce_max_sessions=False``(将由 ``enforce_distributed_session_controls``
异步接管)时仅做 maskingmasking 始终在会话限制检查之后执行,避免
用被伪装后的 session_id 做计数。
"""
scope_key = str(ctx.scope_key or "").strip()
if not scope_key:
return
# 先基于真实 session_id 做会话限制检查。
if enforce_max_sessions and ctx.max_sessions and ctx.max_sessions > 0:
user_id = _get_metadata_user_id(request_body)
if user_id:
session_id = _extract_session_id(user_id) or user_id
allowed, active_count = _register_or_reject_session(
scope_key=scope_key,
session_id=session_id,
max_sessions=ctx.max_sessions,
idle_timeout_minutes=ctx.session_idle_timeout_minutes,
)
if not allowed:
raise _build_session_limit_error(
max_sessions=ctx.max_sessions,
active_count=active_count,
key_id=ctx.key_id,
)
if not enforce_max_sessions:
# 分布式模式下 masking 延迟到 enforce_distributed_session_controls 中执行,
# 避免 wrap_request 提前改写 user_id 导致分布式检查拿到伪装后的 session_id。
return
# 仅在本地模式下立即 masking。
if ctx.session_id_masking_enabled:
_apply_session_id_masking(request_body, scope_key=scope_key)
def _is_distributed_session_control_available() -> bool:
try:
return get_redis_client_sync() is not None
except Exception:
return False
async def enforce_distributed_session_controls(
request_body: dict[str, Any],
ctx: ClaudeCodeRequestContext | None,
) -> None:
"""异步执行会话限制 + masking。
优先使用 Redis多实例共享Redis 不可用时回退到进程内计数。
masking 在会话限制检查通过后执行,确保计数使用真实 session_id。
"""
if ctx is None:
return
scope_key = str(ctx.scope_key or "").strip()
if not scope_key:
# 即使无 scope_key 也无法做 masking需要 scope_key 作为 key直接返回。
return
if not ctx.max_sessions or ctx.max_sessions <= 0:
# 无会话限制,仅做 masking。
if ctx.session_id_masking_enabled:
_apply_session_id_masking(request_body, scope_key=scope_key)
return
# 基于真实 user_id 提取 session_id 做限制检查。
user_id = _get_metadata_user_id(request_body)
if not user_id:
return
session_id = _extract_session_id(user_id) or user_id
idle_seconds = max(60, int(ctx.session_idle_timeout_minutes * 60))
redis_ttl = idle_seconds + 300
now = int(time.time())
expire_before = now - idle_seconds
redis_client = await get_redis_client(require_redis=False)
if redis_client is not None:
try:
raw_result = await redis_client.eval(
_REDIS_SESSION_RESERVE_LUA,
1,
_redis_session_key(scope_key),
session_id,
str(now),
str(expire_before),
str(ctx.max_sessions),
str(redis_ttl),
)
parsed = _parse_redis_session_result(raw_result)
if parsed is None:
raise ValueError(f"invalid redis eval result: {raw_result!r}")
allowed, active_count = parsed
if not allowed:
raise _build_session_limit_error(
max_sessions=ctx.max_sessions,
active_count=active_count,
key_id=ctx.key_id,
)
# 会话限制通过后再做 masking。
if ctx.session_id_masking_enabled:
_apply_session_id_masking(request_body, scope_key=scope_key)
return
except ConcurrencyLimitError:
raise
except Exception as exc:
logger.warning("Claude Code 分布式会话控制失败,回退本地计数: {}", str(exc))
allowed, active_count = _register_or_reject_session(
scope_key=scope_key,
session_id=session_id,
max_sessions=ctx.max_sessions,
idle_timeout_minutes=ctx.session_idle_timeout_minutes,
)
if not allowed:
raise _build_session_limit_error(
max_sessions=ctx.max_sessions,
active_count=active_count,
key_id=ctx.key_id,
)
# 会话限制通过后再做 masking。
if ctx.session_id_masking_enabled:
_apply_session_id_masking(request_body, scope_key=scope_key)
class ClaudeCodeEnvelope:
"""Provider envelope hooks for Claude Code OAuth upstream."""
name = "claude:cli"
def extra_headers(self) -> dict[str, str] | None:
ctx = get_claude_code_request_context()
is_stream = bool(ctx.is_stream) if ctx else False
headers = dict(CLAUDE_CODE_DEFAULT_HEADERS)
headers["Accept"] = DEFAULT_ACCEPT
headers["anthropic-version"] = DEFAULT_ANTHROPIC_VERSION
headers["anthropic-beta"] = merge_anthropic_beta_tokens(None)
if is_stream:
headers["x-stainless-helper-method"] = STREAM_HELPER_METHOD
ua = str(getattr(config, "internal_user_agent_claude_cli", "") or "").strip()
if ua:
headers["User-Agent"] = ua
return headers
def wrap_request(
self,
request_body: dict[str, Any],
*,
model: str, # noqa: ARG002
url_model: str | None,
decrypted_auth_config: dict[str, Any] | None, # noqa: ARG002
) -> tuple[dict[str, Any], str | None]:
raw_stream = request_body.get("stream", False)
is_stream = _parse_stream_flag(raw_stream)
ctx = get_claude_code_request_context()
if ctx is None:
ctx = ClaudeCodeRequestContext()
# Extract session_uuid from metadata.user_id for pool sticky session.
session_uuid: str | None = None
user_id = _get_metadata_user_id(request_body)
if user_id:
session_uuid = _extract_session_id(user_id)
ctx = replace(ctx, is_stream=is_stream, session_uuid=session_uuid)
set_claude_code_request_context(ctx)
_sanitize_thinking_blocks(request_body)
_enforce_session_controls(
request_body,
ctx,
enforce_max_sessions=not _is_distributed_session_control_available(),
)
return request_body, url_model
def unwrap_response(self, data: Any) -> Any:
return data
def postprocess_unwrapped_response(self, *, model: str, data: Any) -> None: # noqa: ARG002
return
def capture_selected_base_url(self) -> str | None:
return None
def on_http_status(self, *, base_url: str | None, status_code: int) -> None: # noqa: ARG002
return
def on_connection_error(self, *, base_url: str | None, exc: Exception) -> None: # noqa: ARG002
return
def force_stream_rewrite(self) -> bool:
return False
# ------------------------------------------------------------------
# Optional lifecycle hooks
# ------------------------------------------------------------------
def prepare_context(
self,
*,
provider_config: Any,
key_id: str,
is_stream: bool,
provider_id: str | None = None,
) -> str | None:
from src.services.provider.adapters.claude_code.context import (
build_and_set_claude_code_request_context,
)
_ctx, tls_profile = build_and_set_claude_code_request_context(
provider_config=provider_config,
key_id=key_id,
is_stream=is_stream,
provider_id=provider_id,
)
return tls_profile
async def post_wrap_request(self, request_body: dict[str, Any]) -> None:
await enforce_distributed_session_controls(
request_body,
get_claude_code_request_context(),
)
def excluded_beta_tokens(self) -> frozenset[str]:
return frozenset({BETA_CONTEXT_1M})
claude_code_envelope = ClaudeCodeEnvelope()
__all__ = [
"ClaudeCodeEnvelope",
"claude_code_envelope",
"enforce_distributed_session_controls",
"merge_anthropic_beta_tokens",
]

View File

@@ -0,0 +1,62 @@
"""Claude Code provider plugin."""
from __future__ import annotations
from typing import Any
from urllib.parse import urlencode
from src.services.provider.adapters.claude_code.constants import CLAUDE_MESSAGES_PATH
from src.services.provider.preset_models import create_preset_models_fetcher
fetch_models_claude_code = create_preset_models_fetcher("claude_code")
def build_claude_code_url(
endpoint: Any,
*,
is_stream: bool,
effective_query_params: dict[str, Any],
) -> str:
"""Build Claude Code upstream URL and avoid duplicate /v1/messages suffix."""
_ = is_stream
base = str(getattr(endpoint, "base_url", "") or "").rstrip("/")
if base.endswith(CLAUDE_MESSAGES_PATH) or base.endswith("/messages"):
url = base
elif base.endswith("/v1"):
url = f"{base}/messages"
else:
url = f"{base}{CLAUDE_MESSAGES_PATH}"
if effective_query_params:
query_string = urlencode(effective_query_params, doseq=True)
if query_string:
url = f"{url}?{query_string}"
return url
def register_all() -> None:
"""Register Claude Code hooks into shared registries."""
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
from src.services.provider.adapters.claude_code.envelope import claude_code_envelope
from src.services.provider.envelope import register_envelope
from src.services.provider.transport import register_transport_hook
register_envelope("claude_code", "claude:cli", claude_code_envelope)
register_envelope("claude_code", "", claude_code_envelope)
register_transport_hook("claude_code", "claude:cli", build_claude_code_url)
UpstreamModelsFetcherRegistry.register(
provider_types=["claude_code"],
fetcher=fetch_models_claude_code,
)
from src.services.provider.adapters.claude_code.pool_hook import claude_code_pool_hook
from src.services.provider.pool.hooks import register_pool_hook
register_pool_hook("claude_code", claude_code_pool_hook)
__all__ = ["build_claude_code_url", "fetch_models_claude_code", "register_all"]

View File

@@ -236,6 +236,7 @@ async def get_provider_auth(
key: "ProviderAPIKey",
*,
force_refresh: bool = False,
refresh_skew: int | None = None,
) -> ProviderAuthInfo | None:
"""
获取 Provider 的认证信息
@@ -261,7 +262,12 @@ async def get_provider_auth(
if auth_type == "oauth":
# OAuth token 保存在 key.api_key加密refresh_token/expires_at 等在 auth_config加密 JSON中。
# 在请求前做一次懒刷新:接近过期时刷新 access_token并用 Redis lock 避免并发风暴。
encrypted_auth_config = getattr(key, "auth_config", None)
# 先解密 auth_config -- 下游 build_provider_url 等依赖 decrypted_auth_config
# 中的 provider_type / project_id / region 等元数据,即使 access_token 命中缓存
# 也不能跳过。
if encrypted_auth_config:
try:
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
@@ -271,16 +277,51 @@ async def get_provider_auth(
else:
token_meta = {}
decrypted_auth_config: dict[str, Any] | None = (
token_meta if isinstance(token_meta, dict) and token_meta else None
)
# 快路径:查 Redis token 缓存,命中则跳过 refresh 和 api_key 解密。
# 注意token_meta/decrypted_auth_config 已在上方解密,此处只是跳过后续刷新逻辑。
if not force_refresh and encrypted_auth_config:
try:
from src.services.provider.pool.oauth_cache import get_cached_token
_cached = await get_cached_token(str(key.id))
if _cached:
return ProviderAuthInfo(
auth_header="Authorization",
auth_value=f"Bearer {_cached}",
decrypted_auth_config=decrypted_auth_config,
)
except Exception:
logger.debug("OAuth token cache lookup failed for key {}", str(key.id)[:8])
expires_at = token_meta.get("expires_at")
refresh_token = token_meta.get("refresh_token")
provider_type = str(token_meta.get("provider_type") or "")
cached_access_token = str(token_meta.get("access_token") or "").strip()
# 120s skew (or force refresh when upstream returns 401)
# Refresh skew: providers with pool config use configurable
# proactive_refresh_seconds (default 180 s), others use 120 s.
# Prefer the caller-supplied value to avoid ORM lazy-load on key.provider.
_refresh_skew = refresh_skew if refresh_skew is not None else 120
if refresh_skew is None:
try:
from src.services.provider.pool.config import parse_pool_config
provider_obj = getattr(key, "provider", None)
pcfg = getattr(provider_obj, "config", None) if provider_obj else None
pool_cfg = parse_pool_config(pcfg) if pcfg else None
if pool_cfg is not None:
_refresh_skew = pool_cfg.proactive_refresh_seconds
except Exception:
pass
should_refresh = False
try:
if expires_at is not None:
should_refresh = int(time.time()) >= int(expires_at) - 120
should_refresh = int(time.time()) >= int(expires_at) - _refresh_skew
except Exception:
should_refresh = False
@@ -294,6 +335,7 @@ async def get_provider_auth(
elif crypto_service.decrypt(key.api_key) == "__placeholder__":
should_refresh = True
_refreshed = False
if should_refresh and refresh_token and provider_type:
try:
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
@@ -313,6 +355,7 @@ async def get_provider_auth(
token_meta = await _refresh_generic_oauth_token(
key, endpoint, template, provider_type, refresh_token, token_meta
)
_refreshed = True
finally:
if got_lock:
await _release_refresh_lock(redis, key.id)
@@ -328,7 +371,20 @@ async def get_provider_auth(
else:
effective_token = crypto_service.decrypt(key.api_key)
decrypted_auth_config: dict[str, Any] | None = None
# 刷新成功后写入 Redis token 缓存(所有 OAuth key 均可受益)
if _refreshed and effective_token:
try:
from src.services.provider.pool.oauth_cache import cache_token
new_expires_at = token_meta.get("expires_at")
if new_expires_at is not None:
remaining = int(new_expires_at) - int(time.time())
if remaining > 0:
await cache_token(str(key.id), effective_token, remaining)
except Exception:
logger.debug("OAuth token cache write failed for key {}", str(key.id)[:8])
# 刷新可能更新了 token_meta同步 decrypted_auth_config
if isinstance(token_meta, dict) and token_meta:
decrypted_auth_config = token_meta

View File

@@ -50,6 +50,39 @@ class ProviderEnvelope(Protocol):
def force_stream_rewrite(self) -> bool:
"""Whether streaming should always go through the rewrite/conversion path."""
# ------------------------------------------------------------------
# Optional lifecycle hooks (checked via hasattr before calling)
# ------------------------------------------------------------------
def prepare_context(
self,
*,
provider_config: Any,
key_id: str,
is_stream: bool,
provider_id: str | None = None,
) -> str | None:
"""Pre-wrap hook: build provider-specific request context.
Called before wrap_request(). Returns tls_profile (or None).
Implementations typically set contextvars that wrap_request()
and extra_headers() will read.
"""
async def post_wrap_request(self, request_body: dict[str, Any]) -> None:
"""Post-wrap hook: async processing after wrap_request().
Called after wrap_request() completes. Use for async operations
like distributed session control that cannot run in sync wrap_request().
"""
def excluded_beta_tokens(self) -> frozenset[str]:
"""Beta tokens to strip from the merged anthropic-beta header.
Called by the request builder after merging envelope extra_headers
with client original headers. Return an empty frozenset to keep all.
"""
# ---------------------------------------------------------------------------
# Envelope Registry
@@ -119,10 +152,14 @@ def ensure_providers_bootstrapped() -> None:
from src.services.provider.adapters.antigravity.plugin import (
register_all as _reg_antigravity,
)
from src.services.provider.adapters.claude_code.plugin import (
register_all as _reg_claude_code,
)
from src.services.provider.adapters.codex.plugin import register_all as _reg_codex
from src.services.provider.adapters.kiro.plugin import register_all as _reg_kiro
_reg_antigravity()
_reg_claude_code()
_reg_codex()
_reg_kiro()

View File

@@ -60,6 +60,39 @@ PRESET_MODELS: dict[str, list[dict[str, Any]]] = {
"display_name": "Claude Haiku 4.5",
},
],
# Claude Code (Claude CLI OAuth 反代)
"claude_code": [
{
"id": "claude-opus-4-5-20251101",
"object": "model",
"owned_by": "anthropic",
"display_name": "Claude Opus 4.5",
},
{
"id": "claude-opus-4-6",
"object": "model",
"owned_by": "anthropic",
"display_name": "Claude Opus 4.6",
},
{
"id": "claude-sonnet-4-6",
"object": "model",
"owned_by": "anthropic",
"display_name": "Claude Sonnet 4.6",
},
{
"id": "claude-sonnet-4-5-20250929",
"object": "model",
"owned_by": "anthropic",
"display_name": "Claude Sonnet 4.5",
},
{
"id": "claude-haiku-4-5-20251001",
"object": "model",
"owned_by": "anthropic",
"display_name": "Claude Haiku 4.5",
},
],
# Codex (OpenAI CLI 反代)
"codex": [
{

View File

@@ -361,6 +361,7 @@ class CacheAwareScheduler:
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
request_body: dict | None = None,
) -> tuple[list[ProviderCandidate], str, int]:
"""
预先获取所有可用的 Provider/Endpoint/Key 组合
@@ -519,6 +520,7 @@ class CacheAwareScheduler:
is_stream=is_stream,
capability_requirements=capability_requirements,
global_conversion_enabled=global_conversion_enabled,
request_body=request_body,
)
# 3. 应用优先级模式排序 + 调度模式排序

View File

@@ -35,12 +35,20 @@ from src.services.scheduling.utils import release_db_connection_before_await
if TYPE_CHECKING:
from src.models.database import GlobalModel
from src.services.provider.pool.config import PoolConfig
from src.services.scheduling.protocols import CandidateSorterProtocol
from src.services.scheduling.schemas import ProviderCandidate
from src.services.cache.model_cache import ModelCacheService
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
"""Return parsed PoolConfig if the provider has pool enabled, else None."""
from src.services.provider.pool.config import PoolConfig, parse_pool_config
return parse_pool_config(getattr(provider, "config", None))
def _sort_endpoints_by_family_priority(
eps: Sequence[ProviderEndpoint],
) -> list[ProviderEndpoint]:
@@ -359,6 +367,7 @@ class CandidateBuilder:
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
global_conversion_enabled: bool = True,
request_body: dict | None = None,
) -> "list[ProviderCandidate]":
"""
构建候选列表
@@ -417,6 +426,8 @@ class CandidateBuilder:
] = {}
exact_candidates: list[ProviderCandidate] = []
convertible_candidates: list[ProviderCandidate] = []
pool_has_usable = False
pool_cfg = _get_pool_config(provider)
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
@@ -553,21 +564,34 @@ class CandidateBuilder:
if not active_keys:
continue
# 检查是否所有 Key 都是 TTL=0轮换模式
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
if use_random and len(active_keys) > 1:
logger.debug(
" Provider {} 启用 Key 轮换模式 (endpoint_format={}, {} keys)",
provider.name,
endpoint_format_str,
len(active_keys),
# --- Pool branch: select a single key internally ------
if pool_cfg is not None:
selected_key = await self._pool_select_key(
db, provider, pool_cfg, active_keys, request_body
)
if selected_key is None:
logger.debug(
"Pool[{}]: no schedulable key for endpoint {}",
str(provider.id)[:8],
endpoint_format_str,
)
continue
keys_to_check: list[ProviderAPIKey] = [selected_key]
else:
# --- Normal branch: check all keys ----
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
if use_random and len(active_keys) > 1:
logger.debug(
" Provider {} 启用 Key 轮换模式 (endpoint_format={}, {} keys)",
provider.name,
endpoint_format_str,
len(active_keys),
)
keys_to_check = self._sorter.shuffle_keys_by_internal_priority(
active_keys, affinity_key, use_random
)
keys = self._sorter.shuffle_keys_by_internal_priority(
active_keys, affinity_key, use_random
)
for key in keys:
for key in keys_to_check:
# Key 级别检查(健康度/熔断按 provider_format bucket
# 传入 provider_model_names 作为 candidate_models
# 用于检查 Key 的 allowed_models 是否支持 Provider 定义的模型名称
@@ -600,6 +624,13 @@ class CandidateBuilder:
else:
exact_candidates.append(candidate)
if is_available:
pool_has_usable = True
# Pool mode: stop after the first endpoint that produced a usable candidate.
if pool_cfg is not None and pool_has_usable:
break
candidates.extend(exact_candidates)
candidates.extend(convertible_candidates)
@@ -608,3 +639,22 @@ class CandidateBuilder:
candidates = candidates[:max_candidates]
return candidates
async def _pool_select_key(
self,
db: Session,
provider: Provider,
pool_cfg: "PoolConfig",
active_keys: list[ProviderAPIKey],
request_body: dict | None,
) -> ProviderAPIKey | None:
"""Select a single key via pool scheduling (sticky -> cooldown/cost -> LRU)."""
from src.services.provider.pool.hooks import get_pool_hook
from src.services.provider.pool.manager import PoolManager
provider_type = str(getattr(provider, "provider_type", "") or "")
hook = get_pool_hook(provider_type)
session_uuid = hook.extract_session_uuid(request_body) if hook and request_body else None
mgr = PoolManager(str(provider.id), pool_cfg)
release_db_connection_before_await(db)
return await mgr.select_key(session_uuid, active_keys)

View File

@@ -33,6 +33,9 @@ class ExecutionResult:
attempt_count: int = 0
request_candidate_id: str | None = None
# pool scheduling summary (populated when pool mode is active)
pool_summary: dict[str, Any] | None = None
# failure
error_type: str | None = None
error_message: str | None = None

View File

@@ -186,6 +186,151 @@ class TaskService:
request_body=request_body,
)
@staticmethod
def _extract_session_uuid(
provider_type: str, request_body: dict[str, Any] | None
) -> str | None:
"""Extract a session UUID from the request body (provider-type aware)."""
if not isinstance(request_body, dict):
return None
from src.services.provider.pool.hooks import get_pool_hook
hook = get_pool_hook(provider_type)
if hook is not None:
return hook.extract_session_uuid(request_body)
return None
@staticmethod
async def _apply_pool_reorder(
candidates: list[Any],
request_body: dict[str, Any] | None,
) -> tuple[list[Any], list[Any]]:
"""Apply Account Pool reordering when applicable.
Groups candidates by provider_id and applies pool reordering
independently per provider, then reassembles in original group order.
Non-pool providers are left in their original order.
Returns:
Tuple of (reordered_candidates, pool_traces) where pool_traces
is a list of :class:`PoolSchedulingTrace` objects (one per
pooled provider group, may be empty).
"""
if not candidates:
return candidates, []
pool_traces: list[Any] = []
try:
from collections import OrderedDict
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.manager import PoolManager
# Group candidates by provider_id while preserving order.
groups: OrderedDict[str, list[Any]] = OrderedDict()
for c in candidates:
pid = str(getattr(c.provider, "id", "") or "")
groups.setdefault(pid, []).append(c)
result: list[Any] = []
for pid, group in groups.items():
provider = group[0].provider
pool_cfg = parse_pool_config(getattr(provider, "config", None))
if pool_cfg is None or not pid:
result.extend(group)
continue
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
mgr = PoolManager(pid, pool_cfg)
reordered = await mgr.reorder_candidates(session_uuid, group)
result.extend(reordered)
# Extract trace attached by PoolManager.reorder_candidates
if reordered:
trace = getattr(reordered[0], "_pool_scheduling_trace", None)
if trace is not None:
pool_traces.append(trace)
return result, pool_traces
except Exception:
from src.core.logger import logger
logger.opt(exception=True).debug("Pool reorder failed, using original order")
return candidates, []
@staticmethod
async def _pool_on_success(
candidate: Any,
request_body: dict[str, Any] | None,
) -> None:
"""Notify the pool manager about a successful request (sticky + LRU)."""
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.manager import PoolManager
provider = candidate.provider
provider_config = getattr(provider, "config", None)
pool_cfg = parse_pool_config(provider_config)
if pool_cfg is None:
return
provider_id = str(getattr(provider, "id", "") or "")
key_id = str(getattr(candidate.key, "id", "") or "")
if not provider_id or not key_id:
return
provider_type = str(getattr(provider, "provider_type", "") or "")
session_uuid = TaskService._extract_session_uuid(provider_type, request_body)
mgr = PoolManager(provider_id, pool_cfg)
await mgr.on_request_success(
session_uuid=session_uuid,
key_id=key_id,
)
except Exception:
logger.opt(exception=True).debug("Pool on_request_success failed (non-blocking)")
@staticmethod
async def _pool_on_error(
provider: Any,
key: Any,
status_code: int,
cause: Any,
) -> None:
"""Notify the pool manager about an upstream error (health policy)."""
try:
from src.services.provider.pool.config import parse_pool_config
from src.services.provider.pool.health_policy import apply_health_policy
pool_cfg = parse_pool_config(getattr(provider, "config", None))
if pool_cfg is None:
return
error_text = ""
resp_headers: dict[str, str] = {}
if getattr(cause, "response", None) is not None:
try:
error_text = (cause.response.text or "")[:4000]
except Exception:
pass
try:
resp_headers = dict(cause.response.headers)
except Exception:
pass
await apply_health_policy(
provider_id=str(provider.id),
key_id=str(key.id),
status_code=status_code,
error_body=error_text,
response_headers=resp_headers,
config=pool_cfg,
)
except Exception:
pass
async def _execute_sync_unified(
self,
*,
@@ -287,6 +432,12 @@ class TaskService:
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body=request_body,
)
# Account Pool: reorder candidates for claude_code providers.
all_candidates, pool_traces = await self._apply_pool_reorder(
all_candidates, request_body=request_body
)
candidate_record_map = candidate_resolver.create_candidate_records(
@@ -350,6 +501,9 @@ class TaskService:
)
_ = (attempt_id, _provider_name, _provider_id, _endpoint_id, _key_id)
# Account Pool: on success, update sticky binding + LRU.
await self._pool_on_success(candidate, request_body)
if is_stream:
return AttemptResult(
kind=AttemptKind.STREAM,
@@ -441,6 +595,16 @@ class TaskService:
)
if result.success:
# Build pool scheduling summary from traces collected during reorder.
if pool_traces and result.key_id:
try:
for pt in pool_traces:
summary = pt.build_summary(result.key_id)
if summary:
result.pool_summary = summary
break
except Exception:
pass
return result
self._raise_all_failed_exception(
@@ -933,6 +1097,9 @@ class TaskService:
attempt=attempt,
)
# Account Pool: apply health policy (cooldown/disable).
await self._pool_on_error(provider, key, status_code, cause)
converted_error = extra_data.get("converted_error")
serializable_extra_data = {
k: v for k, v in extra_data.items() if k != "converted_error"

View File

@@ -38,6 +38,7 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
"billing_snapshot",
"billing_updated_at",
"perf",
"pool_summary",
"_metadata_truncated",
}
)

View File

@@ -7,14 +7,23 @@ import ssl
from loguru import logger
try:
import certifi
_SSL_CONTEXT = ssl.create_default_context(cafile=certifi.where())
except ImportError:
def _create_default_ssl_context() -> ssl.SSLContext:
try:
import certifi
return ssl.create_default_context(cafile=certifi.where())
except ImportError:
return ssl.create_default_context()
try:
_SSL_CONTEXT = _create_default_ssl_context()
except Exception:
_SSL_CONTEXT = ssl.create_default_context()
_PROXY_SSL_CONTEXT: ssl.SSLContext | None = None
_PROFILE_SSL_CONTEXTS: dict[str, ssl.SSLContext] = {}
def get_ssl_context() -> ssl.SSLContext:
@@ -55,3 +64,54 @@ def get_proxy_ssl_context(expected_fingerprint: str | None = None) -> ssl.SSLCon
_PROXY_SSL_CONTEXT = ctx
# TODO: 实现基于 expected_fingerprint 的证书指纹校验
return _PROXY_SSL_CONTEXT
def _build_claude_code_ssl_context() -> ssl.SSLContext:
"""构建 Claude Code best-effort TLS 配置。
说明Python/OpenSSL 无法完整模拟 Node.js ClientHello。
这里仅做可控项的尽力对齐ALPN/TLS 版本/常见 cipher 偏好)。
"""
ctx = _create_default_ssl_context()
ctx.minimum_version = ssl.TLSVersion.TLSv1_2
try:
ctx.maximum_version = ssl.TLSVersion.TLSv1_3
except Exception:
pass
try:
ctx.set_alpn_protocols(["h2", "http/1.1"])
except Exception:
pass
try:
ctx.set_ciphers(
"ECDHE-ECDSA-AES128-GCM-SHA256:"
"ECDHE-RSA-AES128-GCM-SHA256:"
"ECDHE-ECDSA-AES256-GCM-SHA384:"
"ECDHE-RSA-AES256-GCM-SHA384:"
"ECDHE-ECDSA-CHACHA20-POLY1305:"
"ECDHE-RSA-CHACHA20-POLY1305"
)
except Exception:
pass
return ctx
def get_ssl_context_for_profile(tls_profile: str | None = None) -> ssl.SSLContext:
"""按 profile 返回 SSL 上下文。"""
profile = str(tls_profile or "").strip().lower()
if not profile:
return get_ssl_context()
if profile in _PROFILE_SSL_CONTEXTS:
logger.debug("复用 TLS profile SSL context: {}", profile)
return _PROFILE_SSL_CONTEXTS[profile]
if profile == "claude_code_nodejs":
logger.info("启用 TLS profile: {}best-effort", profile)
ctx = _build_claude_code_ssl_context()
else:
logger.warning("未知 TLS profile: {},回退默认 SSL context", profile)
ctx = get_ssl_context()
_PROFILE_SSL_CONTEXTS[profile] = ctx
return ctx