mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
test: migrate snippet and vector sessions to SQLite (#40089)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Byron Wang <byron@dify.ai>
This commit is contained in:
co-authored by
autofix-ci[bot]
Byron Wang
parent
cd0e88c680
commit
a64a24b39c
@@ -4,6 +4,7 @@ from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.workflow.snippet_start import SNIPPET_VIRTUAL_START_NODE_ID
|
||||
from models.workflow import Workflow, WorkflowKind, WorkflowType
|
||||
@@ -277,7 +278,10 @@ def test_run_draft_node_raises_when_draft_workflow_missing(monkeypatch: pytest.M
|
||||
)
|
||||
|
||||
|
||||
def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch: pytest.MonkeyPatch):
|
||||
def test_generate_single_iteration_delegates_to_workflow_generator(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
workflow = _workflow({"nodes": [{"id": "iteration-1", "data": {"type": "iteration"}}], "edges": []})
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
user = SimpleNamespace(id="user-1")
|
||||
@@ -292,13 +296,12 @@ def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch:
|
||||
)
|
||||
monkeypatch.setattr("services.snippet_generate_service.WorkflowAppGenerator", workflow_generator_class)
|
||||
|
||||
session = Mock()
|
||||
result = SnippetGenerateService.generate_single_iteration(
|
||||
snippet=snippet,
|
||||
user=user,
|
||||
node_id="iteration-1",
|
||||
args={"inputs": {"items": [1]}},
|
||||
session_maker=_session_maker(session),
|
||||
session_maker=sqlite_session_factory,
|
||||
)
|
||||
|
||||
assert list(result) == ["event"]
|
||||
@@ -309,7 +312,7 @@ def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch:
|
||||
assert kwargs["node_id"] == "iteration-1"
|
||||
assert kwargs["user"] is user
|
||||
assert kwargs["streaming"] is True
|
||||
assert kwargs["session"] is session
|
||||
assert isinstance(kwargs["session"], Session)
|
||||
workflow_generator_class.convert_to_event_stream.assert_called_once_with(response)
|
||||
|
||||
|
||||
@@ -329,7 +332,10 @@ def test_generate_single_iteration_raises_when_draft_workflow_missing(monkeypatc
|
||||
)
|
||||
|
||||
|
||||
def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch: pytest.MonkeyPatch):
|
||||
def test_generate_single_loop_delegates_to_workflow_generator(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
workflow = _workflow({"nodes": [{"id": "loop-1", "data": {"type": "loop"}}], "edges": []})
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
user = SimpleNamespace(id="user-1")
|
||||
@@ -344,13 +350,12 @@ def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch: pytes
|
||||
)
|
||||
monkeypatch.setattr("services.snippet_generate_service.WorkflowAppGenerator", workflow_generator_class)
|
||||
|
||||
session = Mock()
|
||||
result = SnippetGenerateService.generate_single_loop(
|
||||
snippet=snippet,
|
||||
user=user,
|
||||
node_id="loop-1",
|
||||
args=SimpleNamespace(inputs={"items": [1]}),
|
||||
session_maker=_session_maker(session),
|
||||
session_maker=sqlite_session_factory,
|
||||
)
|
||||
|
||||
assert list(result) == ["event"]
|
||||
@@ -361,7 +366,7 @@ def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch: pytes
|
||||
assert kwargs["node_id"] == "loop-1"
|
||||
assert kwargs["user"] is user
|
||||
assert kwargs["streaming"] is True
|
||||
assert kwargs["session"] is session
|
||||
assert isinstance(kwargs["session"], Session)
|
||||
workflow_generator_class.convert_to_event_stream.assert_called_once_with(response)
|
||||
|
||||
|
||||
|
||||
@@ -1,42 +1,33 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from enums import DeploymentEdition
|
||||
from models.snippet import SnippetType
|
||||
from models.workflow import Workflow, WorkflowKind, WorkflowType
|
||||
from extensions.storage.storage_type import StorageType
|
||||
from graphon.variables.segments import StringSegment
|
||||
from graphon.variables.types import SegmentType
|
||||
from models.agent import Agent, AgentScope, AgentSource, AgentStatus
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import UploadFile
|
||||
from models.snippet import CustomizedSnippet, SnippetType
|
||||
from models.workflow import (
|
||||
Workflow,
|
||||
WorkflowDraftVariable,
|
||||
WorkflowDraftVariableFile,
|
||||
WorkflowKind,
|
||||
WorkflowType,
|
||||
)
|
||||
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
|
||||
from services.snippet_service import SnippetService
|
||||
|
||||
|
||||
class _SessionWithoutNameLookup:
|
||||
def __init__(self) -> None:
|
||||
self.add = Mock()
|
||||
self.commit = Mock()
|
||||
|
||||
def query(self, *args, **kwargs):
|
||||
raise AssertionError("snippet name uniqueness lookup should not be used")
|
||||
|
||||
|
||||
class _SessionContext:
|
||||
def __init__(self, session) -> None:
|
||||
self._session = session
|
||||
|
||||
def __enter__(self):
|
||||
return self._session
|
||||
|
||||
def __exit__(self, *args) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _session_maker(session):
|
||||
return lambda: _SessionContext(session)
|
||||
|
||||
|
||||
def _create_workflow(*, workflow_id: str, version: str, graph: dict, features: dict) -> Workflow:
|
||||
return Workflow(
|
||||
id=workflow_id,
|
||||
@@ -54,13 +45,26 @@ def _create_workflow(*, workflow_id: str, version: str, graph: dict, features: d
|
||||
)
|
||||
|
||||
|
||||
def test_create_snippet_allows_duplicate_names(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
session = _SessionWithoutNameLookup()
|
||||
account = SimpleNamespace(id="account-1")
|
||||
def _snippet() -> CustomizedSnippet:
|
||||
return CustomizedSnippet(
|
||||
id="snippet-1",
|
||||
tenant_id="tenant-1",
|
||||
name="Snippet",
|
||||
description="",
|
||||
type=SnippetType.NODE,
|
||||
created_by="account-1",
|
||||
)
|
||||
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session = None
|
||||
service._session_maker = _session_maker(session)
|
||||
|
||||
def test_create_snippet_allows_duplicate_names(
|
||||
sqlite_session_factory: sessionmaker[Session], sqlite_session: Session
|
||||
) -> None:
|
||||
account = SimpleNamespace(id="account-1")
|
||||
existing = _snippet()
|
||||
existing.name = "shared name"
|
||||
sqlite_session.add(existing)
|
||||
sqlite_session.commit()
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
|
||||
snippet = service.create_snippet(
|
||||
tenant_id="tenant-1",
|
||||
@@ -73,8 +77,12 @@ def test_create_snippet_allows_duplicate_names(monkeypatch: pytest.MonkeyPatch)
|
||||
)
|
||||
|
||||
assert snippet.name == "shared name"
|
||||
session.add.assert_called_once_with(snippet)
|
||||
session.commit.assert_called_once()
|
||||
stored = sqlite_session.scalars(
|
||||
select(CustomizedSnippet).where(
|
||||
CustomizedSnippet.tenant_id == "tenant-1", CustomizedSnippet.name == "shared name"
|
||||
)
|
||||
).all()
|
||||
assert {item.id for item in stored} == {existing.id, snippet.id}
|
||||
|
||||
|
||||
def test_validate_snippet_graph_forbidden_nodes_ignores_malformed_nodes() -> None:
|
||||
@@ -95,30 +103,35 @@ def test_validate_snippet_graph_forbidden_nodes_raises_with_node_details() -> No
|
||||
SnippetService.validate_snippet_graph_forbidden_nodes({"nodes": [{"id": "start-1", "data": {"type": "start"}}]})
|
||||
|
||||
|
||||
def test_get_snippets_returns_empty_when_tag_filter_has_no_targets(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
session = _SessionWithoutNameLookup()
|
||||
def test_get_snippets_returns_empty_when_tag_filter_has_no_targets(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
get_target_ids = Mock(return_value=[])
|
||||
monkeypatch.setattr("services.snippet_service.TagService.get_target_ids_by_tag_ids", get_target_ids)
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
|
||||
result = service.get_snippets(tenant_id="tenant-1", session=session, tag_ids=["tag-1"])
|
||||
result = service.get_snippets(tenant_id="tenant-1", session=sqlite_session, tag_ids=["tag-1"])
|
||||
|
||||
assert result == ([], 0, False)
|
||||
get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], session, match_all=True)
|
||||
get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], sqlite_session, match_all=True)
|
||||
|
||||
|
||||
def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
snippets = [
|
||||
SimpleNamespace(id="snippet-1"),
|
||||
SimpleNamespace(id="snippet-2"),
|
||||
SimpleNamespace(id="snippet-3"),
|
||||
]
|
||||
session = SimpleNamespace(
|
||||
scalar=Mock(return_value=3),
|
||||
scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=snippets))),
|
||||
)
|
||||
def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
snippets = []
|
||||
for index in range(3):
|
||||
snippet = CustomizedSnippet(
|
||||
id=f"snippet-{index + 1}",
|
||||
tenant_id="tenant-1",
|
||||
name=f"search {index}",
|
||||
description="search result",
|
||||
type=SnippetType.NODE,
|
||||
created_by="account-1",
|
||||
is_published=True,
|
||||
)
|
||||
sqlite_session.add(snippet)
|
||||
snippets.append(snippet)
|
||||
sqlite_session.flush()
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session_maker = _session_maker(session)
|
||||
get_target_ids = Mock(return_value=["snippet-1", "snippet-2", "snippet-3"])
|
||||
monkeypatch.setattr(
|
||||
"services.snippet_service.TagService.get_target_ids_by_tag_ids",
|
||||
@@ -127,7 +140,7 @@ def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPa
|
||||
|
||||
result, total, has_more = service.get_snippets(
|
||||
tenant_id="tenant-1",
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
page=2,
|
||||
limit=2,
|
||||
keyword="search",
|
||||
@@ -136,26 +149,23 @@ def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPa
|
||||
tag_ids=["tag-1"],
|
||||
)
|
||||
|
||||
assert result == snippets[:2]
|
||||
assert {snippet.id for snippet in result} <= {snippet.id for snippet in snippets}
|
||||
assert len(result) == 1
|
||||
assert total == 3
|
||||
assert has_more is True
|
||||
get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], session, match_all=True)
|
||||
session.scalar.assert_called_once()
|
||||
session.scalars.assert_called_once()
|
||||
assert has_more is False
|
||||
get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], sqlite_session, match_all=True)
|
||||
|
||||
|
||||
def test_update_snippet_allows_duplicate_names() -> None:
|
||||
session = _SessionWithoutNameLookup()
|
||||
snippet = SimpleNamespace(
|
||||
id="snippet-1",
|
||||
tenant_id="tenant-1",
|
||||
name="old name",
|
||||
description="",
|
||||
icon_info=None,
|
||||
def test_update_snippet_allows_duplicate_names(sqlite_session: Session) -> None:
|
||||
snippet = _snippet()
|
||||
other = CustomizedSnippet(
|
||||
id="snippet-2", tenant_id="tenant-1", name="shared name", description="", type=SnippetType.NODE
|
||||
)
|
||||
sqlite_session.add_all([snippet, other])
|
||||
sqlite_session.flush()
|
||||
|
||||
result = SnippetService.update_snippet(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=snippet,
|
||||
account_id="account-1",
|
||||
data={"name": "shared name"},
|
||||
@@ -163,21 +173,18 @@ def test_update_snippet_allows_duplicate_names() -> None:
|
||||
|
||||
assert result is snippet
|
||||
assert snippet.name == "shared name"
|
||||
session.add.assert_called_once_with(snippet)
|
||||
sqlite_session.flush()
|
||||
assert sqlite_session.get(CustomizedSnippet, snippet.id).name == "shared name"
|
||||
|
||||
|
||||
def test_update_snippet_updates_optional_fields() -> None:
|
||||
session = _SessionWithoutNameLookup()
|
||||
snippet = SimpleNamespace(
|
||||
id="snippet-1",
|
||||
tenant_id="tenant-1",
|
||||
name="old name",
|
||||
description="old description",
|
||||
icon_info=None,
|
||||
)
|
||||
def test_update_snippet_updates_optional_fields(sqlite_session: Session) -> None:
|
||||
snippet = _snippet()
|
||||
snippet.description = "old description"
|
||||
sqlite_session.add(snippet)
|
||||
sqlite_session.flush()
|
||||
|
||||
result = SnippetService.update_snippet(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=snippet,
|
||||
account_id="account-1",
|
||||
data={"description": "new description", "icon_info": {"icon": "star"}},
|
||||
@@ -187,23 +194,20 @@ def test_update_snippet_updates_optional_fields() -> None:
|
||||
assert snippet.description == "new description"
|
||||
assert snippet.icon_info == {"icon": "star"}
|
||||
assert snippet.updated_by == "account-1"
|
||||
session.add.assert_called_once_with(snippet)
|
||||
sqlite_session.flush()
|
||||
stored = sqlite_session.get(CustomizedSnippet, snippet.id)
|
||||
assert stored is not None
|
||||
assert stored.description == "new description"
|
||||
|
||||
|
||||
def test_sync_draft_workflow_creates_draft_and_updates_input_fields(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session = None
|
||||
def test_sync_draft_workflow_creates_draft_and_updates_input_fields(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=None))
|
||||
session = Mock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
service._session_maker = _session_maker(session)
|
||||
snippet = SimpleNamespace(
|
||||
id="snippet-1",
|
||||
tenant_id="tenant-1",
|
||||
input_fields=None,
|
||||
updated_by=None,
|
||||
updated_at=None,
|
||||
)
|
||||
snippet = _snippet()
|
||||
account = SimpleNamespace(id="account-1")
|
||||
|
||||
workflow = service.sync_draft_workflow(
|
||||
@@ -217,14 +221,18 @@ def test_sync_draft_workflow_creates_draft_and_updates_input_fields(monkeypatch:
|
||||
assert workflow.app_id == snippet.id
|
||||
assert workflow.kind == WorkflowKind.SNIPPET
|
||||
assert json.loads(snippet.input_fields) == [{"variable": "query"}]
|
||||
session.add.assert_any_call(workflow)
|
||||
session.add.assert_any_call(snippet)
|
||||
session.commit.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
stored_workflow = sqlite_session.scalar(select(Workflow).where(Workflow.id == workflow.id))
|
||||
stored_snippet = sqlite_session.get(CustomizedSnippet, snippet.id)
|
||||
assert stored_workflow is not None
|
||||
assert stored_snippet is not None
|
||||
assert stored_snippet.input_fields_list == [{"variable": "query"}]
|
||||
|
||||
|
||||
def test_sync_draft_workflow_raises_when_hash_mismatches() -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session_maker = _session_maker(SimpleNamespace(commit=Mock(), add=Mock()))
|
||||
def test_sync_draft_workflow_raises_when_hash_mismatches(
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
service.get_draft_workflow = Mock(return_value=SimpleNamespace(unique_hash="server-hash"))
|
||||
|
||||
with pytest.raises(WorkflowHashNotEqualError):
|
||||
@@ -236,9 +244,12 @@ def test_sync_draft_workflow_raises_when_hash_mismatches() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session = None
|
||||
def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
workflow = _create_workflow(
|
||||
workflow_id="workflow-1",
|
||||
version=Workflow.VERSION_DRAFT,
|
||||
@@ -246,19 +257,9 @@ def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(monkeyp
|
||||
features={},
|
||||
)
|
||||
unique_hash = workflow.unique_hash
|
||||
snippet = SimpleNamespace(
|
||||
id="snippet-1",
|
||||
tenant_id="tenant-1",
|
||||
input_fields=None,
|
||||
updated_by=None,
|
||||
updated_at=None,
|
||||
)
|
||||
snippet = _snippet()
|
||||
account = SimpleNamespace(id="account-1")
|
||||
session = Mock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
|
||||
monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=workflow))
|
||||
service._session_maker = _session_maker(session)
|
||||
|
||||
result = service.sync_draft_workflow(
|
||||
snippet=snippet,
|
||||
@@ -276,18 +277,26 @@ def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(monkeyp
|
||||
assert workflow.environment_variables == []
|
||||
assert workflow.conversation_variables == []
|
||||
assert json.loads(snippet.input_fields) == [{"variable": "query"}]
|
||||
session.commit.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(Workflow, workflow.id) is not None
|
||||
assert sqlite_session.get(CustomizedSnippet, snippet.id) is not None
|
||||
|
||||
|
||||
def test_update_workflow_updates_marked_fields() -> None:
|
||||
def test_update_workflow_updates_marked_fields(sqlite_session: Session) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
workflow = SimpleNamespace(marked_name="", marked_comment="", updated_by=None, updated_at=None)
|
||||
session = SimpleNamespace(scalar=Mock(return_value=workflow), add=Mock())
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
workflow = _create_workflow(
|
||||
workflow_id="workflow-1",
|
||||
version="2026-01-01 00:00:00",
|
||||
graph={"nodes": []},
|
||||
features={},
|
||||
)
|
||||
snippet = _snippet()
|
||||
sqlite_session.add_all([snippet, workflow])
|
||||
sqlite_session.flush()
|
||||
account = SimpleNamespace(id="account-1")
|
||||
|
||||
result = service.update_workflow(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=snippet,
|
||||
workflow_id="workflow-1",
|
||||
account=account,
|
||||
@@ -298,16 +307,17 @@ def test_update_workflow_updates_marked_fields() -> None:
|
||||
assert workflow.marked_name == "v1"
|
||||
assert workflow.marked_comment == "first version"
|
||||
assert workflow.updated_by == "account-1"
|
||||
session.scalar.assert_called_once()
|
||||
session.add.assert_called_once_with(workflow)
|
||||
sqlite_session.flush()
|
||||
stored = sqlite_session.get(Workflow, workflow.id)
|
||||
assert stored is not None
|
||||
assert stored.marked_name == "v1"
|
||||
|
||||
|
||||
def test_update_workflow_returns_none_when_missing() -> None:
|
||||
def test_update_workflow_returns_none_when_missing(sqlite_session: Session) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
session = SimpleNamespace(scalar=Mock(return_value=None), add=Mock())
|
||||
|
||||
result = service.update_workflow(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"),
|
||||
workflow_id="missing-workflow",
|
||||
account=SimpleNamespace(id="account-1"),
|
||||
@@ -315,7 +325,6 @@ def test_update_workflow_returns_none_when_missing() -> None:
|
||||
)
|
||||
|
||||
assert result is None
|
||||
session.add.assert_not_called()
|
||||
|
||||
|
||||
def test_get_default_block_configs_skips_empty_defaults(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@@ -362,8 +371,10 @@ def test_get_default_block_config_returns_none_for_empty_default(monkeypatch: py
|
||||
|
||||
def test_restore_published_snippet_workflow_to_draft_copies_source_snapshot(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
snippet = _snippet()
|
||||
account = SimpleNamespace(id="account-2")
|
||||
source_graph = {"nodes": [{"id": "llm-1", "data": {"type": "llm"}}], "edges": []}
|
||||
source_features = {"opening_statement": "hello"}
|
||||
@@ -379,11 +390,7 @@ def test_restore_published_snippet_workflow_to_draft_copies_source_snapshot(
|
||||
graph={"nodes": [], "edges": []},
|
||||
features={},
|
||||
)
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session = None
|
||||
session = Mock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
service._session_maker = _session_maker(session)
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
|
||||
monkeypatch.setattr(service, "get_published_workflow_by_id", Mock(return_value=source_workflow))
|
||||
monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=draft_workflow))
|
||||
@@ -398,17 +405,19 @@ def test_restore_published_snippet_workflow_to_draft_copies_source_snapshot(
|
||||
assert draft_workflow.graph_dict == source_graph
|
||||
assert draft_workflow.features_dict == source_features
|
||||
assert draft_workflow.updated_by == account.id
|
||||
session.add.assert_called_once_with(draft_workflow)
|
||||
session.commit.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
stored = sqlite_session.get(Workflow, draft_workflow.id)
|
||||
assert stored is not None
|
||||
assert stored.graph_dict == source_graph
|
||||
|
||||
|
||||
def test_restore_published_snippet_workflow_to_draft_raises_when_source_missing(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
account = SimpleNamespace(id="account-2")
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session_maker = _session_maker(SimpleNamespace(add=Mock(), commit=Mock()))
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
|
||||
monkeypatch.setattr(service, "get_published_workflow_by_id", Mock(return_value=None))
|
||||
|
||||
@@ -420,8 +429,12 @@ def test_restore_published_snippet_workflow_to_draft_raises_when_source_missing(
|
||||
)
|
||||
|
||||
|
||||
def test_restore_published_snippet_workflow_to_draft_adds_new_draft(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
def test_restore_published_snippet_workflow_to_draft_adds_new_draft(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
snippet = _snippet()
|
||||
account = SimpleNamespace(id="account-2")
|
||||
source_workflow = _create_workflow(
|
||||
workflow_id="published-workflow",
|
||||
@@ -435,11 +448,7 @@ def test_restore_published_snippet_workflow_to_draft_adds_new_draft(monkeypatch:
|
||||
graph={"nodes": [], "edges": []},
|
||||
features={},
|
||||
)
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session = None
|
||||
session = Mock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
service._session_maker = _session_maker(session)
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
|
||||
monkeypatch.setattr(service, "get_published_workflow_by_id", Mock(return_value=source_workflow))
|
||||
monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=None))
|
||||
@@ -455,8 +464,8 @@ def test_restore_published_snippet_workflow_to_draft_adds_new_draft(monkeypatch:
|
||||
)
|
||||
|
||||
assert result is new_draft_workflow
|
||||
session.add.assert_called_once_with(new_draft_workflow)
|
||||
session.commit.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(Workflow, new_draft_workflow.id) is not None
|
||||
|
||||
|
||||
def test_get_published_workflow_returns_none_without_workflow_id() -> None:
|
||||
@@ -467,12 +476,15 @@ def test_get_published_workflow_returns_none_without_workflow_id() -> None:
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_get_published_workflow_by_id_raises_for_draft(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
draft_workflow = SimpleNamespace(version=Workflow.VERSION_DRAFT)
|
||||
session = SimpleNamespace(scalar=Mock(return_value=draft_workflow))
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
service._session = None
|
||||
service._session_maker = _session_maker(session)
|
||||
def test_get_published_workflow_by_id_raises_for_draft(
|
||||
sqlite_session_factory: sessionmaker[Session], sqlite_session: Session
|
||||
) -> None:
|
||||
draft_workflow = _create_workflow(
|
||||
workflow_id="workflow-1", version=Workflow.VERSION_DRAFT, graph={"nodes": []}, features={}
|
||||
)
|
||||
sqlite_session.add(draft_workflow)
|
||||
sqlite_session.commit()
|
||||
service = SnippetService(session_maker=sqlite_session_factory)
|
||||
|
||||
with pytest.raises(IsDraftWorkflowError):
|
||||
service.get_published_workflow_by_id(
|
||||
@@ -481,19 +493,20 @@ def test_get_published_workflow_by_id_raises_for_draft(monkeypatch: pytest.Monke
|
||||
)
|
||||
|
||||
|
||||
def test_publish_workflow_raises_when_draft_missing() -> None:
|
||||
def test_publish_workflow_raises_when_draft_missing(sqlite_session: Session) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
session = SimpleNamespace(scalar=Mock(return_value=None))
|
||||
|
||||
with pytest.raises(ValueError, match="No valid workflow found"):
|
||||
service.publish_workflow(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"),
|
||||
account=SimpleNamespace(id="account-1"),
|
||||
)
|
||||
|
||||
|
||||
def test_publish_workflow_creates_snapshot_and_updates_snippet(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_publish_workflow_creates_snapshot_and_updates_snippet(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
draft_workflow = _create_workflow(
|
||||
workflow_id="draft-workflow",
|
||||
@@ -513,22 +526,16 @@ def test_publish_workflow_creates_snapshot_and_updates_snippet(monkeypatch: pyte
|
||||
},
|
||||
features={"opening_statement": "hello"},
|
||||
)
|
||||
snippet = SimpleNamespace(
|
||||
id="snippet-1",
|
||||
tenant_id="tenant-1",
|
||||
version=1,
|
||||
is_published=False,
|
||||
workflow_id=None,
|
||||
updated_by=None,
|
||||
)
|
||||
session = SimpleNamespace(scalar=Mock(return_value=draft_workflow), add=Mock())
|
||||
snippet = _snippet()
|
||||
sqlite_session.add_all([draft_workflow, snippet])
|
||||
sqlite_session.flush()
|
||||
monkeypatch.setattr(
|
||||
"services.agent.workflow_publish_service.WorkflowAgentPublishService.copy_agent_node_bindings_to_published",
|
||||
Mock(return_value=set()),
|
||||
)
|
||||
|
||||
result, retirement_candidates = service.publish_workflow(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=snippet,
|
||||
account=SimpleNamespace(id="account-1"),
|
||||
)
|
||||
@@ -538,15 +545,17 @@ def test_publish_workflow_creates_snapshot_and_updates_snippet(monkeypatch: pyte
|
||||
assert snippet.is_published is True
|
||||
assert snippet.workflow_id == result.id
|
||||
assert snippet.updated_by == "account-1"
|
||||
assert session.add.call_args_list[-1].args == (snippet,)
|
||||
sqlite_session.flush()
|
||||
assert sqlite_session.get(Workflow, result.id) is result
|
||||
assert sqlite_session.get(CustomizedSnippet, snippet.id).workflow_id == result.id
|
||||
assert retirement_candidates == set()
|
||||
|
||||
|
||||
def test_get_all_published_workflows_returns_empty_without_current_workflow() -> None:
|
||||
def test_get_all_published_workflows_returns_empty_without_current_workflow(unbound_session: Session) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
|
||||
result = service.get_all_published_workflows(
|
||||
session=SimpleNamespace(),
|
||||
session=unbound_session,
|
||||
snippet=SimpleNamespace(id="snippet-1", workflow_id=None),
|
||||
page=1,
|
||||
limit=20,
|
||||
@@ -555,73 +564,75 @@ def test_get_all_published_workflows_returns_empty_without_current_workflow() ->
|
||||
assert result == ([], False)
|
||||
|
||||
|
||||
def test_get_all_published_workflows_paginates() -> None:
|
||||
def test_get_all_published_workflows_paginates(sqlite_session: Session) -> None:
|
||||
service = SnippetService.__new__(SnippetService)
|
||||
workflows = [SimpleNamespace(id="workflow-1"), SimpleNamespace(id="workflow-2"), SimpleNamespace(id="workflow-3")]
|
||||
session = SimpleNamespace(scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=workflows))))
|
||||
workflows = [
|
||||
_create_workflow(
|
||||
workflow_id=f"workflow-{index}",
|
||||
version=f"2026-01-0{index} 00:00:00",
|
||||
graph={"nodes": []},
|
||||
features={},
|
||||
)
|
||||
for index in range(1, 4)
|
||||
]
|
||||
sqlite_session.add_all(workflows)
|
||||
sqlite_session.flush()
|
||||
|
||||
result, has_more = service.get_all_published_workflows(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=SimpleNamespace(id="snippet-1", workflow_id="workflow-current"),
|
||||
page=1,
|
||||
limit=2,
|
||||
)
|
||||
|
||||
assert result == workflows[:2]
|
||||
assert [workflow.id for workflow in result] == ["workflow-3", "workflow-2"]
|
||||
assert has_more is True
|
||||
session.scalars.assert_called_once()
|
||||
|
||||
|
||||
def test_delete_snippet_removes_related_records() -> None:
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
session = SimpleNamespace(
|
||||
execute=Mock(),
|
||||
scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=[]))),
|
||||
delete=Mock(),
|
||||
def test_delete_snippet_removes_related_records(
|
||||
sqlite_session: Session, sqlite_session_factory: sessionmaker[Session]
|
||||
) -> None:
|
||||
snippet = _snippet()
|
||||
workflow = _create_workflow(
|
||||
workflow_id="workflow-1", version=Workflow.VERSION_DRAFT, graph={"nodes": []}, features={}
|
||||
)
|
||||
sqlite_session.add_all([snippet, workflow])
|
||||
sqlite_session.flush()
|
||||
|
||||
result = SnippetService.delete_snippet(session=session, snippet=snippet)
|
||||
result = SnippetService.delete_snippet(session=sqlite_session, snippet=snippet)
|
||||
|
||||
assert result is True
|
||||
executed_sql = "\n".join(str(call.args[0]) for call in session.execute.call_args_list)
|
||||
assert "workflow_draft_variables" in executed_sql
|
||||
assert "tool_workflow_providers" in executed_sql
|
||||
assert "workflow_app_logs" in executed_sql
|
||||
assert "workflow_archive_logs" in executed_sql
|
||||
assert "workflow_node_executions" in executed_sql
|
||||
assert "workflow_runs" in executed_sql
|
||||
assert "workflows" in executed_sql
|
||||
assert "kind" in executed_sql
|
||||
assert "tag_bindings" in executed_sql
|
||||
session.delete.assert_called_once_with(snippet)
|
||||
sqlite_session.commit()
|
||||
with sqlite_session_factory() as observer:
|
||||
assert observer.get(CustomizedSnippet, snippet.id) is None
|
||||
assert observer.get(Workflow, workflow.id) is None
|
||||
|
||||
|
||||
def test_delete_snippet_archives_owned_agents_and_schedules_backing_app_cleanup(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
agent = SimpleNamespace(
|
||||
snippet = _snippet()
|
||||
agent = Agent(
|
||||
id="agent-1",
|
||||
tenant_id=snippet.tenant_id,
|
||||
name="Snippet agent",
|
||||
description="",
|
||||
role="",
|
||||
scope=AgentScope.WORKFLOW_ONLY,
|
||||
source=AgentSource.WORKFLOW,
|
||||
app_id=snippet.id,
|
||||
backing_app_id="backing-app-1",
|
||||
status="active",
|
||||
archived_by=None,
|
||||
archived_at=None,
|
||||
status=AgentStatus.ACTIVE,
|
||||
updated_by="creator-1",
|
||||
updated_at=None,
|
||||
)
|
||||
scalar_results = [
|
||||
SimpleNamespace(all=Mock(return_value=[])),
|
||||
SimpleNamespace(all=Mock(return_value=[agent])),
|
||||
]
|
||||
session = SimpleNamespace(
|
||||
execute=Mock(),
|
||||
scalars=Mock(side_effect=scalar_results),
|
||||
delete=Mock(),
|
||||
)
|
||||
listen = Mock()
|
||||
monkeypatch.setattr("services.snippet_service.event.listen", listen)
|
||||
sqlite_session.add_all([snippet, agent])
|
||||
sqlite_session.flush()
|
||||
cleanup_delay = Mock()
|
||||
monkeypatch.setattr("tasks.remove_app_and_related_data_task.remove_app_and_related_data_task.delay", cleanup_delay)
|
||||
|
||||
result = SnippetService.delete_snippet(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
snippet=snippet,
|
||||
account_id="account-1",
|
||||
)
|
||||
@@ -631,34 +642,62 @@ def test_delete_snippet_archives_owned_agents_and_schedules_backing_app_cleanup(
|
||||
assert agent.archived_by == "account-1"
|
||||
assert agent.archived_at is not None
|
||||
assert agent.updated_by == "account-1"
|
||||
executed_sql = "\n".join(str(call.args[0]) for call in session.execute.call_args_list)
|
||||
assert "DELETE FROM apps" in executed_sql
|
||||
listen.assert_called_once_with(session, "after_commit", listen.call_args.args[2], once=True)
|
||||
sqlite_session.commit()
|
||||
assert sqlite_session.get(Agent, agent.id).status == AgentStatus.ARCHIVED
|
||||
cleanup_delay.assert_called_once_with(tenant_id=snippet.tenant_id, app_id="backing-app-1")
|
||||
|
||||
|
||||
def test_delete_draft_variable_files_removes_storage_objects(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_delete_draft_variable_files_removes_storage_objects(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
from extensions.ext_storage import storage
|
||||
|
||||
snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1")
|
||||
snippet = _snippet()
|
||||
storage_delete = Mock()
|
||||
monkeypatch.setattr(storage, "delete", storage_delete)
|
||||
session = SimpleNamespace(
|
||||
scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=["file-1"]))),
|
||||
execute=Mock(
|
||||
side_effect=[
|
||||
SimpleNamespace(all=Mock(return_value=[("file-1", "upload-1", "storage-key")])),
|
||||
None,
|
||||
None,
|
||||
]
|
||||
),
|
||||
upload_file = UploadFile(
|
||||
tenant_id=snippet.tenant_id,
|
||||
storage_type=StorageType.LOCAL,
|
||||
key="storage-key",
|
||||
name="value.txt",
|
||||
size=10,
|
||||
extension=".txt",
|
||||
mime_type="text/plain",
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by="account-1",
|
||||
created_at=datetime(2025, 1, 1),
|
||||
used=True,
|
||||
)
|
||||
variable_file = WorkflowDraftVariableFile(
|
||||
tenant_id=snippet.tenant_id,
|
||||
app_id=snippet.id,
|
||||
user_id="account-1",
|
||||
upload_file_id=upload_file.id,
|
||||
size=10,
|
||||
length=None,
|
||||
value_type=SegmentType.STRING,
|
||||
)
|
||||
variable = WorkflowDraftVariable.new_node_variable(
|
||||
app_id=snippet.id,
|
||||
user_id="account-1",
|
||||
node_id="node-1",
|
||||
name="value",
|
||||
value=StringSegment(value="truncated"),
|
||||
node_execution_id="execution-1",
|
||||
file_id=variable_file.id,
|
||||
)
|
||||
sqlite_session.add_all([snippet, upload_file, variable_file, variable])
|
||||
sqlite_session.flush()
|
||||
|
||||
SnippetService._delete_draft_variable_files(session=session, snippet=snippet)
|
||||
SnippetService._delete_draft_variable_files(session=sqlite_session, snippet=snippet)
|
||||
|
||||
storage_delete.assert_called_once_with("storage-key")
|
||||
executed_sql = "\n".join(str(call.args[0]) for call in session.execute.call_args_list)
|
||||
assert "upload_files" in executed_sql
|
||||
assert "workflow_draft_variable_files" in executed_sql
|
||||
sqlite_session.commit()
|
||||
with sqlite_session_factory() as observer:
|
||||
assert observer.get(UploadFile, upload_file.id) is None
|
||||
assert observer.get(WorkflowDraftVariableFile, variable_file.id) is None
|
||||
|
||||
|
||||
def test_delete_archived_workflow_run_files_removes_prefixed_objects(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@@ -743,11 +782,14 @@ def test_workflow_run_node_executions_returns_empty_when_run_missing() -> None:
|
||||
service._node_execution_service_repo.get_executions_by_workflow_run.assert_not_called()
|
||||
|
||||
|
||||
def test_increment_use_count_adds_updated_snippet() -> None:
|
||||
snippet = SimpleNamespace(use_count=2)
|
||||
session = SimpleNamespace(add=Mock())
|
||||
def test_increment_use_count_adds_updated_snippet(sqlite_session: Session) -> None:
|
||||
snippet = _snippet()
|
||||
snippet.use_count = 2
|
||||
sqlite_session.add(snippet)
|
||||
sqlite_session.flush()
|
||||
|
||||
SnippetService.increment_use_count(session=session, snippet=snippet)
|
||||
SnippetService.increment_use_count(session=sqlite_session, snippet=snippet)
|
||||
|
||||
assert snippet.use_count == 3
|
||||
session.add.assert_called_once_with(snippet)
|
||||
sqlite_session.flush()
|
||||
assert sqlite_session.get(CustomizedSnippet, snippet.id).use_count == 3
|
||||
|
||||
@@ -4,22 +4,24 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import event
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import services.vector_service as vector_service_module
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
|
||||
from extensions.storage.storage_type import StorageType
|
||||
from models import UploadFile
|
||||
from models.dataset import ChildChunk, DatasetProcessRule, SegmentAttachmentBinding
|
||||
from models.dataset import Document as DatasetDocument
|
||||
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, ProcessRuleMode
|
||||
from services.vector_service import VectorService
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _UploadFileStub:
|
||||
id: str
|
||||
name: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ChildDocStub:
|
||||
page_content: str
|
||||
@@ -77,21 +79,27 @@ def _make_segment(
|
||||
return segment
|
||||
|
||||
|
||||
def _mock_db_session_for_update_multimodel(*, upload_files: list[_UploadFileStub] | None) -> MagicMock:
|
||||
session = MagicMock(name="session")
|
||||
|
||||
# db.session.execute() is used for delete(SegmentAttachmentBinding).where(...)
|
||||
session.execute = MagicMock(name="execute")
|
||||
|
||||
# db.session.scalars(select(UploadFile).where(...)).all() returns upload files
|
||||
session.scalars.return_value.all.return_value = upload_files or []
|
||||
|
||||
db_mock = MagicMock(name="db")
|
||||
db_mock.session = session
|
||||
return db_mock
|
||||
def _upload_file(*, file_id: str = "file-1", name: str = "img.png") -> UploadFile:
|
||||
upload_file = UploadFile(
|
||||
tenant_id="tenant-1",
|
||||
storage_type=StorageType.LOCAL,
|
||||
key=f"uploads/{file_id}",
|
||||
name=name,
|
||||
size=10,
|
||||
extension="png",
|
||||
mime_type="image/png",
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by="account-1",
|
||||
created_at=datetime(2026, 1, 1),
|
||||
used=False,
|
||||
)
|
||||
upload_file.id = file_id
|
||||
return upload_file
|
||||
|
||||
|
||||
def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(is_multimodal=False)
|
||||
segment = _make_segment()
|
||||
|
||||
@@ -101,7 +109,7 @@ def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(mo
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
VectorService.create_segments_vector(
|
||||
[["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=MagicMock()
|
||||
[["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=sqlite_session
|
||||
)
|
||||
|
||||
index_processor.load.assert_called_once()
|
||||
@@ -113,7 +121,9 @@ def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(mo
|
||||
assert kwargs["keywords_list"] == [["k1"]]
|
||||
|
||||
|
||||
def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_segments_vector_regular_indexing_loads_multimodal_documents(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(is_multimodal=True)
|
||||
segment = _make_segment(
|
||||
attachments=[
|
||||
@@ -127,9 +137,8 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk
|
||||
factory_instance.init_index_processor.return_value = index_processor
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.create_segments_vector(
|
||||
[["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=session
|
||||
[["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=sqlite_session
|
||||
)
|
||||
|
||||
assert index_processor.load.call_count == 2
|
||||
@@ -143,43 +152,62 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk
|
||||
assert second_args[1] == []
|
||||
assert len(second_args[2]) == 2
|
||||
assert second_kwargs["with_keywords"] is False
|
||||
segment.get_attachments.assert_called_once_with(session=session)
|
||||
segment.get_attachments.assert_called_once_with(session=sqlite_session)
|
||||
|
||||
|
||||
def test_create_segments_vector_with_no_segments_does_not_load(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_segments_vector_with_no_segments_does_not_load(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset()
|
||||
index_processor = MagicMock(name="index_processor")
|
||||
factory_instance = MagicMock()
|
||||
factory_instance.init_index_processor.return_value = index_processor
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, session=MagicMock())
|
||||
VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, session=sqlite_session)
|
||||
index_processor.load.assert_not_called()
|
||||
|
||||
|
||||
def _mock_parent_child_queries(
|
||||
def _persist_parent_child_rows(
|
||||
session: Session,
|
||||
*,
|
||||
dataset_document: object | None,
|
||||
processing_rule: object | None,
|
||||
) -> MagicMock:
|
||||
session = MagicMock(name="session")
|
||||
|
||||
get_dispatch: dict[object, object | None] = {
|
||||
vector_service_module.DatasetDocument: dataset_document,
|
||||
vector_service_module.DatasetProcessRule: processing_rule,
|
||||
}
|
||||
|
||||
def get_side_effect(model: object, pk: object) -> object | None:
|
||||
return get_dispatch.get(model)
|
||||
|
||||
session.get.side_effect = get_side_effect
|
||||
db_mock = MagicMock(name="db")
|
||||
db_mock.session = session
|
||||
return db_mock
|
||||
segment: MagicMock,
|
||||
include_document: bool = True,
|
||||
include_rule: bool = True,
|
||||
) -> tuple[DatasetDocument | None, DatasetProcessRule | None]:
|
||||
document = None
|
||||
rule = None
|
||||
if include_document:
|
||||
document = DatasetDocument(
|
||||
id=segment.document_id,
|
||||
tenant_id=segment.tenant_id,
|
||||
dataset_id=segment.dataset_id,
|
||||
position=1,
|
||||
data_source_type=DataSourceType.UPLOAD_FILE,
|
||||
dataset_process_rule_id="rule-1",
|
||||
batch="batch-1",
|
||||
name="Document",
|
||||
created_from=DocumentCreatedFrom.API,
|
||||
created_by="user-1",
|
||||
doc_language="en",
|
||||
)
|
||||
session.add(document)
|
||||
if include_rule:
|
||||
rule = DatasetProcessRule(
|
||||
dataset_id=segment.dataset_id,
|
||||
mode=ProcessRuleMode.HIERARCHICAL,
|
||||
rules='{"parent_mode":"full-doc"}',
|
||||
created_by="user-1",
|
||||
)
|
||||
rule.id = "rule-1"
|
||||
session.add(rule)
|
||||
session.flush()
|
||||
return document, rule
|
||||
|
||||
|
||||
def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_explicit_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(
|
||||
doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
@@ -188,16 +216,9 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex
|
||||
)
|
||||
segment = _make_segment()
|
||||
|
||||
dataset_document = MagicMock(name="dataset_document")
|
||||
dataset_document.id = segment.document_id
|
||||
dataset_document.dataset_process_rule_id = "rule-1"
|
||||
dataset_document.doc_language = "en"
|
||||
dataset_document.created_by = "user-1"
|
||||
|
||||
processing_rule = MagicMock(name="processing_rule")
|
||||
processing_rule.to_dict.return_value = {"rules": {}}
|
||||
|
||||
db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule)
|
||||
dataset_document, processing_rule = _persist_parent_child_rows(sqlite_session, segment=segment)
|
||||
assert dataset_document is not None
|
||||
assert processing_rule is not None
|
||||
|
||||
embedding_model_instance = MagicMock(name="embedding_model_instance")
|
||||
model_manager_instance = MagicMock(name="model_manager_instance")
|
||||
@@ -215,18 +236,23 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
VectorService.create_segments_vector(
|
||||
None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, session=db_mock.session
|
||||
None,
|
||||
[segment],
|
||||
dataset,
|
||||
vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
model_manager_instance.get_model_instance.assert_called_once()
|
||||
generate_child_chunks_mock.assert_called_once_with(
|
||||
segment, dataset_document, dataset, embedding_model_instance, processing_rule, False, session=db_mock.session
|
||||
segment, dataset_document, dataset, embedding_model_instance, processing_rule, False, session=sqlite_session
|
||||
)
|
||||
index_processor.load.assert_not_called()
|
||||
|
||||
|
||||
def test_create_segments_vector_parent_child_uses_default_embedding_model_when_provider_missing(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(
|
||||
doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
@@ -235,15 +261,7 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p
|
||||
)
|
||||
segment = _make_segment()
|
||||
|
||||
dataset_document = MagicMock()
|
||||
dataset_document.dataset_process_rule_id = "rule-1"
|
||||
dataset_document.doc_language = "en"
|
||||
dataset_document.created_by = "user-1"
|
||||
|
||||
processing_rule = MagicMock()
|
||||
processing_rule.to_dict.return_value = {"rules": {}}
|
||||
|
||||
db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule)
|
||||
_persist_parent_child_rows(sqlite_session, segment=segment)
|
||||
|
||||
embedding_model_instance = MagicMock()
|
||||
model_manager_instance = MagicMock()
|
||||
@@ -261,7 +279,11 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
VectorService.create_segments_vector(
|
||||
None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, session=db_mock.session
|
||||
None,
|
||||
[segment],
|
||||
dataset,
|
||||
vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
model_manager_instance.get_default_model_instance.assert_called_once()
|
||||
@@ -271,13 +293,11 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p
|
||||
def test_create_segments_vector_parent_child_missing_document_logs_warning_and_continues(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX)
|
||||
segment = _make_segment()
|
||||
|
||||
processing_rule = MagicMock()
|
||||
db_mock = _mock_parent_child_queries(dataset_document=None, processing_rule=processing_rule)
|
||||
|
||||
index_processor = MagicMock()
|
||||
factory_instance = MagicMock()
|
||||
factory_instance.init_index_processor.return_value = index_processor
|
||||
@@ -289,19 +309,19 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c
|
||||
[segment],
|
||||
dataset,
|
||||
vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
session=db_mock.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
assert any(r.levelno >= logging.WARNING for r in caplog.records)
|
||||
index_processor.load.assert_not_called()
|
||||
|
||||
|
||||
def test_create_segments_vector_parent_child_missing_processing_rule_raises(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_segments_vector_parent_child_missing_processing_rule_raises(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX)
|
||||
segment = _make_segment()
|
||||
|
||||
dataset_document = MagicMock()
|
||||
dataset_document.dataset_process_rule_id = "rule-1"
|
||||
db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=None)
|
||||
_persist_parent_child_rows(sqlite_session, segment=segment, include_rule=False)
|
||||
|
||||
with pytest.raises(ValueError, match="No processing rule found"):
|
||||
VectorService.create_segments_vector(
|
||||
@@ -309,20 +329,19 @@ def test_create_segments_vector_parent_child_missing_processing_rule_raises(monk
|
||||
[segment],
|
||||
dataset,
|
||||
vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
session=db_mock.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
|
||||
def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_segments_vector_parent_child_non_high_quality_raises(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(
|
||||
doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
indexing_technique=IndexTechniqueType.ECONOMY,
|
||||
)
|
||||
segment = _make_segment()
|
||||
dataset_document = MagicMock()
|
||||
dataset_document.dataset_process_rule_id = "rule-1"
|
||||
processing_rule = MagicMock()
|
||||
db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule)
|
||||
_persist_parent_child_rows(sqlite_session, segment=segment)
|
||||
|
||||
with pytest.raises(ValueError, match="not high quality"):
|
||||
VectorService.create_segments_vector(
|
||||
@@ -330,11 +349,13 @@ def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch
|
||||
[segment],
|
||||
dataset,
|
||||
vector_service_module.IndexStructureType.PARENT_CHILD_INDEX,
|
||||
session=db_mock.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
|
||||
def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_segment_vector_high_quality_uses_vector(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY)
|
||||
segment = _make_segment()
|
||||
|
||||
@@ -342,10 +363,9 @@ def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.Monk
|
||||
vector_cls = MagicMock(return_value=vector_instance)
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.update_segment_vector(["k"], segment, dataset, session=session)
|
||||
VectorService.update_segment_vector(["k"], segment, dataset, session=sqlite_session)
|
||||
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session)
|
||||
vector_instance.delete_by_ids.assert_called_once_with([segment.index_node_id])
|
||||
vector_instance.add_texts.assert_called_once()
|
||||
add_args, add_kwargs = vector_instance.add_texts.call_args
|
||||
@@ -353,41 +373,45 @@ def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.Monk
|
||||
assert add_kwargs["duplicate_check"] is True
|
||||
|
||||
|
||||
def test_update_segment_vector_economy_uses_keyword_with_keywords_list(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_segment_vector_economy_uses_keyword_with_keywords_list(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY)
|
||||
segment = _make_segment()
|
||||
|
||||
keyword_instance = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Keyword", MagicMock(return_value=keyword_instance))
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.update_segment_vector(["a", "b"], segment, dataset, session=session)
|
||||
VectorService.update_segment_vector(["a", "b"], segment, dataset, session=sqlite_session)
|
||||
|
||||
keyword_instance.delete_by_ids.assert_called_once_with([segment.index_node_id], session)
|
||||
keyword_instance.delete_by_ids.assert_called_once_with([segment.index_node_id], sqlite_session)
|
||||
keyword_instance.add_texts.assert_called_once()
|
||||
args, kwargs = keyword_instance.add_texts.call_args
|
||||
assert len(args[0]) == 1
|
||||
assert args[1] is session
|
||||
assert args[1] is sqlite_session
|
||||
assert kwargs["keywords_list"] == [["a", "b"]]
|
||||
|
||||
|
||||
def test_update_segment_vector_economy_uses_keyword_without_keywords_list(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_segment_vector_economy_uses_keyword_without_keywords_list(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY)
|
||||
segment = _make_segment()
|
||||
|
||||
keyword_instance = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Keyword", MagicMock(return_value=keyword_instance))
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.update_segment_vector(None, segment, dataset, session=session)
|
||||
VectorService.update_segment_vector(None, segment, dataset, session=sqlite_session)
|
||||
keyword_instance.add_texts.assert_called_once()
|
||||
args, kwargs = keyword_instance.add_texts.call_args
|
||||
assert len(args[0]) == 1
|
||||
assert args[1] is session
|
||||
assert args[1] is sqlite_session
|
||||
assert "keywords_list" not in kwargs
|
||||
|
||||
|
||||
def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_generate_child_chunks_regenerate_cleans_then_saves_children(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(doc_form=IndexStructureType.PARAGRAPH_INDEX, tenant_id="tenant-1", dataset_id="dataset-1")
|
||||
segment = _make_segment(segment_id="seg-1")
|
||||
|
||||
@@ -409,13 +433,6 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch
|
||||
factory_instance.init_index_processor.return_value = index_processor
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
child_chunk_ctor = MagicMock(side_effect=lambda **kwargs: kwargs)
|
||||
monkeypatch.setattr(vector_service_module, "ChildChunk", child_chunk_ctor)
|
||||
|
||||
db_mock = MagicMock()
|
||||
db_mock.session.add = MagicMock()
|
||||
db_mock.session.flush = MagicMock()
|
||||
|
||||
VectorService.generate_child_chunks(
|
||||
segment=segment,
|
||||
dataset_document=dataset_document,
|
||||
@@ -423,18 +440,20 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch
|
||||
embedding_model_instance=MagicMock(),
|
||||
processing_rule=processing_rule,
|
||||
regenerate=True,
|
||||
session=db_mock.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
index_processor.clean.assert_called_once()
|
||||
_, transform_kwargs = index_processor.transform.call_args
|
||||
assert transform_kwargs["process_rule"]["rules"]["parent_mode"] == vector_service_module.ParentMode.FULL_DOC
|
||||
index_processor.load.assert_called_once()
|
||||
assert db_mock.session.add.call_count == 2
|
||||
db_mock.session.flush.assert_called_once()
|
||||
stored = sqlite_session.query(ChildChunk).order_by(ChildChunk.position).all()
|
||||
assert [chunk.content for chunk in stored] == ["c1", "c2"]
|
||||
|
||||
|
||||
def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_generate_child_chunks_flushes_even_when_no_children(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(doc_form=IndexStructureType.PARAGRAPH_INDEX)
|
||||
segment = _make_segment()
|
||||
dataset_document = MagicMock()
|
||||
@@ -450,8 +469,6 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest
|
||||
factory_instance.init_index_processor.return_value = index_processor
|
||||
monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance))
|
||||
|
||||
db_mock = MagicMock()
|
||||
|
||||
VectorService.generate_child_chunks(
|
||||
segment=segment,
|
||||
dataset_document=dataset_document,
|
||||
@@ -459,15 +476,16 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest
|
||||
embedding_model_instance=MagicMock(),
|
||||
processing_rule=processing_rule,
|
||||
regenerate=False,
|
||||
session=db_mock.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
index_processor.load.assert_not_called()
|
||||
db_mock.session.add.assert_not_called()
|
||||
db_mock.session.flush.assert_called_once()
|
||||
assert sqlite_session.query(ChildChunk).count() == 0
|
||||
|
||||
|
||||
def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_child_chunk_vector_high_quality_adds_texts(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY)
|
||||
child_chunk = MagicMock()
|
||||
child_chunk.content = "child"
|
||||
@@ -480,13 +498,12 @@ def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.M
|
||||
vector_cls = MagicMock(return_value=vector_instance)
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.create_child_chunk_vector(child_chunk, dataset, session=session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=session)
|
||||
VectorService.create_child_chunk_vector(child_chunk, dataset, session=sqlite_session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session)
|
||||
vector_instance.add_texts.assert_called_once()
|
||||
|
||||
|
||||
def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY)
|
||||
vector_cls = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
@@ -498,11 +515,13 @@ def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch)
|
||||
child_chunk.document_id = "doc-1"
|
||||
child_chunk.dataset_id = "dataset-1"
|
||||
|
||||
VectorService.create_child_chunk_vector(child_chunk, dataset, session=MagicMock())
|
||||
VectorService.create_child_chunk_vector(child_chunk, dataset, session=sqlite_session)
|
||||
vector_cls.assert_not_called()
|
||||
|
||||
|
||||
def test_update_child_chunk_vector_high_quality_updates_vector(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_child_chunk_vector_high_quality_updates_vector(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY)
|
||||
|
||||
new_chunk = MagicMock()
|
||||
@@ -526,25 +545,24 @@ def test_update_child_chunk_vector_high_quality_updates_vector(monkeypatch: pyte
|
||||
vector_cls = MagicMock(return_value=vector_instance)
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.update_child_chunk_vector([new_chunk], [upd_chunk], [del_chunk], dataset, session=session)
|
||||
VectorService.update_child_chunk_vector([new_chunk], [upd_chunk], [del_chunk], dataset, session=sqlite_session)
|
||||
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session)
|
||||
vector_instance.delete_by_ids.assert_called_once_with(["uid", "did"])
|
||||
vector_instance.add_texts.assert_called_once()
|
||||
docs = vector_instance.add_texts.call_args.args[0]
|
||||
assert len(docs) == 2
|
||||
|
||||
|
||||
def test_update_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY)
|
||||
vector_cls = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
VectorService.update_child_chunk_vector([], [], [], dataset, session=MagicMock())
|
||||
VectorService.update_child_chunk_vector([], [], [], dataset, session=sqlite_session)
|
||||
vector_cls.assert_not_called()
|
||||
|
||||
|
||||
def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
dataset = _make_dataset()
|
||||
child_chunk = MagicMock()
|
||||
child_chunk.index_node_id = "cid"
|
||||
@@ -553,9 +571,8 @@ def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch
|
||||
vector_cls = MagicMock(return_value=vector_instance)
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
session = MagicMock()
|
||||
VectorService.delete_child_chunk_vector(child_chunk, dataset, session=session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=session)
|
||||
VectorService.delete_child_chunk_vector(child_chunk, dataset, session=sqlite_session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session)
|
||||
vector_instance.delete_by_ids.assert_called_once_with(["cid"])
|
||||
|
||||
|
||||
@@ -564,156 +581,159 @@ def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_update_multimodel_vector_returns_when_not_high_quality(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_multimodel_vector_returns_when_not_high_quality(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY, is_multimodal=True)
|
||||
segment = _make_segment(tenant_id="t", attachments=[{"id": "a"}])
|
||||
|
||||
vector_cls = MagicMock()
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[])
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
VectorService.update_multimodel_vector(
|
||||
session=db_mock.session, segment=segment, attachment_ids=["a"], dataset=dataset
|
||||
session=sqlite_session, segment=segment, attachment_ids=["a"], dataset=dataset
|
||||
)
|
||||
vector_cls.assert_not_called()
|
||||
db_mock.session.query.assert_not_called()
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
|
||||
def test_update_multimodel_vector_returns_when_no_actual_change(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_multimodel_vector_returns_when_no_actual_change(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True)
|
||||
segment = _make_segment(tenant_id="t", attachments=[{"id": "a"}, {"id": "b"}])
|
||||
|
||||
vector_cls = MagicMock()
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[])
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
VectorService.update_multimodel_vector(
|
||||
session=db_mock.session, segment=segment, attachment_ids=["b", "a"], dataset=dataset
|
||||
session=sqlite_session, segment=segment, attachment_ids=["b", "a"], dataset=dataset
|
||||
)
|
||||
vector_cls.assert_not_called()
|
||||
db_mock.session.query.assert_not_called()
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
|
||||
def test_update_multimodel_vector_deletes_bindings_and_commits_on_empty_new_ids(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True)
|
||||
segment = _make_segment(tenant_id="tenant-1", attachments=[{"id": "old-1"}, {"id": "old-2"}])
|
||||
|
||||
vector_instance = MagicMock(name="vector_instance")
|
||||
vector_cls = MagicMock(return_value=vector_instance)
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[])
|
||||
sqlite_session.add_all(
|
||||
[
|
||||
SegmentAttachmentBinding(
|
||||
tenant_id="tenant-1",
|
||||
dataset_id="dataset-1",
|
||||
document_id="doc-1",
|
||||
segment_id="seg-1",
|
||||
attachment_id=attachment_id,
|
||||
)
|
||||
for attachment_id in ("old-1", "old-2")
|
||||
]
|
||||
)
|
||||
sqlite_session.flush()
|
||||
|
||||
monkeypatch.setattr(vector_service_module, "Vector", vector_cls)
|
||||
|
||||
VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=db_mock.session)
|
||||
VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=sqlite_session)
|
||||
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=db_mock.session)
|
||||
vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session)
|
||||
vector_instance.delete_by_ids.assert_called_once_with(["old-1", "old-2"])
|
||||
db_mock.session.execute.assert_called_once()
|
||||
db_mock.session.flush.assert_called_once()
|
||||
db_mock.session.add_all.assert_not_called()
|
||||
assert sqlite_session.query(SegmentAttachmentBinding).count() == 0
|
||||
vector_instance.add_texts.assert_not_called()
|
||||
|
||||
|
||||
def test_update_multimodel_vector_commits_when_no_upload_files_found(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_update_multimodel_vector_flushes_when_no_upload_files_found(
|
||||
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True)
|
||||
segment = _make_segment(tenant_id="tenant-1", attachments=[{"id": "old-1"}])
|
||||
|
||||
vector_instance = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance))
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[])
|
||||
|
||||
VectorService.update_multimodel_vector(
|
||||
session=db_mock.session, segment=segment, attachment_ids=["new-1"], dataset=dataset
|
||||
session=sqlite_session, segment=segment, attachment_ids=["new-1"], dataset=dataset
|
||||
)
|
||||
|
||||
db_mock.session.flush.assert_called_once()
|
||||
db_mock.session.add_all.assert_not_called()
|
||||
assert sqlite_session.query(SegmentAttachmentBinding).count() == 0
|
||||
vector_instance.add_texts.assert_not_called()
|
||||
|
||||
|
||||
def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_upload_files(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True)
|
||||
segment = _make_segment(segment_id="seg-1", tenant_id="tenant-1", attachments=[{"id": "old-1"}])
|
||||
|
||||
vector_instance = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance))
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")])
|
||||
|
||||
binding_ctor = MagicMock(side_effect=lambda **kwargs: kwargs)
|
||||
monkeypatch.setattr(vector_service_module, "SegmentAttachmentBinding", binding_ctor)
|
||||
monkeypatch.setattr(vector_service_module, "delete", MagicMock())
|
||||
monkeypatch.setattr(vector_service_module, "select", MagicMock())
|
||||
sqlite_session.add(_upload_file())
|
||||
sqlite_session.flush()
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="services.vector_service"):
|
||||
VectorService.update_multimodel_vector(
|
||||
session=db_mock.session, segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset
|
||||
session=sqlite_session, segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset
|
||||
)
|
||||
assert any(r.levelno >= logging.WARNING for r in caplog.records)
|
||||
db_mock.session.add_all.assert_called_once()
|
||||
bindings = db_mock.session.add_all.call_args.args[0]
|
||||
bindings = sqlite_session.query(SegmentAttachmentBinding).all()
|
||||
assert len(bindings) == 1
|
||||
assert bindings[0]["attachment_id"] == "file-1"
|
||||
assert bindings[0].attachment_id == "file-1"
|
||||
|
||||
vector_instance.create_multimodal.assert_called_once()
|
||||
documents = vector_instance.create_multimodal.call_args.args[0]
|
||||
assert len(documents) == 1
|
||||
assert documents[0].page_content == "img.png"
|
||||
assert documents[0].metadata["doc_id"] == "file-1"
|
||||
db_mock.session.flush.assert_called_once()
|
||||
|
||||
|
||||
def test_update_multimodel_vector_updates_bindings_without_multimodal_vector_ops(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=False)
|
||||
segment = _make_segment(tenant_id="tenant-1", attachments=[{"id": "old-1"}])
|
||||
|
||||
vector_instance = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance))
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")])
|
||||
monkeypatch.setattr(
|
||||
vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs)
|
||||
)
|
||||
monkeypatch.setattr(vector_service_module, "delete", MagicMock())
|
||||
monkeypatch.setattr(vector_service_module, "select", MagicMock())
|
||||
sqlite_session.add(_upload_file())
|
||||
sqlite_session.flush()
|
||||
|
||||
VectorService.update_multimodel_vector(
|
||||
session=db_mock.session, segment=segment, attachment_ids=["file-1"], dataset=dataset
|
||||
session=sqlite_session, segment=segment, attachment_ids=["file-1"], dataset=dataset
|
||||
)
|
||||
|
||||
vector_instance.delete_by_ids.assert_not_called()
|
||||
vector_instance.add_texts.assert_not_called()
|
||||
db_mock.session.add_all.assert_called_once()
|
||||
db_mock.session.flush.assert_called_once()
|
||||
binding = sqlite_session.query(SegmentAttachmentBinding).one()
|
||||
assert binding.attachment_id == "file-1"
|
||||
|
||||
|
||||
def test_update_multimodel_vector_rolls_back_and_reraises_on_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True)
|
||||
segment = _make_segment(segment_id="seg-1", tenant_id="tenant-1", attachments=[{"id": "old-1"}])
|
||||
|
||||
vector_instance = MagicMock()
|
||||
monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance))
|
||||
db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")])
|
||||
db_mock.session.flush.side_effect = RuntimeError("boom")
|
||||
monkeypatch.setattr(
|
||||
vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs)
|
||||
)
|
||||
monkeypatch.setattr(vector_service_module, "delete", MagicMock())
|
||||
monkeypatch.setattr(vector_service_module, "select", MagicMock())
|
||||
sqlite_session.add(_upload_file())
|
||||
sqlite_session.flush()
|
||||
rollback_events: list[str] = []
|
||||
event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback"))
|
||||
monkeypatch.setattr(sqlite_session, "flush", MagicMock(side_effect=RuntimeError("boom")))
|
||||
|
||||
with caplog.at_level(logging.ERROR, logger="services.vector_service"):
|
||||
with pytest.raises(RuntimeError, match="boom"):
|
||||
VectorService.update_multimodel_vector(
|
||||
session=db_mock.session, segment=segment, attachment_ids=["file-1"], dataset=dataset
|
||||
session=sqlite_session, segment=segment, attachment_ids=["file-1"], dataset=dataset
|
||||
)
|
||||
|
||||
assert any(r.levelno >= logging.ERROR for r in caplog.records)
|
||||
db_mock.session.rollback.assert_called_once()
|
||||
assert rollback_events == ["rollback"]
|
||||
|
||||
Reference in New Issue
Block a user