test: use sqlite3 session in test_file_factory (#38758)

This commit is contained in:
Asuka Minato 2026-07-14 15:37:25 +09:00 committed by GitHub
parent 67cce8eb86
commit 7da854555e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,11 +1,18 @@
import re import re
from datetime import UTC, datetime
from unittest.mock import MagicMock from unittest.mock import MagicMock
from uuid import uuid4
import pytest import pytest
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
from extensions.storage.storage_type import StorageType
from factories.file_factory import builders from factories.file_factory import builders
from factories.file_factory.remote import extract_filename, get_remote_file_info from factories.file_factory.remote import extract_filename, get_remote_file_info
from graphon.file import FileTransferMethod from graphon.file import FileTransferMethod
from models import UploadFile
from models.enums import CreatorUserRole
class _FakeResponse: class _FakeResponse:
@ -298,30 +305,36 @@ class TestExtractFilename:
assert result == "file.txt" assert result == "file.txt"
@pytest.mark.parametrize("sqlite_session", [(UploadFile,)], indirect=True)
class TestBuildFromDatasourceFile: class TestBuildFromDatasourceFile:
"""Tests for _build_from_datasource_file extension handling.""" """Tests for _build_from_datasource_file extension handling."""
@staticmethod @staticmethod
def _patch_session(monkeypatch: pytest.MonkeyPatch, datasource_file): def _bind_session_factory(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
"""Stub session_factory.create_session() so it returns the given UploadFile-shaped record.""" """Bind builder-owned sessions to the test database."""
session = MagicMock() factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
session.scalar.return_value = datasource_file monkeypatch.setattr(builders.session_factory, "create_session", factory)
ctx = MagicMock()
ctx.__enter__ = MagicMock(return_value=session)
ctx.__exit__ = MagicMock(return_value=False)
monkeypatch.setattr(builders.session_factory, "create_session", lambda: ctx)
def _make_datasource_file(self, *, key: str, mime_type: str = "text/csv"): @staticmethod
f = MagicMock() def _make_datasource_file(*, key: str, mime_type: str = "text/csv") -> UploadFile:
f.id = "file-id" return UploadFile(
f.key = key tenant_id=str(uuid4()),
f.name = key.split("/")[-1] storage_type=StorageType.LOCAL,
f.mime_type = mime_type key=key,
f.size = 123 name=key.split("/")[-1],
f.source_url = f"https://example.com/{key}" size=123,
return f extension=key.rsplit(".", maxsplit=1)[-1] if "." in key else "",
mime_type=mime_type,
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
created_at=datetime.now(UTC),
used=False,
source_url=f"https://example.com/{key}",
)
def test_extension_passed_without_doubled_dot(self, monkeypatch: pytest.MonkeyPatch): def test_extension_passed_without_doubled_dot(
self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
):
"""Regression: standardize_file_type must receive the extension exactly once-prefixed. """Regression: standardize_file_type must receive the extension exactly once-prefixed.
Previously the call was ``standardize_file_type(extension="." + extension, ...)`` while Previously the call was ``standardize_file_type(extension="." + extension, ...)`` while
@ -341,14 +354,16 @@ class TestBuildFromDatasourceFile:
monkeypatch.setattr(builders, "standardize_file_type", fake_standardize) monkeypatch.setattr(builders, "standardize_file_type", fake_standardize)
datasource_file = self._make_datasource_file(key="folder/data.csv", mime_type="text/csv") datasource_file = self._make_datasource_file(key="folder/data.csv", mime_type="text/csv")
self._patch_session(monkeypatch, datasource_file) sqlite_session.add(datasource_file)
sqlite_session.commit()
self._bind_session_factory(monkeypatch, sqlite_engine)
access_controller = MagicMock() access_controller = MagicMock()
access_controller.apply_upload_file_filters = lambda stmt: stmt access_controller.apply_upload_file_filters = lambda stmt: stmt
file = builders._build_from_datasource_file( file = builders._build_from_datasource_file(
mapping={"datasource_file_id": "file-id", "transfer_method": "datasource_file"}, mapping={"datasource_file_id": datasource_file.id, "transfer_method": "datasource_file"},
tenant_id="tenant-id", tenant_id=datasource_file.tenant_id,
transfer_method=FileTransferMethod.DATASOURCE_FILE, transfer_method=FileTransferMethod.DATASOURCE_FILE,
access_controller=access_controller, access_controller=access_controller,
) )
@ -360,7 +375,9 @@ class TestBuildFromDatasourceFile:
assert file.extension == ".csv" assert file.extension == ".csv"
assert file.transfer_method == FileTransferMethod.DATASOURCE_FILE assert file.transfer_method == FileTransferMethod.DATASOURCE_FILE
def test_extension_falls_back_to_bin_when_key_has_no_dot(self, monkeypatch: pytest.MonkeyPatch): def test_extension_falls_back_to_bin_when_key_has_no_dot(
self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
):
captured: dict = {} captured: dict = {}
def fake_standardize(*, extension: str = "", mime_type: str = ""): def fake_standardize(*, extension: str = "", mime_type: str = ""):
@ -372,14 +389,16 @@ class TestBuildFromDatasourceFile:
monkeypatch.setattr(builders, "standardize_file_type", fake_standardize) monkeypatch.setattr(builders, "standardize_file_type", fake_standardize)
datasource_file = self._make_datasource_file(key="dotless-key", mime_type="application/octet-stream") datasource_file = self._make_datasource_file(key="dotless-key", mime_type="application/octet-stream")
self._patch_session(monkeypatch, datasource_file) sqlite_session.add(datasource_file)
sqlite_session.commit()
self._bind_session_factory(monkeypatch, sqlite_engine)
access_controller = MagicMock() access_controller = MagicMock()
access_controller.apply_upload_file_filters = lambda stmt: stmt access_controller.apply_upload_file_filters = lambda stmt: stmt
file = builders._build_from_datasource_file( file = builders._build_from_datasource_file(
mapping={"datasource_file_id": "file-id", "transfer_method": "datasource_file"}, mapping={"datasource_file_id": datasource_file.id, "transfer_method": "datasource_file"},
tenant_id="tenant-id", tenant_id=datasource_file.tenant_id,
transfer_method=FileTransferMethod.DATASOURCE_FILE, transfer_method=FileTransferMethod.DATASOURCE_FILE,
access_controller=access_controller, access_controller=access_controller,
) )
@ -387,3 +406,22 @@ class TestBuildFromDatasourceFile:
assert captured["extension"] == ".bin" assert captured["extension"] == ".bin"
assert file.extension == ".bin" assert file.extension == ".bin"
assert file.transfer_method == FileTransferMethod.DATASOURCE_FILE assert file.transfer_method == FileTransferMethod.DATASOURCE_FILE
def test_datasource_file_is_scoped_to_tenant(
self, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
):
datasource_file = self._make_datasource_file(key="folder/data.csv")
sqlite_session.add(datasource_file)
sqlite_session.commit()
self._bind_session_factory(monkeypatch, sqlite_engine)
access_controller = MagicMock()
access_controller.apply_upload_file_filters = lambda stmt: stmt
with pytest.raises(ValueError, match=f"DatasourceFile {datasource_file.id} not found"):
builders._build_from_datasource_file(
mapping={"datasource_file_id": datasource_file.id, "transfer_method": "datasource_file"},
tenant_id=str(uuid4()),
transfer_method=FileTransferMethod.DATASOURCE_FILE,
access_controller=access_controller,
)