diff --git a/openviking_cli/client/_http_compat.py b/openviking_cli/client/_http_compat.py index 1190b60a22..9fe606c261 100644 --- a/openviking_cli/client/_http_compat.py +++ b/openviking_cli/client/_http_compat.py @@ -6,6 +6,7 @@ import json import os +from dataclasses import asdict, is_dataclass from pathlib import Path from typing import Any, Dict @@ -60,6 +61,20 @@ } +def _message_part_to_payload(part: Any) -> Any: + if not is_dataclass(part): + return part + + serialized_part = asdict(part) + if serialized_part.get("type") != "image_url": + return serialized_part + + image_url = {"url": serialized_part.get("url", "")} + if serialized_part.get("detail") is not None: + image_url["detail"] = serialized_part["detail"] + return {"type": "image_url", "image_url": image_url} + + def _timeout_configured_outside_call() -> bool: if os.getenv("OPENVIKING_TIMEOUT"): return True @@ -171,7 +186,7 @@ async def add_message( session_id: str, role: str, content: str | None = None, - parts: list[dict] | None = None, + parts: list[Any] | None = None, created_at: str | None = None, peer_id: str | None = None, telemetry: Any = False, @@ -181,7 +196,7 @@ async def add_message( ) -> Dict[str, Any]: payload: Dict[str, Any] = {"role": role} if parts is not None: - payload["parts"] = parts + payload["parts"] = [_message_part_to_payload(part) for part in parts] elif content is not None: payload["content"] = content else: @@ -244,7 +259,7 @@ def add_message( session_id: str, role: str, content: str | None = None, - parts: list[dict] | None = None, + parts: list[Any] | None = None, created_at: str | None = None, peer_id: str | None = None, telemetry: Any = False, diff --git a/sdk/python/openviking_sdk/client.py b/sdk/python/openviking_sdk/client.py index eddacc3522..470865c711 100644 --- a/sdk/python/openviking_sdk/client.py +++ b/sdk/python/openviking_sdk/client.py @@ -7,6 +7,7 @@ import tempfile import uuid import zipfile +from dataclasses import asdict, is_dataclass from enum import Enum from pathlib import Path from typing import Any, Callable, Dict, List, Optional, Union @@ -67,6 +68,19 @@ GATEWAY_TOKEN_HEADER = "X-Gateway-Token" +def _message_part_to_payload(part: Any) -> Any: + if not is_dataclass(part): + return part + + serialized_part = asdict(part) + if serialized_part.get("type") != "image_url": + return serialized_part + + image_url = {"url": serialized_part.get("url", "")} + if serialized_part.get("detail") is not None: + image_url["detail"] = serialized_part["detail"] + return {"type": "image_url", "image_url": image_url} + def _image_mime_type(file_name: str = "") -> str: mime_type, _ = mimetypes.guess_type(file_name or "") @@ -137,7 +151,7 @@ async def add_message( self, role: str, content: str | None = None, - parts: list[dict] | None = None, + parts: list[Any] | None = None, created_at: str | None = None, peer_id: str | None = None, turn_id: str | None = None, @@ -213,7 +227,7 @@ def add_message( self, role: str, content: str | None = None, - parts: list[dict] | None = None, + parts: list[Any] | None = None, created_at: str | None = None, peer_id: str | None = None, turn_id: str | None = None, @@ -1467,7 +1481,7 @@ async def add_message( session_id: str, role: str, content: str | None = None, - parts: list[dict] | None = None, + parts: list[Any] | None = None, created_at: str | None = None, peer_id: str | None = None, telemetry: Any = False, @@ -1477,7 +1491,7 @@ async def add_message( ) -> Dict[str, Any]: payload: Dict[str, Any] = {"role": role} if parts is not None: - payload["parts"] = parts + payload["parts"] = [_message_part_to_payload(part) for part in parts] elif content is not None: payload["content"] = content else: @@ -2443,7 +2457,7 @@ def add_message( session_id: str, role: str, content: str | None = None, - parts: list[dict] | None = None, + parts: list[Any] | None = None, created_at: str | None = None, peer_id: str | None = None, telemetry: Any = False, diff --git a/sdk/python/tests/test_async_client_behaviors.py b/sdk/python/tests/test_async_client_behaviors.py index 9bad6fed2e..66c50cb0fa 100644 --- a/sdk/python/tests/test_async_client_behaviors.py +++ b/sdk/python/tests/test_async_client_behaviors.py @@ -1,14 +1,38 @@ import inspect +import json +from dataclasses import dataclass from pathlib import Path from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch +import httpx import pytest from openviking_sdk import AsyncHTTPClient, SyncHTTPClient from openviking_sdk.client import Session, SyncSession from openviking_sdk.errors import NotFoundError +@dataclass +class DataclassTextPart: + text: str + type: str = "text" + + +@dataclass +class DataclassImagePart: + url: str + detail: str | None = None + type: str = "image_url" + + +@dataclass +class DataclassToolPart: + tool_id: str + tool_name: str + tool_input: dict | None = None + type: str = "tool" + + def test_add_resource_signatures_keep_telemetry_position(): for func in (AsyncHTTPClient.add_resource, SyncHTTPClient.add_resource): params = list(inspect.signature(func).parameters) @@ -137,6 +161,65 @@ async def test_async_http_client_sends_message_semantics_and_turn_retention(): } +def test_sync_http_client_converts_dataclass_message_parts_to_payload(): + request_payloads = [] + + def handle_request(request): + request_payloads.append(json.loads(request.content)) + return httpx.Response( + 200, + json={"status": "success", "result": {"message_id": "msg-1"}}, + ) + + client = SyncHTTPClient(url="http://localhost:1933") + client._async_client._http = httpx.AsyncClient( + base_url="http://localhost:1933", + transport=httpx.MockTransport(handle_request), + ) + try: + result = client.add_message( + "demo-session", + "user", + parts=[ + DataclassTextPart(text="Hello world!"), + DataclassImagePart( + url="https://example.com/image.png", + detail="high", + ), + DataclassToolPart( + tool_id="call-1", + tool_name="search", + tool_input={"query": "hello"}, + ), + ], + ) + finally: + client.close() + + assert result == {"message_id": "msg-1"} + assert request_payloads == [ + { + "role": "user", + "parts": [ + {"text": "Hello world!", "type": "text"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/image.png", + "detail": "high", + }, + }, + { + "tool_id": "call-1", + "tool_name": "search", + "tool_input": {"query": "hello"}, + "type": "tool", + }, + ], + } + ] + + @pytest.mark.asyncio async def test_async_http_client_reindex_posts_content_reindex(): client = AsyncHTTPClient(url="http://localhost:1933") diff --git a/sdk/python/tests/test_main_package_exports.py b/sdk/python/tests/test_main_package_exports.py index ef8c2332da..cb6bce8ec8 100644 --- a/sdk/python/tests/test_main_package_exports.py +++ b/sdk/python/tests/test_main_package_exports.py @@ -1,6 +1,9 @@ +import json import sys from pathlib import Path +import httpx + SDK_ROOT = Path(__file__).resolve().parents[1] REPO_ROOT = Path(__file__).resolve().parents[3] @@ -80,3 +83,54 @@ def test_openviking_http_client_preserves_legacy_exception_types(): assert exc.code == "CONFLICT" else: raise AssertionError("expected ConflictError") + + +def test_openviking_sync_http_client_converts_message_parts_to_payload(): + _purge_openviking_modules() + import openviking + from openviking.message import ImagePart, TextPart, ToolPart + + request_payloads = [] + + def handle_request(request): + request_payloads.append(json.loads(request.content)) + return httpx.Response( + 200, + json={"status": "success", "result": {"message_id": "msg-1"}}, + ) + + client = openviking.SyncHTTPClient(url="http://localhost:1933") + client._async_client._http = httpx.AsyncClient( + base_url="http://localhost:1933", + transport=httpx.MockTransport(handle_request), + ) + try: + result = client.add_message( + "demo-session", + "user", + parts=[ + TextPart(text="Hello world!"), + ImagePart(url="https://example.com/image.png", detail="high"), + ToolPart( + tool_id="call-1", + tool_name="search", + tool_input={"query": "hello"}, + ), + ], + ) + finally: + client.close() + + assert result == {"message_id": "msg-1"} + assert request_payloads[0]["role"] == "user" + assert request_payloads[0]["parts"][0] == {"text": "Hello world!", "type": "text"} + assert request_payloads[0]["parts"][1] == { + "type": "image_url", + "image_url": { + "url": "https://example.com/image.png", + "detail": "high", + }, + } + assert request_payloads[0]["parts"][2]["type"] == "tool" + assert request_payloads[0]["parts"][2]["tool_id"] == "call-1" + assert request_payloads[0]["parts"][2]["tool_input"] == {"query": "hello"}