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

@@ -8,7 +8,7 @@
from __future__ import annotations
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 .stream_events import InternalStreamEvent
@@ -41,8 +41,20 @@ class FormatNormalizer(ABC):
raise NotImplementedError
@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
# ============ 流式转换(可选) ============

View File

@@ -241,7 +241,12 @@ class ClaudeNormalizer(FormatNormalizer):
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"
if not cid.startswith("msg_"):
cid = f"msg_{cid}"
@@ -294,11 +299,14 @@ class ClaudeNormalizer(FormatNormalizer):
if 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 {
"id": cid,
"type": "message",
"role": "assistant",
"model": internal.model,
"model": model_name,
"content": content,
"stop_reason": stop_reason,
"stop_sequence": None,
@@ -329,9 +337,11 @@ class ClaudeNormalizer(FormatNormalizer):
message_raw = chunk.get("message")
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
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.model = model or state.model
if not state.model:
state.model = model
ss["message_started"] = True
ss.setdefault("block_index_to_tool_id", {})
events.append(MessageStartEvent(message_id=msg_id, model=model))
@@ -440,7 +450,9 @@ class ClaudeNormalizer(FormatNormalizer):
if isinstance(event, MessageStartEvent):
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", {})
message_obj: Dict[str, Any] = {
"id": state.message_id or "msg_stream",

View File

@@ -333,7 +333,12 @@ class GeminiNormalizer(FormatNormalizer):
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]] = []
for b in internal.content:
if isinstance(b, TextBlock):
@@ -386,9 +391,12 @@ class GeminiNormalizer(FormatNormalizer):
if finish_reason is not None:
candidate["finishReason"] = finish_reason
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else (internal.model or "gemini")
out: Dict[str, Any] = {
"candidates": [candidate],
"modelVersion": internal.model or "gemini",
"modelVersion": model_name,
}
if usage_metadata:
@@ -553,7 +561,9 @@ class GeminiNormalizer(FormatNormalizer):
if isinstance(event, MessageStartEvent):
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", {})
return out

View File

@@ -281,13 +281,20 @@ class OpenAINormalizer(FormatNormalizer):
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
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else internal.model
out: Dict[str, Any] = {
"id": internal.id or "chatcmpl-unknown",
"object": "chat.completion",
"created": int(time.time()),
"model": internal.model,
"model": model_name,
"choices": [],
}
@@ -452,7 +459,9 @@ class OpenAINormalizer(FormatNormalizer):
if isinstance(event, MessageStartEvent):
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"}))
return out

View File

@@ -187,7 +187,12 @@ class OpenAICliNormalizer(FormatNormalizer):
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)
output_message = {
@@ -204,11 +209,14 @@ class OpenAICliNormalizer(FormatNormalizer):
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
}
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else (internal.model or "")
return {
"id": internal.id or "resp",
"object": "response",
"created": int(time.time()),
"model": internal.model or "",
"model": model_name,
"status": "completed",
"output": [output_message],
"usage": usage_obj,
@@ -241,10 +249,12 @@ class OpenAICliNormalizer(FormatNormalizer):
resp_obj = chunk.get("response")
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 "")
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:
state.message_id = msg_id
state.model = model
if not state.model:
state.model = model
ss["message_started"] = True
ss.setdefault("text_block_started", False)
ss.setdefault("text_block_stopped", False)
@@ -301,11 +311,13 @@ class OpenAICliNormalizer(FormatNormalizer):
# response.in_progress状态更新不产生内容事件
if etype == "response.in_progress":
# 更新 state 中的元数据(如果有)
# 注意model 保持初始值(客户端请求的模型),不被上游覆盖
resp_obj = chunk.get("response")
if isinstance(resp_obj, dict):
if 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"))
return []
@@ -390,7 +402,9 @@ class OpenAICliNormalizer(FormatNormalizer):
if isinstance(event, MessageStartEvent):
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", "")
out.append(
event_block(

View File

@@ -88,8 +88,28 @@ class FormatConversionRegistry:
response: Dict[str, Any],
source_format: str,
target_format: str,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
"""转换响应格式
Args:
response: 原始响应
source_format: 源格式
target_format: 目标格式
requested_model: 用户请求的原始模型名(可选)。
如果提供,响应中的 model 字段将使用此值,
而不是上游返回的映射后模型名。
"""
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
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()):
try:
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:
raise FormatConversionError(source_format, target_format, str(e)) from e
@@ -151,6 +171,12 @@ class FormatConversionRegistry:
)
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()
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):