From bec8e4ff4de85266e45f1ffea188b7a3e9a18567 Mon Sep 17 00:00:00 2001 From: alcholiclg Date: Fri, 18 Sep 2026 11:25:44 +0800 Subject: [PATCH 1/2] fix(llm): keep async agents responsive during model requests --- ms_agent/agent/llm_agent.py | 74 ++- ms_agent/llm/anthropic_llm.py | 110 ++-- ms_agent/llm/io.py | 67 +++ ms_agent/llm/openai_llm.py | 33 +- ms_agent/llm/transport/anthropic_messages.py | 10 +- ms_agent/llm/transport/openai_compat.py | 40 +- tests/agent/test_llm_io.py | 502 +++++++++++++++++++ 7 files changed, 748 insertions(+), 88 deletions(-) create mode 100644 ms_agent/llm/io.py create mode 100644 tests/agent/test_llm_io.py diff --git a/ms_agent/agent/llm_agent.py b/ms_agent/agent/llm_agent.py index c48a4303d..11d8f2d09 100644 --- a/ms_agent/agent/llm_agent.py +++ b/ms_agent/agent/llm_agent.py @@ -20,6 +20,7 @@ from ms_agent.callbacks import Callback, callbacks_mapping from ms_agent.knowledge_search import SirchmunkSearch from ms_agent.llm import multimodal +from ms_agent.llm.io import run_in_llm_executor from ms_agent.llm.llm import LLM from ms_agent.llm.message_text import (append_text, flatten_message_text, prepend_text) @@ -2040,6 +2041,49 @@ def _append_task_notifications(self, messages.append(Message(role='user', content=body)) return messages + @staticmethod + def _close_llm_stream(llm, response=None) -> None: + """Best-effort cleanup, including a response returned after cancellation.""" + for close in (getattr(llm, 'interrupt', + None), getattr(response, 'close', None)): + if callable(close): + try: + close() + except Exception: # teardown must not mask the original error + pass + + async def _generate_response(self, messages, tools): + """Open a synchronous model request without blocking the event loop. + + Cancellation cannot stop a Python worker thread. If the request returns + after its caller has gone, the worker closes that response instead of + abandoning it. The lock covers only ownership transfer, never I/O. + """ + lock = threading.Lock() + llm = self.llm + abandoned = False + response = None + + def generate(): + nonlocal response + result = llm.generate(messages, tools=tools) + with lock: + discard = abandoned + if not discard: + response = result + if discard: + self._close_llm_stream(llm, result) + return result + + try: + return await run_in_llm_executor(generate) + except asyncio.CancelledError: + with lock: + abandoned = True + result = response + self._close_llm_stream(llm, result) + raise + # retry_if: a hard 4xx (bad payload, content filter, auth) is a verdict on # the request, not a transient fault — retrying it 5× only adds ~40s of # backoff before the same failure surfaces. @@ -2118,9 +2162,10 @@ async def step( # ui.events.ToolCallComposing). _composing: Dict[int, int] = {} _reported_images = False - _gen = self.llm.generate(messages, tools=tools) - _loop = asyncio.get_running_loop() + _llm = self.llm + _gen = await self._generate_response(messages, tools) _NO_MORE = object() + _stopped = threading.Event() def _next_chunk(_g=_gen): # Step the BLOCKING sync LLM stream off the event loop, so @@ -2133,10 +2178,15 @@ def _next_chunk(_g=_gen): return next(_g) except StopIteration: return _NO_MORE + finally: + # A cancelled await leaves next() running in its worker. + # Close the iterator there once that read has unwound. + if _stopped.is_set(): + self._close_llm_stream(_llm, _g) try: while True: - _chunk = await _loop.run_in_executor(None, _next_chunk) + _chunk = await run_in_llm_executor(_next_chunk) if _chunk is _NO_MORE: break _response_message = _chunk @@ -2189,24 +2239,13 @@ def _next_chunk(_g=_gen): messages[-1] = _response_message yield messages finally: + _stopped.set() if not _reported_images: # A turn that produced nothing still attached images, and # what became of them is still worth saying. self._record_image_deliveries(messages) _reported_images = True - # Turn abandoned mid-stream (client disconnect / stop): ask - # the provider to close the live upstream response so the - # server stops generating, instead of leaving it to run to - # completion into a dropped connection. Only the data-driven - # provider layer implements interrupt(); the legacy LLM does - # not, so this is a no-op there (unchanged). Harmless on a - # normal finish (the stream is already exhausted). - _interrupt = getattr(self.llm, 'interrupt', None) - if callable(_interrupt): - try: - _interrupt() - except Exception: # noqa: BLE001 - teardown never raises - pass + self._close_llm_stream(_llm, _gen) if self.stream_output: if _printed_reasoning_header and not _printed_reasoning_footer: self._emit_reasoning_end() @@ -2222,7 +2261,8 @@ def _next_chunk(_g=_gen): self._emit_content_end() else: - _response_message = self.llm.generate(messages, tools=tools) + _response_message = await self._generate_response( + messages, tools) if self.show_reasoning: reasoning_text = ( getattr(_response_message, 'reasoning_content', '') diff --git a/ms_agent/llm/anthropic_llm.py b/ms_agent/llm/anthropic_llm.py index 728f99016..bafa995f9 100644 --- a/ms_agent/llm/anthropic_llm.py +++ b/ms_agent/llm/anthropic_llm.py @@ -5,6 +5,7 @@ from typing import Any, Dict, Generator, Iterator, List, Optional, Union from ms_agent.llm import LLM +from ms_agent.llm.io import OpenedStream, interrupt_stream from ms_agent.llm.thinking import create_with_thinking_fallback from ms_agent.llm.utils import Message, Tool, ToolCall from ms_agent.utils import assert_package_exist, get_logger, retry @@ -167,6 +168,14 @@ def __init__( self.args: Dict = OmegaConf.to_container( getattr(config, 'generation_config', DictConfig({}))) + self._active_stream = None + + def interrupt(self) -> None: + """Close the owned HTTP stream, including before iteration starts.""" + stream = self._active_stream + interrupt_stream(stream) + if self._active_stream is stream: + self._active_stream = None def format_tools(self, tools: Optional[List[Tool]]) -> Optional[List[Dict]]: @@ -284,7 +293,9 @@ def _call_llm(self, def _send(**call): call.setdefault('model', self.model) if stream: - return self.client.messages.stream(**call) + opened = OpenedStream(self.client.messages.stream(**call)) + self._active_stream = opened + return opened return self.client.messages.create(**call) # This legacy engine owned no repair at all: a model that cannot think @@ -333,53 +344,58 @@ def _stream_format_output_message(self, ) tool_call_id_map = {} # index -> tool_call_id (用于去重 yield) with stream_manager as stream: - full_content = '' - full_thinking = '' - for event in stream: - event_type = getattr(event, 'type') - if event_type == 'message_start': - msg = event.message - current_message.id = msg.id - tool_call_id_map = {} - yield current_message - elif event_type == 'content_block_delta': - if event.delta.type == 'thinking_delta': - full_thinking += event.delta.thinking - current_message.reasoning_content = full_thinking - elif event.delta.type == 'text_delta': - full_content += event.delta.text + self._active_stream = stream + try: + full_content = '' + full_thinking = '' + for event in stream: + event_type = getattr(event, 'type') + if event_type == 'message_start': + msg = event.message + current_message.id = msg.id + tool_call_id_map = {} + yield current_message + elif event_type == 'content_block_delta': + if event.delta.type == 'thinking_delta': + full_thinking += event.delta.thinking + current_message.reasoning_content = full_thinking + elif event.delta.type == 'text_delta': + full_content += event.delta.text + current_message.content = full_content + yield current_message + elif event_type == 'message_stop': + final_msg = getattr(event, 'message') + full_content = '' + used_tool_call_ids = set() + for idx, block in enumerate(event.message.content): + if block is None: + continue + if block.type == 'text': + full_content += block.text + elif block.type == 'tool_use': + tool_call_id = tool_call_id_map.get(idx) + tool_call = ToolCall( + id=tool_call_id, + index=len(current_message.tool_calls), + type='function', + tool_name=block.name, + arguments=block.input, + ) + current_message.tool_calls.append(tool_call) + used_tool_call_ids.add(tool_call_id) current_message.content = full_content - yield current_message - elif event_type == 'message_stop': - final_msg = getattr(event, 'message') - full_content = '' - used_tool_call_ids = set() - for idx, block in enumerate(event.message.content): - if block is None: - continue - if block.type == 'text': - full_content += block.text - elif block.type == 'tool_use': - tool_call_id = tool_call_id_map.get(idx) - tool_call = ToolCall( - id=tool_call_id, - index=len(current_message.tool_calls), - type='function', - tool_name=block.name, - arguments=block.input, - ) - current_message.tool_calls.append(tool_call) - used_tool_call_ids.add(tool_call_id) - current_message.content = full_content - current_message.partial = False - current_message.completion_tokens = getattr( - final_msg.usage, 'output_tokens', - current_message.completion_tokens) - current_message.prompt_tokens = getattr( - final_msg.usage, 'input_tokens', - current_message.prompt_tokens) - - yield current_message + current_message.partial = False + current_message.completion_tokens = getattr( + final_msg.usage, 'output_tokens', + current_message.completion_tokens) + current_message.prompt_tokens = getattr( + final_msg.usage, 'input_tokens', + current_message.prompt_tokens) + + yield current_message + finally: + if self._active_stream is stream: + self._active_stream = None @staticmethod def _format_output_message(completion) -> Message: diff --git a/ms_agent/llm/io.py b/ms_agent/llm/io.py new file mode 100644 index 000000000..f398105f4 --- /dev/null +++ b/ms_agent/llm/io.py @@ -0,0 +1,67 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Keep blocking model I/O out of the application's default executor.""" +import asyncio +import contextvars +import socket +from concurrent.futures import ThreadPoolExecutor + +# Shared, bounded, and started lazily by ThreadPoolExecutor. Its lifetime is +# the process's; the executor joins its workers on process exit. Slow models +# must not occupy the host's executor for session storage and cancellation. +_executor = ThreadPoolExecutor(thread_name_prefix='ms-agent-llm') + + +async def run_in_llm_executor(func, *args): + context = contextvars.copy_context() + return await asyncio.get_running_loop().run_in_executor( + _executor, context.run, func, *args) + + +def interrupt_stream(stream): + """Interrupt an owned HTTP/1 response before closing its SDK stream. + + A socket close in another thread need not wake a blocked recv(). Shutdown + does. Never shut down an HTTP/2 connection, which may carry other requests. + Custom transports without HTTPX's network extension retain normal close. + """ + try: + response = getattr(stream, 'response', None) + if (response is not None and not response.is_closed + and response.http_version in ('HTTP/1.0', 'HTTP/1.1')): + network = response.extensions.get('network_stream') + sock = network.get_extra_info('socket') if network else None + if sock is not None: + sock.shutdown(socket.SHUT_RDWR) + except Exception: # shutdown is best effort; always attempt SDK cleanup + pass + try: + if stream is not None: + stream.close() + except Exception: + pass + + +class OpenedStream: + """Open a lazy stream manager while the request still owns cancellation. + + Anthropic sends the HTTP request in __enter__, not messages.stream(). + Opening it during generate() lets LLMAgent close a late response before + reading any body. The consumer retains the manager's context protocol. + """ + + def __init__(self, manager): + self._manager = manager + self._stream = manager.__enter__() + + def __enter__(self): + return self._stream + + @property + def response(self): + return getattr(self._stream, 'response', None) + + def __exit__(self, *exc): + return self._manager.__exit__(*exc) + + def close(self): + self.__exit__(None, None, None) diff --git a/ms_agent/llm/openai_llm.py b/ms_agent/llm/openai_llm.py index ef1892895..730205e2a 100644 --- a/ms_agent/llm/openai_llm.py +++ b/ms_agent/llm/openai_llm.py @@ -12,6 +12,7 @@ from typing import Any, Dict, Generator, Iterable, List, Optional from ms_agent.llm import LLM, multimodal +from ms_agent.llm.io import interrupt_stream from ms_agent.llm.thinking import apply_effort, create_with_thinking_fallback from ms_agent.llm.utils import Message, Tool, ToolCall from ms_agent.llm.vision import create_with_vision_fallback @@ -101,6 +102,7 @@ def __init__( float(_read_timeout), connect=float(_connect_timeout)), ) self.base_url = base_url or '' + self._active_stream = None # Image attachments (legacy non-router path). Resolution mirrors the # router's: an explicit per-model switch wins, else the service's @@ -318,6 +320,18 @@ def _mark_images_degraded(self, reason: str) -> None: for d in self._last_deliveries ] + def _track_stream(self, stream): + # Retain the HTTP stream, not the lazy thinking/vision wrapper: closing + # an unstarted or currently executing generator cannot close its socket. + self._active_stream = stream + return stream + + def interrupt(self) -> None: + stream = getattr(self, '_active_stream', None) + interrupt_stream(stream) + if getattr(self, '_active_stream', None) is stream: + self._active_stream = None + def _call_llm(self, messages: List[Message], tools: Optional[List[Tool]] = None, @@ -348,17 +362,21 @@ def _call_llm(self, sent_images = any( multimodal.has_image_blocks(m.get('content')) for m in messages if isinstance(m, dict)) + + def create(messages, **params): + result = self.client.chat.completions.create( + model=self.model, messages=messages, tools=tools, **params) + return self._track_stream(result) if is_streaming else result + return create_with_vision_fallback( lambda messages, **kw: create_with_thinking_fallback( - lambda **kw2: self.client.chat.completions.create( - model=self.model, messages=messages, tools=tools, **kw2), - self.client, self.model, logger, **kw), + lambda **kw2: create(messages=messages, **kw2), self.client, + self.model, logger, **kw), base_url=getattr(self.client, 'base_url', ''), model=self.model, messages=messages, sent_images=sent_images, - max_edge=getattr( - getattr(self, '_vision', None), 'max_edge', 0), + max_edge=getattr(getattr(self, '_vision', None), 'max_edge', 0), on_degrade=self._mark_images_degraded, logger_=logger, **kwargs) @@ -995,8 +1013,9 @@ def _responses_stream_generate(self, kwargs['tools'] = resp_tools stream = create_with_thinking_fallback( - lambda **kw: self._responses_client.responses.create( - model=self.model, input=input_items, **kw), + lambda **kw: self._track_stream( + self._responses_client.responses.create( + model=self.model, input=input_items, **kw)), self._responses_client, self.model, logger, diff --git a/ms_agent/llm/transport/anthropic_messages.py b/ms_agent/llm/transport/anthropic_messages.py index 425af6973..e63464987 100644 --- a/ms_agent/llm/transport/anthropic_messages.py +++ b/ms_agent/llm/transport/anthropic_messages.py @@ -17,6 +17,7 @@ from typing import Any, Dict, Generator, Iterator, List, Optional, Union from ms_agent.llm import multimodal +from ms_agent.llm.io import OpenedStream, interrupt_stream from ms_agent.llm.thinking import apply_effort, create_with_thinking_fallback from ms_agent.llm.transport.base import Transport from ms_agent.llm.utils import Message, Tool, ToolCall @@ -349,7 +350,9 @@ def _send(messages, **kw): # than forwarded — the API call has to name it again itself. call['model'] = self.model if stream: - return self.client.messages.stream(**call) + opened = OpenedStream(self.client.messages.stream(**call)) + self._active_stream = opened + return opened return self.client.messages.create(**call) # Thinking is the OTHER per-model hard-400, and this transport used to @@ -507,7 +510,10 @@ def interrupt(self) -> None: from a different thread than the one iterating the stream: closing the underlying HTTP response unblocks that read. A no-op when nothing streams. """ - self._close_stream(self._active_stream) + stream = self._active_stream + interrupt_stream(stream) + if self._active_stream is stream: + self._active_stream = None @staticmethod def _format_output_message(completion) -> Message: diff --git a/ms_agent/llm/transport/openai_compat.py b/ms_agent/llm/transport/openai_compat.py index 1b323f026..4dbe214ee 100644 --- a/ms_agent/llm/transport/openai_compat.py +++ b/ms_agent/llm/transport/openai_compat.py @@ -24,6 +24,7 @@ from typing import Any, Dict, Generator, Iterable, List, Optional, Union from ms_agent.llm import multimodal +from ms_agent.llm.io import interrupt_stream from ms_agent.llm.thinking import apply_effort, create_with_thinking_fallback from ms_agent.llm.transport.base import Transport from ms_agent.llm.utils import Message, Tool, ToolCall @@ -413,7 +414,10 @@ def interrupt(self) -> None: underlying HTTP response unblocks that read (it raises inside ``next()``, which the caller discards). A no-op when nothing is streaming. """ - self._close_stream(self._active_stream) + stream = self._active_stream + interrupt_stream(stream) + if self._active_stream is stream: + self._active_stream = None # ------------------------------------------------------------------ # # inline handling (e.g. MiniMax M-series) @@ -478,19 +482,28 @@ def _call_llm(self, # fallback so the two compose: a request can be retried for thinking # and, independently, for images. sent_images = any( - multimodal.has_image_blocks(m.get('content')) - for m in messages if isinstance(m, dict)) + multimodal.has_image_blocks(m.get('content')) for m in messages + if isinstance(m, dict)) + + def create(messages, **params): + result = self.client.chat.completions.create( + model=self.model, messages=messages, tools=tools, **params) + if is_streaming: + # Keep the actual HTTP stream, before the fallback layers wrap + # it in generators. Closing a wrapper before its first next() + # (or while next() is running) cannot close the connection. + self._active_stream = result + return result + return create_with_vision_fallback( lambda messages, **kw: create_with_thinking_fallback( - lambda **kw2: self.client.chat.completions.create( - model=self.model, messages=messages, tools=tools, **kw2), - self.client, self.model, logger, **kw), + lambda **kw2: create(messages=messages, **kw2), self.client, + self.model, logger, **kw), base_url=getattr(self.client, 'base_url', ''), model=self.model, messages=messages, sent_images=sent_images, - max_edge=getattr( - getattr(self, '_vision', None), 'max_edge', 0), + max_edge=getattr(getattr(self, '_vision', None), 'max_edge', 0), on_degrade=self._mark_images_degraded, logger_=logger, **kwargs) @@ -592,10 +605,8 @@ def _stream_continue_generate(self, **kwargs) -> Generator[Message, None, None]: flag = self._continue_flag message = None - # Track the stream so interrupt() can close it from another thread; a - # continuation rebinds it below. The finally releases it on every exit - # path (normal end, error, or GeneratorExit when the consumer stops). - self._active_stream = completion + # _call_llm tracks the raw HTTP stream for interruption. Also close its + # fallback wrappers on normal end, error, or GeneratorExit. try: for chunk in completion: message_chunk = self._stream_format_output_message(chunk) @@ -633,7 +644,6 @@ def _stream_continue_generate(self, f'continue generate.') completion = self._call_llm_for_continue_gen( messages, message, tools, **kwargs) - self._active_stream = completion for chunk in self._stream_continue_generate( messages, completion, tools, max_runs - 1 if max_runs is not None else None, @@ -650,8 +660,8 @@ def _stream_continue_generate(self, yield message finally: self._close_stream(completion) - if self._active_stream is completion: - self._active_stream = None + self.interrupt() + self._active_stream = None @staticmethod def _stream_format_output_message(completion_chunk) -> Message: diff --git a/tests/agent/test_llm_io.py b/tests/agent/test_llm_io.py new file mode 100644 index 000000000..852806f13 --- /dev/null +++ b/tests/agent/test_llm_io.py @@ -0,0 +1,502 @@ +"""Model I/O must not starve other tasks, including before response headers.""" +import asyncio +import json +import threading +from concurrent.futures import ThreadPoolExecutor +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import httpx +import openai +import pytest +from omegaconf import OmegaConf + +from ms_agent.agent.llm_agent import LLMAgent +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 + + +class ResponseBody(httpx.SyncByteStream): + + def __init__(self, stream, *, wait_for_body=False, protocol='openai'): + self.stream = stream + self.protocol = protocol + self.entered = threading.Event() + self.release = threading.Event() + self.closed = threading.Event() + self.reader_thread = None + if not wait_for_body: + self.release.set() + + def __iter__(self): + self.reader_thread = threading.get_ident() + self.entered.set() + assert self.release.wait(10), 'test did not release response body' + if self.protocol == 'anthropic': + message = { + 'id': 'probe', 'type': 'message', 'role': 'assistant', + 'model': 'probe', 'content': [], 'stop_reason': None, + 'stop_sequence': None, + 'usage': {'input_tokens': 1, 'output_tokens': 1}, + } + if self.stream: + events = [ + {'type': 'message_start', 'message': message}, + {'type': 'content_block_start', 'index': 0, + 'content_block': {'type': 'text', 'text': ''}}, + {'type': 'content_block_delta', 'index': 0, + 'delta': {'type': 'text_delta', 'text': 'OK'}}, + {'type': 'content_block_stop', 'index': 0}, + {'type': 'message_delta', + 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, + 'usage': {'output_tokens': 1}}, + {'type': 'message_stop'}, + ] + for event in events: + yield (f"event: {event['type']}\ndata: " + + json.dumps(event) + '\n\n').encode() + else: + message.update(content=[{'type': 'text', 'text': 'OK'}], + stop_reason='end_turn') + yield json.dumps(message).encode() + return + if self.stream: + for delta, reason in [({'content': 'OK'}, None), ({}, 'stop')]: + chunk = { + 'id': 'probe', 'object': 'chat.completion.chunk', + 'created': 0, 'model': 'probe', + 'choices': [{'index': 0, 'delta': delta, + 'finish_reason': reason}], + } + yield ('data: ' + json.dumps(chunk) + '\n\n').encode() + yield b'data: [DONE]\n\n' + else: + yield json.dumps({ + 'id': 'probe', 'object': 'chat.completion', 'created': 0, + 'model': 'probe', 'choices': [{'index': 0, + 'message': {'role': 'assistant', 'content': 'OK'}, + 'finish_reason': 'stop'}], + 'usage': {'prompt_tokens': 1, 'completion_tokens': 1, + 'total_tokens': 2}, + }).encode() + + def close(self): + self.closed.set() + self.release.set() + + +def make_agent(handler, stream=True, legacy=False, protocol='openai'): + config = OmegaConf.create({ + 'llm': {'model': 'probe'}, 'generation_config': {'stream': stream} + }) + if protocol == 'anthropic': + 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 + else: + llm = OpenAICompatTransport( + model='probe', api_key='local-test', base_url='http://model.test/v1', + generation_config={'stream': stream}) + client_cls = openai.OpenAI + 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))) + agent = LLMAgent.__new__(LLMAgent) + agent.llm = llm + agent.config = OmegaConf.create({ + 'generation_config': {'stream': stream, 'stream_output': False}}) + agent.task_manager = None + agent.load_cache = False + agent._event_sink = None + agent.on_generate_response = AsyncMock() + agent.tool_manager = SimpleNamespace(get_tools=AsyncMock(return_value=[])) + agent.on_tool_call = AsyncMock() + for method in ('log_output', 'handle_new_response', 'save_history', + '_emit_tool_composing', '_record_image_deliveries'): + setattr(agent, method, Mock()) + 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')])] + + +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) + + +@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()) + + +@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()) + + +@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 + loop = asyncio.get_running_loop() + + def generate(*args, **kwargs): + result = original(*args, **kwargs) + loop.call_soon_threadsafe(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()) + + +@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) + iterator_closed = threading.Event() + generate = agent.llm.generate + + def tracked_generate(*args, **kwargs): + chunks = generate(*args, **kwargs) + + def tracked_chunks(): + try: + yield from 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()) + + +@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()) + + +@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 + 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()) + + +@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()) + + +@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 + + release, reading = threading.Event(), threading.Event() + + class Handler(BaseHTTPRequestHandler): + + def log_message(self, *args): + pass + + def do_POST(self): + self.rfile.read(int(self.headers['Content-Length'])) + self.send_response(200) + self.send_header('Content-Type', 'text/event-stream') + self.send_header('Connection', 'close') + self.end_headers() + self.wfile.flush() + release.wait(5) + + class MarkRead(httpx.SyncByteStream): + + def __init__(self, inner): + self.inner = inner + + def __iter__(self): + reading.set() + yield from self.inner + + def close(self): + self.inner.close() + + def mark_response(response): + response.stream = MarkRead(response.stream) + + server = ThreadingHTTPServer(('127.0.0.1', 0), Handler) + 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()) + finally: + release.set() + server.shutdown() + server.server_close() + thread.join(timeout=2) + + +@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}) + if closed: + response.close() + stream = SimpleNamespace(response=response, close=Mock()) + interrupt_stream(stream) + network.get_extra_info.assert_not_called() + stream.close.assert_called_once() From f77644ee8e2129fe62af0422fd31fec7a68542c5 Mon Sep 17 00:00:00 2001 From: alcholiclg Date: Fri, 18 Sep 2026 14:44:49 +0800 Subject: [PATCH 2/2] test(llm): simplify model I/O regression tests --- tests/agent/test_llm_io.py | 448 ++++++++++++++----------------------- 1 file changed, 165 insertions(+), 283 deletions(-) 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()