mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
test: use SQLite sessions in datasource files (#39048)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user