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:
fawney19
2026-02-09 01:05:48 +08:00
parent e324ffdcd6
commit b34dd12863
82 changed files with 6498 additions and 840 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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