diff --git a/src/exo/api/main.py b/src/exo/api/main.py index 5635a409..cf9bef8c 100644 --- a/src/exo/api/main.py +++ b/src/exo/api/main.py @@ -169,10 +169,20 @@ from exo.shared.types.events import ( ChunkGenerated, Event, IndexedEvent, + InstanceDeleted, TracesMerged, ) from exo.shared.types.memory import Memory from exo.shared.types.state import State +from exo.shared.types.tasks import ( + ImageEdits as ImageEditsTask, +) +from exo.shared.types.tasks import ( + ImageGeneration as ImageGenerationTask, +) +from exo.shared.types.tasks import ( + TextGeneration as TextGenerationTask, +) from exo.shared.types.text_generation import TextGenerationTaskParams from exo.shared.types.worker.downloads import DownloadCompleted from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta @@ -1814,9 +1824,25 @@ class API: await queue.send(event.chunk) except BrokenResourceError: self._text_generation_queues.pop(event.command_id, None) + if isinstance(event, InstanceDeleted): + self._close_streams_for_instance(event.instance_id) if isinstance(event, TracesMerged): self._save_merged_trace(event) + def _close_streams_for_instance(self, instance_id: InstanceId) -> None: + """Close any active generation streams for commands running on the given instance.""" + for task in self.state.tasks.values(): + if task.instance_id != instance_id: + continue + if not isinstance( + task, (TextGenerationTask, ImageGenerationTask, ImageEditsTask) + ): + continue + if sender := self._text_generation_queues.pop(task.command_id, None): + sender.close() + if sender := self._image_generation_queues.pop(task.command_id, None): + sender.close() + def _save_merged_trace(self, event: TracesMerged) -> None: traces = [ TraceEvent( diff --git a/src/exo/api/tests/test_instance_deleted_stream_cleanup.py b/src/exo/api/tests/test_instance_deleted_stream_cleanup.py new file mode 100644 index 00000000..843554b9 --- /dev/null +++ b/src/exo/api/tests/test_instance_deleted_stream_cleanup.py @@ -0,0 +1,93 @@ +# pyright: reportUnusedFunction=false, reportAny=false +"""Tests that InstanceDeleted events close active generation streams.""" + +from unittest.mock import MagicMock + +from exo.api.main import API +from exo.api.types import ImageGenerationTaskParams +from exo.shared.types.common import CommandId, ModelId +from exo.shared.types.state import State +from exo.shared.types.tasks import ImageGeneration, TextGeneration +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams +from exo.shared.types.worker.instances import InstanceId + + +def _make_api_with_state(state: State) -> API: + """Create a minimal API instance with pre-set state.""" + api = object.__new__(API) + api.state = state + api._text_generation_queues = {} # pyright: ignore[reportPrivateUsage] + api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage] + return api + + +def _make_text_gen_task( + instance_id: InstanceId, command_id: CommandId +) -> TextGeneration: + return TextGeneration( + instance_id=instance_id, + command_id=command_id, + task_params=TextGenerationTaskParams( + model=ModelId("test-model"), + input=[InputMessage(role="user", content="hello")], + ), + ) + + +def test_close_streams_for_deleted_instance() -> None: + """Deleting an instance closes the text generation sender for commands on that instance.""" + instance_id = InstanceId("inst-1") + command_id = CommandId("cmd-1") + task = _make_text_gen_task(instance_id, command_id) + + state = State(tasks={task.task_id: task}) + api = _make_api_with_state(state) + + sender = MagicMock() + api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage] + + api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage] + + sender.close.assert_called_once() + assert command_id not in api._text_generation_queues # pyright: ignore[reportPrivateUsage] + + +def test_close_streams_ignores_unrelated_instances() -> None: + """Deleting an instance does NOT close streams for commands on other instances.""" + target_id = InstanceId("inst-delete") + other_id = InstanceId("inst-keep") + other_cmd = CommandId("cmd-keep") + other_task = _make_text_gen_task(other_id, other_cmd) + + state = State(tasks={other_task.task_id: other_task}) + api = _make_api_with_state(state) + + sender = MagicMock() + api._text_generation_queues[other_cmd] = sender # pyright: ignore[reportPrivateUsage] + + api._close_streams_for_instance(target_id) # pyright: ignore[reportPrivateUsage] + + sender.close.assert_not_called() + assert other_cmd in api._text_generation_queues # pyright: ignore[reportPrivateUsage] + + +def test_close_streams_for_deleted_instance_image_generation() -> None: + """Deleting an instance closes the image generation sender for commands on that instance.""" + instance_id = InstanceId("inst-img") + command_id = CommandId("cmd-img") + task = ImageGeneration( + instance_id=instance_id, + command_id=command_id, + task_params=ImageGenerationTaskParams(prompt="a cat", model="test-model"), + ) + + state = State(tasks={task.task_id: task}) + api = _make_api_with_state(state) + + sender = MagicMock() + api._image_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage] + + api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage] + + sender.close.assert_called_once() + assert command_id not in api._image_generation_queues # pyright: ignore[reportPrivateUsage]