"""流式 LLM 在 provider 无新分片时仍能快速响应用户停止。""" from __future__ import annotations import threading import time import unittest from unittest.mock import patch from core.llm import LLM class _BlockingStream: def __init__(self) -> None: self.reading = threading.Event() self.released = threading.Event() self.closed = False def __iter__(self): return self def __next__(self): self.reading.set() self.released.wait(30) raise StopIteration def close(self) -> None: self.closed = True self.released.set() class LlmStreamCancelTests(unittest.TestCase): def test_interruptible_stream_preserves_normal_chunks(self) -> None: llm = object.__new__(LLM) llm._build_kwargs = lambda *args, **kwargs: {"model": "test"} with patch("core.llm.litellm.completion", return_value=iter(["a", "b"])): chunks = list(llm.chat_stream([], cancel_check=lambda: False)) self.assertEqual(chunks, ["a", "b"]) def test_cancel_does_not_wait_for_next_provider_chunk(self) -> None: llm = object.__new__(LLM) llm._build_kwargs = lambda *args, **kwargs: {"model": "test"} raw = _BlockingStream() cancelled = threading.Event() def trigger_cancel() -> None: self.assertTrue(raw.reading.wait(2)) cancelled.set() trigger = threading.Thread(target=trigger_cancel, daemon=True) trigger.start() started = time.monotonic() with patch("core.llm.litellm.completion", return_value=raw): chunks = list(llm.chat_stream([], cancel_check=cancelled.is_set)) elapsed = time.monotonic() - started trigger.join(2) self.assertEqual(chunks, []) self.assertTrue(raw.closed) self.assertLess(elapsed, 1.0) if __name__ == "__main__": unittest.main()