test: migrate workspace snippet sessions and ORM models to SQLite (#40524)

This commit is contained in:
Asuka Minato 2026-09-08 10:17:08 +00:00 committed by GitHub
parent e4eb3d0d92
commit bcda6b4570
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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)