diff --git a/features/steps/database_handler_crud_steps.py b/features/steps/database_handler_crud_steps.py index 5c6141424..a1ed08ffa 100644 --- a/features/steps/database_handler_crud_steps.py +++ b/features/steps/database_handler_crud_steps.py @@ -406,6 +406,7 @@ def step_then_write_message_contains(context: Context, text: str) -> None: ) +@then('the SQLite table "{table_name}" should contain {count:d} rows') @then('the SQLite table "{table_name}" should contain {count:d} row') def step_then_sqlite_table_row_count( context: Context, table_name: str, count: int diff --git a/features/steps/tui_session_export_import_steps.py b/features/steps/tui_session_export_import_steps.py index 229b0a79b..95e08eb8d 100644 --- a/features/steps/tui_session_export_import_steps.py +++ b/features/steps/tui_session_export_import_steps.py @@ -14,7 +14,7 @@ import tempfile from dataclasses import dataclass, field from datetime import datetime from pathlib import Path -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock from behave import given, then, when from behave.runner import Context @@ -268,10 +268,15 @@ def step_exported_md_file_contains(context: Context, text: str) -> None: def _make_tui_router_with_export_mock( session_id: str, - export_data: dict, + export_data: dict[str, object], session_obj: Session, -) -> tuple[TuiCommandRouter, MagicMock]: - """Build a TuiCommandRouter with a mocked container/service for export.""" +) -> TuiCommandRouter: + """Build a TuiCommandRouter with a mocked container/service for export. + + The mock container is injected via the ``container_factory`` constructor + parameter so the mock survives ``multiprocessing.fork()`` boundaries + used by the parallel test runner. + """ mock_service = MagicMock() mock_service.export_session.return_value = export_data mock_service.get.return_value = session_obj @@ -281,8 +286,12 @@ def _make_tui_router_with_export_mock( registry = FakePersonaRegistry() state = FakePersonaState() - router = TuiCommandRouter(persona_registry=registry, persona_state=state) - return router, mock_container + router = TuiCommandRouter( + persona_registry=registry, + persona_state=state, + container_factory=lambda: mock_container, + ) + return router @given("a TUI command router with mocked session service for export") @@ -297,11 +306,9 @@ def step_tui_router_for_export(context: Context) -> None: session = _make_session(session_id=session_id, messages=msgs) export_data = session.as_export_dict() - router, mock_container = _make_tui_router_with_export_mock( + context.tui_router = _make_tui_router_with_export_mock( session_id, export_data, session ) - context.tui_router = router - context.tui_mock_container = mock_container @given("a TUI command router with mocked session service for import") @@ -318,9 +325,12 @@ def step_tui_router_for_import(context: Context) -> None: registry = FakePersonaRegistry() state = FakePersonaState() - router = TuiCommandRouter(persona_registry=registry, persona_state=state) + router = TuiCommandRouter( + persona_registry=registry, + persona_state=state, + container_factory=lambda: mock_container, + ) context.tui_router = router - context.tui_mock_container = mock_container context.tui_imported_session = imported_session @@ -346,37 +356,25 @@ def step_invalid_json_for_tui(context: Context) -> None: @when('I call TUI handle with "{command}" for the current session') def step_tui_handle_command(context: Context, command: str) -> None: - with patch( - "cleveragents.tui.commands.get_container", - return_value=context.tui_mock_container, - ): - context.tui_result = context.tui_router.handle( - command, session_id=context.tui_session_id - ) + context.tui_result = context.tui_router.handle( + command, session_id=context.tui_session_id + ) @when('I call TUI handle with "session import " for the import file') def step_tui_handle_import_valid(context: Context) -> None: command = f"session import {context.tui_import_path}" - with patch( - "cleveragents.tui.commands.get_container", - return_value=context.tui_mock_container, - ): - context.tui_result = context.tui_router.handle( - command, session_id=context.tui_session_id - ) + context.tui_result = context.tui_router.handle( + command, session_id=context.tui_session_id + ) @when('I call TUI handle with "session import " for the import file') def step_tui_handle_import_invalid(context: Context) -> None: command = f"session import {context.tui_invalid_json_path}" - with patch( - "cleveragents.tui.commands.get_container", - return_value=context.tui_mock_container, - ): - context.tui_result = context.tui_router.handle( - command, session_id=context.tui_session_id - ) + context.tui_result = context.tui_router.handle( + command, session_id=context.tui_session_id + ) @then('the TUI handle result should contain "{text}"') diff --git a/src/cleveragents/tui/commands.py b/src/cleveragents/tui/commands.py index b205d0d4a..34d4d22bb 100644 --- a/src/cleveragents/tui/commands.py +++ b/src/cleveragents/tui/commands.py @@ -3,8 +3,10 @@ from __future__ import annotations import json -from dataclasses import dataclass +from collections.abc import Callable +from dataclasses import dataclass, field from pathlib import Path +from typing import Any from cleveragents.application.container import get_container from cleveragents.tui.app import CleverAgentsTuiApp, textual_available @@ -14,10 +16,32 @@ from cleveragents.tui.persona.state import PersonaState @dataclass(slots=True) class TuiCommandRouter: - """Slash command router used by the TUI prompt.""" + """Slash command router used by the TUI prompt. + + Parameters + ---------- + persona_registry: + Registry of available personas. + persona_state: + Mutable state tracking the active persona per session. + container_factory: + Optional callable that returns the DI container. Defaults to + :func:`get_container` when *None*. Accepting this parameter + enables constructor-based dependency injection, which is + essential for test isolation under ``multiprocessing.fork()`` + where ``unittest.mock.patch`` context managers do not propagate + to child workers. + """ persona_registry: PersonaRegistry persona_state: PersonaState + container_factory: Callable[[], Any] | None = field(default=None, repr=False) + + def _resolve_container(self) -> Any: + """Return the DI container, using the injected factory or the global default.""" + if self.container_factory is not None: + return self.container_factory() + return get_container() def handle(self, raw: str, *, session_id: str) -> str: tokens = raw.strip().split() @@ -75,7 +99,7 @@ class TuiCommandRouter: return f"Invalid format: {fmt!r}. Use 'json' or 'md'." try: - container = get_container() + container = self._resolve_container() service = container.session_service() export_data = service.export_session(session_id) @@ -133,7 +157,7 @@ class TuiCommandRouter: return f"Invalid JSON: {exc}" try: - container = get_container() + container = self._resolve_container() service = container.session_service() session = service.import_session(data) return f"Session imported: {session.session_id}"