zcbot/tests/test_llm_stream_cancel.py

64 lines
1.9 KiB
Python

"""流式 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()