diff --git a/api/tests/unit_tests/services/test_snippet_generate_service.py b/api/tests/unit_tests/services/test_snippet_generate_service.py index 6a373db041a..83568a382ae 100644 --- a/api/tests/unit_tests/services/test_snippet_generate_service.py +++ b/api/tests/unit_tests/services/test_snippet_generate_service.py @@ -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) diff --git a/api/tests/unit_tests/services/test_snippet_service.py b/api/tests/unit_tests/services/test_snippet_service.py index 61ab2ff93ba..21a9deb824e 100644 --- a/api/tests/unit_tests/services/test_snippet_service.py +++ b/api/tests/unit_tests/services/test_snippet_service.py @@ -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 diff --git a/api/tests/unit_tests/services/test_vector_service.py b/api/tests/unit_tests/services/test_vector_service.py index eb7bd57e720..75d371c8fa9 100644 --- a/api/tests/unit_tests/services/test_vector_service.py +++ b/api/tests/unit_tests/services/test_vector_service.py @@ -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"]