fix: 响应中保留用户请求的原始模型名

- 非流式响应:通过 requested_model 参数传递用户请求的模型名
- 流式响应:StreamState 初始化时使用 ctx.model 而非 ctx.mapped_model
- 流式转换:normalizers 保留初始 model 值,不被上游事件覆盖
- parsed_chunks:格式转换时记录转换后数据,确保与客户端收到的一致
This commit is contained in:
fawney19
2026-01-29 10:32:36 +08:00
parent 2e9c9e80d1
commit 557e64c77f
10 changed files with 141 additions and 34 deletions

View File

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

View File

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

View File

@@ -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 记录转换后的数据)
if not skip_record:
ctx.parsed_chunks.append(data) 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,

View File

@@ -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
# ============ 流式转换(可选) ============ # ============ 流式转换(可选) ============

View File

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

View File

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

View File

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

View File

@@ -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,9 +249,11 @@ 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
if not state.model:
state.model = model state.model = model
ss["message_started"] = True ss["message_started"] = True
ss.setdefault("text_block_started", False) ss.setdefault("text_block_started", 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(

View File

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

View File

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