diff --git a/tests/agent/test_llm_io.py b/tests/agent/test_llm_io.py index 852806f13..08b2c25d1 100644 --- a/tests/agent/test_llm_io.py +++ b/tests/agent/test_llm_io.py @@ -3,6 +3,8 @@ import json import threading from concurrent.futures import ThreadPoolExecutor +from contextlib import asynccontextmanager +from functools import wraps from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from types import SimpleNamespace from unittest.mock import AsyncMock, Mock @@ -13,6 +15,7 @@ from omegaconf import OmegaConf from ms_agent.agent.llm_agent import LLMAgent +from ms_agent.llm import io from ms_agent.llm.openai_llm import OpenAI as LegacyOpenAI from ms_agent.llm.transport.openai_compat import OpenAICompatTransport from ms_agent.llm.utils import Message @@ -87,38 +90,26 @@ def close(self): self.release.set() -def make_agent(handler, stream=True, legacy=False, protocol='openai'): +def make_agent(handler=None, *, stream=True, legacy=False, protocol='openai', + base_url='http://model.test/v1', http_client=None): config = OmegaConf.create({ - 'llm': {'model': 'probe'}, 'generation_config': {'stream': stream} - }) + 'llm': {'model': 'probe'}, 'generation_config': {'stream': stream}}) if protocol == 'anthropic': - anthropic = pytest.importorskip('anthropic') + pytest.importorskip('anthropic') from ms_agent.llm.anthropic_llm import Anthropic from ms_agent.llm.transport.anthropic_messages import ( AnthropicMessagesTransport) - if legacy: - llm = Anthropic(config, api_key='local-test', - base_url='http://model.test/v1') - else: - llm = AnthropicMessagesTransport( - model='probe', api_key='local-test', - base_url='http://model.test/v1', - generation_config={'stream': stream}) - client_cls = anthropic.Anthropic - elif legacy: - llm = LegacyOpenAI(OmegaConf.create({ - 'llm': {'model': 'probe'}, 'generation_config': {'stream': stream} - }), api_key='local-test', base_url='http://model.test/v1') - client_cls = openai.OpenAI + cls = Anthropic if legacy else AnthropicMessagesTransport else: - llm = OpenAICompatTransport( - model='probe', api_key='local-test', base_url='http://model.test/v1', - generation_config={'stream': stream}) - client_cls = openai.OpenAI + cls = LegacyOpenAI if legacy else OpenAICompatTransport + credentials = {'api_key': 'local-test', 'base_url': base_url} + llm = (cls(config, **credentials) if legacy else cls( + model='probe', generation_config={'stream': stream}, **credentials)) llm.client.close() - llm.client = client_cls( - api_key='local-test', base_url='http://model.test/v1', max_retries=0, - http_client=httpx.Client(transport=httpx.MockTransport(handler))) + llm.client = type(llm.client)( + **credentials, max_retries=0, timeout=2, + http_client=http_client or httpx.Client( + transport=httpx.MockTransport(handler))) agent = LLMAgent.__new__(LLMAgent) agent.llm = llm agent.config = OmegaConf.create({ @@ -135,132 +126,129 @@ def make_agent(handler, stream=True, legacy=False, protocol='openai'): return agent -async def run_step(agent): - # Test the real step body without unrelated retry/backoff policy. - return [rows async for rows in LLMAgent.step.__wrapped__( - agent, [Message(role='user', content='Respond with OK')])] +def run_async(test): + # Keep SDK tests independent of the WebUI's pytest-asyncio dependency. + @wraps(test) + def run(*args, **kwargs): + return asyncio.run(test(*args, **kwargs)) + return run async def reached(event): - assert await asyncio.wait_for(asyncio.to_thread(event.wait, 5), 6) - - -def response(body): - return httpx.Response(200, headers={ - 'Content-Type': 'text/event-stream' if body.stream else 'application/json' - }, stream=body) + async def wait(): + while not event.is_set(): + await asyncio.sleep(.01) + # Do not use the executor being tested to wait for its progress. + await asyncio.wait_for(wait(), 5) + + +async def cancel(task): + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, .5) + + +@asynccontextmanager +async def running(agent, *releases): + async def step(): + # Exercise the real step body without unrelated retry/backoff policy. + return [rows async for rows in LLMAgent.step.__wrapped__( + agent, [Message(role='user', content='Respond with OK')])] + llm = agent.llm # Cleanup must retain this client if a test replaces it. + task = asyncio.create_task(step()) + try: + yield task + finally: + for event in releases: + event.set() + task.cancel() + await asyncio.gather(task, return_exceptions=True) + llm.client.close() + + +@asynccontextmanager +async def mock_request(protocol, *, stream=True, legacy=False, + wait_for='headers', error=None): + body = ResponseBody(stream, wait_for_body=True, protocol=protocol) + request = SimpleNamespace(body=body, headers_entered=threading.Event(), + release_headers=threading.Event(), threads=[]) + if wait_for == 'body': + request.release_headers.set() + + def handle(http_request): + request.threads.append(threading.get_ident()) + request.headers_entered.set() + assert request.release_headers.wait(10), 'headers not released' + if error: + raise error('simulated timeout', request=http_request) + return httpx.Response(200, stream=body, headers={ + 'Content-Type': 'text/event-stream' if stream else 'application/json'}) + + request.agent = make_agent(handle, stream=stream, legacy=legacy, + protocol=protocol) + async with running(request.agent, request.release_headers, body.release) as task: + request.task = task + yield request @pytest.mark.parametrize('stream', [True, False]) @pytest.mark.parametrize('legacy', [False, True]) @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_first_request_leaves_event_loop_responsive(stream, legacy, protocol): - - async def check(): - entered, release = threading.Event(), threading.Event() - body = ResponseBody(stream, protocol=protocol) - request_threads = [] - - def handle(request): - request_threads.append(threading.get_ident()) - entered.set() - assert release.wait(10), 'test did not release response headers' - return response(body) - - agent = make_agent(handle, stream, legacy=legacy, protocol=protocol) - task = asyncio.create_task(run_step(agent)) - try: - await reached(entered) - assert request_threads[0] != threading.get_ident() - assert not task.done() - # A separate coroutine runs while the HTTP request is still waiting. - tick = asyncio.Event() - asyncio.get_running_loop().call_soon(tick.set) - await asyncio.wait_for(tick.wait(), .5) - release.set() - await asyncio.wait_for(task, 2) - assert agent.handle_new_response.call_args.args[1].content == 'OK' - assert body.closed.is_set() - finally: - release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - - asyncio.run(check()) +@run_async +async def test_first_request_leaves_event_loop_responsive(stream, legacy, protocol): + async with mock_request(protocol, stream=stream, legacy=legacy) as r: + await reached(r.headers_entered) + assert r.threads[0] != threading.get_ident() + assert not r.task.done() + tick = asyncio.Event() + asyncio.get_running_loop().call_soon(tick.set) + await asyncio.wait_for(tick.wait(), .5) + r.release_headers.set() + r.body.release.set() + await asyncio.wait_for(r.task, 2) + assert r.agent.handle_new_response.call_args.args[1].content == 'OK' + assert r.body.closed.is_set() @pytest.mark.parametrize('legacy', [False, True]) @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_cancel_before_headers_closes_late_response_without_reading_body(legacy, protocol): - - async def check(): - entered, release = threading.Event(), threading.Event() - body = ResponseBody(True, wait_for_body=True, protocol=protocol) - - def handle(request): - entered.set() - assert release.wait(10), 'test did not release response headers' - return response(body) - - agent = make_agent(handle, legacy=legacy, protocol=protocol) - task = asyncio.create_task(run_step(agent)) - try: - await reached(entered) - task.cancel() - with pytest.raises(asyncio.CancelledError): - await asyncio.wait_for(task, .5) - release.set() - await reached(body.closed) - assert not body.entered.is_set() - agent.handle_new_response.assert_not_called() - finally: - release.set() - body.release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - - asyncio.run(check()) +@run_async +async def test_cancel_before_headers_closes_late_response_without_reading_body(legacy, protocol): + async with mock_request(protocol, legacy=legacy) as r: + await reached(r.headers_entered) + await cancel(r.task) + r.release_headers.set() + await reached(r.body.closed) + assert not r.body.entered.is_set() + r.agent.handle_new_response.assert_not_called() @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_cancel_as_headers_return_does_not_lose_response_ownership(protocol): - - async def check(): - body = ResponseBody(True, wait_for_body=True, protocol=protocol) - agent = make_agent(lambda request: response(body), protocol=protocol) - original = agent.llm.generate +@run_async +async def test_cancel_as_headers_return_does_not_lose_response_ownership(protocol): + async with mock_request(protocol, wait_for='body') as r: + generate = r.agent.llm.generate loop = asyncio.get_running_loop() - def generate(*args, **kwargs): - result = original(*args, **kwargs) - loop.call_soon_threadsafe(task.cancel) + def cancel_on_return(*args, **kwargs): + result = generate(*args, **kwargs) + loop.call_soon_threadsafe(r.task.cancel) return result - agent.llm.generate = generate - task = asyncio.create_task(run_step(agent)) - try: - with pytest.raises(asyncio.CancelledError): - await asyncio.wait_for(task, 2) - await reached(body.closed) - assert not body.entered.is_set() - finally: - body.release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - - asyncio.run(check()) + r.agent.llm.generate = cancel_on_return + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(r.task, 2) + await reached(r.body.closed) + assert not r.body.entered.is_set() @pytest.mark.parametrize('legacy', [False, True]) @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_cancel_during_stream_read_interrupts_and_closes_iterator(legacy, protocol): - - async def check(): - body = ResponseBody(True, wait_for_body=True, protocol=protocol) - agent = make_agent(lambda request: response(body), legacy=legacy, - protocol=protocol) +@run_async +async def test_cancel_during_stream_read_interrupts_and_closes_iterator(legacy, protocol): + async with mock_request(protocol, legacy=legacy, wait_for='body') as r: iterator_closed = threading.Event() - generate = agent.llm.generate + generate = r.agent.llm.generate def tracked_generate(*args, **kwargs): chunks = generate(*args, **kwargs) @@ -271,146 +259,61 @@ def tracked_chunks(): finally: chunks.close() iterator_closed.set() - return tracked_chunks() - agent.llm.generate = tracked_generate - task = asyncio.create_task(run_step(agent)) - try: - await reached(body.entered) - assert body.reader_thread != threading.get_ident() - task.cancel() - with pytest.raises(asyncio.CancelledError): - await asyncio.wait_for(task, .5) - await reached(body.closed) - # Cancellation returns before the worker; wait for its iterator's - # finally block as well as the immediate HTTP-response close. - await reached(iterator_closed) - assert agent.llm._active_stream is None - agent.handle_new_response.assert_not_called() - finally: - body.release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - - asyncio.run(check()) + r.agent.llm.generate = tracked_generate + await reached(r.body.entered) + assert r.body.reader_thread != threading.get_ident() + await cancel(r.task) + await reached(r.body.closed) + # Wait for the worker's iterator finally, not just HTTP close. + await reached(iterator_closed) + assert r.agent.llm._active_stream is None + r.agent.handle_new_response.assert_not_called() @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_request_timeout_propagates_without_blocking_event_loop(protocol): - - async def check(): - entered, release = threading.Event(), threading.Event() - request_threads = [] - - def handle(request): - request_threads.append(threading.get_ident()) - entered.set() - assert release.wait(10), 'test did not release request timeout' - raise httpx.ReadTimeout('simulated timeout', request=request) - - agent = make_agent(handle, protocol=protocol) - task = asyncio.create_task(run_step(agent)) - try: - await reached(entered) - assert request_threads[0] != threading.get_ident() - release.set() - with pytest.raises((openai if protocol == 'openai' else - pytest.importorskip('anthropic')).APITimeoutError): - await asyncio.wait_for(task, 2) - agent.handle_new_response.assert_not_called() - finally: - release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - - asyncio.run(check()) +@run_async +async def test_request_timeout_propagates_without_blocking_event_loop(protocol): + async with mock_request(protocol, error=httpx.ReadTimeout) as r: + await reached(r.headers_entered) + assert r.threads[0] != threading.get_ident() + r.release_headers.set() + sdk = openai if protocol == 'openai' else pytest.importorskip('anthropic') + with pytest.raises(sdk.APITimeoutError): + await asyncio.wait_for(r.task, 2) + r.agent.handle_new_response.assert_not_called() @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_cancelled_request_cleanup_does_not_interrupt_replacement_llm(protocol): - """Reusing an SDK agent after cancellation can install a new transport.""" - - async def check(): - entered, release = threading.Event(), threading.Event() - body = ResponseBody(True, wait_for_body=True, protocol=protocol) - - def handle(request): - entered.set() - assert release.wait(10) - return response(body) - - agent = make_agent(handle, protocol=protocol) - original_llm = agent.llm +@run_async +async def test_cancelled_request_cleanup_does_not_interrupt_replacement_llm(protocol): + async with mock_request(protocol) as r: + await reached(r.headers_entered) + await cancel(r.task) replacement = Mock() - task = asyncio.create_task(run_step(agent)) - try: - await reached(entered) - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - # run() prepares a new LLM on the next invocation. The abandoned - # worker still owns the old transport, never this replacement. - agent.llm = replacement - release.set() - await reached(body.closed) - replacement.interrupt.assert_not_called() - assert body.closed.is_set() - finally: - release.set() - body.release.set() - await asyncio.gather(task, return_exceptions=True) - original_llm.client.close() - - asyncio.run(check()) + r.agent.llm = replacement + r.release_headers.set() + await reached(r.body.closed) + replacement.interrupt.assert_not_called() @pytest.mark.parametrize('wait_for_body', [False, True]) @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) -def test_slow_model_does_not_starve_host_executor(wait_for_body, protocol): - """Session storage must still run when all model workers are waiting.""" - - async def check(): - entered, release = threading.Event(), threading.Event() - body = ResponseBody(True, wait_for_body=wait_for_body, protocol=protocol) - - def handle(request): - entered.set() - if not wait_for_body: - assert release.wait(10) - return response(body) - - agent = make_agent(handle, protocol=protocol) - loop = asyncio.get_running_loop() - loop.set_default_executor(ThreadPoolExecutor(max_workers=1)) - task = asyncio.create_task(run_step(agent)) - try: - target = body.entered if wait_for_body else entered - async def wait_for_request(): - while not target.is_set(): - await asyncio.sleep(.01) - await asyncio.wait_for(wait_for_request(), 3) - assert await asyncio.wait_for( - asyncio.to_thread(lambda: 'session storage'), .5 - ) == 'session storage' - finally: - release.set() - body.release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - - asyncio.run(check()) +@run_async +async def test_slow_model_does_not_starve_host_executor(wait_for_body, protocol): + asyncio.get_running_loop().set_default_executor(ThreadPoolExecutor(max_workers=1)) + async with mock_request(protocol, wait_for='body' if wait_for_body else 'headers') as r: + await reached(r.body.entered if wait_for_body else r.headers_entered) + assert await asyncio.wait_for( + asyncio.to_thread(lambda: 'session storage'), .5) == 'session storage' @pytest.mark.parametrize('protocol', ['openai', 'anthropic']) @pytest.mark.parametrize('legacy', [False, True]) -def test_cancel_stalled_tcp_read_releases_model_worker(monkeypatch, protocol, - legacy): - """A real blocked socket must exit; a mock close that releases an Event - cannot detect cross-thread close() leaving recv() alive until timeout. - """ - from ms_agent.llm import io - +@run_async +async def test_cancel_stalled_tcp_read_releases_model_worker(monkeypatch, protocol, legacy): + """A mock close can release an Event while a real socket stays in recv().""" release, reading = threading.Event(), threading.Event() class Handler(BaseHTTPRequestHandler): @@ -446,37 +349,19 @@ def mark_response(response): server.daemon_threads = True thread = threading.Thread(target=server.serve_forever, daemon=True) thread.start() - - async def check(): - agent = make_agent(lambda r: None, legacy=legacy, protocol=protocol) - client_cls = type(agent.llm.client) - agent.llm.client.close() - agent.llm.client = client_cls( - api_key='local-test', - base_url=f'http://127.0.0.1:{server.server_port}', - max_retries=0, timeout=2, - http_client=httpx.Client(event_hooks={'response': [mark_response]})) - task = asyncio.create_task(run_step(agent)) - try: - await reached(reading) - # Give the worker time to enter recv(), not just the stream wrapper. - await asyncio.sleep(.05) - task.cancel() - with pytest.raises(asyncio.CancelledError): - await asyncio.wait_for(task, .5) - assert await asyncio.wait_for( - io.run_in_llm_executor(lambda: 'next request'), .5 - ) == 'next request' - assert not release.is_set() - finally: - release.set() - await asyncio.gather(task, return_exceptions=True) - agent.llm.client.close() - try: with ThreadPoolExecutor(max_workers=1) as pool: monkeypatch.setattr(io, '_executor', pool) - asyncio.run(check()) + agent = make_agent(legacy=legacy, protocol=protocol, + base_url=f'http://127.0.0.1:{server.server_port}', + http_client=httpx.Client(event_hooks={'response': [mark_response]})) + async with running(agent, release) as task: + await reached(reading) + await asyncio.sleep(.05) # Let the worker enter recv(). + await cancel(task) + assert await asyncio.wait_for( + io.run_in_llm_executor(lambda: 'next request'), .5) == 'next request' + assert not release.is_set() finally: release.set() server.shutdown() @@ -487,16 +372,13 @@ async def check(): @pytest.mark.parametrize('http_version,closed', [('HTTP/2', False), ('HTTP/1.1', True)]) def test_interrupt_preserves_shared_or_released_connection(http_version, closed): - from ms_agent.llm.io import interrupt_stream - network = Mock() response = httpx.Response( 200, stream=ResponseBody(True), - extensions={'http_version': http_version.encode(), - 'network_stream': network}) + extensions={'http_version': http_version.encode(), 'network_stream': network}) if closed: response.close() stream = SimpleNamespace(response=response, close=Mock()) - interrupt_stream(stream) + io.interrupt_stream(stream) network.get_extra_info.assert_not_called() stream.close.assert_called_once()