diff --git a/tests/test_agent.py b/tests/test_agent.py index 2f76220..cffad05 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -1,5 +1,6 @@ """Tests for the Agent resource.""" +import json import pytest from datetime import datetime from pydantic import BaseModel @@ -12,6 +13,7 @@ CreditUsage, ) from vlmrun.client.agent import Agent +from vlmrun.types import MessageContent class SampleInputModel(BaseModel): @@ -296,3 +298,88 @@ class MockClient: agent = Agent(MockClient()) result = agent._process_inputs(None) assert result is None + + def test_process_inputs_dict_with_nested_basemodel(self): + """Test that dict inputs with nested BaseModel values are JSON-serializable.""" + + class MockClient: + api_key = "test-key" + base_url = "https://api.vlm.run/v1" + timeout = 120.0 + max_retries = 1 + + agent = Agent(MockClient()) + inputs = { + "file": MessageContent( + type="input_file", file_id="d9f74779-5e5f-4bec-a901-97c55598c56e" + ), + } + result = agent._process_inputs(inputs) + + assert isinstance(result, dict) + assert isinstance(result["file"], dict) + assert result["file"]["type"] == "input_file" + assert result["file"]["file_id"] == "d9f74779-5e5f-4bec-a901-97c55598c56e" + json.dumps(result) + + def test_process_inputs_dict_with_list_of_basemodels(self): + """Test that dict inputs with lists of BaseModel values are serialized.""" + + class MockClient: + api_key = "test-key" + base_url = "https://api.vlm.run/v1" + timeout = 120.0 + max_retries = 1 + + agent = Agent(MockClient()) + inputs = { + "files": [ + MessageContent(type="input_file", file_id="aaa"), + MessageContent(type="input_file", file_id="bbb"), + ], + } + result = agent._process_inputs(inputs) + + assert isinstance(result["files"], list) + assert all(isinstance(item, dict) for item in result["files"]) + assert result["files"][0]["file_id"] == "aaa" + assert result["files"][1]["file_id"] == "bbb" + json.dumps(result) + + def test_process_inputs_dict_plain_values_unchanged(self): + """Test that dict inputs with plain string/int values pass through unchanged.""" + + class MockClient: + api_key = "test-key" + base_url = "https://api.vlm.run/v1" + timeout = 120.0 + max_retries = 1 + + agent = Agent(MockClient()) + inputs = {"url": "https://example.com/image.jpg", "count": 3, "flag": True} + result = agent._process_inputs(inputs) + + assert result == {"url": "https://example.com/image.jpg", "count": 3, "flag": True} + json.dumps(result) + + def test_process_inputs_dict_with_nested_dict_containing_basemodel(self): + """Test that deeply nested BaseModel values inside dicts are serialized.""" + + class MockClient: + api_key = "test-key" + base_url = "https://api.vlm.run/v1" + timeout = 120.0 + max_retries = 1 + + agent = Agent(MockClient()) + inputs = { + "metadata": { + "content": MessageContent(type="text", text="hello world"), + }, + } + result = agent._process_inputs(inputs) + + assert isinstance(result["metadata"]["content"], dict) + assert result["metadata"]["content"]["type"] == "text" + assert result["metadata"]["content"]["text"] == "hello world" + json.dumps(result) diff --git a/vlmrun/client/agent.py b/vlmrun/client/agent.py index 7f859f4..727b785 100644 --- a/vlmrun/client/agent.py +++ b/vlmrun/client/agent.py @@ -33,6 +33,17 @@ def __init__(self, client: "VLMRunProtocol") -> None: self._client = client self._requestor = APIRequestor(client) + @staticmethod + def _serialize_value(value: Any) -> Any: + """Recursively serialize a value, converting BaseModel instances to dicts.""" + if isinstance(value, BaseModel): + return value.model_dump(exclude_none=True) + elif isinstance(value, dict): + return {k: Agent._serialize_value(v) for k, v in value.items()} + elif isinstance(value, list): + return [Agent._serialize_value(item) for item in value] + return value + def _process_inputs( self, inputs: Union[dict[str, Any], BaseModel, None] ) -> Optional[dict[str, Any]]: @@ -53,6 +64,7 @@ def _process_inputs( DeprecationWarning, stacklevel=3, ) + return {k: self._serialize_value(v) for k, v in inputs.items()} return inputs def get( diff --git a/vlmrun/version.py b/vlmrun/version.py index a779a44..1cc82e6 100644 --- a/vlmrun/version.py +++ b/vlmrun/version.py @@ -1 +1 @@ -__version__ = "0.5.6" +__version__ = "0.5.7"