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:
Asuka Minato
2026-08-11 08:51:07 +00:00
committed by GitHub
co-authored by autofix-ci[bot] Byron Wang
parent cd0e88c680
commit a64a24b39c
3 changed files with 479 additions and 412 deletions
@@ -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"]