diff --git a/backend/app/services/oasis_profile_generator.py b/backend/app/services/oasis_profile_generator.py index 7704a627..5847bcd1 100644 --- a/backend/app/services/oasis_profile_generator.py +++ b/backend/app/services/oasis_profile_generator.py @@ -21,6 +21,7 @@ from zep_cloud.client import Zep from ..config import Config from ..utils.logger import get_logger from ..utils.locale import get_language_instruction, get_locale, set_locale, t +from ..utils.openai_chat_compat import create_chat_completion, extract_chat_completion_text from .zep_entity_reader import EntityNode, ZepEntityReader logger = get_logger('mirofish.oasis_profile') @@ -527,18 +528,19 @@ class OasisProfileGenerator: for attempt in range(max_attempts): try: - response = self.client.chat.completions.create( + response = create_chat_completion( + self.client, model=self.model_name, messages=[ {"role": "system", "content": self._get_system_prompt(is_individual)}, {"role": "user", "content": prompt} ], response_format={"type": "json_object"}, - temperature=0.7 - (attempt * 0.1) # 每次重试降低温度 + temperature=0.7 - (attempt * 0.1), # 每次重试降低温度 # 不设置max_tokens,让LLM自由发挥 ) - content = response.choices[0].message.content + content = extract_chat_completion_text(response) # 检查是否被截断(finish_reason不是'stop') finish_reason = response.choices[0].finish_reason diff --git a/backend/app/services/simulation_config_generator.py b/backend/app/services/simulation_config_generator.py index cb77f6b6..2c6fbbf1 100644 --- a/backend/app/services/simulation_config_generator.py +++ b/backend/app/services/simulation_config_generator.py @@ -21,6 +21,7 @@ from openai import OpenAI from ..config import Config from ..utils.logger import get_logger from ..utils.locale import get_language_instruction, t +from ..utils.openai_chat_compat import create_chat_completion, extract_chat_completion_text from .zep_entity_reader import EntityNode, ZepEntityReader logger = get_logger('mirofish.simulation_config') @@ -440,18 +441,19 @@ class SimulationConfigGenerator: for attempt in range(max_attempts): try: - response = self.client.chat.completions.create( + response = create_chat_completion( + self.client, model=self.model_name, messages=[ {"role": "system", "content": system_prompt}, {"role": "user", "content": prompt} ], response_format={"type": "json_object"}, - temperature=0.7 - (attempt * 0.1) # 每次重试降低温度 + temperature=0.7 - (attempt * 0.1), # 每次重试降低温度 # 不设置max_tokens,让LLM自由发挥 ) - content = response.choices[0].message.content + content = extract_chat_completion_text(response) finish_reason = response.choices[0].finish_reason # 检查是否被截断 diff --git a/backend/app/utils/llm_client.py b/backend/app/utils/llm_client.py index 6c1a81f4..fc316fce 100644 --- a/backend/app/utils/llm_client.py +++ b/backend/app/utils/llm_client.py @@ -9,6 +9,7 @@ from typing import Optional, Dict, Any, List from openai import OpenAI from ..config import Config +from .openai_chat_compat import create_chat_completion, extract_chat_completion_text class LLMClient: @@ -51,18 +52,15 @@ class LLMClient: Returns: 模型响应文本 """ - kwargs = { - "model": self.model, - "messages": messages, - "temperature": temperature, - "max_tokens": max_tokens, - } - - if response_format: - kwargs["response_format"] = response_format - - response = self.client.chat.completions.create(**kwargs) - content = response.choices[0].message.content + response = create_chat_completion( + self.client, + model=self.model, + messages=messages, + temperature=temperature, + max_tokens=max_tokens, + response_format=response_format, + ) + content = extract_chat_completion_text(response) # 部分模型(如MiniMax M2.5)会在content中包含思考内容,需要移除 content = re.sub(r'[\s\S]*?', '', content).strip() return content diff --git a/backend/app/utils/openai_chat_compat.py b/backend/app/utils/openai_chat_compat.py new file mode 100644 index 00000000..f0ce9374 --- /dev/null +++ b/backend/app/utils/openai_chat_compat.py @@ -0,0 +1,101 @@ +""" +OpenAI Chat Completions compatibility helpers. + +This module keeps existing behavior for legacy models/providers while +gracefully adapting request parameters for GPT-5 family models. +""" + +from __future__ import annotations + +from typing import Any, Dict, List, Optional + + +def is_gpt5_family(model: Optional[str]) -> bool: + """Return True when model belongs to GPT-5 family aliases/snapshots.""" + if not model: + return False + return model.strip().lower().startswith("gpt-5") + + +def create_chat_completion( + client: Any, + *, + model: str, + messages: List[Dict[str, Any]], + temperature: Optional[float] = None, + max_tokens: Optional[int] = None, + response_format: Optional[Dict[str, Any]] = None, +) -> Any: + """ + Create a chat completion with model-specific request parameters. + + Compatibility strategy: + - For GPT-5 family, avoid sending temperature by default. + - For token limit, use `max_completion_tokens` on GPT-5, `max_tokens` otherwise. + - Preserve the legacy request shape for every non-GPT-5 model/provider. + - Propagate provider errors unchanged instead of guessing from message text. + """ + kwargs: Dict[str, Any] = { + "model": model, + "messages": messages, + } + + if response_format is not None: + kwargs["response_format"] = response_format + + gpt5_family = is_gpt5_family(model) + + if temperature is not None and not gpt5_family: + kwargs["temperature"] = temperature + + if max_tokens is not None: + if gpt5_family: + kwargs["max_completion_tokens"] = max_tokens + else: + kwargs["max_tokens"] = max_tokens + + return client.chat.completions.create(**kwargs) + + +def extract_chat_completion_text(response: Any) -> str: + """Extract plain text from chat completion response across SDK content shapes.""" + choices = getattr(response, "choices", None) or [] + if not choices: + return "" + + message = getattr(choices[0], "message", None) + if message is None: + return "" + + content = getattr(message, "content", "") + + if isinstance(content, str): + return content + + if isinstance(content, list): + chunks: List[str] = [] + for item in content: + if isinstance(item, dict): + text_obj = item.get("text") + if isinstance(text_obj, dict): + text_obj = text_obj.get("value") + if isinstance(text_obj, str): + chunks.append(text_obj) + elif isinstance(item.get("content"), str): + chunks.append(item["content"]) + continue + + text_obj = getattr(item, "text", None) + if isinstance(text_obj, dict): + text_obj = text_obj.get("value") + if isinstance(text_obj, str): + chunks.append(text_obj) + continue + + content_obj = getattr(item, "content", None) + if isinstance(content_obj, str): + chunks.append(content_obj) + + return "".join(chunks).strip() + + return str(content or "") diff --git a/backend/tests/test_openai_chat_compat.py b/backend/tests/test_openai_chat_compat.py new file mode 100644 index 00000000..fe08a62a --- /dev/null +++ b/backend/tests/test_openai_chat_compat.py @@ -0,0 +1,123 @@ +from types import SimpleNamespace + +import pytest + +from app.utils.openai_chat_compat import ( + create_chat_completion, + extract_chat_completion_text, + is_gpt5_family, +) + + +class CompletionRecorder: + def __init__(self, result=None, error=None): + self.result = result or object() + self.error = error + self.calls = [] + + def create(self, **kwargs): + self.calls.append(kwargs) + if self.error is not None: + raise self.error + return self.result + + +def client_for(recorder): + return SimpleNamespace(chat=SimpleNamespace(completions=recorder)) + + +def test_gpt5_uses_completion_token_limit_without_temperature(): + recorder = CompletionRecorder() + messages = [{"role": "user", "content": "hello"}] + + result = create_chat_completion( + client_for(recorder), + model="gpt-5-2025-08-07", + messages=messages, + temperature=0.2, + max_tokens=123, + response_format={"type": "json_object"}, + ) + + assert result is recorder.result + assert recorder.calls == [ + { + "model": "gpt-5-2025-08-07", + "messages": messages, + "max_completion_tokens": 123, + "response_format": {"type": "json_object"}, + } + ] + + +def test_legacy_model_preserves_original_request_shape(): + recorder = CompletionRecorder() + messages = [{"role": "user", "content": "hello"}] + + create_chat_completion( + client_for(recorder), + model="third-party-chat-model", + messages=messages, + temperature=0.7, + max_tokens=456, + response_format={"type": "json_object"}, + ) + + assert recorder.calls == [ + { + "model": "third-party-chat-model", + "messages": messages, + "temperature": 0.7, + "max_tokens": 456, + "response_format": {"type": "json_object"}, + } + ] + + +def test_provider_error_is_propagated_without_guessing_or_retrying(): + provider_error = RuntimeError("unsupported max_tokens due to a server outage") + recorder = CompletionRecorder(error=provider_error) + + with pytest.raises(RuntimeError) as captured: + create_chat_completion( + client_for(recorder), + model="legacy-model", + messages=[], + max_tokens=10, + ) + + assert captured.value is provider_error + assert len(recorder.calls) == 1 + + +@pytest.mark.parametrize( + ("model", "expected"), + [ + ("gpt-5", True), + (" GPT-5.1-mini ", True), + ("gpt-4.1", False), + ("my-gpt-5-proxy", False), + (None, False), + ], +) +def test_gpt5_family_detection(model, expected): + assert is_gpt5_family(model) is expected + + +def test_extracts_text_from_supported_content_shapes(): + response = SimpleNamespace( + choices=[ + SimpleNamespace( + message=SimpleNamespace( + content=[ + {"text": {"value": "first"}}, + {"content": " second"}, + SimpleNamespace(text=" third"), + ] + ) + ) + ] + ) + + assert extract_chat_completion_text(response) == "first second third" + assert extract_chat_completion_text(SimpleNamespace(choices=[])) == ""