mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
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:
@@ -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():
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
3
src/services/provider/adapters/claude_code/__init__.py
Normal file
3
src/services/provider/adapters/claude_code/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
"""Claude Code provider adapter."""
|
||||
|
||||
__all__ = []
|
||||
52
src/services/provider/adapters/claude_code/constants.py
Normal file
52
src/services/provider/adapters/claude_code/constants.py
Normal 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",
|
||||
]
|
||||
141
src/services/provider/adapters/claude_code/context.py
Normal file
141
src/services/provider/adapters/claude_code/context.py
Normal 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",
|
||||
]
|
||||
526
src/services/provider/adapters/claude_code/envelope.py
Normal file
526
src/services/provider/adapters/claude_code/envelope.py
Normal 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``
|
||||
异步接管)时仅做 masking;masking 始终在会话限制检查之后执行,避免
|
||||
用被伪装后的 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",
|
||||
]
|
||||
62
src/services/provider/adapters/claude_code/plugin.py
Normal file
62
src/services/provider/adapters/claude_code/plugin.py
Normal 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"]
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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": [
|
||||
{
|
||||
|
||||
@@ -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. 应用优先级模式排序 + 调度模式排序
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -38,6 +38,7 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
|
||||
"billing_snapshot",
|
||||
"billing_updated_at",
|
||||
"perf",
|
||||
"pool_summary",
|
||||
"_metadata_truncated",
|
||||
}
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user