mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: 响应中保留用户请求的原始模型名
- 非流式响应:通过 requested_model 参数传递用户请求的模型名 - 流式响应:StreamState 初始化时使用 ctx.model 而非 ctx.mapped_model - 流式转换:normalizers 保留初始 model 值,不被上游事件覆盖 - parsed_chunks:格式转换时记录转换后数据,确保与客户端收到的一致
This commit is contained in:
@@ -1104,6 +1104,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
response_json,
|
response_json,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
|
requested_model=model, # 使用用户请求的原始模型名
|
||||||
)
|
)
|
||||||
|
|
||||||
return response_json if isinstance(response_json, dict) else {}
|
return response_json if isinstance(response_json, dict) else {}
|
||||||
|
|||||||
@@ -2350,6 +2350,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
response_json,
|
response_json,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
api_format,
|
api_format,
|
||||||
|
requested_model=model, # 使用用户请求的原始模型名
|
||||||
)
|
)
|
||||||
logger.debug(f"非流式响应格式转换完成: {provider_api_format} -> {api_format}")
|
logger.debug(f"非流式响应格式转换完成: {provider_api_format} -> {api_format}")
|
||||||
except Exception as conv_err:
|
except Exception as conv_err:
|
||||||
|
|||||||
@@ -107,6 +107,8 @@ class StreamProcessor:
|
|||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
event_name: Optional[str],
|
event_name: Optional[str],
|
||||||
data_str: str,
|
data_str: str,
|
||||||
|
*,
|
||||||
|
skip_record: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
处理单个 SSE 事件
|
处理单个 SSE 事件
|
||||||
@@ -117,6 +119,7 @@ class StreamProcessor:
|
|||||||
ctx: 流式上下文
|
ctx: 流式上下文
|
||||||
event_name: 事件名称
|
event_name: 事件名称
|
||||||
data_str: 事件数据字符串
|
data_str: 事件数据字符串
|
||||||
|
skip_record: 是否跳过记录到 parsed_chunks(当需要格式转换时应为 True)
|
||||||
"""
|
"""
|
||||||
if not data_str:
|
if not data_str:
|
||||||
return
|
return
|
||||||
@@ -130,13 +133,13 @@ class StreamProcessor:
|
|||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
return
|
return
|
||||||
|
|
||||||
ctx.data_count += 1
|
|
||||||
|
|
||||||
if not isinstance(data, dict):
|
if not isinstance(data, dict):
|
||||||
return
|
return
|
||||||
|
|
||||||
# 收集原始 chunk 数据
|
# 收集原始 chunk 数据(当需要格式转换时跳过,由 _emit_converted_line 记录转换后的数据)
|
||||||
ctx.parsed_chunks.append(data)
|
if not skip_record:
|
||||||
|
ctx.parsed_chunks.append(data)
|
||||||
|
ctx.data_count += 1
|
||||||
|
|
||||||
# 根据 Provider 格式选择解析器
|
# 根据 Provider 格式选择解析器
|
||||||
parser = self.get_parser_for_provider(ctx)
|
parser = self.get_parser_for_provider(ctx)
|
||||||
@@ -557,7 +560,13 @@ class StreamProcessor:
|
|||||||
|
|
||||||
skip_next_blank_line = True
|
skip_next_blank_line = True
|
||||||
out: list[bytes] = []
|
out: list[bytes] = []
|
||||||
|
|
||||||
for evt in converted_events:
|
for evt in converted_events:
|
||||||
|
# 记录转换后的数据到 parsed_chunks(这是客户端实际收到的格式)
|
||||||
|
if isinstance(evt, dict):
|
||||||
|
ctx.parsed_chunks.append(evt)
|
||||||
|
ctx.data_count += 1
|
||||||
|
|
||||||
# 统一使用 SSE 格式输出(Gemini streamGenerateContent 也使用 SSE)
|
# 统一使用 SSE 格式输出(Gemini streamGenerateContent 也使用 SSE)
|
||||||
# 参考: https://ai.google.dev/api/generate-content
|
# 参考: https://ai.google.dev/api/generate-content
|
||||||
out.append(
|
out.append(
|
||||||
@@ -580,7 +589,8 @@ class StreamProcessor:
|
|||||||
line = ""
|
line = ""
|
||||||
|
|
||||||
if line:
|
if line:
|
||||||
self._process_line(ctx, sse_parser, line)
|
# 需要格式转换时,跳过记录原始数据(由 _emit_converted_line 记录转换后的数据)
|
||||||
|
self._process_line(ctx, sse_parser, line, skip_record=True)
|
||||||
normalized_line = line.rstrip("\r\n") if line else ""
|
normalized_line = line.rstrip("\r\n") if line else ""
|
||||||
out_chunks = _emit_converted_line(normalized_line)
|
out_chunks = _emit_converted_line(normalized_line)
|
||||||
if not out_chunks:
|
if not out_chunks:
|
||||||
@@ -613,7 +623,8 @@ class StreamProcessor:
|
|||||||
line = ""
|
line = ""
|
||||||
|
|
||||||
if line:
|
if line:
|
||||||
self._process_line(ctx, sse_parser, line)
|
# 需要格式转换时,跳过记录原始数据(由 _emit_converted_line 记录转换后的数据)
|
||||||
|
self._process_line(ctx, sse_parser, line, skip_record=True)
|
||||||
normalized_line = line.rstrip("\r\n") if line else ""
|
normalized_line = line.rstrip("\r\n") if line else ""
|
||||||
out_chunks = _emit_converted_line(normalized_line)
|
out_chunks = _emit_converted_line(normalized_line)
|
||||||
if not out_chunks:
|
if not out_chunks:
|
||||||
@@ -642,7 +653,8 @@ class StreamProcessor:
|
|||||||
)
|
)
|
||||||
line = ""
|
line = ""
|
||||||
if line:
|
if line:
|
||||||
self._process_line(ctx, sse_parser, line)
|
# 需要格式转换时,跳过记录原始数据
|
||||||
|
self._process_line(ctx, sse_parser, line, skip_record=True)
|
||||||
normalized_line = line.rstrip("\r\n")
|
normalized_line = line.rstrip("\r\n")
|
||||||
out_chunks = _emit_converted_line(normalized_line)
|
out_chunks = _emit_converted_line(normalized_line)
|
||||||
for out in out_chunks:
|
for out in out_chunks:
|
||||||
@@ -729,6 +741,8 @@ class StreamProcessor:
|
|||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
sse_parser: SSEEventParser,
|
sse_parser: SSEEventParser,
|
||||||
line: str,
|
line: str,
|
||||||
|
*,
|
||||||
|
skip_record: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
处理单行数据
|
处理单行数据
|
||||||
@@ -737,6 +751,7 @@ class StreamProcessor:
|
|||||||
ctx: 流式上下文
|
ctx: 流式上下文
|
||||||
sse_parser: SSE 解析器
|
sse_parser: SSE 解析器
|
||||||
line: 原始行数据
|
line: 原始行数据
|
||||||
|
skip_record: 是否跳过记录到 parsed_chunks(当需要格式转换时应为 True)
|
||||||
"""
|
"""
|
||||||
# SSEEventParser 以"去掉换行符"的单行文本作为输入;这里统一剔除 CR/LF,
|
# SSEEventParser 以"去掉换行符"的单行文本作为输入;这里统一剔除 CR/LF,
|
||||||
# 避免把空行误判成 "\n" 并导致事件边界解析错误。
|
# 避免把空行误判成 "\n" 并导致事件边界解析错误。
|
||||||
@@ -747,7 +762,9 @@ class StreamProcessor:
|
|||||||
ctx.chunk_count += 1
|
ctx.chunk_count += 1
|
||||||
|
|
||||||
for event in events:
|
for event in events:
|
||||||
self.handle_sse_event(ctx, event.get("event"), event.get("data") or "")
|
self.handle_sse_event(
|
||||||
|
ctx, event.get("event"), event.get("data") or "", skip_record=skip_record
|
||||||
|
)
|
||||||
|
|
||||||
async def create_monitored_stream(
|
async def create_monitored_stream(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -8,7 +8,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Any, Dict, List
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
|
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
|
||||||
from .stream_events import InternalStreamEvent
|
from .stream_events import InternalStreamEvent
|
||||||
@@ -41,8 +41,20 @@ class FormatNormalizer(ABC):
|
|||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
def response_from_internal(
|
||||||
"""将内部表示转换为格式特定响应"""
|
self,
|
||||||
|
internal: InternalResponse,
|
||||||
|
*,
|
||||||
|
requested_model: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""将内部表示转换为格式特定响应
|
||||||
|
|
||||||
|
Args:
|
||||||
|
internal: 内部响应表示
|
||||||
|
requested_model: 用户请求的原始模型名(可选)。
|
||||||
|
如果提供,响应中的 model 字段将使用此值,
|
||||||
|
而不是上游返回的映射后模型名。
|
||||||
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
# ============ 流式转换(可选) ============
|
# ============ 流式转换(可选) ============
|
||||||
|
|||||||
@@ -241,7 +241,12 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
def response_from_internal(
|
||||||
|
self,
|
||||||
|
internal: InternalResponse,
|
||||||
|
*,
|
||||||
|
requested_model: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
cid = internal.id or "unknown"
|
cid = internal.id or "unknown"
|
||||||
if not cid.startswith("msg_"):
|
if not cid.startswith("msg_"):
|
||||||
cid = f"msg_{cid}"
|
cid = f"msg_{cid}"
|
||||||
@@ -294,11 +299,14 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if internal.usage.cache_write_tokens:
|
if internal.usage.cache_write_tokens:
|
||||||
usage["cache_creation_input_tokens"] = int(internal.usage.cache_write_tokens)
|
usage["cache_creation_input_tokens"] = int(internal.usage.cache_write_tokens)
|
||||||
|
|
||||||
|
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||||
|
model_name = requested_model if requested_model else internal.model
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"id": cid,
|
"id": cid,
|
||||||
"type": "message",
|
"type": "message",
|
||||||
"role": "assistant",
|
"role": "assistant",
|
||||||
"model": internal.model,
|
"model": model_name,
|
||||||
"content": content,
|
"content": content,
|
||||||
"stop_reason": stop_reason,
|
"stop_reason": stop_reason,
|
||||||
"stop_sequence": None,
|
"stop_sequence": None,
|
||||||
@@ -329,9 +337,11 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
message_raw = chunk.get("message")
|
message_raw = chunk.get("message")
|
||||||
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
||||||
msg_id = str(message.get("id") or "")
|
msg_id = str(message.get("id") or "")
|
||||||
model = str(message.get("model") or "")
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||||
|
model = state.model or str(message.get("model") or "")
|
||||||
state.message_id = msg_id or state.message_id
|
state.message_id = msg_id or state.message_id
|
||||||
state.model = model or state.model
|
if not state.model:
|
||||||
|
state.model = model
|
||||||
ss["message_started"] = True
|
ss["message_started"] = True
|
||||||
ss.setdefault("block_index_to_tool_id", {})
|
ss.setdefault("block_index_to_tool_id", {})
|
||||||
events.append(MessageStartEvent(message_id=msg_id, model=model))
|
events.append(MessageStartEvent(message_id=msg_id, model=model))
|
||||||
@@ -440,7 +450,9 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if isinstance(event, MessageStartEvent):
|
if isinstance(event, MessageStartEvent):
|
||||||
state.message_id = event.message_id or state.message_id
|
state.message_id = event.message_id or state.message_id
|
||||||
state.model = event.model or state.model
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用事件值
|
||||||
|
if not state.model:
|
||||||
|
state.model = event.model or ""
|
||||||
ss.setdefault("block_index_to_tool_id", {})
|
ss.setdefault("block_index_to_tool_id", {})
|
||||||
message_obj: Dict[str, Any] = {
|
message_obj: Dict[str, Any] = {
|
||||||
"id": state.message_id or "msg_stream",
|
"id": state.message_id or "msg_stream",
|
||||||
|
|||||||
@@ -333,7 +333,12 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
def response_from_internal(
|
||||||
|
self,
|
||||||
|
internal: InternalResponse,
|
||||||
|
*,
|
||||||
|
requested_model: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
parts: List[Dict[str, Any]] = []
|
parts: List[Dict[str, Any]] = []
|
||||||
for b in internal.content:
|
for b in internal.content:
|
||||||
if isinstance(b, TextBlock):
|
if isinstance(b, TextBlock):
|
||||||
@@ -386,9 +391,12 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if finish_reason is not None:
|
if finish_reason is not None:
|
||||||
candidate["finishReason"] = finish_reason
|
candidate["finishReason"] = finish_reason
|
||||||
|
|
||||||
|
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||||
|
model_name = requested_model if requested_model else (internal.model or "gemini")
|
||||||
|
|
||||||
out: Dict[str, Any] = {
|
out: Dict[str, Any] = {
|
||||||
"candidates": [candidate],
|
"candidates": [candidate],
|
||||||
"modelVersion": internal.model or "gemini",
|
"modelVersion": model_name,
|
||||||
}
|
}
|
||||||
|
|
||||||
if usage_metadata:
|
if usage_metadata:
|
||||||
@@ -553,7 +561,9 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if isinstance(event, MessageStartEvent):
|
if isinstance(event, MessageStartEvent):
|
||||||
state.message_id = event.message_id or state.message_id
|
state.message_id = event.message_id or state.message_id
|
||||||
state.model = event.model or state.model
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用事件值
|
||||||
|
if not state.model:
|
||||||
|
state.model = event.model or ""
|
||||||
ss.setdefault("tool_blocks", {})
|
ss.setdefault("tool_blocks", {})
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|||||||
@@ -281,13 +281,20 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
def response_from_internal(
|
||||||
|
self,
|
||||||
|
internal: InternalResponse,
|
||||||
|
*,
|
||||||
|
requested_model: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
# OpenAI Chat Completions response envelope
|
# OpenAI Chat Completions response envelope
|
||||||
|
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||||
|
model_name = requested_model if requested_model else internal.model
|
||||||
out: Dict[str, Any] = {
|
out: Dict[str, Any] = {
|
||||||
"id": internal.id or "chatcmpl-unknown",
|
"id": internal.id or "chatcmpl-unknown",
|
||||||
"object": "chat.completion",
|
"object": "chat.completion",
|
||||||
"created": int(time.time()),
|
"created": int(time.time()),
|
||||||
"model": internal.model,
|
"model": model_name,
|
||||||
"choices": [],
|
"choices": [],
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -452,7 +459,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if isinstance(event, MessageStartEvent):
|
if isinstance(event, MessageStartEvent):
|
||||||
state.message_id = event.message_id or state.message_id
|
state.message_id = event.message_id or state.message_id
|
||||||
state.model = event.model or state.model
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用事件值
|
||||||
|
if not state.model:
|
||||||
|
state.model = event.model or ""
|
||||||
out.append(base_chunk({"role": "assistant"}))
|
out.append(base_chunk({"role": "assistant"}))
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|||||||
@@ -187,7 +187,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
extra=extra,
|
extra=extra,
|
||||||
)
|
)
|
||||||
|
|
||||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
def response_from_internal(
|
||||||
|
self,
|
||||||
|
internal: InternalResponse,
|
||||||
|
*,
|
||||||
|
requested_model: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
text = self._collapse_internal_text(internal.content)
|
text = self._collapse_internal_text(internal.content)
|
||||||
|
|
||||||
output_message = {
|
output_message = {
|
||||||
@@ -204,11 +209,14 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
|
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||||
|
model_name = requested_model if requested_model else (internal.model or "")
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"id": internal.id or "resp",
|
"id": internal.id or "resp",
|
||||||
"object": "response",
|
"object": "response",
|
||||||
"created": int(time.time()),
|
"created": int(time.time()),
|
||||||
"model": internal.model or "",
|
"model": model_name,
|
||||||
"status": "completed",
|
"status": "completed",
|
||||||
"output": [output_message],
|
"output": [output_message],
|
||||||
"usage": usage_obj,
|
"usage": usage_obj,
|
||||||
@@ -241,10 +249,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
resp_obj = chunk.get("response")
|
resp_obj = chunk.get("response")
|
||||||
resp_obj = resp_obj if isinstance(resp_obj, dict) else {}
|
resp_obj = resp_obj if isinstance(resp_obj, dict) else {}
|
||||||
msg_id = str(resp_obj.get("id") or chunk.get("id") or state.message_id or "")
|
msg_id = str(resp_obj.get("id") or chunk.get("id") or state.message_id or "")
|
||||||
model = str(resp_obj.get("model") or chunk.get("model") or state.model or "")
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||||
|
model = state.model or str(resp_obj.get("model") or chunk.get("model") or "")
|
||||||
if msg_id or model or etype:
|
if msg_id or model or etype:
|
||||||
state.message_id = msg_id
|
state.message_id = msg_id
|
||||||
state.model = model
|
if not state.model:
|
||||||
|
state.model = model
|
||||||
ss["message_started"] = True
|
ss["message_started"] = True
|
||||||
ss.setdefault("text_block_started", False)
|
ss.setdefault("text_block_started", False)
|
||||||
ss.setdefault("text_block_stopped", False)
|
ss.setdefault("text_block_stopped", False)
|
||||||
@@ -301,11 +311,13 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
# response.in_progress:状态更新,不产生内容事件
|
# response.in_progress:状态更新,不产生内容事件
|
||||||
if etype == "response.in_progress":
|
if etype == "response.in_progress":
|
||||||
# 更新 state 中的元数据(如果有)
|
# 更新 state 中的元数据(如果有)
|
||||||
|
# 注意:model 保持初始值(客户端请求的模型),不被上游覆盖
|
||||||
resp_obj = chunk.get("response")
|
resp_obj = chunk.get("response")
|
||||||
if isinstance(resp_obj, dict):
|
if isinstance(resp_obj, dict):
|
||||||
if resp_obj.get("id"):
|
if resp_obj.get("id"):
|
||||||
state.message_id = str(resp_obj.get("id"))
|
state.message_id = str(resp_obj.get("id"))
|
||||||
if resp_obj.get("model"):
|
# 仅在 model 为空时才用上游值
|
||||||
|
if not state.model and resp_obj.get("model"):
|
||||||
state.model = str(resp_obj.get("model"))
|
state.model = str(resp_obj.get("model"))
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@@ -390,7 +402,9 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if isinstance(event, MessageStartEvent):
|
if isinstance(event, MessageStartEvent):
|
||||||
state.message_id = event.message_id or state.message_id or "resp_stream"
|
state.message_id = event.message_id or state.message_id or "resp_stream"
|
||||||
state.model = event.model or state.model or ""
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用事件值
|
||||||
|
if not state.model:
|
||||||
|
state.model = event.model or ""
|
||||||
ss.setdefault("collected_text", "")
|
ss.setdefault("collected_text", "")
|
||||||
out.append(
|
out.append(
|
||||||
event_block(
|
event_block(
|
||||||
|
|||||||
@@ -88,8 +88,28 @@ class FormatConversionRegistry:
|
|||||||
response: Dict[str, Any],
|
response: Dict[str, Any],
|
||||||
source_format: str,
|
source_format: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
|
*,
|
||||||
|
requested_model: Optional[str] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
|
"""转换响应格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
response: 原始响应
|
||||||
|
source_format: 源格式
|
||||||
|
target_format: 目标格式
|
||||||
|
requested_model: 用户请求的原始模型名(可选)。
|
||||||
|
如果提供,响应中的 model 字段将使用此值,
|
||||||
|
而不是上游返回的映射后模型名。
|
||||||
|
"""
|
||||||
if str(source_format).upper() == str(target_format).upper():
|
if str(source_format).upper() == str(target_format).upper():
|
||||||
|
# 即使格式相同,也需要替换 model 字段
|
||||||
|
if requested_model and isinstance(response, dict):
|
||||||
|
response = dict(response) # 避免修改原始响应
|
||||||
|
# 支持不同格式的 model 字段名
|
||||||
|
if "model" in response:
|
||||||
|
response["model"] = requested_model
|
||||||
|
elif "modelVersion" in response:
|
||||||
|
response["modelVersion"] = requested_model
|
||||||
return response
|
return response
|
||||||
|
|
||||||
src = self._require_normalizer(source_format)
|
src = self._require_normalizer(source_format)
|
||||||
@@ -98,7 +118,7 @@ class FormatConversionRegistry:
|
|||||||
with _track_conversion_metrics("response", str(source_format).upper(), str(target_format).upper()):
|
with _track_conversion_metrics("response", str(source_format).upper(), str(target_format).upper()):
|
||||||
try:
|
try:
|
||||||
internal = src.response_to_internal(response)
|
internal = src.response_to_internal(response)
|
||||||
return tgt.response_from_internal(internal)
|
return tgt.response_from_internal(internal, requested_model=requested_model)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||||
|
|
||||||
@@ -151,6 +171,12 @@ class FormatConversionRegistry:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if state is None:
|
if state is None:
|
||||||
|
# 调用方应提供预初始化的 state(包含 model/message_id),
|
||||||
|
# 这里仅作为防御性回退,可能导致响应中 model 字段为空
|
||||||
|
logger.debug(
|
||||||
|
f"convert_stream_chunk: state is None, creating empty StreamState "
|
||||||
|
f"(source={source_format}, target={target_format})"
|
||||||
|
)
|
||||||
state = StreamState()
|
state = StreamState()
|
||||||
|
|
||||||
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
|
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
|
||||||
|
|||||||
@@ -54,9 +54,10 @@ class MockCliHandler:
|
|||||||
return [line]
|
return [line]
|
||||||
|
|
||||||
# 初始化流式转换状态
|
# 初始化流式转换状态
|
||||||
|
# 使用客户端请求的模型名(ctx.model),而非映射后的模型名(ctx.mapped_model)
|
||||||
if ctx.stream_conversion_state is None:
|
if ctx.stream_conversion_state is None:
|
||||||
ctx.stream_conversion_state = StreamState(
|
ctx.stream_conversion_state = StreamState(
|
||||||
model=ctx.mapped_model or ctx.model,
|
model=ctx.model,
|
||||||
message_id=ctx.response_id or ctx.request_id,
|
message_id=ctx.response_id or ctx.request_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -143,12 +144,16 @@ class TestConvertSseLineWithMockConverter:
|
|||||||
assert json.loads(result[0][6:]) == chunk
|
assert json.loads(result[0][6:]) == chunk
|
||||||
|
|
||||||
def test_state_initialization(self) -> None:
|
def test_state_initialization(self) -> None:
|
||||||
"""测试状态自动初始化"""
|
"""测试状态自动初始化
|
||||||
|
|
||||||
|
流式转换状态应使用用户请求的原始模型名(ctx.model),
|
||||||
|
而非映射后的模型名(ctx.mapped_model),确保返回给客户端的响应使用原始模型名。
|
||||||
|
"""
|
||||||
handler = MockCliHandler()
|
handler = MockCliHandler()
|
||||||
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
|
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
|
||||||
ctx.provider_api_format = "OPENAI"
|
ctx.provider_api_format = "OPENAI"
|
||||||
ctx.client_api_format = "OPENAI"
|
ctx.client_api_format = "OPENAI"
|
||||||
ctx.mapped_model = "claude-3-5-sonnet"
|
ctx.mapped_model = "claude-3-5-sonnet" # 映射后的模型名(发给上游的)
|
||||||
ctx.request_id = "req_123"
|
ctx.request_id = "req_123"
|
||||||
|
|
||||||
chunk = {"choices": [{"delta": {"content": "test"}}]}
|
chunk = {"choices": [{"delta": {"content": "test"}}]}
|
||||||
@@ -156,9 +161,9 @@ class TestConvertSseLineWithMockConverter:
|
|||||||
|
|
||||||
handler._convert_sse_line(ctx, line, [])
|
handler._convert_sse_line(ctx, line, [])
|
||||||
|
|
||||||
# 验证状态已初始化
|
# 验证状态已初始化,使用用户请求的原始模型名
|
||||||
assert ctx.stream_conversion_state is not None
|
assert ctx.stream_conversion_state is not None
|
||||||
assert ctx.stream_conversion_state.model == "claude-3-5-sonnet"
|
assert ctx.stream_conversion_state.model == "gpt-4" # 应使用原始模型名,非 mapped_model
|
||||||
assert ctx.stream_conversion_state.message_id == "req_123"
|
assert ctx.stream_conversion_state.message_id == "req_123"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user