feat: Antigravity 和 Codex 服务支持

- 新增 Antigravity 服务:签名缓存、URL 可用性检测、信封处理
- 新增 Codex 服务:信封处理、元数据收集器
- 重构 provider transport 支持新的服务架构
- 新增 stream_bridge 和 upstream_stream_bridge 处理流式响应
- 优化 OAuth 工具函数
- 添加相关测试用例
This commit is contained in:
fawney19
2026-02-05 15:57:52 +08:00
parent ed2ff5c1d7
commit 440721368f
44 changed files with 3498 additions and 134 deletions

View File

@@ -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)

View File

@@ -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=Falsedata_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

View File

@@ -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

View File

@@ -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:

View File

@@ -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:

View 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",
]

View File

@@ -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"]

View File

@@ -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