test: migrate residual model sessions and ORM models to SQLite (#40531)

This commit is contained in:
Asuka Minato 2026-08-18 12:23:02 +00:00 committed by GitHub
parent 303ce70ac4
commit e2e3521874
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
14 changed files with 772 additions and 445 deletions

View File

@ -2,10 +2,8 @@ import logging
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.orm import Session
import core.db.session_factory as session_factory_module
from core.indexing_runner import DocumentIsPausedError
from events.event_handlers import create_document_index as handler_module
from models.dataset import Document
@ -14,13 +12,8 @@ from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
@pytest.fixture
def persisted_document(
sqlite_engine: Engine,
sqlite_session: Session,
monkeypatch: pytest.MonkeyPatch,
) -> Document:
real_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
monkeypatch.setattr(session_factory_module, "_session_maker", real_session_maker)
document = Document(
id="doc-1",
tenant_id="tenant-1",

View File

@ -1,43 +1,64 @@
import logging
from types import SimpleNamespace
import pytest
from core.tools.errors import ToolProviderNotFoundError
from events.event_handlers import delete_tool_parameters_cache_when_sync_draft_workflow as handler_module
from graphon.nodes.tool.entities import ToolEntity, ToolProviderType
from models.model import App, AppMode, IconType
from models.workflow import Workflow
def test_missing_tool_provider_does_not_log_error_traceback(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
):
app = SimpleNamespace(id="workflow-id", tenant_id="tenant-id")
workflow = SimpleNamespace(
graph_dict={
"nodes": [
{
"id": "node-id",
"data": {
"type": "tool",
},
}
]
}
app = App(
id="workflow-id",
tenant_id="tenant-id",
name="Workflow app",
description="",
mode=AppMode.WORKFLOW,
icon_type=IconType.EMOJI,
icon="workflow",
icon_background="#FFFFFF",
enable_site=False,
enable_api=False,
max_active_requests=0,
)
tool_entity = SimpleNamespace(
provider_type=SimpleNamespace(value="mcp"),
workflow = Workflow(
tenant_id=app.tenant_id,
app_id=app.id,
type="workflow",
version="draft",
graph='{"nodes": [{"id": "node-id", "data": {"type": "tool"}}]}',
features="{}",
created_by="account-id",
environment_variables=[],
conversation_variables=[],
)
tool_entity = ToolEntity(
provider_type=ToolProviderType.MCP,
provider_id="my-test-mcp-server",
provider_name="my-test-mcp-server",
tool_name="echo",
tool_label="Echo",
tool_configurations={},
credential_id=None,
)
monkeypatch.setattr(handler_module, "adapt_node_config_for_graph", lambda node_data: {"data": node_data["data"]})
monkeypatch.setattr(handler_module.ToolEntity, "model_validate", lambda data: tool_entity)
monkeypatch.setattr(
handler_module,
"adapt_node_config_for_graph",
lambda node_data: {
"data": node_data["data"],
},
)
monkeypatch.setattr(handler_module.ToolEntity, "model_validate", lambda _data: tool_entity)
monkeypatch.setattr(
handler_module.ToolManager,
"get_tool_runtime",
lambda **kwargs: (_ for _ in ()).throw(ToolProviderNotFoundError("mcp provider not found")),
lambda **_kwargs: (_ for _ in ()).throw(ToolProviderNotFoundError("mcp provider not found")),
)
with caplog.at_level(logging.INFO, logger=handler_module.logger.name):

View File

@ -1,10 +1,16 @@
import logging
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from events.event_handlers import queue_default_plugin_install_when_tenant_created as handler_module
from models.account import Tenant
def _tenant() -> Tenant:
tenant = Tenant(name="Test tenant")
tenant.id = "tenant-1"
return tenant
def test_handle_skips_when_no_default_plugins_are_configured(monkeypatch: pytest.MonkeyPatch) -> None:
@ -12,7 +18,7 @@ def test_handle_skips_when_no_default_plugins_are_configured(monkeypatch: pytest
monkeypatch.setattr(handler_module.dify_config, "NEW_USER_DEFAULT_PLUGIN_IDS", "")
monkeypatch.setattr(handler_module.install_default_plugins_task, "delay", delay)
handler_module.handle(SimpleNamespace(id="tenant-1"))
handler_module.handle(_tenant())
delay.assert_not_called()
@ -26,7 +32,7 @@ def test_handle_queues_configured_plugins(monkeypatch: pytest.MonkeyPatch) -> No
monkeypatch.setattr(handler_module.dify_config, "NEW_USER_DEFAULT_PLUGIN_IDS", ",".join(plugins))
monkeypatch.setattr(handler_module.install_default_plugins_task, "delay", delay)
handler_module.handle(SimpleNamespace(id="tenant-1"))
handler_module.handle(_tenant())
delay.assert_called_once_with("tenant-1", plugins)
@ -47,6 +53,6 @@ def test_handle_does_not_fail_tenant_creation_when_queue_is_unavailable(
)
with caplog.at_level(logging.ERROR, logger=handler_module.logger.name):
handler_module.handle(SimpleNamespace(id="tenant-1"))
handler_module.handle(_tenant())
assert "Failed to queue default plugin installation for tenant tenant-1" in caplog.text

View File

@ -1,11 +1,10 @@
import json
from collections.abc import Iterator
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from unittest.mock import patch
from uuid import uuid4
import pytest
from sqlalchemy import event
from sqlalchemy import event, select
from sqlalchemy.orm import Session
from events.app_event import app_was_deleted, app_was_updated
@ -223,57 +222,105 @@ class TestAppWasUpdatedSignal:
class TestAppModelConfigWasUpdatedSignal:
def test_requires_caller_session(self) -> None:
def test_requires_caller_session(self, app_model: App) -> None:
from events.event_handlers.update_app_dataset_join_when_app_model_config_updated import handle
with pytest.raises(TypeError, match="session"):
handle(SimpleNamespace(id="app-1"), app_model_config=None)
handle(app_model, app_model_config=None)
def test_reuses_provided_session_without_committing(self) -> None:
@pytest.mark.parametrize("sqlite_session", [(App, Account, AppModelConfig, AppDatasetJoin)], indirect=True)
def test_reuses_provided_session_without_committing(self, app_model: App, sqlite_session: Session) -> None:
from events.event_handlers.update_app_dataset_join_when_app_model_config_updated import handle
session = MagicMock()
session.scalars.return_value.all.return_value = []
app_model_config = AppModelConfig(app_id="app-1", created_by="user-1", updated_by="user-1")
app_model_config = AppModelConfig(app_id=app_model.id, created_by="user-1", updated_by="user-1")
app_model_config.dataset_configs = json.dumps(
{
"retrieval_model": "multiple",
"datasets": {"datasets": [{"dataset": {"id": "dataset-1"}}]},
}
)
sqlite_session.add_all(
[
app_model_config,
AppDatasetJoin(app_id="other-app", dataset_id="dataset-1"),
]
)
sqlite_session.flush()
commits: list[str] = []
handle(SimpleNamespace(id="app-1"), app_model_config=app_model_config, session=session)
def after_commit(_session: Session) -> None:
commits.append("commit")
added_join = session.add.call_args.args[0]
assert isinstance(added_join, AppDatasetJoin)
assert added_join.app_id == "app-1"
assert added_join.dataset_id == "dataset-1"
session.commit.assert_not_called()
event.listen(sqlite_session, "after_commit", after_commit)
try:
handle(app_model, app_model_config=app_model_config, session=sqlite_session)
finally:
event.remove(sqlite_session, "after_commit", after_commit)
added_joins = [row for row in sqlite_session.new if isinstance(row, AppDatasetJoin)]
assert len(added_joins) == 1
assert added_joins[0].app_id == app_model.id
assert added_joins[0].dataset_id == "dataset-1"
assert commits == []
sqlite_session.flush()
persisted = sqlite_session.scalars(select(AppDatasetJoin).where(AppDatasetJoin.app_id == app_model.id)).all()
assert [(row.app_id, row.dataset_id) for row in persisted] == [(app_model.id, "dataset-1")]
@pytest.mark.parametrize("sqlite_session", [(App, Account, InstalledApp)], indirect=True)
class TestCreateInstalledAppWhenAppCreated:
def test_skips_existing_installation(self) -> None:
def test_skips_existing_installation(self, app_model: App, sqlite_session: Session) -> None:
from events.event_handlers.create_installed_app_when_app_created import handle
session = MagicMock()
session.scalar.return_value = "installed-app-1"
existing = InstalledApp(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
app_owner_tenant_id=app_model.tenant_id,
)
sqlite_session.add(existing)
sqlite_session.commit()
handle(SimpleNamespace(id="app-1", tenant_id="tenant-1"), session=session)
handle(app_model, session=sqlite_session)
session.add.assert_not_called()
session.flush.assert_not_called()
installations = sqlite_session.scalars(
select(InstalledApp).where(
InstalledApp.tenant_id == app_model.tenant_id,
InstalledApp.app_id == app_model.id,
)
).all()
assert installations == [existing]
def test_adds_missing_installation_without_committing(self) -> None:
def test_adds_missing_installation_without_committing(self, app_model: App, sqlite_session: Session) -> None:
from events.event_handlers.create_installed_app_when_app_created import handle
session = MagicMock()
session.scalar.return_value = None
sqlite_session.add(
InstalledApp(
tenant_id="other-tenant",
app_id=app_model.id,
app_owner_tenant_id="other-tenant",
)
)
sqlite_session.commit()
commits: list[str] = []
handle(SimpleNamespace(id="app-1", tenant_id="tenant-1"), session=session)
def after_commit(_session: Session) -> None:
commits.append("commit")
event.listen(sqlite_session, "after_commit", after_commit)
try:
handle(app_model, session=sqlite_session)
finally:
event.remove(sqlite_session, "after_commit", after_commit)
installed_app = sqlite_session.scalar(
select(InstalledApp).where(
InstalledApp.tenant_id == app_model.tenant_id,
InstalledApp.app_id == app_model.id,
)
)
installed_app = session.add.call_args.args[0]
assert isinstance(installed_app, InstalledApp)
assert installed_app.app_id == "app-1"
assert installed_app.tenant_id == "tenant-1"
session.flush.assert_called_once_with()
session.commit.assert_not_called()
assert installed_app.app_id == app_model.id
assert installed_app.tenant_id == app_model.tenant_id
assert commits == []

View File

@ -5,7 +5,6 @@ from uuid import uuid4
import pytest
from sqlalchemy import select
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker
from core.app.entities.app_invoke_entities import AgentAppGenerateEntity, ChatAppGenerateEntity
@ -17,11 +16,10 @@ from models.provider import ProviderType
@pytest.fixture
def credit_pool_session_factory(sqlite_engine: Engine) -> Iterator[sessionmaker[Session]]:
def credit_pool_session_factory(sqlite_session_factory: sessionmaker[Session]) -> Iterator[sessionmaker[Session]]:
"""Bind message-created accounting to fixture-owned SQLite sessions."""
session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
with patch("events.event_handlers.update_provider_when_message_created.db.session", session_factory):
yield session_factory
with patch("events.event_handlers.update_provider_when_message_created.db.session", sqlite_session_factory):
yield sqlite_session_factory
def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_insufficient(

View File

@ -1,7 +1,6 @@
from __future__ import annotations
from datetime import UTC, datetime
from types import SimpleNamespace
import pytest
@ -30,7 +29,7 @@ def test_conversation_variable_response_normalizes_callable_exposed_type() -> No
{
"id": "550e8400-e29b-41d4-a716-446655440000",
"name": "foo",
"value_type": SimpleNamespace(exposed_type=lambda: SegmentType.STRING.exposed_type()),
"value_type": SegmentType.STRING,
}
)

View File

@ -1,7 +1,10 @@
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
from sqlalchemy.orm import Session
from fields.dataset_fields import DatasetDetailResponse, dataset_detail_response_source
from models.account import Account
from models.dataset import AppDatasetJoin, Dataset
from models.model import App, AppMode, IconType
def _dataset_detail_payload(**overrides):
@ -184,45 +187,59 @@ def test_dataset_detail_expands_missing_weighted_score_nested_fields():
}
def test_dataset_detail_response_source_uses_caller_session_for_database_fields():
session = Mock()
getter_mocks = {
"get_app_count": Mock(return_value=3),
"get_document_count": Mock(return_value=4),
"get_word_count": Mock(return_value=500),
"get_author_name": Mock(return_value="Ada"),
"get_tags": Mock(return_value=[{"id": "tag-1", "name": "Tag", "type": "knowledge"}]),
"get_doc_form": Mock(return_value="paragraph"),
"get_external_knowledge_info": Mock(
return_value={
"external_knowledge_id": "knowledge-id",
"external_knowledge_api_id": "api-id",
"external_knowledge_api_name": "api",
"external_knowledge_api_endpoint": "https://example.com",
}
),
"get_doc_metadata": Mock(return_value=[{"id": "metadata-1", "name": "Metadata", "type": "string"}]),
"get_is_published": Mock(return_value=True),
"get_total_documents": Mock(return_value=4),
"get_total_available_documents": Mock(return_value=2),
}
dataset = SimpleNamespace(**_dataset_detail_payload(), **getter_mocks)
@pytest.mark.parametrize("sqlite_session", [(Dataset, Account, App, AppDatasetJoin)], indirect=True)
def test_dataset_detail_response_source_uses_caller_session_for_database_fields(sqlite_session: Session):
account = Account(name="Ada", email="ada@example.com")
account.id = "account-1"
dataset = Dataset(
id="ds-1",
tenant_id="tenant-1",
name="Dataset",
description="desc",
provider="vendor",
permission="only_me",
data_source_type=None,
indexing_technique="economy",
created_by=account.id,
retrieval_model=_dataset_detail_payload()["retrieval_model_dict"],
summary_index_setting=_dataset_detail_payload()["summary_index_setting"],
built_in_field_enabled=False,
icon_info=_dataset_detail_payload()["icon_info"],
runtime_mode="general",
enable_api=False,
is_multimodal=False,
)
dataset.embedding_available = True
decoy_app = App(
id="decoy-app",
tenant_id="tenant-1",
name="Decoy app",
description="",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="app",
icon_background="#FFFFFF",
enable_site=False,
enable_api=False,
max_active_requests=0,
)
decoy_join = AppDatasetJoin(app_id=decoy_app.id, dataset_id="other-dataset")
sqlite_session.add_all([account, dataset, decoy_app, decoy_join])
sqlite_session.flush()
response = DatasetDetailResponse.model_validate(
dataset_detail_response_source(dataset, session=session),
dataset_detail_response_source(dataset, session=sqlite_session),
from_attributes=True,
)
assert response.app_count == 3
assert response.document_count == 4
assert response.word_count == 500
assert response.app_count == 0
assert response.document_count == 0
assert response.word_count == 0
assert response.author_name == "Ada"
assert response.tags[0].id == "tag-1"
assert response.doc_form == "paragraph"
assert response.external_knowledge_info.external_knowledge_api_id == "api-id"
assert response.doc_metadata[0].id == "metadata-1"
assert response.is_published is True
assert response.total_documents == 4
assert response.total_available_documents == 2
for getter in getter_mocks.values():
getter.assert_called_once_with(session=session)
assert response.tags == []
assert response.doc_form is None
assert response.external_knowledge_info.external_knowledge_api_id is None
assert response.doc_metadata == []
assert response.is_published is False
assert response.total_documents == 0
assert response.total_available_documents == 0

View File

@ -1,19 +1,59 @@
from unittest.mock import MagicMock
import json
from datetime import datetime
import pytest
from sqlalchemy.orm import Session
from extensions.storage.storage_type import StorageType
from fields.document_fields import DocumentWithSession
from models.dataset import Document, DocumentSegment
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom
from models.model import UploadFile
def test_document_with_session_uses_explicit_getters() -> None:
session = MagicMock()
document = MagicMock()
document.get_data_source_detail_dict.return_value = {"source": "detail"}
document.get_hit_count.return_value = 3
document.get_doc_metadata_details.return_value = [{"name": "author"}]
source = DocumentWithSession(document=document, session=session)
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentSegment, UploadFile)], indirect=True)
def test_document_with_session_uses_explicit_getters(sqlite_session: Session) -> None:
upload = UploadFile(
tenant_id="tenant-1",
storage_type=StorageType.LOCAL,
key="documents/source.txt",
name="source.txt",
size=12,
extension=".txt",
mime_type="text/plain",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-1",
created_at=datetime(2024, 1, 1),
used=True,
)
document = Document(
id="document-1",
tenant_id="tenant-1",
dataset_id="dataset-1",
position=1,
data_source_type=DataSourceType.UPLOAD_FILE,
data_source_info=json.dumps({"upload_file_id": upload.id}),
batch="batch-1",
name="source.txt",
created_from=DocumentCreatedFrom.WEB,
created_by="account-1",
doc_metadata=None,
)
segment = DocumentSegment(
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=1,
content="hello world",
word_count=2,
tokens=2,
created_by="account-1",
hit_count=3,
)
sqlite_session.add_all([upload, document, segment])
sqlite_session.flush()
source = DocumentWithSession(document=document, session=sqlite_session)
assert source.data_source_detail_dict == {"source": "detail"}
assert source.data_source_detail_dict["upload_file"]["name"] == "source.txt"
assert source.hit_count == 3
assert source.doc_metadata_details == [{"name": "author"}]
document.get_data_source_detail_dict.assert_called_once_with(session=session)
document.get_hit_count.assert_called_once_with(session=session)
document.get_doc_metadata_details.assert_called_once_with(session=session)
assert source.doc_metadata_details is None

View File

@ -1,34 +1,40 @@
from __future__ import annotations
from datetime import datetime
from types import SimpleNamespace
import pytest
from core.workflow.file_reference import build_file_reference
from extensions.storage.storage_type import StorageType
from fields import conversation_fields, message_fields
from fields.file_fields import FileResponse, FileWithSignedUrl, RemoteFileInfo, UploadConfig
from graphon.file import File, FileTransferMethod, FileType
from models.enums import CreatorUserRole
from models.model import UploadFile
def test_file_response_serializes_datetime() -> None:
created_at = datetime(2024, 1, 1, 12, 0, 0)
file_obj = SimpleNamespace(
id="file-1",
file_obj = UploadFile(
tenant_id="tenant-1",
storage_type=StorageType.LOCAL,
key="key-1",
name="example.txt",
size=1024,
extension="txt",
mime_type="text/plain",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user-1",
created_at=created_at,
preview_url="https://preview",
source_url="https://source",
original_url="https://origin",
user_id="user-1",
tenant_id="tenant-1",
conversation_id="conv-1",
file_key="key-1",
used=False,
)
file_obj.id = "file-1"
file_obj.preview_url = "https://preview"
file_obj.original_url = "https://origin"
file_obj.user_id = "user-1"
file_obj.conversation_id = "conv-1"
file_obj.file_key = "key-1"
serialized = FileResponse.model_validate(file_obj, from_attributes=True).model_dump(mode="json")

View File

@ -1,13 +1,22 @@
from datetime import UTC, datetime
from types import SimpleNamespace
import pytest
from sqlalchemy.orm import Session
from fields.snippet_fields import SnippetListItemResponse
from libs.helper import dump_response
from models import snippet as snippet_module
from models.account import Account
from models.snippet import CustomizedSnippet
def test_snippet_list_fields_include_author_name() -> None:
snippet = SimpleNamespace(
@pytest.mark.parametrize("sqlite_session", [(CustomizedSnippet, Account)], indirect=True)
def test_snippet_list_fields_include_author_name(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> None:
account = Account(name="Alice", email="alice@example.com")
account.id = "account-1"
snippet = CustomizedSnippet(
id="snippet-1",
tenant_id="tenant-1",
name="Snippet",
description="Reusable node",
type="node",
@ -15,13 +24,14 @@ def test_snippet_list_fields_include_author_name() -> None:
use_count=0,
is_published=False,
icon_info=None,
tags=[],
created_by="account-1",
author_name="Alice",
created_at=datetime.fromtimestamp(1704067200, tz=UTC),
updated_by="account-1",
updated_at=datetime.fromtimestamp(1704067201, tz=UTC),
)
sqlite_session.add_all([account, snippet])
sqlite_session.flush()
monkeypatch.setattr(snippet_module.db, "session", sqlite_session)
result = dump_response(SnippetListItemResponse, snippet)

View File

@ -12,7 +12,7 @@ This test suite covers:
import base64
import secrets
from datetime import UTC, datetime
from unittest.mock import MagicMock, patch
from unittest.mock import patch
from uuid import uuid4
import pytest
@ -335,36 +335,67 @@ class TestTenantRelationshipIntegrity:
# Assert
assert tenant_id_none is None
def test_set_current_tenant_with_session_uses_caller_session(self):
@pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin)], indirect=True)
def test_set_current_tenant_with_session_uses_caller_session(self, sqlite_session: Session):
account = Account(name="Test User", email="test@example.com")
account.id = str(uuid4())
tenant = Tenant(name="Test Tenant")
tenant.id = str(uuid4())
join = MagicMock(role=TenantAccountRole.OWNER)
session = MagicMock(spec=Session)
session.scalar.return_value = join
session.scalars.return_value.one.return_value = tenant
decoy_tenant = Tenant(name="Decoy Tenant")
decoy_tenant.id = str(uuid4())
sqlite_session.add_all(
[
account,
tenant,
decoy_tenant,
TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=TenantAccountRole.OWNER,
),
TenantAccountJoin(
tenant_id=decoy_tenant.id,
account_id=account.id,
role=TenantAccountRole.NORMAL,
),
]
)
sqlite_session.flush()
with patch("models.account.Session") as session_class:
account.set_current_tenant_with_session(tenant, session=session)
account.set_current_tenant_with_session(tenant, session=sqlite_session)
session_class.assert_not_called()
assert account.current_tenant is tenant
assert account.role == TenantAccountRole.OWNER
def test_set_tenant_id_with_session_uses_caller_session(self):
@pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin)], indirect=True)
def test_set_tenant_id_with_session_uses_caller_session(self, sqlite_session: Session):
account = Account(name="Test User", email="test@example.com")
account.id = str(uuid4())
decoy_account = Account(name="Decoy User", email="decoy@example.com")
decoy_account.id = str(uuid4())
tenant = Tenant(name="Test Tenant")
tenant.id = str(uuid4())
join = MagicMock(role=TenantAccountRole.ADMIN)
session = MagicMock(spec=Session)
session.execute.return_value.first.return_value = (tenant, join)
sqlite_session.add_all(
[
account,
decoy_account,
tenant,
TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=TenantAccountRole.ADMIN,
),
TenantAccountJoin(
tenant_id=tenant.id,
account_id=decoy_account.id,
role=TenantAccountRole.OWNER,
),
]
)
sqlite_session.flush()
with patch("models.account.Session") as session_class:
account.set_tenant_id_with_session(tenant.id, session=session)
account.set_tenant_id_with_session(tenant.id, session=sqlite_session)
session_class.assert_not_called()
assert account.current_tenant is tenant
assert account.role == TenantAccountRole.ADMIN

View File

@ -12,17 +12,17 @@ This test suite covers:
import json
import pickle
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import Mock, call, patch
from unittest.mock import patch
from urllib.parse import parse_qs, urlparse
from uuid import uuid4
import pytest
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from core.rag.entities import ParentMode
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from extensions.storage.storage_type import StorageType
from models import dataset as dataset_module
from models.account import Account
from models.dataset import (
AppDatasetJoin,
@ -36,6 +36,7 @@ from models.dataset import (
Embedding,
ExternalKnowledgeApis,
ExternalKnowledgeBindings,
SegmentAttachmentBinding,
)
from models.enums import (
CreatorUserRole,
@ -45,7 +46,84 @@ from models.enums import (
ProcessRuleMode,
SegmentStatus,
)
from models.model import UploadFile
from models.model import App, AppMode, IconType, UploadFile
def _make_dataset(
*,
dataset_id: str = "dataset-1",
tenant_id: str = "tenant-1",
created_by: str = "account-1",
provider: str = "vendor",
) -> Dataset:
return Dataset(
id=dataset_id,
tenant_id=tenant_id,
name=f"Dataset {dataset_id}",
data_source_type=DataSourceType.UPLOAD_FILE,
created_by=created_by,
provider=provider,
built_in_field_enabled=False,
)
def _make_document(
*,
document_id: str = "document-1",
dataset_id: str = "dataset-1",
tenant_id: str = "tenant-1",
process_rule_id: str | None = None,
position: int = 1,
word_count: int | None = None,
indexing_status: IndexingStatus = IndexingStatus.WAITING,
) -> Document:
return Document(
id=document_id,
tenant_id=tenant_id,
dataset_id=dataset_id,
position=position,
data_source_type=DataSourceType.UPLOAD_FILE,
dataset_process_rule_id=process_rule_id,
batch="batch-1",
name=f"{document_id}.txt",
created_from=DocumentCreatedFrom.WEB,
created_by="account-1",
word_count=word_count,
indexing_status=indexing_status,
)
def _make_app(*, app_id: str, tenant_id: str = "tenant-1") -> App:
return App(
id=app_id,
tenant_id=tenant_id,
name=f"App {app_id}",
description="",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="app",
icon_background="#FFFFFF",
enable_site=False,
enable_api=False,
max_active_requests=0,
)
def _make_segments(document: Document, hit_counts: list[int]) -> list[DocumentSegment]:
return [
DocumentSegment(
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=position,
content=f"Segment {position}",
word_count=2,
tokens=2,
created_by="account-1",
hit_count=hit_count,
)
for position, hit_count in enumerate(hit_counts, start=1)
]
class TestDatasetModelValidation:
@ -92,104 +170,97 @@ class TestDatasetModelValidation:
assert dataset.embedding_model == "text-embedding-ada-002"
assert dataset.embedding_model_provider == "openai"
def test_session_aware_dataset_getters_use_caller_session(self):
dataset = Dataset(
tenant_id=str(uuid4()),
name="Test Dataset",
data_source_type=DataSourceType.UPLOAD_FILE,
created_by=str(uuid4()),
id=str(uuid4()),
@pytest.mark.parametrize("sqlite_session", [(Dataset, Account, DatasetProcessRule, Document)], indirect=True)
def test_session_aware_dataset_getters_use_caller_session(self, sqlite_session: Session):
account = Account(name="Ada", email="ada@example.com")
account.id = "account-1"
dataset = _make_dataset(created_by=account.id)
process_rule = DatasetProcessRule(
dataset_id=dataset.id,
mode=ProcessRuleMode.CUSTOM,
rules=json.dumps({"segmentation": {"max_tokens": 100}}),
created_by=account.id,
)
account = Mock()
process_rule = Mock()
session = Mock()
session.get.return_value = account
session.scalar.side_effect = [process_rule, IndexStructureType.PARAGRAPH_INDEX]
document = _make_document(dataset_id=dataset.id)
document.doc_form = IndexStructureType.PARAGRAPH_INDEX
sqlite_session.add_all([dataset, account, process_rule, document])
sqlite_session.flush()
assert dataset.get_created_by_account(session=session) is account
assert dataset.get_latest_process_rule(session=session) is process_rule
assert dataset.get_doc_form(session=session) == IndexStructureType.PARAGRAPH_INDEX
session.get.assert_called_once_with(Account, dataset.created_by)
assert session.scalar.call_count == 2
assert dataset.get_created_by_account(session=sqlite_session) is account
assert dataset.get_latest_process_rule(session=sqlite_session) is process_rule
assert dataset.get_doc_form(session=sqlite_session) == IndexStructureType.PARAGRAPH_INDEX
@pytest.mark.parametrize("sqlite_session", [(Dataset, Document)], indirect=True)
def test_get_doc_form_ignores_foreign_tenant_document(self, sqlite_session: Session) -> None:
dataset_id = str(uuid4())
tenant_id = str(uuid4())
created_by = str(uuid4())
dataset = Dataset(
id=dataset_id,
tenant_id=tenant_id,
name="Dataset",
data_source_type=DataSourceType.UPLOAD_FILE,
created_by=created_by,
)
foreign_document = Document(
tenant_id=str(uuid4()),
dataset_id=dataset_id,
position=1,
data_source_type=DataSourceType.UPLOAD_FILE,
batch="foreign",
name="Foreign",
created_from=DocumentCreatedFrom.WEB,
created_by=created_by,
doc_form=IndexStructureType.PARENT_CHILD_INDEX,
dataset = _make_dataset()
foreign_document = _make_document(
dataset_id=dataset.id,
tenant_id="tenant-2",
)
foreign_document.doc_form = IndexStructureType.PARENT_CHILD_INDEX
sqlite_session.add_all([dataset, foreign_document])
sqlite_session.flush()
assert dataset.get_doc_form(session=sqlite_session) is None
def test_get_dataset_keyword_table_uses_caller_session(self):
dataset = Dataset(
tenant_id=str(uuid4()),
name="Test Dataset",
data_source_type=DataSourceType.UPLOAD_FILE,
created_by=str(uuid4()),
id=str(uuid4()),
@pytest.mark.parametrize("sqlite_session", [(Dataset, DatasetKeywordTable)], indirect=True)
def test_get_dataset_keyword_table_uses_caller_session(self, sqlite_session: Session):
dataset = _make_dataset()
keyword_table = DatasetKeywordTable(
dataset_id=dataset.id,
keyword_table=json.dumps({"keyword": ["node-1"]}),
)
keyword_table = Mock()
session = Mock()
session.scalar.return_value = keyword_table
sqlite_session.add_all([dataset, keyword_table])
sqlite_session.flush()
result = dataset.get_dataset_keyword_table(session=session)
result = dataset.get_dataset_keyword_table(session=sqlite_session)
assert result is keyword_table
session.scalar.assert_called_once()
def test_dataset_detail_getters_use_caller_session(self):
dataset = Dataset(
tenant_id=str(uuid4()),
name="Test Dataset",
data_source_type=DataSourceType.UPLOAD_FILE,
created_by=str(uuid4()),
provider="vendor",
built_in_field_enabled=False,
id=str(uuid4()),
@pytest.mark.parametrize("sqlite_session", [(Dataset, Account, App, AppDatasetJoin, Document)], indirect=True)
def test_dataset_detail_getters_use_caller_session(self, sqlite_session: Session):
account = Account(name="Ada", email="ada@example.com")
account.id = "account-1"
dataset = _make_dataset(created_by=account.id)
available_document = _make_document(
document_id="document-1",
dataset_id=dataset.id,
word_count=200,
indexing_status=IndexingStatus.COMPLETED,
)
account = Mock(name="account", name_value="Ada")
account.name = "Ada"
session = Mock()
session.get.return_value = account
session.scalar.side_effect = [2, 1, 3, 2, 500, IndexStructureType.PARAGRAPH_INDEX]
session.scalars.return_value.all.return_value = []
available_document.enabled = True
available_document.archived = False
available_document.doc_form = IndexStructureType.PARAGRAPH_INDEX
waiting_document = _make_document(
document_id="document-2",
dataset_id=dataset.id,
position=2,
word_count=300,
)
app = _make_app(app_id="app-1")
sqlite_session.add_all(
[
account,
dataset,
available_document,
waiting_document,
app,
AppDatasetJoin(app_id=app.id, dataset_id=dataset.id),
]
)
sqlite_session.flush()
with patch("models.dataset.db") as mock_db:
assert dataset.get_total_documents(session=session) == 2
assert dataset.get_total_available_documents(session=session) == 1
assert dataset.get_app_count(session=session) == 3
assert dataset.get_document_count(session=session) == 2
assert dataset.get_word_count(session=session) == 500
assert dataset.get_author_name(session=session) == "Ada"
assert dataset.get_tags(session=session) == []
assert dataset.get_doc_form(session=session) == IndexStructureType.PARAGRAPH_INDEX
assert dataset.get_external_knowledge_info(session=session) is None
assert dataset.get_doc_metadata(session=session) == []
assert dataset.get_is_published(session=session) is False
assert session.scalar.call_count == 6
assert session.scalars.call_count == 2
mock_db.session.scalar.assert_not_called()
mock_db.session.scalars.assert_not_called()
assert dataset.get_total_documents(session=sqlite_session) == 2
assert dataset.get_total_available_documents(session=sqlite_session) == 1
assert dataset.get_app_count(session=sqlite_session) == 1
assert dataset.get_document_count(session=sqlite_session) == 2
assert dataset.get_word_count(session=sqlite_session) == 500
assert dataset.get_author_name(session=sqlite_session) == "Ada"
assert dataset.get_tags(session=sqlite_session) == []
assert dataset.get_doc_form(session=sqlite_session) == IndexStructureType.PARAGRAPH_INDEX
assert dataset.get_external_knowledge_info(session=sqlite_session) is None
assert dataset.get_doc_metadata(session=sqlite_session) == []
assert dataset.get_is_published(session=sqlite_session) is False
def test_dataset_indexing_technique_validation(self):
"""Test dataset indexing technique values."""
@ -292,56 +363,82 @@ class TestDatasetModelValidation:
assert result["top_k"] == 2
assert result["score_threshold"] == 0.0
def test_dataset_external_knowledge_info_returns_none_for_cross_tenant_template(self):
@pytest.mark.parametrize(
"sqlite_session", [(Dataset, ExternalKnowledgeBindings, ExternalKnowledgeApis)], indirect=True
)
def test_dataset_external_knowledge_info_returns_none_for_cross_tenant_template(self, sqlite_session: Session):
"""Test external datasets fail closed when the bound template is outside the tenant."""
dataset = Dataset(
tenant_id=str(uuid4()),
name="External Dataset",
data_source_type=DataSourceType.UPLOAD_FILE,
created_by=str(uuid4()),
provider="external",
dataset = _make_dataset(provider="external")
external_api = ExternalKnowledgeApis(
tenant_id="other-tenant",
created_by="account-1",
updated_by=None,
name="Other tenant API",
description="",
settings=json.dumps({"endpoint": "https://example.com"}),
)
binding = ExternalKnowledgeBindings(
tenant_id="tenant-id",
external_knowledge_api_id=str(uuid4()),
dataset_id="dataset-id",
tenant_id=dataset.tenant_id,
external_knowledge_api_id=external_api.id,
dataset_id=dataset.id,
external_knowledge_id="knowledge-1",
created_by="account-id",
)
sqlite_session.add_all([dataset, external_api, binding])
sqlite_session.flush()
session = Mock()
session.scalar.side_effect = [binding, None]
with patch("models.dataset.db") as mock_db:
assert dataset.get_external_knowledge_info(session=session) is None
assert dataset.get_external_knowledge_info(session=sqlite_session) is None
assert session.scalar.call_count == 2
mock_db.session.scalar.assert_not_called()
def test_external_knowledge_api_dataset_bindings_use_caller_session(self):
@pytest.mark.parametrize(
"sqlite_session", [(ExternalKnowledgeApis, ExternalKnowledgeBindings, Dataset)], indirect=True
)
def test_external_knowledge_api_dataset_bindings_use_caller_session(self, sqlite_session: Session):
external_api = ExternalKnowledgeApis(
tenant_id=str(uuid4()),
tenant_id="tenant-1",
created_by=str(uuid4()),
updated_by=None,
name="External API",
description="",
settings=None,
)
binding = Mock(dataset_id="dataset-1")
dataset = SimpleNamespace(id="dataset-1", name="Dataset")
session = Mock()
session.scalars.side_effect = [
Mock(all=Mock(return_value=[binding])),
Mock(all=Mock(return_value=[dataset])),
]
other_api = ExternalKnowledgeApis(
tenant_id="tenant-1",
created_by="account-1",
updated_by=None,
name="Other API",
description="",
settings=None,
)
dataset = _make_dataset()
decoy_dataset = _make_dataset(dataset_id="dataset-2")
sqlite_session.add_all([external_api, other_api, dataset, decoy_dataset])
sqlite_session.flush()
sqlite_session.add_all(
[
ExternalKnowledgeBindings(
tenant_id="tenant-1",
external_knowledge_api_id=external_api.id,
dataset_id=dataset.id,
external_knowledge_id="knowledge-1",
created_by="account-1",
),
ExternalKnowledgeBindings(
tenant_id="tenant-1",
external_knowledge_api_id=other_api.id,
dataset_id=decoy_dataset.id,
external_knowledge_id="knowledge-2",
created_by="account-1",
),
]
)
sqlite_session.flush()
with patch("models.dataset.db") as mock_db:
result = external_api.get_dataset_bindings(session=session)
result = external_api.get_dataset_bindings(session=sqlite_session)
assert result == [{"id": "dataset-1", "name": "Dataset"}]
assert session.scalars.call_count == 2
mock_db.session.scalars.assert_not_called()
assert result == [{"id": dataset.id, "name": dataset.name}]
def test_dataset_query_get_queries_uses_caller_session(self):
@pytest.mark.parametrize("sqlite_session", [(DatasetQuery, UploadFile)], indirect=True)
def test_dataset_query_get_queries_uses_caller_session(self, sqlite_session: Session):
dataset_query = DatasetQuery(
dataset_id=str(uuid4()),
content=json.dumps([{"content_type": "image_query", "content": "file-1"}]),
@ -350,21 +447,25 @@ class TestDatasetModelValidation:
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
)
upload_file = SimpleNamespace(
id="file-1",
upload_file = UploadFile(
tenant_id="tenant-1",
storage_type=StorageType.LOCAL,
key="image.png",
name="image.png",
size=10,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-1",
created_at=datetime(2024, 1, 1),
used=False,
)
session = Mock()
session.scalar.return_value = upload_file
upload_file.id = "file-1"
sqlite_session.add_all([dataset_query, upload_file])
sqlite_session.flush()
with (
patch("models.dataset.db") as mock_db,
patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"),
):
queries = dataset_query.get_queries(session=session)
with patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"):
queries = dataset_query.get_queries(session=sqlite_session)
assert queries == [
{
@ -380,8 +481,6 @@ class TestDatasetModelValidation:
},
}
]
session.scalar.assert_called_once()
mock_db.session.scalar.assert_not_called()
def test_dataset_retrieval_model_dict_property(self):
"""Test retrieval_model_dict property with default values."""
@ -479,29 +578,35 @@ class TestDocumentModelRelationships:
assert "notion_import" in Document.DATA_SOURCES
assert "website_crawl" in Document.DATA_SOURCES
def test_session_aware_document_getters_use_caller_session(self):
document = Document(
tenant_id=str(uuid4()),
dataset_id=str(uuid4()),
position=1,
data_source_type=DataSourceType.UPLOAD_FILE,
batch="batch_001",
name="test.pdf",
created_from=DocumentCreatedFrom.WEB,
created_by=str(uuid4()),
dataset_process_rule_id=str(uuid4()),
@pytest.mark.parametrize("sqlite_session", [(Document, DatasetProcessRule, DocumentSegment)], indirect=True)
def test_session_aware_document_getters_use_caller_session(self, sqlite_session: Session):
process_rule = DatasetProcessRule(
dataset_id="dataset-1",
mode=ProcessRuleMode.CUSTOM,
rules=None,
created_by="account-1",
)
process_rule = Mock()
session = Mock()
session.get.return_value = process_rule
session.scalar.side_effect = [3, 7]
document = _make_document(process_rule_id=process_rule.id)
segments = [
DocumentSegment(
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=position,
content=f"Segment {position}",
word_count=2,
tokens=2,
created_by="account-1",
hit_count=hit_count,
)
for position, hit_count in [(1, 2), (2, 1), (3, 4)]
]
sqlite_session.add_all([process_rule, document, *segments])
sqlite_session.flush()
assert document.get_dataset_process_rule(session=session) is process_rule
assert document.get_segment_count(session=session) == 3
assert document.get_hit_count(session=session) == 7
session.get.assert_called_once_with(DatasetProcessRule, document.dataset_process_rule_id)
assert session.scalar.call_count == 2
assert document.get_dataset_process_rule(session=sqlite_session) is process_rule
assert document.get_segment_count(session=sqlite_session) == 3
assert document.get_hit_count(session=sqlite_session) == 7
def test_document_display_status_queuing(self):
"""Test document display_status property for queuing state."""
@ -706,25 +811,22 @@ class TestDocumentModelRelationships:
# Assert
assert result == {}
def test_document_get_dataset_uses_caller_session(self):
document = Document(
tenant_id=str(uuid4()),
dataset_id=str(uuid4()),
position=1,
data_source_type=DataSourceType.UPLOAD_FILE,
batch="batch_001",
name="test.pdf",
created_from=DocumentCreatedFrom.WEB,
created_by=str(uuid4()),
)
dataset = Dataset()
session = Mock()
session.get.return_value = dataset
@pytest.mark.parametrize("sqlite_session", [(Document, Dataset)], indirect=True)
def test_document_get_dataset_uses_caller_session(self, sqlite_session: Session):
dataset = _make_dataset()
document = _make_document(dataset_id=dataset.id)
sqlite_session.add_all([dataset, document])
sqlite_session.flush()
assert document.get_dataset(session=session) is dataset
session.get.assert_called_once_with(Dataset, document.dataset_id)
assert document.get_dataset(session=sqlite_session) is dataset
def test_document_average_segment_length(self):
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentSegment)], indirect=True)
def test_document_average_segment_length(
self,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
monkeypatch: pytest.MonkeyPatch,
):
"""Test average_segment_length property calculation."""
# Arrange
document = Document(
@ -738,14 +840,17 @@ class TestDocumentModelRelationships:
created_by=str(uuid4()),
word_count=1000,
)
sqlite_session.add(document)
sqlite_session.flush()
sqlite_session.add_all(_make_segments(document, [0] * 10))
sqlite_session.commit()
monkeypatch.setattr(dataset_module.db, "session", scoped_session(sqlite_session_factory))
# Mock segment_count property
with patch.object(Document, "segment_count", new_callable=lambda: property(lambda self: 10)):
# Act
result = document.average_segment_length
# Act
result = document.average_segment_length
# Assert
assert result == 100
# Assert
assert result == 100
def test_document_average_segment_length_zero(self):
"""Test average_segment_length property when word_count is zero."""
@ -772,107 +877,101 @@ class TestDocumentModelRelationships:
class TestDocumentSegmentIndexing:
"""Test suite for DocumentSegment model indexing and operations."""
def test_get_child_chunks_uses_caller_session(self):
@pytest.mark.parametrize(
"sqlite_session", [(DocumentSegment, Document, DatasetProcessRule, ChildChunk)], indirect=True
)
def test_get_child_chunks_uses_caller_session(self, sqlite_session: Session):
process_rule = DatasetProcessRule(
dataset_id="dataset-1",
mode=ProcessRuleMode.HIERARCHICAL,
rules=json.dumps({"parent_mode": ParentMode.PARAGRAPH}),
created_by="account-1",
)
document = _make_document(process_rule_id=process_rule.id)
segment = DocumentSegment(
tenant_id=str(uuid4()),
dataset_id=str(uuid4()),
document_id=str(uuid4()),
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=1,
content="Test content",
word_count=2,
tokens=5,
created_by=str(uuid4()),
)
document = Document()
process_rule = Mock(mode="hierarchical", rules_dict={"parent_mode": "paragraph"})
child_chunk = ChildChunk(
tenant_id="tenant-id",
dataset_id="dataset-id",
document_id="document-id",
segment_id="segment-id",
tenant_id=segment.tenant_id,
dataset_id=segment.dataset_id,
document_id=segment.document_id,
segment_id=segment.id,
position=1,
content="",
word_count=0,
created_by="account-id",
)
session = Mock()
session.get.return_value = document
session.scalars.return_value.all.return_value = [child_chunk]
with (
patch.object(Document, "get_dataset_process_rule", return_value=process_rule) as get_process_rule,
patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode="paragraph")),
):
result = segment.get_child_chunks(session=session)
sqlite_session.add_all([process_rule, document, segment, child_chunk])
sqlite_session.flush()
result = segment.get_child_chunks(session=sqlite_session)
assert result == [child_chunk]
session.get.assert_called_once_with(Document, segment.document_id)
get_process_rule.assert_called_once_with(session=session)
session.scalars.assert_called_once()
def test_get_child_chunks_includes_full_doc_unless_explicitly_hidden(self):
@pytest.mark.parametrize(
"sqlite_session", [(DocumentSegment, Document, DatasetProcessRule, ChildChunk)], indirect=True
)
def test_get_child_chunks_includes_full_doc_unless_explicitly_hidden(self, sqlite_session: Session):
process_rule = DatasetProcessRule(
dataset_id="dataset-1",
mode=ProcessRuleMode.HIERARCHICAL,
rules=json.dumps({"parent_mode": ParentMode.FULL_DOC}),
created_by="account-1",
)
document = _make_document(process_rule_id=process_rule.id)
segment = DocumentSegment(
tenant_id=str(uuid4()),
dataset_id=str(uuid4()),
document_id=str(uuid4()),
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=1,
content="Test content",
word_count=2,
tokens=5,
created_by=str(uuid4()),
)
document = Document()
process_rule = Mock(
mode="hierarchical",
rules_dict={"parent_mode": ParentMode.FULL_DOC},
)
session = Mock()
session.get.return_value = document
child_chunk = ChildChunk(
tenant_id="tenant-id",
dataset_id="dataset-id",
document_id="document-id",
segment_id="segment-id",
tenant_id=segment.tenant_id,
dataset_id=segment.dataset_id,
document_id=segment.document_id,
segment_id=segment.id,
position=1,
content="",
word_count=0,
created_by="account-id",
)
session.scalars.return_value.all.return_value = [child_chunk]
with (
patch.object(Document, "get_dataset_process_rule", return_value=process_rule),
patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode=ParentMode.FULL_DOC)),
):
result = segment.get_child_chunks(session=session)
response_result = segment.get_child_chunks(session=session, include_full_doc=False)
sqlite_session.add_all([process_rule, document, segment, child_chunk])
sqlite_session.flush()
result = segment.get_child_chunks(session=sqlite_session)
response_result = segment.get_child_chunks(session=sqlite_session, include_full_doc=False)
assert result == [child_chunk]
assert response_result == []
session.scalars.assert_called_once()
def test_relationship_getters_use_caller_session(self):
@pytest.mark.parametrize("sqlite_session", [(Dataset, Document, DocumentSegment)], indirect=True)
def test_relationship_getters_use_caller_session(self, sqlite_session: Session):
dataset = _make_dataset()
document = _make_document(dataset_id=dataset.id)
segment = DocumentSegment(
tenant_id=str(uuid4()),
dataset_id=str(uuid4()),
document_id=str(uuid4()),
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
document_id=document.id,
position=1,
content="Test content",
word_count=2,
tokens=5,
created_by=str(uuid4()),
)
dataset = Dataset()
document = Document()
session = Mock()
session.get.side_effect = [dataset, document]
sqlite_session.add_all([dataset, document, segment])
sqlite_session.flush()
assert segment.get_dataset(session=session) is dataset
assert segment.get_document(session=session) is document
assert session.get.call_args_list == [
call(Dataset, segment.dataset_id),
call(Document, segment.document_id),
]
assert segment.get_dataset(session=sqlite_session) is dataset
assert segment.get_document(session=sqlite_session) is document
def test_document_segment_creation_with_required_fields(self):
"""Test creating a document segment with all required fields."""
@ -1029,7 +1128,10 @@ class TestDocumentSegmentIndexing:
# Assert
assert segment.hit_count == 5
def test_document_segment_attachments_prefers_files_url_for_source_url(self, monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("sqlite_session", [(DocumentSegment, UploadFile, SegmentAttachmentBinding)], indirect=True)
def test_document_segment_attachments_prefers_files_url_for_source_url(
self, sqlite_session: Session, monkeypatch: pytest.MonkeyPatch
):
"""Test attachment source URLs use FILES_URL before falling back to CONSOLE_API_URL."""
# Arrange
segment = DocumentSegment(
@ -1057,6 +1159,15 @@ class TestDocumentSegmentIndexing:
used=False,
)
attachment.id = "upload-1"
binding = SegmentAttachmentBinding(
tenant_id=segment.tenant_id,
dataset_id=segment.dataset_id,
document_id=segment.document_id,
segment_id=segment.id,
attachment_id=attachment.id,
)
sqlite_session.add_all([segment, attachment, binding])
sqlite_session.flush()
monkeypatch.setattr("models.dataset.time.time", lambda: 1700000000)
monkeypatch.setattr("models.dataset.os.urandom", lambda _: b"\x01" * 16)
@ -1064,11 +1175,8 @@ class TestDocumentSegmentIndexing:
monkeypatch.setattr("models.dataset.dify_config.FILES_URL", "https://files.example.com")
monkeypatch.setattr("models.dataset.dify_config.CONSOLE_API_URL", "https://console.example.com")
session = Mock()
session.execute.return_value.all.return_value = [(Mock(), attachment)]
# Act
attachments = segment.get_attachments(session=session)
attachments = segment.get_attachments(session=sqlite_session)
# Assert
assert len(attachments) == 1
@ -1080,7 +1188,6 @@ class TestDocumentSegmentIndexing:
assert query["timestamp"] == ["1700000000"]
assert query["nonce"] == ["01010101010101010101010101010101"]
assert query["sign"][0]
session.execute.assert_called_once()
def test_document_segment_error_tracking(self):
"""Test document segment error tracking."""
@ -1308,20 +1415,20 @@ class TestDatasetKeywordTable:
# Assert
assert keyword_table.data_source_type == "file"
def test_get_keyword_table_dict_from_database_uses_caller_session(self):
dataset = Mock(tenant_id="tenant-1")
session = Mock()
session.scalar.return_value = dataset
@pytest.mark.parametrize("sqlite_session", [(Dataset, DatasetKeywordTable)], indirect=True)
def test_get_keyword_table_dict_from_database_uses_caller_session(self, sqlite_session: Session):
dataset = _make_dataset()
keyword_table = DatasetKeywordTable(
dataset_id="dataset-1",
dataset_id=dataset.id,
keyword_table=json.dumps({"__data__": {"table": {"keyword": ["node-1"]}}}),
data_source_type="database",
)
sqlite_session.add_all([dataset, keyword_table])
sqlite_session.flush()
result = keyword_table.get_keyword_table_dict(session=session)
result = keyword_table.get_keyword_table_dict(session=sqlite_session)
assert result == {"__data__": {"table": {"keyword": {"node-1"}}}}
session.scalar.assert_called_once()
class TestAppDatasetJoin:
@ -1462,7 +1569,13 @@ class TestModelIntegration:
assert document.word_count == 100
assert segment.status == SegmentStatus.COMPLETED
def test_document_to_dict_serialization(self):
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentSegment)], indirect=True)
def test_document_to_dict_serialization(
self,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
monkeypatch: pytest.MonkeyPatch,
):
"""Test document to_dict method for serialization."""
# Arrange
tenant_id = str(uuid4())
@ -1481,20 +1594,20 @@ class TestModelIntegration:
word_count=100,
indexing_status=IndexingStatus.COMPLETED,
)
sqlite_session.add(document)
sqlite_session.flush()
sqlite_session.add_all(_make_segments(document, [2, 2, 2, 2, 2]))
sqlite_session.commit()
monkeypatch.setattr(dataset_module.db, "session", scoped_session(sqlite_session_factory))
# Mock segment_count and hit_count
with (
patch.object(Document, "segment_count", new_callable=lambda: property(lambda self: 5)),
patch.object(Document, "hit_count", new_callable=lambda: property(lambda self: 10)),
):
# Act
result = document.to_dict()
# Act
result = document.to_dict()
# Assert
assert result["tenant_id"] == tenant_id
assert result["dataset_id"] == dataset_id
assert result["name"] == "test.pdf"
assert result["word_count"] == 100
assert result["indexing_status"] == IndexingStatus.COMPLETED
assert result["segment_count"] == 5
assert result["hit_count"] == 10
# Assert
assert result["tenant_id"] == tenant_id
assert result["dataset_id"] == dataset_id
assert result["name"] == "test.pdf"
assert result["word_count"] == 100
assert result["indexing_status"] == IndexingStatus.COMPLETED
assert result["segment_count"] == 5
assert result["hit_count"] == 10

View File

@ -1,12 +1,12 @@
import importlib
import types
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import pytest
from sqlalchemy.orm import Session
from core.workflow.file_reference import build_file_reference
from graphon.file import FILE_MODEL_IDENTITY, FileTransferMethod
from models.model import Conversation, Message
from models import model as model_module
from models.model import App, AppMode, Conversation, IconType, Message
@pytest.fixture(autouse=True)
@ -14,10 +14,11 @@ def patch_file_helpers(monkeypatch: pytest.MonkeyPatch):
"""
Patch file_helpers.get_signed_file_url to a deterministic stub.
"""
model_module = importlib.import_module("models.model")
dummy = types.SimpleNamespace(get_signed_file_url=lambda fid: f"https://signed.example/{fid}")
# Inject/override file_helpers on models.model
monkeypatch.setattr(model_module, "file_helpers", dummy, raising=False)
monkeypatch.setattr(
model_module.file_helpers,
"get_signed_file_url",
lambda file_id: f"https://signed.example/{file_id}",
)
def _wrap_md(url: str) -> str:
@ -124,17 +125,43 @@ def test_inputs_restore_external_remote_url_file_mappings(owner_cls: type[Conver
assert restored_file.remote_url == "https://example.com/report.pdf"
def test_message_inputs_resolve_file_tenant_with_caller_session() -> None:
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_message_inputs_resolve_file_tenant_with_caller_session(sqlite_session: Session) -> None:
app = App(
id="app-1",
tenant_id="tenant-1",
name="File owner",
description="",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="file",
icon_background="#FFFFFF",
enable_site=False,
enable_api=False,
max_active_requests=0,
)
decoy = App(
id="other-app",
tenant_id="other-tenant",
name="Decoy",
description="",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="file",
icon_background="#FFFFFF",
enable_site=False,
enable_api=False,
max_active_requests=0,
)
sqlite_session.add_all([decoy, app])
sqlite_session.flush()
message = Message(app_id="app-1")
message.inputs = {"file": _build_local_file_mapping("upload-1")}
session = MagicMock()
session.scalar.return_value = "tenant-1"
with patch(
"models.model.build_file_from_input_mapping",
side_effect=lambda **kwargs: kwargs["tenant_resolver"](),
):
inputs = message.inputs_with_session(session=session)
inputs = message.inputs_with_session(session=sqlite_session)
assert inputs["file"] == "tenant-1"
session.scalar.assert_called_once()

View File

@ -3,6 +3,9 @@ import json
from unittest import mock
from uuid import uuid4
import pytest
from sqlalchemy.orm import Session
from constants import HIDDEN_VALUE
from core.helper import encrypter
from core.workflow.file_reference import build_file_reference
@ -12,6 +15,7 @@ from graphon.file import File, FileTransferMethod, FileType
from graphon.variables import FloatVariable, IntegerVariable, SecretVariable, StringVariable
from graphon.variables.segments import IntegerSegment, Segment
from models.account import Account
from models.tools import WorkflowToolProvider
from models.workflow import (
Workflow,
WorkflowDraftVariable,
@ -162,7 +166,14 @@ def test_to_dict():
assert workflow_dict["environment_variables"][1]["value"] == "text"
def test_workflow_account_getters_use_caller_session():
@pytest.mark.parametrize("sqlite_session", [(Workflow, Account)], indirect=True)
def test_workflow_account_getters_use_caller_session(sqlite_session: Session):
created_account = Account(name="Created Account", email="created@example.com")
created_account.id = "created-account-id"
updated_account = Account(name="Updated Account", email="updated@example.com")
updated_account.id = "updated-account-id"
decoy_account = Account(name="Decoy Account", email="decoy@example.com")
decoy_account.id = "decoy-account-id"
workflow = Workflow(
tenant_id="tenant_id",
app_id="app_id",
@ -175,23 +186,15 @@ def test_workflow_account_getters_use_caller_session():
conversation_variables=[],
updated_by="updated-account-id",
)
created_account = Account(name="Test Account", email="test@example.com")
updated_account = Account(name="Test Account", email="test@example.com")
session = mock.Mock()
session.get.side_effect = [created_account, updated_account]
sqlite_session.add_all([decoy_account, updated_account, workflow, created_account])
sqlite_session.flush()
with mock.patch("models.workflow.db") as mock_db:
assert workflow.get_created_by_account(session=session) is created_account
assert workflow.get_updated_by_account(session=session) is updated_account
assert session.get.call_args_list == [
mock.call(Account, "created-account-id"),
mock.call(Account, "updated-account-id"),
]
mock_db.session.get.assert_not_called()
assert workflow.get_created_by_account(session=sqlite_session) is created_account
assert workflow.get_updated_by_account(session=sqlite_session) is updated_account
def test_workflow_tool_published_getter_uses_caller_session():
@pytest.mark.parametrize("sqlite_session", [(Workflow, WorkflowToolProvider)], indirect=True)
def test_workflow_tool_published_getter_uses_caller_session(sqlite_session: Session):
workflow = Workflow(
tenant_id="tenant_id",
app_id="app_id",
@ -203,14 +206,30 @@ def test_workflow_tool_published_getter_uses_caller_session():
environment_variables=[],
conversation_variables=[],
)
session = mock.Mock()
session.execute.return_value.scalar_one.return_value = True
matching_provider = WorkflowToolProvider(
name="matching-provider",
label="Matching provider",
icon="tool",
app_id=workflow.app_id,
version="1",
user_id="account-id",
tenant_id=workflow.tenant_id,
description="Matching workflow tool",
)
decoy_provider = WorkflowToolProvider(
name="decoy-provider",
label="Decoy provider",
icon="tool",
app_id="other-app",
version="1",
user_id="account-id",
tenant_id=workflow.tenant_id,
description="Different app",
)
sqlite_session.add_all([decoy_provider, workflow, matching_provider])
sqlite_session.flush()
with mock.patch("models.workflow.db") as mock_db:
assert workflow.get_tool_published(session=session) is True
session.execute.assert_called_once()
mock_db.session.execute.assert_not_called()
assert workflow.get_tool_published(session=sqlite_session) is True
def test_normalize_environment_variable_mappings_converts_full_mask_to_hidden_value():