diff --git a/ollama/_client.py b/ollama/_client.py index 8dfce824..274349d9 100644 --- a/ollama/_client.py +++ b/ollama/_client.py @@ -181,18 +181,21 @@ def _request( if stream: def inner(): - with self._client.stream(*args, **kwargs) as r: - try: - r.raise_for_status() - except httpx.HTTPStatusError as e: - e.response.read() - raise ResponseError(e.response.text, e.response.status_code) from None - - for line in r.iter_lines(): - part = json.loads(line) - if err := part.get('error'): - raise ResponseError(err) - yield cls(**part) + try: + with self._client.stream(*args, **kwargs) as r: + try: + r.raise_for_status() + except httpx.HTTPStatusError as e: + e.response.read() + raise ResponseError(e.response.text, e.response.status_code) from None + + for line in r.iter_lines(): + part = json.loads(line) + if err := part.get('error'): + raise ResponseError(err) + yield cls(**part) + except httpx.ConnectError: + raise ConnectionError(CONNECTION_ERROR_MESSAGE) from None return inner() @@ -774,18 +777,21 @@ async def _request( if stream: async def inner(): - async with self._client.stream(*args, **kwargs) as r: - try: - r.raise_for_status() - except httpx.HTTPStatusError as e: - await e.response.aread() - raise ResponseError(e.response.text, e.response.status_code) from None - - async for line in r.aiter_lines(): - part = json.loads(line) - if err := part.get('error'): - raise ResponseError(err) - yield cls(**part) + try: + async with self._client.stream(*args, **kwargs) as r: + try: + r.raise_for_status() + except httpx.HTTPStatusError as e: + await e.response.aread() + raise ResponseError(e.response.text, e.response.status_code) from None + + async for line in r.aiter_lines(): + part = json.loads(line) + if err := part.get('error'): + raise ResponseError(err) + yield cls(**part) + except httpx.ConnectError: + raise ConnectionError(CONNECTION_ERROR_MESSAGE) from None return inner() diff --git a/tests/test_client.py b/tests/test_client.py index 7b7ab38e..4a9b26fa 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -8,13 +8,14 @@ from typing import Any import pytest +from httpx import ConnectError, MockTransport from httpx import Response as httpxResponse from pydantic import BaseModel from pytest_httpserver import HTTPServer, URIPattern from werkzeug.wrappers import Request, Response from ollama._client import CONNECTION_ERROR_MESSAGE, AsyncClient, Client, _copy_tools -from ollama._types import Image, Message +from ollama._types import Image, Message, ResponseError PNG_BASE64 = 'iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGNgYGAAAAAEAAH2FzhVAAAAAElFTkSuQmCC' PNG_BYTES = base64.b64decode(PNG_BASE64) @@ -1338,6 +1339,68 @@ def test_client_connection_error(): client.show('model') +@pytest.fixture +def disconnected_transport(): + def handler(request): + raise ConnectError('connection refused', request=request) + + return MockTransport(handler) + + +@pytest.mark.parametrize('method', ['chat', 'generate', 'pull', 'push', 'create']) +@pytest.mark.parametrize('stream', [False, True]) +def test_client_connection_error_streaming_parity(disconnected_transport, method, stream): + client = Client(transport=disconnected_transport) + try: + with pytest.raises(ConnectionError) as exc_info: + response = getattr(client, method)('model', stream=stream) + if stream: + list(response) + assert str(exc_info.value) == CONNECTION_ERROR_MESSAGE + finally: + client.close() + + +@pytest.mark.parametrize('method', ['chat', 'generate', 'pull', 'push', 'create']) +@pytest.mark.parametrize('stream', [False, True]) +async def test_async_client_connection_error_streaming_parity(disconnected_transport, method, stream): + client = AsyncClient(transport=disconnected_transport) + try: + with pytest.raises(ConnectionError) as exc_info: + response = await getattr(client, method)('model', stream=stream) + if stream: + async for _ in response: + pass + assert str(exc_info.value) == CONNECTION_ERROR_MESSAGE + finally: + await client.close() + + +@pytest.mark.parametrize('status_code', [200, 500]) +def test_client_stream_response_error(status_code): + transport = MockTransport(lambda request: httpxResponse(status_code, json={'error': 'model failed'})) + client = Client(transport=transport) + try: + with pytest.raises(ResponseError, match='model failed') as exc_info: + list(client.generate('model', stream=True)) + assert exc_info.value.status_code == (-1 if status_code == 200 else status_code) + finally: + client.close() + + +@pytest.mark.parametrize('status_code', [200, 500]) +async def test_async_client_stream_response_error(status_code): + transport = MockTransport(lambda request: httpxResponse(status_code, json={'error': 'model failed'})) + client = AsyncClient(transport=transport) + try: + with pytest.raises(ResponseError, match='model failed') as exc_info: + async for _ in await client.generate('model', stream=True): + pass + assert exc_info.value.status_code == (-1 if status_code == 200 else status_code) + finally: + await client.close() + + async def test_async_client_connection_error(): client = AsyncClient('http://localhost:1234') with pytest.raises(ConnectionError) as exc_info: