64 lines
1.9 KiB
Python
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()
|