from __future__ import annotations import asyncio import pytest from yuxi.channels.services.context import ChatAbortEntry, ChatRunBuffer, GatewayRequestContext class TestChatAbortEntry: def test_abort_cancels_task(self): async def dummy(): await asyncio.sleep(10) loop = asyncio.new_event_loop() task = loop.create_task(dummy()) entry = ChatAbortEntry(task=task) entry.abort() loop.run_until_complete(asyncio.sleep(0)) assert task.cancelled() loop.close() def test_abort_with_none_task_does_not_raise(self): entry = ChatAbortEntry(task=None) entry.abort() class TestChatRunBuffer: def test_append_and_get_full_text(self): buffer = ChatRunBuffer(run_id="test-1") buffer.append_chunk("Hello ") buffer.append_chunk("World") assert buffer.get_full_text() == "Hello World" def test_is_finished_defaults_false(self): buffer = ChatRunBuffer(run_id="test-2") assert not buffer.is_finished() def test_mark_finished(self): buffer = ChatRunBuffer(run_id="test-3") buffer.mark_finished() assert buffer.is_finished() def test_empty_buffer_returns_empty_string(self): buffer = ChatRunBuffer(run_id="test-4") assert buffer.get_full_text() == "" class TestGatewayRequestContext: def test_default_values(self): ctx = GatewayRequestContext() assert ctx.runtime_config == {} assert ctx.start_channel is None assert ctx.stop_channel is None assert ctx.broadcast_fn is None assert ctx.node_send_to_session_fn is None assert ctx.get_snapshot is None assert ctx.chat_abort_controllers == {} assert ctx.chat_run_buffers == {} assert ctx.runtime_snapshot is None assert ctx.dedupe is None assert ctx.mark_channel_logged_out is None def test_get_runtime_config(self): ctx = GatewayRequestContext(runtime_config={"key": "value"}) assert ctx.get_runtime_config("key") == "value" assert ctx.get_runtime_config("missing", "default") == "default" def test_set_runtime_config(self): ctx = GatewayRequestContext() ctx.set_runtime_config("new_key", "new_value") assert ctx.get_runtime_config("new_key") == "new_value" def test_chat_abort_controllers_mutable(self): ctx = GatewayRequestContext() entry = ChatAbortEntry() ctx.chat_abort_controllers["run-1"] = entry assert ctx.chat_abort_controllers["run-1"] is entry def test_chat_run_buffers_mutable(self): ctx = GatewayRequestContext() buffer = ChatRunBuffer(run_id="run-1") ctx.chat_run_buffers["run-1"] = buffer assert ctx.chat_run_buffers["run-1"] is buffer def test_add_chat_run(self): ctx = GatewayRequestContext() entry = ChatAbortEntry() buffer = ChatRunBuffer(run_id="run-1") ctx.add_chat_run("run-1", entry, buffer) assert ctx.chat_abort_controllers["run-1"] is entry assert ctx.chat_run_buffers["run-1"] is buffer def test_remove_chat_run(self): ctx = GatewayRequestContext() entry = ChatAbortEntry() ctx.chat_abort_controllers["run-1"] = entry ctx.chat_run_buffers["run-1"] = ChatRunBuffer(run_id="run-1") ctx.remove_chat_run("run-1") assert "run-1" not in ctx.chat_abort_controllers assert "run-1" not in ctx.chat_run_buffers def test_abort_chat_run(self): ctx = GatewayRequestContext() entry = ChatAbortEntry() ctx.chat_abort_controllers["run-1"] = entry assert ctx.abort_chat_run("run-1") is True assert entry.is_aborted() def test_abort_chat_run_unknown(self): ctx = GatewayRequestContext() assert ctx.abort_chat_run("nonexistent") is False def test_get_chat_buffer(self): ctx = GatewayRequestContext() buffer = ChatRunBuffer(run_id="run-1") ctx.chat_run_buffers["run-1"] = buffer assert ctx.get_chat_buffer("run-1") is buffer def test_get_chat_buffer_unknown(self): ctx = GatewayRequestContext() assert ctx.get_chat_buffer("nonexistent") is None @pytest.mark.asyncio async def test_callbacks_are_awaitable(self): called = False async def demo_start(channel_id: str, config: dict): nonlocal called called = True ctx = GatewayRequestContext(start_channel=demo_start) await ctx.start_channel("test", {}) assert called @pytest.mark.asyncio async def test_broadcast_method_delegates(self): received = [] async def mock_broadcast(event: str, payload): received.append((event, payload)) ctx = GatewayRequestContext(broadcast_fn=mock_broadcast) await ctx.broadcast("test.event", {"data": 1}) assert received == [("test.event", {"data": 1})] @pytest.mark.asyncio async def test_broadcast_without_fn_does_not_raise(self): ctx = GatewayRequestContext() await ctx.broadcast("test.event", {}) await ctx.node_send_to_session("s1", "test", {})