mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor(cli): 提取 _build_upstream_request 统一流式/非流式的上游请求构建逻辑
将 cli_stream_mixin 和 cli_sync_mixin 中重复的上游请求构建代码 (provider behavior / stream policy / envelope / auth / RequestBuilder / URL 构建) 提取到 cli_request_mixin._build_upstream_request,返回 CliUpstreamRequestResult dataclass。
This commit is contained in:
@@ -24,6 +24,7 @@ if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.base_handler import MessageTelemetry
|
||||
from src.api.handlers.base.cli_request_mixin import CliUpstreamRequestResult
|
||||
from src.api.handlers.base.request_builder import RequestBuilder
|
||||
from src.api.handlers.base.response_parser import ResponseParser
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
@@ -180,6 +181,24 @@ class CliHandlerProtocol(Protocol):
|
||||
output_limit: int | None = ...,
|
||||
) -> tuple[dict[str, Any], str]: ...
|
||||
|
||||
async def _build_upstream_request(
|
||||
self,
|
||||
*,
|
||||
provider: Any,
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None,
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
fallback_model: str,
|
||||
mapped_model: str | None,
|
||||
client_is_stream: bool,
|
||||
needs_conversion: bool = ...,
|
||||
output_limit: int | None = ...,
|
||||
) -> CliUpstreamRequestResult: ...
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: dict[str, Any],
|
||||
|
||||
@@ -2,19 +2,44 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
)
|
||||
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.api.handlers.base.utils import get_format_converter_registry
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.prompt_cache import maybe_patch_request_with_prompt_cache_key
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
get_upstream_stream_policy,
|
||||
resolve_upstream_is_stream,
|
||||
)
|
||||
from src.services.provider.transport import build_provider_url
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
|
||||
from src.core.api_format import EndpointDefinition
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class CliUpstreamRequestResult:
|
||||
"""Final outbound request artifacts for a selected Provider candidate."""
|
||||
|
||||
payload: dict[str, Any]
|
||||
headers: dict[str, str]
|
||||
url: str
|
||||
url_model: str
|
||||
envelope: Any
|
||||
upstream_is_stream: bool
|
||||
tls_profile: str | None = None
|
||||
selected_base_url: str | None = None
|
||||
|
||||
|
||||
class CliRequestMixin:
|
||||
"""请求准备相关方法的 Mixin"""
|
||||
|
||||
@@ -166,6 +191,152 @@ class CliRequestMixin:
|
||||
|
||||
return request_body
|
||||
|
||||
async def _build_upstream_request(
|
||||
self: CliHandlerProtocol,
|
||||
*,
|
||||
provider: Any,
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None,
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
fallback_model: str,
|
||||
mapped_model: str | None,
|
||||
client_is_stream: bool,
|
||||
needs_conversion: bool = False,
|
||||
output_limit: int | None = None,
|
||||
) -> CliUpstreamRequestResult:
|
||||
"""Build the final outbound URL/body/headers for the selected upstream."""
|
||||
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
target_variant = behavior.same_format_variant
|
||||
conversion_variant = behavior.cross_format_variant
|
||||
|
||||
upstream_policy = get_upstream_stream_policy(
|
||||
endpoint,
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
upstream_is_stream = resolve_upstream_is_stream(
|
||||
client_is_stream=client_is_stream,
|
||||
policy=upstream_policy,
|
||||
)
|
||||
|
||||
envelope_tls_profile: str | None = None
|
||||
if envelope and hasattr(envelope, "prepare_context"):
|
||||
envelope_tls_profile = envelope.prepare_context(
|
||||
provider_config=getattr(provider, "config", None),
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
|
||||
is_stream=upstream_is_stream,
|
||||
provider_id=str(getattr(provider, "id", "") or ""),
|
||||
key=key,
|
||||
)
|
||||
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = await self._convert_request_for_cross_format(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
fallback_model,
|
||||
is_stream=upstream_is_stream,
|
||||
target_variant=conversion_variant,
|
||||
output_limit=output_limit,
|
||||
)
|
||||
else:
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
url_model = (
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or fallback_model
|
||||
)
|
||||
if target_variant and provider_api_format:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
upstream_is_stream=upstream_is_stream,
|
||||
)
|
||||
|
||||
request_body = maybe_patch_request_with_prompt_cache_key(
|
||||
request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_type=provider_type,
|
||||
base_url=getattr(endpoint, "base_url", None),
|
||||
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
|
||||
request_headers=original_headers,
|
||||
)
|
||||
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
if envelope:
|
||||
request_body, url_model = envelope.wrap_request(
|
||||
request_body,
|
||||
model=url_model or fallback_model or "",
|
||||
url_model=url_model,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
if hasattr(envelope, "post_wrap_request"):
|
||||
await envelope.post_wrap_request(request_body)
|
||||
|
||||
extra_headers: dict[str, str] = {}
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=upstream_is_stream,
|
||||
extra_headers=extra_headers if extra_headers else None,
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
envelope=envelope,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=upstream_is_stream,
|
||||
key=key,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
selected_base_url = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
return CliUpstreamRequestResult(
|
||||
payload=provider_payload,
|
||||
headers=provider_headers,
|
||||
url=str(url),
|
||||
url_model=str(url_model or fallback_model or ""),
|
||||
envelope=envelope,
|
||||
upstream_is_stream=upstream_is_stream,
|
||||
tls_profile=envelope_tls_profile,
|
||||
selected_base_url=selected_base_url,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_format_metadata(format_id: str) -> "EndpointDefinition | None":
|
||||
"""获取 endpoint 元数据(解析失败返回 None)"""
|
||||
|
||||
@@ -19,9 +19,7 @@ from src.api.handlers.base.base_handler import (
|
||||
wait_for_with_disconnect_detection,
|
||||
)
|
||||
from src.api.handlers.base.parsers import get_parser_for_format
|
||||
from src.api.handlers.base.request_builder import (
|
||||
get_provider_auth,
|
||||
)
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.utils import (
|
||||
build_sse_headers,
|
||||
@@ -41,15 +39,6 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.prompt_cache import (
|
||||
maybe_patch_request_with_prompt_cache_key,
|
||||
)
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
get_upstream_stream_policy,
|
||||
resolve_upstream_is_stream,
|
||||
)
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.utils.sse_parser import SSEEventParser
|
||||
@@ -325,143 +314,32 @@ class CliStreamMixin:
|
||||
)
|
||||
ctx.needs_conversion = needs_conversion
|
||||
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
target_variant = behavior.same_format_variant
|
||||
# 跨格式转换也允许变体(Antigravity 需要保留/翻译 Claude thinking 块)
|
||||
conversion_variant = behavior.cross_format_variant
|
||||
|
||||
# Upstream streaming policy (per-endpoint): may force upstream to sync/stream mode.
|
||||
upstream_policy = get_upstream_stream_policy(
|
||||
endpoint,
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
upstream_is_stream = resolve_upstream_is_stream(
|
||||
client_is_stream=True,
|
||||
policy=upstream_policy,
|
||||
)
|
||||
# Envelope lifecycle: prepare_context (pre-wrap hook).
|
||||
envelope_tls_profile: str | None = None
|
||||
if envelope and hasattr(envelope, "prepare_context"):
|
||||
envelope_tls_profile = envelope.prepare_context(
|
||||
provider_config=getattr(provider, "config", None),
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
|
||||
is_stream=upstream_is_stream,
|
||||
provider_id=str(getattr(provider, "id", "") or ""),
|
||||
key=key,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = await self._convert_request_for_cross_format(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
ctx.model,
|
||||
is_stream=upstream_is_stream,
|
||||
target_variant=conversion_variant,
|
||||
output_limit=candidate.output_limit if candidate else None,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
url_model = (
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or ctx.model
|
||||
)
|
||||
# 同格式 Provider 仍可能声明 target_variant
|
||||
if target_variant and provider_api_format:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
upstream_request = await self._build_upstream_request(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body=request_body,
|
||||
original_headers=original_headers,
|
||||
query_params=query_params,
|
||||
client_api_format=client_api_format,
|
||||
provider_api_format=provider_api_format,
|
||||
fallback_model=ctx.model,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=provider_api_format,
|
||||
client_is_stream=True,
|
||||
needs_conversion=needs_conversion,
|
||||
output_limit=candidate.output_limit if candidate else None,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
upstream_is_stream=upstream_is_stream,
|
||||
)
|
||||
|
||||
request_body = maybe_patch_request_with_prompt_cache_key(
|
||||
request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_type=provider_type,
|
||||
base_url=getattr(endpoint, "base_url", None),
|
||||
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
|
||||
request_headers=original_headers,
|
||||
)
|
||||
|
||||
# 获取认证信息(处理 Service Account 等异步认证场景)
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
|
||||
# Provider envelope: wrap request after auth is available and before RequestBuilder.build().
|
||||
if envelope:
|
||||
request_body, url_model = envelope.wrap_request(
|
||||
request_body,
|
||||
model=url_model or ctx.model or "",
|
||||
url_model=url_model,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
# Envelope lifecycle: post_wrap_request (post-wrap hook).
|
||||
if hasattr(envelope, "post_wrap_request"):
|
||||
await envelope.post_wrap_request(request_body)
|
||||
|
||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||
extra_headers: dict[str, str] = {}
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
# 使用 RequestBuilder 构建请求体和请求头
|
||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||
# 上游始终使用 header 认证,不跟随客户端的 query 方式
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=upstream_is_stream,
|
||||
extra_headers=extra_headers if extra_headers else None,
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
envelope=envelope,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
provider_headers = upstream_request.headers
|
||||
provider_payload = upstream_request.payload
|
||||
url = upstream_request.url
|
||||
envelope = upstream_request.envelope
|
||||
upstream_is_stream = upstream_request.upstream_is_stream
|
||||
envelope_tls_profile = upstream_request.tls_profile
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
ctx.provider_request_headers = provider_headers
|
||||
ctx.provider_request_body = provider_payload
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=upstream_is_stream,
|
||||
key=key,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
# Capture the selected base_url from transport (used by some envelopes for failover).
|
||||
ctx.selected_base_url = envelope.capture_selected_base_url() if envelope else None
|
||||
ctx.selected_base_url = upstream_request.selected_base_url
|
||||
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import get_proxy_label as _gpl
|
||||
@@ -670,7 +548,7 @@ class CliStreamMixin:
|
||||
f" └─ [{self.request_id}] 发送流式请求: "
|
||||
f"Provider={provider.name}, Endpoint={endpoint.id[:8] if endpoint.id else 'N/A'}..., "
|
||||
f"Key=***{key.api_key[-4:] if key.api_key else 'N/A'}, "
|
||||
f"原始模型={ctx.model}, 映射后={mapped_model or '无映射'}, URL模型={url_model}, "
|
||||
f"原始模型={ctx.model}, 映射后={mapped_model or '无映射'}, URL模型={upstream_request.url_model}, "
|
||||
f"timeout={request_timeout}s, 代理={_proxy_label}"
|
||||
)
|
||||
|
||||
|
||||
@@ -11,9 +11,6 @@ import httpx
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from src.api.handlers.base.parsers import get_parser_for_format
|
||||
from src.api.handlers.base.request_builder import (
|
||||
get_provider_auth,
|
||||
)
|
||||
from src.api.handlers.base.stream_context import extract_proxy_timing, is_format_converted
|
||||
from src.api.handlers.base.upstream_stream_bridge import (
|
||||
aggregate_upstream_stream_to_internal_response,
|
||||
@@ -35,16 +32,6 @@ from src.core.exceptions import (
|
||||
ThinkingSignatureException,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.prompt_cache import (
|
||||
maybe_patch_request_with_prompt_cache_key,
|
||||
)
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
get_upstream_stream_policy,
|
||||
resolve_upstream_is_stream,
|
||||
)
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -146,143 +133,30 @@ class CliSyncMixin:
|
||||
)
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
target_variant = behavior.same_format_variant
|
||||
# 跨格式转换也允许变体(Antigravity 需要保留/翻译 Claude thinking 块)
|
||||
conversion_variant = behavior.cross_format_variant
|
||||
|
||||
# Upstream streaming policy (per-endpoint).
|
||||
upstream_policy = get_upstream_stream_policy(
|
||||
endpoint,
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
upstream_is_stream = resolve_upstream_is_stream(
|
||||
client_is_stream=False,
|
||||
policy=upstream_policy,
|
||||
)
|
||||
# Envelope lifecycle: prepare_context (pre-wrap hook).
|
||||
envelope_tls_profile: str | None = None
|
||||
if envelope and hasattr(envelope, "prepare_context"):
|
||||
envelope_tls_profile = envelope.prepare_context(
|
||||
provider_config=getattr(provider, "config", None),
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
|
||||
is_stream=upstream_is_stream,
|
||||
provider_id=str(getattr(provider, "id", "") or ""),
|
||||
key=key,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = await self._convert_request_for_cross_format(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
model,
|
||||
is_stream=upstream_is_stream,
|
||||
target_variant=conversion_variant,
|
||||
output_limit=candidate.output_limit if candidate else None,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
url_model = (
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or model
|
||||
)
|
||||
# 同格式 Provider 仍可能声明 target_variant
|
||||
if target_variant and provider_api_format:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
upstream_request = await self._build_upstream_request(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body=request_body,
|
||||
original_headers=original_headers,
|
||||
query_params=query_params,
|
||||
client_api_format=client_api_format,
|
||||
provider_api_format=provider_api_format,
|
||||
fallback_model=model,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=provider_api_format,
|
||||
client_is_stream=False,
|
||||
needs_conversion=needs_conversion,
|
||||
output_limit=candidate.output_limit if candidate else None,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
upstream_is_stream=upstream_is_stream,
|
||||
)
|
||||
|
||||
request_body = maybe_patch_request_with_prompt_cache_key(
|
||||
request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_type=provider_type,
|
||||
base_url=getattr(endpoint, "base_url", None),
|
||||
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
|
||||
request_headers=original_headers,
|
||||
)
|
||||
|
||||
# 获取认证信息(处理 Service Account 等异步认证场景)
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
|
||||
# Provider envelope: wrap request after auth is available and before RequestBuilder.build().
|
||||
if envelope:
|
||||
request_body, url_model = envelope.wrap_request(
|
||||
request_body,
|
||||
model=url_model or model or "",
|
||||
url_model=url_model,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
# Envelope lifecycle: post_wrap_request (post-wrap hook).
|
||||
if hasattr(envelope, "post_wrap_request"):
|
||||
await envelope.post_wrap_request(request_body)
|
||||
|
||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||
extra_headers: dict[str, str] = {}
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
# 使用 RequestBuilder 构建请求体和请求头
|
||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||
# 上游始终使用 header 认证,不跟随客户端的 query 方式
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=upstream_is_stream,
|
||||
extra_headers=extra_headers if extra_headers else None,
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
envelope=envelope,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
provider_headers = upstream_request.headers
|
||||
provider_payload = upstream_request.payload
|
||||
provider_request_headers = provider_headers
|
||||
provider_request_body = provider_payload
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=upstream_is_stream, # sync handler may still force upstream streaming
|
||||
key=key,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
|
||||
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
|
||||
url = upstream_request.url
|
||||
envelope = upstream_request.envelope
|
||||
upstream_is_stream = upstream_request.upstream_is_stream
|
||||
envelope_tls_profile = upstream_request.tls_profile
|
||||
selected_base_url_cached = upstream_request.selected_base_url
|
||||
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
@@ -299,7 +173,7 @@ class CliSyncMixin:
|
||||
f" └─ [{self.request_id}] 发送{'上游流式(聚合)' if upstream_is_stream else '非流式'}请求: "
|
||||
f"Provider={provider.name}, Endpoint={endpoint.id[:8] if endpoint.id else 'N/A'}..., "
|
||||
f"Key=***{key.api_key[-4:] if key.api_key else 'N/A'}, "
|
||||
f"原始模型={model}, 映射后={mapped_model or '无映射'}, URL模型={url_model}, "
|
||||
f"原始模型={model}, 映射后={mapped_model or '无映射'}, URL模型={upstream_request.url_model}, "
|
||||
f"代理={_proxy_label}"
|
||||
)
|
||||
|
||||
@@ -377,7 +251,12 @@ class CliSyncMixin:
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
byte_iter = stream_resp.aiter_bytes()
|
||||
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
_provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
if (
|
||||
_provider_type == "kiro"
|
||||
and envelope
|
||||
and envelope.force_stream_rewrite()
|
||||
):
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
@@ -6,7 +6,8 @@ from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
import src.api.handlers.base.cli_stream_mixin as mixmod
|
||||
import src.api.handlers.base.cli_request_mixin as request_mixmod
|
||||
from src.api.handlers.base.cli_request_mixin import CliRequestMixin
|
||||
from src.api.handlers.base.cli_stream_mixin import CliStreamMixin
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
|
||||
@@ -33,7 +34,7 @@ class _CaptureBuilder:
|
||||
raise _StopBuild()
|
||||
|
||||
|
||||
class _DummyCliStreamHandler(CliStreamMixin):
|
||||
class _DummyCliStreamHandler(CliRequestMixin, CliStreamMixin):
|
||||
FORMAT_ID = "openai:cli"
|
||||
|
||||
def __init__(self) -> None:
|
||||
@@ -79,9 +80,9 @@ async def test_execute_stream_request_does_not_mutate_original_request_body(
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
monkeypatch.setattr(request_mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
monkeypatch.setattr(
|
||||
mixmod,
|
||||
request_mixmod,
|
||||
"get_provider_behavior",
|
||||
lambda **kwargs: SimpleNamespace(
|
||||
envelope=None,
|
||||
@@ -89,15 +90,19 @@ async def test_execute_stream_request_does_not_mutate_original_request_body(
|
||||
cross_format_variant=None,
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(mixmod, "get_upstream_stream_policy", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(request_mixmod, "get_upstream_stream_policy", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
mixmod,
|
||||
request_mixmod,
|
||||
"resolve_upstream_is_stream",
|
||||
lambda *, client_is_stream, policy: client_is_stream,
|
||||
)
|
||||
monkeypatch.setattr(mixmod, "enforce_stream_mode_for_upstream", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
mixmod,
|
||||
request_mixmod,
|
||||
"enforce_stream_mode_for_upstream",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
request_mixmod,
|
||||
"maybe_patch_request_with_prompt_cache_key",
|
||||
lambda request_body, **kwargs: request_body,
|
||||
)
|
||||
@@ -107,7 +112,12 @@ async def test_execute_stream_request_does_not_mutate_original_request_body(
|
||||
ctx.client_api_format = "openai:cli"
|
||||
|
||||
provider = SimpleNamespace(name="provider", id="provider-1", provider_type="", proxy=None)
|
||||
endpoint = SimpleNamespace(id="endpoint-1", api_format="openai:cli", base_url="https://x")
|
||||
endpoint = SimpleNamespace(
|
||||
id="endpoint-1",
|
||||
api_format="openai:cli",
|
||||
base_url="https://x",
|
||||
custom_path=None,
|
||||
)
|
||||
key = SimpleNamespace(id="key-1", proxy=None)
|
||||
candidate = SimpleNamespace(
|
||||
mapping_matched_model=None, needs_conversion=False, output_limit=None
|
||||
|
||||
248
tests/api/handlers/base/test_cli_upstream_request_builder.py
Normal file
248
tests/api/handlers/base/test_cli_upstream_request_builder.py
Normal file
@@ -0,0 +1,248 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
import src.api.handlers.base.cli_request_mixin as mixmod
|
||||
from src.api.handlers.base.cli_request_mixin import CliRequestMixin
|
||||
from src.api.handlers.base.request_builder import PassthroughRequestBuilder
|
||||
from src.core.api_format.metadata import CODEX_DEFAULT_BODY_RULES
|
||||
from src.services.provider.adapters.codex.context import (
|
||||
CodexRequestContext,
|
||||
set_codex_request_context,
|
||||
)
|
||||
from src.services.provider.prompt_cache import build_stable_codex_prompt_cache_key
|
||||
|
||||
|
||||
class _DummyAuthInfo:
|
||||
auth_header = "Authorization"
|
||||
auth_value = "Bearer upstream-token"
|
||||
decrypted_auth_config = None
|
||||
|
||||
def as_tuple(self) -> tuple[str, str]:
|
||||
return self.auth_header, self.auth_value
|
||||
|
||||
|
||||
class _DummyCliRequestHandler(CliRequestMixin):
|
||||
FORMAT_ID = "openai:cli"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.api_key = SimpleNamespace(id="user-key-123")
|
||||
self._request_builder = PassthroughRequestBuilder()
|
||||
|
||||
|
||||
def _build_codex_provider() -> Any:
|
||||
return SimpleNamespace(
|
||||
id="provider-1",
|
||||
provider_type="codex",
|
||||
config=None,
|
||||
proxy=None,
|
||||
)
|
||||
|
||||
|
||||
def _build_codex_endpoint(*, api_format: str) -> Any:
|
||||
provider = _build_codex_provider()
|
||||
return SimpleNamespace(
|
||||
id="endpoint-1",
|
||||
api_family="openai",
|
||||
endpoint_kind="cli",
|
||||
api_format=api_format,
|
||||
base_url="https://chatgpt.com/backend-api/codex",
|
||||
custom_path=None,
|
||||
body_rules=list(CODEX_DEFAULT_BODY_RULES),
|
||||
header_rules=None,
|
||||
provider=provider,
|
||||
)
|
||||
|
||||
|
||||
def _build_key() -> Any:
|
||||
return SimpleNamespace(id="key-1", api_key="unused", proxy=None)
|
||||
|
||||
|
||||
def _build_headers() -> dict[str, str]:
|
||||
return {
|
||||
"accept": "application/json",
|
||||
"content-type": "application/json",
|
||||
"user-agent": "Codex Desktop/0.108.0-alpha.12",
|
||||
"originator": "Codex Desktop",
|
||||
"x-codex-turn-metadata": '{"turn_id":"abc"}',
|
||||
"host": "aether.hetunai.cn",
|
||||
"content-length": "123",
|
||||
"x-forwarded-scheme": "https",
|
||||
}
|
||||
|
||||
|
||||
def _assert_common_codex_headers(headers: dict[str, str]) -> None:
|
||||
assert headers["accept"] == "application/json"
|
||||
assert headers["content-type"] == "application/json"
|
||||
assert headers["user-agent"] == "Codex Desktop/0.108.0-alpha.12"
|
||||
assert headers["originator"] == "Codex Desktop"
|
||||
assert headers["x-codex-turn-metadata"] == '{"turn_id":"abc"}'
|
||||
assert headers["Authorization"] == "Bearer upstream-token"
|
||||
assert "host" not in headers
|
||||
assert "content-length" not in headers
|
||||
assert "x-forwarded-scheme" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_upstream_request_codex_cli_injects_prompt_cache_and_forces_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
|
||||
handler = _DummyCliRequestHandler()
|
||||
endpoint = _build_codex_endpoint(api_format="openai:cli")
|
||||
provider = endpoint.provider
|
||||
key = _build_key()
|
||||
|
||||
result = await handler._build_upstream_request(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body={
|
||||
"model": "gpt-5",
|
||||
"input": [],
|
||||
"stream": False,
|
||||
"max_output_tokens": 4096,
|
||||
"temperature": 0.7,
|
||||
"top_p": 0.8,
|
||||
},
|
||||
original_headers=_build_headers(),
|
||||
query_params=None,
|
||||
client_api_format="openai:cli",
|
||||
provider_api_format="openai:cli",
|
||||
fallback_model="gpt-5",
|
||||
mapped_model=None,
|
||||
client_is_stream=False,
|
||||
)
|
||||
|
||||
assert result.url == "https://chatgpt.com/backend-api/codex/responses"
|
||||
assert result.upstream_is_stream is True
|
||||
assert result.payload["stream"] is True
|
||||
assert result.payload["instructions"] == "You are GPT-5."
|
||||
assert result.payload["store"] is False
|
||||
assert result.payload["prompt_cache_key"] == build_stable_codex_prompt_cache_key("user-key-123")
|
||||
assert "max_output_tokens" not in result.payload
|
||||
assert "temperature" not in result.payload
|
||||
assert "top_p" not in result.payload
|
||||
_assert_common_codex_headers(result.headers)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_upstream_request_codex_compact_drops_stream_and_skips_prompt_cache(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
|
||||
handler = _DummyCliRequestHandler()
|
||||
endpoint = _build_codex_endpoint(api_format="openai:compact")
|
||||
provider = endpoint.provider
|
||||
key = _build_key()
|
||||
|
||||
result = await handler._build_upstream_request(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body={"model": "gpt-5", "input": [], "stream": True},
|
||||
original_headers=_build_headers(),
|
||||
query_params=None,
|
||||
client_api_format="openai:cli",
|
||||
provider_api_format="openai:compact",
|
||||
fallback_model="gpt-5",
|
||||
mapped_model=None,
|
||||
client_is_stream=False,
|
||||
)
|
||||
|
||||
assert result.url == "https://chatgpt.com/backend-api/codex/responses/compact"
|
||||
assert result.upstream_is_stream is False
|
||||
assert "stream" not in result.payload
|
||||
assert "prompt_cache_key" not in result.payload
|
||||
assert result.payload["instructions"] == "You are GPT-5."
|
||||
assert result.payload["store"] is False
|
||||
_assert_common_codex_headers(result.headers)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_upstream_request_legacy_codex_compact_context_uses_compact_url(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
|
||||
handler = _DummyCliRequestHandler()
|
||||
endpoint = _build_codex_endpoint(api_format="openai:cli")
|
||||
provider = endpoint.provider
|
||||
key = _build_key()
|
||||
|
||||
try:
|
||||
set_codex_request_context(CodexRequestContext(is_compact=True))
|
||||
result = await handler._build_upstream_request(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body={"model": "gpt-5", "input": [], "stream": True},
|
||||
original_headers=_build_headers(),
|
||||
query_params=None,
|
||||
client_api_format="openai:cli",
|
||||
provider_api_format="openai:cli",
|
||||
fallback_model="gpt-5",
|
||||
mapped_model=None,
|
||||
client_is_stream=False,
|
||||
)
|
||||
finally:
|
||||
set_codex_request_context(None)
|
||||
|
||||
assert result.url == "https://chatgpt.com/backend-api/codex/responses/compact"
|
||||
assert result.upstream_is_stream is False
|
||||
assert "stream" not in result.payload
|
||||
assert "prompt_cache_key" not in result.payload
|
||||
assert result.payload["instructions"] == "You are GPT-5."
|
||||
_assert_common_codex_headers(result.headers)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_upstream_request_codex_cli_preserves_explicit_prompt_cache_key(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
|
||||
handler = _DummyCliRequestHandler()
|
||||
endpoint = _build_codex_endpoint(api_format="openai:cli")
|
||||
provider = endpoint.provider
|
||||
key = _build_key()
|
||||
|
||||
result = await handler._build_upstream_request(
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body={
|
||||
"model": "gpt-5",
|
||||
"input": [],
|
||||
"prompt_cache_key": "client-cache-key",
|
||||
},
|
||||
original_headers=_build_headers(),
|
||||
query_params=None,
|
||||
client_api_format="openai:cli",
|
||||
provider_api_format="openai:cli",
|
||||
fallback_model="gpt-5",
|
||||
mapped_model=None,
|
||||
client_is_stream=False,
|
||||
)
|
||||
|
||||
assert result.url == "https://chatgpt.com/backend-api/codex/responses"
|
||||
assert result.payload["prompt_cache_key"] == "client-cache-key"
|
||||
assert result.payload["stream"] is True
|
||||
_assert_common_codex_headers(result.headers)
|
||||
Reference in New Issue
Block a user