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

@@ -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"])