diff --git a/src/askui/model_providers/__init__.py b/src/askui/model_providers/__init__.py index 5d6a034e..4700d61c 100644 --- a/src/askui/model_providers/__init__.py +++ b/src/askui/model_providers/__init__.py @@ -35,6 +35,7 @@ from askui.model_providers.openai_image_qa_provider import OpenAIImageQAProvider from askui.model_providers.openai_vlm_provider import OpenAIVlmProvider from askui.model_providers.vlm_provider import VlmProvider +from askui.models.openai.messages_api import MessageTransform from askui.models.shared.coordinate_space import ( NormalizedCoordinateSpace, PixelCoordinateSpace, @@ -59,6 +60,7 @@ "ImageQAProvider", "ContainedImageScaler", "ImageScaler", + "MessageTransform", "ModelPricing", "PatchOptimizedImageScaler", "NormalizedCoordinateSpace", diff --git a/src/askui/model_providers/openai_compatible_vlm_provider.py b/src/askui/model_providers/openai_compatible_vlm_provider.py index 2f6fb9f9..2e73ab44 100644 --- a/src/askui/model_providers/openai_compatible_vlm_provider.py +++ b/src/askui/model_providers/openai_compatible_vlm_provider.py @@ -4,6 +4,7 @@ from openai import OpenAI from askui.model_providers.openai_vlm_provider import OpenAIVlmProvider +from askui.models.openai.messages_api import MessageTransform from askui.models.shared.coordinate_space import ( PixelCoordinateSpace, VlmCoordinateSpace, @@ -37,6 +38,10 @@ class OpenAICompatibleVlmProvider(OpenAIVlmProvider): is not provided. Reads ``ASKUI_VLM_MAX_IMAGE_EDGE`` from the environment if not provided. Inherits the default from `OpenAIVlmProvider` (1024). + message_transform (`MessageTransform` | None, optional): Hook to + post-process the OpenAI-format ``messages`` list right before it is + sent — for gateways that deviate from the stock chat spec. See + `OpenAIVlmProvider`. Example: ```python @@ -61,6 +66,7 @@ def __init__( coordinate_space: VlmCoordinateSpace = _DEFAULT_COORDINATE_SPACE, image_scaler: ImageScaler | None = None, image_edge_max: int | None = None, + message_transform: MessageTransform | None = None, ) -> None: def _rewrite_url(request: httpx.Request) -> None: request.url = httpx.URL(endpoint_url) @@ -79,4 +85,5 @@ def _rewrite_url(request: httpx.Request) -> None: coordinate_space=coordinate_space, image_scaler=image_scaler, image_edge_max=image_edge_max, + message_transform=message_transform, ) diff --git a/src/askui/model_providers/openai_vlm_provider.py b/src/askui/model_providers/openai_vlm_provider.py index 789ee573..4dc50999 100644 --- a/src/askui/model_providers/openai_vlm_provider.py +++ b/src/askui/model_providers/openai_vlm_provider.py @@ -8,7 +8,7 @@ from typing_extensions import override from askui.model_providers.vlm_provider import VlmProvider -from askui.models.openai.messages_api import OpenAIMessagesApi +from askui.models.openai.messages_api import MessageTransform, OpenAIMessagesApi from askui.models.shared.agent_message_param import ( MessageParam, ThinkingConfigParam, @@ -53,6 +53,12 @@ class OpenAIVlmProvider(VlmProvider): for screenshots sent to the model. Only used when ``image_scaler`` is not provided. Reads ``ASKUI_VLM_MAX_IMAGE_EDGE`` from the environment if not provided. Defaults to 1024. + message_transform (`MessageTransform` | None, optional): Hook to + post-process the OpenAI-format ``messages`` list right before it is + sent. Receives and returns the list of OpenAI message dicts. Use it + for OpenAI-compatible gateways that deviate from the stock chat spec + (e.g. stricter message-ordering or content rules). ``None`` (default) + sends the messages unchanged. Example: ```python @@ -81,6 +87,7 @@ def __init__( output_cost_per_million_tokens: float | None = None, cache_write_cost_per_million_tokens: float | None = None, cache_read_cost_per_million_tokens: float | None = None, + message_transform: MessageTransform | None = None, ) -> None: self._model_id_value = ( model_id or os.environ.get("VLM_PROVIDER_MODEL_ID") or _DEFAULT_MODEL_ID @@ -103,6 +110,7 @@ def __init__( api_key=api_key, base_url=base_url, ) + self._message_transform = message_transform self._pricing = ModelPricing.for_model( self._model_id_value, @@ -135,7 +143,10 @@ def image_scaler(self) -> ImageScaler: @cached_property def _messages_api(self) -> OpenAIMessagesApi: """Lazily initialise the `OpenAIMessagesApi` on first use.""" - return OpenAIMessagesApi(client=self._client) + return OpenAIMessagesApi( + client=self._client, + message_transform=self._message_transform, + ) @override def augment_system_prompt(self, system: SystemPrompt) -> SystemPrompt: diff --git a/src/askui/models/openai/messages_api.py b/src/askui/models/openai/messages_api.py index 51eacc0d..f1e3c5eb 100644 --- a/src/askui/models/openai/messages_api.py +++ b/src/askui/models/openai/messages_api.py @@ -2,6 +2,7 @@ import json import logging +from collections.abc import Callable from typing import Any from openai import OpenAI @@ -303,11 +304,23 @@ def _from_openai_response(response: ChatCompletion) -> MessageParam: ) +#: A hook to post-process the OpenAI-format ``messages`` list right before it is +#: sent to the chat/completions endpoint. Receives (and returns) the list of +#: OpenAI message dicts. Useful for OpenAI-compatible gateways that deviate from +#: the stock chat spec (e.g. stricter message-ordering or content rules). +MessageTransform = Callable[[list[dict[str, Any]]], list[dict[str, Any]]] + + class OpenAIMessagesApi(MessagesApi): """MessagesApi implementation for any OpenAI-compatible chat API.""" - def __init__(self, client: OpenAI) -> None: + def __init__( + self, + client: OpenAI, + message_transform: MessageTransform | None = None, + ) -> None: self._client = client + self._message_transform = message_transform @override def create_message( @@ -341,6 +354,8 @@ def create_message( The model's response as a `MessageParam`. """ openai_messages = _to_openai_messages(messages, system) + if self._message_transform is not None: + openai_messages = self._message_transform(openai_messages) kwargs: dict[str, Any] = { "model": model_id, diff --git a/tests/unit/models/openai/test_message_transform.py b/tests/unit/models/openai/test_message_transform.py new file mode 100644 index 00000000..6562189a --- /dev/null +++ b/tests/unit/models/openai/test_message_transform.py @@ -0,0 +1,54 @@ +"""Unit tests for the OpenAIMessagesApi message_transform hook.""" + +from typing import Any +from unittest.mock import MagicMock + +from askui.models.openai.messages_api import OpenAIMessagesApi +from askui.models.shared.agent_message_param import MessageParam + + +def _fake_completion(content: str = "ok") -> MagicMock: + message = MagicMock() + message.content = content + message.tool_calls = None + choice = MagicMock() + choice.message = message + choice.finish_reason = "stop" + response = MagicMock() + response.choices = [choice] + response.usage = None + return response + + +def _client_capturing(captured: dict[str, Any]) -> MagicMock: + client = MagicMock() + + def create(**kwargs: Any) -> MagicMock: + captured["messages"] = kwargs["messages"] + return _fake_completion() + + client.chat.completions.create.side_effect = create + return client + + +def test_message_transform_is_applied_before_send() -> None: + captured: dict = {} + client = _client_capturing(captured) + + def transform(messages: list[dict]) -> list[dict]: + return [*messages, {"role": "user", "content": "SHIM"}] + + api = OpenAIMessagesApi(client=client, message_transform=transform) + api.create_message(messages=[MessageParam(role="user", content="hi")], model_id="m") + + assert captured["messages"][-1] == {"role": "user", "content": "SHIM"} + + +def test_no_transform_sends_messages_unchanged() -> None: + captured: dict = {} + client = _client_capturing(captured) + + api = OpenAIMessagesApi(client=client) + api.create_message(messages=[MessageParam(role="user", content="hi")], model_id="m") + + assert captured["messages"] == [{"role": "user", "content": "hi"}]