diff --git a/src/agentpool/models/__init__.py b/src/agentpool/models/__init__.py index 60ea571cf..093e79024 100644 --- a/src/agentpool/models/__init__.py +++ b/src/agentpool/models/__init__.py @@ -5,6 +5,10 @@ from agentpool.models.acp_agents import ACPAgentConfig, ACPAgentConfigTypes, BaseACPAgentConfig from agentpool.models.agents import AnyToolConfig, NativeAgentConfig # noqa: F401 from agentpool.models.manifest import AgentsManifest, AnyAgentConfig +from agentpool.models.openai_compatible import ( + OpenAICompatibleModel, + OpenAICompatibleModelProfile, +) from agentpool.models.pending_interaction import PendingPermission, PendingQuestion @@ -15,6 +19,8 @@ "AnyAgentConfig", "BaseACPAgentConfig", "NativeAgentConfig", + "OpenAICompatibleModel", + "OpenAICompatibleModelProfile", "PendingPermission", "PendingQuestion", ] diff --git a/src/agentpool/models/openai_compatible.py b/src/agentpool/models/openai_compatible.py new file mode 100644 index 000000000..9e09bf1f2 --- /dev/null +++ b/src/agentpool/models/openai_compatible.py @@ -0,0 +1,323 @@ +"""OpenAI-compatible model with native list tool return support. + +Subclass of pydantic-ai's ``OpenAIChatModel`` that optionally emits native list +content for tool return messages instead of JSON-serialized strings. This is +useful for OpenAI-compatible models (GLM-5, vLLM, etc.) whose chat templates +natively render list-type tool message content. + +Example YAML manifest usage: + +```yaml +model_variants: + glm-5: + type: import + model: agentpool.models.openai_compatible.OpenAICompatibleModel + kw_args: + model_name: "glm-5" + base_url: "https://open.bigmodel.cn/api/paas/v4/" + api_key: "${OPENAI_API_KEY}" + tool_return_as_list: "true" + openai_system_prompt_role: "developer" + openai_supports_strict_tool_definition: "false" +``` +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, TypedDict, cast, override + +from pydantic_ai.messages import ( + ModelRequest, + RetryPromptPart, + SystemPromptPart, + ToolReturnPart, + UserPromptPart, +) +from pydantic_ai.models.openai import ( # type: ignore[attr-defined] + OpenAIChatModel, + _guard_tool_call_id, +) +from pydantic_ai.profiles import ModelProfile, ModelProfileSpec +from pydantic_ai.profiles.openai import OpenAIModelProfile +from pydantic_ai.providers.openai import OpenAIProvider + + +try: + from openai.types import chat + from openai.types.chat import ChatCompletionContentPartTextParam +except ImportError: # pragma: no cover + pass + +if TYPE_CHECKING: + from collections.abc import AsyncIterator + + from openai import AsyncOpenAI + from pydantic_ai.providers import Provider + from pydantic_ai.settings import ModelSettings + + +class OpenAICompatibleModelProfile(TypedDict, total=False): + """Profile dict for :class:`OpenAICompatibleModel`. + + Extends ``OpenAIModelProfile`` fields with the additional + ``openai_chat_tool_return_as_list`` key. + """ + + openai_chat_tool_return_as_list: bool + + +_OPENAI_BOOL_PROFILE_KEYS: frozenset[str] = frozenset({ + "openai_chat_tool_return_as_list", + "openai_supports_strict_tool_definition", + "openai_supports_sampling_settings", + "openai_supports_tool_choice_required", + "openai_chat_supports_multiple_system_messages", + "openai_chat_supports_web_search", + "openai_chat_supports_file_urls", + "openai_supports_encrypted_reasoning_content", + "openai_supports_reasoning", + "openai_supports_reasoning_effort_none", + "openai_responses_requires_function_call_status_none", + "openai_supports_phase", + "supports_inline_system_prompts", +}) +"""Profile keys whose values should be coerced from string to bool.""" + +_OPENAI_PROFILE_FIELDS: frozenset[str] = frozenset({ + "openai_chat_thinking_field", + "openai_chat_send_back_thinking_parts", + "openai_supports_strict_tool_definition", + "openai_supports_sampling_settings", + "openai_unsupported_model_settings", + "openai_supports_tool_choice_required", + "openai_system_prompt_role", + "supports_inline_system_prompts", + "openai_chat_supports_multiple_system_messages", + "openai_chat_supports_web_search", + "openai_chat_audio_input_encoding", + "openai_chat_supports_file_urls", + "openai_supports_encrypted_reasoning_content", + "openai_supports_reasoning", + "openai_supports_reasoning_effort_none", + "openai_responses_requires_function_call_status_none", + "openai_supports_phase", +}) + + +def _coerce_profile_value(key: str, value: Any) -> Any: + """Coerce string values to bool for known boolean profile keys. + + Args: + key: The profile key name. + value: The raw value (typically a string from YAML kw_args). + + Returns: + The coerced value — ``bool`` for known boolean keys, original value otherwise. + """ + if key in _OPENAI_BOOL_PROFILE_KEYS and isinstance(value, str): + return value.lower() in ("true", "1", "yes") + return value + + +def _profile_to_dict(profile: ModelProfile) -> dict[str, Any]: + """Extract non-default OpenAI profile fields into a dict.""" + result: dict[str, Any] = {} + for field_name in _OPENAI_PROFILE_FIELDS: + val = getattr(profile, field_name, None) + if val is not None: + result[field_name] = val + return result + + +class OpenAICompatibleModel(OpenAIChatModel): + """An ``OpenAIChatModel`` subclass that can emit native list tool return content. + + When the ``openai_chat_tool_return_as_list`` profile flag is ``True`` and a + ``ToolReturnPart`` has non-empty list content with no multimodal files, the + tool message ``content`` is emitted as + ``list[ChatCompletionContentPartTextParam]`` instead of a JSON-serialized + string. This matches the expectation of chat templates that branch on + ``m.content is string`` vs list (e.g. GLM-5). + + All other behavior is inherited from :class:`OpenAIChatModel`. + """ + + def __init__( + self, + model_name: str, + *, + base_url: str | None = None, + api_key: str | None = None, + provider: Provider[AsyncOpenAI] | None = None, + tool_return_as_list: str | bool = False, + profile: ModelProfileSpec | None = None, + settings: ModelSettings | None = None, + **profile_overrides: str, + ) -> None: + """Initialize an OpenAI-compatible model. + + Args: + model_name: The name of the model to use. + base_url: Base URL for the OpenAI-compatible API. Ignored if + ``provider`` is given. + api_key: API key for authentication. Ignored if ``provider`` is + given. + provider: A pre-built ``Provider[AsyncOpenAI]``. When ``None``, an + ``OpenAIProvider`` is constructed from ``base_url`` and + ``api_key``. + tool_return_as_list: Whether to emit native list tool return + content. Accepts ``str`` (``"true"``/``"false"``) or ``bool`` + for YAML convenience. + profile: The model profile spec to use. + settings: Default model settings for this model instance. + **profile_overrides: Arbitrary ``openai_*`` prefixed keys merged + into the profile. Known boolean keys are coerced from + string to ``bool``. + """ + # Validate openai_* kwargs before constructing OpenAIProvider, + # so invalid kwargs raise TypeError before any network/auth setup. + real_overrides: dict[str, Any] = {} + for key, value in profile_overrides.items(): + if key.startswith("openai_"): + coerced = _coerce_profile_value(key, value) + if key in _OPENAI_PROFILE_FIELDS: + real_overrides[key] = coerced + else: + msg = ( + f"Unexpected keyword argument '{key}'. " + f"Only 'openai_*' prefixed keys are accepted as profile overrides." + ) + raise TypeError(msg) + + if provider is None: + provider = OpenAIProvider(base_url=base_url, api_key=api_key) + + # Coerce tool_return_as_list to bool and store on instance + self._tool_return_as_list_enabled: bool = ( + tool_return_as_list + if isinstance(tool_return_as_list, bool) + else isinstance(tool_return_as_list, str) + and tool_return_as_list.lower() in ("true", "1", "yes") + ) + + # Build the merged profile spec + merged_profile = self._build_merged_profile(profile, real_overrides) + + super().__init__( + model_name=model_name, + provider=provider, + profile=merged_profile, + settings=settings, + ) + + @staticmethod + def _build_merged_profile( + profile: ModelProfileSpec | None, + real_overrides: dict[str, Any], + ) -> ModelProfileSpec | None: + """Merge openai_* overrides into the profile spec.""" + if not real_overrides: + return profile + + if profile is None: + return OpenAIModelProfile(**real_overrides) + + if isinstance(profile, ModelProfile): + base_dict = _profile_to_dict(profile) + base_dict.update(real_overrides) + return OpenAIModelProfile(**base_dict) + + # Callable profile: wrap to inject overrides post-call + original_fn = profile + + def _wrapped(model_name: str) -> ModelProfile | None: + result = original_fn(model_name) + if result is None: + return None + result_dict = _profile_to_dict(result) if isinstance(result, ModelProfile) else {} + result_dict.update(real_overrides) + return OpenAIModelProfile(**result_dict) + + return _wrapped + + @property + def _resolved_profile(self) -> OpenAICompatibleModelProfile: + """Return the resolved profile as a typed dict for flag access.""" + result: dict[str, Any] = {} + profile = OpenAIModelProfile.from_profile(self.profile) + for field_name in ( + "openai_supports_strict_tool_definition", + "openai_system_prompt_role", + ): + val = getattr(profile, field_name, None) + if val is not None: + result[field_name] = val + result["openai_chat_tool_return_as_list"] = self._tool_return_as_list_enabled + return cast(OpenAICompatibleModelProfile, result) + + @override + async def _map_user_message( + self, message: ModelRequest + ) -> AsyncIterator[chat.ChatCompletionMessageParam]: + if not self._tool_return_as_list_enabled: + # Flag disabled: delegate entirely to parent + async for item in super()._map_user_message(message): + yield item + return + + # Flag enabled: duplicate parent logic, replacing ToolReturnPart branch + file_content: list[Any] = [] + for part in message.parts: + if isinstance(part, SystemPromptPart): + system_prompt_role = OpenAIModelProfile.from_profile( + self.profile + ).openai_system_prompt_role + if system_prompt_role == "developer": + yield chat.ChatCompletionDeveloperMessageParam( + role="developer", content=part.content + ) + elif system_prompt_role == "user": + yield chat.ChatCompletionUserMessageParam(role="user", content=part.content) + else: + yield chat.ChatCompletionSystemMessageParam(role="system", content=part.content) + elif isinstance(part, UserPromptPart): + yield await self._map_user_prompt(part) + elif isinstance(part, ToolReturnPart): + if isinstance(part.content, list) and part.content and not part.files: + # Native list content: emit list[ChatCompletionContentPartTextParam] + content_parts: list[ChatCompletionContentPartTextParam] = [ + ChatCompletionContentPartTextParam(type="text", text=item) + for item in part.content_items(mode="str") + if isinstance(item, str) + ] + yield chat.ChatCompletionToolMessageParam( + role="tool", + tool_call_id=_guard_tool_call_id(t=part), + content=content_parts, + ) + else: + # String content, empty list, or files: use parent behavior + tool_text, tool_file_content = part.model_response_str_and_user_content() + file_content.extend(tool_file_content) + yield chat.ChatCompletionToolMessageParam( + role="tool", + tool_call_id=_guard_tool_call_id(t=part), + content=tool_text, + ) + elif isinstance(part, RetryPromptPart): + if part.tool_name is None: + yield chat.ChatCompletionUserMessageParam( + role="user", content=part.model_response() + ) + else: + yield chat.ChatCompletionToolMessageParam( + role="tool", + tool_call_id=_guard_tool_call_id(t=part), + content=part.model_response(), + ) + else: + from typing import assert_never + + assert_never(part) + if file_content: + yield await self._map_user_prompt(UserPromptPart(content=file_content)) diff --git a/tests/models/__init__.py b/tests/models/__init__.py new file mode 100644 index 000000000..744975a84 --- /dev/null +++ b/tests/models/__init__.py @@ -0,0 +1 @@ +"""Tests for model classes.""" diff --git a/tests/models/test_openai_compatible.py b/tests/models/test_openai_compatible.py new file mode 100644 index 000000000..aeccfc1d9 --- /dev/null +++ b/tests/models/test_openai_compatible.py @@ -0,0 +1,490 @@ +"""Tests for OpenAICompatibleModel.""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import MagicMock, patch + +from pydantic_ai.messages import ( + ModelRequest, + RetryPromptPart, + SystemPromptPart, + ToolReturnPart, + UserPromptPart, +) +from pydantic_ai.profiles.openai import OpenAIModelProfile +import pytest + +from agentpool.models.openai_compatible import ( + OpenAICompatibleModel, + _coerce_profile_value, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _mock_provider() -> MagicMock: + """Create a mock provider with a real OpenAIModelProfile.""" + provider = MagicMock() + provider.model_profile = OpenAIModelProfile() + return provider + + +async def _collect(async_iter: Any) -> list[Any]: + """Collect all items from an async iterator into a list.""" + return [item async for item in async_iter] + + +# --------------------------------------------------------------------------- +# 1. Subclass relationship +# --------------------------------------------------------------------------- + + +def test_is_subclass_of_openai_chat_model() -> None: + """OpenAICompatibleModel should be a subclass of OpenAIChatModel.""" + from pydantic_ai.models.openai import OpenAIChatModel + + assert issubclass(OpenAICompatibleModel, OpenAIChatModel) + + +# --------------------------------------------------------------------------- +# 2. Constructor and profile handling +# --------------------------------------------------------------------------- + + +def test_constructor_accepts_base_url_and_api_key() -> None: + """Constructor should accept base_url, api_key, and tool_return_as_list.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_provider_cls: + mock_provider_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test-model", + base_url="https://api.example.com/v1", + api_key="test-key", + tool_return_as_list="true", + ) + mock_provider_cls.assert_called_once_with( + base_url="https://api.example.com/v1", api_key="test-key" + ) + assert model.model_name == "test-model" + + +def test_tool_return_as_list_string_true() -> None: + """tool_return_as_list='true' should set the flag to True.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test", + tool_return_as_list="true", + ) + assert model._resolved_profile.get("openai_chat_tool_return_as_list") is True + + +def test_tool_return_as_list_string_false() -> None: + """tool_return_as_list='false' should set the flag to False.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test", + tool_return_as_list="false", + ) + assert model._resolved_profile.get("openai_chat_tool_return_as_list") is False + + +def test_tool_return_as_list_bool() -> None: + """tool_return_as_list=True (bool) should set the flag to True.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test", + tool_return_as_list=True, + ) + assert model._resolved_profile.get("openai_chat_tool_return_as_list") is True + + +def test_tool_return_as_list_default_false() -> None: + """Default (no flag) should be False.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel(model_name="test") + assert model._resolved_profile.get("openai_chat_tool_return_as_list") is False + + +def test_openai_profile_overrides_passed_through() -> None: + """openai_* prefixed kwargs should be merged into the profile.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test", + tool_return_as_list="true", + openai_system_prompt_role="developer", + openai_supports_strict_tool_definition="false", + ) + profile = OpenAIModelProfile.from_profile(model.profile) + assert profile.openai_system_prompt_role == "developer" + assert profile.openai_supports_strict_tool_definition is False + + +def test_non_openai_kwarg_raises_type_error() -> None: + """Non-openai_* unknown kwargs should raise TypeError.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + with pytest.raises(TypeError, match="Unexpected keyword argument"): + OpenAICompatibleModel( + model_name="test", + foo_bar="baz", # type: ignore[call-arg] + ) + + +def test_profile_dict_merged_with_overrides() -> None: + """When profile is a ModelProfile, overrides should be merged in.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test", + tool_return_as_list="true", + profile=OpenAIModelProfile(openai_system_prompt_role="developer"), + ) + profile = OpenAIModelProfile.from_profile(model.profile) + assert profile.openai_system_prompt_role == "developer" + assert model._resolved_profile.get("openai_chat_tool_return_as_list") is True + + +# --------------------------------------------------------------------------- +# 3. _coerce_profile_value helper +# --------------------------------------------------------------------------- + + +def test_coerce_bool_true() -> None: + """String 'true' should coerce to True for boolean keys.""" + assert _coerce_profile_value("openai_supports_strict_tool_definition", "true") is True + + +def test_coerce_bool_false() -> None: + """String 'false' should coerce to False for boolean keys.""" + assert _coerce_profile_value("openai_supports_strict_tool_definition", "false") is False + + +def test_coerce_non_bool_key_unchanged() -> None: + """Non-boolean keys should return the value unchanged.""" + assert _coerce_profile_value("openai_system_prompt_role", "developer") == "developer" + + +def test_coerce_non_string_value_unchanged() -> None: + """Non-string values should be returned unchanged even for boolean keys.""" + assert _coerce_profile_value("openai_supports_strict_tool_definition", True) is True + assert _coerce_profile_value("openai_supports_strict_tool_definition", 1) == 1 + + +# --------------------------------------------------------------------------- +# 4. _map_user_message behavior +# --------------------------------------------------------------------------- + + +@pytest.fixture +def model_flag_disabled() -> OpenAICompatibleModel: + """Model with tool_return_as_list disabled (default).""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + return OpenAICompatibleModel(model_name="test", tool_return_as_list="false") + + +@pytest.fixture +def model_flag_enabled() -> OpenAICompatibleModel: + """Model with tool_return_as_list enabled.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + return OpenAICompatibleModel(model_name="test", tool_return_as_list="true") + + +async def test_flag_disabled_delegates_to_super( + model_flag_disabled: OpenAICompatibleModel, +) -> None: + """When flag is False, _map_user_message should delegate to parent.""" + part = ToolReturnPart( + tool_name="test_tool", + content=["item1", "item2"], + tool_call_id="call_123", + ) + message = ModelRequest(parts=[part]) + + # Mock the parent's _map_user_message to verify delegation + with ( + patch.object( + OpenAIChatModel, + "_map_user_message", + return_value=AsyncIteratorMock([MagicMock()]), + ) as mock_super, + ): + await _collect(model_flag_disabled._map_user_message(message)) + mock_super.assert_called_once_with(message) + + +async def test_flag_enabled_list_string_items( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + list of strings -> content is list of text parts.""" + part = ToolReturnPart( + tool_name="test_tool", + content=["result1", "result2"], + tool_call_id="call_123", + ) + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + tool_msg = results[0] + assert tool_msg["role"] == "tool" + assert tool_msg["tool_call_id"] == "call_123" + content = tool_msg["content"] + assert isinstance(content, list) + assert len(content) == 2 + assert content[0]["type"] == "text" + assert content[0]["text"] == "result1" + assert content[1]["type"] == "text" + assert content[1]["text"] == "result2" + + +async def test_flag_enabled_list_non_string_items( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + list with non-string items -> each JSON-serialized and wrapped.""" + part = ToolReturnPart( + tool_name="test_tool", + content=[{"key": "value"}, 42], + tool_call_id="call_123", + ) + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + tool_msg = results[0] + content = tool_msg["content"] + assert isinstance(content, list) + assert len(content) == 2 + assert content[0]["type"] == "text" + # Non-string items are JSON-serialized via content_items(mode='str') + assert '"key"' in content[0]["text"] + assert "value" in content[0]["text"] + assert content[1]["type"] == "text" + assert content[1]["text"] == "42" + + +async def test_flag_enabled_string_content( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + string content -> content remains plain string.""" + part = ToolReturnPart( + tool_name="test_tool", + content="plain string", + tool_call_id="call_123", + ) + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + tool_msg = results[0] + content = tool_msg["content"] + assert isinstance(content, str) + assert content == "plain string" + + +async def test_flag_enabled_empty_list( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + empty list -> falls back to parent (empty string).""" + part = ToolReturnPart( + tool_name="test_tool", + content=[], + tool_call_id="call_123", + ) + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + tool_msg = results[0] + content = tool_msg["content"] + # Empty list falls back to parent behavior which serializes to '' + assert isinstance(content, str) + + +async def test_flag_enabled_system_prompt_part( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + SystemPromptPart -> mapped identically to parent.""" + part = SystemPromptPart(content="You are a helpful assistant.") + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + sys_msg = results[0] + assert sys_msg["role"] == "system" + assert sys_msg["content"] == "You are a helpful assistant." + + +async def test_flag_enabled_system_prompt_developer_role( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + SystemPromptPart with developer role -> developer message.""" + with patch("agentpool.models.openai_compatible.OpenAIProvider") as mock_cls: + mock_cls.return_value = _mock_provider() + model = OpenAICompatibleModel( + model_name="test", + tool_return_as_list="true", + openai_system_prompt_role="developer", + ) + part = SystemPromptPart(content="You are a developer.") + message = ModelRequest(parts=[part]) + + results = await _collect(model._map_user_message(message)) + + assert len(results) == 1 + assert results[0]["role"] == "developer" + + +async def test_flag_enabled_retry_prompt_with_tool_name( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + RetryPromptPart with tool_name -> tool message.""" + part = RetryPromptPart( + tool_name="test_tool", + content="Retry this", + tool_call_id="call_123", + ) + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + tool_msg = results[0] + assert tool_msg["role"] == "tool" + + +async def test_flag_enabled_retry_prompt_without_tool_name( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + RetryPromptPart without tool_name -> user message.""" + part = RetryPromptPart( + tool_name=None, + content="Retry this", + ) + message = ModelRequest(parts=[part]) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + assert len(results) == 1 + user_msg = results[0] + assert user_msg["role"] == "user" + + +async def test_flag_enabled_mixed_message( + model_flag_enabled: OpenAICompatibleModel, +) -> None: + """Flag True + mixed message parts -> only ToolReturnPart with list is modified.""" + from pydantic_ai.messages import ModelRequestPart + + parts: list[ModelRequestPart] = [ + UserPromptPart(content="Run the tool"), + ToolReturnPart( + tool_name="test_tool", + content=["result1", "result2"], + tool_call_id="call_1", + ), + ToolReturnPart( + tool_name="other_tool", + content="string result", + tool_call_id="call_2", + ), + ] + message = ModelRequest(parts=parts) + + results = await _collect(model_flag_enabled._map_user_message(message)) + + # Should have: user prompt, tool msg 1 (list), tool msg 2 (string) + assert len(results) == 3 + # First: user message + assert results[0]["role"] == "user" + # Second: tool message with list content + assert results[1]["role"] == "tool" + assert isinstance(results[1]["content"], list) + assert len(results[1]["content"]) == 2 + # Third: tool message with string content + assert results[2]["role"] == "tool" + assert isinstance(results[2]["content"], str) + + +# --------------------------------------------------------------------------- +# 5. Integration: ImportModelConfig resolution +# --------------------------------------------------------------------------- + + +def test_import_model_config_resolves_model() -> None: + """ImportModelConfig should resolve OpenAICompatibleModel from YAML.""" + from llmling_models_config import ImportModelConfig + + config = ImportModelConfig( + model="agentpool.models.openai_compatible.OpenAICompatibleModel", + kw_args={ + "model_name": "glm-5", + "base_url": "https://open.bigmodel.cn/api/paas/v4/", + "api_key": "test-key", + "tool_return_as_list": "true", + "openai_system_prompt_role": "developer", + "openai_supports_strict_tool_definition": "false", + }, + ) + model = config.get_model() + assert isinstance(model, OpenAICompatibleModel) + assert model._resolved_profile.get("openai_chat_tool_return_as_list") is True + profile = OpenAIModelProfile.from_profile(model.profile) + assert profile.openai_system_prompt_role == "developer" + assert profile.openai_supports_strict_tool_definition is False + + +def test_import_model_config_non_openai_kwarg_raises() -> None: + """ImportModelConfig with non-openai_* kwarg should raise TypeError.""" + from llmling_models_config import ImportModelConfig + + config = ImportModelConfig( + model="agentpool.models.openai_compatible.OpenAICompatibleModel", + kw_args={ + "model_name": "test", + "foo_bar": "baz", + }, + ) + with pytest.raises(TypeError, match="Unexpected keyword argument"): + config.get_model() + + +# --------------------------------------------------------------------------- +# Helper: AsyncIteratorMock +# --------------------------------------------------------------------------- + + +# Import at module level for patch target +from pydantic_ai.models.openai import OpenAIChatModel # noqa: E402 + + +class AsyncIteratorMock: + """Mock that acts as an async iterator yielding provided items.""" + + def __init__(self, items: list[Any]) -> None: + self._items = items + self._index = 0 + + def __aiter__(self) -> AsyncIteratorMock: + return self + + async def __anext__(self) -> Any: + if self._index >= len(self._items): + raise StopAsyncIteration + item = self._items[self._index] + self._index += 1 + return item