feat: stream provider responses
This commit is contained in:
+60
-1
@@ -19,6 +19,22 @@ class _FakeProvider:
|
||||
def chat(self, messages: list[ChatMessage]) -> str:
|
||||
return f"Echo: {messages[-1].content}"
|
||||
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
yield "Echo: "
|
||||
yield messages[-1].content
|
||||
|
||||
|
||||
class _MarkdownProvider(_FakeProvider):
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
yield "## **Ant"
|
||||
yield "wort**\n`SQLite`"
|
||||
|
||||
|
||||
class _InterruptingProvider(_FakeProvider):
|
||||
def stream_chat(self, messages: list[ChatMessage]):
|
||||
yield "angefangene Antwort"
|
||||
raise KeyboardInterrupt
|
||||
|
||||
|
||||
class _FakeHybridProvider(_FakeProvider):
|
||||
name = "hybrid"
|
||||
@@ -59,10 +75,53 @@ class CliTests(unittest.TestCase):
|
||||
)
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertTrue(any(line == "Javis: Echo: Hallo" for line in output))
|
||||
self.assertIn("Javis: Echo: Hallo\n", "".join(output))
|
||||
self.assertTrue(any("Nachrichten" in line for line in output))
|
||||
self.assertEqual(output[-1], "Chat beendet.")
|
||||
|
||||
def test_streaming_cleans_terminal_markdown_but_persists_raw_answer(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
service = ChatService(store, _MarkdownProvider())
|
||||
inputs = iter(["Frage", "/exit"])
|
||||
output: list[str] = []
|
||||
streamed: list[str] = []
|
||||
|
||||
result = run_chat(
|
||||
service,
|
||||
input_fn=lambda _prompt: next(inputs),
|
||||
output=output.append,
|
||||
stream_output=streamed.append,
|
||||
)
|
||||
session = service.list_sessions()[0]
|
||||
messages = service.load_session(session.id).messages
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertEqual("".join(streamed), "Javis: Antwort\nSQLite\n")
|
||||
self.assertEqual(messages[-1].content, "## **Antwort**\n`SQLite`")
|
||||
|
||||
def test_keyboard_interrupt_drops_partial_answer_and_returns_to_prompt(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
service = ChatService(store, _InterruptingProvider())
|
||||
inputs = iter(["Frage", "/exit"])
|
||||
output: list[str] = []
|
||||
streamed: list[str] = []
|
||||
|
||||
result = run_chat(
|
||||
service,
|
||||
input_fn=lambda _prompt: next(inputs),
|
||||
output=output.append,
|
||||
stream_output=streamed.append,
|
||||
)
|
||||
session = service.list_sessions()[0]
|
||||
messages = service.load_session(session.id).messages
|
||||
|
||||
self.assertEqual(result, 0)
|
||||
self.assertEqual(messages, [])
|
||||
self.assertTrue(any("nicht gespeichert" in line for line in output))
|
||||
self.assertEqual(output[-1], "Chat beendet.")
|
||||
|
||||
def test_provider_privacy_and_status_commands_are_sanitized(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
store = SQLiteSessionStore(Path(directory) / "sessions.sqlite3")
|
||||
|
||||
Reference in New Issue
Block a user