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:
fawney19
2026-03-18 13:13:47 +08:00
parent 1af3067303
commit 53ef35ec80
6 changed files with 506 additions and 301 deletions
+19
View File
@@ -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],
+171
View File
@@ -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)"""
+22 -144
View File
@@ -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}"
)
+27 -148
View File
@@ -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,
)