refactor: 代理节点架构重构与功能增强

aether-proxy:
- 重构 main.rs,拆分为 app/state/hardware/net 模块
- setup.rs 拆分为 setup/tui.rs + setup/service.rs,支持 systemd 服务管理子命令
- 新增 delegate 端点,支持后端通过代理节点转发请求而非传统 CONNECT 代理
- 注册时上报硬件信息(CPU/内存/fd_limit)和估算最大并发数
- 心跳上报活跃连接数,支持远程下发 node_name 配置
- HTTP 转发时剥离 X-Forwarded-* 等敏感头部
- 切换到 rustls-tls,降低日志级别减少噪音

后端:
- 从 http_client.py 提取代理相关逻辑至 proxy_node/resolver.py
- 从 routes.py 提取业务逻辑至 proxy_node/service.py
- handler 支持 delegate 模式(通过代理节点 HTTP 端点转发而非 CONNECT 隧道)
- ProxyNode 模型新增 hardware_info 和 estimated_max_concurrency 字段

前端:
- 新增 HardwareTooltip 组件展示节点硬件信息
- 远程配置支持下发 node_name
This commit is contained in:
fawney19
2026-02-08 13:33:08 +08:00
parent 254d30d32d
commit 519ad67eb1
44 changed files with 3339 additions and 1709 deletions

View File

@@ -6,9 +6,7 @@
from __future__ import annotations
import ipaddress
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from fastapi import APIRouter, Depends, Query, Request
@@ -18,51 +16,17 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.exceptions import InvalidRequestException
from src.database import get_db
from src.models.database import ProxyNode, ProxyNodeStatus, SystemConfig
from src.services.proxy_node.service import ProxyNodeService, node_to_dict
router = APIRouter(prefix="/api/admin/proxy-nodes", tags=["Admin - Proxy Nodes"])
pipeline = ApiRequestPipeline()
def _mask_password(password: str | None) -> str | None:
"""脱敏密码仅显示前2位和后2位长度不足 8 时全部遮蔽)"""
if not password:
return None
if len(password) < 8:
return "****"
return password[:2] + "****" + password[-2:]
def _node_to_dict(node: ProxyNode) -> dict[str, Any]:
d = {
"id": node.id,
"name": node.name,
"ip": node.ip,
"port": node.port,
"region": node.region,
"status": node.status.value if node.status else None,
"is_manual": bool(node.is_manual),
"registered_by": node.registered_by,
"last_heartbeat_at": node.last_heartbeat_at,
"heartbeat_interval": node.heartbeat_interval,
"active_connections": node.active_connections,
"total_requests": node.total_requests,
"avg_latency_ms": node.avg_latency_ms,
"tls_enabled": bool(node.tls_enabled),
"tls_cert_fingerprint": node.tls_cert_fingerprint,
"remote_config": node.remote_config,
"config_version": node.config_version,
"created_at": node.created_at,
"updated_at": node.updated_at,
}
# 手动节点附带代理配置(密码脱敏)
if node.is_manual:
d["proxy_url"] = node.proxy_url
d["proxy_username"] = node.proxy_username
d["proxy_password"] = _mask_password(node.proxy_password)
return d
# ---------------------------------------------------------------------------
# Pydantic 请求模型
# ---------------------------------------------------------------------------
class ProxyNodeRegisterRequest(BaseModel):
@@ -83,6 +47,10 @@ class ProxyNodeRegisterRequest(BaseModel):
None, max_length=128, description="TLS 证书 SHA-256 指纹"
)
# 硬件信息
hardware_info: dict | None = Field(None, description="硬件信息 JSON")
estimated_max_concurrency: int | None = Field(None, ge=0, description="估算最大并发连接数")
@field_validator("ip")
@classmethod
def validate_ip(cls, v: str) -> str:
@@ -110,6 +78,7 @@ class ProxyNodeUnregisterRequest(BaseModel):
class ProxyNodeRemoteConfigRequest(BaseModel):
"""管理端远程配置 — 通过心跳下发给 aether-proxy"""
node_name: str | None = Field(None, min_length=1, max_length=100, description="节点名称")
allowed_ports: list[int] | None = Field(None, description="允许代理的目标端口")
log_level: str | None = Field(None, description="日志级别 (trace/debug/info/warn/error)")
heartbeat_interval: int | None = Field(None, ge=5, le=600, description="心跳间隔(秒)")
@@ -186,6 +155,11 @@ class ManualProxyNodeUpdateRequest(BaseModel):
return v
# ---------------------------------------------------------------------------
# 路由端点
# ---------------------------------------------------------------------------
@router.post("/register")
async def register_proxy_node(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = AdminRegisterProxyNodeAdapter()
@@ -250,6 +224,11 @@ async def update_proxy_node_config(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
# ---------------------------------------------------------------------------
# 辅助函数
# ---------------------------------------------------------------------------
def _format_validation_error(exc: ValidationError) -> str:
parts: list[str] = []
for err in exc.errors():
@@ -259,6 +238,11 @@ def _format_validation_error(exc: ValidationError) -> str:
return "; ".join(parts) or "输入验证失败"
# ---------------------------------------------------------------------------
# Adapter 实现
# ---------------------------------------------------------------------------
@dataclass
class AdminRegisterProxyNodeAdapter(AdminApiAdapter):
name: str = "admin_register_proxy_node"
@@ -270,50 +254,22 @@ class AdminRegisterProxyNodeAdapter(AdminApiAdapter):
except ValidationError as exc:
raise InvalidRequestException("输入验证失败: " + _format_validation_error(exc))
now = datetime.now(timezone.utc)
node = (
context.db.query(ProxyNode)
.filter(ProxyNode.ip == req.ip, ProxyNode.port == req.port)
.first()
node = ProxyNodeService.register_node(
context.db,
name=req.name,
ip=req.ip,
port=req.port,
region=req.region,
heartbeat_interval=req.heartbeat_interval,
tls_enabled=req.tls_enabled,
tls_cert_fingerprint=req.tls_cert_fingerprint,
hardware_info=req.hardware_info,
estimated_max_concurrency=req.estimated_max_concurrency,
active_connections=req.active_connections,
total_requests=req.total_requests,
avg_latency_ms=req.avg_latency_ms,
registered_by=context.user.id if context.user else None,
)
if node:
node.name = req.name
node.region = req.region
node.status = ProxyNodeStatus.ONLINE
node.last_heartbeat_at = now
node.heartbeat_interval = req.heartbeat_interval
node.tls_enabled = req.tls_enabled
node.tls_cert_fingerprint = req.tls_cert_fingerprint
if req.active_connections is not None:
node.active_connections = req.active_connections
if req.total_requests is not None:
node.total_requests = req.total_requests
if req.avg_latency_ms is not None:
node.avg_latency_ms = req.avg_latency_ms
else:
node = ProxyNode(
id=str(uuid.uuid4()),
name=req.name,
ip=req.ip,
port=req.port,
region=req.region,
status=ProxyNodeStatus.ONLINE,
registered_by=context.user.id if context.user else None,
last_heartbeat_at=now,
heartbeat_interval=req.heartbeat_interval,
active_connections=req.active_connections or 0,
total_requests=req.total_requests or 0,
avg_latency_ms=req.avg_latency_ms,
tls_enabled=req.tls_enabled,
tls_cert_fingerprint=req.tls_cert_fingerprint,
created_at=now,
updated_at=now,
)
context.db.add(node)
context.db.commit()
context.db.refresh(node)
context.add_audit_metadata(
action="proxy_node_register",
@@ -322,7 +278,7 @@ class AdminRegisterProxyNodeAdapter(AdminApiAdapter):
proxy_node_port=node.port,
)
return {"node_id": node.id, "node": _node_to_dict(node)}
return {"node_id": node.id, "node": node_to_dict(node)}
@dataclass
@@ -336,31 +292,21 @@ class AdminHeartbeatProxyNodeAdapter(AdminApiAdapter):
except ValidationError as exc:
raise InvalidRequestException("输入验证失败: " + _format_validation_error(exc))
node = context.db.query(ProxyNode).filter(ProxyNode.id == req.node_id).first()
if not node:
raise NotFoundException(f"ProxyNode {req.node_id} 不存在", "proxy_node")
now = datetime.now(timezone.utc)
node.status = ProxyNodeStatus.ONLINE
node.last_heartbeat_at = now
if req.heartbeat_interval is not None:
node.heartbeat_interval = req.heartbeat_interval
if req.active_connections is not None:
node.active_connections = req.active_connections
if req.total_requests is not None:
node.total_requests = req.total_requests
if req.avg_latency_ms is not None:
node.avg_latency_ms = req.avg_latency_ms
context.db.commit()
context.db.refresh(node)
node = ProxyNodeService.heartbeat(
context.db,
node_id=req.node_id,
heartbeat_interval=req.heartbeat_interval,
active_connections=req.active_connections,
total_requests=req.total_requests,
avg_latency_ms=req.avg_latency_ms,
)
context.add_audit_metadata(
action="proxy_node_heartbeat",
proxy_node_id=node.id,
)
return {"message": "heartbeat ok", "node": _node_to_dict(node)}
return {"message": "heartbeat ok", "node": node_to_dict(node)}
@dataclass
@@ -374,13 +320,7 @@ class AdminUnregisterProxyNodeAdapter(AdminApiAdapter):
except ValidationError as exc:
raise InvalidRequestException("输入验证失败: " + _format_validation_error(exc))
node = context.db.query(ProxyNode).filter(ProxyNode.id == req.node_id).first()
if not node:
raise NotFoundException(f"ProxyNode {req.node_id} 不存在", "proxy_node")
node.status = ProxyNodeStatus.OFFLINE
node.updated_at = datetime.now(timezone.utc)
context.db.commit()
node = ProxyNodeService.unregister_node(context.db, node_id=req.node_id)
context.add_audit_metadata(
action="proxy_node_unregister",
@@ -398,20 +338,11 @@ class AdminListProxyNodesAdapter(AdminApiAdapter):
limit: int = 100
async def handle(self, context: ApiRequestContext) -> Any:
query = context.db.query(ProxyNode)
if self.status:
normalized = self.status.strip().lower()
allowed = {"online", "unhealthy", "offline"}
if normalized not in allowed:
raise InvalidRequestException(f"status 必须是以下之一: {sorted(allowed)}", "status")
query = query.filter(ProxyNode.status == ProxyNodeStatus(normalized))
total = query.count()
nodes = (
query.order_by(ProxyNode.updated_at.desc()).offset(self.skip).limit(self.limit).all()
nodes, total = ProxyNodeService.list_nodes(
context.db, status=self.status, skip=self.skip, limit=self.limit
)
return {
"items": [_node_to_dict(n) for n in nodes],
"items": [node_to_dict(n) for n in nodes],
"total": total,
"skip": self.skip,
"limit": self.limit,
@@ -424,93 +355,21 @@ class AdminDeleteProxyNodeAdapter(AdminApiAdapter):
node_id: str = ""
async def handle(self, context: ApiRequestContext) -> Any:
node = context.db.query(ProxyNode).filter(ProxyNode.id == self.node_id).first()
if not node:
raise NotFoundException(f"ProxyNode {self.node_id} 不存在", "proxy_node")
result = ProxyNodeService.delete_node(context.db, node_id=self.node_id)
context.add_audit_metadata(
action="proxy_node_delete",
proxy_node_id=node.id,
proxy_node_ip=node.ip,
proxy_node_port=node.port,
proxy_node_id=self.node_id,
**result.get("node_info", {}),
)
# 若该节点是系统默认代理,自动清除引用
was_system_proxy = False
sys_cfg = (
context.db.query(SystemConfig)
.filter(SystemConfig.key == "system_proxy_node_id")
.first()
)
if sys_cfg and sys_cfg.value == self.node_id:
sys_cfg.value = None
was_system_proxy = True
context.db.delete(node)
context.db.commit()
if was_system_proxy:
from src.clients.http_client import invalidate_system_proxy_cache
invalidate_system_proxy_cache()
msg = "deleted"
if was_system_proxy:
msg = "deleted, system default proxy cleared"
return {"message": msg, "node_id": self.node_id, "cleared_system_proxy": was_system_proxy}
def _parse_host_port(proxy_url: str) -> tuple[str, int]:
"""从代理 URL 中解析 host 和 port含协议前缀避免唯一约束冲突"""
from urllib.parse import urlparse
parsed = urlparse(proxy_url)
host = parsed.hostname or "manual"
default_ports = {"https": 443, "socks5": 1080}
port = parsed.port or default_ports.get((parsed.scheme or "").lower(), 80)
# 添加协议前缀区分同 host:port 不同协议的场景
scheme = (parsed.scheme or "http").lower()
if scheme != "http":
host = f"{scheme}://{host}"
return host, port
def _sanitize_proxy_error(err: Exception) -> str:
"""去除异常消息中可能包含的代理 URL 凭据(如 HMAC 签名)"""
import re
return re.sub(r"://[^@/]+@", "://***@", str(err))
def _build_test_proxy_url(node: ProxyNode) -> str:
"""为测试连通性构建代理 URL无需节点在线"""
if node.is_manual:
proxy_url = node.proxy_url
if not proxy_url:
raise InvalidRequestException("手动节点缺少 proxy_url")
if node.proxy_username:
from urllib.parse import quote, urlparse
parsed = urlparse(proxy_url)
encoded_username = quote(node.proxy_username, safe="")
encoded_password = quote(node.proxy_password, safe="") if node.proxy_password else ""
host_part = parsed.hostname or "localhost"
if parsed.port:
host_part = f"{host_part}:{parsed.port}"
if encoded_password:
proxy_url = f"{parsed.scheme}://{encoded_username}:{encoded_password}@{host_part}"
else:
proxy_url = f"{parsed.scheme}://{encoded_username}@{host_part}"
if parsed.path:
proxy_url += parsed.path
return proxy_url
else:
# aether-proxy: 使用 HMAC 认证构建代理 URL
from src.clients.http_client import _build_hmac_proxy_url
return _build_hmac_proxy_url(
node.ip, node.port, node.id, tls_enabled=bool(node.tls_enabled)
)
was_system_proxy = result["cleared_system_proxy"]
msg = "deleted, system default proxy cleared" if was_system_proxy else "deleted"
return {
"message": msg,
"node_id": self.node_id,
"cleared_system_proxy": was_system_proxy,
}
@dataclass
@@ -524,48 +383,22 @@ class AdminCreateManualProxyNodeAdapter(AdminApiAdapter):
except ValidationError as exc:
raise InvalidRequestException("输入验证失败: " + _format_validation_error(exc))
host, port = _parse_host_port(req.proxy_url)
now = datetime.now(timezone.utc)
node = ProxyNode(
id=str(uuid.uuid4()),
node = ProxyNodeService.create_manual_node(
context.db,
name=req.name,
ip=host,
port=port,
region=req.region,
is_manual=True,
proxy_url=req.proxy_url,
proxy_username=req.username,
proxy_password=req.password,
status=ProxyNodeStatus.ONLINE,
username=req.username,
password=req.password,
region=req.region,
registered_by=context.user.id if context.user else None,
last_heartbeat_at=None,
heartbeat_interval=0,
active_connections=0,
total_requests=0,
avg_latency_ms=None,
created_at=now,
updated_at=now,
)
# 检查是否已存在同地址的节点
existing = (
context.db.query(ProxyNode).filter(ProxyNode.ip == host, ProxyNode.port == port).first()
)
if existing:
raise InvalidRequestException(
f"已存在相同地址的代理节点: {existing.name} ({existing.ip}:{existing.port})"
)
context.db.add(node)
context.db.commit()
context.db.refresh(node)
context.add_audit_metadata(
action="proxy_node_manual_create",
proxy_node_id=node.id,
)
return {"node_id": node.id, "node": _node_to_dict(node)}
return {"node_id": node.id, "node": node_to_dict(node)}
@dataclass
@@ -574,53 +407,28 @@ class AdminUpdateManualProxyNodeAdapter(AdminApiAdapter):
node_id: str = ""
async def handle(self, context: ApiRequestContext) -> Any:
node = context.db.query(ProxyNode).filter(ProxyNode.id == self.node_id).first()
if not node:
raise NotFoundException(f"ProxyNode {self.node_id} 不存在", "proxy_node")
if not node.is_manual:
raise InvalidRequestException("只能编辑手动添加的代理节点")
payload = context.ensure_json_body()
try:
req = ManualProxyNodeUpdateRequest.model_validate(payload)
except ValidationError as exc:
raise InvalidRequestException("输入验证失败: " + _format_validation_error(exc))
if req.name is not None:
node.name = req.name
if req.proxy_url is not None:
node.proxy_url = req.proxy_url
host, port = _parse_host_port(req.proxy_url)
# 检查新地址是否与其他节点冲突
existing = (
context.db.query(ProxyNode)
.filter(ProxyNode.ip == host, ProxyNode.port == port, ProxyNode.id != node.id)
.first()
)
if existing:
raise InvalidRequestException(
f"已存在相同地址的代理节点: {existing.name} ({existing.ip}:{existing.port})"
)
node.ip = host
node.port = port
if req.username is not None:
node.proxy_username = req.username
# password: None=不发送(保留原值), ""=清空, 非空=更新
if req.password is not None:
node.proxy_password = req.password or None
if req.region is not None:
node.region = req.region
node.updated_at = datetime.now(timezone.utc)
context.db.commit()
context.db.refresh(node)
node = ProxyNodeService.update_manual_node(
context.db,
node_id=self.node_id,
name=req.name,
proxy_url=req.proxy_url,
username=req.username,
password=req.password,
region=req.region,
)
context.add_audit_metadata(
action="proxy_node_manual_update",
proxy_node_id=node.id,
)
return {"node_id": node.id, "node": _node_to_dict(node)}
return {"node_id": node.id, "node": node_to_dict(node)}
@dataclass
@@ -631,81 +439,7 @@ class AdminTestProxyNodeAdapter(AdminApiAdapter):
node_id: str = ""
async def handle(self, context: ApiRequestContext) -> Any:
import time as _time
import httpx
node = context.db.query(ProxyNode).filter(ProxyNode.id == self.node_id).first()
if not node:
raise NotFoundException(f"ProxyNode {self.node_id} 不存在", "proxy_node")
# 构建代理 URL
try:
proxy_url = _build_test_proxy_url(node)
except Exception as exc:
return {"success": False, "latency_ms": None, "exit_ip": None, "error": str(exc)}
test_url = "https://1.1.1.1/cdn-cgi/trace"
start = _time.monotonic()
# TLS 代理需要 proxy_ssl_context
from src.clients.http_client import _make_proxy_param
proxy_param = _make_proxy_param(proxy_url)
try:
async with httpx.AsyncClient(
proxy=proxy_param,
timeout=httpx.Timeout(15.0, connect=10.0),
) as client:
response = await client.get(test_url)
elapsed_ms = round((_time.monotonic() - start) * 1000, 1)
exit_ip = None
if response.status_code == 200:
for line in response.text.splitlines():
if line.startswith("ip="):
exit_ip = line.split("=", 1)[1].strip()
break
return {
"success": True,
"latency_ms": elapsed_ms,
"exit_ip": exit_ip,
"error": None,
}
except httpx.ProxyError as exc:
elapsed_ms = round((_time.monotonic() - start) * 1000, 1)
return {
"success": False,
"latency_ms": elapsed_ms,
"exit_ip": None,
"error": f"代理连接失败: {_sanitize_proxy_error(exc)}",
}
except httpx.ConnectError as exc:
elapsed_ms = round((_time.monotonic() - start) * 1000, 1)
return {
"success": False,
"latency_ms": elapsed_ms,
"exit_ip": None,
"error": f"连接失败: {_sanitize_proxy_error(exc)}",
}
except httpx.TimeoutException:
elapsed_ms = round((_time.monotonic() - start) * 1000, 1)
return {
"success": False,
"latency_ms": elapsed_ms,
"exit_ip": None,
"error": "连接超时15秒",
}
except Exception as exc:
elapsed_ms = round((_time.monotonic() - start) * 1000, 1)
return {
"success": False,
"latency_ms": elapsed_ms,
"exit_ip": None,
"error": _sanitize_proxy_error(exc),
}
return await ProxyNodeService.test_node(context.db, node_id=self.node_id)
@dataclass
@@ -716,12 +450,6 @@ class AdminUpdateProxyNodeConfigAdapter(AdminApiAdapter):
node_id: str = ""
async def handle(self, context: ApiRequestContext) -> Any:
node = context.db.query(ProxyNode).filter(ProxyNode.id == self.node_id).first()
if not node:
raise NotFoundException(f"ProxyNode {self.node_id} 不存在", "proxy_node")
if node.is_manual:
raise InvalidRequestException("手动节点不支持远程配置下发")
payload = context.ensure_json_body()
try:
req = ProxyNodeRemoteConfigRequest.model_validate(payload)
@@ -729,27 +457,21 @@ class AdminUpdateProxyNodeConfigAdapter(AdminApiAdapter):
raise InvalidRequestException("输入验证失败: " + _format_validation_error(exc))
# Build config dict with only the supplied fields
config: dict[str, Any] = {}
config_updates: dict[str, Any] = {}
if req.node_name is not None:
config_updates["node_name"] = req.node_name
if req.allowed_ports is not None:
config["allowed_ports"] = req.allowed_ports
config_updates["allowed_ports"] = req.allowed_ports
if req.log_level is not None:
config["log_level"] = req.log_level
config_updates["log_level"] = req.log_level
if req.heartbeat_interval is not None:
config["heartbeat_interval"] = req.heartbeat_interval
config_updates["heartbeat_interval"] = req.heartbeat_interval
if req.timestamp_tolerance is not None:
config["timestamp_tolerance"] = req.timestamp_tolerance
config_updates["timestamp_tolerance"] = req.timestamp_tolerance
# Merge with existing config (so partial updates are preserved)
# Copy to a new dict so SQLAlchemy detects the change on the JSON column
existing = dict(node.remote_config) if node.remote_config else {}
existing.update(config)
node.remote_config = existing
node.config_version = (node.config_version or 0) + 1
node.updated_at = datetime.now(timezone.utc)
context.db.commit()
context.db.refresh(node)
node = ProxyNodeService.update_node_config(
context.db, node_id=self.node_id, config_updates=config_updates
)
context.add_audit_metadata(
action="proxy_node_config_update",
@@ -761,5 +483,5 @@ class AdminUpdateProxyNodeConfigAdapter(AdminApiAdapter):
"node_id": node.id,
"config_version": node.config_version,
"remote_config": node.remote_config,
"node": _node_to_dict(node),
"node": node_to_dict(node),
}

View File

@@ -914,7 +914,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
ctx.selected_base_url = envelope.capture_selected_base_url() if envelope else None
# 记录代理信息
from src.clients.http_client import get_proxy_label, resolve_proxy_info
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
ctx.proxy_info = resolve_proxy_info(provider.proxy)
proxy_label = get_proxy_label(ctx.proxy_info)
@@ -928,19 +928,23 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# simulate streaming to the client (sync -> stream bridge).
if not upstream_is_stream:
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import build_post_kwargs, resolve_delegate_config
request_timeout_sync = provider.request_timeout or config.http_request_timeout
http_client = await HTTPClientPool.get_proxy_client(
proxy_config=provider.proxy,
delegate_cfg = resolve_delegate_config(provider.proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=provider.proxy
)
try:
resp = await http_client.post(
url,
json=provider_payload,
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout_sync),
payload=provider_payload,
timeout=request_timeout_sync,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
@@ -971,12 +975,15 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
ctx.provider_request_headers = provider_headers
# retry once
resp = await http_client.post(
url,
json=provider_payload,
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout_sync),
payload=provider_payload,
timeout=request_timeout_sync,
refresh_auth=True,
)
resp = await http_client.post(**_pkw)
ctx.status_code = resp.status_code
ctx.response_headers = dict(resp.headers)
if envelope:
@@ -1120,10 +1127,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 创建 HTTP 客户端(支持代理配置,从 Provider 读取)
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import build_stream_kwargs, resolve_delegate_config
http_client = HTTPClientPool.create_client_with_proxy(
proxy_config=provider.proxy,
timeout=timeout_config,
delegate_cfg = resolve_delegate_config(provider.proxy)
http_client = HTTPClientPool.create_upstream_stream_client(
delegate_cfg, proxy_config=provider.proxy, timeout=timeout_config
)
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
@@ -1134,9 +1142,18 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
async def _connect_and_prefetch() -> None:
"""建立连接并预读首字节(受整体超时控制)"""
nonlocal byte_iterator, prefetched_chunks, response_ctx
response_ctx = http_client.stream(
"POST", url, json=provider_payload, headers=provider_headers
_skw = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
payload=provider_payload,
timeout=(
provider.request_timeout or config.http_request_timeout
if delegate_cfg
else None
),
)
response_ctx = http_client.stream(**_skw)
stream_response = await response_ctx.__aenter__()
ctx.status_code = stream_response.status_code
@@ -1547,7 +1564,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
# 记录代理信息
from src.clients.http_client import get_proxy_label, resolve_proxy_info
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
sync_proxy_info = resolve_proxy_info(provider.proxy)
_proxy_label = get_proxy_label(sync_proxy_info)
@@ -1562,12 +1579,19 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 获取复用的 HTTP 客户端(支持代理配置,从 Provider 读取)
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import (
build_post_kwargs,
build_stream_kwargs,
resolve_delegate_config,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
http_client = await HTTPClientPool.get_proxy_client(
proxy_config=provider.proxy,
delegate_cfg = resolve_delegate_config(provider.proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=provider.proxy
)
# 注意:不使用 async with因为复用的客户端不应该被关闭
@@ -1575,12 +1599,14 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
resp: httpx.Response | None = None
if not upstream_is_stream:
try:
resp = await http_client.post(
url,
json=provider_payload,
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_hdrs,
timeout=httpx.Timeout(request_timeout),
payload=provider_payload,
timeout=request_timeout,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
@@ -1596,13 +1622,14 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
)
try:
async with http_client.stream(
"POST",
url,
json=provider_payload,
_stream_args = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_hdrs,
timeout=httpx.Timeout(request_timeout),
) as stream_resp:
payload=provider_payload,
timeout=request_timeout,
)
async with http_client.stream(**_stream_args) as stream_resp:
resp = stream_resp
status_code = stream_resp.status_code

View File

@@ -891,8 +891,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
ctx.selected_base_url = envelope.capture_selected_base_url() if envelope else None
# 记录代理信息sync-bridge 路径,早于流式路径执行)
from src.clients.http_client import get_proxy_label as _gpl
from src.clients.http_client import resolve_proxy_info as _rpi
from src.services.proxy_node.resolver import get_proxy_label as _gpl
from src.services.proxy_node.resolver import resolve_proxy_info as _rpi
ctx.proxy_info = _rpi(provider.proxy)
@@ -900,19 +900,23 @@ class CliMessageHandlerBase(BaseMessageHandler):
# simulate streaming to the client (sync -> stream bridge).
if not upstream_is_stream:
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import build_post_kwargs, resolve_delegate_config
request_timeout_sync = provider.request_timeout or config.http_request_timeout
http_client = await HTTPClientPool.get_proxy_client(
proxy_config=provider.proxy,
delegate_cfg = resolve_delegate_config(provider.proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=provider.proxy
)
try:
resp = await http_client.post(
url,
json=provider_payload,
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout_sync),
payload=provider_payload,
timeout=request_timeout_sync,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
@@ -943,12 +947,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
ctx.provider_request_headers = provider_headers
# retry once
resp = await http_client.post(
url,
json=provider_payload,
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout_sync),
payload=provider_payload,
timeout=request_timeout_sync,
refresh_auth=True,
)
resp = await http_client.post(**_pkw)
ctx.status_code = resp.status_code
ctx.response_headers = dict(resp.headers)
if envelope:
@@ -1091,10 +1098,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 创建 HTTP 客户端(支持代理配置,从 Provider 读取)
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import build_stream_kwargs, resolve_delegate_config
http_client = HTTPClientPool.create_client_with_proxy(
proxy_config=provider.proxy,
timeout=timeout_config,
delegate_cfg = resolve_delegate_config(provider.proxy)
http_client = HTTPClientPool.create_upstream_stream_client(
delegate_cfg, proxy_config=provider.proxy, timeout=timeout_config
)
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
@@ -1105,9 +1113,18 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def _connect_and_prefetch() -> None:
"""建立连接并预读首字节(受整体超时控制)"""
nonlocal byte_iterator, prefetched_chunks, response_ctx
response_ctx = http_client.stream(
"POST", url, json=provider_payload, headers=provider_headers
_skw = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
payload=provider_payload,
timeout=(
provider.request_timeout or config.http_request_timeout
if delegate_cfg
else None
),
)
response_ctx = http_client.stream(**_skw)
stream_response = await response_ctx.__aenter__()
ctx.status_code = stream_response.status_code
@@ -2962,7 +2979,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
# 记录代理信息
from src.clients.http_client import get_proxy_label, resolve_proxy_info
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
sync_proxy_info = resolve_proxy_info(provider.proxy)
_proxy_label = get_proxy_label(sync_proxy_info)
@@ -2978,12 +2995,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 获取复用的 HTTP 客户端(支持代理配置,从 Provider 读取)
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import (
build_post_kwargs,
build_stream_kwargs,
resolve_delegate_config,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
http_client = await HTTPClientPool.get_proxy_client(
proxy_config=provider.proxy,
delegate_cfg = resolve_delegate_config(provider.proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=provider.proxy
)
# 注意:不使用 async with因为复用的客户端不应该被关闭
@@ -2991,12 +3015,14 @@ class CliMessageHandlerBase(BaseMessageHandler):
resp: httpx.Response | None = None
if not upstream_is_stream:
try:
resp = await http_client.post(
url,
json=provider_payload,
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout),
payload=provider_payload,
timeout=request_timeout,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
@@ -3013,13 +3039,14 @@ class CliMessageHandlerBase(BaseMessageHandler):
)
try:
async with http_client.stream(
"POST",
url,
json=provider_payload,
_stream_args = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
timeout=httpx.Timeout(request_timeout),
) as stream_resp:
payload=provider_payload,
timeout=request_timeout,
)
async with http_client.stream(**_stream_args) as stream_resp:
resp = stream_resp
status_code = stream_resp.status_code