diff --git a/tests/test_client.py b/tests/test_client.py index 2f6c8c1..6afbe70 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1,7 +1,6 @@ import os import io import json -import types import pytest import tempfile from pathlib import Path @@ -81,9 +80,11 @@ def test_client_chat_stream(httpserver: HTTPServer): client = Client(httpserver.url_for('/')) response = client.chat('dummy', messages=[{'role': 'user', 'content': 'Why is the sky blue?'}], stream=True) + + it = iter(['I ', "don't ", 'know.']) for part in response: assert part['message']['role'] in 'assistant' - assert part['message']['content'] in ['I ', "don't ", 'know.'] + assert part['message']['content'] == next(it) def test_client_chat_images(httpserver: HTTPServer): @@ -187,9 +188,11 @@ def test_client_generate_stream(httpserver: HTTPServer): client = Client(httpserver.url_for('/')) response = client.generate('dummy', 'Why is the sky blue?', stream=True) + + it = iter(['Because ', 'it ', 'is.']) for part in response: assert part['model'] == 'dummy' - assert part['response'] in ['Because ', 'it ', 'is.'] + assert part['response'] == next(it) def test_client_generate_images(httpserver: HTTPServer): @@ -263,11 +266,14 @@ def test_client_pull_stream(httpserver: HTTPServer): 'insecure': False, 'stream': True, }, - ).respond_with_json({}) + ).respond_with_handler(stream_handler) client = Client(httpserver.url_for('/')) response = client.pull('dummy', stream=True) - assert isinstance(response, types.GeneratorType) + + it = iter(['pulling manifest', 'verifying sha256 digest', 'writing manifest', 'removing any unused layers', 'success']) + for part in response: + assert part['status'] == next(it) def test_client_push(httpserver: HTTPServer): @@ -287,6 +293,14 @@ def test_client_push(httpserver: HTTPServer): def test_client_push_stream(httpserver: HTTPServer): + def stream_handler(_: Request): + def generate(): + yield json.dumps({'status': 'retrieving manifest'}) + '\n' + yield json.dumps({'status': 'pushing manifest'}) + '\n' + yield json.dumps({'status': 'success'}) + '\n' + + return Response(generate()) + httpserver.expect_ordered_request( '/api/push', method='POST', @@ -295,11 +309,14 @@ def test_client_push_stream(httpserver: HTTPServer): 'insecure': False, 'stream': True, }, - ).respond_with_json({}) + ).respond_with_handler(stream_handler) client = Client(httpserver.url_for('/')) response = client.push('dummy', stream=True) - assert isinstance(response, types.GeneratorType) + + it = iter(['retrieving manifest', 'pushing manifest', 'success']) + for part in response: + assert part['status'] == next(it) def test_client_create_path(httpserver: HTTPServer): @@ -458,6 +475,24 @@ async def test_async_client_chat(httpserver: HTTPServer): @pytest.mark.asyncio async def test_async_client_chat_stream(httpserver: HTTPServer): + def stream_handler(_: Request): + def generate(): + for message in ['I ', "don't ", 'know.']: + yield ( + json.dumps( + { + 'model': 'dummy', + 'message': { + 'role': 'assistant', + 'content': message, + }, + } + ) + + '\n' + ) + + return Response(generate()) + httpserver.expect_ordered_request( '/api/chat', method='POST', @@ -468,11 +503,15 @@ async def test_async_client_chat_stream(httpserver: HTTPServer): 'format': '', 'options': {}, }, - ).respond_with_json({}) + ).respond_with_handler(stream_handler) client = AsyncClient(httpserver.url_for('/')) response = await client.chat('dummy', messages=[{'role': 'user', 'content': 'Why is the sky blue?'}], stream=True) - assert isinstance(response, types.AsyncGeneratorType) + + it = iter(['I ', "don't ", 'know.']) + async for part in response: + assert part['message']['role'] == 'assistant' + assert part['message']['content'] == next(it) @pytest.mark.asyncio @@ -529,6 +568,21 @@ async def test_async_client_generate(httpserver: HTTPServer): @pytest.mark.asyncio async def test_async_client_generate_stream(httpserver: HTTPServer): + def stream_handler(_: Request): + def generate(): + for message in ['Because ', 'it ', 'is.']: + yield ( + json.dumps( + { + 'model': 'dummy', + 'response': message, + } + ) + + '\n' + ) + + return Response(generate()) + httpserver.expect_ordered_request( '/api/generate', method='POST', @@ -544,11 +598,15 @@ async def test_async_client_generate_stream(httpserver: HTTPServer): 'format': '', 'options': {}, }, - ).respond_with_json({}) + ).respond_with_handler(stream_handler) client = AsyncClient(httpserver.url_for('/')) response = await client.generate('dummy', 'Why is the sky blue?', stream=True) - assert isinstance(response, types.AsyncGeneratorType) + + it = iter(['Because ', 'it ', 'is.']) + async for part in response: + assert part['model'] == 'dummy' + assert part['response'] == next(it) @pytest.mark.asyncio @@ -597,6 +655,16 @@ async def test_async_client_pull(httpserver: HTTPServer): @pytest.mark.asyncio async def test_async_client_pull_stream(httpserver: HTTPServer): + def stream_handler(_: Request): + def generate(): + yield json.dumps({'status': 'pulling manifest'}) + '\n' + yield json.dumps({'status': 'verifying sha256 digest'}) + '\n' + yield json.dumps({'status': 'writing manifest'}) + '\n' + yield json.dumps({'status': 'removing any unused layers'}) + '\n' + yield json.dumps({'status': 'success'}) + '\n' + + return Response(generate()) + httpserver.expect_ordered_request( '/api/pull', method='POST', @@ -605,11 +673,14 @@ async def test_async_client_pull_stream(httpserver: HTTPServer): 'insecure': False, 'stream': True, }, - ).respond_with_json({}) + ).respond_with_handler(stream_handler) client = AsyncClient(httpserver.url_for('/')) response = await client.pull('dummy', stream=True) - assert isinstance(response, types.AsyncGeneratorType) + + it = iter(['pulling manifest', 'verifying sha256 digest', 'writing manifest', 'removing any unused layers', 'success']) + async for part in response: + assert part['status'] == next(it) @pytest.mark.asyncio @@ -631,6 +702,14 @@ async def test_async_client_push(httpserver: HTTPServer): @pytest.mark.asyncio async def test_async_client_push_stream(httpserver: HTTPServer): + def stream_handler(_: Request): + def generate(): + yield json.dumps({'status': 'retrieving manifest'}) + '\n' + yield json.dumps({'status': 'pushing manifest'}) + '\n' + yield json.dumps({'status': 'success'}) + '\n' + + return Response(generate()) + httpserver.expect_ordered_request( '/api/push', method='POST', @@ -639,11 +718,14 @@ async def test_async_client_push_stream(httpserver: HTTPServer): 'insecure': False, 'stream': True, }, - ).respond_with_json({}) + ).respond_with_handler(stream_handler) client = AsyncClient(httpserver.url_for('/')) response = await client.push('dummy', stream=True) - assert isinstance(response, types.AsyncGeneratorType) + + it = iter(['retrieving manifest', 'pushing manifest', 'success']) + async for part in response: + assert part['status'] == next(it) @pytest.mark.asyncio