mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
feat: Antigravity 和 Codex 服务支持
- 新增 Antigravity 服务:签名缓存、URL 可用性检测、信封处理 - 新增 Codex 服务:信封处理、元数据收集器 - 重构 provider transport 支持新的服务架构 - 新增 stream_bridge 和 upstream_stream_bridge 处理流式响应 - 优化 OAuth 工具函数 - 添加相关测试用例
This commit is contained in:
@@ -43,12 +43,18 @@ from src.api.handlers.base.response_parser import ResponseParser
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.stream_processor import StreamProcessor
|
||||
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
|
||||
from src.api.handlers.base.upstream_stream_bridge import (
|
||||
aggregate_upstream_stream_to_internal_response,
|
||||
)
|
||||
from src.api.handlers.base.utils import (
|
||||
build_sse_headers,
|
||||
filter_proxy_response_headers,
|
||||
get_format_converter_registry,
|
||||
)
|
||||
from src.config.settings import config
|
||||
from src.core.api_format.conversion.stream_bridge import (
|
||||
iter_internal_response_as_stream_events,
|
||||
)
|
||||
from src.core.error_utils import extract_client_error_message
|
||||
from src.core.exceptions import (
|
||||
EmbeddedErrorException,
|
||||
@@ -68,6 +74,12 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
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,
|
||||
get_vertex_ai_effective_format,
|
||||
@@ -760,9 +772,25 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(original_request_body)
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
same_format_variant = behavior.same_format_variant
|
||||
cross_format_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=str(provider_api_format),
|
||||
)
|
||||
upstream_is_stream = resolve_upstream_is_stream(
|
||||
client_is_stream=True,
|
||||
policy=upstream_policy,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
registry = get_format_converter_registry()
|
||||
@@ -771,7 +799,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
request_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
target_variant=target_variant,
|
||||
target_variant=cross_format_variant,
|
||||
)
|
||||
# 格式转换后,为需要 model 字段的格式设置模型名
|
||||
self._set_model_after_conversion(
|
||||
@@ -785,50 +813,223 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
request_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
is_stream=True,
|
||||
is_stream=upstream_is_stream,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
if same_format_variant:
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
str(provider_api_format),
|
||||
str(provider_api_format),
|
||||
target_variant=target_variant,
|
||||
target_variant=same_format_variant,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
request_body,
|
||||
provider_api_format=str(provider_api_format),
|
||||
upstream_is_stream=upstream_is_stream,
|
||||
)
|
||||
|
||||
# 获取 URL 模型名
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or ctx.model
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||
extra_headers: dict[str, str] = {}
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=True,
|
||||
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,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when in streaming mode.
|
||||
provider_headers["Accept"] = "text/event-stream"
|
||||
|
||||
ctx.provider_request_headers = provider_headers
|
||||
ctx.provider_request_body = provider_payload
|
||||
|
||||
# 获取 URL 模型名
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or ctx.model
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=True,
|
||||
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
|
||||
|
||||
logger.debug(
|
||||
f" [{self.request_id}] 发送流式请求: Provider={provider.name}, "
|
||||
f"模型={ctx.model} -> {mapped_model or '无映射'}"
|
||||
)
|
||||
|
||||
# If upstream is forced to non-stream mode, we execute a sync request and then
|
||||
# simulate streaming to the client (sync -> stream bridge).
|
||||
if not upstream_is_stream:
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
|
||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||
http_client = await HTTPClientPool.get_proxy_client(
|
||||
proxy_config=provider.proxy,
|
||||
)
|
||||
|
||||
try:
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout_sync),
|
||||
)
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
ctx.status_code = resp.status_code
|
||||
ctx.response_headers = dict(resp.headers)
|
||||
if envelope:
|
||||
envelope.on_http_status(base_url=ctx.selected_base_url, status_code=ctx.status_code)
|
||||
|
||||
# Reuse HTTPStatusError classification path (handled by TaskService/error_classifier).
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
# Safe JSON parsing.
|
||||
try:
|
||||
response_json = resp.json()
|
||||
except (UnicodeDecodeError, json.JSONDecodeError) as e:
|
||||
raw_content = ""
|
||||
try:
|
||||
raw_content = resp.text[:500] if resp.text else "(empty)"
|
||||
except Exception:
|
||||
try:
|
||||
raw_content = repr(resp.content[:500]) if resp.content else "(empty)"
|
||||
except Exception:
|
||||
raw_content = "(unable to read)"
|
||||
raise ProviderNotAvailableException(
|
||||
"上游服务返回了无效的响应",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=f"json_decode_error={type(e).__name__}: {raw_content}",
|
||||
)
|
||||
|
||||
if envelope:
|
||||
response_json = envelope.unwrap_response(response_json)
|
||||
envelope.postprocess_unwrapped_response(model=ctx.model, data=response_json)
|
||||
|
||||
# Embedded error detection (HTTP 200 but error body).
|
||||
if isinstance(response_json, dict):
|
||||
parser = get_parser_for_format(provider_api_format)
|
||||
if parser.is_error_response(response_json):
|
||||
parsed = parser.parse_response(response_json, 200)
|
||||
raise EmbeddedErrorException(
|
||||
provider_name=str(provider.name),
|
||||
error_code=parsed.embedded_status_code,
|
||||
error_message=parsed.error_message,
|
||||
error_status=parsed.error_type,
|
||||
)
|
||||
|
||||
# Convert sync JSON -> InternalResponse, then InternalResponse -> client stream events.
|
||||
src_norm = (
|
||||
registry.get_normalizer(str(provider_api_format)) if provider_api_format else None
|
||||
)
|
||||
if src_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {provider_api_format}")
|
||||
|
||||
internal_resp = src_norm.response_to_internal(
|
||||
response_json if isinstance(response_json, dict) else {}
|
||||
)
|
||||
internal_resp.model = str(ctx.model or internal_resp.model or "")
|
||||
if internal_resp.id:
|
||||
ctx.response_id = internal_resp.id
|
||||
|
||||
if internal_resp.usage:
|
||||
ctx.input_tokens = int(internal_resp.usage.input_tokens or 0)
|
||||
ctx.output_tokens = int(internal_resp.usage.output_tokens or 0)
|
||||
ctx.cached_tokens = int(internal_resp.usage.cache_read_tokens or 0)
|
||||
ctx.cache_creation_tokens = int(internal_resp.usage.cache_write_tokens or 0)
|
||||
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
tgt_norm = (
|
||||
registry.get_normalizer(str(client_api_format)) if client_api_format else None
|
||||
)
|
||||
if tgt_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
|
||||
|
||||
state = StreamState(
|
||||
model=str(ctx.model or ""),
|
||||
message_id=str(ctx.response_id or ctx.request_id or self.request_id or ""),
|
||||
)
|
||||
|
||||
output_state = {"started": False}
|
||||
|
||||
async def _streamified() -> AsyncGenerator[bytes]:
|
||||
for ev in iter_internal_response_as_stream_events(internal_resp):
|
||||
converted_events = tgt_norm.stream_event_from_internal(ev, state)
|
||||
if not converted_events:
|
||||
continue
|
||||
for evt in converted_events:
|
||||
if isinstance(evt, dict):
|
||||
ctx.data_count += 1
|
||||
if ctx.record_parsed_chunks:
|
||||
ctx.parsed_chunks.append(evt)
|
||||
payload = json.dumps(evt, ensure_ascii=False)
|
||||
ctx.chunk_count += 1
|
||||
if not output_state["started"]:
|
||||
ctx.record_first_byte_time(self.start_time)
|
||||
if stream_processor.on_streaming_start:
|
||||
stream_processor.on_streaming_start()
|
||||
output_state["started"] = True
|
||||
yield f"data: {payload}\n\n".encode("utf-8")
|
||||
|
||||
# OpenAI chat clients expect a final [DONE] marker.
|
||||
if str(client_api_format or "").strip().lower() == "openai:chat":
|
||||
if not output_state["started"]:
|
||||
ctx.record_first_byte_time(self.start_time)
|
||||
if stream_processor.on_streaming_start:
|
||||
stream_processor.on_streaming_start()
|
||||
output_state["started"] = True
|
||||
ctx.chunk_count += 1
|
||||
yield b"data: [DONE]\n\n"
|
||||
ctx.has_completion = True
|
||||
|
||||
return _streamified()
|
||||
|
||||
# 配置 HTTP 超时
|
||||
# 注意:read timeout 用于检测连接断开,不是整体请求超时
|
||||
# 整体请求超时由 asyncio.wait_for 控制,使用全局配置
|
||||
@@ -866,6 +1067,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
ctx.status_code = stream_response.status_code
|
||||
ctx.response_headers = dict(stream_response.headers)
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=ctx.selected_base_url,
|
||||
status_code=ctx.status_code,
|
||||
)
|
||||
|
||||
stream_response.raise_for_status()
|
||||
|
||||
@@ -925,6 +1131,22 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
timeout=int(request_timeout),
|
||||
)
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
# 连接/读写超时:清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_text = await self._extract_error_text(e)
|
||||
logger.error(f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}")
|
||||
@@ -1099,9 +1321,25 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(request_body_ref["body"])
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
same_format_variant = behavior.same_format_variant
|
||||
cross_format_variant = behavior.cross_format_variant
|
||||
|
||||
# Upstream streaming policy (per-endpoint).
|
||||
upstream_policy = get_upstream_stream_policy(
|
||||
endpoint,
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=str(provider_api_format),
|
||||
)
|
||||
upstream_is_stream = resolve_upstream_is_stream(
|
||||
client_is_stream=False,
|
||||
policy=upstream_policy,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
registry = get_format_converter_registry()
|
||||
@@ -1110,7 +1348,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
target_variant=cross_format_variant,
|
||||
)
|
||||
# 格式转换后,为需要 model 字段的格式设置模型名
|
||||
self._set_model_after_conversion(
|
||||
@@ -1124,48 +1362,76 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
is_stream=False,
|
||||
is_stream=upstream_is_stream,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
if same_format_variant:
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
target_variant=same_format_variant,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
request_body,
|
||||
provider_api_format=str(provider_api_format),
|
||||
upstream_is_stream=upstream_is_stream,
|
||||
)
|
||||
|
||||
# 获取 URL 模型名(兜底使用外层的 model,确保 Gemini 等格式能正确构建 URL)
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or model
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
|
||||
extra_headers: dict[str, str] = {}
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||
provider_payload, provider_hdrs = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=False,
|
||||
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,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when forced to streaming mode.
|
||||
provider_hdrs["Accept"] = "text/event-stream"
|
||||
|
||||
provider_request_headers = provider_hdrs
|
||||
provider_request_body = provider_payload
|
||||
|
||||
# 获取 URL 模型名(兜底使用外层的 model,确保 Gemini 等格式能正确构建 URL)
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or model
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=False,
|
||||
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
|
||||
|
||||
logger.info(
|
||||
f" [{self.request_id}] 发送非流式请求: Provider={provider.name}, "
|
||||
f"模型={model} -> {mapped_model or '无映射'}"
|
||||
f" [{self.request_id}] 发送{'上游流式(聚合)' if upstream_is_stream else '非流式'}请求: "
|
||||
f"Provider={provider.name}, 模型={model} -> {mapped_model or '无映射'}"
|
||||
)
|
||||
logger.debug(f" [{self.request_id}] 请求URL: {redact_url_for_log(url)}")
|
||||
|
||||
@@ -1182,16 +1448,93 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||
# 超时通过 timeout 参数控制
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_hdrs,
|
||||
timeout=httpx.Timeout(request_timeout),
|
||||
)
|
||||
resp: httpx.Response | None = None
|
||||
if not upstream_is_stream:
|
||||
try:
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_hdrs,
|
||||
timeout=httpx.Timeout(request_timeout),
|
||||
)
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
|
||||
if selected_base_url_cached:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
|
||||
)
|
||||
raise
|
||||
else:
|
||||
# Forced upstream streaming: aggregate SSE to a sync JSON response.
|
||||
provider_parser = (
|
||||
get_parser_for_format(provider_api_format) if provider_api_format else None
|
||||
)
|
||||
|
||||
try:
|
||||
async with http_client.stream(
|
||||
"POST",
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_hdrs,
|
||||
timeout=httpx.Timeout(request_timeout),
|
||||
) as stream_resp:
|
||||
resp = stream_resp
|
||||
|
||||
status_code = stream_resp.status_code
|
||||
response_headers = dict(stream_resp.headers)
|
||||
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=selected_base_url_cached,
|
||||
status_code=status_code,
|
||||
)
|
||||
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
internal_resp = await aggregate_upstream_stream_to_internal_response(
|
||||
stream_resp.aiter_bytes(),
|
||||
provider_api_format=provider_api_format,
|
||||
provider_name=str(provider.name),
|
||||
model=str(model or ""),
|
||||
request_id=str(self.request_id or ""),
|
||||
envelope=envelope,
|
||||
provider_parser=provider_parser,
|
||||
)
|
||||
|
||||
tgt_norm = (
|
||||
registry.get_normalizer(client_api_format)
|
||||
if client_api_format
|
||||
else None
|
||||
)
|
||||
if tgt_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
|
||||
|
||||
response_json = tgt_norm.response_from_internal(
|
||||
internal_resp,
|
||||
requested_model=model,
|
||||
)
|
||||
response_json = response_json if isinstance(response_json, dict) else {}
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
|
||||
if selected_base_url_cached:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
status_code = resp.status_code
|
||||
response_headers = dict(resp.headers)
|
||||
|
||||
if envelope:
|
||||
envelope.on_http_status(base_url=selected_base_url_cached, status_code=status_code)
|
||||
|
||||
# Forced upstream streaming already built response_json via aggregator.
|
||||
if upstream_is_stream:
|
||||
return response_json if isinstance(response_json, dict) else {}
|
||||
|
||||
# 统一使用 HTTPStatusError,让 TaskService/error_classifier 负责分类(客户端错误/兼容性错误/限流等)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
@@ -1233,6 +1576,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
upstream_response=raw_content,
|
||||
)
|
||||
|
||||
if envelope:
|
||||
response_json = envelope.unwrap_response(response_json)
|
||||
envelope.postprocess_unwrapped_response(model=model, data=response_json)
|
||||
|
||||
# 检查响应体中的嵌套错误(HTTP 200 但响应体包含错误)
|
||||
if isinstance(response_json, dict):
|
||||
parser = get_parser_for_format(provider_api_format)
|
||||
|
||||
@@ -44,6 +44,9 @@ from src.api.handlers.base.response_parser import (
|
||||
ResponseParser,
|
||||
)
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.upstream_stream_bridge import (
|
||||
aggregate_upstream_stream_to_internal_response,
|
||||
)
|
||||
from src.api.handlers.base.utils import (
|
||||
build_sse_headers,
|
||||
check_html_response,
|
||||
@@ -53,6 +56,9 @@ from src.api.handlers.base.utils import (
|
||||
)
|
||||
from src.config.constants import StreamDefaults
|
||||
from src.config.settings import config
|
||||
from src.core.api_format.conversion.stream_bridge import (
|
||||
iter_internal_response_as_stream_events,
|
||||
)
|
||||
from src.core.error_utils import extract_client_error_message
|
||||
from src.core.exceptions import (
|
||||
EmbeddedErrorException,
|
||||
@@ -72,6 +78,12 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
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.system.config import SystemConfigService
|
||||
from src.utils.sse_parser import SSEEventParser
|
||||
@@ -702,6 +714,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx.final_response = None
|
||||
ctx.response_id = None
|
||||
ctx.response_metadata = {} # 重置 Provider 响应元数据
|
||||
ctx.selected_base_url = None # 重置本次请求选用的 base_url(重试时避免污染)
|
||||
|
||||
# 记录 Provider 信息
|
||||
ctx.provider_name = str(provider.name)
|
||||
@@ -740,9 +753,26 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
)
|
||||
ctx.needs_conversion = needs_conversion
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
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,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
@@ -752,8 +782,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
ctx.model,
|
||||
is_stream=True,
|
||||
target_variant=target_variant,
|
||||
is_stream=upstream_is_stream,
|
||||
target_variant=conversion_variant,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
@@ -762,7 +792,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or ctx.model
|
||||
)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
if target_variant and provider_api_format:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
@@ -771,9 +801,31 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
# 获取认证信息(处理 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,
|
||||
)
|
||||
|
||||
# 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 方式
|
||||
@@ -782,9 +834,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=True,
|
||||
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,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when in streaming mode.
|
||||
provider_headers["Accept"] = "text/event-stream"
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
ctx.provider_request_headers = provider_headers
|
||||
@@ -794,10 +850,146 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=True, # CLI handler 处理流式请求
|
||||
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
|
||||
|
||||
# If upstream is forced to non-stream mode, we execute a sync request and then
|
||||
# simulate streaming to the client (sync -> stream bridge).
|
||||
if not upstream_is_stream:
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
|
||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||
http_client = await HTTPClientPool.get_proxy_client(
|
||||
proxy_config=provider.proxy,
|
||||
)
|
||||
|
||||
try:
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout_sync),
|
||||
)
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
ctx.status_code = resp.status_code
|
||||
ctx.response_headers = dict(resp.headers)
|
||||
if envelope:
|
||||
envelope.on_http_status(base_url=ctx.selected_base_url, status_code=ctx.status_code)
|
||||
|
||||
# Reuse HTTPStatusError classification path (handled by TaskService/error_classifier).
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
# Safe JSON parsing.
|
||||
try:
|
||||
response_json = resp.json()
|
||||
except (UnicodeDecodeError, json.JSONDecodeError) as e:
|
||||
raw_content = ""
|
||||
try:
|
||||
raw_content = resp.text[:500] if resp.text else "(empty)"
|
||||
except Exception:
|
||||
raw_content = "(unable to read)"
|
||||
raise ProviderNotAvailableException(
|
||||
"上游服务返回了无效的响应",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=f"json_decode_error={type(e).__name__}: {raw_content}",
|
||||
)
|
||||
|
||||
if envelope:
|
||||
response_json = envelope.unwrap_response(response_json)
|
||||
envelope.postprocess_unwrapped_response(model=ctx.model, data=response_json)
|
||||
|
||||
# Embedded error detection (HTTP 200 but error body).
|
||||
if isinstance(response_json, dict) and provider_api_format:
|
||||
parser = get_parser_for_format(provider_api_format)
|
||||
if parser.is_error_response(response_json):
|
||||
parsed = parser.parse_response(response_json, 200)
|
||||
raise EmbeddedErrorException(
|
||||
provider_name=str(provider.name),
|
||||
error_code=parsed.embedded_status_code,
|
||||
error_message=parsed.error_message,
|
||||
error_status=parsed.error_type,
|
||||
)
|
||||
|
||||
# Extract Provider response metadata (best-effort).
|
||||
if isinstance(response_json, dict):
|
||||
ctx.response_metadata = self._extract_response_metadata(response_json)
|
||||
|
||||
# Convert sync JSON -> InternalResponse, then InternalResponse -> client stream events.
|
||||
registry = get_format_converter_registry()
|
||||
src_norm = registry.get_normalizer(provider_api_format) if provider_api_format else None
|
||||
if src_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {provider_api_format}")
|
||||
|
||||
internal_resp = src_norm.response_to_internal(
|
||||
response_json if isinstance(response_json, dict) else {}
|
||||
)
|
||||
internal_resp.model = str(ctx.model or internal_resp.model or "")
|
||||
if internal_resp.id:
|
||||
ctx.response_id = internal_resp.id
|
||||
|
||||
if internal_resp.usage:
|
||||
ctx.input_tokens = int(internal_resp.usage.input_tokens or 0)
|
||||
ctx.output_tokens = int(internal_resp.usage.output_tokens or 0)
|
||||
ctx.cached_tokens = int(internal_resp.usage.cache_read_tokens or 0)
|
||||
ctx.cache_creation_tokens = int(internal_resp.usage.cache_write_tokens or 0)
|
||||
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
tgt_norm = registry.get_normalizer(client_api_format) if client_api_format else None
|
||||
if tgt_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
|
||||
|
||||
state = StreamState(
|
||||
model=str(ctx.model or ""),
|
||||
message_id=str(ctx.response_id or ctx.request_id or self.request_id or ""),
|
||||
)
|
||||
output_state = {"first_yield": True, "streaming_updated": False}
|
||||
|
||||
async def _streamified() -> AsyncGenerator[bytes]:
|
||||
for ev in iter_internal_response_as_stream_events(internal_resp):
|
||||
converted_events = tgt_norm.stream_event_from_internal(ev, state)
|
||||
if not converted_events:
|
||||
continue
|
||||
self._record_converted_chunks(ctx, converted_events)
|
||||
for sse_line in _format_converted_events_to_sse(
|
||||
converted_events, client_api_format
|
||||
):
|
||||
if not sse_line:
|
||||
continue
|
||||
ctx.chunk_count += 1
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield (sse_line + "\n").encode("utf-8")
|
||||
|
||||
# OpenAI chat clients expect a final [DONE] marker.
|
||||
if str(client_api_format or "").strip().lower() == "openai:chat":
|
||||
ctx.chunk_count += 1
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield b"data: [DONE]\n\n"
|
||||
ctx.has_completion = True
|
||||
|
||||
return _streamified()
|
||||
|
||||
# 配置 HTTP 超时
|
||||
# 注意:read timeout 用于检测连接断开,不是整体请求超时
|
||||
@@ -847,6 +1039,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
logger.debug(f" └─ 收到响应: status={stream_response.status_code}")
|
||||
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=ctx.selected_base_url,
|
||||
status_code=ctx.status_code,
|
||||
)
|
||||
|
||||
stream_response.raise_for_status()
|
||||
|
||||
# 使用字节流迭代器(避免 aiter_lines 的性能问题, aiter_bytes 会自动解压 gzip/deflate)
|
||||
@@ -871,7 +1069,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
|
||||
except TimeoutError:
|
||||
except TimeoutError as e:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
@@ -879,6 +1077,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
await http_client.aclose()
|
||||
logger.warning(
|
||||
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
|
||||
@@ -901,6 +1101,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_text = await self._extract_error_text(e)
|
||||
logger.error(
|
||||
@@ -958,6 +1168,14 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 使用已设置的 ctx.needs_conversion(由候选筛选阶段根据端点配置判断)
|
||||
# 不再调用 _needs_format_conversion,它只检查格式差异,不检查端点配置
|
||||
needs_conversion = ctx.needs_conversion
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=ctx.provider_type,
|
||||
endpoint_sig=ctx.provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
if envelope and envelope.force_stream_rewrite():
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
async for chunk in stream_response.aiter_bytes():
|
||||
buffer += chunk
|
||||
@@ -1321,6 +1539,14 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 使用已设置的 ctx.needs_conversion(由候选筛选阶段根据端点配置判断)
|
||||
# 不再调用 _needs_format_conversion,它只检查格式差异,不检查端点配置
|
||||
needs_conversion = ctx.needs_conversion
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=ctx.provider_type,
|
||||
endpoint_sig=ctx.provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
if envelope and envelope.force_stream_rewrite():
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
# 先处理预读的字节块
|
||||
for chunk in prefetched_chunks:
|
||||
@@ -1568,18 +1794,31 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except json.JSONDecodeError:
|
||||
return
|
||||
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=ctx.provider_type,
|
||||
endpoint_sig=ctx.provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
if envelope:
|
||||
data = envelope.unwrap_response(data)
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
|
||||
# 当不需要格式转换时,更新 data_count;需要记录时再写入 parsed_chunks。
|
||||
# 当需要格式转换时(record_chunk=False),data_count 由 _record_converted_chunks 更新
|
||||
if record_chunk and isinstance(data, dict):
|
||||
if record_chunk:
|
||||
ctx.data_count += 1
|
||||
if ctx.record_parsed_chunks:
|
||||
ctx.parsed_chunks.append(data)
|
||||
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
|
||||
event_type = event_name or data.get("type", "")
|
||||
|
||||
if envelope:
|
||||
envelope.postprocess_unwrapped_response(model=ctx.model, data=data)
|
||||
|
||||
# 调用格式特定的处理逻辑
|
||||
# 注意:跨格式转换时,_process_event_data 会自动选择正确的 Provider 解析器
|
||||
self._process_event_data(ctx, event_type, data)
|
||||
@@ -1928,6 +2167,17 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
logger.warning(f"[{ctx.request_id}] 流式请求失败,未选中提供商")
|
||||
return
|
||||
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=ctx.provider_type,
|
||||
endpoint_sig=ctx.provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=ctx.selected_base_url,
|
||||
status_code=ctx.status_code,
|
||||
)
|
||||
|
||||
# 获取新的 DB session
|
||||
db_gen = get_db()
|
||||
bg_db = next(db_gen)
|
||||
@@ -2326,9 +2576,26 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
)
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
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,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
@@ -2338,8 +2605,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
model,
|
||||
is_stream=False,
|
||||
target_variant=target_variant,
|
||||
is_stream=upstream_is_stream,
|
||||
target_variant=conversion_variant,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
@@ -2348,7 +2615,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or model
|
||||
)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
if target_variant and provider_api_format:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
@@ -2357,9 +2624,31 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
# 获取认证信息(处理 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,
|
||||
)
|
||||
|
||||
# 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 方式
|
||||
@@ -2368,9 +2657,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
original_headers,
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=False,
|
||||
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,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when forced to streaming mode.
|
||||
provider_headers["Accept"] = "text/event-stream"
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
provider_request_headers = provider_headers
|
||||
@@ -2380,13 +2673,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
path_params={"model": url_model},
|
||||
is_stream=False, # 非流式请求
|
||||
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
|
||||
|
||||
logger.info(
|
||||
f" └─ [{self.request_id}] 发送非流式请求: "
|
||||
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}"
|
||||
@@ -2405,49 +2700,106 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||
# 超时通过 timeout 参数控制
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout),
|
||||
)
|
||||
resp: httpx.Response | None = None
|
||||
if not upstream_is_stream:
|
||||
try:
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout),
|
||||
)
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
|
||||
if selected_base_url_cached:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
|
||||
)
|
||||
raise
|
||||
else:
|
||||
# Forced upstream streaming: aggregate SSE to a sync JSON response.
|
||||
registry = get_format_converter_registry()
|
||||
provider_parser = (
|
||||
get_parser_for_format(provider_api_format) if provider_api_format else None
|
||||
)
|
||||
|
||||
try:
|
||||
async with http_client.stream(
|
||||
"POST",
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout),
|
||||
) as stream_resp:
|
||||
resp = stream_resp
|
||||
|
||||
status_code = stream_resp.status_code
|
||||
response_headers = dict(stream_resp.headers)
|
||||
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=selected_base_url_cached,
|
||||
status_code=status_code,
|
||||
)
|
||||
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
internal_resp = await aggregate_upstream_stream_to_internal_response(
|
||||
stream_resp.aiter_bytes(),
|
||||
provider_api_format=provider_api_format,
|
||||
provider_name=str(provider.name),
|
||||
model=str(model or ""),
|
||||
request_id=str(self.request_id or ""),
|
||||
envelope=envelope,
|
||||
provider_parser=provider_parser,
|
||||
)
|
||||
|
||||
tgt_norm = (
|
||||
registry.get_normalizer(client_api_format)
|
||||
if client_api_format
|
||||
else None
|
||||
)
|
||||
if tgt_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
|
||||
|
||||
response_json = tgt_norm.response_from_internal(
|
||||
internal_resp,
|
||||
requested_model=model,
|
||||
)
|
||||
response_json = response_json if isinstance(response_json, dict) else {}
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
|
||||
if selected_base_url_cached:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
status_code = resp.status_code
|
||||
response_headers = dict(resp.headers)
|
||||
|
||||
if resp.status_code == 401:
|
||||
raise ProviderAuthException(str(provider.name))
|
||||
elif resp.status_code == 429:
|
||||
raise ProviderRateLimitException(
|
||||
"请求过于频繁,请稍后重试",
|
||||
provider_name=str(provider.name),
|
||||
response_headers=response_headers,
|
||||
retry_after=int(resp.headers.get("retry-after", 0)) or None,
|
||||
)
|
||||
elif resp.status_code >= 500:
|
||||
error_text = resp.text
|
||||
raise ProviderNotAvailableException(
|
||||
f"上游服务暂时不可用 (HTTP {resp.status_code})",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=error_text,
|
||||
)
|
||||
elif 300 <= resp.status_code < 400:
|
||||
redirect_url = resp.headers.get("location", "unknown")
|
||||
raise ProviderNotAvailableException(
|
||||
"上游服务返回重定向响应",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=f"重定向 {resp.status_code} -> {redirect_url}",
|
||||
)
|
||||
elif resp.status_code != 200:
|
||||
error_text = resp.text
|
||||
raise ProviderNotAvailableException(
|
||||
f"上游服务返回错误 (HTTP {resp.status_code})",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=error_text,
|
||||
)
|
||||
if envelope:
|
||||
envelope.on_http_status(base_url=selected_base_url_cached, status_code=status_code)
|
||||
|
||||
# Forced upstream streaming already built response_json via aggregator.
|
||||
if upstream_is_stream:
|
||||
response_metadata_result = self._extract_response_metadata(response_json or {})
|
||||
return response_json if isinstance(response_json, dict) else {}
|
||||
|
||||
# Reuse HTTPStatusError classification path (handled by TaskService/error_classifier).
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
# 安全解析 JSON 响应,处理可能的编码错误
|
||||
try:
|
||||
@@ -2483,6 +2835,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
upstream_response=raw_content,
|
||||
)
|
||||
|
||||
if envelope:
|
||||
response_json = envelope.unwrap_response(response_json)
|
||||
envelope.postprocess_unwrapped_response(model=model, data=response_json)
|
||||
|
||||
# 提取 Provider 响应元数据(子类可覆盖)
|
||||
response_metadata_result = self._extract_response_metadata(response_json)
|
||||
|
||||
@@ -2531,12 +2887,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
if response_json is None:
|
||||
response_json = {}
|
||||
|
||||
# 检查是否需要格式转换(同族格式无需转换,如 CLAUDE 和 CLAUDE_CLI)
|
||||
from src.core.api_format.utils import get_base_format
|
||||
|
||||
provider_base = get_base_format(provider_api_format) if provider_api_format else None
|
||||
client_base = get_base_format(api_format) if api_format else None
|
||||
if provider_base and client_base and provider_base != client_base:
|
||||
# 跨格式:响应转换回 client_format(失败不触发 failover,保守回退为原始响应)
|
||||
if (
|
||||
needs_conversion
|
||||
and provider_api_format
|
||||
and api_format
|
||||
and isinstance(response_json, dict)
|
||||
):
|
||||
try:
|
||||
registry = get_format_converter_registry()
|
||||
response_json = registry.convert_response(
|
||||
@@ -2835,6 +3192,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
if status == "invalid" or status == "passthrough":
|
||||
return [line], []
|
||||
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=ctx.provider_type,
|
||||
endpoint_sig=ctx.provider_api_format,
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
if envelope:
|
||||
data_obj = envelope.unwrap_response(data_obj)
|
||||
envelope.postprocess_unwrapped_response(model=ctx.model, data=data_obj)
|
||||
|
||||
# 初始化流式转换状态
|
||||
if ctx.stream_conversion_state is None:
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
@@ -586,7 +586,16 @@ async def get_provider_auth(
|
||||
pass
|
||||
|
||||
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||
return ProviderAuthInfo(auth_header="Authorization", auth_value=f"Bearer {decrypted_key}")
|
||||
|
||||
decrypted_auth_config: dict[str, Any] | None = None
|
||||
if isinstance(token_meta, dict) and token_meta:
|
||||
decrypted_auth_config = token_meta
|
||||
|
||||
return ProviderAuthInfo(
|
||||
auth_header="Authorization",
|
||||
auth_value=f"Bearer {decrypted_key}",
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
|
||||
if auth_type == "vertex_ai":
|
||||
from src.core.vertex_auth import VertexAuthError, VertexAuthService
|
||||
|
||||
@@ -40,6 +40,8 @@ class StreamContext:
|
||||
provider_name: str | None = None
|
||||
provider_id: str | None = None
|
||||
provider_type: str | None = None # Provider 类型(如 codex),用于元数据采集
|
||||
# Transport 层选中的 base_url(用于 URL 可用性更新/故障转移等场景)
|
||||
selected_base_url: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
attempt_id: str | None = None
|
||||
@@ -130,6 +132,7 @@ class StreamContext:
|
||||
self.final_response = None
|
||||
self.stream_conversion_state = None
|
||||
self.needs_conversion = False
|
||||
self.selected_base_url = None
|
||||
|
||||
@property
|
||||
def collected_text(self) -> str:
|
||||
|
||||
@@ -44,6 +44,7 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.models.database import Provider, ProviderEndpoint
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.utils.perf import PerfRecorder
|
||||
from src.utils.sse_parser import SSEEventParser
|
||||
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
||||
@@ -213,6 +214,11 @@ class StreamProcessor:
|
||||
"""
|
||||
prefetched_chunks: list = []
|
||||
parser = self.get_parser_for_provider(ctx)
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=str(getattr(ctx, "provider_type", "") or ""),
|
||||
endpoint_sig=str(getattr(ctx, "provider_api_format", "") or ""),
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
buffer = b""
|
||||
line_count = 0
|
||||
should_stop = False
|
||||
@@ -291,6 +297,14 @@ class StreamProcessor:
|
||||
break
|
||||
continue
|
||||
|
||||
# Provider envelope: unwrap SSE data chunk before error detection / trial conversion.
|
||||
if envelope and isinstance(data, dict):
|
||||
data = envelope.unwrap_response(data)
|
||||
envelope.postprocess_unwrapped_response(
|
||||
model=str(ctx.model or ""),
|
||||
data=data,
|
||||
)
|
||||
|
||||
# 使用解析器检查是否为错误响应
|
||||
if isinstance(data, dict) and parser.is_error_response(data):
|
||||
parsed = parser.parse_response(data, 200)
|
||||
@@ -431,6 +445,14 @@ class StreamProcessor:
|
||||
) or "unknown"
|
||||
# 使用 handler 层预计算的 needs_conversion(由 candidate 决定)
|
||||
needs_conversion = ctx.needs_conversion
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=str(getattr(ctx, "provider_type", "") or ""),
|
||||
endpoint_sig=str(getattr(ctx, "provider_api_format", "") or ""),
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
if envelope and envelope.force_stream_rewrite():
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
# 安全检查:needs_conversion 为 True 时,provider_format 必须有值
|
||||
if needs_conversion and not provider_format:
|
||||
|
||||
202
src/api/handlers/base/upstream_stream_bridge.py
Normal file
202
src/api/handlers/base/upstream_stream_bridge.py
Normal file
@@ -0,0 +1,202 @@
|
||||
"""Upstream stream bridging helpers (handler layer).
|
||||
|
||||
This module provides small utilities used when handler-layer policies force an
|
||||
upstream request to be streaming (SSE) even when the client asked for sync.
|
||||
|
||||
It intentionally stays lightweight and works with:
|
||||
- standard SSE `data: {...}` lines (OpenAI/Claude/etc.)
|
||||
- Gemini CLI JSON-array lines (best-effort)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import codecs
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.response_parser import ResponseParser
|
||||
from src.api.handlers.base.utils import get_format_converter_registry
|
||||
from src.core.api_format.conversion.internal import InternalResponse
|
||||
from src.core.api_format.conversion.stream_bridge import InternalStreamAggregator
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
from src.core.exceptions import EmbeddedErrorException
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.envelope import ProviderEnvelope
|
||||
|
||||
|
||||
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
|
||||
"""Parse `data: {...}` as JSON."""
|
||||
payload = line[5:].strip()
|
||||
if not payload:
|
||||
return None, "empty"
|
||||
try:
|
||||
return json.loads(payload), "ok"
|
||||
except json.JSONDecodeError:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
|
||||
"""Parse `event: xxx data: {...}` as JSON (best-effort)."""
|
||||
# Split only on the first " data:" occurrence.
|
||||
try:
|
||||
_, data_part = line.split(" data:", 1)
|
||||
except ValueError:
|
||||
return None, "invalid"
|
||||
payload = data_part.strip()
|
||||
if not payload:
|
||||
return None, "empty"
|
||||
try:
|
||||
return json.loads(payload), "ok"
|
||||
except json.JSONDecodeError:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
|
||||
"""Parse Gemini CLI JSON-array streaming line (best-effort).
|
||||
|
||||
Gemini CLI may stream objects in a JSON array form like:
|
||||
- "[{...},"
|
||||
- " {...},"
|
||||
- " {...}]"
|
||||
"""
|
||||
stripped = (line or "").strip()
|
||||
if not stripped:
|
||||
return None, "empty"
|
||||
|
||||
# Quick filter: must contain a JSON object boundary.
|
||||
if "{" not in stripped:
|
||||
return None, "skip"
|
||||
|
||||
candidate = stripped.lstrip(",").rstrip(",").strip()
|
||||
# Drop array brackets on edges.
|
||||
if candidate.startswith("["):
|
||||
candidate = candidate[1:].strip()
|
||||
if candidate.endswith("]"):
|
||||
candidate = candidate[:-1].strip()
|
||||
candidate = candidate.lstrip(",").rstrip(",").strip()
|
||||
|
||||
if not candidate:
|
||||
return None, "empty"
|
||||
try:
|
||||
return json.loads(candidate), "ok"
|
||||
except json.JSONDecodeError:
|
||||
logger.debug(f"Gemini JSON-array line skip: {stripped[:50]}")
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def parse_provider_stream_line_to_json(
|
||||
line: str,
|
||||
provider_format: str,
|
||||
) -> tuple[Any | None, str]:
|
||||
"""Best-effort parse for upstream streaming lines (SSE or Gemini JSON-array)."""
|
||||
|
||||
if not line:
|
||||
return None, "skip"
|
||||
|
||||
normalized = line.rstrip("\r").strip("\n")
|
||||
if not normalized or normalized.strip() == "":
|
||||
return None, "skip"
|
||||
|
||||
# Standard SSE data line.
|
||||
if normalized.startswith("data:"):
|
||||
# `data: [DONE]` is a sentinel.
|
||||
if normalized[5:].strip() == "[DONE]":
|
||||
return None, "skip"
|
||||
return _parse_sse_data_line(normalized)
|
||||
|
||||
# event + data on same line.
|
||||
if normalized.startswith("event:") and " data:" in normalized:
|
||||
return _parse_sse_event_data_line(normalized)
|
||||
|
||||
# Other control lines.
|
||||
if normalized.startswith(("event:", "id:", "retry:")):
|
||||
return None, "skip"
|
||||
|
||||
# Gemini JSON-array/chunked streaming (no SSE prefix).
|
||||
if str(provider_format or "").strip().lower().startswith("gemini"):
|
||||
return _parse_gemini_json_array_line(normalized)
|
||||
|
||||
return None, "skip"
|
||||
|
||||
|
||||
async def aggregate_upstream_stream_to_internal_response(
|
||||
byte_iter: AsyncIterator[bytes],
|
||||
*,
|
||||
provider_api_format: str,
|
||||
provider_name: str,
|
||||
model: str,
|
||||
request_id: str,
|
||||
envelope: ProviderEnvelope | None = None,
|
||||
provider_parser: ResponseParser | None = None,
|
||||
) -> InternalResponse:
|
||||
"""Aggregate upstream SSE/streaming bytes into an InternalResponse (best-effort)."""
|
||||
|
||||
registry = get_format_converter_registry()
|
||||
src_norm = registry.get_normalizer(provider_api_format) if provider_api_format else None
|
||||
if src_norm is None:
|
||||
raise RuntimeError(f"未注册 Normalizer: {provider_api_format}")
|
||||
if not getattr(src_norm, "capabilities", None) or not src_norm.capabilities.supports_stream:
|
||||
raise RuntimeError(f"上游格式不支持流式: {provider_api_format}")
|
||||
|
||||
state = StreamState(model=str(model or ""), message_id=str(request_id or ""))
|
||||
aggregator = InternalStreamAggregator(
|
||||
fallback_id=str(request_id or "resp"),
|
||||
fallback_model=str(model or ""),
|
||||
)
|
||||
|
||||
buffer = b""
|
||||
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
|
||||
|
||||
def _feed_line(normalized_line: str) -> None:
|
||||
data_obj, st = parse_provider_stream_line_to_json(normalized_line, provider_api_format)
|
||||
if st != "ok" or data_obj is None:
|
||||
return
|
||||
if not isinstance(data_obj, dict):
|
||||
return
|
||||
|
||||
if envelope:
|
||||
unwrapped = envelope.unwrap_response(data_obj)
|
||||
if not isinstance(unwrapped, dict):
|
||||
return
|
||||
data_obj = unwrapped
|
||||
envelope.postprocess_unwrapped_response(model=model, data=data_obj)
|
||||
|
||||
if provider_parser and provider_parser.is_error_response(data_obj):
|
||||
parsed = provider_parser.parse_response(data_obj, 200)
|
||||
raise EmbeddedErrorException(
|
||||
provider_name=str(provider_name),
|
||||
error_code=parsed.embedded_status_code,
|
||||
error_message=parsed.error_message,
|
||||
error_status=parsed.error_type,
|
||||
)
|
||||
|
||||
internal_events = src_norm.stream_chunk_to_internal(data_obj, state)
|
||||
aggregator.feed(internal_events)
|
||||
|
||||
async for chunk in byte_iter:
|
||||
buffer += chunk
|
||||
while b"\n" in buffer:
|
||||
line_bytes, buffer = buffer.split(b"\n", 1)
|
||||
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
|
||||
normalized_line = line.rstrip("\r")
|
||||
|
||||
_feed_line(normalized_line)
|
||||
|
||||
# Flush remaining buffered bytes (in case upstream doesn't end with newline).
|
||||
if buffer:
|
||||
try:
|
||||
tail = decoder.decode(buffer, True)
|
||||
except Exception:
|
||||
tail = ""
|
||||
normalized_tail = (tail or "").rstrip("\r\n")
|
||||
if normalized_tail:
|
||||
_feed_line(normalized_tail)
|
||||
|
||||
return aggregator.build()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"aggregate_upstream_stream_to_internal_response",
|
||||
"parse_provider_stream_line_to_json",
|
||||
]
|
||||
@@ -96,7 +96,12 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
# 优先使用映射后的模型名,否则使用请求体中的
|
||||
return mapped_model or request_body.get("model")
|
||||
|
||||
def _extract_usage_from_event(self, event: dict[str, Any]) -> dict[str, int]:
|
||||
def _extract_usage_from_event(
|
||||
self,
|
||||
event: dict[str, Any],
|
||||
*,
|
||||
provider_type: str | None = None,
|
||||
) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 事件中提取 token 使用情况
|
||||
|
||||
@@ -104,10 +109,14 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
Args:
|
||||
event: Gemini 流式响应事件
|
||||
provider_type: Provider 类型(用于 Antigravity 特判)
|
||||
|
||||
Returns:
|
||||
包含 input_tokens, output_tokens, cached_tokens 的字典
|
||||
"""
|
||||
if str(provider_type or "").lower() == "antigravity":
|
||||
return self._extract_antigravity_usage(event)
|
||||
|
||||
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
||||
|
||||
usage = GeminiStreamParser().extract_usage(event)
|
||||
@@ -125,6 +134,35 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
"cached_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def _extract_antigravity_usage(self, event: dict[str, Any]) -> dict[str, int]:
|
||||
"""Antigravity 专用 usage 提取(宽松 + 边界保护)。
|
||||
|
||||
Antigravity 的 usageMetadata 可能缺少 totalTokenCount,因此不能依赖
|
||||
GeminiStreamParser.extract_usage 的“totalTokenCount 必须存在”的严格判断。
|
||||
"""
|
||||
usage_metadata = event.get("usageMetadata", {})
|
||||
if not isinstance(usage_metadata, dict) or not usage_metadata:
|
||||
return {"input_tokens": 0, "output_tokens": 0, "cached_tokens": 0}
|
||||
|
||||
def _as_int(v: Any) -> int:
|
||||
try:
|
||||
return int(v or 0)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
prompt = _as_int(usage_metadata.get("promptTokenCount"))
|
||||
cached = _as_int(usage_metadata.get("cachedContentTokenCount"))
|
||||
candidates = _as_int(usage_metadata.get("candidatesTokenCount"))
|
||||
thoughts = _as_int(usage_metadata.get("thoughtsTokenCount"))
|
||||
|
||||
return {
|
||||
# 注意:计费层会根据 api_family(GEMINI) 扣除 cache_read_tokens,
|
||||
# 因此这里保持 Gemini 口径:input_tokens=promptTokenCount(含缓存)。
|
||||
"input_tokens": max(0, prompt),
|
||||
"output_tokens": max(0, candidates + thoughts),
|
||||
"cached_tokens": max(0, cached),
|
||||
}
|
||||
|
||||
def _process_event_data(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
@@ -172,7 +210,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
ctx.final_response = data
|
||||
|
||||
# 提取使用量信息(复用 GeminiStreamParser.extract_usage)
|
||||
usage = self._extract_usage_from_event(data)
|
||||
usage = self._extract_usage_from_event(data, provider_type=ctx.provider_type)
|
||||
if usage["input_tokens"] > 0 or usage["output_tokens"] > 0:
|
||||
ctx.input_tokens = usage["input_tokens"]
|
||||
ctx.output_tokens = usage["output_tokens"]
|
||||
|
||||
@@ -6,7 +6,6 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -126,10 +125,10 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
|
||||
# 仅 Codex 端点添加特定头部
|
||||
if base_url and is_codex_url(base_url):
|
||||
headers["x-oai-web-search-eligible"] = "true"
|
||||
headers["session_id"] = str(uuid.uuid4())
|
||||
headers["accept"] = "text/event-stream"
|
||||
headers["originator"] = "codex_cli_rs"
|
||||
# 与运行时路径保持一致:使用 Codex envelope 的 best-effort headers。
|
||||
from src.services.codex.envelope import codex_oauth_envelope
|
||||
|
||||
headers.update(codex_oauth_envelope.extra_headers() or {})
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
Reference in New Issue
Block a user