mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/ - 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层 - 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构 - 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations) - 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image - 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
"""Claude Code provider adapter."""
|
||||
|
||||
__all__ = []
|
||||
@@ -0,0 +1,96 @@
|
||||
"""Claude Code CLI client restriction.
|
||||
|
||||
When cli_only_enabled is True, only requests from genuine Claude Code CLI
|
||||
clients are allowed. Non-CLI traffic receives a 403 response.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextvars
|
||||
from typing import Any
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
# Contextvar to carry the original request headers into the envelope layer.
|
||||
_original_request_headers: contextvars.ContextVar[dict[str, str] | None] = contextvars.ContextVar(
|
||||
"claude_code_original_request_headers",
|
||||
default=None,
|
||||
)
|
||||
|
||||
|
||||
def set_original_request_headers(headers: dict[str, str] | None) -> None:
|
||||
_original_request_headers.set(headers)
|
||||
|
||||
|
||||
def get_original_request_headers() -> dict[str, str] | None:
|
||||
return _original_request_headers.get()
|
||||
|
||||
|
||||
# Known Claude Code CLI User-Agent patterns.
|
||||
_CLI_USER_AGENT_PATTERNS = (
|
||||
"claude-code",
|
||||
"claudecode",
|
||||
"claude_code",
|
||||
)
|
||||
|
||||
# Known originator / x-app values indicating CLI usage.
|
||||
_CLI_APP_VALUES = {"cli"}
|
||||
|
||||
|
||||
def is_claude_code_client(headers: dict[str, Any]) -> bool:
|
||||
"""Detect whether the request originates from a Claude Code CLI client.
|
||||
|
||||
Detection signals (any match is sufficient):
|
||||
1. User-Agent contains a known Claude Code CLI pattern
|
||||
2. x-app header equals "cli"
|
||||
"""
|
||||
# Normalize header keys to lowercase for case-insensitive matching.
|
||||
lower_headers = {k.lower(): v for k, v in headers.items()}
|
||||
|
||||
# Check User-Agent
|
||||
ua = str(lower_headers.get("user-agent", "")).lower()
|
||||
for pattern in _CLI_USER_AGENT_PATTERNS:
|
||||
if pattern in ua:
|
||||
return True
|
||||
|
||||
# Check x-app header
|
||||
x_app = str(lower_headers.get("x-app", "")).strip().lower()
|
||||
if x_app in _CLI_APP_VALUES:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def enforce_cli_only(cli_only_enabled: bool) -> None:
|
||||
"""Enforce CLI-only restriction if enabled.
|
||||
|
||||
Reads original request headers from contextvar, checks whether the
|
||||
client is a Claude Code CLI, and raises HTTPException(403) if not.
|
||||
"""
|
||||
if not cli_only_enabled:
|
||||
return
|
||||
|
||||
headers = get_original_request_headers()
|
||||
if headers is None:
|
||||
# No headers available; skip enforcement (should not happen in normal flow).
|
||||
logger.debug("CLI-only check skipped: no request headers in context")
|
||||
return
|
||||
|
||||
if is_claude_code_client(headers):
|
||||
return
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
logger.info("CLI-only restriction: rejected non-CLI client")
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="This endpoint only accepts requests from Claude Code CLI clients.",
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"enforce_cli_only",
|
||||
"get_original_request_headers",
|
||||
"is_claude_code_client",
|
||||
"set_original_request_headers",
|
||||
]
|
||||
@@ -0,0 +1,49 @@
|
||||
"""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 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",
|
||||
]
|
||||
@@ -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
|
||||
|
||||
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
|
||||
session_id_masking_enabled: bool = False
|
||||
cache_ttl_override_enabled: bool = False
|
||||
cache_ttl_override_target: str = "ephemeral"
|
||||
cli_only_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
|
||||
)
|
||||
session_id_masking_enabled = (
|
||||
bool(advanced_config.session_id_masking_enabled) if advanced_config else False
|
||||
)
|
||||
cache_ttl_override_enabled = (
|
||||
bool(advanced_config.cache_ttl_override_enabled) if advanced_config else False
|
||||
)
|
||||
cache_ttl_override_target = (
|
||||
str(advanced_config.cache_ttl_override_target or "ephemeral")
|
||||
if advanced_config
|
||||
else "ephemeral"
|
||||
)
|
||||
cli_only_enabled = bool(advanced_config.cli_only_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,
|
||||
session_id_masking_enabled=session_id_masking_enabled,
|
||||
cache_ttl_override_enabled=cache_ttl_override_enabled,
|
||||
cache_ttl_override_target=cache_ttl_override_target,
|
||||
cli_only_enabled=cli_only_enabled,
|
||||
provider_id=str(provider_id or "").strip() or None,
|
||||
pool_config=pool_cfg,
|
||||
)
|
||||
|
||||
|
||||
def build_and_set_claude_code_request_context(
|
||||
*,
|
||||
provider_config: Any,
|
||||
key_id: str | None,
|
||||
is_stream: bool,
|
||||
provider_id: str | None = None,
|
||||
) -> ClaudeCodeRequestContext:
|
||||
"""构建并写入 Claude Code 上下文。"""
|
||||
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
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_and_set_claude_code_request_context",
|
||||
"build_claude_code_request_context",
|
||||
"ClaudeCodeRequestContext",
|
||||
"get_claude_code_request_context",
|
||||
"set_claude_code_request_context",
|
||||
]
|
||||
@@ -0,0 +1,677 @@
|
||||
"""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,
|
||||
)
|
||||
from src.services.provider.request_context import get_current_fingerprint
|
||||
|
||||
_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]] = {}
|
||||
# 上限保护:scope_key 总数超过此值时触发全局清理
|
||||
_MAX_SCOPE_KEYS = 5000
|
||||
# 全局清理间隔(monotonic 秒),避免高频请求时每次都做全局扫描
|
||||
_LAST_GLOBAL_CLEANUP: float = 0.0
|
||||
_GLOBAL_CLEANUP_INTERVAL = 300.0 # 5 分钟
|
||||
_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,
|
||||
)
|
||||
|
||||
|
||||
# -- Cache TTL Override -------------------------------------------------------
|
||||
|
||||
_VALID_CACHE_TTL_TARGETS = {"ephemeral", "1h"}
|
||||
|
||||
|
||||
def _override_cache_control_in_blocks(blocks: list[Any], target: str) -> int:
|
||||
"""Override cache_control TTL in a list of content blocks. Returns count of overrides.
|
||||
|
||||
According to Anthropic API docs, cache_control format is:
|
||||
{"type": "ephemeral", "ttl": "5m" | "1h"}
|
||||
``type`` is always "ephemeral"; the ``ttl`` field controls the actual duration.
|
||||
When ttl is absent, the default is 5m (ephemeral).
|
||||
"""
|
||||
count = 0
|
||||
for block in blocks:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
cc = block.get("cache_control")
|
||||
if not isinstance(cc, dict):
|
||||
continue
|
||||
# Ensure type is always "ephemeral"
|
||||
if cc.get("type") != "ephemeral":
|
||||
cc["type"] = "ephemeral"
|
||||
if target == "ephemeral":
|
||||
# Target is 5m (default) -- remove explicit ttl so it falls back to default
|
||||
if "ttl" in cc:
|
||||
del cc["ttl"]
|
||||
count += 1
|
||||
else:
|
||||
# Target is "1h" -- set ttl explicitly
|
||||
if cc.get("ttl") != target:
|
||||
cc["ttl"] = target
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
def _apply_cache_ttl_override(request_body: dict[str, Any], target: str) -> None:
|
||||
"""Force all cache_control entries to use a unified TTL type.
|
||||
|
||||
Prevents multi-user behavioral fingerprinting when sharing an OAuth account.
|
||||
"""
|
||||
if target not in _VALID_CACHE_TTL_TARGETS:
|
||||
return
|
||||
|
||||
overridden = 0
|
||||
|
||||
# system prompt (can be string or list of blocks)
|
||||
system = request_body.get("system")
|
||||
if isinstance(system, list):
|
||||
overridden += _override_cache_control_in_blocks(system, target)
|
||||
|
||||
# messages
|
||||
messages = request_body.get("messages")
|
||||
if isinstance(messages, list):
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
content = msg.get("content")
|
||||
if isinstance(content, list):
|
||||
overridden += _override_cache_control_in_blocks(content, target)
|
||||
|
||||
# tools
|
||||
tools = request_body.get("tools")
|
||||
if isinstance(tools, list):
|
||||
overridden += _override_cache_control_in_blocks(tools, target)
|
||||
|
||||
if overridden:
|
||||
logger.debug("Cache TTL override: {} block(s) -> {}", overridden, target)
|
||||
|
||||
|
||||
def _cleanup_stale_scope_keys(now: float, idle_seconds: int) -> None:
|
||||
"""清理空 bucket 和过期的 _masked_sessions 条目(调用方须持有锁)。"""
|
||||
global _LAST_GLOBAL_CLEANUP
|
||||
|
||||
if now - _LAST_GLOBAL_CLEANUP < _GLOBAL_CLEANUP_INTERVAL:
|
||||
return
|
||||
_LAST_GLOBAL_CLEANUP = now
|
||||
|
||||
# 清理所有 scope_key 下的过期 session,并删除空 bucket
|
||||
stale_keys = []
|
||||
for sk, bucket in _active_sessions.items():
|
||||
expired = [sid for sid, ts in bucket.items() if now - ts > idle_seconds]
|
||||
for sid in expired:
|
||||
bucket.pop(sid, None)
|
||||
if not bucket:
|
||||
stale_keys.append(sk)
|
||||
for sk in stale_keys:
|
||||
_active_sessions.pop(sk, None)
|
||||
|
||||
# 清理过期的 masked session 条目
|
||||
expired_masked = [sk for sk, (_, exp) in _masked_sessions.items() if exp <= now]
|
||||
for sk in expired_masked:
|
||||
_masked_sessions.pop(sk, None)
|
||||
|
||||
total = len(_active_sessions) + len(_masked_sessions)
|
||||
if total > 0 and (stale_keys or expired_masked):
|
||||
logger.debug(
|
||||
"Session 全局清理: 移除 {} 个 active scope + {} 个 masked scope, 剩余 {}",
|
||||
len(stale_keys),
|
||||
len(expired_masked),
|
||||
total,
|
||||
)
|
||||
|
||||
|
||||
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:
|
||||
# scope_key 总数超限或达到清理间隔时,执行全局清理
|
||||
if (
|
||||
len(_active_sessions) + len(_masked_sessions) > _MAX_SCOPE_KEYS
|
||||
or now - _LAST_GLOBAL_CLEANUP >= _GLOBAL_CLEANUP_INTERVAL
|
||||
):
|
||||
_cleanup_stale_scope_keys(now, idle_seconds)
|
||||
|
||||
bucket = _active_sessions.setdefault(scope_key, {})
|
||||
|
||||
# 先清理当前 bucket 的过期会话,避免误判占用。
|
||||
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
|
||||
fp = get_current_fingerprint()
|
||||
|
||||
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
|
||||
|
||||
if fp:
|
||||
headers["X-Stainless-Package-Version"] = fp.stainless_package_version
|
||||
headers["X-Stainless-OS"] = fp.stainless_os
|
||||
headers["X-Stainless-Arch"] = fp.stainless_arch
|
||||
headers["X-Stainless-Runtime-Version"] = fp.stainless_runtime_version
|
||||
headers["X-Stainless-Timeout"] = fp.stainless_timeout
|
||||
headers["User-Agent"] = fp.user_agent
|
||||
else:
|
||||
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()
|
||||
|
||||
# CLI-only restriction: reject non-CLI clients early.
|
||||
if ctx.cli_only_enabled:
|
||||
from src.services.provider.adapters.claude_code.client_restriction import (
|
||||
enforce_cli_only,
|
||||
)
|
||||
|
||||
enforce_cli_only(ctx.cli_only_enabled)
|
||||
|
||||
# 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)
|
||||
|
||||
# Cache TTL override: unify cache_control types to prevent behavioral fingerprinting.
|
||||
if ctx.cache_ttl_override_enabled:
|
||||
_apply_cache_ttl_override(request_body, ctx.cache_ttl_override_target)
|
||||
|
||||
_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,
|
||||
user_api_key_id: str | None = None, # noqa: ARG002
|
||||
is_stream: bool,
|
||||
provider_id: str | None = None,
|
||||
key: Any = None,
|
||||
) -> str | None:
|
||||
from src.services.provider.adapters.claude_code.context import (
|
||||
build_and_set_claude_code_request_context,
|
||||
)
|
||||
|
||||
# 在 envelope 层设置指纹 context var(仅 Claude Code 需要指纹注入)
|
||||
if key is not None:
|
||||
from src.services.provider.fingerprint import ensure_key_fingerprint
|
||||
from src.services.provider.request_context import set_current_fingerprint
|
||||
|
||||
set_current_fingerprint(ensure_key_fingerprint(key, persist_if_missing=True))
|
||||
|
||||
build_and_set_claude_code_request_context(
|
||||
provider_config=provider_config,
|
||||
key_id=key_id,
|
||||
is_stream=is_stream,
|
||||
provider_id=provider_id,
|
||||
)
|
||||
fp = get_current_fingerprint()
|
||||
if fp:
|
||||
return fp.impersonate
|
||||
return None
|
||||
|
||||
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",
|
||||
]
|
||||
@@ -0,0 +1,63 @@
|
||||
"""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],
|
||||
**_kwargs: 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"]
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Backward compat shim -- canonical definitions moved to src.services.provider.pool.config."""
|
||||
|
||||
from src.services.provider.pool.config import * # noqa: F401,F403
|
||||
from src.services.provider.pool.config import PoolConfig, UnschedulableRule, parse_pool_config
|
||||
|
||||
__all__ = ["PoolConfig", "UnschedulableRule", "parse_pool_config"]
|
||||
@@ -0,0 +1,10 @@
|
||||
"""Backward compat shim -- canonical definitions moved to src.services.provider.pool.cost_tracker."""
|
||||
|
||||
from src.services.provider.pool.cost_tracker import ( # noqa: F401
|
||||
get_window_usage,
|
||||
is_approaching_limit,
|
||||
is_at_limit,
|
||||
record_usage,
|
||||
)
|
||||
|
||||
__all__ = ["record_usage", "get_window_usage", "is_at_limit", "is_approaching_limit"]
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Backward compat shim -- canonical definitions moved to src.services.provider.pool.health_policy."""
|
||||
|
||||
from src.services.provider.pool.health_policy import * # noqa: F401,F403
|
||||
from src.services.provider.pool.health_policy import apply_health_policy # noqa: F811
|
||||
|
||||
__all__ = ["apply_health_policy"]
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Claude Code pool scheduling hook."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
class ClaudeCodePoolHook:
|
||||
"""Pool scheduling hook for Claude Code providers.
|
||||
|
||||
Extracts the session UUID from ``metadata.user_id`` which follows the
|
||||
pattern ``<user>_session_<uuid>``.
|
||||
"""
|
||||
|
||||
name = "claude_code"
|
||||
|
||||
def extract_session_uuid(self, request_body: dict[str, Any]) -> str | None:
|
||||
metadata = request_body.get("metadata")
|
||||
if isinstance(metadata, dict):
|
||||
user_id = metadata.get("user_id")
|
||||
if isinstance(user_id, str) and "_session_" in user_id:
|
||||
idx = user_id.rfind("_session_")
|
||||
return user_id[idx + len("_session_") :].strip() or None
|
||||
return None
|
||||
|
||||
|
||||
claude_code_pool_hook = ClaudeCodePoolHook()
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Backward compat shim -- canonical definitions moved to src.services.provider.pool.manager."""
|
||||
|
||||
from src.services.provider.pool.manager import * # noqa: F401,F403
|
||||
from src.services.provider.pool.manager import PoolManager
|
||||
|
||||
ClaudeCodePoolManager = PoolManager # noqa: F811
|
||||
|
||||
__all__ = ["PoolManager", "ClaudeCodePoolManager"]
|
||||
@@ -0,0 +1,9 @@
|
||||
"""Backward compat shim -- canonical definitions moved to src.services.provider.pool.oauth_cache."""
|
||||
|
||||
from src.services.provider.pool.oauth_cache import ( # noqa: F401
|
||||
cache_token,
|
||||
get_cached_token,
|
||||
invalidate_token,
|
||||
)
|
||||
|
||||
__all__ = ["get_cached_token", "cache_token", "invalidate_token"]
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Backward compat shim -- canonical definitions moved to src.services.provider.pool.redis_ops."""
|
||||
|
||||
from src.services.provider.pool.redis_ops import * # noqa: F401,F403
|
||||
from src.services.provider.pool.redis_ops import (
|
||||
add_cost_entry,
|
||||
batch_get_cooldowns,
|
||||
cache_oauth_token,
|
||||
clear_cooldown,
|
||||
clear_cost,
|
||||
delete_sticky_binding,
|
||||
get_cached_oauth_token,
|
||||
get_cooldown,
|
||||
get_cooldown_ttl,
|
||||
get_cost_window_total,
|
||||
get_key_sticky_count,
|
||||
get_lru_scores,
|
||||
get_sticky_binding,
|
||||
get_sticky_session_count,
|
||||
invalidate_oauth_token_cache,
|
||||
set_cooldown,
|
||||
set_sticky_binding,
|
||||
touch_lru,
|
||||
)
|
||||
Reference in New Issue
Block a user