mirror of
https://github.com/langgenius/dify.git
synced 2026-09-05 00:31:19 +08:00
test: migrate residual model sessions and ORM models to SQLite (#40531)
This commit is contained in:
parent
303ce70ac4
commit
e2e3521874
@ -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",
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 == []
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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")
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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():
|
||||
|
||||
Loading…
Reference in New Issue
Block a user