diff --git a/ollama/_client.py b/ollama/_client.py index 8dfce824..823282ec 100644 --- a/ollama/_client.py +++ b/ollama/_client.py @@ -1313,7 +1313,10 @@ async def ps(self) -> ProcessResponse: ) -def _copy_images(images: Optional[Sequence[Union[Image, Any]]]) -> Iterator[Image]: +def _copy_images(images: Optional[Union[Image, Sequence[Union[Image, Any]]]]) -> Iterator[Image]: + if isinstance(images, (str, bytes, PathLike, Image)): + images = [images] + for image in images or []: yield image if isinstance(image, Image) else Image(value=image) diff --git a/tests/test_client.py b/tests/test_client.py index 7b7ab38e..bd3525d8 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -191,6 +191,60 @@ def test_client_chat_images(httpserver: HTTPServer, message_format: str, file_st assert response['message']['content'] == "I don't know." +@pytest.mark.parametrize('image_content', (PNG_BASE64, PNG_BYTES)) +def test_client_chat_single_image_not_in_list(httpserver: HTTPServer, image_content): + httpserver.expect_ordered_request( + '/api/chat', + method='POST', + json={ + 'model': 'dummy', + 'messages': [ + { + 'role': 'user', + 'content': 'Why is the sky blue?', + 'images': [PNG_BASE64], + }, + ], + 'tools': [], + 'stream': False, + }, + ).respond_with_json( + { + 'model': 'dummy', + 'message': { + 'role': 'assistant', + 'content': "I don't know.", + }, + } + ) + + client = Client(httpserver.url_for('/')) + response = client.chat('dummy', messages=[{'role': 'user', 'content': 'Why is the sky blue?', 'images': image_content}]) + assert response['message']['content'] == "I don't know." + + +def test_client_generate_single_image_not_in_list(httpserver: HTTPServer): + httpserver.expect_ordered_request( + '/api/generate', + method='POST', + json={ + 'model': 'dummy', + 'prompt': 'What is in this image?', + 'images': [PNG_BASE64], + 'stream': False, + }, + ).respond_with_json( + { + 'model': 'dummy', + 'response': 'A single pixel.', + } + ) + + client = Client(httpserver.url_for('/')) + response = client.generate('dummy', 'What is in this image?', images=PNG_BASE64) + assert response['response'] == 'A single pixel.' + + def test_client_chat_format_json(httpserver: HTTPServer): httpserver.expect_ordered_request( '/api/chat',