mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
feat: OAuth 导入导出、提供商筛选、Gemini 图像生成支持与流式处理增强
- OAuth: 支持通过 Refresh Token 导入账号(文件拖拽/粘贴),OAuth Key 可导出为 JSON - OAuth: 所有 OAuth 端点添加 require_admin 鉴权 - 提供商管理: 新增状态/API格式/模型三级筛选,后端返回 global_model_ids - Gemini: 新增图像生成模型适配(finalize_provider_request 钩子 + envelope 跳过不兼容字段) - 流式处理: buffer 残留数据 flush 与 token 兜底估算 - 上游元数据: 提取 merge_upstream_metadata,配额耗尽模型保留与深度合并 - Antigravity 配额: 无 quotaInfo 时视为耗尽,移除 Other 兜底分组 - README: 新增升级备份与回滚指南
This commit is contained in:
@@ -409,6 +409,32 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
"""
|
||||
return request_body
|
||||
|
||||
def finalize_provider_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: str | None,
|
||||
provider_api_format: str | None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
格式转换完成后、envelope 之前的模型感知后处理钩子 - 子类可覆盖
|
||||
|
||||
用于根据目标模型的特性对请求体做最终调整,例如:
|
||||
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
|
||||
- 特定模型需要注入/移除某些字段
|
||||
|
||||
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
|
||||
|
||||
Args:
|
||||
request_body: 已完成格式转换的请求体
|
||||
mapped_model: 映射后的目标模型名
|
||||
provider_api_format: Provider 侧 API 格式标识
|
||||
|
||||
Returns:
|
||||
调整后的请求体
|
||||
"""
|
||||
return request_body
|
||||
|
||||
def _set_model_after_conversion(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
@@ -827,6 +853,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
target_variant=same_format_variant,
|
||||
)
|
||||
|
||||
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=str(provider_api_format) if provider_api_format else None,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
@@ -1440,6 +1473,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
target_variant=same_format_variant,
|
||||
)
|
||||
|
||||
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=str(provider_api_format) if provider_api_format else None,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
|
||||
@@ -368,6 +368,32 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"""
|
||||
return request_body
|
||||
|
||||
def finalize_provider_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: str | None,
|
||||
provider_api_format: str | None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
格式转换完成后、envelope 之前的模型感知后处理钩子 - 子类可覆盖
|
||||
|
||||
用于根据目标模型的特性对请求体做最终调整,例如:
|
||||
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
|
||||
- 特定模型需要注入/移除某些字段
|
||||
|
||||
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
|
||||
|
||||
Args:
|
||||
request_body: 已完成格式转换的请求体
|
||||
mapped_model: 映射后的目标模型名
|
||||
provider_api_format: Provider 侧 API 格式标识
|
||||
|
||||
Returns:
|
||||
调整后的请求体
|
||||
"""
|
||||
return request_body
|
||||
|
||||
@staticmethod
|
||||
def _get_format_metadata(format_id: str) -> "EndpointDefinition | None":
|
||||
"""获取 endpoint 元数据(解析失败返回 None)"""
|
||||
@@ -801,6 +827,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
@@ -2812,6 +2845,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
|
||||
request_body = self.finalize_provider_request(
|
||||
request_body,
|
||||
mapped_model=mapped_model,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
|
||||
# Force upstream stream/sync mode in request body (best-effort).
|
||||
if provider_api_format:
|
||||
enforce_stream_mode_for_upstream(
|
||||
|
||||
@@ -708,6 +708,7 @@ class StreamProcessor:
|
||||
f"[{self.request_id}] 处理剩余缓冲区失败: {e}, bytes={buffer[:50]!r}"
|
||||
)
|
||||
line = ""
|
||||
buffer = b"" # 标记已消费,避免 finally 中重复处理
|
||||
if line:
|
||||
# 需要格式转换时,跳过记录原始数据
|
||||
_process_line_with_perf(line, skip_record=True)
|
||||
@@ -782,8 +783,26 @@ class StreamProcessor:
|
||||
logger.warning(
|
||||
f"[{self.request_id}] 处理剩余缓冲区失败: {e}, bytes={buffer[:50]!r}"
|
||||
)
|
||||
buffer = b"" # 标记已消费,避免下方重复处理
|
||||
|
||||
# 处理剩余事件
|
||||
# flush 残留的字节 buffer(异常中断时 buffer 可能仍有未解析的数据,
|
||||
# 如包含 usage 的 message_delta/response.completed 事件)
|
||||
# 正常结束时 buffer 已在上方被消费为空,此处为 no-op
|
||||
if buffer:
|
||||
try:
|
||||
remaining = decoder.decode(buffer, True)
|
||||
for line in remaining.split("\n"):
|
||||
stripped = line.rstrip("\r\n")
|
||||
if stripped:
|
||||
events = sse_parser.feed_line(stripped)
|
||||
for event in events:
|
||||
self.handle_sse_event(
|
||||
ctx, event.get("event"), event.get("data") or ""
|
||||
)
|
||||
except Exception:
|
||||
pass # best-effort: 不应因 flush 失败影响后续流程
|
||||
|
||||
# flush SSE parser 内部累积的未完成事件
|
||||
for event in sse_parser.flush():
|
||||
self.handle_sse_event(ctx, event.get("event"), event.get("data") or "")
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
@@ -104,6 +105,19 @@ class StreamTelemetryRecorder:
|
||||
if writer is None:
|
||||
return
|
||||
actual_request_body = ctx.provider_request_body or original_request_body
|
||||
|
||||
# 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算
|
||||
# 覆盖 Chat Handler 路径(CLI Handler 在更早的位置已做估算,
|
||||
# 若已估算过则 token > 0,此处条件不会触发)
|
||||
if (
|
||||
ctx.is_success()
|
||||
and not ctx.has_completion
|
||||
and ctx.data_count > 0
|
||||
and ctx.input_tokens == 0
|
||||
and ctx.output_tokens == 0
|
||||
):
|
||||
self._estimate_tokens_for_incomplete_stream(ctx, actual_request_body)
|
||||
|
||||
should_log_body = SystemConfigService.should_log_body(bg_db)
|
||||
include_bodies = (
|
||||
writer.include_bodies
|
||||
@@ -536,6 +550,59 @@ class StreamTelemetryRecorder:
|
||||
return "cancelled"
|
||||
return "failed"
|
||||
|
||||
@staticmethod
|
||||
def _estimate_tokens_for_incomplete_stream(
|
||||
ctx: StreamContext,
|
||||
request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
流未正常完成(无 response.completed)且 token 均为 0 时的兜底估算。
|
||||
|
||||
从已收集的输出文本和请求体粗略估算 token 数,确保 usage 记录不为 0。
|
||||
估算采用 ~4 字符/token 的保守比例。
|
||||
"""
|
||||
# 输出 tokens:从已收集的文本估算
|
||||
collected = ctx.collected_text
|
||||
if collected:
|
||||
ctx.output_tokens = max(1, len(collected) // 4)
|
||||
|
||||
# 输入 tokens:从请求体文本内容估算
|
||||
try:
|
||||
total_input_len = 0
|
||||
instructions = request_body.get("instructions")
|
||||
if isinstance(instructions, str):
|
||||
total_input_len += len(instructions)
|
||||
# OpenAI Responses API 使用 input 字段;Claude 使用 messages
|
||||
input_items = request_body.get("input") or request_body.get("messages") or []
|
||||
if isinstance(input_items, list):
|
||||
for item in input_items:
|
||||
if isinstance(item, str):
|
||||
total_input_len += len(item)
|
||||
elif isinstance(item, dict):
|
||||
content = item.get("content", "")
|
||||
if isinstance(content, str):
|
||||
total_input_len += len(content)
|
||||
elif isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict):
|
||||
text = block.get("text", "")
|
||||
if isinstance(text, str):
|
||||
total_input_len += len(text)
|
||||
if total_input_len > 0:
|
||||
ctx.input_tokens = max(1, total_input_len // 4)
|
||||
else:
|
||||
# fallback: 整个请求体 JSON 大小
|
||||
body_str = json.dumps(request_body, ensure_ascii=False)
|
||||
ctx.input_tokens = max(1, len(body_str) // 4)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if ctx.input_tokens > 0 or ctx.output_tokens > 0:
|
||||
logger.warning(
|
||||
f"[{ctx.request_id}] 流未正常完成 (has_completion=False, data_count={ctx.data_count}), "
|
||||
f"使用估算 tokens: in={ctx.input_tokens}, out={ctx.output_tokens}"
|
||||
)
|
||||
|
||||
def _build_db_writer(self, bg_db: Session) -> DbTelemetryWriter | None:
|
||||
user = bg_db.query(User).filter(User.id == self.user_id).first()
|
||||
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
|
||||
|
||||
@@ -157,6 +157,22 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def finalize_provider_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: str | None,
|
||||
provider_api_format: str | None, # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
from src.api.handlers.gemini.image_gen import (
|
||||
adapt_request_for_image_gen,
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化 Gemini 响应
|
||||
|
||||
45
src/api/handlers/gemini/image_gen.py
Normal file
45
src/api/handlers/gemini/image_gen.py
Normal file
@@ -0,0 +1,45 @@
|
||||
"""
|
||||
Gemini 图像生成模型请求适配
|
||||
|
||||
- 图像生成模型不支持 tools / system_instruction,需要移除
|
||||
- responseModalities / responseMimeType 与 imageConfig 冲突,需要移除
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
def is_image_gen_model(model: str | None) -> bool:
|
||||
"""判断是否为图像生成模型(模式匹配,覆盖 gemini-*-image / imagen-* 系列)"""
|
||||
if not model:
|
||||
return False
|
||||
m = model.lower()
|
||||
return "image" in m and ("gemini" in m or "imagen" in m)
|
||||
|
||||
|
||||
def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]:
|
||||
"""为图像生成模型清理不兼容字段"""
|
||||
# 移除图像生成不支持的顶层字段
|
||||
for key in ("tools", "tool_config", "toolConfig", "system_instruction", "systemInstruction"):
|
||||
if key in body:
|
||||
body.pop(key)
|
||||
|
||||
# 处理 generationConfig
|
||||
gc_key = "generationConfig" if "generationConfig" in body else "generation_config"
|
||||
gc = body.get(gc_key)
|
||||
if not isinstance(gc, dict):
|
||||
gc = {}
|
||||
body[gc_key] = gc
|
||||
|
||||
# 移除与图像生成冲突的字段
|
||||
for key in (
|
||||
"responseMimeType",
|
||||
"response_mime_type",
|
||||
"responseModalities",
|
||||
"response_modalities",
|
||||
):
|
||||
gc.pop(key, None)
|
||||
|
||||
# 设置输出模态
|
||||
gc["responseModalities"] = ["TEXT", "IMAGE"]
|
||||
|
||||
return body
|
||||
@@ -78,6 +78,22 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
result.pop("model", None)
|
||||
return result
|
||||
|
||||
def finalize_provider_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: str | None,
|
||||
provider_api_format: str | None, # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
from src.api.handlers.gemini.image_gen import (
|
||||
adapt_request_for_image_gen,
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
|
||||
Reference in New Issue
Block a user