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,