diff --git a/src/agentscope/_version.py b/src/agentscope/_version.py index 2a3386b3ef..984adfd408 100644 --- a/src/agentscope/_version.py +++ b/src/agentscope/_version.py @@ -1,4 +1,4 @@ # -*- coding: utf-8 -*- """The version of agentscope.""" -__version__ = "1.0.20" +__version__ = "1.0.21" diff --git a/src/agentscope/formatter/__init__.py b/src/agentscope/formatter/__init__.py index 3202298135..37d7c3ef8d 100644 --- a/src/agentscope/formatter/__init__.py +++ b/src/agentscope/formatter/__init__.py @@ -29,6 +29,11 @@ ) from ._a2a_formatter import A2AChatFormatter +from ._openai_response_formatter import ( + OpenAIResponseChatFormatter, + OpenAIResponseMultiAgentFormatter, +) + __all__ = [ "FormatterBase", "TruncatedFormatterBase", @@ -45,4 +50,6 @@ "DeepSeekChatFormatter", "DeepSeekMultiAgentFormatter", "A2AChatFormatter", + "OpenAIResponseChatFormatter", + "OpenAIResponseMultiAgentFormatter", ] diff --git a/src/agentscope/formatter/_anthropic_formatter.py b/src/agentscope/formatter/_anthropic_formatter.py index 7b8db2b435..7a49e0ad59 100644 --- a/src/agentscope/formatter/_anthropic_formatter.py +++ b/src/agentscope/formatter/_anthropic_formatter.py @@ -171,7 +171,7 @@ async def _format( elif typ == "tool_result": output = block.get("output") if output is None: - content_value = [{"type": "text", "text": None}] + content_value = [{"type": "text", "text": ""}] elif isinstance(output, list): content_value = [ _format_anthropic_image_block(item) @@ -207,7 +207,7 @@ async def _format( msg_anthropic = { "role": role, - "content": content_blocks or None, + "content": content_blocks, } # When both content and tool_calls are None, skipped diff --git a/src/agentscope/formatter/_gemini_formatter.py b/src/agentscope/formatter/_gemini_formatter.py index 81d2dfbd22..ed1fb82d47 100644 --- a/src/agentscope/formatter/_gemini_formatter.py +++ b/src/agentscope/formatter/_gemini_formatter.py @@ -191,23 +191,30 @@ async def _format( for block in msg.get_content_blocks(): typ = block.get("type") if typ == "text": - parts.append( - { - "text": block.get("text"), - }, - ) + text_part: dict[str, Any] = { + "text": block.get("text"), + } + text_sig = block.get("thought_signature") + if text_sig is not None: + text_part["thought_signature"] = base64.b64decode( + text_sig, # type: ignore[arg-type] + ) + parts.append(text_part) elif typ == "tool_use": - parts.append( - { - "function_call": { - "id": None, - "name": block["name"], - "args": block["input"], - }, - "thought_signature": block.get("id", None), + fc_part: dict[str, Any] = { + "function_call": { + "id": None, + "name": block["name"], + "args": block["input"], }, - ) + } + thought_sig = block.get("thought_signature") + if thought_sig is not None: + fc_part["thought_signature"] = base64.b64decode( + thought_sig, # type: ignore[arg-type] + ) + parts.append(fc_part) elif typ == "tool_result": ( diff --git a/src/agentscope/formatter/_openai_response_formatter.py b/src/agentscope/formatter/_openai_response_formatter.py new file mode 100644 index 0000000000..01061a2af7 --- /dev/null +++ b/src/agentscope/formatter/_openai_response_formatter.py @@ -0,0 +1,452 @@ +# -*- coding: utf-8 -*- +# pylint: disable=too-many-branches, too-many-nested-blocks +"""The OpenAI response formatter for agentscope.""" +import json +from typing import Any + +from ._openai_formatter import _to_openai_image_url, _to_openai_audio_data +from ._truncated_formatter_base import TruncatedFormatterBase +from .._logging import logger +from ..message import ( + Msg, + TextBlock, + ImageBlock, + ToolUseBlock, + ToolResultBlock, +) +from ..token import TokenCounterBase + + +def _format_openai_response_image_block( + image_block: ImageBlock, +) -> dict[str, Any]: + """Format an image block for OpenAI response API. + + Args: + image_block (`ImageBlock`): + The image block to format. + + Returns: + `dict[str, Any]`: + A dictionary with "type" and "image_url" keys in OpenAI + response format. + + Raises: + `ValueError`: + If the source type is not supported. + """ + source = image_block["source"] + if source["type"] == "url": + url = _to_openai_image_url(source["url"]) + elif source["type"] == "base64": + data = source["data"] + media_type = source["media_type"] + url = f"data:{media_type};base64,{data}" + else: + raise ValueError( + f"Unsupported image source type: {source['type']}", + ) + + return { + "type": "input_image", + "image_url": url, + } + + +class OpenAIResponseChatFormatter(TruncatedFormatterBase): + """The OpenAI response formatter class for chatbot scenario, where only + a user and an agent are involved. We use the `name` field in OpenAI + response API to identify different entities in the conversation. + """ + + support_tools_api: bool = True + """Whether support tools API""" + + support_multiagent: bool = True + """Whether support multi-agent conversation""" + + support_vision: bool = True + """Whether support vision models""" + + supported_blocks: list[type] = [ + TextBlock, + ImageBlock, + ToolUseBlock, + ToolResultBlock, + ] + """Supported message blocks for OpenAI response API""" + + def __init__( + self, + promote_tool_result_images: bool = False, + token_counter: TokenCounterBase | None = None, + max_tokens: int | None = None, + ) -> None: + """Initialize the OpenAI response chat formatter. + + Args: + promote_tool_result_images (`bool`, defaults to `False`): + Whether to promote images from tool results to user messages. + Most LLM APIs don't support images in tool result blocks, but + do support them in user message blocks. When `True`, images are + extracted and appended as a separate user message with + explanatory text indicating their source. + token_counter (`TokenCounterBase | None`, optional): + A token counter instance used to count tokens in the messages. + If not provided, the formatter will format the messages + without considering token limits. + max_tokens (`int | None`, optional): + The maximum number of tokens allowed in the formatted + messages. If not provided, the formatter will not truncate + the messages. + """ + super().__init__(token_counter=token_counter, max_tokens=max_tokens) + self.promote_tool_result_images = promote_tool_result_images + + async def _format( + self, + msgs: list[Msg], + ) -> list[dict[str, Any]]: + """Format message objects into OpenAI response API required format. + + Args: + msgs (`list[Msg]`): + The list of Msg objects to format. + + Returns: + `list[dict[str, Any]]`: + A list of dictionaries, where each dictionary has "name", + "role", and "content" keys. + """ + self.assert_list_of_msgs(msgs) + + messages: list[dict] = [] + i = 0 + while i < len(msgs): + msg = msgs[i] + content_blocks = [] + # Responses API treats function_call / function_call_output as + # top-level input items (not nested inside a message). Collect + # them here and flush after the current message item. + trailing_items: list[dict] = [] + # Assistant text must use output_text in Responses API; user / + # system messages use input_text. + text_type = ( + "output_text" if msg.role == "assistant" else "input_text" + ) + + for block in msg.get_content_blocks(): + typ = block.get("type") + if typ == "text": + content_blocks.append( + { + "type": text_type, + "text": block.get("text"), + }, + ) + + elif typ == "tool_use": + trailing_items.append( + { + "type": "function_call", + "call_id": block.get("id"), + "name": block.get("name"), + "arguments": json.dumps( + block.get("input", {}), + ensure_ascii=False, + ), + }, + ) + + elif typ == "tool_result": + ( + textual_output, + multimodal_data, + ) = self.convert_tool_result_to_string(block["output"]) + + trailing_items.append( + { + "type": "function_call_output", + "call_id": block.get("id"), + "output": textual_output, + }, + ) + + # Then, handle the multimodal data if any + promoted_content: list = [] + for url, multimodal_block in multimodal_data: + if ( + multimodal_block["type"] == "image" + and self.promote_tool_result_images + ): + promoted_content.extend( + [ + { + "type": "input_text", + "text": ( + f"\n- The image from " f"'{url}': " + ), + }, + { + "type": "input_image", + "image_url": ( + _to_openai_image_url( + url, + ) + ), + }, + ], + ) + + if promoted_content: + messages.append( + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": ( + "The following" + " are the image contents " + "from the tool result of " + f"'{block['name']}':" + ), + }, + *promoted_content, + { + "type": "input_text", + "text": "", + }, + ], + }, + ) + + elif typ == "image": + content_blocks.append( + _format_openai_response_image_block( + block, # type: ignore[arg-type] + ), + ) + + elif typ == "audio": + # Filter out audio content when the multimodal model + # outputs both text and audio, to prevent errors in + # subsequent model calls + if msg.role == "assistant": + continue + input_audio = _to_openai_audio_data( + block["source"], + ) + content_blocks.append( + { + "type": "input_audio", + "input_audio": input_audio, + }, + ) + + else: + logger.warning( + "Unsupported block type %s in the message, skipped.", + typ, + ) + + if content_blocks: + messages.append( + { + "role": msg.role, + "content": content_blocks, + }, + ) + + # Append function_call / function_call_output items (if any) + # as separate top-level items after the message. + messages.extend(trailing_items) + + # Move to next message + i += 1 + + return messages + + +class OpenAIResponseMultiAgentFormatter(TruncatedFormatterBase): + """ + OpenAI response formatter for multi-agent conversations, where more than + a user and an agent are involved. + + .. tip:: This formatter is compatible with OpenAI response API and + OpenAI-response-compatible services like vLLM, Azure OpenAI, and others. + """ + + support_tools_api: bool = True + """Whether support tools API""" + + support_multiagent: bool = True + """Whether support multi-agent conversation""" + + support_vision: bool = True + """Whether support vision models""" + + supported_blocks: list[type] = [ + TextBlock, + ImageBlock, + ToolUseBlock, + ToolResultBlock, + ] + """Supported message blocks for OpenAI response API""" + + def __init__( + self, + conversation_history_prompt: str = ( + "# Conversation History\n" + "The content between tags contains " + "your conversation history\n" + ), + promote_tool_result_images: bool = False, + token_counter: TokenCounterBase | None = None, + max_tokens: int | None = None, + ) -> None: + """Initialize the OpenAI response multi-agent formatter. + + Args: + conversation_history_prompt (`str`): + The prompt to use for the conversation history section. + promote_tool_result_images (`bool`, defaults to `False`): + Whether to promote images from tool results to user messages. + Most LLM APIs don't support images in tool result blocks, but + do support them in user message blocks. When `True`, images are + extracted and appended as a separate user message with + explanatory text indicating their source. + token_counter (`TokenCounterBase | None`, optional): + A token counter instance used to count tokens in the messages. + If not provided, the formatter will format the messages + without considering token limits. + max_tokens (`int | None`, optional): + The maximum number of tokens allowed in the formatted + messages. If not provided, the formatter will not truncate + the messages. + """ + super().__init__(token_counter=token_counter, max_tokens=max_tokens) + self.conversation_history_prompt = conversation_history_prompt + self.promote_tool_result_images = promote_tool_result_images + + async def _format_system_message( + self, + msg: Msg, + ) -> dict[str, Any]: + """Format system message using ``input_text`` block type.""" + return { + "role": "system", + "content": [ + {"type": "input_text", "text": block["text"]} + for block in msg.get_content_blocks("text") + ], + } + + async def _format_tool_sequence( + self, + msgs: list[Msg], + ) -> list[dict[str, Any]]: + """Given a sequence of tool call/result messages, format them into + the required format for the OpenAI response API.""" + return await OpenAIResponseChatFormatter( + promote_tool_result_images=self.promote_tool_result_images, + ).format(msgs) + + async def _format_agent_message( + self, + msgs: list[Msg], + is_first: bool = True, + ) -> list[dict[str, Any]]: + """Given a sequence of messages without tool calls/results, format + them into the required format for the OpenAI response API.""" + + if is_first: + conversation_history_prompt = self.conversation_history_prompt + else: + conversation_history_prompt = "" + + # Format into required OpenAI response format + formatted_msgs: list[dict] = [] + + conversation_blocks: list = [] + accumulated_text = [] + images = [] + audios = [] + + for msg in msgs: + for block in msg.get_content_blocks(): + if block["type"] == "text": + accumulated_text.append(f"{msg.name}: {block['text']}") + + elif block["type"] == "image": + images.append(_format_openai_response_image_block(block)) + elif block["type"] == "audio": + # Filter out audio content when the multimodal model + # outputs both text and audio, to prevent errors in + # subsequent model calls + if msg.role == "assistant": + continue + input_audio = _to_openai_audio_data( + block["source"], + ) + audios.append( + { + "type": "input_audio", + "input_audio": input_audio, + }, + ) + + if accumulated_text: + conversation_blocks.append( + {"text": "\n".join(accumulated_text)}, + ) + + if conversation_blocks: + if conversation_blocks[0].get("text"): + conversation_blocks[0]["text"] = ( + conversation_history_prompt + + "\n" + + conversation_blocks[0]["text"] + ) + + else: + conversation_blocks.insert( + 0, + { + "text": conversation_history_prompt + "\n", + }, + ) + + if conversation_blocks[-1].get("text"): + conversation_blocks[-1]["text"] += "\n" + + else: + conversation_blocks.append({"text": ""}) + + conversation_blocks_text = "\n".join( + conversation_block.get("text", "") + for conversation_block in conversation_blocks + ) + + content_list: list[dict[str, Any]] = [] + if conversation_blocks_text: + content_list.append( + { + "type": "input_text", + "text": conversation_blocks_text, + }, + ) + if images: + content_list.extend(images) + if audios: + content_list.extend(audios) + + user_message = { + "role": "user", + "content": content_list, + } + + if content_list: + formatted_msgs.append(user_message) + + return formatted_msgs diff --git a/src/agentscope/model/__init__.py b/src/agentscope/model/__init__.py index 7cd0b7a8f7..0fb2d59093 100644 --- a/src/agentscope/model/__init__.py +++ b/src/agentscope/model/__init__.py @@ -9,6 +9,7 @@ from ._ollama_model import OllamaChatModel from ._gemini_model import GeminiChatModel from ._trinity_model import TrinityChatModel +from ._openai_response_model import OpenAIResponseModel __all__ = [ "ChatModelBase", @@ -19,4 +20,5 @@ "OllamaChatModel", "GeminiChatModel", "TrinityChatModel", + "OpenAIResponseModel", ] diff --git a/src/agentscope/model/_gemini_model.py b/src/agentscope/model/_gemini_model.py index 0a393b69d8..e0e7db8691 100644 --- a/src/agentscope/model/_gemini_model.py +++ b/src/agentscope/model/_gemini_model.py @@ -368,6 +368,7 @@ async def _parse_gemini_stream_generation_response( tool_calls: list[ToolUseBlock] = [] metadata: dict | None = None response_id: str | None = None + text_thought_signature: str | None = None async for chunk in response: if ( chunk.candidates @@ -391,26 +392,46 @@ async def _parse_gemini_stream_generation_response( # requires the thought_signature for some # llms like gemini-3-pro + thought_sig_b64: str | None = None if part.thought_signature: - call_id = base64.b64encode( + thought_sig_b64 = base64.b64encode( part.thought_signature, ).decode("utf-8") - else: - call_id = part.function_call.id - - tool_calls.append( - ToolUseBlock( - type="tool_use", - id=call_id, - name=part.function_call.name, - input=keyword_args, - raw_input=json.dumps( - keyword_args, - ensure_ascii=False, - ), + + call_id = ( + thought_sig_b64 + or part.function_call.id + or f"call_{len(tool_calls)}" + ) + + tool_call = ToolUseBlock( + type="tool_use", + id=call_id, + name=part.function_call.name, + input=keyword_args, + raw_input=json.dumps( + keyword_args, + ensure_ascii=False, ), ) + if thought_sig_b64: + tool_call["thought_signature"] = thought_sig_b64 + + tool_calls.append(tool_call) + + # Capture thought_signature on non-FC parts for + # preserving reasoning context in text responses. + # May arrive on a part with empty text during streaming. + if ( + not part.function_call + and getattr(part, "thought_signature", None) + and isinstance(part.thought_signature, bytes) + ): + text_thought_signature = base64.b64encode( + part.thought_signature, + ).decode("utf-8") + # Text parts if text and structured_model: metadata = _json_loads_with_repair(text) @@ -429,12 +450,13 @@ async def _parse_gemini_stream_generation_response( ) if text: - content_blocks.append( - TextBlock( - type="text", - text=text, - ), + text_block = TextBlock( + type="text", + text=text, ) + if text_thought_signature: + text_block["thought_signature"] = text_thought_signature + content_blocks.append(text_block) if response_id is None: response_id = getattr(chunk, "response_id", None) @@ -477,6 +499,7 @@ def _parse_gemini_generation_response( content_blocks: List[TextBlock | ToolUseBlock | ThinkingBlock] = [] metadata: dict | None = None tool_calls: list = [] + text_thought_signature: str | None = None if ( response.candidates @@ -509,26 +532,52 @@ def _parse_gemini_generation_response( # someday, but Gemini requires the thought_signature # for some llms like gemini-3-pro + thought_sig_b64: str | None = None if part.thought_signature: - call_id = base64.b64encode( + thought_sig_b64 = base64.b64encode( part.thought_signature, ).decode("utf-8") - else: - call_id = part.function_call.id - tool_calls.append( - ToolUseBlock( - type="tool_use", - id=call_id, - name=part.function_call.name, - input=keyword_args, - raw_input=json.dumps( - keyword_args, - ensure_ascii=False, - ), + call_id = ( + thought_sig_b64 + or part.function_call.id + or f"call_{len(tool_calls)}" + ) + + tool_call = ToolUseBlock( + type="tool_use", + id=call_id, + name=part.function_call.name, + input=keyword_args, + raw_input=json.dumps( + keyword_args, + ensure_ascii=False, ), ) + if thought_sig_b64: + tool_call["thought_signature"] = thought_sig_b64 + + tool_calls.append(tool_call) + + # Capture thought_signature on non-FC parts for + # preserving reasoning context in text responses. + if ( + not part.function_call + and getattr(part, "thought_signature", None) + and isinstance(part.thought_signature, bytes) + ): + text_thought_signature = base64.b64encode( + part.thought_signature, + ).decode("utf-8") + + # Apply captured text thought_signature to the last TextBlock + if text_thought_signature: + for block in reversed(content_blocks): + if block.get("type") == "text": + block["thought_signature"] = text_thought_signature + break + # For the structured output case if response.text and structured_model: metadata = _json_loads_with_repair(response.text) diff --git a/src/agentscope/model/_openai_model.py b/src/agentscope/model/_openai_model.py index 21f02747cf..ab06e4e8fd 100644 --- a/src/agentscope/model/_openai_model.py +++ b/src/agentscope/model/_openai_model.py @@ -14,7 +14,7 @@ ) from collections import OrderedDict -from pydantic import BaseModel +from pydantic import BaseModel, ValidationError from . import ChatResponse from ._model_base import ChatModelBase @@ -76,7 +76,16 @@ def __init__( model_name: str, api_key: str | None = None, stream: bool = True, - reasoning_effort: Literal["low", "medium", "high"] | None = None, + reasoning_effort: Literal[ + "none", + "minimal", + "low", + "medium", + "high", + "xhigh", + ] + | str + | None = None, organization: str = None, stream_tool_parsing: bool = True, client_type: Literal["openai", "azure"] = "openai", @@ -94,12 +103,15 @@ def __init__( be read from the environment variable `OPENAI_API_KEY`. stream (`bool`, default `True`): Whether to use streaming output or not. - reasoning_effort (`Literal["low", "medium", "high"] | None`, \ - optional): + reasoning_effort (`Literal["none", "minimal", "low", "medium", \ + "high", "xhigh"] | str | None`, optional): Reasoning effort, supported for o3, o4, etc. Please refer to `OpenAI documentation `_ - for more details. + for more details. The type also accepts arbitrary ``str`` + values so that OpenAI-compatible providers using non-standard + levels can be used directly. For example, DeepSeek supports + ``"max"`` and ``"high"``. organization (`str`, default `None`): The organization ID for OpenAI API. If not specified, it will be read from the environment variable `OPENAI_ORGANIZATION`. @@ -297,6 +309,26 @@ async def __call__( response = await self.client.chat.completions.parse( **kwargs, ) + # Provider-agnostic post-parse validation: some + # OpenAI-compatible endpoints (e.g. DashScope) reply + # HTTP 200 with a body that does not conform to the + # requested schema. The SDK may either raise + # ``ValidationError`` mid-parse (caught below) or + # silently leave ``message.parsed`` as ``None`` / + # a non-matching type (caught here). Both must + # drop to the tool-call fallback rather than + # surfacing a confusing error to the caller. + mismatch = self._structured_parse_mismatch_reason( + response, + structured_model, + ) + if mismatch is not None: + response = await self._fallback_to_tool_call( + kwargs, + structured_model, + start_datetime, + mismatch, + ) else: response = self.client.chat.completions.stream( **kwargs, @@ -307,19 +339,12 @@ async def __call__( structured_model, kwargs, ) - except openai.BadRequestError as e: - logger.warning( - "response_format structured output failed (%s: %s), " - "falling back to tool-call based structured output. " - "Subsequent calls will use tool-call directly.", - type(e).__name__, - e, - ) - self._structured_output_fallback = True - response = await self._structured_via_tool_call( + except (openai.BadRequestError, ValidationError) as e: + response = await self._fallback_to_tool_call( kwargs, structured_model, start_datetime, + f"{type(e).__name__}: {e}", ) if isinstance(response, AsyncGenerator): return response @@ -680,6 +705,69 @@ def _parse_openai_completion_response( return ChatResponse(**resp_kwargs) + @staticmethod + def _structured_parse_mismatch_reason( + response: Any, + structured_model: Type[BaseModel], + ) -> str | None: + """Return a human-readable reason string if ``response.choices[0] + .message.parsed`` is not an instance of ``structured_model``, else + ``None``. + + Covers the case where an endpoint returns HTTP 200 with a body that + does not match the requested ``response_format`` schema and the + OpenAI SDK silently leaves ``message.parsed`` as ``None`` or a + non-matching type instead of raising ``pydantic.ValidationError``. + """ + if not getattr(response, "choices", None): + return "response_format returned no choices" + parsed = getattr(response.choices[0].message, "parsed", None) + if isinstance(parsed, structured_model): + return None + return ( + "response_format returned a body that does not conform to the " + f"requested schema (parsed type: {type(parsed).__name__})" + ) + + @staticmethod + def _is_tool_choice_rejection(err: BaseException) -> bool: + """Return ``True`` if ``err`` looks like an endpoint rejecting a + forced ``tool_choice`` value. + + Centralises the heuristic used by every tool-call retry path: + non-streaming ``BadRequestError`` from ``create()``, streaming + ``BadRequestError`` from ``create()``, and lazy ``APIError`` from + the first ``__anext__``. Substring match is intentional -- vendor + error messages vary (e.g. DashScope: ``"The tool_choice parameter + does not support being set to required..."``). If the message + does not mention ``tool_choice`` we re-raise instead of looping. + """ + return "tool_choice" in str(err) + + async def _fallback_to_tool_call( + self, + kwargs: dict, + structured_model: Type[BaseModel], + start_datetime: datetime, + reason: str, + ) -> Any: + """Log a warning, latch ``_structured_output_fallback``, and route + through ``_structured_via_tool_call``. Used by both the sync ``parse`` + path and the streaming wrapper. + """ + logger.warning( + "response_format structured output failed (%s), falling back " + "to tool-call based structured output. Subsequent calls will " + "use tool-call directly.", + reason, + ) + self._structured_output_fallback = True + return await self._structured_via_tool_call( + kwargs, + structured_model, + start_datetime, + ) + async def _structured_stream_with_fallback( self, start_datetime: datetime, @@ -696,7 +784,16 @@ async def _structured_stream_with_fallback( ``_structured_output_fallback`` is never set. This wrapper catches such errors during stream consumption and - transparently falls back to the tool-call approach. + transparently falls back to the tool-call approach. We catch both: + + * ``openai.BadRequestError`` -- the endpoint rejects + ``response_format`` outright (e.g. DeepSeek), surfaced lazily on + first stream read. + * ``pydantic.ValidationError`` -- the endpoint replies HTTP 200 + with a body that the SDK parses successfully as JSON but fails + to validate against the requested schema (e.g. DashScope + OpenAI-compat returns ``{"message": "Hi!"}`` for a schema with + required fields). """ import openai @@ -707,19 +804,12 @@ async def _structured_stream_with_fallback( structured_model, ): yield chunk - except openai.BadRequestError as e: - logger.warning( - "response_format structured output failed during streaming " - "(%s: %s), falling back to tool-call based structured " - "output. Subsequent calls will use tool-call directly.", - type(e).__name__, - e, - ) - self._structured_output_fallback = True - fallback = await self._structured_via_tool_call( + except (openai.BadRequestError, ValidationError) as e: + fallback = await self._fallback_to_tool_call( kwargs, structured_model, datetime.now(), + f"streaming {type(e).__name__}: {e}", ) if isinstance(fallback, AsyncGenerator): async for chunk in fallback: @@ -737,7 +827,16 @@ async def _structured_via_tool_call( Falls back to this when the API endpoint does not support json_schema response_format (e.g. DashScope, DeepSeek). + + Some reasoning / "thinking" models (e.g. DashScope qwen3.6 thinking + mode) reject any forced ``tool_choice`` with HTTP 400. When that + happens we retry once with ``tool_choice="auto"``; since the tools + list contains only the schema-derived tool, the model is highly + likely to call it anyway, which keeps the structured-output + contract intact for the caller. """ + import openai + kwargs.pop("response_format", None) format_tool = _create_tool_from_base_model(structured_model) kwargs["tools"] = self._format_tools_json_schemas([format_tool]) @@ -747,14 +846,96 @@ async def _structured_via_tool_call( if self.stream: kwargs["stream"] = True kwargs["stream_options"] = {"include_usage": True} - response = await self.client.chat.completions.create(**kwargs) - if self.stream: - return self._parse_openai_stream_response( + return self._tool_call_stream_with_choice_retry( + kwargs, + structured_model, + start_datetime, + ) + try: + return await self.client.chat.completions.create(**kwargs) + except openai.BadRequestError as e: + if not self._is_tool_choice_rejection(e): + raise + logger.warning( + "tool_choice rejected by endpoint (%s: %s); retrying with " + "tool_choice='auto' for structured output fallback.", + type(e).__name__, + e, + ) + kwargs["tool_choice"] = "auto" + return await self.client.chat.completions.create(**kwargs) + + async def _tool_call_stream_with_choice_retry( + self, + kwargs: dict, + structured_model: Type[BaseModel], + start_datetime: datetime, + ) -> AsyncGenerator[ChatResponse, None]: + """Stream-mode counterpart of the ``tool_choice`` retry in + :meth:`_structured_via_tool_call`. + + DashScope thinking-mode rejects forced ``tool_choice`` in two + observed shapes during streaming: + + * **Synchronous** ``openai.BadRequestError`` raised from + ``client.chat.completions.create(**kwargs)`` itself (e.g. + ``qwen3.6-plus`` / ``qwen3.6-flash``). + * **Lazy** ``openai.APIError`` surfaced from an SSE error event + during the first ``__anext__`` (e.g. + ``qwen3-235b-a22b-thinking-2507``). + + Both shapes are caught here -- but only when no chunk has been + yielded yet, so retrying with ``tool_choice="auto"`` is safe and + the caller never sees duplicated output. + """ + import openai + + try: + response = await self.client.chat.completions.create(**kwargs) + except openai.BadRequestError as e: + if not self._is_tool_choice_rejection(e): + raise + logger.warning( + "tool_choice rejected by endpoint on streaming create " + "(%s: %s); retrying with tool_choice='auto' for " + "structured output fallback.", + type(e).__name__, + e, + ) + kwargs["tool_choice"] = "auto" + response = await self.client.chat.completions.create(**kwargs) + + yielded = False + captured: openai.APIError | None = None + try: + async for chunk in self._parse_openai_stream_response( start_datetime, response, structured_model, - ) - return response + ): + yielded = True + yield chunk + return + except openai.APIError as e: + if yielded or not self._is_tool_choice_rejection(e): + raise + captured = e + + logger.warning( + "tool_choice rejected during streaming iteration (%s: %s); " + "retrying with tool_choice='auto' for structured output " + "fallback.", + type(captured).__name__, + captured, + ) + kwargs["tool_choice"] = "auto" + retry_response = await self.client.chat.completions.create(**kwargs) + async for chunk in self._parse_openai_stream_response( + start_datetime, + retry_response, + structured_model, + ): + yield chunk def _format_tools_json_schemas( self, diff --git a/src/agentscope/model/_openai_response_model.py b/src/agentscope/model/_openai_response_model.py new file mode 100644 index 0000000000..0b2d19cde4 --- /dev/null +++ b/src/agentscope/model/_openai_response_model.py @@ -0,0 +1,581 @@ +# -*- coding: utf-8 -*- +# pylint: disable=too-many-branches +"""OpenAI Response API Chat model class.""" +import json +from datetime import datetime +from typing import ( + Any, + List, + AsyncGenerator, + Literal, + Type, +) + +from pydantic import BaseModel + +from . import ChatResponse +from ._model_base import ChatModelBase +from ._model_usage import ChatUsage +from .._logging import logger +from .._utils._common import _json_loads_with_repair +from ..message import ( + ToolUseBlock, + TextBlock, + ThinkingBlock, +) +from ..tracing import trace_llm +from ..types import JSONSerializableObject + + +class OpenAIResponseModel(ChatModelBase): + """Chat model using the OpenAI Responses API + (``client.responses.create``). + + Compared with the Chat Completions API, the Responses API provides + first-class streaming events for reasoning / thinking, text output + and function-call arguments, which makes it a natural fit for models + that expose chain-of-thought reasoning (e.g. ``o3``, ``o4-mini``). + + Compatible with any OpenAI-compatible endpoint by passing a custom + ``base_url`` via ``client_kwargs``. + """ + + def __init__( + self, + model_name: str, + api_key: str | None = None, + stream: bool = True, + reasoning_effort: Literal["minimal", "low", "medium", "high"] + | None = None, + reasoning_summary: Literal[ + "auto", + "concise", + "detailed", + ] + | None = None, + organization: str | None = None, + stream_tool_parsing: bool = True, + client_kwargs: dict[str, JSONSerializableObject] | None = None, + generate_kwargs: dict[str, JSONSerializableObject] | None = None, + **kwargs: Any, + ) -> None: + """Initialize the OpenAI Response API client. + + Args: + model_name (`str`): + The name of the model to use (e.g. ``"qwen3.5-plus"``). + api_key (`str`, optional): + API key. Falls back to ``OPENAI_API_KEY`` env var. + stream (`bool`, default ``True``): + Whether to use streaming output. + reasoning_effort (`Literal["minimal", "low", "medium", \ + "high"]`, optional): + Reasoning effort level. + reasoning_summary (`Literal["auto", "concise", "detailed"]`, \ + optional): + Controls how reasoning summaries are returned in streaming + mode. Defaults to ``"auto"`` when ``reasoning_effort`` + is set. + organization (`str`, optional): + OpenAI organization ID. + stream_tool_parsing (`bool`, default ``True``): + Whether to parse incomplete tool-call JSON during + streaming with auto-repair. + client_kwargs (`dict`, optional): + Extra keyword arguments forwarded to + ``openai.AsyncClient`` (e.g. ``base_url``). + generate_kwargs (`dict`, optional): + Extra keyword arguments forwarded to + ``client.responses.create`` on every call + (e.g. ``temperature``, ``top_p``). + **kwargs: + Ignored (with a warning). + """ + if kwargs: + logger.warning( + "Unknown keyword arguments: %s. These will be ignored.", + list(kwargs.keys()), + ) + + super().__init__(model_name, stream) + + import openai + + self.client = openai.AsyncClient( + api_key=api_key, + organization=organization, + **(client_kwargs or {}), + ) + + self.reasoning_effort = reasoning_effort + self.reasoning_summary = reasoning_summary + self.stream_tool_parsing = stream_tool_parsing + self.generate_kwargs = generate_kwargs or {} + + @trace_llm + async def __call__( + self, + messages: list[dict], + tools: list[dict] | None = None, + tool_choice: Literal["auto", "none", "required"] + | str + | list + | None = None, + structured_model: Type[BaseModel] | None = None, + **kwargs: Any, + ) -> ChatResponse | AsyncGenerator[ChatResponse, None]: + """Call the OpenAI Responses API. + + Args: + messages (`list[dict]`): + A list of message dicts with at least ``role`` and + ``content`` keys. Passed as the ``input`` parameter to + the API. + tools (`list[dict]`, optional): + Tool JSON schemas (Chat-Completions format accepted; + they are automatically converted to the Responses API + format). + tool_choice (`Literal["auto", "none", "required"] | str | list`, + optional): + ``"auto"``, ``"none"``, ``"required"``, a specific + tool name, or a list of tool names. + structured_model (`Type[BaseModel]`, optional): + A Pydantic BaseModel class for structured output. + When provided, the model is instructed to return JSON + conforming to the schema via the ``text.format`` + parameter. ``tools`` and ``tool_choice`` are ignored. + **kwargs: + Forwarded to ``client.responses.create``. + + Returns: + `ChatResponse | AsyncGenerator[ChatResponse, None]` + """ + if not isinstance(messages, list): + raise ValueError( + "OpenAI Response API `messages` field expected type `list`, " + f"got `{type(messages)}` instead.", + ) + + api_kwargs: dict[str, Any] = { + "model": self.model_name, + "input": messages, + "stream": self.stream, + **self.generate_kwargs, + **kwargs, + } + + if self.reasoning_effort and "reasoning" not in api_kwargs: + reasoning_cfg: dict[str, str | None] = { + "effort": self.reasoning_effort, + } + if self.reasoning_summary: + reasoning_cfg["summary"] = self.reasoning_summary + api_kwargs["reasoning"] = reasoning_cfg + + if structured_model: + if tools or tool_choice: + logger.warning( + "structured_model is provided. Both 'tools' and " + "'tool_choice' parameters will be overridden and " + "ignored. The model will only perform structured output " + "generation without calling any other tools.", + ) + api_kwargs.pop("tools", None) + api_kwargs.pop("tool_choice", None) + api_kwargs["text"] = { + "format": { + "type": "json_schema", + "name": structured_model.__name__, + "schema": structured_model.model_json_schema(), + "strict": True, + }, + } + else: + if tools: + api_kwargs["tools"] = self._format_tools(tools) + + if tool_choice: + self._validate_tool_choice(tool_choice, tools) + api_kwargs["tool_choice"] = self._format_tool_choice( + tool_choice, + ) + + start_datetime = datetime.now() + + response = await self.client.responses.create(**api_kwargs) + + if self.stream: + return self._parse_stream_response( + start_datetime, + response, + structured_model, + ) + + return self._parse_response( + start_datetime, + response, + structured_model, + ) + + # ------------------------------------------------------------------ + # Streaming + # ------------------------------------------------------------------ + + async def _parse_stream_response( + self, + start_datetime: datetime, + response: Any, + structured_model: Type[BaseModel] | None = None, + ) -> AsyncGenerator[ChatResponse, None]: + """Parse the event stream produced by the Responses API. + + Recognised event types (``event.type``): + + * ``response.reasoning_summary_text.delta`` – thinking delta + * ``response.output_text.delta`` – text delta + * ``response.output_item.added`` – new output item (may be a + ``function_call``) + * ``response.function_call_arguments.delta`` – tool-call arg delta + * ``response.completed`` – final event carrying usage info + """ + usage: ChatUsage | None = None + response_id: str | None = None + text = "" + thinking = "" + tool_calls: dict[str, dict[str, Any]] = {} + last_input_objs: dict[str, Any] = {} + metadata: dict | None = None + + last_contents = None + + async for event in response: + event_type = event.type + + # ---- capture response id from the first event that has it + if response_id is None: + resp_obj = getattr(event, "response", None) + if resp_obj is not None: + response_id = getattr(resp_obj, "id", None) + + # ---- reasoning / thinking -------------------------------- + if event_type == "response.reasoning_summary_text.delta": + thinking += event.delta + + # ---- text output ----------------------------------------- + elif event_type == "response.output_text.delta": + text += event.delta + + # ---- function call: register new tool call --------------- + elif event_type == "response.output_item.added": + item = event.item + if getattr(item, "type", None) == "function_call": + # NOTE: two distinct ids are in play here. + # * ``item.id`` is the Responses-API stream item id; + # subsequent ``function_call_arguments.delta`` events + # reference it via ``event.item_id``, so it must be + # the dict key. + # * ``call_id`` is the public identifier used to pair + # the call with its later ``function_call_output``; + # it's what we expose on ``ToolUseBlock.id``. + # Do not collapse them — downstream tool-result + # matching relies on ``call_id``. + call_id = getattr(item, "call_id", None) or getattr( + item, + "id", + "", + ) + tool_calls[item.id] = { + "type": "tool_use", + "id": call_id, + "name": getattr(item, "name", ""), + "input": "", + } + + # ---- function call: argument deltas ---------------------- + elif event_type == "response.function_call_arguments.delta": + item_id = event.item_id + if item_id in tool_calls: + tool_calls[item_id]["input"] += event.delta + + # ---- completion (usage) ---------------------------------- + elif event_type == "response.completed": + resp = event.response + if response_id is None: + response_id = getattr(resp, "id", None) + if resp.usage: + usage = ChatUsage( + input_tokens=resp.usage.input_tokens, + output_tokens=resp.usage.output_tokens, + time=(datetime.now() - start_datetime).total_seconds(), + metadata=resp.usage, + ) + + # ---- build content blocks and yield ---------------------- + contents = self._build_content_blocks( + thinking, + text, + tool_calls, + last_input_objs, + ) + + if structured_model and text: + metadata = _json_loads_with_repair(text) + + if contents: + chat_resp_kwargs: dict[str, Any] = { + "content": contents, + "usage": usage, + "metadata": metadata, + } + if response_id: + chat_resp_kwargs["id"] = response_id + yield ChatResponse(**chat_resp_kwargs) + last_contents = [dict(b) for b in contents] + + # When stream_tool_parsing is disabled, yield a final response + # with properly parsed tool-call inputs after the stream ends. + if not self.stream_tool_parsing and tool_calls and last_contents: + for block in last_contents: + if block.get("type") == "tool_use": + block["input"] = _json_loads_with_repair( + str(block.get("raw_input") or "{}"), + ) + final_kwargs: dict[str, Any] = { + "content": last_contents, + "usage": usage, + "metadata": metadata, + } + if response_id: + final_kwargs["id"] = response_id + yield ChatResponse(**final_kwargs) + + # ------------------------------------------------------------------ + # Non-streaming + # ------------------------------------------------------------------ + + def _parse_response( + self, + start_datetime: datetime, + response: Any, + structured_model: Type[BaseModel] | None = None, + ) -> ChatResponse: + """Parse a non-streaming ``Response`` object.""" + content_blocks: List[TextBlock | ToolUseBlock | ThinkingBlock] = [] + metadata: dict | None = None + + for item in response.output: + item_type = getattr(item, "type", None) + + if item_type == "reasoning": + for summary in getattr(item, "summary", []): + summary_text = getattr(summary, "text", "") + if summary_text: + content_blocks.append( + ThinkingBlock( + type="thinking", + thinking=summary_text, + ), + ) + + elif item_type == "message": + for part in getattr(item, "content", []): + if getattr(part, "type", None) == "output_text": + content_blocks.append( + TextBlock(type="text", text=part.text), + ) + if structured_model: + metadata = _json_loads_with_repair(part.text) + + elif item_type == "function_call": + call_id = getattr(item, "call_id", None) or getattr( + item, + "id", + "", + ) + content_blocks.append( + ToolUseBlock( + type="tool_use", + id=call_id, + name=item.name, + input=_json_loads_with_repair( + getattr(item, "arguments", "") or "{}", + ), + ), + ) + + usage = None + if response.usage: + usage = ChatUsage( + input_tokens=response.usage.input_tokens, + output_tokens=response.usage.output_tokens, + time=(datetime.now() - start_datetime).total_seconds(), + metadata=response.usage, + ) + + resp_kwargs: dict[str, Any] = { + "content": content_blocks, + "usage": usage, + "metadata": metadata, + } + response_id = getattr(response, "id", None) + if response_id: + resp_kwargs["id"] = response_id + + return ChatResponse(**resp_kwargs) + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + + def _build_content_blocks( + self, + thinking: str, + text: str, + tool_calls: dict[str, dict[str, Any]], + last_input_objs: dict[str, Any], + ) -> List[TextBlock | ToolUseBlock | ThinkingBlock]: + """Assemble content blocks from accumulated state.""" + contents: List[TextBlock | ToolUseBlock | ThinkingBlock] = [] + + if thinking: + contents.append( + ThinkingBlock(type="thinking", thinking=thinking), + ) + + if text: + contents.append(TextBlock(type="text", text=text)) + + for tc in tool_calls.values(): + input_str = tc["input"] + + if self.stream_tool_parsing: + repaired = _json_loads_with_repair(input_str or "{}") + last = last_input_objs.get(tc["id"], {}) + if len(json.dumps(last)) > len(json.dumps(repaired)): + repaired = last + last_input_objs[tc["id"]] = repaired + else: + repaired = {} + + contents.append( + ToolUseBlock( + type="tool_use", + id=tc["id"], + name=tc["name"], + input=repaired, + raw_input=input_str, + ), + ) + + return contents + + @staticmethod + def _format_tools( + schemas: list[dict[str, Any]], + ) -> list[dict[str, Any]]: + """Format the tools JSON schema into OpenAI realtime model format. + + Args: + schemas (`list[dict[str, Any]]`): + The tool schemas. + + Returns: + `list[dict[str, Any]]`: + The formatted tools for OpenAI realtime model. + + .. note:: + The OpenAI Realtime API uses a different tool format compared to + the regular Chat Completions API. While the Chat API expects tools + to be wrapped in ``{"type": "function", "function": {...}}``, the + Realtime API expects a flattened structure where the function + definition is directly at the top level with an added ``"type": + "function"`` field. + """ + formatted: list[dict[str, Any]] = [] + for tool in schemas: + # Accept both Chat-Completions wrapped form + # ({"type": "function", "function": {...}}) and already-flat + # Responses-API form ({"type": "function", "name": ..., ...}). + if "function" in tool and isinstance(tool["function"], dict): + formatted.append({"type": "function", **tool["function"]}) + else: + formatted.append({"type": "function", **tool}) + return formatted + + def _validate_tool_choice( + self, + tool_choice: str | list, + tools: list[dict] | None, + ) -> None: + """Validate tool_choice parameter, supporting list of tool names. + + Extends the base class validation to additionally accept a list of + tool names for OpenAI's ``allowed_tools`` feature. + + Args: + tool_choice (`str | list`): + Tool choice mode, function name, or a list of function names. + tools (`list[dict] | None`): + Available tools list. + Raises: + TypeError: If tool_choice type is invalid. + ValueError: If tool_choice value is invalid. + """ + if isinstance(tool_choice, list): + if not tool_choice: + raise ValueError( + "tool_choice list must not be empty.", + ) + if not all(isinstance(name, str) for name in tool_choice): + raise TypeError( + "All elements in tool_choice list must be str.", + ) + if not tools: + raise ValueError( + "tools must be provided when tool_choice is a list.", + ) + available_functions = [tool["function"]["name"] for tool in tools] + for name in tool_choice: + if name not in available_functions: + raise ValueError( + f"Invalid tool name '{name}' in tool_choice list. " + f"Available functions: " + f"{', '.join(sorted(available_functions))}", + ) + return + + super()._validate_tool_choice(tool_choice, tools) + + def _format_tool_choice( + self, + tool_choice: Literal["auto", "none", "required"] | str | list | None, + ) -> str | dict | None: + """Format tool_choice parameter for API compatibility. + + Args: + tool_choice (`Literal["auto", "none", "required"] | str \ + | list | None`, default `None`): + Controls which (if any) tool is called by the model. + Can be "auto", "none", "required", a specific tool name, + or a list of tool names. For more details, please refer to + https://platform.openai.com/docs/api-reference/responses/create#responses_create-tool_choice + Returns: + `str | dict | None`: + The formatted tool choice configuration, or None if + tool_choice is None. + """ + if tool_choice is None: + return None + + if isinstance(tool_choice, list): + return { + "type": "allowed_tools", + "mode": "auto", + "tools": [ + {"type": "function", "name": name} for name in tool_choice + ], + } + + if tool_choice in ("auto", "none", "required"): + return tool_choice + return {"type": "function", "name": tool_choice} diff --git a/src/agentscope/tracing/_trace.py b/src/agentscope/tracing/_trace.py index 9c0ecd6053..834e4b2c58 100644 --- a/src/agentscope/tracing/_trace.py +++ b/src/agentscope/tracing/_trace.py @@ -1,6 +1,7 @@ # -*- coding: utf-8 -*- """The tracing decorators for agent, formatter, toolkit, chat and embedding models.""" +import asyncio import inspect from functools import wraps from typing import ( @@ -89,13 +90,13 @@ def _set_span_success_status(span: Span) -> None: span.end() -def _set_span_error_status(span: Span, e: Exception) -> None: +def _set_span_error_status(span: Span, e: BaseException) -> None: """Set the status of the span. Args: span (`Span`): The OpenTelemetry span to be used for tracing. - e (`Exception`): - The exception to be recorded. + e (`BaseException`): + The BaseException to be recorded. """ from opentelemetry import trace as trace_api @@ -155,7 +156,7 @@ async def _trace_async_generator_wrapper( last_chunk = chunk yield chunk - except Exception as e: + except (asyncio.CancelledError, Exception) as e: has_error = True _set_span_error_status(span, e) raise e from None @@ -263,7 +264,7 @@ async def wrapper( _set_span_success_status(span) return res - except Exception as e: + except (asyncio.CancelledError, Exception) as e: _set_span_error_status(span, e) raise e from None @@ -356,9 +357,8 @@ async def wrapper( # Return the wrapped generator return _trace_async_generator_wrapper(res, span) - except Exception as e: + except (asyncio.CancelledError, Exception) as e: _set_span_error_status(span, e) - span.end() raise e from None return wrapper @@ -426,7 +426,7 @@ async def wrapper( _set_span_success_status(span) return res - except Exception as e: + except (asyncio.CancelledError, Exception) as e: _set_span_error_status(span, e) raise e from None @@ -486,7 +486,7 @@ async def wrapper( _set_span_success_status(span) return res - except Exception as e: + except (asyncio.CancelledError, Exception) as e: _set_span_error_status(span, e) raise e from None @@ -557,7 +557,7 @@ async def wrapper( _set_span_success_status(span) return res - except Exception as e: + except (asyncio.CancelledError, Exception) as e: _set_span_error_status(span, e) raise e from None @@ -639,7 +639,7 @@ async def async_wrapper( _set_span_success_status(span) return res - except Exception as e: + except (asyncio.CancelledError, Exception) as e: _set_span_error_status(span, e) raise e from None diff --git a/src/agentscope/tts/_dashscope_realtime_tts_model.py b/src/agentscope/tts/_dashscope_realtime_tts_model.py index df02bdd2b0..9b5b69d17b 100644 --- a/src/agentscope/tts/_dashscope_realtime_tts_model.py +++ b/src/agentscope/tts/_dashscope_realtime_tts_model.py @@ -1,11 +1,16 @@ # -*- coding: utf-8 -*- +# pylint: disable=too-many-branches, too-many-statements """DashScope Realtime TTS model implementation.""" +import asyncio import threading from typing import Any, Literal, TYPE_CHECKING, AsyncGenerator +from websocket import WebSocketConnectionClosedException + from ._tts_base import TTSModelBase from ._tts_response import TTSResponse +from .._logging import logger from ..message import Msg, AudioBlock, Base64Source from ..types import JSONSerializableObject @@ -82,6 +87,26 @@ def on_event(self, response: dict[str, Any]) -> None: traceback.print_exc() self.finish_event.set() + def on_close(self, close_status_code: int, close_msg: str) -> None: + """Called when the WebSocket connection is closed. + + Args: + close_status_code (`int`): + The close status code. + close_msg (`str`): + The close message. + """ + # Unblock waiting operations to prevent deadlock + self.finish_event.set() + self.chunk_event.set() + + if close_status_code: + logger.warning( + "TTS WebSocket connection closed with code %s: %s", + close_status_code, + close_msg, + ) + async def get_audio_data(self, block: bool) -> TTSResponse: """Get the current accumulated audio data as base64 string so far. @@ -164,6 +189,10 @@ async def _reset(self) -> None: self.chunk_event.clear() self._audio_data = "" + def has_audio_data(self) -> bool: + """Check if audio data has been received.""" + return bool(self._audio_data) + return _DashScopeRealtimeTTSCallback @@ -196,6 +225,8 @@ def __init__( cold_start_words: int | None = None, client_kwargs: dict[str, JSONSerializableObject] | None = None, generate_kwargs: dict[str, JSONSerializableObject] | None = None, + max_retries: int = 3, + retry_delay: float = 5.0, ) -> None: """Initialize the DashScope TTS model by specifying the model, voice, and other parameters. @@ -240,6 +271,10 @@ def __init__( optional): The extra keyword arguments used in DashScope realtime tts API generation. + max_retries (`int`, defaults to 3): + The maximum number of retry attempts when TTS synthesis fails. + retry_delay (`float`, defaults to 5.0): + The delay in seconds before retrying. Uses exponential backoff. """ super().__init__(model_name=model_name, stream=stream) @@ -255,6 +290,8 @@ def __init__( self.cold_start_words = cold_start_words self.client_kwargs = client_kwargs or {} self.generate_kwargs = generate_kwargs or {} + self.max_retries = max_retries + self.retry_delay = retry_delay # Initialize TTS client # Save callback reference (for DashScope SDK) @@ -298,9 +335,29 @@ async def close(self) -> None: self._connected = False - self._tts_client.finish() self._tts_client.close() + async def _reconnect(self) -> None: + """Reconnect to TTS service by recreating the client.""" + from dashscope.audio.qwen_tts_realtime import QwenTtsRealtime + + try: + self._tts_client.close() + except Exception: + pass + + self._dashscope_callback = _get_qwen_tts_realtime_callback_class()() + self._tts_client = QwenTtsRealtime( + model=self.model_name, + callback=self._dashscope_callback, + **self.client_kwargs, + ) + self._connected = False + self._first_send = True + self._current_msg_id = None + self._current_prefix = "" + await self.connect() + async def push( self, msg: Msg, @@ -362,7 +419,12 @@ async def push( delta_to_send = text.removeprefix(self._current_prefix) if delta_to_send: - self._tts_client.append_text(delta_to_send) + try: + self._tts_client.append_text(delta_to_send) + except WebSocketConnectionClosedException: + # Connection closed, return empty response + # synthesize() will handle retry + return TTSResponse(content=None) # Record sent prefix self._current_prefix += delta_to_send @@ -399,7 +461,11 @@ async def synthesize( "TTS model is not connected. Call `connect()` first.", ) - if self._current_msg_id is not None and self._current_msg_id != msg.id: + if ( + self._current_msg_id is not None + and msg + and self._current_msg_id != msg.id + ): raise RuntimeError( "DashScopeRealtimeTTSModel can only handle one streaming " "input request at a time. Please ensure that all chunks " @@ -416,19 +482,85 @@ async def synthesize( self._current_prefix, ) - # Determine if we should send text based on cold start settings only - # for the first input chunk and not the last chunk - if delta_to_send: - self._tts_client.append_text(delta_to_send) + full_text = (msg.get_text_content() or "") if msg else "" - # To keep correct prefix tracking - self._current_prefix += delta_to_send - self._first_send = False + # Synthesize with retry - if we have text but get no audio, retry + delay = self.retry_delay - # We need to block until synthesis is complete to get all audio - self._tts_client.commit() - self._tts_client.finish() + for attempt in range(self.max_retries): + try: + # Send remaining text if any + if delta_to_send: + self._tts_client.append_text(delta_to_send) + self._current_prefix += delta_to_send + self._first_send = False + + # Commit and finish + self._tts_client.commit() + self._tts_client.finish() + + # Wait for synthesis to complete + self._dashscope_callback.finish_event.wait() + + # Check if we got audio (only retry if we had text but no + # audio) + has_audio = self._dashscope_callback.has_audio_data() + if full_text and not has_audio: + if attempt < self.max_retries - 1: + logger.warning( + "TTS: no audio received, retrying (%d/%d) in " + "%.1fs...", + attempt + 1, + self.max_retries, + delay, + ) + await asyncio.sleep(delay) + await self._reconnect() + # After reconnect, need to resend full text + delta_to_send = full_text + delay *= 2 + continue + logger.error( + "TTS: no audio after %d attempts.", + self.max_retries, + ) + # Reset state before raising + self._current_msg_id = None + self._first_send = True + self._current_prefix = "" + raise RuntimeError( + f"TTS synthesis failed: no audio after" + f" {self.max_retries} attempts", + ) + + # Success + break + + except WebSocketConnectionClosedException: + if attempt < self.max_retries - 1: + logger.warning( + "TTS failed, retrying (%d/%d) in %.1fs...", + attempt + 1, + self.max_retries, + delay, + ) + await asyncio.sleep(delay) + await self._reconnect() + # After reconnect, need to resend full text + delta_to_send = full_text + delay *= 2 + else: + logger.error( + "TTS failed after %d attempts.", + self.max_retries, + ) + # Reset state before raising + self._current_msg_id = None + self._first_send = True + self._current_prefix = "" + raise + # Get result if self.stream: # Return an async generator for audio chunks res = self._dashscope_callback.get_audio_chunk() diff --git a/tests/formatter_gemini_test.py b/tests/formatter_gemini_test.py index bb4d49aae5..a9aa5e8bd1 100644 --- a/tests/formatter_gemini_test.py +++ b/tests/formatter_gemini_test.py @@ -1,5 +1,7 @@ # -*- coding: utf-8 -*- +# pylint: disable=too-many-lines """The gemini formatter unittests.""" +import base64 import os from unittest.async_case import IsolatedAsyncioTestCase from unittest.mock import patch, MagicMock @@ -93,6 +95,9 @@ async def asyncSetUp(self) -> None: ), ] + self.thought_sig_1 = base64.b64encode(b"sig_1").decode("utf-8") + self.thought_sig_2 = base64.b64encode(b"sig_2").decode("utf-8") + self.msgs_tools = [ Msg( "assistant", @@ -102,6 +107,7 @@ async def asyncSetUp(self) -> None: id="1", name="get_capital", input={"country": "Japan"}, + thought_signature=self.thought_sig_1, ), ], "assistant", @@ -162,6 +168,7 @@ async def asyncSetUp(self) -> None: id="2", name="get_capital", input={"country": "South Korea"}, + thought_signature=self.thought_sig_2, ), ], "assistant", @@ -277,7 +284,7 @@ async def asyncSetUp(self) -> None: "country": "Japan", }, }, - "thought_signature": "1", + "thought_signature": b"sig_1", }, ], }, @@ -360,7 +367,7 @@ async def asyncSetUp(self) -> None: "country": "Japan", }, }, - "thought_signature": "1", + "thought_signature": b"sig_1", }, ], }, @@ -413,7 +420,7 @@ async def asyncSetUp(self) -> None: "country": "Japan", }, }, - "thought_signature": "1", + "thought_signature": b"sig_1", }, ], }, @@ -499,7 +506,7 @@ async def asyncSetUp(self) -> None: "country": "Japan", }, }, - "thought_signature": "1", + "thought_signature": b"sig_1", }, ], }, @@ -542,7 +549,7 @@ async def asyncSetUp(self) -> None: "country": "South Korea", }, }, - "thought_signature": "2", + "thought_signature": b"sig_2", }, ], }, @@ -711,7 +718,7 @@ async def test_chat_formatter_with_extract_image_blocks( "country": "Japan", }, }, - "thought_signature": "1", + "thought_signature": b"sig_1", }, ], }, @@ -934,7 +941,7 @@ async def test_multi_agent_formatter_with_promote_tool_result_images( "country": "Japan", }, }, - "thought_signature": "1", + "thought_signature": b"sig_1", }, ], }, diff --git a/tests/formatter_openai_response_test.py b/tests/formatter_openai_response_test.py new file mode 100644 index 0000000000..5b5e1f853a --- /dev/null +++ b/tests/formatter_openai_response_test.py @@ -0,0 +1,379 @@ +# -*- coding: utf-8 -*- +"""The OpenAI Response formatter unittests.""" +import os +from unittest.async_case import IsolatedAsyncioTestCase +from unittest.mock import patch, MagicMock + +from agentscope.formatter._openai_response_formatter import ( + OpenAIResponseChatFormatter, + OpenAIResponseMultiAgentFormatter, +) +from agentscope.message import ( + Msg, + TextBlock, + ImageBlock, + AudioBlock, + URLSource, + ToolResultBlock, + ToolUseBlock, + Base64Source, +) + + +class TestOpenAIResponseFormatter(IsolatedAsyncioTestCase): + """OpenAI Response formatter unittests.""" + + async def asyncSetUp(self) -> None: + """Set up the test environment.""" + self.image_path = os.path.abspath("./image_resp.png") + with open(self.image_path, "wb") as f: + f.write(b"fake image content") + + self.mock_audio_path = ( + "/var/folders/gf/krg8x_ws409cpw_46b2s6rjc0000gn/T/tmpfymnv2w9.wav" + ) + + self.audio_path = os.path.abspath("./audio_resp.wav") + with open(self.audio_path, "wb") as f: + f.write(b"fake audio content") + + self.msgs_system = [ + Msg("system", "You're a helpful assistant.", "system"), + ] + + self.msgs_conversation = [ + Msg( + "user", + [ + TextBlock( + type="text", + text="What is the capital of France?", + ), + ImageBlock( + type="image", + source=URLSource(type="url", url=self.image_path), + ), + ], + "user", + ), + Msg("assistant", "The capital of France is Paris.", "assistant"), + Msg( + "user", + [ + TextBlock( + type="text", + text="What is the capital of Germany?", + ), + AudioBlock( + type="audio", + source=URLSource(type="url", url=self.audio_path), + ), + ], + "user", + ), + Msg( + "assistant", + "The capital of Germany is Berlin.", + "assistant", + ), + Msg("user", "What is the capital of Japan?", "user"), + ] + + self.msgs_tools = [ + Msg( + "assistant", + [ + ToolUseBlock( + type="tool_use", + id="1", + name="get_capital", + input={"country": "Japan"}, + ), + ], + "assistant", + ), + Msg( + "system", + [ + ToolResultBlock( + type="tool_result", + id="1", + name="get_capital", + output=[ + TextBlock( + type="text", + text="The capital of Japan is Tokyo.", + ), + ImageBlock( + type="image", + source=URLSource( + type="url", + url=self.image_path, + ), + ), + AudioBlock( + type="audio", + source=Base64Source( + type="base64", + media_type="audio/wav", + data="ZmFrZSBhdWRpbyBjb250ZW50", + ), + ), + ], + ), + ], + "system", + ), + Msg("assistant", "The capital of Japan is Tokyo.", "assistant"), + ] + + tool_result_text = ( + "- The capital of Japan is Tokyo.\n" + "- The returned image can be found at: " + f"{self.image_path}\n" + "- The returned audio can be found at: " + f"{self.mock_audio_path}" + ) + + self.ground_truth_chat = [ + { + "role": "system", + "content": [ + { + "type": "input_text", + "text": "You're a helpful assistant.", + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "What is the capital of France?", + }, + { + "type": "input_image", + "image_url": "data:image/png;" + "base64,ZmFrZSBpbWFnZSBjb250ZW50", + }, + ], + }, + { + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "The capital of France is Paris.", + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "What is the capital of Germany?", + }, + { + "type": "input_audio", + "input_audio": { + "data": "ZmFrZSBhdWRpbyBjb250ZW50", + "format": "wav", + }, + }, + ], + }, + { + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "The capital of Germany is Berlin.", + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "What is the capital of Japan?", + }, + ], + }, + { + "type": "function_call", + "call_id": "1", + "name": "get_capital", + "arguments": '{"country": "Japan"}', + }, + { + "type": "function_call_output", + "call_id": "1", + "output": tool_result_text, + }, + { + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "The capital of Japan is Tokyo.", + }, + ], + }, + ] + + self.ground_truth_multiagent = [ + { + "role": "system", + "content": [ + { + "type": "input_text", + "text": "You're a helpful assistant.", + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "# Conversation History\n" + "The content between tags contains" + " your conversation history\n" + "\n" + "user: What is the capital of France?\n" + "assistant: The capital of France is Paris.\n" + "user: What is the capital of Germany?\n" + "assistant: The capital of Germany is Berlin.\n" + "user: What is the capital of Japan?\n" + "", + }, + { + "type": "input_image", + "image_url": "data:image/png;base64," + "ZmFrZSBpbWFnZSBjb250ZW50", + }, + { + "type": "input_audio", + "input_audio": { + "data": "ZmFrZSBhdWRpbyBjb250ZW50", + "format": "wav", + }, + }, + ], + }, + { + "type": "function_call", + "call_id": "1", + "name": "get_capital", + "arguments": '{"country": "Japan"}', + }, + { + "type": "function_call_output", + "call_id": "1", + "output": tool_result_text, + }, + { + "role": "user", + "content": [ + { + "type": "input_text", + "text": "\n" + "assistant: The capital of Japan is Tokyo.\n" + "", + }, + ], + }, + ] + + @patch("agentscope.formatter._formatter_base._save_base64_data") + async def test_chat_formatter( + self, + mock_save_base64_data: MagicMock, + ) -> None: + """Test the OpenAI Response chat formatter with full history.""" + mock_save_base64_data.return_value = self.mock_audio_path + formatter = OpenAIResponseChatFormatter() + + res = await formatter.format( + [*self.msgs_system, *self.msgs_conversation, *self.msgs_tools], + ) + self.assertListEqual(res, self.ground_truth_chat) + + @patch("agentscope.formatter._formatter_base._save_base64_data") + async def test_chat_formatter_with_promote_images( + self, + mock_save_base64_data: MagicMock, + ) -> None: + """Test chat formatter with promote_tool_result_images=True.""" + mock_save_base64_data.return_value = self.mock_audio_path + formatter = OpenAIResponseChatFormatter( + promote_tool_result_images=True, + ) + + res = await formatter.format( + [*self.msgs_system, *self.msgs_conversation, *self.msgs_tools], + ) + + # Verify tool result format (Responses API: function_call_output item) + tool_results = [ + m for m in res if m.get("type") == "function_call_output" + ] + self.assertEqual(len(tool_results), 1) + self.assertEqual(tool_results[0]["call_id"], "1") + + # Verify promoted image block is inserted as a separate user message + promoted_msgs = [ + m + for m in res + if m.get("role") == "user" + and any( + "" in (b.get("text", "") or "") + for b in (m.get("content") or []) + ) + ] + self.assertEqual(len(promoted_msgs), 1) + + promoted_content = promoted_msgs[0]["content"] + img_blocks = [ + b for b in promoted_content if b.get("type") == "input_image" + ] + self.assertEqual(len(img_blocks), 1) + + @patch("agentscope.formatter._formatter_base._save_base64_data") + async def test_multiagent_formatter( + self, + mock_save_base64_data: MagicMock, + ) -> None: + """Test the OpenAI Response multi-agent formatter.""" + mock_save_base64_data.return_value = self.mock_audio_path + formatter = OpenAIResponseMultiAgentFormatter() + + res = await formatter.format( + [*self.msgs_system, *self.msgs_conversation, *self.msgs_tools], + ) + self.assertListEqual(res, self.ground_truth_multiagent) + + @patch("agentscope.formatter._formatter_base._save_base64_data") + async def test_multiagent_system_message_uses_input_text( + self, + mock_save_base64_data: MagicMock, + ) -> None: + """Verify multi-agent formatter system message uses 'input_text'.""" + mock_save_base64_data.return_value = self.mock_audio_path + formatter = OpenAIResponseMultiAgentFormatter() + + res = await formatter.format(self.msgs_system) + system_msg = res[0] + self.assertEqual(system_msg["role"], "system") + for block in system_msg["content"]: + self.assertEqual(block["type"], "input_text") + + async def asyncTearDown(self) -> None: + """Clean up the test environment.""" + if os.path.exists(self.image_path): + os.remove(self.image_path) + if os.path.exists(self.audio_path): + os.remove(self.audio_path) diff --git a/tests/model_gemini_test.py b/tests/model_gemini_test.py index 3ebcce5946..f848b95c0d 100644 --- a/tests/model_gemini_test.py +++ b/tests/model_gemini_test.py @@ -35,6 +35,7 @@ def __init__( part.text = text part.thought = False part.function_call = None + part.thought_signature = None first_candidate = Mock() first_candidate.content = Mock() @@ -60,7 +61,7 @@ def _create_usage_mock(self, usage_data: dict) -> Mock: class GeminiFunctionCallMock: """Mock class for Gemini function calls.""" - def __init__(self, call_id: str, name: str, args: dict = None): + def __init__(self, call_id: str | None, name: str, args: dict = None): self.id = call_id self.name = name self.args = args or {} @@ -360,6 +361,159 @@ async def test_streaming_response_processing(self) -> None: ] self.assertEqual(final_response.content, expected_content) + async def test_tool_call_with_thought_signature(self) -> None: + """Test that thought_signature is captured in ToolUseBlock.""" + with patch("google.genai.Client") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + model = GeminiChatModel( + model_name="gemini-3-pro", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + messages = [{"role": "user", "content": "Check the weather"}] + + sig_bytes = b"test_thought_signature_bytes" + fc_part = GeminiPartMock() + fc_part.function_call = GeminiFunctionCallMock( + call_id=None, + name="get_weather", + args={"location": "Paris"}, + ) + fc_part.thought_signature = sig_bytes + + candidate = GeminiCandidateMock(parts=[fc_part]) + mock_response = GeminiResponseMock( + candidates=[candidate], + usage_metadata={ + "prompt_token_count": 10, + "total_token_count": 30, + }, + ) + + mock_client.aio.models.generate_content = AsyncMock( + return_value=mock_response, + ) + result = await model(messages) + + import base64 + + expected_sig = base64.b64encode(sig_bytes).decode("utf-8") + tool_blocks = [ + b for b in result.content if b.get("type") == "tool_use" + ] + self.assertEqual(len(tool_blocks), 1) + self.assertEqual(tool_blocks[0]["id"], expected_sig) + self.assertEqual( + tool_blocks[0]["thought_signature"], + expected_sig, + ) + + async def test_parallel_tool_calls_thought_signature(self) -> None: + """Test parallel FCs: only the first gets thought_signature.""" + with patch("google.genai.Client") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + model = GeminiChatModel( + model_name="gemini-3-pro", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + sig_bytes = b"sig_for_first_fc" + fc1 = GeminiPartMock() + fc1.function_call = GeminiFunctionCallMock( + call_id=None, + name="get_weather", + args={"location": "Paris"}, + ) + fc1.thought_signature = sig_bytes + + fc2 = GeminiPartMock() + fc2.function_call = GeminiFunctionCallMock( + call_id=None, + name="get_weather", + args={"location": "London"}, + ) + fc2.thought_signature = None + + candidate = GeminiCandidateMock(parts=[fc1, fc2]) + mock_response = GeminiResponseMock( + candidates=[candidate], + usage_metadata={ + "prompt_token_count": 10, + "total_token_count": 30, + }, + ) + + mock_client.aio.models.generate_content = AsyncMock( + return_value=mock_response, + ) + result = await model([{"role": "user", "content": "Weather?"}]) + + tool_blocks = [ + b for b in result.content if b.get("type") == "tool_use" + ] + self.assertEqual(len(tool_blocks), 2) + + import base64 + + expected_sig = base64.b64encode(sig_bytes).decode("utf-8") + self.assertEqual( + tool_blocks[0].get("thought_signature"), + expected_sig, + ) + self.assertNotIn("thought_signature", tool_blocks[1]) + self.assertEqual(tool_blocks[1]["id"], "call_1") + + async def test_text_thought_signature(self) -> None: + """Test thought_signature on text parts (non-FC responses).""" + with patch("google.genai.Client") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + + model = GeminiChatModel( + model_name="gemini-3-pro", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + sig_bytes = b"text_reasoning_signature" + text_part = GeminiPartMock(text="Here is the answer.") + text_part.thought_signature = sig_bytes + + candidate = GeminiCandidateMock(parts=[text_part]) + mock_response = GeminiResponseMock( + candidates=[candidate], + usage_metadata={ + "prompt_token_count": 10, + "total_token_count": 30, + }, + ) + + mock_client.aio.models.generate_content = AsyncMock( + return_value=mock_response, + ) + result = await model([{"role": "user", "content": "Question?"}]) + + import base64 + + expected_sig = base64.b64encode(sig_bytes).decode("utf-8") + text_blocks = [ + b for b in result.content if b.get("type") == "text" + ] + self.assertEqual(len(text_blocks), 1) + self.assertEqual( + text_blocks[0].get("thought_signature"), + expected_sig, + ) + async def test_generate_kwargs_integration(self) -> None: """Test integration of generate_kwargs.""" with patch("google.genai.Client") as mock_client_class: diff --git a/tests/model_openai_response_test.py b/tests/model_openai_response_test.py new file mode 100644 index 0000000000..e2120ed75f --- /dev/null +++ b/tests/model_openai_response_test.py @@ -0,0 +1,530 @@ +# -*- coding: utf-8 -*- +"""Unit tests for OpenAI Response API model class.""" +from typing import Any +from unittest.async_case import IsolatedAsyncioTestCase +from unittest.mock import Mock, patch, AsyncMock + +from pydantic import BaseModel + +from agentscope.model import ChatResponse +from agentscope.model._openai_response_model import OpenAIResponseModel +from agentscope.message import TextBlock, ToolUseBlock, ThinkingBlock + + +class SampleModel(BaseModel): + """Sample Pydantic model for testing structured output.""" + + name: str + age: int + + +class TestOpenAIResponseModel(IsolatedAsyncioTestCase): + """Test cases for OpenAIResponseModel.""" + + def test_init_default_params(self) -> None: + """Test initialization with default parameters.""" + with patch("openai.AsyncClient") as mock_client: + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + ) + self.assertEqual(model.model_name, "o3") + self.assertTrue(model.stream) + self.assertIsNone(model.reasoning_effort) + self.assertEqual(model.generate_kwargs, {}) + mock_client.assert_called_once_with( + api_key="test_key", + organization=None, + ) + + def test_init_with_custom_params(self) -> None: + """Test initialization with custom parameters.""" + generate_kwargs = {"temperature": 0.7} + client_kwargs = {"timeout": 30, "base_url": "https://custom.api/v1"} + with patch("openai.AsyncClient") as mock_client: + model = OpenAIResponseModel( + model_name="o4-mini", + api_key="test_key", + stream=False, + reasoning_effort="high", + reasoning_summary="concise", + organization="org-123", + client_kwargs=client_kwargs, + generate_kwargs=generate_kwargs, + ) + self.assertFalse(model.stream) + self.assertEqual(model.reasoning_effort, "high") + self.assertEqual(model.reasoning_summary, "concise") + mock_client.assert_called_once_with( + api_key="test_key", + organization="org-123", + timeout=30, + base_url="https://custom.api/v1", + ) + + async def test_non_streaming_text_response(self) -> None: + """Test non-streaming with a simple text response.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + mock_response = self._create_mock_response( + text="Hello! How can I help?", + ) + mock_client.responses.create = AsyncMock( + return_value=mock_response, + ) + + result = await model( + [{"role": "user", "content": "Hello"}], + ) + + call_args = mock_client.responses.create.call_args[1] + self.assertEqual(call_args["model"], "o3") + self.assertFalse(call_args["stream"]) + self.assertIsInstance(result, ChatResponse) + expected = [TextBlock(type="text", text="Hello! How can I help?")] + self.assertEqual(result.content, expected) + + async def test_non_streaming_with_reasoning(self) -> None: + """Test non-streaming response with reasoning/thinking.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=False, + reasoning_effort="high", + reasoning_summary="concise", + ) + model.client = mock_client + + mock_response = self._create_mock_response( + text="The answer is 42.", + reasoning_summary="Let me think step by step...", + ) + mock_client.responses.create = AsyncMock( + return_value=mock_response, + ) + + result = await model( + [{"role": "user", "content": "Think hard"}], + ) + + call_args = mock_client.responses.create.call_args[1] + self.assertEqual( + call_args["reasoning"], + {"effort": "high", "summary": "concise"}, + ) + expected = [ + ThinkingBlock( + type="thinking", + thinking="Let me think step by step...", + ), + TextBlock(type="text", text="The answer is 42."), + ] + self.assertEqual(result.content, expected) + + async def test_non_streaming_with_function_call(self) -> None: + """Test non-streaming response with a function call.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + mock_response = self._create_mock_response( + function_calls=[ + { + "call_id": "call_abc", + "name": "get_weather", + "arguments": '{"city": "Beijing"}', + }, + ], + ) + mock_client.responses.create = AsyncMock( + return_value=mock_response, + ) + + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object"}, + }, + }, + ] + result = await model( + [{"role": "user", "content": "Weather?"}], + tools=tools, + tool_choice="auto", + ) + + call_args = mock_client.responses.create.call_args[1] + self.assertEqual(call_args["tool_choice"], "auto") + expected = [ + ToolUseBlock( + type="tool_use", + id="call_abc", + name="get_weather", + input={"city": "Beijing"}, + ), + ] + self.assertEqual(result.content, expected) + + async def test_streaming_text_response(self) -> None: + """Test streaming with text deltas.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=True, + ) + model.client = mock_client + + stream = self._create_stream_mock( + [ + {"type": "response.output_text.delta", "delta": "Hello"}, + {"type": "response.output_text.delta", "delta": " world"}, + { + "type": "response.completed", + "input_tokens": 10, + "output_tokens": 5, + }, + ], + ) + mock_client.responses.create = AsyncMock(return_value=stream) + + result = await model( + [{"role": "user", "content": "Hi"}], + ) + + responses = [] + async for resp in result: + responses.append(resp) + + final = responses[-1] + expected = [TextBlock(type="text", text="Hello world")] + self.assertEqual(final.content, expected) + self.assertIsNotNone(final.usage) + + async def test_streaming_with_reasoning(self) -> None: + """Test streaming with reasoning summary deltas.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=True, + reasoning_effort="high", + ) + model.client = mock_client + + stream = self._create_stream_mock( + [ + { + "type": "response.reasoning_summary_text.delta", + "delta": "Thinking...", + }, + { + "type": "response.output_text.delta", + "delta": "Answer", + }, + { + "type": "response.completed", + "input_tokens": 10, + "output_tokens": 20, + }, + ], + ) + mock_client.responses.create = AsyncMock(return_value=stream) + + result = await model( + [{"role": "user", "content": "Think"}], + ) + + responses = [] + async for resp in result: + responses.append(resp) + + final = responses[-1] + self.assertEqual(final.content[0]["type"], "thinking") + self.assertEqual(final.content[0]["thinking"], "Thinking...") + self.assertEqual(final.content[1]["type"], "text") + self.assertEqual(final.content[1]["text"], "Answer") + + async def test_streaming_function_call(self) -> None: + """Test streaming with function call events.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=True, + ) + model.client = mock_client + + stream = self._create_stream_mock( + [ + { + "type": "response.output_item.added", + "item_id": "item_1", + "item_type": "function_call", + "call_id": "call_xyz", + "name": "get_weather", + }, + { + "type": "response.function_call_arguments.delta", + "item_id": "item_1", + "delta": '{"city": "Beijing"}', + }, + { + "type": "response.completed", + "input_tokens": 10, + "output_tokens": 15, + }, + ], + ) + mock_client.responses.create = AsyncMock(return_value=stream) + + result = await model( + [{"role": "user", "content": "Weather?"}], + ) + + responses = [] + async for resp in result: + responses.append(resp) + + final = responses[-1] + tool_blocks = [b for b in final.content if b["type"] == "tool_use"] + self.assertEqual(len(tool_blocks), 1) + self.assertEqual(tool_blocks[0]["name"], "get_weather") + self.assertEqual(tool_blocks[0]["input"], {"city": "Beijing"}) + + async def test_non_streaming_structured_model(self) -> None: + """Test non-streaming with structured_model.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + mock_response = self._create_mock_response( + text='{"name": "Alice", "age": 30}', + ) + mock_client.responses.create = AsyncMock( + return_value=mock_response, + ) + + result = await model( + [{"role": "user", "content": "Generate a person"}], + structured_model=SampleModel, + ) + + call_args = mock_client.responses.create.call_args[1] + self.assertIn("text", call_args) + self.assertEqual( + call_args["text"]["format"]["type"], + "json_schema", + ) + self.assertEqual( + call_args["text"]["format"]["name"], + "SampleModel", + ) + self.assertNotIn("tools", call_args) + self.assertNotIn("tool_choice", call_args) + self.assertIsInstance(result, ChatResponse) + self.assertEqual( + result.metadata, + {"name": "Alice", "age": 30}, + ) + + async def test_streaming_structured_model(self) -> None: + """Test streaming with structured_model.""" + with patch("openai.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_client_cls.return_value = mock_client + + model = OpenAIResponseModel( + model_name="o3", + api_key="test_key", + stream=True, + ) + model.client = mock_client + + stream = self._create_stream_mock( + [ + { + "type": "response.output_text.delta", + "delta": '{"name": "Bob",', + }, + { + "type": "response.output_text.delta", + "delta": ' "age": 25}', + }, + { + "type": "response.completed", + "input_tokens": 5, + "output_tokens": 10, + }, + ], + ) + mock_client.responses.create = AsyncMock(return_value=stream) + + result = await model( + [{"role": "user", "content": "Generate a person"}], + structured_model=SampleModel, + ) + + responses = [] + async for resp in result: + responses.append(resp) + + final = responses[-1] + self.assertEqual( + final.metadata, + {"name": "Bob", "age": 25}, + ) + + # ------------------------------------------------------------------ + # Mock helpers + # ------------------------------------------------------------------ + + @staticmethod + def _create_mock_response( + text: str = "", + reasoning_summary: str = "", + function_calls: list | None = None, + input_tokens: int = 10, + output_tokens: int = 20, + ) -> Mock: + """Create a mock non-streaming Response object.""" + output_items = [] + + if reasoning_summary: + reasoning_item = Mock() + reasoning_item.type = "reasoning" + summary_obj = Mock() + summary_obj.text = reasoning_summary + reasoning_item.summary = [summary_obj] + output_items.append(reasoning_item) + + if text: + msg_item = Mock() + msg_item.type = "message" + text_part = Mock() + text_part.type = "output_text" + text_part.text = text + msg_item.content = [text_part] + output_items.append(msg_item) + + for fc in function_calls or []: + fc_item = Mock() + fc_item.type = "function_call" + fc_item.call_id = fc["call_id"] + fc_item.id = f"fc_{fc['call_id']}" + fc_item.name = fc["name"] + fc_item.arguments = fc["arguments"] + output_items.append(fc_item) + + response = Mock() + response.output = output_items + response.id = "resp_test" + + usage = Mock() + usage.input_tokens = input_tokens + usage.output_tokens = output_tokens + response.usage = usage + + return response + + @staticmethod + def _create_stream_mock(events_data: list) -> Any: + """Create a mock async event stream for the Response API.""" + + class MockResponseStream: + """Mock stream that yields Response API events.""" + + def __init__(self, events_data: list) -> None: + self.events_data = events_data + self.index = 0 + + def __aiter__(self) -> "MockResponseStream": + return self + + async def __anext__(self) -> Any: + if self.index >= len(self.events_data): + raise StopAsyncIteration + data = self.events_data[self.index] + self.index += 1 + + event = Mock() + event.type = data["type"] + + if data["type"] == "response.output_text.delta": + event.delta = data["delta"] + event.response = None + + elif data["type"] == ("response.reasoning_summary_text.delta"): + event.delta = data["delta"] + event.response = None + + elif data["type"] == "response.output_item.added": + item = Mock() + item.type = data.get("item_type", "message") + item.id = data.get("item_id", "") + item.call_id = data.get("call_id") + item.name = data.get("name", "") + event.item = item + event.response = None + + elif data["type"] == ( + "response.function_call_arguments.delta" + ): + event.item_id = data["item_id"] + event.delta = data["delta"] + event.response = None + + elif data["type"] == "response.completed": + resp = Mock() + resp.id = "resp_completed" + usage = Mock() + usage.input_tokens = data.get("input_tokens", 0) + usage.output_tokens = data.get("output_tokens", 0) + resp.usage = usage + event.response = resp + + else: + event.response = None + + return event + + return MockResponseStream(events_data) diff --git a/tests/model_openai_test.py b/tests/model_openai_test.py index bd5e18b996..b6994a4e70 100644 --- a/tests/model_openai_test.py +++ b/tests/model_openai_test.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# pylint: disable=protected-access """Unit tests for OpenAI API model class.""" from typing import AsyncGenerator, Any from unittest.async_case import IsolatedAsyncioTestCase @@ -365,8 +366,10 @@ def _create_mock_response_with_reasoning( def _create_mock_response_with_structured_data(self, data: dict) -> Mock: """Create a mock response with structured data.""" message = Mock() - message.parsed = Mock() - message.parsed.model_dump.return_value = data + # `message.parsed` must be a real instance of the schema so that the + # provider-agnostic ``isinstance(parsed, structured_model)`` guard + # in ``__call__`` accepts the response. + message.parsed = SampleModel(**data) message.content = None message.reasoning_content = None message.tool_calls = [] @@ -380,6 +383,441 @@ def _create_mock_response_with_structured_data(self, data: dict) -> Mock: return response + async def test_structured_sync_fallback_on_validation_error( + self, + ) -> None: + """``.parse()`` raising ``ValidationError`` (HTTP 200 + body that + does not match the schema) must trigger the tool-call fallback. + + Regression for agentscope-ai/agentscope#1631: DashScope + OpenAI-compat returns arbitrary JSON when messages contain the + word ``json``; the OpenAI SDK then raises + ``pydantic.ValidationError`` which the previous code did not + catch, crashing ``ReActAgent.CompressionConfig``. + """ + from pydantic import ValidationError + + with patch("openai.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + model = OpenAIChatModel( + model_name="qwen3.6-flash", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + # Construct a real ValidationError to mirror SDK behaviour. + try: + SampleModel.model_validate({}) + except ValidationError as exc: + ve = exc + mock_client.chat.completions.parse = AsyncMock(side_effect=ve) + + fallback_response = self._create_mock_response_with_tools( + "", + [ + { + "id": "call_1", + "name": "SampleModel", + "arguments": '{"name": "Alice", "age": 7}', + }, + ], + ) + mock_client.chat.completions.create = AsyncMock( + return_value=fallback_response, + ) + + result = await model( + [{"role": "user", "content": "x"}], + structured_model=SampleModel, + ) + + self.assertTrue(model._structured_output_fallback) + self.assertTrue(mock_client.chat.completions.create.called) + self.assertIsInstance(result, ChatResponse) + self.assertEqual( + result.metadata, + {"name": "Alice", "age": 7}, + ) + + async def test_structured_sync_fallback_on_silent_mismatch( + self, + ) -> None: + """``.parse()`` returning ``message.parsed=None`` (silent + non-conforming body) must also trigger the tool-call fallback, + not silently propagate ``metadata=None`` to the caller.""" + with patch("openai.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + model = OpenAIChatModel( + model_name="qwen3.6-flash", + api_key="test_key", + stream=False, + ) + model.client = mock_client + + # `.parse()` returns successfully but parsed is None + bad_response = self._create_mock_response("") + bad_response.choices[0].message.parsed = None + mock_client.chat.completions.parse = AsyncMock( + return_value=bad_response, + ) + + fallback_response = self._create_mock_response_with_tools( + "", + [ + { + "id": "call_1", + "name": "SampleModel", + "arguments": '{"name": "Bob", "age": 9}', + }, + ], + ) + mock_client.chat.completions.create = AsyncMock( + return_value=fallback_response, + ) + + result = await model( + [{"role": "user", "content": "x"}], + structured_model=SampleModel, + ) + + self.assertTrue(model._structured_output_fallback) + self.assertTrue(mock_client.chat.completions.create.called) + self.assertEqual( + result.metadata, + {"name": "Bob", "age": 9}, + ) + + async def test_structured_stream_fallback_on_validation_error( + self, + ) -> None: + """Streaming path: ``ValidationError`` raised during stream + consumption must trigger transparent fallback to tool-call. + """ + from pydantic import ValidationError + + with patch("openai.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + model = OpenAIChatModel( + model_name="qwen3.6-flash", + api_key="test_key", + stream=True, + ) + model.client = mock_client + + try: + SampleModel.model_validate({}) + except ValidationError as exc: + ve = exc + + class _RaisingStream: + """Mock that raises ValidationError mid-stream consumption, + mirroring openai-python's stream parser behaviour when the + accumulated body fails schema validation.""" + + async def __aenter__(self) -> "_RaisingStream": + return self + + async def __aexit__( + self, + exc_type: Any, + exc_val: Any, + exc_tb: Any, + ) -> None: + pass + + def __aiter__(self) -> "_RaisingStream": + return self + + async def __anext__(self) -> Any: + raise ve + + mock_client.chat.completions.stream = Mock( + return_value=_RaisingStream(), + ) + + fallback_stream = self._create_stream_mock( + [ + { + "tool_calls": [ + { + "id": "call_1", + "name": "SampleModel", + "arguments": '{"name":"C","age":3}', + }, + ], + }, + ], + ) + mock_client.chat.completions.create = AsyncMock( + return_value=fallback_stream, + ) + + result = await model( + [{"role": "user", "content": "x"}], + structured_model=SampleModel, + ) + + chunks = [] + async for chunk in result: + chunks.append(chunk) + + self.assertTrue(model._structured_output_fallback) + self.assertTrue(mock_client.chat.completions.create.called) + self.assertTrue(len(chunks) >= 1) + # In streaming fallback the structured payload is carried on the + # tool_use block (the tool name is the schema's tool name). + tool_blocks = [ + b for b in chunks[-1].content if b.get("type") == "tool_use" + ] + self.assertEqual(len(tool_blocks), 1) + self.assertEqual( + tool_blocks[0]["input"], + {"name": "C", "age": 3}, + ) + + async def test_structured_via_tool_call_retries_on_tool_choice_400( + self, + ) -> None: + """``_structured_via_tool_call`` must retry with + ``tool_choice='auto'`` when the endpoint rejects forced + ``tool_choice`` (e.g. DashScope thinking-mode models).""" + import openai + + with patch("openai.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + model = OpenAIChatModel( + model_name="qwen3.6-flash", + api_key="test_key", + stream=False, + ) + model.client = mock_client + # Skip the response_format attempt; go straight to fallback. + model._structured_output_fallback = True + + http_response = Mock() + http_response.status_code = 400 + http_response.request = Mock() + http_response.headers = {} + bad_request = openai.BadRequestError( + message=( + "tool_choice required is not supported in thinking " "mode" + ), + response=http_response, + body=None, + ) + + success_response = self._create_mock_response_with_tools( + "", + [ + { + "id": "call_1", + "name": "SampleModel", + "arguments": '{"name": "Dan", "age": 4}', + }, + ], + ) + mock_client.chat.completions.create = AsyncMock( + side_effect=[bad_request, success_response], + ) + + result = await model( + [{"role": "user", "content": "x"}], + structured_model=SampleModel, + ) + + self.assertEqual( + mock_client.chat.completions.create.call_count, + 2, + ) + # Second call must downgrade to tool_choice='auto' + second_call_kwargs = ( + mock_client.chat.completions.create.call_args_list[1][1] + ) + self.assertEqual(second_call_kwargs["tool_choice"], "auto") + self.assertEqual( + result.metadata, + {"name": "Dan", "age": 4}, + ) + + async def test_structured_via_tool_call_stream_retries_on_lazy_400( + self, + ) -> None: + """Streaming tool-call fallback: when the endpoint surfaces a + ``tool_choice``-related error lazily as ``openai.APIError`` during + stream iteration (DashScope thinking mode), the wrapper must retry + once with ``tool_choice='auto'`` before any chunk is yielded. + """ + import openai + + with patch("openai.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + model = OpenAIChatModel( + model_name="qwen3.6-flash", + api_key="test_key", + stream=True, + ) + model.client = mock_client + # Skip response_format; go straight to tool-call fallback. + model._structured_output_fallback = True + + http_request = Mock() + + class _LazyApiErrorStream: + """Mirrors openai SDK behaviour where SSE error events + surface as ``APIError`` during stream iteration.""" + + async def __aenter__(self) -> "_LazyApiErrorStream": + return self + + async def __aexit__( + self, + exc_type: Any, + exc_val: Any, + exc_tb: Any, + ) -> None: + pass + + def __aiter__(self) -> "_LazyApiErrorStream": + return self + + async def __anext__(self) -> Any: + raise openai.APIError( + message=( + "The tool_choice parameter does not support " + "being set to required or object in thinking " + "mode" + ), + request=http_request, + body=None, + ) + + success_stream = self._create_stream_mock( + [ + { + "tool_calls": [ + { + "id": "call_1", + "name": "SampleModel", + "arguments": '{"name":"Eve","age":5}', + }, + ], + }, + ], + ) + mock_client.chat.completions.create = AsyncMock( + side_effect=[_LazyApiErrorStream(), success_stream], + ) + + result = await model( + [{"role": "user", "content": "x"}], + structured_model=SampleModel, + ) + + chunks = [] + async for chunk in result: + chunks.append(chunk) + + self.assertEqual( + mock_client.chat.completions.create.call_count, + 2, + ) + second_call_kwargs = ( + mock_client.chat.completions.create.call_args_list[1][1] + ) + self.assertEqual(second_call_kwargs["tool_choice"], "auto") + tool_blocks = [ + b for b in chunks[-1].content if b.get("type") == "tool_use" + ] + self.assertEqual(len(tool_blocks), 1) + self.assertEqual( + tool_blocks[0]["input"], + {"name": "Eve", "age": 5}, + ) + + async def test_structured_via_tool_call_stream_retries_on_sync_400( + self, + ) -> None: + """Streaming tool-call fallback: when the endpoint rejects + forced ``tool_choice`` *synchronously* on + ``client.chat.completions.create()`` (e.g. DashScope qwen3.6-plus + / qwen3.6-flash), the wrapper must retry once with + ``tool_choice='auto'`` before consuming the stream. + """ + import openai + + with patch("openai.AsyncClient") as mock_client_class: + mock_client = AsyncMock() + mock_client_class.return_value = mock_client + model = OpenAIChatModel( + model_name="qwen3.6-plus", + api_key="test_key", + stream=True, + ) + model.client = mock_client + model._structured_output_fallback = True + + http_response = Mock() + http_response.status_code = 400 + http_response.request = Mock() + http_response.headers = {} + sync_400 = openai.BadRequestError( + message=( + "The tool_choice parameter does not support being " + "set to required or object in thinking mode" + ), + response=http_response, + body=None, + ) + + success_stream = self._create_stream_mock( + [ + { + "tool_calls": [ + { + "id": "call_1", + "name": "SampleModel", + "arguments": '{"name":"Fay","age":6}', + }, + ], + }, + ], + ) + mock_client.chat.completions.create = AsyncMock( + side_effect=[sync_400, success_stream], + ) + + result = await model( + [{"role": "user", "content": "x"}], + structured_model=SampleModel, + ) + + chunks = [] + async for chunk in result: + chunks.append(chunk) + + self.assertEqual( + mock_client.chat.completions.create.call_count, + 2, + ) + second_call_kwargs = ( + mock_client.chat.completions.create.call_args_list[1][1] + ) + self.assertEqual(second_call_kwargs["tool_choice"], "auto") + tool_blocks = [ + b for b in chunks[-1].content if b.get("type") == "tool_use" + ] + self.assertEqual(len(tool_blocks), 1) + self.assertEqual( + tool_blocks[0]["input"], + {"name": "Fay", "age": 6}, + ) + async def test_streaming_response_with_none_delta(self) -> None: """Test streaming response when a chunk has delta = None.""" with patch("openai.AsyncClient") as mock_client_class: diff --git a/tests/tracing_cancelled_error_test.py b/tests/tracing_cancelled_error_test.py new file mode 100644 index 0000000000..8a701292ec --- /dev/null +++ b/tests/tracing_cancelled_error_test.py @@ -0,0 +1,78 @@ +# -*- coding: utf-8 -*- +"""The unittests for tracing handling CancelledError.""" +import asyncio +from unittest import IsolatedAsyncioTestCase +from unittest.mock import patch + +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, +) +from opentelemetry.trace import StatusCode + +from agentscope import _config +from agentscope.tracing import trace + + +class TracingCancelledErrorTest(IsolatedAsyncioTestCase): + """Test the tracing for handling CancelledError in async functions.""" + + async def asyncSetUp(self) -> None: + """Set up the test case""" + + self._original_trace_enabled = _config.trace_enabled + _config.trace_enabled = True + self.exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(self.exporter)) + self.tracer = provider.get_tracer( + "tests.tracing_cancelled_error_test", + ) + self.tracer_patcher = patch( + "agentscope.tracing._trace._get_tracer", + return_value=self.tracer, + ) + self.tracer_patcher.start() + self.addCleanup(self.tracer_patcher.stop) + + async def test_trace_ends_span_for_normal_exception(self) -> None: + """Test that normal exceptions end the span with error status.""" + + @trace(name="normal_exception_case") + async def raise_value_error() -> None: + raise ValueError("normal exception") + + with self.assertRaises(ValueError): + await raise_value_error() + + finished_spans = self.exporter.get_finished_spans() + + self.assertEqual(len(finished_spans), 1) + self.assertEqual( + finished_spans[0].status.status_code, + StatusCode.ERROR, + ) + + async def test_trace_should_end_span_for_cancelled_error(self) -> None: + """Test that CancelledError ends the span with error status.""" + + @trace(name="cancelled_error_case") + async def raise_cancelled_error() -> None: + raise asyncio.CancelledError("set cancelled error") + + with self.assertRaises(asyncio.CancelledError): + await raise_cancelled_error() + + finished_spans = self.exporter.get_finished_spans() + + self.assertEqual(len(finished_spans), 1) + self.assertEqual( + finished_spans[0].status.status_code, + StatusCode.ERROR, + ) + + async def asyncTearDown(self) -> None: + """Restore tracing configuration after each test.""" + + _config.trace_enabled = self._original_trace_enabled diff --git a/tests/tts_dashscope_test.py b/tests/tts_dashscope_test.py index 8a37f58bd7..dc1f15a6c8 100644 --- a/tests/tts_dashscope_test.py +++ b/tests/tts_dashscope_test.py @@ -132,6 +132,13 @@ async def test_synthesize_non_streaming(self) -> None: api_key=self.api_key, stream=False, ) as model: + # Mock finish_event to not block + model._dashscope_callback.finish_event = Mock() + model._dashscope_callback.finish_event.wait = Mock() + # Mock has_audio_data to return True (skip retry) + model._dashscope_callback.has_audio_data = Mock( + return_value=True, + ) model._dashscope_callback.get_audio_data = AsyncMock( return_value=TTSResponse( content=AudioBlock( @@ -169,6 +176,13 @@ async def test_synthesize_streaming(self) -> None: api_key=self.api_key, stream=True, ) as model: + # Mock finish_event to not block + model._dashscope_callback.finish_event = Mock() + model._dashscope_callback.finish_event.wait = Mock() + # Mock has_audio_data to return True (skip retry) + model._dashscope_callback.has_audio_data = Mock( + return_value=True, + ) async def mock_generator() -> AsyncGenerator[ TTSResponse,