Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/askui/model_providers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -59,6 +60,7 @@
"ImageQAProvider",
"ContainedImageScaler",
"ImageScaler",
"MessageTransform",
"ModelPricing",
"PatchOptimizedImageScaler",
"NormalizedCoordinateSpace",
Expand Down
7 changes: 7 additions & 0 deletions src/askui/model_providers/openai_compatible_vlm_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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,
)
15 changes: 13 additions & 2 deletions src/askui/model_providers/openai_vlm_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
17 changes: 16 additions & 1 deletion src/askui/models/openai/messages_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import json
import logging
from collections.abc import Callable
from typing import Any

from openai import OpenAI
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down
54 changes: 54 additions & 0 deletions tests/unit/models/openai/test_message_transform.py
Original file line number Diff line number Diff line change
@@ -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"}]
Loading