mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: 新增 Kiro 适配器、OAuth 改进与多项功能增强
- 新增 Kiro provider 适配器(EventStream 协议解析、令牌管理、用量追踪) - 重构 OAuth 账户管理与统一配额机制 - 重构 Handler 基类(CLI adapter/handler、请求构建器、流处理器) - 增强缓存监控后端 API 与前端可视化 - 改进 Gemini 格式标准化器与请求头处理 - Antigravity/Codex 适配器更新,移除旧 metadata_collector - 新增数据库迁移:proxy provider API keys - 前端 UI 多项优化 Co-Authored-By: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
@@ -422,6 +422,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
用于根据目标模型的特性对请求体做最终调整,例如:
|
||||
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
|
||||
- 特定模型需要注入/移除某些字段
|
||||
- Gemini 格式:清理无效 parts 和合并连续同角色 contents
|
||||
|
||||
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
|
||||
|
||||
@@ -433,6 +434,18 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
Returns:
|
||||
调整后的请求体
|
||||
"""
|
||||
# Gemini 格式请求:清理无效 parts 和合并连续同角色 contents
|
||||
# 跨格式转换(如 Claude → Gemini)可能产生 thinking 等无法表示的块,
|
||||
# 导致 parts 为空或缺少有效 data-oneof 字段,被 Google API 拒绝。
|
||||
if provider_api_format and "gemini" in str(provider_api_format).lower():
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
return request_body
|
||||
|
||||
def _set_model_after_conversion(
|
||||
@@ -896,8 +909,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
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"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
ctx.provider_request_headers = provider_headers
|
||||
ctx.provider_request_body = provider_payload
|
||||
@@ -913,10 +927,15 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 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
|
||||
|
||||
# 记录代理信息
|
||||
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
get_proxy_label,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info,
|
||||
)
|
||||
|
||||
ctx.proxy_info = resolve_proxy_info(provider.proxy)
|
||||
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
|
||||
ctx.proxy_info = resolve_proxy_info(effective_proxy)
|
||||
proxy_label = get_proxy_label(ctx.proxy_info)
|
||||
|
||||
logger.debug(
|
||||
@@ -931,9 +950,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
from src.services.proxy_node.resolver import build_post_kwargs, resolve_delegate_config
|
||||
|
||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=effective_proxy
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1125,13 +1144,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 优先使用 Provider 配置,否则使用全局配置
|
||||
request_timeout = provider.stream_first_byte_timeout or config.stream_first_byte_timeout
|
||||
|
||||
# 创建 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 创建 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import build_stream_kwargs, resolve_delegate_config
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = HTTPClientPool.create_upstream_stream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy, timeout=timeout_config
|
||||
delegate_cfg, proxy_config=effective_proxy, timeout=timeout_config
|
||||
)
|
||||
|
||||
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
||||
@@ -1546,8 +1565,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
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"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_hdrs)
|
||||
|
||||
provider_request_headers = provider_hdrs
|
||||
provider_request_body = provider_payload
|
||||
@@ -1563,10 +1583,15 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
|
||||
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
# 记录代理信息
|
||||
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
get_proxy_label,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info,
|
||||
)
|
||||
|
||||
sync_proxy_info = resolve_proxy_info(provider.proxy)
|
||||
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
|
||||
sync_proxy_info = resolve_proxy_info(_effective_proxy)
|
||||
_proxy_label = get_proxy_label(sync_proxy_info)
|
||||
|
||||
logger.info(
|
||||
@@ -1576,7 +1601,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
)
|
||||
logger.debug(f" [{self.request_id}] 请求URL: {redact_url_for_log(url)}")
|
||||
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import (
|
||||
@@ -1589,9 +1614,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 优先使用 Provider 配置,否则使用全局配置
|
||||
request_timeout = provider.request_timeout or config.http_request_timeout
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=_effective_proxy
|
||||
)
|
||||
|
||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||
@@ -1643,8 +1668,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
byte_iter = stream_resp.aiter_bytes()
|
||||
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
|
||||
|
||||
internal_resp = await aggregate_upstream_stream_to_internal_response(
|
||||
stream_resp.aiter_bytes(),
|
||||
byte_iter,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_name=str(provider.name),
|
||||
model=str(model or ""),
|
||||
|
||||
@@ -643,10 +643,23 @@ class CliAdapterBase(ApiAdapter):
|
||||
from src.core.provider_types import ProviderType
|
||||
|
||||
is_antigravity = provider_type == ProviderType.ANTIGRAVITY
|
||||
is_kiro = provider_type == ProviderType.KIRO
|
||||
is_oauth = auth_type == "oauth"
|
||||
|
||||
# ---- URL ----
|
||||
if is_antigravity:
|
||||
if is_kiro:
|
||||
# Kiro 需要替换 base_url 中的 {region} 占位符,并使用专用路径
|
||||
from src.services.provider.adapters.kiro.constants import (
|
||||
DEFAULT_REGION,
|
||||
KIRO_GENERATE_ASSISTANT_PATH,
|
||||
)
|
||||
|
||||
region = (decrypted_auth_config or {}).get("region") or DEFAULT_REGION
|
||||
effective_base_url = (
|
||||
base_url.replace("{region}", region) if "{region}" in base_url else base_url
|
||||
)
|
||||
url = f"{str(effective_base_url).rstrip('/')}{KIRO_GENERATE_ASSISTANT_PATH}"
|
||||
elif is_antigravity:
|
||||
# Antigravity 走 v1internal 端点,模型名在请求体 envelope 中,不在 URL 路径里
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
V1INTERNAL_PATH_TEMPLATE,
|
||||
@@ -673,6 +686,25 @@ class CliAdapterBase(ApiAdapter):
|
||||
if is_antigravity:
|
||||
merged_extra["User-Agent"] = _get_antigravity_ua()
|
||||
|
||||
# Kiro 需要特定的请求头
|
||||
if is_kiro:
|
||||
from src.services.provider.adapters.kiro.headers import build_generate_assistant_headers
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import generate_machine_id
|
||||
|
||||
kiro_cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {})
|
||||
region = kiro_cfg.region or DEFAULT_REGION
|
||||
machine_id = generate_machine_id(kiro_cfg)
|
||||
kiro_headers = build_generate_assistant_headers(
|
||||
host=f"q.{region}.amazonaws.com",
|
||||
access_token=api_key,
|
||||
machine_id=machine_id,
|
||||
kiro_version=kiro_cfg.kiro_version,
|
||||
system_version=kiro_cfg.system_version,
|
||||
node_version=kiro_cfg.node_version,
|
||||
)
|
||||
merged_extra.update(kiro_headers)
|
||||
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
|
||||
# OAuth 统一处理:替换端点默认认证头为 Authorization: Bearer
|
||||
@@ -701,6 +733,21 @@ class CliAdapterBase(ApiAdapter):
|
||||
model=effective_model,
|
||||
)
|
||||
|
||||
# Kiro:用 conversationState envelope 包装请求体
|
||||
if is_kiro:
|
||||
from src.services.provider.adapters.kiro.converter import (
|
||||
convert_claude_messages_to_conversation_state,
|
||||
)
|
||||
|
||||
effective_model = model_name or request_data.get("model", "")
|
||||
conversation_state = convert_claude_messages_to_conversation_state(
|
||||
body,
|
||||
model=effective_model,
|
||||
)
|
||||
body = {"conversationState": conversation_state}
|
||||
if isinstance(kiro_cfg.profile_arn, str) and kiro_cfg.profile_arn.strip():
|
||||
body["profileArn"] = kiro_cfg.profile_arn.strip()
|
||||
|
||||
# ---- Header Rules ----
|
||||
if header_rules:
|
||||
from src.core.api_format import get_auth_config_for_endpoint as _get_auth_cfg
|
||||
|
||||
@@ -381,6 +381,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
用于根据目标模型的特性对请求体做最终调整,例如:
|
||||
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
|
||||
- 特定模型需要注入/移除某些字段
|
||||
- Gemini 格式:清理无效 parts 和合并连续同角色 contents
|
||||
|
||||
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
|
||||
|
||||
@@ -392,6 +393,18 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
Returns:
|
||||
调整后的请求体
|
||||
"""
|
||||
# Gemini 格式请求:清理无效 parts 和合并连续同角色 contents
|
||||
# 跨格式转换(如 Claude → Gemini)可能产生 thinking 等无法表示的块,
|
||||
# 导致 parts 为空或缺少有效 data-oneof 字段,被 Google API 拒绝。
|
||||
if provider_api_format and "gemini" in str(provider_api_format).lower():
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
return request_body
|
||||
|
||||
@staticmethod
|
||||
@@ -872,8 +885,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
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"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
ctx.provider_request_headers = provider_headers
|
||||
@@ -890,11 +904,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 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
|
||||
|
||||
# 记录代理信息(sync-bridge 路径,早于流式路径执行)
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import get_proxy_label as _gpl
|
||||
from src.services.proxy_node.resolver import resolve_effective_proxy as _rep
|
||||
from src.services.proxy_node.resolver import resolve_proxy_info as _rpi
|
||||
|
||||
ctx.proxy_info = _rpi(provider.proxy)
|
||||
effective_proxy = _rep(provider.proxy, getattr(key, "proxy", None))
|
||||
ctx.proxy_info = _rpi(effective_proxy)
|
||||
|
||||
# If upstream is forced to non-stream mode, we execute a sync request and then
|
||||
# simulate streaming to the client (sync -> stream bridge).
|
||||
@@ -903,9 +919,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
from src.services.proxy_node.resolver import build_post_kwargs, resolve_delegate_config
|
||||
|
||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=effective_proxy
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1096,13 +1112,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
f"timeout={request_timeout}s, 代理={_proxy_label}"
|
||||
)
|
||||
|
||||
# 创建 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 创建 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import build_stream_kwargs, resolve_delegate_config
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = HTTPClientPool.create_upstream_stream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy, timeout=timeout_config
|
||||
delegate_cfg, proxy_config=effective_proxy, timeout=timeout_config
|
||||
)
|
||||
|
||||
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
||||
@@ -1297,7 +1313,25 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
async for chunk in stream_response.aiter_bytes():
|
||||
# Kiro 特殊处理:AWS Event Stream 二进制流需要重写为 SSE
|
||||
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
|
||||
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
chunk_source: AsyncGenerator[bytes, None] = apply_kiro_stream_rewrite(
|
||||
stream_response.aiter_bytes(),
|
||||
model=str(ctx.model or ""),
|
||||
input_tokens=int(ctx.input_tokens or 0),
|
||||
)
|
||||
# Kiro 重写后输出的是 Claude SSE 格式,不需要再进行格式转换
|
||||
needs_conversion = False
|
||||
ctx.needs_conversion = False
|
||||
else:
|
||||
chunk_source = stream_response.aiter_bytes()
|
||||
|
||||
async for chunk in chunk_source:
|
||||
buffer += chunk
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
@@ -1307,8 +1341,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"[{self.request_id}] UTF-8 解码失败: {e}, "
|
||||
f"bytes={line_bytes[:50]!r}"
|
||||
"[{}] UTF-8 解码失败: {}, bytes={!r}",
|
||||
self.request_id,
|
||||
e,
|
||||
line_bytes[:50],
|
||||
)
|
||||
continue
|
||||
|
||||
@@ -1333,11 +1369,14 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
|
||||
elapsed = time.time() - last_data_time
|
||||
if elapsed > self.DATA_TIMEOUT:
|
||||
logger.warning(f"Provider '{ctx.provider_name}' 流超时且无数据")
|
||||
# 设置错误状态用于后续记录
|
||||
logger.warning("Provider '{}' 流超时且无数据", ctx.provider_name)
|
||||
ctx.status_code = 504
|
||||
ctx.error_message = "流式响应超时,未收到有效数据"
|
||||
ctx.upstream_response = f"流超时: Provider={ctx.provider_name}, elapsed={elapsed:.1f}s, chunk_count={ctx.chunk_count}, data_count=0"
|
||||
ctx.upstream_response = (
|
||||
f"流超时: Provider={ctx.provider_name}, "
|
||||
f"elapsed={elapsed:.1f}s, "
|
||||
f"chunk_count={ctx.chunk_count}, data_count=0"
|
||||
)
|
||||
error_event = {
|
||||
"type": "error",
|
||||
"error": {
|
||||
@@ -1364,16 +1403,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield (line + "\n").encode("utf-8")
|
||||
|
||||
for event in events:
|
||||
self._handle_sse_event(
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
for event in events:
|
||||
self._handle_sse_event(
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
if ctx.data_count > 0:
|
||||
last_data_time = time.time()
|
||||
if ctx.data_count > 0:
|
||||
last_data_time = time.time()
|
||||
|
||||
# 处理剩余事件
|
||||
for event in sse_parser.flush():
|
||||
@@ -1386,13 +1425,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 检查是否收到数据
|
||||
if ctx.data_count == 0:
|
||||
# 流已开始,无法抛出异常进行故障转移
|
||||
# 发送错误事件并记录日志
|
||||
logger.warning(f"Provider '{ctx.provider_name}' 返回空流式响应")
|
||||
# 设置错误状态用于后续记录
|
||||
logger.warning("Provider '{}' 返回空流式响应", ctx.provider_name)
|
||||
ctx.status_code = 503
|
||||
ctx.error_message = "上游服务返回了空的流式响应"
|
||||
ctx.upstream_response = f"空流式响应: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0"
|
||||
ctx.upstream_response = (
|
||||
f"空流式响应: Provider={ctx.provider_name}, "
|
||||
f"chunk_count={ctx.chunk_count}, data_count=0"
|
||||
)
|
||||
error_event = {
|
||||
"type": "error",
|
||||
"error": {
|
||||
@@ -1792,6 +1831,26 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
# Kiro 特殊处理:AWS Event Stream 二进制流需要重写为 SSE
|
||||
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
|
||||
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iterator = apply_kiro_stream_rewrite(
|
||||
byte_iterator,
|
||||
model=str(ctx.model or ""),
|
||||
input_tokens=int(ctx.input_tokens or 0),
|
||||
prefetched_chunks=list(prefetched_chunks) if prefetched_chunks else None,
|
||||
)
|
||||
prefetched_chunks = []
|
||||
|
||||
# Kiro 重写后输出的是 Claude SSE 格式
|
||||
# 客户端也是 Claude CLI,不需要再进行格式转换
|
||||
needs_conversion = False
|
||||
ctx.needs_conversion = False
|
||||
|
||||
# 先处理预读的字节块
|
||||
for chunk in prefetched_chunks:
|
||||
buffer += chunk
|
||||
@@ -2461,10 +2520,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
try:
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
|
||||
# 采集上游元数据(仅成功请求)
|
||||
if ctx.is_success():
|
||||
self._collect_upstream_metadata(bg_db, ctx)
|
||||
|
||||
user = bg_db.query(User).filter(User.id == ctx.user_id).first()
|
||||
api_key = bg_db.query(ApiKeyModel).filter(ApiKeyModel.id == ctx.api_key_id).first()
|
||||
|
||||
@@ -2719,19 +2774,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except Exception as e:
|
||||
logger.exception("记录流式统计信息时出错")
|
||||
|
||||
@staticmethod
|
||||
def _collect_upstream_metadata(db: Session, ctx: StreamContext) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)"""
|
||||
from src.services.provider.metadata_collectors import collect_and_save_upstream_metadata
|
||||
|
||||
collect_and_save_upstream_metadata(
|
||||
db,
|
||||
provider_type=ctx.provider_type or "",
|
||||
key_id=ctx.key_id or "",
|
||||
response_headers=ctx.response_headers or {},
|
||||
request_id=ctx.request_id or "",
|
||||
)
|
||||
|
||||
async def _record_stream_failure(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
@@ -2960,8 +3002,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
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"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
provider_request_headers = provider_headers
|
||||
@@ -2978,10 +3021,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
|
||||
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
# 记录代理信息
|
||||
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
get_proxy_label,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info,
|
||||
)
|
||||
|
||||
sync_proxy_info = resolve_proxy_info(provider.proxy)
|
||||
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
|
||||
sync_proxy_info = resolve_proxy_info(_effective_proxy)
|
||||
_proxy_label = get_proxy_label(sync_proxy_info)
|
||||
|
||||
logger.info(
|
||||
@@ -2992,7 +3040,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
f"代理={_proxy_label}"
|
||||
)
|
||||
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import (
|
||||
@@ -3005,9 +3053,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 优先使用 Provider 配置,否则使用全局配置
|
||||
request_timeout = provider.request_timeout or config.http_request_timeout
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=_effective_proxy
|
||||
)
|
||||
|
||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||
@@ -3060,8 +3108,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
byte_iter = stream_resp.aiter_bytes()
|
||||
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
|
||||
|
||||
internal_resp = await aggregate_upstream_stream_to_internal_response(
|
||||
stream_resp.aiter_bytes(),
|
||||
byte_iter,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_name=str(provider.name),
|
||||
model=str(model or ""),
|
||||
|
||||
@@ -57,6 +57,73 @@ class ProviderAuthInfo:
|
||||
return (self.auth_header, self.auth_value)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh helpers
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _acquire_refresh_lock(key_id: str) -> tuple[Any, bool]:
|
||||
"""尝试获取 OAuth refresh 分布式锁。
|
||||
|
||||
返回 ``(redis_client | None, got_lock)``。调用方在刷新完成后
|
||||
必须调用 :func:`_release_refresh_lock` 释放锁。
|
||||
"""
|
||||
redis = await get_redis_client(require_redis=False)
|
||||
lock_key = f"provider_oauth_refresh_lock:{key_id}"
|
||||
got_lock = False
|
||||
if redis is not None:
|
||||
try:
|
||||
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
|
||||
except Exception:
|
||||
got_lock = False
|
||||
return redis, got_lock
|
||||
|
||||
|
||||
async def _release_refresh_lock(redis: Any, key_id: str) -> None:
|
||||
"""释放 OAuth refresh 分布式锁(best-effort)。"""
|
||||
if redis is not None:
|
||||
try:
|
||||
await redis.delete(f"provider_oauth_refresh_lock:{key_id}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _persist_refreshed_token(
|
||||
key: Any,
|
||||
access_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> None:
|
||||
"""将刷新后的 access_token 和 auth_config 持久化到数据库。"""
|
||||
key.api_key = crypto_service.encrypt(access_token)
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
sess = object_session(key)
|
||||
if sess is not None:
|
||||
sess.add(key)
|
||||
sess.commit()
|
||||
else:
|
||||
logger.warning(
|
||||
"[OAUTH_REFRESH] key {} refreshed but cannot persist (no session); "
|
||||
"next request will refresh again",
|
||||
key.id,
|
||||
)
|
||||
|
||||
|
||||
def _get_proxy_config(key: Any, endpoint: Any = None) -> Any:
|
||||
"""获取有效代理配置(Key 级别优先于 Provider 级别)。"""
|
||||
try:
|
||||
from src.services.proxy_node.resolver import resolve_effective_proxy
|
||||
|
||||
provider = getattr(key, "provider", None) or (
|
||||
getattr(endpoint, "provider", None) if endpoint else None
|
||||
)
|
||||
provider_proxy = getattr(provider, "proxy", None)
|
||||
key_proxy = getattr(key, "proxy", None)
|
||||
return resolve_effective_proxy(provider_proxy, key_proxy)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# 统一的头部配置常量
|
||||
# ==============================================================================
|
||||
@@ -591,6 +658,142 @@ def build_passthrough_request(
|
||||
)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh logic (Kiro / Generic)
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _refresh_kiro_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Kiro OAuth refresh: validate + call Kiro-specific refresh endpoint."""
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import (
|
||||
refresh_access_token,
|
||||
validate_refresh_token,
|
||||
)
|
||||
|
||||
cfg = KiroAuthConfig.from_dict(token_meta or {})
|
||||
if not (cfg.refresh_token or "").strip():
|
||||
raise InvalidRequestException(
|
||||
"Kiro auth_config missing refresh_token; please re-import credentials."
|
||||
)
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
access_token, new_cfg = await refresh_access_token(
|
||||
cfg,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
new_meta = new_cfg.to_dict()
|
||||
new_meta["updated_at"] = int(time.time())
|
||||
|
||||
_persist_refreshed_token(key, access_token, new_meta)
|
||||
return new_meta
|
||||
|
||||
|
||||
async def _refresh_generic_oauth_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
template: Any,
|
||||
provider_type: str,
|
||||
refresh_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Generic OAuth refresh via template (Codex, Antigravity, ClaudeCode, etc.)."""
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
scopes = getattr(template.oauth, "scopes", None) or []
|
||||
scope_str = " ".join(scopes) if scopes else ""
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if scope_str:
|
||||
body["scope"] = scope_str
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if scope_str:
|
||||
form["scope"] = scope_str
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=30.0,
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
_persist_refreshed_token(key, access_token, token_meta)
|
||||
else:
|
||||
logger.warning(
|
||||
"OAuth token refresh failed: provider={}, key_id={}, status={}",
|
||||
provider_type,
|
||||
getattr(key, "id", "?"),
|
||||
resp.status_code,
|
||||
)
|
||||
|
||||
return token_meta
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Service Account 认证支持
|
||||
# ==============================================================================
|
||||
@@ -660,113 +863,19 @@ async def get_provider_auth(
|
||||
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
|
||||
except Exception:
|
||||
template = None
|
||||
if template:
|
||||
redis = await get_redis_client(require_redis=False)
|
||||
lock_key = f"provider_oauth_refresh_lock:{key.id}"
|
||||
got_lock = False
|
||||
if redis is not None:
|
||||
try:
|
||||
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
|
||||
except Exception:
|
||||
got_lock = False
|
||||
|
||||
if got_lock or redis is None:
|
||||
try:
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
proxy_config = None
|
||||
try:
|
||||
provider = getattr(key, "provider", None)
|
||||
proxy_config = getattr(provider, "proxy", None)
|
||||
except Exception:
|
||||
proxy_config = None
|
||||
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=30.0,
|
||||
redis, got_lock = await _acquire_refresh_lock(key.id)
|
||||
if got_lock or redis is None:
|
||||
try:
|
||||
if provider_type == ProviderType.KIRO.value:
|
||||
token_meta = await _refresh_kiro_token(key, endpoint, token_meta)
|
||||
elif template:
|
||||
token_meta = await _refresh_generic_oauth_token(
|
||||
key, endpoint, template, provider_type, refresh_token, token_meta
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
key.api_key = crypto_service.encrypt(access_token)
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
# 持久化:key 实体来自 DB session 时,尝试直接提交更新。
|
||||
sess = object_session(key)
|
||||
if sess is not None:
|
||||
sess.add(key)
|
||||
sess.commit()
|
||||
else:
|
||||
logger.warning(
|
||||
"[OAUTH_REFRESH] key {} 刷新成功但无法持久化(无绑定 session),"
|
||||
"下次请求将重新刷新",
|
||||
key.id,
|
||||
)
|
||||
finally:
|
||||
if got_lock and redis is not None:
|
||||
try:
|
||||
await redis.delete(lock_key)
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
if got_lock:
|
||||
await _release_refresh_lock(redis, key.id)
|
||||
except Exception:
|
||||
# 刷新失败不阻断请求;后续由上游返回 401 再触发管理端处理
|
||||
pass
|
||||
@@ -782,7 +891,6 @@ async def get_provider_auth(
|
||||
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
|
||||
|
||||
|
||||
@@ -383,6 +383,10 @@ class StreamProcessor:
|
||||
endpoint_sig=str(getattr(ctx, "provider_api_format", "") or ""),
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
ctx_provider_type = str(getattr(ctx, "provider_type", "") or "").strip().lower()
|
||||
kiro_binary_stream = (
|
||||
ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite()
|
||||
)
|
||||
buffer = b""
|
||||
line_count = 0
|
||||
should_stop = False
|
||||
@@ -402,6 +406,11 @@ class StreamProcessor:
|
||||
)
|
||||
prefetched_chunks.append(first_chunk)
|
||||
total_prefetched_bytes += len(first_chunk)
|
||||
|
||||
# Kiro upstream uses AWS Event Stream (binary). Do not attempt to split/decode lines here;
|
||||
# we only enforce TTFB and let StreamProcessor rewrite bytes later.
|
||||
if kiro_binary_stream:
|
||||
return prefetched_chunks
|
||||
buffer += first_chunk
|
||||
|
||||
# 继续读取剩余的预读数据
|
||||
@@ -619,6 +628,26 @@ class StreamProcessor:
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
ctx_provider_type = str(getattr(ctx, "provider_type", "") or "").strip().lower()
|
||||
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iterator = apply_kiro_stream_rewrite(
|
||||
byte_iterator,
|
||||
model=str(ctx.model or ""),
|
||||
input_tokens=int(ctx.input_tokens or 0),
|
||||
prefetched_chunks=list(prefetched_chunks) if prefetched_chunks else None,
|
||||
)
|
||||
prefetched_chunks = None
|
||||
|
||||
# Kiro 重写后输出的是 Claude SSE 格式(data: {...}\n\n)
|
||||
# 如果客户端也是 Claude 格式,则不需要再进行格式转换
|
||||
if client_family == "claude":
|
||||
needs_conversion = False
|
||||
ctx.needs_conversion = False
|
||||
|
||||
# 安全检查:needs_conversion 为 True 时,provider_format 必须有值
|
||||
if needs_conversion and not provider_format:
|
||||
logger.warning(
|
||||
|
||||
@@ -97,10 +97,6 @@ class StreamTelemetryRecorder:
|
||||
bg_db = next(db_gen)
|
||||
|
||||
try:
|
||||
# 采集上游元数据(仅成功请求,放在 writer 获取之前以确保执行)
|
||||
if ctx.is_success():
|
||||
self._collect_upstream_metadata(bg_db, ctx)
|
||||
|
||||
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
|
||||
if writer is None:
|
||||
return
|
||||
@@ -537,19 +533,6 @@ class StreamTelemetryRecorder:
|
||||
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _collect_upstream_metadata(db: Session, ctx: StreamContext) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)"""
|
||||
from src.services.provider.metadata_collectors import collect_and_save_upstream_metadata
|
||||
|
||||
collect_and_save_upstream_metadata(
|
||||
db,
|
||||
provider_type=ctx.provider_type or "",
|
||||
key_id=ctx.key_id or "",
|
||||
response_headers=ctx.response_headers or {},
|
||||
request_id=ctx.request_id or "",
|
||||
)
|
||||
|
||||
def _get_status_from_ctx(self, ctx: StreamContext) -> str:
|
||||
"""根据上下文获取状态字符串"""
|
||||
if ctx.is_success():
|
||||
|
||||
@@ -169,6 +169,17 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
# Sanitize Gemini contents: strip parts without a valid data-oneof
|
||||
# field and merge consecutive same-role entries. This catches cases
|
||||
# missed by the normalizer (passthrough) or the antigravity envelope.
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
@@ -90,6 +90,17 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
# Sanitize Gemini contents: strip parts without a valid data-oneof
|
||||
# field and merge consecutive same-role entries. This catches cases
|
||||
# missed by the normalizer (passthrough) or the antigravity envelope.
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
Reference in New Issue
Block a user