diff --git a/api/tests/unit_tests/controllers/console/workspace/test_snippets.py b/api/tests/unit_tests/controllers/console/workspace/test_snippets.py index d664fb3eb9a..75dc6180e90 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_snippets.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_snippets.py @@ -6,7 +6,8 @@ from unittest.mock import Mock import pytest from flask import Flask from pydantic import ValidationError -from sqlalchemy.orm import Session +from sqlalchemy import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker from werkzeug.exceptions import BadRequest, NotFound from controllers.console.workspace import snippets as snippets_module @@ -16,23 +17,19 @@ from services.snippet_dsl_service import ImportStatus, SnippetImportInfo @pytest.fixture(autouse=True) -def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch): +def _patch_snippet_service_factory( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +): def factory(): return snippets_module.SnippetService.__new__(snippets_module.SnippetService) + database_session = scoped_session(sqlite_session_factory) monkeypatch.setattr(snippets_module, "_snippet_service", factory) - - -class _SessionContext: - def __init__(self, engine, *args, **kwargs): - self.engine = engine - self.session = kwargs.pop("session", None) - - def __enter__(self): - return self.session - - def __exit__(self, exc_type, exc, tb): - return False + monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=sqlite_engine, session=database_session)) + yield + database_session.remove() def _account(account_id: str = "account-1") -> Account: @@ -440,17 +437,10 @@ def _persisted_name(session: Session) -> str | None: def test_delete_snippet_deletes_and_commits(app: Flask, monkeypatch: pytest.MonkeyPatch): snippet = _snippet() user = _account() - session = SimpleNamespace(merge=Mock(return_value=snippet), commit=Mock()) delete_snippet = Mock() - class SessionContext(_SessionContext): - def __init__(self, engine, *args, **kwargs): - super().__init__(engine, *args, session=session, **kwargs) - monkeypatch.setattr(snippets_module.SnippetService, "get_snippet_by_id", Mock(return_value=snippet)) monkeypatch.setattr(snippets_module.SnippetService, "delete_snippet", delete_snippet) - monkeypatch.setattr(snippets_module, "Session", SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) api = snippets_module.CustomizedSnippetDetailApi() handler = unwrap(api.delete) @@ -460,18 +450,14 @@ def test_delete_snippet_deletes_and_commits(app: Flask, monkeypatch: pytest.Monk assert status_code == 204 assert response == "" - delete_snippet.assert_called_once_with(session=session, snippet=snippet, account_id=user.id) - session.commit.assert_called_once() + assert delete_snippet.call_args.kwargs["account_id"] == user.id + assert isinstance(delete_snippet.call_args.kwargs["session"], Session) + assert isinstance(delete_snippet.call_args.kwargs["snippet"], CustomizedSnippet) def test_export_snippet_returns_yaml_attachment(app: Flask, monkeypatch: pytest.MonkeyPatch): snippet = _snippet(name="Snippet One") export_snippet_dsl = Mock(return_value="version: 0.1.0\nkind: snippet\n") - session = SimpleNamespace() - - class SessionContext(_SessionContext): - def __init__(self, engine, *args, **kwargs): - super().__init__(engine, *args, session=session, **kwargs) monkeypatch.setattr(snippets_module.SnippetService, "get_snippet_by_id", Mock(return_value=snippet)) monkeypatch.setattr( @@ -479,8 +465,6 @@ def test_export_snippet_returns_yaml_attachment(app: Flask, monkeypatch: pytest. "SnippetDslService", Mock(return_value=SimpleNamespace(export_snippet_dsl=export_snippet_dsl)), ) - monkeypatch.setattr(snippets_module, "Session", SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) api = snippets_module.CustomizedSnippetExportApi() handler = unwrap(api.get) @@ -510,9 +494,6 @@ def test_export_snippet_raises_not_found_for_missing_workflow(app: Flask, monkey ) ), ) - monkeypatch.setattr(snippets_module, "Session", _SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) - api = snippets_module.CustomizedSnippetExportApi() handler = unwrap(api.get) @@ -521,24 +502,12 @@ def test_export_snippet_raises_not_found_for_missing_workflow(app: Flask, monkey handler(api, "tenant-1", snippet_id="snippet-1") -def test_import_snippet_returns_202_for_pending_confirmation(app: Flask, monkeypatch: pytest.MonkeyPatch): +def test_import_snippet_returns_202_for_pending_confirmation( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): user = _account("account-1") result = SnippetImportInfo(id="import-1", status=ImportStatus.PENDING, imported_dsl_version="999.0.0") import_snippet = Mock(return_value=result) - session = SimpleNamespace(commit=Mock(), rollback=Mock()) - - class _SessionContext: - def __init__(self, engine): - self.engine = engine - - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(snippets_module, "Session", _SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) monkeypatch.setattr( snippets_module, "SnippetDslService", @@ -555,27 +524,19 @@ def test_import_snippet_returns_202_for_pending_confirmation(app: Flask, monkeyp method="POST", json={"mode": "yaml-content", "yaml_content": "kind: snippet"}, ): - response, status_code = handler(api, req_data, session, user) + response, status_code = handler(api, req_data, sqlite_session, user) assert status_code == 202 assert response["status"] == ImportStatus.PENDING.value import_snippet.assert_called_once() - session.rollback.assert_not_called() - session.commit.assert_not_called() -def test_import_snippet_returns_400_for_failed_import(app: Flask, monkeypatch: pytest.MonkeyPatch): +def test_import_snippet_returns_400_for_failed_import( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): user = _account("account-1") result = SnippetImportInfo(id="import-1", status=ImportStatus.FAILED, error="Invalid DSL") import_snippet = Mock(return_value=result) - session = SimpleNamespace(commit=Mock(), rollback=Mock()) - - class SessionContext(_SessionContext): - def __init__(self, engine, *args, **kwargs): - super().__init__(engine, *args, session=session, **kwargs) - - monkeypatch.setattr(snippets_module, "Session", SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) monkeypatch.setattr( snippets_module, "SnippetDslService", @@ -592,26 +553,18 @@ def test_import_snippet_returns_400_for_failed_import(app: Flask, monkeypatch: p method="POST", json={"mode": "yaml-content", "yaml_content": "kind: snippet"}, ): - response, status_code = handler(api, req_data, session, user) + response, status_code = handler(api, req_data, sqlite_session, user) assert status_code == 400 assert response["error"] == "Invalid DSL" - session.rollback.assert_not_called() - session.commit.assert_not_called() -def test_import_confirm_returns_200_for_completed_import(app: Flask, monkeypatch: pytest.MonkeyPatch): +def test_import_confirm_returns_200_for_completed_import( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): user = _account("account-1") result = SnippetImportInfo(id="import-1", status=ImportStatus.COMPLETED, snippet_id="snippet-1") confirm_import = Mock(return_value=result) - session = SimpleNamespace(commit=Mock(), rollback=Mock()) - - class SessionContext(_SessionContext): - def __init__(self, engine, *args, **kwargs): - super().__init__(engine, *args, session=session, **kwargs) - - monkeypatch.setattr(snippets_module, "Session", SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) monkeypatch.setattr( snippets_module, "SnippetDslService", @@ -625,12 +578,11 @@ def test_import_confirm_returns_200_for_completed_import(app: Flask, monkeypatch "/workspaces/current/customized-snippets/imports/import-1/confirm", method="POST", ): - response, status_code = handler(api, session, user, import_id="import-1") + response, status_code = handler(api, sqlite_session, user, import_id="import-1") assert status_code == 200 assert response["snippet_id"] == "snippet-1" confirm_import.assert_called_once_with(import_id="import-1", account=user) - session.commit.assert_not_called() def test_check_dependencies_raises_when_snippet_missing(app: Flask, monkeypatch: pytest.MonkeyPatch): @@ -647,15 +599,8 @@ def test_check_dependencies_raises_when_snippet_missing(app: Flask, monkeypatch: def test_check_dependencies_returns_dependency_result(app: Flask, monkeypatch: pytest.MonkeyPatch): snippet = _snippet() check_dependencies = Mock(return_value=SimpleNamespace(model_dump=Mock(return_value={"leaked_dependencies": []}))) - session = SimpleNamespace() - - class SessionContext(_SessionContext): - def __init__(self, engine, *args, **kwargs): - super().__init__(engine, *args, session=session, **kwargs) monkeypatch.setattr(snippets_module.SnippetService, "get_snippet_by_id", Mock(return_value=snippet)) - monkeypatch.setattr(snippets_module, "Session", SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) monkeypatch.setattr( snippets_module, "SnippetDslService", @@ -688,25 +633,15 @@ def test_increment_use_count_raises_when_snippet_missing(app: Flask, monkeypatch def test_increment_use_count_returns_refreshed_count(app: Flask, monkeypatch: pytest.MonkeyPatch): - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1", use_count=2) - merged_snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1", use_count=3) - session = SimpleNamespace(merge=Mock(return_value=merged_snippet), commit=Mock(), refresh=Mock()) + snippet = _snippet(use_count=2) - class _SessionContext: - def __init__(self, engine): - self.engine = engine + def increment_use_count(*, session: Session, snippet: CustomizedSnippet) -> None: + assert isinstance(session, Session) + snippet.use_count += 1 - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - increment_use_count = Mock() + increment_use_count_mock = Mock(side_effect=increment_use_count) monkeypatch.setattr(snippets_module.SnippetService, "get_snippet_by_id", Mock(return_value=snippet)) - monkeypatch.setattr(snippets_module.SnippetService, "increment_use_count", increment_use_count) - monkeypatch.setattr(snippets_module, "Session", _SessionContext) - monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(snippets_module.SnippetService, "increment_use_count", increment_use_count_mock) api = snippets_module.CustomizedSnippetUseCountIncrementApi() handler = unwrap(api.post) @@ -719,6 +654,5 @@ def test_increment_use_count_returns_refreshed_count(app: Flask, monkeypatch: py assert status_code == 200 assert response == {"result": "success", "use_count": 3} - increment_use_count.assert_called_once_with(session=session, snippet=merged_snippet) - session.commit.assert_called_once() - session.refresh.assert_called_once_with(merged_snippet) + assert isinstance(increment_use_count_mock.call_args.kwargs["session"], Session) + assert isinstance(increment_use_count_mock.call_args.kwargs["snippet"], CustomizedSnippet)