test: use SQLite sessions in datasource files (#39048)

This commit is contained in:
Asuka Minato
2026-08-06 08:28:27 +00:00
committed by GitHub
parent 1a69435fff
commit 78129cf6c3
2 changed files with 114 additions and 154 deletions
+29 -20
View File
@@ -12,8 +12,8 @@ from uuid import uuid4
import httpx
from configs import dify_config
from core.db.session_factory import session_factory
from core.file import remote_fetcher
from extensions.ext_database import db
from extensions.ext_storage import storage
from extensions.storage.storage_type import StorageType
from models.enums import CreatorUserRole
@@ -54,6 +54,7 @@ class DatasourceFileManager:
mimetype: str,
filename: str | None = None,
) -> UploadFile:
"""Persist an uploaded datasource file and its storage payload."""
extension = guess_extension(mimetype) or ".bin"
unique_name = uuid4().hex
unique_filename = f"{unique_name}{extension}"
@@ -82,9 +83,10 @@ class DatasourceFileManager:
created_at=datetime.now(),
)
db.session.add(upload_file)
db.session.commit()
db.session.refresh(upload_file)
with session_factory.create_session() as session:
session.add(upload_file)
session.commit()
session.refresh(upload_file)
return upload_file
@@ -95,6 +97,7 @@ class DatasourceFileManager:
file_url: str,
conversation_id: str | None = None,
) -> ToolFile:
"""Download a remote file and persist its tool-file metadata."""
# try to download image
try:
response = remote_fetcher.make_request("GET", file_url)
@@ -125,8 +128,9 @@ class DatasourceFileManager:
size=len(blob),
)
db.session.add(tool_file)
db.session.commit()
with session_factory.create_session() as session:
session.add(tool_file)
session.commit()
return tool_file
@@ -139,7 +143,8 @@ class DatasourceFileManager:
:return: the binary of the file, mime type
"""
upload_file: UploadFile | None = db.session.get(UploadFile, id)
with session_factory.create_session() as session:
upload_file: UploadFile | None = session.get(UploadFile, id)
if not upload_file:
return None
@@ -157,21 +162,24 @@ class DatasourceFileManager:
:return: the binary of the file, mime type
"""
message_file: MessageFile | None = db.session.get(MessageFile, id)
with session_factory.create_session() as session:
message_file: MessageFile | None = session.get(MessageFile, id)
# Check if message_file is not None
if message_file is not None:
# get tool file id
if message_file.url is not None:
tool_file_id = message_file.url.split("/")[-1]
# trim extension
tool_file_id = tool_file_id.split(".")[0]
# Check if message_file is not None
if message_file is not None:
# get tool file id
if message_file.url is not None:
tool_file_id = message_file.url.split("/")[-1]
# trim extension
tool_file_id = tool_file_id.split(".")[0]
else:
tool_file_id = None
else:
tool_file_id = None
else:
tool_file_id = None
tool_file: ToolFile | None = db.session.get(ToolFile, tool_file_id)
if not tool_file_id:
return None
tool_file: ToolFile | None = session.get(ToolFile, tool_file_id)
if not tool_file:
return None
@@ -185,11 +193,12 @@ class DatasourceFileManager:
"""
get file binary
:param tool_file_id: the id of the tool file
:param upload_file_id: the id of the upload file
:return: the binary of the file, mime type
"""
upload_file: UploadFile | None = db.session.get(UploadFile, upload_file_id)
with session_factory.create_session() as session:
upload_file: UploadFile | None = session.get(UploadFile, upload_file_id)
if not upload_file:
return None, None
@@ -1,14 +1,62 @@
from datetime import datetime
from datetime import UTC, datetime
from unittest.mock import MagicMock, patch
import httpx
import pytest
from sqlalchemy.orm import Session
from core.datasource.datasource_file_manager import DatasourceFileManager
from extensions.storage.storage_type import StorageType
from models.enums import CreatorUserRole
from models.model import MessageFile, UploadFile
from models.tools import ToolFile
def _upload_file(id: str, *, key: str, mime_type: str) -> UploadFile:
upload_file = UploadFile(
tenant_id="tenant-1",
storage_type=StorageType.LOCAL,
key=key,
name="file.png",
size=4,
extension=".png",
mime_type=mime_type,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user-1",
created_at=datetime.now(UTC).replace(tzinfo=None),
used=False,
)
upload_file.id = id
return upload_file
def _tool_file(id: str, *, key: str = "tool_key", mimetype: str = "image/png") -> ToolFile:
tool_file = ToolFile(
tenant_id="tenant-1",
user_id="user-1",
conversation_id=None,
file_key=key,
mimetype=mimetype,
name="tool.png",
size=4,
)
tool_file.id = id
return tool_file
def _message_file(id: str, *, url: str | None) -> MessageFile:
message_file = MessageFile(
message_id="message-1",
type="image",
transfer_method="remote_url",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user-1",
url=url,
)
message_file.id = id
return message_file
class TestDatasourceFileManager:
@patch("core.datasource.datasource_file_manager.time.time")
@patch("core.datasource.datasource_file_manager.os.urandom")
@@ -32,11 +80,10 @@ class TestDatasourceFileManager:
assert f"nonce={mock_urandom.return_value.hex()}" in signed_url
assert "sign=" in signed_url
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
@patch("core.datasource.datasource_file_manager.dify_config")
def test_create_file_by_raw(self, mock_config, mock_uuid, mock_storage, mock_db):
def test_create_file_by_raw(self, mock_config, mock_uuid, mock_storage, sqlite_session: Session):
# Setup
mock_uuid.return_value = MagicMock(hex="unique_hex")
mock_config.STORAGE_TYPE = "local"
@@ -64,14 +111,14 @@ class TestDatasourceFileManager:
assert upload_file.key == f"datasources/{tenant_id}/unique_hex.png"
mock_storage.save.assert_called_once_with(upload_file.key, file_binary)
mock_db.session.add.assert_called_once()
mock_db.session.commit.assert_called_once()
persisted_file = sqlite_session.get(UploadFile, upload_file.id)
assert persisted_file is not None
assert persisted_file.key == upload_file.key
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
@patch("core.datasource.datasource_file_manager.dify_config")
def test_create_file_by_raw_filename_no_extension(self, mock_config, mock_uuid, mock_storage, mock_db):
def test_create_file_by_raw_filename_no_extension(self, mock_config, mock_uuid, mock_storage):
# Setup
mock_uuid.return_value = MagicMock(hex="unique_hex")
mock_config.STORAGE_TYPE = "local"
@@ -94,12 +141,11 @@ class TestDatasourceFileManager:
# Verify
assert upload_file.name == "test.png" # Should append extension
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
@patch("core.datasource.datasource_file_manager.dify_config")
@patch("core.datasource.datasource_file_manager.guess_extension")
def test_create_file_by_raw_unknown_extension(self, mock_guess_ext, mock_config, mock_uuid, mock_storage, mock_db):
def test_create_file_by_raw_unknown_extension(self, mock_guess_ext, mock_config, mock_uuid, mock_storage):
# Setup
mock_guess_ext.return_value = None # Cannot guess
mock_uuid.return_value = MagicMock(hex="unique_hex")
@@ -118,11 +164,10 @@ class TestDatasourceFileManager:
assert upload_file.extension == ".bin"
assert upload_file.name == "unique_hex.bin"
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
@patch("core.datasource.datasource_file_manager.dify_config")
def test_create_file_by_raw_no_filename(self, mock_config, mock_uuid, mock_storage, mock_db):
def test_create_file_by_raw_no_filename(self, mock_config, mock_uuid, mock_storage):
# Setup
mock_uuid.return_value = MagicMock(hex="unique_hex")
mock_config.STORAGE_TYPE = "local"
@@ -141,10 +186,9 @@ class TestDatasourceFileManager:
assert upload_file.extension == ".pdf"
@patch("core.datasource.datasource_file_manager.remote_fetcher")
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
def test_create_file_by_url_mimetype_from_guess(self, mock_uuid, mock_storage, mock_db, mock_ssrf):
def test_create_file_by_url_mimetype_from_guess(self, mock_uuid, mock_storage, mock_ssrf):
# Setup
mock_uuid.return_value = MagicMock(hex="unique_hex")
mock_response = MagicMock()
@@ -154,17 +198,18 @@ class TestDatasourceFileManager:
# Execute
tool_file = DatasourceFileManager.create_file_by_url(
user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/photo.png"
user_id="user_123",
tenant_id="tenant_456",
file_url="https://example.com/photo.png",
)
# Verify
assert tool_file.mimetype == "image/png" # Guessed from .png in URL
@patch("core.datasource.datasource_file_manager.remote_fetcher")
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
def test_create_file_by_url_mimetype_default(self, mock_uuid, mock_storage, mock_db, mock_ssrf):
def test_create_file_by_url_mimetype_default(self, mock_uuid, mock_storage, mock_ssrf):
# Setup
mock_uuid.return_value = MagicMock(hex="unique_hex")
mock_response = MagicMock()
@@ -183,10 +228,9 @@ class TestDatasourceFileManager:
assert tool_file.mimetype == "application/octet-stream"
@patch("core.datasource.datasource_file_manager.remote_fetcher")
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
@patch("core.datasource.datasource_file_manager.uuid4")
def test_create_file_by_url_success(self, mock_uuid, mock_storage, mock_db, mock_ssrf):
def test_create_file_by_url_success(self, mock_uuid, mock_storage, mock_ssrf):
# Setup
mock_uuid.return_value = MagicMock(hex="unique_hex")
mock_response = MagicMock()
@@ -196,7 +240,9 @@ class TestDatasourceFileManager:
# Execute
tool_file = DatasourceFileManager.create_file_by_url(
user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/photo.jpg"
user_id="user_123",
tenant_id="tenant_456",
file_url="https://example.com/photo.jpg",
)
# Verify
@@ -216,152 +262,59 @@ class TestDatasourceFileManager:
user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/large.file"
)
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
def test_get_file_binary(self, mock_storage, mock_db):
# Setup
mock_upload_file = UploadFile(
tenant_id="tenant-id",
storage_type="opendal",
key="some_key",
name="test.txt",
size=0,
extension="txt",
mime_type="image/png",
created_by_role="account",
created_by="account-id",
created_at=datetime.now(),
used=False,
)
mock_db.session.get.return_value = mock_upload_file
def test_get_file_binary(self, mock_storage, sqlite_session: Session):
sqlite_session.add(_upload_file("file_id", key="some_key", mime_type="image/png"))
sqlite_session.commit()
mock_storage.load_once.return_value = b"file content"
# Execute
result = DatasourceFileManager.get_file_binary("file_id")
# Verify
assert result == (b"file content", "image/png")
# Case: Not found
mock_db.session.get.return_value = None
assert DatasourceFileManager.get_file_binary("unknown") is None
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
def test_get_file_binary_by_message_file_id(self, mock_storage, mock_db):
# Setup
mock_message_file = MessageFile(
message_id="message-id",
type="image",
transfer_method="local_file",
created_by_role="user",
created_by="user-id",
url="http://localhost/files/tools/tool_id.png",
def test_get_file_binary_by_message_file_id(self, mock_storage, sqlite_session: Session):
sqlite_session.add_all(
[
_message_file("msg_file_id", url="http://localhost/files/tools/tool_id.png"),
_tool_file("tool_id"),
]
)
mock_tool_file = ToolFile(
user_id="user-id",
tenant_id="tenant-id",
conversation_id="conversation-id",
file_key="tool_key",
mimetype="image/png",
)
def mock_get(model, id):
if model == MessageFile:
return mock_message_file
elif model == ToolFile:
return mock_tool_file
return None
mock_db.session.get.side_effect = mock_get
sqlite_session.commit()
mock_storage.load_once.return_value = b"tool content"
# Execute
result = DatasourceFileManager.get_file_binary_by_message_file_id("msg_file_id")
# Verify
assert result == (b"tool content", "image/png")
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
def test_get_file_binary_by_message_file_id_with_extension(self, mock_storage, mock_db):
# Test that it correctly parses tool_id even with extension in URL
mock_message_file = MessageFile(
message_id="message-id",
type="image",
transfer_method="local_file",
created_by_role="user",
created_by="user-id",
url="http://localhost/files/tools/abcdef.png",
def test_get_file_binary_by_message_file_id_with_extension(self, mock_storage, sqlite_session: Session):
sqlite_session.add_all(
[_message_file("m", url="http://localhost/files/tools/abcdef.png"), _tool_file("abcdef", key="tk")]
)
mock_tool_file = ToolFile(
user_id="user-id",
tenant_id="tenant-id",
conversation_id="conversation-id",
file_key="tk",
mimetype="image/png",
)
mock_tool_file.id = "abcdef"
def mock_get(model, id):
if model == MessageFile:
return mock_message_file
return mock_tool_file
mock_db.session.get.side_effect = mock_get
sqlite_session.commit()
mock_storage.load_once.return_value = b"bits"
result = DatasourceFileManager.get_file_binary_by_message_file_id("m")
assert result == (b"bits", "image/png")
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
def test_get_file_binary_by_message_file_id_failures(self, mock_storage, mock_db):
# Case 1: Message file not found
mock_db.session.get.return_value = None
def test_get_file_binary_by_message_file_id_failures(self, mock_storage, sqlite_session: Session):
assert DatasourceFileManager.get_file_binary_by_message_file_id("none") is None
# Case 2: Message file found but tool file not found
mock_message_file = MessageFile(
message_id="message-id",
type="image",
transfer_method="local_file",
created_by_role="user",
created_by="user-id",
url=None,
)
def mock_get_v2(model, id):
if model == MessageFile:
return mock_message_file
return None
mock_db.session.get.side_effect = mock_get_v2
sqlite_session.add(_message_file("msg_id", url=None))
sqlite_session.commit()
assert DatasourceFileManager.get_file_binary_by_message_file_id("msg_id") is None
@patch("core.datasource.datasource_file_manager.db")
@patch("core.datasource.datasource_file_manager.storage")
def test_get_file_generator_by_upload_file_id(self, mock_storage, mock_db):
# Setup
mock_upload_file = UploadFile(
tenant_id="tenant-id",
storage_type="opendal",
key="upload_key",
name="test.txt",
size=0,
extension="txt",
mime_type="text/plain",
created_by_role="account",
created_by="account-id",
created_at=datetime.now(),
used=False,
)
mock_db.session.get.return_value = mock_upload_file
def test_get_file_generator_by_upload_file_id(self, mock_storage, sqlite_session: Session):
sqlite_session.add(_upload_file("upload_id", key="upload_key", mime_type="text/plain"))
sqlite_session.commit()
mock_storage.load_stream.return_value = iter([b"chunk1", b"chunk2"])
@@ -372,8 +325,6 @@ class TestDatasourceFileManager:
assert mimetype == "text/plain"
assert list(stream) == [b"chunk1", b"chunk2"]
# Case: Not found
mock_db.session.get.return_value = None
stream, mimetype = DatasourceFileManager.get_file_generator_by_upload_file_id("none")
assert stream is None
assert mimetype is None