test: migrate core RAG sessions to SQLite (#40085)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato
2026-08-11 16:05:26 +09:00
committed by GitHub
parent 44c5f1c011
commit 500e37e2fd
15 changed files with 831 additions and 584 deletions
@@ -172,7 +172,6 @@ class TestApiModeration:
with pytest.raises(ValueError, match="API-based Extension not found"):
api_moderation._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_INPUT, {})
@pytest.mark.parametrize("sqlite_session", [(APIBasedExtension,)], indirect=True)
def test_get_api_based_extension(self, sqlite_session: Session) -> None:
target = APIBasedExtension(
tenant_id="tenant-1",
@@ -1,30 +1,80 @@
import json
from types import SimpleNamespace
from unittest.mock import MagicMock
from typing import cast
import pytest
from pydantic import BaseModel
from pytest_mock import MockerFixture
from sqlalchemy.dialects import postgresql
from sqlalchemy import Engine, event
from sqlalchemy.orm import Session, sessionmaker
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig
from core.plugin.backwards_invocation.app import PluginAppBackwardsInvocation
from core.plugin.backwards_invocation.base import BaseBackwardsInvocation
from models.model import AppMode
from models import Account, Tenant, TenantAccountJoin
from models.enums import EndUserType
from models.model import App, AppMode, AppModelConfig, EndUser
from models.workflow import Workflow, WorkflowType
class _Chunk(BaseModel):
value: int
def _build_app_model_config(result: dict | None = None):
app_model_config = MagicMock()
app_model_config.app_id = "app-1"
app_model_config.to_dict.return_value = result or {
"user_input_form": [{"name": "bar"}],
"annotation_reply": {"enabled": False},
}
return app_model_config
class _DatabaseWithEngine:
def __init__(self, engine: Engine) -> None:
self.engine = engine
def _app(
*,
app_id: str = "app-1",
tenant_id: str = "tenant-1",
mode: AppMode = AppMode.WORKFLOW,
workflow_id: str | None = None,
app_model_config_id: str | None = None,
) -> App:
return App(
id=app_id,
tenant_id=tenant_id,
name="Plugin app",
description="",
mode=mode,
enable_site=False,
enable_api=False,
workflow_id=workflow_id,
app_model_config_id=app_model_config_id,
)
def _workflow(*, workflow_id: str = "workflow-1", app_id: str = "app-1", tenant_id: str = "tenant-1") -> Workflow:
return Workflow(
id=workflow_id,
tenant_id=tenant_id,
app_id=app_id,
type=WorkflowType.WORKFLOW,
version=Workflow.VERSION_DRAFT,
graph="{}",
_features="{}",
created_by="account-1",
)
def _end_user(
*,
user_id: str = "user-1",
tenant_id: str = "tenant-1",
app_id: str = "app-1",
session_id: str = "browser-session",
) -> EndUser:
return EndUser(
id=user_id,
tenant_id=tenant_id,
app_id=app_id,
type=EndUserType.BROWSER,
session_id=session_id,
name="Browser user",
is_anonymous=True,
)
class TestBaseBackwardsInvocation:
@@ -53,23 +103,25 @@ class TestBaseBackwardsInvocation:
class TestPluginAppBackwardsInvocation:
def patch_create_session(self, mocker: MockerFixture, *, return_value=None, side_effect=None):
session = MagicMock()
if side_effect is not None:
session.scalar.side_effect = side_effect
else:
session.scalar.return_value = return_value
session_ctx = MagicMock()
session_ctx.__enter__.return_value = session
session_ctx.__exit__.return_value = None
mocker.patch("core.plugin.backwards_invocation.app.create_session", return_value=session_ctx)
return session
@pytest.fixture(autouse=True)
def _real_sessions(
self,
mocker: MockerFixture,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
sqlite_engine: Engine,
) -> None:
self.session = sqlite_session
self.session_factory = sqlite_session_factory
self.sqlite_engine = sqlite_engine
mocker.patch("core.plugin.backwards_invocation.app.create_session", side_effect=sqlite_session_factory)
def test_fetch_app_info_workflow_path(self, mocker: MockerFixture):
workflow = MagicMock()
workflow.features_dict = {"feature": "v"}
workflow.user_input_form.return_value = [{"name": "foo"}]
app = MagicMock(mode=AppMode.WORKFLOW)
variable = {"type": "text-input", "variable": "foo", "label": "Foo", "required": False}
workflow = _workflow()
workflow.features = json.dumps({"feature": "v"})
workflow.graph = json.dumps({"nodes": [{"data": {"type": "start", "variables": [variable]}}]})
app = _app(mode=AppMode.WORKFLOW)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow)
mapper = mocker.patch(
@@ -80,11 +132,11 @@ class TestPluginAppBackwardsInvocation:
result = PluginAppBackwardsInvocation.fetch_app_info("app-1", "tenant-1")
assert result == {"data": {"mapped": True}}
mapper.assert_called_once_with(features_dict={"feature": "v"}, user_input_form=[{"name": "foo"}])
mapper.assert_called_once_with(features_dict={"feature": "v"}, user_input_form=[{"text-input": variable}])
def test_fetch_app_info_model_config_path(self, mocker: MockerFixture):
model_config_dict = {"user_input_form": [{"name": "bar"}], "k": "v"}
app = MagicMock(mode=AppMode.COMPLETION)
app = _app(mode=AppMode.COMPLETION)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app_model_config_dict", return_value=model_config_dict)
mocker.patch(
@@ -107,9 +159,9 @@ class TestPluginAppBackwardsInvocation:
],
)
def test_invoke_app_routes_by_mode(self, mocker: MockerFixture, mode, route_method):
app = MagicMock(mode=mode)
user = MagicMock()
workflow = MagicMock()
app = _app(mode=mode)
user = _end_user()
workflow = _workflow()
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=user)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow)
@@ -124,16 +176,16 @@ class TestPluginAppBackwardsInvocation:
stream=False,
inputs={"x": 1},
files=[],
session=MagicMock(),
session=self.session,
)
assert result == {"routed": True}
assert route.call_count == 1
def test_invoke_app_uses_end_user_when_user_id_missing(self, mocker: MockerFixture):
app = MagicMock(mode=AppMode.WORKFLOW)
end_user = MagicMock()
workflow = MagicMock()
app = _app(mode=AppMode.WORKFLOW)
end_user = _end_user()
workflow = _workflow()
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow)
get_or_create = mocker.patch(
@@ -151,7 +203,7 @@ class TestPluginAppBackwardsInvocation:
stream=True,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
assert result == {"ok": True}
@@ -160,8 +212,8 @@ class TestPluginAppBackwardsInvocation:
assert route.call_args.args[2] is end_user
def test_invoke_app_missing_query_for_chat_raises(self, mocker: MockerFixture):
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=MagicMock(mode=AppMode.CHAT))
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=MagicMock())
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=_app(mode=AppMode.CHAT))
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=_end_user())
with pytest.raises(ValueError, match="missing query"):
PluginAppBackwardsInvocation.invoke_app(
@@ -173,12 +225,16 @@ class TestPluginAppBackwardsInvocation:
stream=False,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
def test_invoke_app_unexpected_mode_raises(self, mocker: MockerFixture):
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=MagicMock(mode="other"))
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=MagicMock())
mocker.patch.object(
PluginAppBackwardsInvocation,
"_get_app",
return_value=_app(mode=cast(AppMode, "other")),
)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=_end_user())
with pytest.raises(ValueError, match="unexpected app type"):
PluginAppBackwardsInvocation.invoke_app(
@@ -190,7 +246,7 @@ class TestPluginAppBackwardsInvocation:
stream=False,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
@pytest.mark.parametrize(
@@ -201,44 +257,43 @@ class TestPluginAppBackwardsInvocation:
],
)
def test_invoke_chat_app_agent_and_chat(self, mocker: MockerFixture, mode, generator_path):
app = MagicMock(mode=mode, workflow=None)
app = _app(mode=mode)
spy = mocker.patch(generator_path, return_value={"result": "ok"})
result = PluginAppBackwardsInvocation.invoke_chat_app(
app=app,
user=MagicMock(),
user=_end_user(),
conversation_id="conv-1",
query="hello",
stream=False,
inputs={"k": "v"},
files=[],
session=MagicMock(),
session=self.session,
)
assert result == {"result": "ok"}
assert spy.call_count == 1
def test_invoke_chat_app_advanced_chat_injects_pause_state_config(self, mocker: MockerFixture):
workflow = MagicMock()
workflow = _workflow()
workflow.created_by = "owner-id"
app = MagicMock()
app.mode = AppMode.ADVANCED_CHAT
app = _app(mode=AppMode.ADVANCED_CHAT)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow)
mocker.patch(
"core.plugin.backwards_invocation.app.db",
SimpleNamespace(engine=MagicMock()),
_DatabaseWithEngine(self.sqlite_engine),
)
generator_spy = mocker.patch(
"core.plugin.backwards_invocation.app.AdvancedChatAppGenerator.generate",
return_value={"result": "ok"},
)
session = MagicMock()
session = self.session
result = PluginAppBackwardsInvocation.invoke_chat_app(
app=app,
user=MagicMock(),
user=_end_user(),
conversation_id="conv-1",
query="hello",
stream=False,
@@ -255,44 +310,43 @@ class TestPluginAppBackwardsInvocation:
assert pause_state_config.state_owner_user_id == "owner-id"
def test_invoke_chat_app_advanced_chat_without_workflow_raises(self, mocker: MockerFixture):
app = MagicMock(mode=AppMode.ADVANCED_CHAT)
app = _app(mode=AppMode.ADVANCED_CHAT)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=None)
with pytest.raises(ValueError, match="unexpected app type"):
PluginAppBackwardsInvocation.invoke_chat_app(
app=app,
user=MagicMock(),
user=_end_user(),
conversation_id="conv-1",
query="hello",
stream=False,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
def test_invoke_chat_app_unexpected_mode_raises(self):
app = MagicMock(mode="invalid")
app = _app(mode=cast(AppMode, "invalid"))
with pytest.raises(ValueError, match="unexpected app type"):
PluginAppBackwardsInvocation.invoke_chat_app(
app=app,
user=MagicMock(),
user=_end_user(),
conversation_id="conv-1",
query="hello",
stream=False,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
def test_invoke_workflow_app_injects_pause_state_config(self, mocker: MockerFixture):
workflow = MagicMock()
workflow = _workflow()
workflow.created_by = "owner-id"
app = MagicMock()
app.mode = AppMode.WORKFLOW
app = _app(mode=AppMode.WORKFLOW)
mocker.patch(
"core.plugin.backwards_invocation.app.db",
SimpleNamespace(engine=MagicMock()),
_DatabaseWithEngine(self.sqlite_engine),
)
generator_spy = mocker.patch(
"core.plugin.backwards_invocation.app.WorkflowAppGenerator.generate",
@@ -302,7 +356,7 @@ class TestPluginAppBackwardsInvocation:
result = PluginAppBackwardsInvocation.invoke_workflow_app(
app=app,
workflow=workflow,
user=MagicMock(),
user=_end_user(),
stream=False,
inputs={"k": "v"},
files=[],
@@ -315,9 +369,9 @@ class TestPluginAppBackwardsInvocation:
assert pause_state_config.state_owner_user_id == "owner-id"
def test_invoke_app_workflow_without_workflow_raises(self, mocker: MockerFixture):
app = MagicMock(mode=AppMode.WORKFLOW)
app = _app(mode=AppMode.WORKFLOW)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=MagicMock())
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=_end_user())
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=None)
with pytest.raises(ValueError, match="unexpected app type"):
PluginAppBackwardsInvocation.invoke_app(
@@ -329,77 +383,140 @@ class TestPluginAppBackwardsInvocation:
stream=False,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
def test_invoke_completion_app(self, mocker: MockerFixture):
spy = mocker.patch(
"core.plugin.backwards_invocation.app.CompletionAppGenerator.generate", return_value={"ok": 1}
)
app = MagicMock(mode=AppMode.COMPLETION)
app = _app(mode=AppMode.COMPLETION)
result = PluginAppBackwardsInvocation.invoke_completion_app(app, MagicMock(), False, {"x": 1}, [], MagicMock())
result = PluginAppBackwardsInvocation.invoke_completion_app(app, _end_user(), False, {"x": 1}, [], self.session)
assert result == {"ok": 1}
assert spy.call_count == 1
def test_get_user_returns_end_user(self, mocker: MockerFixture):
session = self.patch_create_session(mocker, side_effect=[MagicMock(id="end-user")])
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
def test_get_user_returns_end_user(self):
app = _app()
end_user = EndUser(
id="uid",
tenant_id=app.tenant_id,
app_id=app.id,
type=EndUserType.BROWSER,
session_id="browser-session",
name="Browser user",
is_anonymous=True,
)
self.session.add(end_user)
self.session.commit()
user = PluginAppBackwardsInvocation._get_user("uid", app)
assert user.id == "end-user"
stmt = session.scalar.call_args_list[0].args[0]
compiled = str(stmt.compile(dialect=postgresql.dialect()))
assert "end_users.id" in compiled
assert "end_users.tenant_id" in compiled
assert "end_users.app_id" in compiled
assert stmt.compile().params == {"id_1": "uid", "tenant_id_1": "tenant-1", "app_id_1": "app-1"}
assert user.id == "uid"
assert user.tenant_id == app.tenant_id
assert user.app_id == app.id
def test_get_user_returns_end_user_by_session_id(self, mocker: MockerFixture):
session = self.patch_create_session(mocker, side_effect=[None, MagicMock(id="session-user")])
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
def test_get_user_returns_end_user_by_session_id(self):
app = _app()
end_user = EndUser(
id="session-user",
tenant_id=app.tenant_id,
app_id=app.id,
type=EndUserType.BROWSER,
session_id="wecom-sender-1",
name="External user",
is_anonymous=True,
)
self.session.add(end_user)
self.session.commit()
user = PluginAppBackwardsInvocation._get_user("wecom-sender-1", app)
assert user.id == "session-user"
stmt = session.scalar.call_args_list[1].args[0]
compiled = str(stmt.compile(dialect=postgresql.dialect()))
assert "end_users.session_id" in compiled
assert "end_users.tenant_id" in compiled
assert "end_users.app_id" in compiled
assert stmt.compile().params == {
"session_id_1": "wecom-sender-1",
"tenant_id_1": "tenant-1",
"app_id_1": "app-1",
}
def test_get_user_falls_back_to_account_user(self, mocker: MockerFixture):
session = self.patch_create_session(mocker, side_effect=[None, None, MagicMock(id="account-user")])
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
def test_get_user_rejects_end_user_from_another_app(self):
app = _app()
end_user = EndUser(
id="uid",
tenant_id=app.tenant_id,
app_id="other-app",
type=EndUserType.BROWSER,
session_id="browser-session",
name="Browser user",
is_anonymous=True,
)
self.session.add(end_user)
self.session.commit()
user = PluginAppBackwardsInvocation._get_user("uid", app)
with pytest.raises(ValueError, match="user not found"):
PluginAppBackwardsInvocation._get_user("uid", app)
def test_get_user_rejects_nonmatching_session_id(self):
app = _app()
end_user = EndUser(
id="session-user",
tenant_id=app.tenant_id,
app_id=app.id,
type=EndUserType.BROWSER,
session_id="other-session",
name="External user",
is_anonymous=True,
)
self.session.add(end_user)
self.session.commit()
with pytest.raises(ValueError, match="user not found"):
PluginAppBackwardsInvocation._get_user("wecom-sender-1", app)
def test_get_user_falls_back_to_account_user(self):
app = _app()
tenant = Tenant(name="Plugin tenant")
tenant.id = app.tenant_id
account = Account(name="Account user", email="account-user@example.com")
account.id = "account-user"
membership = TenantAccountJoin(tenant_id=tenant.id, account_id=account.id)
self.session.add_all([tenant, account, membership])
self.session.commit()
user = PluginAppBackwardsInvocation._get_user(account.id, app)
assert user.id == "account-user"
stmt = session.scalar.call_args_list[2].args[0]
compiled = str(stmt.compile(dialect=postgresql.dialect()))
assert "accounts.id" in compiled
assert "tenant_account_joins.account_id" in compiled
assert "tenant_account_joins.tenant_id" in compiled
assert stmt.compile().params == {"id_1": "uid", "tenant_id_1": "tenant-1"}
def test_get_user_raises_when_user_not_found(self, mocker: MockerFixture):
self.patch_create_session(mocker, side_effect=[None, None, None])
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
def test_get_user_rejects_account_from_another_tenant(self):
app = _app()
tenant = Tenant(name="Plugin tenant")
tenant.id = "other-tenant"
account = Account(name="Account user", email="account-user@example.com")
account.id = "account-user"
membership = TenantAccountJoin(tenant_id=tenant.id, account_id=account.id)
self.session.add_all([tenant, account, membership])
self.session.commit()
with pytest.raises(ValueError, match="user not found"):
PluginAppBackwardsInvocation._get_user(account.id, app)
def test_get_user_raises_when_user_not_found(self):
app = _app()
other_tenant_user = EndUser(
id="uid",
tenant_id="other-tenant",
app_id=app.id,
type=EndUserType.BROWSER,
session_id="uid",
name="Wrong tenant",
is_anonymous=True,
)
self.session.add(other_tenant_user)
self.session.commit()
with pytest.raises(ValueError, match="user not found"):
PluginAppBackwardsInvocation._get_user("uid", app)
def test_invoke_app_creates_end_user_for_unknown_external_user_id(self, mocker: MockerFixture):
app = MagicMock(mode=AppMode.WORKFLOW)
end_user = MagicMock()
workflow = MagicMock()
app = _app(mode=AppMode.WORKFLOW)
end_user = _end_user(session_id="wecom-sender-1")
workflow = _workflow()
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow)
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", side_effect=ValueError("user not found"))
@@ -418,70 +535,92 @@ class TestPluginAppBackwardsInvocation:
stream=True,
inputs={},
files=[],
session=MagicMock(),
session=self.session,
)
assert result == {"ok": True}
get_or_create.assert_called_once_with(app, user_id="wecom-sender-1")
assert route.call_args.args[2] is end_user
def test_get_app_returns_app(self, mocker: MockerFixture):
app_obj = MagicMock(id="app")
self.patch_create_session(mocker, return_value=app_obj)
def test_get_app_returns_app(self):
app_obj = _app(app_id="app", tenant_id="tenant")
self.session.add(app_obj)
self.session.commit()
assert PluginAppBackwardsInvocation._get_app("app", "tenant") is app_obj
result = PluginAppBackwardsInvocation._get_app("app", "tenant")
assert result.id == app_obj.id
assert result.tenant_id == app_obj.tenant_id
def test_get_app_raises_when_missing(self, mocker: MockerFixture):
self.patch_create_session(mocker, return_value=None)
def test_get_app_raises_when_missing(self):
self.session.add(_app(app_id="app", tenant_id="other-tenant"))
self.session.commit()
with pytest.raises(ValueError, match="app not found"):
PluginAppBackwardsInvocation._get_app("app", "tenant")
def test_get_app_raises_when_query_fails(self, mocker: MockerFixture):
self.patch_create_session(mocker, side_effect=RuntimeError("db down"))
def test_get_app_raises_when_query_fails(self):
def fail_query(*_args, **_kwargs):
raise RuntimeError("db down")
event.listen(self.sqlite_engine, "before_cursor_execute", fail_query, once=True)
with pytest.raises(ValueError, match="app not found"):
PluginAppBackwardsInvocation._get_app("app", "tenant")
def test_get_workflow_stays_inside_app_boundary(self, mocker: MockerFixture):
workflow = MagicMock(id="workflow")
session = self.patch_create_session(mocker, return_value=workflow)
app = SimpleNamespace(id="app-1", tenant_id="tenant-1", workflow_id="workflow-1")
def test_get_workflow_stays_inside_app_boundary(self):
workflow = _workflow()
other_workflow = _workflow(workflow_id="workflow-other", tenant_id="other-tenant")
self.session.add_all([workflow, other_workflow])
self.session.commit()
app = _app(workflow_id="workflow-1")
assert PluginAppBackwardsInvocation._get_workflow(app) is workflow
result = PluginAppBackwardsInvocation._get_workflow(app)
assert result is not None
assert result.id == workflow.id
assert result.tenant_id == app.tenant_id
stmt = session.scalar.call_args.args[0]
compiled = str(stmt.compile(dialect=postgresql.dialect()))
assert "workflows.id" in compiled
assert "workflows.tenant_id" in compiled
assert "workflows.app_id" in compiled
assert stmt.compile().params == {
"id_1": "workflow-1",
"tenant_id_1": "tenant-1",
"app_id_1": "app-1",
"param_1": 1,
}
def test_get_workflow_rejects_workflow_from_another_tenant(self):
workflow = _workflow(tenant_id="other-tenant")
self.session.add(workflow)
self.session.commit()
app = _app(app_id=workflow.app_id, tenant_id="tenant-1", workflow_id=workflow.id)
assert PluginAppBackwardsInvocation._get_workflow(app) is None
def test_get_workflow_rejects_workflow_from_another_app(self):
workflow = _workflow(app_id="other-app")
self.session.add(workflow)
self.session.commit()
app = _app(app_id="app-1", tenant_id=workflow.tenant_id, workflow_id=workflow.id)
assert PluginAppBackwardsInvocation._get_workflow(app) is None
def test_get_app_model_config_dict_uses_explicit_session_for_annotation_reply(self, mocker: MockerFixture):
annotation_reply = {"enabled": False}
app_model_config = _build_app_model_config()
session = self.patch_create_session(mocker, return_value=app_model_config)
app_model_config = AppModelConfig(app_id="app-1", user_input_form=json.dumps([{"name": "bar"}]))
app_model_config.id = "config-1"
self.session.add(app_model_config)
self.session.commit()
load_annotation_reply_config = mocker.patch(
"core.plugin.backwards_invocation.app.load_annotation_reply_config",
return_value=annotation_reply,
)
app = SimpleNamespace(id="app-1", app_model_config_id="config-1")
app = _app(app_model_config_id="config-1")
result = PluginAppBackwardsInvocation._get_app_model_config_dict(app)
assert result is not None
assert result["user_input_form"] == [{"name": "bar"}]
assert result["annotation_reply"] == annotation_reply
load_annotation_reply_config.assert_called_once_with(session, "app-1")
app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply)
queried_session, queried_app_id = load_annotation_reply_config.call_args.args
assert isinstance(queried_session, Session)
assert queried_app_id == "app-1"
stmt = session.scalar.call_args.args[0]
compiled = str(stmt.compile(dialect=postgresql.dialect()))
assert "app_model_configs.id" in compiled
assert "app_model_configs.app_id" in compiled
assert stmt.compile().params == {"id_1": "config-1", "app_id_1": "app-1", "param_1": 1}
def test_get_app_model_config_dict_rejects_config_from_another_app(self):
app_model_config = AppModelConfig(app_id="other-app", user_input_form=json.dumps([{"name": "bar"}]))
app_model_config.id = "config-1"
self.session.add(app_model_config)
self.session.commit()
app = _app(app_model_config_id=app_model_config.id)
assert PluginAppBackwardsInvocation._get_app_model_config_dict(app) is None
@@ -1,5 +1,7 @@
from unittest.mock import MagicMock, patch
from sqlalchemy.orm import Session
from core.rag.data_post_processor.data_post_processor import DataPostProcessor
from core.rag.data_post_processor.reorder import ReorderRunner
from core.rag.index_processor.constant.query_type import QueryType
@@ -14,10 +16,9 @@ def _doc(content: str) -> Document:
class TestDataPostProcessor:
def test_init_sets_rerank_and_reorder_runners(self):
def test_init_sets_rerank_and_reorder_runners(self, unbound_session: Session):
rerank_runner = object()
reorder_runner = object()
session = MagicMock()
with patch.object(DataPostProcessor, "_get_rerank_runner", return_value=rerank_runner) as rerank_mock:
with patch.object(DataPostProcessor, "_get_reorder_runner", return_value=reorder_runner) as reorder_mock:
@@ -27,7 +28,7 @@ class TestDataPostProcessor:
reranking_model={"config": "value"},
weights={"weight": "value"},
reorder_enabled=True,
session=session,
session=unbound_session,
)
assert processor.rerank_runner is rerank_runner
@@ -37,7 +38,7 @@ class TestDataPostProcessor:
"tenant-1",
{"config": "value"},
{"weight": "value"},
session=session,
session=unbound_session,
)
reorder_mock.assert_called_once_with(True)
@@ -79,7 +80,7 @@ class TestDataPostProcessor:
assert processor.invoke(query="query", documents=documents) == documents
def test_get_rerank_runner_for_weighted_score(self):
def test_get_rerank_runner_for_weighted_score(self, unbound_session: Session):
weights_config = {
"vector_setting": {
"vector_weight": 0.7,
@@ -90,7 +91,6 @@ class TestDataPostProcessor:
}
expected_runner = object()
processor = DataPostProcessor.__new__(DataPostProcessor)
session = MagicMock()
with patch(
"core.rag.data_post_processor.data_post_processor.RerankRunnerFactory.create_rerank_runner",
@@ -101,7 +101,7 @@ class TestDataPostProcessor:
tenant_id="tenant-1",
reranking_model=None,
weights=weights_config,
session=session,
session=unbound_session,
)
assert result is expected_runner
@@ -113,13 +113,12 @@ class TestDataPostProcessor:
assert kwargs["weights"].vector_setting.embedding_model_name == "embedding-y"
assert kwargs["weights"].keyword_setting.keyword_weight == 0.3
def test_get_rerank_runner_for_reranking_model_returns_none_without_model_instance(self):
def test_get_rerank_runner_for_reranking_model_returns_none_without_model_instance(self, unbound_session: Session):
processor = DataPostProcessor.__new__(DataPostProcessor)
reranking_model = {
"reranking_provider_name": "provider-x",
"reranking_model_name": "model-y",
}
session = MagicMock()
with patch.object(DataPostProcessor, "_get_rerank_model_instance", return_value=None) as model_mock:
with patch(
@@ -130,18 +129,17 @@ class TestDataPostProcessor:
tenant_id="tenant-1",
reranking_model=reranking_model,
weights=None,
session=session,
session=unbound_session,
)
assert result is None
model_mock.assert_called_once_with("tenant-1", reranking_model)
factory_mock.assert_not_called()
def test_get_rerank_runner_for_reranking_model_creates_runner_with_model_instance(self):
def test_get_rerank_runner_for_reranking_model_creates_runner_with_model_instance(self, unbound_session: Session):
processor = DataPostProcessor.__new__(DataPostProcessor)
model_instance = object()
expected_runner = object()
session = MagicMock()
with patch.object(DataPostProcessor, "_get_rerank_model_instance", return_value=model_instance):
with patch(
@@ -156,22 +154,24 @@ class TestDataPostProcessor:
"reranking_model_name": "model-y",
},
weights=None,
session=session,
session=unbound_session,
)
assert result is expected_runner
factory_mock.assert_called_once_with(
runner_type=RerankMode.RERANKING_MODEL,
rerank_model_instance=model_instance,
session=session,
session=unbound_session,
)
def test_get_rerank_runner_returns_none_for_unsupported_mode(self):
def test_get_rerank_runner_returns_none_for_unsupported_mode(self, unbound_session: Session):
processor = DataPostProcessor.__new__(DataPostProcessor)
session = MagicMock()
assert processor._get_rerank_runner("unsupported", "tenant-1", None, None, session=session) is None
assert processor._get_rerank_runner(RerankMode.WEIGHTED_SCORE, "tenant-1", None, None, session=session) is None
assert processor._get_rerank_runner("unsupported", "tenant-1", None, None, session=unbound_session) is None
assert (
processor._get_rerank_runner(RerankMode.WEIGHTED_SCORE, "tenant-1", None, None, session=unbound_session)
is None
)
def test_get_reorder_runner_by_flag(self):
processor = DataPostProcessor.__new__(DataPostProcessor)
@@ -4,10 +4,13 @@ from typing import Any
from unittest.mock import MagicMock
import pytest
from sqlalchemy import event, select
from sqlalchemy.orm import Session
import core.rag.datasource.keyword.jieba.jieba as jieba_module
from core.rag.datasource.keyword.jieba.jieba import Jieba, dumps_with_sets, set_orjson_default
from core.rag.models.document import Document
from models.dataset import DatasetKeywordTable, DocumentSegment
class _DummyLock:
@@ -18,37 +21,6 @@ class _DummyLock:
return False
class _Field:
def __init__(self, name: str):
self._name = name
def __eq__(self, other):
return ("eq", self._name, other)
def in_(self, values):
return ("in", self._name, tuple(values))
class _FakeExecuteResult:
def __init__(self, segments: list[SimpleNamespace]):
self._segments = segments
def scalars(self):
return self
def all(self):
return self._segments
class _FakeSelect:
def __init__(self):
self.where_conditions: tuple | None = None
def where(self, *conditions):
self.where_conditions = conditions
return self
def _dataset_keyword_table(data_source_type: str = "database", keyword_table_dict: dict[str, Any] | None = None):
return SimpleNamespace(
data_source_type=data_source_type,
@@ -67,8 +39,7 @@ def _dataset(dataset_keyword_table=None, keyword_number=None):
@pytest.fixture
def patched_runtime(monkeypatch: pytest.MonkeyPatch):
session = MagicMock()
def patched_runtime(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session):
storage = MagicMock()
lock = MagicMock(return_value=_DummyLock())
redis_client = SimpleNamespace(lock=lock)
@@ -76,7 +47,28 @@ def patched_runtime(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(jieba_module, "storage", storage)
monkeypatch.setattr(jieba_module, "redis_client", redis_client)
return SimpleNamespace(session=session, storage=storage, lock=lock)
return SimpleNamespace(session=sqlite_session, storage=storage, lock=lock)
def _segment(*, index_node_id: str = "node-2") -> DocumentSegment:
segment = DocumentSegment(
tenant_id="tenant-1",
dataset_id="dataset-1",
document_id="doc-2",
position=1,
content="segment-content",
word_count=1,
tokens=1,
created_by="user-1",
enabled=True,
keywords=[],
answer=None,
index_node_id=index_node_id,
index_node_hash="hash-2",
status="completed",
)
segment.id = "segment-1"
return segment
def test_create_indexes_documents_and_returns_self(monkeypatch: pytest.MonkeyPatch, patched_runtime):
@@ -156,9 +148,11 @@ def test_add_texts_without_keywords_list_always_uses_extractor(monkeypatch: pyte
assert keyword._update_segment_keywords.call_args.args[3] is patched_runtime.session
def test_text_exists_handles_missing_and_existing_keyword_table(monkeypatch: pytest.MonkeyPatch):
def test_text_exists_handles_missing_and_existing_keyword_table(
monkeypatch: pytest.MonkeyPatch, unbound_session: Session
):
keyword = Jieba(_dataset(_dataset_keyword_table(keyword_table_dict=None)))
session = MagicMock()
session = unbound_session
assert keyword.text_exists("node-1", session=session) is False
keyword = Jieba(
@@ -205,24 +199,9 @@ def test_delete_by_ids_saves_none_when_keyword_table_is_missing(monkeypatch: pyt
def test_search_returns_documents_in_rank_order_and_applies_filter(monkeypatch: pytest.MonkeyPatch, patched_runtime):
class _FakeDocumentSegment:
dataset_id = _Field("dataset_id")
index_node_id = _Field("index_node_id")
document_id = _Field("document_id")
keyword = Jieba(_dataset(_dataset_keyword_table()))
patched_runtime.session.scalars.return_value.all.return_value = [
SimpleNamespace(
index_node_id="node-2",
content="segment-content",
index_node_hash="hash-2",
document_id="doc-2",
dataset_id="dataset-1",
)
]
monkeypatch.setattr(jieba_module, "DocumentSegment", _FakeDocumentSegment)
monkeypatch.setattr(jieba_module, "select", lambda *_: _FakeSelect())
patched_runtime.session.add(_segment())
patched_runtime.session.flush()
monkeypatch.setattr(keyword, "_retrieve_ids_by_query", MagicMock(return_value=["node-1", "node-2"]))
documents = keyword.search("query", session=patched_runtime.session, top_k=2, document_ids_filter=["doc-2"])
@@ -233,39 +212,47 @@ def test_search_returns_documents_in_rank_order_and_applies_filter(monkeypatch:
assert documents[0].metadata["doc_hash"] == "hash-2"
def test_delete_removes_keyword_table_and_optional_file(monkeypatch: pytest.MonkeyPatch, patched_runtime):
db_keyword = _dataset_keyword_table(data_source_type="database")
file_keyword = _dataset_keyword_table(data_source_type="object_storage")
def test_delete_removes_keyword_table_and_optional_file(patched_runtime):
db_keyword = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="database")
patched_runtime.session.add(db_keyword)
patched_runtime.session.commit()
commits: list[str] = []
event.listen(patched_runtime.session, "after_commit", lambda _session: commits.append("commit"))
keyword_db = Jieba(_dataset(db_keyword))
keyword_db.delete(session=patched_runtime.session)
patched_runtime.storage.delete.assert_not_called()
assert patched_runtime.session.get(DatasetKeywordTable, db_keyword.id) is None
file_keyword = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="object_storage")
patched_runtime.session.add(file_keyword)
patched_runtime.session.commit()
keyword_file = Jieba(_dataset(file_keyword))
keyword_file.delete(session=patched_runtime.session)
patched_runtime.storage.delete.assert_called_once_with("keyword_files/tenant-1/dataset-1.txt")
assert patched_runtime.session.delete.call_count == 2
assert patched_runtime.session.commit.call_count == 2
assert patched_runtime.session.get(DatasetKeywordTable, file_keyword.id) is None
assert commits == ["commit", "commit", "commit"]
def test_save_dataset_keyword_table_to_database(monkeypatch: pytest.MonkeyPatch, patched_runtime):
dataset_keyword_table = _dataset_keyword_table(data_source_type="database")
def test_save_dataset_keyword_table_to_database(patched_runtime):
dataset_keyword_table = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="database")
patched_runtime.session.add(dataset_keyword_table)
patched_runtime.session.flush()
keyword = Jieba(_dataset(dataset_keyword_table))
patched_runtime.session.scalar.return_value = dataset_keyword_table
keyword._save_dataset_keyword_table({"kw": {"node-1"}}, patched_runtime.session)
assert '"__type__":"keyword_table"' in dataset_keyword_table.keyword_table
assert '"index_id":"dataset-1"' in dataset_keyword_table.keyword_table
patched_runtime.session.flush.assert_called_once()
def test_save_dataset_keyword_table_to_file_storage(monkeypatch: pytest.MonkeyPatch, patched_runtime):
dataset_keyword_table = _dataset_keyword_table(data_source_type="file")
def test_save_dataset_keyword_table_to_file_storage(patched_runtime):
dataset_keyword_table = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="file")
patched_runtime.session.add(dataset_keyword_table)
patched_runtime.session.flush()
keyword = Jieba(_dataset(dataset_keyword_table))
patched_runtime.storage.exists.return_value = True
patched_runtime.session.scalar.return_value = dataset_keyword_table
keyword._save_dataset_keyword_table({"kw": {"node-1"}}, patched_runtime.session)
@@ -276,33 +263,38 @@ def test_save_dataset_keyword_table_to_file_storage(monkeypatch: pytest.MonkeyPa
assert isinstance(save_args[1], bytes)
def test_get_dataset_keyword_table_returns_existing_table_data(monkeypatch: pytest.MonkeyPatch, patched_runtime):
existing = _dataset_keyword_table(
keyword_table_dict={"__type__": "keyword_table", "__data__": {"table": {"kw": ["node-1"]}}}
def test_get_dataset_keyword_table_returns_existing_table_data(patched_runtime):
existing = DatasetKeywordTable(
dataset_id="dataset-1",
keyword_table="",
data_source_type="database",
)
existing.get_keyword_table_dict = MagicMock(
return_value={"__type__": "keyword_table", "__data__": {"table": {"kw": ["node-1"]}}}
)
patched_runtime.session.add(existing)
patched_runtime.session.flush()
keyword = Jieba(_dataset(existing))
patched_runtime.session.scalar.return_value = existing
assert keyword._get_dataset_keyword_table(patched_runtime.session) == {"kw": ["node-1"]}
missing_payload = _dataset_keyword_table(keyword_table_dict=None)
keyword_with_missing_payload = Jieba(_dataset(missing_payload))
patched_runtime.session.scalar.return_value = missing_payload
existing.get_keyword_table_dict = MagicMock(return_value=None)
keyword_with_missing_payload = Jieba(_dataset(existing))
assert keyword_with_missing_payload._get_dataset_keyword_table(patched_runtime.session) == {}
def test_get_dataset_keyword_table_creates_table_when_missing(monkeypatch: pytest.MonkeyPatch, patched_runtime):
keyword = Jieba(_dataset(dataset_keyword_table=None))
monkeypatch.setattr(jieba_module.dify_config, "KEYWORD_DATA_SOURCE_TYPE", "database")
patched_runtime.session.scalar.return_value = None
result = keyword._get_dataset_keyword_table(patched_runtime.session)
assert result == {}
created_table = patched_runtime.session.add.call_args.args[0]
created_table = patched_runtime.session.scalar(
select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == "dataset-1")
)
assert created_table is not None
assert created_table.dataset_id == "dataset-1"
assert created_table.data_source_type == "database"
assert '"index_id":"dataset-1"' in created_table.keyword_table
patched_runtime.session.flush.assert_called_once()
def test_add_and_delete_ids_from_keyword_table_helpers():
@@ -334,31 +326,18 @@ def test_retrieve_ids_by_query_ranks_by_keyword_frequency(monkeypatch: pytest.Mo
assert ranked_ids == ["node-2"]
def test_update_segment_keywords_updates_when_segment_exists(monkeypatch: pytest.MonkeyPatch, patched_runtime):
class _FakeDocumentSegment:
dataset_id = _Field("dataset_id")
index_node_id = _Field("index_node_id")
monkeypatch.setattr(jieba_module, "DocumentSegment", _FakeDocumentSegment)
monkeypatch.setattr(jieba_module, "select", lambda *_: _FakeSelect())
def test_update_segment_keywords_updates_when_segment_exists(patched_runtime):
keyword = Jieba(_dataset(_dataset_keyword_table()))
segment = SimpleNamespace(keywords=[])
patched_runtime.session.scalar.return_value = segment
segment = _segment(index_node_id="node-1")
patched_runtime.session.add(segment)
patched_runtime.session.flush()
keyword._update_segment_keywords("dataset-1", "node-1", ["kw1", "kw2"], patched_runtime.session)
assert segment.keywords == ["kw1", "kw2"]
patched_runtime.session.add.assert_called_once_with(segment)
patched_runtime.session.flush.assert_called_once()
patched_runtime.session.reset_mock()
patched_runtime.session.scalar.return_value = None
keyword._update_segment_keywords("dataset-1", "node-missing", ["kw3"], patched_runtime.session)
patched_runtime.session.add.assert_not_called()
patched_runtime.session.flush.assert_not_called()
assert segment.keywords == ["kw1", "kw2"]
def test_create_segment_keywords_and_update_segment_keywords_index(monkeypatch: pytest.MonkeyPatch, patched_runtime):
@@ -1,8 +1,8 @@
from types import SimpleNamespace
from typing import override
from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from core.rag.datasource.keyword.keyword_base import BaseKeyword
from core.rag.models.document import Document
@@ -64,9 +64,9 @@ class _KeywordForHelpers(BaseKeyword):
return []
def test_abstract_methods_raise_not_implemented():
def test_abstract_methods_raise_not_implemented(unbound_session: Session):
keyword = _KeywordThatRaises(SimpleNamespace(id="dataset-1"))
session = MagicMock()
session = unbound_session
with pytest.raises(NotImplementedError):
keyword.create([], session)
@@ -87,7 +87,7 @@ def test_abstract_methods_raise_not_implemented():
keyword.search("query", session=session)
def test_filter_duplicate_texts_removes_existing_doc_ids():
def test_filter_duplicate_texts_removes_existing_doc_ids(unbound_session: Session):
keyword = _KeywordForHelpers(SimpleNamespace(id="dataset-1"), existing_ids={"duplicate"})
texts = [
Document(page_content="keep", metadata={"doc_id": "keep"}),
@@ -95,7 +95,7 @@ def test_filter_duplicate_texts_removes_existing_doc_ids():
SimpleNamespace(page_content="without-metadata", metadata=None),
]
filtered = keyword._filter_duplicate_texts(texts, session=MagicMock())
filtered = keyword._filter_duplicate_texts(texts, session=unbound_session)
assert [text.metadata["doc_id"] for text in filtered if text.metadata] == ["keep"]
assert any(text.metadata is None for text in filtered)
@@ -4,6 +4,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from core.rag.datasource.keyword.keyword_factory import Keyword
from core.rag.datasource.keyword.keyword_type import KeyWordType
@@ -39,7 +40,7 @@ def test_keyword_initialization_uses_configured_factory(monkeypatch: pytest.Monk
assert keyword._keyword_processor is fake_processor
def test_keyword_methods_forward_to_processor():
def test_keyword_methods_forward_to_processor(unbound_session: Session):
processor = MagicMock()
processor.text_exists.return_value = True
processor.search.return_value = [Document(page_content="matched", metadata={"doc_id": "doc-1"})]
@@ -48,7 +49,7 @@ def test_keyword_methods_forward_to_processor():
keyword._keyword_processor = processor
docs = [Document(page_content="doc", metadata={"doc_id": "doc-1"})]
session = MagicMock()
session = unbound_session
keyword.create(docs, session, foo="bar")
keyword.add_texts(docs, session, batch=True, keywords_list=[["kw"]])
assert keyword.text_exists("doc-1", session=session) is True
@@ -1,12 +1,19 @@
import base64
import sys
import types
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from sqlalchemy.orm import Session
from core.rag.models.document import Document
from extensions.storage.storage_type import StorageType
from models.dataset import Whitelist
from models.enums import CreatorUserRole
from models.model import UploadFile
def _register_fake_factory_module(monkeypatch: pytest.MonkeyPatch, module_path: str, class_name: str):
@@ -145,13 +152,12 @@ def test_get_vector_factory_entry_point_overrides_builtin(vector_factory_module,
assert result_cls is _PluginChromaFactory
def test_vector_init_uses_default_and_custom_attributes(vector_factory_module):
def test_vector_init_uses_default_and_custom_attributes(vector_factory_module, unbound_session: Session):
dataset = SimpleNamespace(id="dataset-1")
session = MagicMock()
with patch.object(vector_factory_module.Vector, "_init_vector", return_value="processor") as init_vector:
default_vector = vector_factory_module.Vector(dataset, session=session)
custom_vector = vector_factory_module.Vector(dataset, attributes=["doc_id"], session=session)
default_vector = vector_factory_module.Vector(dataset, session=unbound_session)
custom_vector = vector_factory_module.Vector(dataset, attributes=["doc_id"], session=unbound_session)
# `is_summary` and `original_chunk_id` must be in the default return-properties
# projection so summary index retrieval works on backends that honor the list
@@ -171,10 +177,10 @@ def test_vector_init_uses_default_and_custom_attributes(vector_factory_module):
# trigger billing/feature-service calls during ``Vector(dataset, session=...)``
# construction. See ``_LazyEmbeddings``.
assert isinstance(default_vector._embeddings, vector_factory_module._LazyEmbeddings)
assert default_vector._session is session
assert custom_vector._session is session
assert default_vector._session is unbound_session
assert custom_vector._session is unbound_session
assert default_vector._vector_processor == "processor"
assert [call.kwargs["session"] for call in init_vector.call_args_list] == [session, session]
assert [call.kwargs["session"] for call in init_vector.call_args_list] == [unbound_session, unbound_session]
def test_lazy_embeddings_defer_real_load_until_first_embed_call(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
@@ -220,7 +226,9 @@ def test_lazy_embeddings_defer_real_load_until_first_embed_call(vector_factory_m
inner_model.embed_documents.assert_called_once_with(["world"])
def test_init_vector_prefers_dataset_index_struct(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
def test_init_vector_prefers_dataset_index_struct(
vector_factory_module, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
):
calls = {"vector_type": None, "init_args": None}
class _Factory:
@@ -241,28 +249,25 @@ def test_init_vector_prefers_dataset_index_struct(vector_factory_module, monkeyp
vector._attributes = ["doc_id"]
vector._embeddings = "embeddings"
result = vector._init_vector(session=MagicMock())
result = vector._init_vector(session=unbound_session)
assert result == "vector-processor"
assert calls["vector_type"] == vector_factory_module.VectorType.UPSTASH
assert calls["init_args"] == (vector._dataset, ["doc_id"], "embeddings")
def test_init_vector_uses_whitelist_override(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
class _Expr:
def __eq__(self, _other):
return "expr"
def test_init_vector_uses_whitelist_override(
vector_factory_module, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
):
calls = {"vector_type": None}
class _Factory:
def init_vector(self, dataset, attributes, embeddings):
return "vector-processor"
monkeypatch.setattr(vector_factory_module, "Whitelist", SimpleNamespace(tenant_id=_Expr(), category=_Expr()))
monkeypatch.setattr(vector_factory_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt"))
session = MagicMock()
session.scalars.return_value.one_or_none.return_value = object()
tenant_id = str(uuid4())
sqlite_session.add(Whitelist(tenant_id=tenant_id, category="vector_db"))
sqlite_session.commit()
monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE", vector_factory_module.VectorType.CHROMA)
monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE_WHITELIST_ENABLE", True)
monkeypatch.setattr(
@@ -272,18 +277,19 @@ def test_init_vector_uses_whitelist_override(vector_factory_module, monkeypatch:
)
vector = vector_factory_module.Vector.__new__(vector_factory_module.Vector)
vector._dataset = SimpleNamespace(index_struct_dict=None, tenant_id="tenant-1")
vector._dataset = SimpleNamespace(index_struct_dict=None, tenant_id=tenant_id)
vector._attributes = ["doc_id"]
vector._embeddings = "embeddings"
result = vector._init_vector(session=session)
result = vector._init_vector(session=sqlite_session)
assert result == "vector-processor"
assert calls["vector_type"] == vector_factory_module.VectorType.TIDB_ON_QDRANT
session.scalars.assert_called_once_with("stmt")
def test_init_vector_raises_when_vector_store_missing(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
def test_init_vector_raises_when_vector_store_missing(
vector_factory_module, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
):
monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE", None)
monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE_WHITELIST_ENABLE", False)
@@ -293,7 +299,7 @@ def test_init_vector_raises_when_vector_store_missing(vector_factory_module, mon
vector._embeddings = "embeddings"
with pytest.raises(ValueError, match="Vector store must be specified"):
vector._init_vector(session=MagicMock())
vector._init_vector(session=unbound_session)
def test_create_batches_texts_and_skips_empty_input(vector_factory_module):
@@ -347,36 +353,41 @@ def test_create_skips_empty_text_documents_before_embedding(vector_factory_modul
vector._vector_processor.create.assert_not_called()
def test_create_multimodal_filters_missing_uploads(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
class _Field:
def in_(self, value):
return value
def __eq__(self, value):
return value
def test_create_multimodal_filters_missing_uploads(
vector_factory_module, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
):
upload_file = UploadFile(
tenant_id=str(uuid4()),
storage_type=StorageType.LOCAL,
key="k-1",
name="image.png",
size=3,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
created_at=datetime.now(UTC),
used=True,
)
sqlite_session.add(upload_file)
sqlite_session.commit()
vector = vector_factory_module.Vector.__new__(vector_factory_module.Vector)
vector._embeddings = MagicMock()
vector._embeddings.embed_multimodal_documents.return_value = [[0.1, 0.2]]
vector._vector_processor = MagicMock()
session = MagicMock()
vector._session = session
session.scalars.return_value = SimpleNamespace(all=lambda: [SimpleNamespace(id="f-1", key="k-1")])
monkeypatch.setattr(vector_factory_module, "UploadFile", SimpleNamespace(id=_Field()))
monkeypatch.setattr(vector_factory_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt"))
vector._session = sqlite_session
monkeypatch.setattr(vector_factory_module.storage, "load_once", MagicMock(return_value=b"abc"))
docs = [
Document(page_content="file-1", metadata={"doc_id": "f-1", "doc_type": "image"}),
Document(page_content="file-2", metadata={"doc_id": "f-2", "doc_type": "image"}),
Document(page_content="file-1", metadata={"doc_id": upload_file.id, "doc_type": "image"}),
Document(page_content="file-2", metadata={"doc_id": str(uuid4()), "doc_type": "image"}),
]
vector.create_multimodal(file_documents=docs, request_id="r-1")
file_base64 = base64.b64encode(b"abc").decode()
vector._embeddings.embed_multimodal_documents.assert_called_once_with(
[{"content": file_base64, "content_type": "image", "file_id": "f-1"}]
[{"content": file_base64, "content_type": "image", "file_id": upload_file.id}]
)
vector._vector_processor.create.assert_called_once_with(
texts=[docs[0]],
@@ -482,30 +493,43 @@ def test_vector_delegation_methods(vector_factory_module):
vector._vector_processor.delete_by_metadata_field.assert_called_once_with("doc_id", "doc-1")
def test_search_by_file_handles_missing_and_existing_upload(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
def test_search_by_file_handles_missing_and_existing_upload(
vector_factory_module, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
):
vector = vector_factory_module.Vector.__new__(vector_factory_module.Vector)
vector._embeddings = MagicMock()
vector._vector_processor = MagicMock()
session = MagicMock()
session.get.return_value = None
vector._session = session
vector._session = sqlite_session
missing_id = str(uuid4())
assert vector.search_by_file("file-1") == []
session.get.assert_called_once_with(vector_factory_module.UploadFile, "file-1")
assert vector.search_by_file(missing_id) == []
session.get.return_value = SimpleNamespace(key="blob-key")
upload_file = UploadFile(
tenant_id=str(uuid4()),
storage_type=StorageType.LOCAL,
key="blob-key",
name="query.png",
size=10,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
created_at=datetime.now(UTC),
used=True,
)
sqlite_session.add(upload_file)
sqlite_session.commit()
monkeypatch.setattr(vector_factory_module.storage, "load_once", MagicMock(return_value=b"file-bytes"))
vector._embeddings.embed_multimodal_query.return_value = [0.3, 0.4]
vector._vector_processor.search_by_vector.return_value = ["hit"]
result = vector.search_by_file("file-2", top_k=2)
result = vector.search_by_file(upload_file.id, top_k=2)
assert result == ["hit"]
session.get.assert_called_with(vector_factory_module.UploadFile, "file-2")
payload = vector._embeddings.embed_multimodal_query.call_args.args[0]
assert payload["content_type"] == vector_factory_module.DocType.IMAGE
assert payload["file_id"] == "file-2"
assert payload["file_id"] == upload_file.id
def test_delete_clears_redis_cache_when_collection_exists(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
@@ -15,9 +15,12 @@ import pytest
from docx import Document
from docx.oxml import OxmlElement
from docx.oxml.ns import qn
from sqlalchemy import event, select
from sqlalchemy.orm import Session
import core.rag.extractor.word_extractor as we
from core.rag.extractor.word_extractor import WordExtractor
from models.model import UploadFile
class _TextOxmlElement(Protocol):
@@ -112,7 +115,7 @@ def test_init_downloads_via_remote_fetcher(monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("inject_session", [False, True])
def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_session: bool):
def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_session: bool, sqlite_session: Session):
external_bytes = b"ext-bytes"
internal_bytes = b"int-bytes"
@@ -124,35 +127,13 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_sessio
monkeypatch.setattr(we, "storage", SimpleNamespace(save=save))
# Patch db.session to record adds/commit
class DummySession:
def __init__(self):
self.added = []
self.committed = False
def add_all(self, objects):
self.added.extend(objects)
def commit(self):
self.committed = True
db_stub = SimpleNamespace(session=DummySession())
db_stub = SimpleNamespace(session=sqlite_session)
monkeypatch.setattr(we, "db", db_stub)
# Patch config values used for URL composition and storage type
monkeypatch.setattr(we.dify_config, "FILES_URL", "http://files.local", raising=False)
monkeypatch.setattr(we.dify_config, "STORAGE_TYPE", "local", raising=False)
# Patch UploadFile to avoid real DB models
class FakeUploadFile:
_i = 0
def __init__(self, **kwargs): # kwargs match the real signature fields
type(self)._i += 1
self.id = f"u{self._i}"
monkeypatch.setattr(we, "UploadFile", FakeUploadFile)
# Patch external image fetcher
def fake_make_request(method: str, url: str, **kwargs):
assert method == "GET"
@@ -176,9 +157,11 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_sessio
doc = SimpleNamespace(part=SimpleNamespace(rels={"rId1": rel_ext, "rId2": rel_int}))
extractor = object.__new__(WordExtractor)
extractor.tenant_id = "t1"
extractor.user_id = "u1"
extractor.tenant_id = "00000000-0000-0000-0000-000000000001"
extractor.user_id = "00000000-0000-0000-0000-000000000002"
extractor._session = db_stub.session if inject_session else None
transaction_events: list[str] = []
event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit"))
image_map = extractor._extract_images_from_docx(doc)
@@ -191,12 +174,13 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_sessio
assert external_bytes in payloads
assert internal_bytes in payloads
# DB interactions should be recorded
assert len(db_stub.session.added) == 2
assert db_stub.session.committed is not inject_session
assert len(sqlite_session.scalars(select(UploadFile)).all()) == 2
assert transaction_events == ([] if inject_session else ["commit"])
def test_extract_images_does_not_stage_partial_files_on_storage_failure(monkeypatch: pytest.MonkeyPatch):
def test_extract_images_does_not_stage_partial_files_on_storage_failure(
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
):
class HashablePart:
def __init__(self, blob: bytes):
self.blob = blob
@@ -222,22 +206,20 @@ def test_extract_images_does_not_stage_partial_files_on_storage_failure(monkeypa
}
)
)
session = MagicMock()
save = MagicMock(side_effect=[None, RuntimeError("storage failure")])
monkeypatch.setattr(we, "storage", SimpleNamespace(save=save))
monkeypatch.setattr(we.dify_config, "FILES_URL", "http://files.local", raising=False)
monkeypatch.setattr(we.dify_config, "STORAGE_TYPE", "local", raising=False)
extractor = object.__new__(WordExtractor)
extractor.tenant_id = "tenant"
extractor.user_id = "user"
extractor._session = session
extractor.tenant_id = "00000000-0000-0000-0000-000000000001"
extractor.user_id = "00000000-0000-0000-0000-000000000002"
extractor._session = sqlite_session
with pytest.raises(RuntimeError, match="storage failure"):
extractor._extract_images_from_docx(doc)
session.add_all.assert_not_called()
session.commit.assert_not_called()
assert sqlite_session.scalars(select(UploadFile)).all() == []
def test_extract_images_from_docx_uses_internal_files_url():
@@ -1,14 +1,40 @@
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import override
from unittest.mock import Mock, patch
from uuid import uuid4
import httpx
import pytest
from sqlalchemy.orm import Session
from core.entities.knowledge_entities import PreviewDetail
from core.rag.index_processor.constant.doc_type import DocType
from core.rag.index_processor.index_processor_base import BaseIndexProcessor
from core.rag.models.document import AttachmentDocument, Document
from extensions.storage.storage_type import StorageType
from models.enums import CreatorUserRole
from models.model import UploadFile
from models.tools import ToolFile
def _persist_upload(session: Session, *, upload_id: str, name: str) -> UploadFile:
upload = UploadFile(
tenant_id=str(uuid4()),
storage_type=StorageType.LOCAL,
key=f"uploads/{name}",
name=name,
size=4,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
created_at=datetime.now(UTC),
used=True,
)
upload.id = upload_id
session.add(upload)
return upload
class _ForwardingBaseIndexProcessor(BaseIndexProcessor):
@@ -68,21 +94,23 @@ class TestBaseIndexProcessor:
def processor(self) -> _ForwardingBaseIndexProcessor:
return _ForwardingBaseIndexProcessor()
def test_abstract_methods_raise_not_implemented(self, processor: _ForwardingBaseIndexProcessor) -> None:
def test_abstract_methods_raise_not_implemented(
self, processor: _ForwardingBaseIndexProcessor, unbound_session: Session
) -> None:
with pytest.raises(NotImplementedError):
processor.extract(Mock(), session=Mock())
processor.extract(Mock(), session=unbound_session)
with pytest.raises(NotImplementedError):
processor.transform([], session=Mock())
processor.transform([], session=unbound_session)
with pytest.raises(NotImplementedError):
processor.generate_summary_preview(
"tenant", [PreviewDetail(content="c")], {"enable": False}, session=Mock()
"tenant", [PreviewDetail(content="c")], {"enable": False}, session=unbound_session
)
with pytest.raises(NotImplementedError):
processor.load(Mock(), [], session=Mock())
processor.load(Mock(), [], session=unbound_session)
with pytest.raises(NotImplementedError):
processor.clean(Mock(), None, session=Mock())
processor.clean(Mock(), None, session=unbound_session)
with pytest.raises(NotImplementedError):
processor.index(Mock(), Mock(), {}, Mock())
processor.index(Mock(), Mock(), {}, unbound_session)
with pytest.raises(NotImplementedError):
processor.format_preview([])
@@ -123,12 +151,14 @@ class TestBaseIndexProcessor:
images = processor._extract_markdown_images(markdown)
assert images == ["https://a/img.png", "/files/123/file-preview"]
def test_get_content_files_without_images_returns_empty(self, processor: _ForwardingBaseIndexProcessor) -> None:
def test_get_content_files_without_images_returns_empty(
self, processor: _ForwardingBaseIndexProcessor, unbound_session: Session
) -> None:
document = Document(page_content="no image markdown", metadata={"document_id": "doc-1", "dataset_id": "ds-1"})
assert processor._get_content_files(document, session=Mock()) == []
assert processor._get_content_files(document, session=unbound_session) == []
def test_get_content_files_handles_all_sources_and_duplicates(
self, processor: _ForwardingBaseIndexProcessor
self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session
) -> None:
document = Document(page_content="ignored", metadata={"document_id": "doc-1", "dataset_id": "ds-1"})
images = [
@@ -138,22 +168,19 @@ class TestBaseIndexProcessor:
"/files/tools/cccccccc-cccc-cccc-cccc-cccccccccccc.png",
"https://example.com/remote.png?x=1",
]
upload_a = SimpleNamespace(id="aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", name="a.png")
upload_b = SimpleNamespace(id="bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb", name="b.png")
upload_tool = SimpleNamespace(id="tool-upload-id", name="tool.png")
upload_remote = SimpleNamespace(id="remote-upload-id", name="remote.png")
scalars_result = Mock()
scalars_result.all.return_value = [upload_a, upload_b, upload_tool, upload_remote]
db_session = Mock()
db_session.scalars.return_value = scalars_result
_persist_upload(sqlite_session, upload_id="aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", name="a.png")
_persist_upload(sqlite_session, upload_id="bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb", name="b.png")
tool_upload = _persist_upload(sqlite_session, upload_id=str(uuid4()), name="tool.png")
remote_upload = _persist_upload(sqlite_session, upload_id=str(uuid4()), name="remote.png")
sqlite_session.commit()
current_user = Mock()
with (
patch.object(processor, "_extract_markdown_images", return_value=images),
patch.object(processor, "_download_tool_file", return_value="tool-upload-id") as mock_tool_download,
patch.object(processor, "_download_image", return_value="remote-upload-id") as mock_image_download,
patch.object(processor, "_download_tool_file", return_value=tool_upload.id) as mock_tool_download,
patch.object(processor, "_download_image", return_value=remote_upload.id) as mock_image_download,
):
files = processor._get_content_files(document, current_user=current_user, session=db_session)
files = processor._get_content_files(document, current_user=current_user, session=sqlite_session)
assert len(files) == 5
assert all(isinstance(file, AttachmentDocument) for file in files)
@@ -165,33 +192,30 @@ class TestBaseIndexProcessor:
mock_tool_download.assert_called_once_with(
"cccccccc-cccc-cccc-cccc-cccccccccccc",
current_user,
session=db_session,
session=sqlite_session,
)
mock_image_download.assert_called_once()
def test_get_content_files_skips_tool_and_remote_download_without_user(
self, processor: _ForwardingBaseIndexProcessor
self, processor: _ForwardingBaseIndexProcessor, unbound_session: Session
) -> None:
document = Document(page_content="ignored", metadata={"document_id": "doc-1", "dataset_id": "ds-1"})
images = ["/files/tools/cccccccc-cccc-cccc-cccc-cccccccccccc.png", "https://example.com/remote.png"]
with patch.object(processor, "_extract_markdown_images", return_value=images):
files = processor._get_content_files(document, current_user=None, session=Mock())
files = processor._get_content_files(document, current_user=None, session=unbound_session)
assert files == []
def test_get_content_files_ignores_missing_upload_records(self, processor: _ForwardingBaseIndexProcessor) -> None:
def test_get_content_files_ignores_missing_upload_records(
self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session
) -> None:
document = Document(page_content="ignored", metadata={"document_id": "doc-1", "dataset_id": "ds-1"})
images = ["/files/aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa/image-preview"]
scalars_result = Mock()
scalars_result.all.return_value = []
db_session = Mock()
db_session.scalars.return_value = scalars_result
with (
patch.object(processor, "_extract_markdown_images", return_value=images),
):
files = processor._get_content_files(document, session=db_session)
files = processor._get_content_files(document, session=sqlite_session)
assert files == []
@@ -270,16 +294,25 @@ class TestBaseIndexProcessor:
):
assert processor._download_image("https://example.com/image.png", current_user=Mock()) is None
def test_download_tool_file_returns_none_when_not_found(self, processor: _ForwardingBaseIndexProcessor) -> None:
db_session = Mock()
db_session.get.return_value = None
def test_download_tool_file_returns_none_when_not_found(
self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session
) -> None:
assert processor._download_tool_file(str(uuid4()), current_user=Mock(), session=sqlite_session) is None
assert processor._download_tool_file("tool-id", current_user=Mock(), session=db_session) is None
def test_download_tool_file_uploads_file_when_found(self, processor: _ForwardingBaseIndexProcessor) -> None:
tool_file = SimpleNamespace(file_key="k1", name="tool.png", mimetype="image/png")
db_session = Mock()
db_session.get.return_value = tool_file
def test_download_tool_file_uploads_file_when_found(
self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session
) -> None:
tool_file = ToolFile(
user_id=str(uuid4()),
tenant_id=str(uuid4()),
conversation_id=None,
file_key="k1",
mimetype="image/png",
name="tool.png",
size=4,
)
sqlite_session.add(tool_file)
sqlite_session.commit()
mock_db = Mock()
mock_db.engine = Mock()
upload_result = SimpleNamespace(id="upload-id")
@@ -290,7 +323,7 @@ class TestBaseIndexProcessor:
patch("services.file_service.FileService") as mock_file_service,
):
mock_file_service.return_value.upload_file.return_value = upload_result
result = processor._download_tool_file("tool-id", current_user=Mock(), session=db_session)
result = processor._download_tool_file(tool_file.id, current_user=Mock(), session=sqlite_session)
assert result == "upload-id"
mock_load.assert_called_once_with("k1")
@@ -1,7 +1,7 @@
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, PropertyMock, call, patch
from unittest.mock import MagicMock, Mock, PropertyMock, patch
import pytest
from flask import has_app_context
@@ -297,9 +297,6 @@ def test_to_system_configuration_never_returns_hosting_credentials_for_package_w
def test_to_system_configuration_uses_owned_session_for_cloud_credit_pools() -> None:
provider_entity = _build_plugin_provider_declaration(PluginInstallationSource.Marketplace)
manager = _build_provider_manager()
owned_session = Mock()
session_context = MagicMock()
session_context.__enter__.return_value = owned_session
trial_pool = SimpleNamespace(quota_used=0, quota_limit=100)
paid_pool = SimpleNamespace(quota_used=0, quota_limit=0)
@@ -313,11 +310,6 @@ def test_to_system_configuration_uses_owned_session_for_cloud_credit_pools() ->
"core.plugin.plugin_service.PluginService.is_plugin_verified",
return_value=True,
),
patch.object(
provider_manager_module.session_factory,
"create_session",
return_value=session_context,
) as create_session,
patch(
"services.credit_pool_service.CreditPoolService.get_pool",
side_effect=[trial_pool, paid_pool],
@@ -333,12 +325,13 @@ def test_to_system_configuration_uses_owned_session_for_cloud_credit_pools() ->
).result()
assert configuration.enabled is True
create_session.assert_called_once_with()
assert get_pool.call_args_list == [
call(tenant_id="tenant-id", pool_type=ProviderQuotaType.TRIAL, session=owned_session),
call(tenant_id="tenant-id", pool_type=ProviderQuotaType.PAID, session=owned_session),
assert [call.kwargs["pool_type"] for call in get_pool.call_args_list] == [
ProviderQuotaType.TRIAL,
ProviderQuotaType.PAID,
]
session_context.__exit__.assert_called_once_with(None, None, None)
owned_sessions = [call.kwargs["session"] for call in get_pool.call_args_list]
assert all(isinstance(session, Session) for session in owned_sessions)
assert owned_sessions[0] is owned_sessions[1]
def test_to_system_configuration_preserves_marketplace_behavior() -> None:
@@ -94,7 +94,6 @@ def _build_tool(runtime: ToolRuntime | None = None) -> DummyTool:
return DummyTool(entity=entity, runtime=runtime)
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_invoke_supports_single_message_and_parameter_casting(sqlite_session: Session):
runtime = ToolRuntime(
tenant_id="tenant-1",
@@ -133,7 +132,6 @@ def test_invoke_supports_single_message_and_parameter_casting(sqlite_session: Se
}
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_invoke_preserves_multiple_select_values(sqlite_session: Session):
tool = _build_tool()
parameter = ToolParameter.get_simple_instance(
@@ -154,7 +152,6 @@ def test_invoke_preserves_multiple_select_values(sqlite_session: Session):
tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={"choice": "a"})
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_invoke_supports_list_and_generator_results(sqlite_session: Session):
tool = _build_tool()
tool.result = [tool.create_text_message("a"), tool.create_text_message("b")]
@@ -375,7 +372,6 @@ def test_message_factory_helpers():
assert variable_message.message.stream is False
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_base_abstract_invoke_placeholder_returns_none(sqlite_session: Session):
tool = _build_tool()
assert Tool._invoke(tool, session=sqlite_session, user_id="u", tool_parameters={}) is None
@@ -241,7 +241,6 @@ def test_do_http_request_builds_arguments_and_handles_invalid_method(monkeypatch
invalid_method_tool.do_http_request("https://api.example.com", "TRACE", headers={}, parameters={})
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session):
openapi = {
"parameters": [],
@@ -2,15 +2,20 @@ from __future__ import annotations
import uuid
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import Mock, patch
from datetime import datetime
from unittest.mock import patch
import pytest
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from yaml import YAMLError
from core.app.app_config.entities import DatasetRetrieveConfigEntity
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.model_manager import ModelInstance, ModelManager
from core.provider_manager import ProviderManager
from core.rag.embedding.retrieval import RetrievalSegments
from core.rag.models.document import Document as RagDocument
from core.rag.rerank.rerank_model import RerankModelRunner
from core.tools.utils.dataset_retriever import dataset_multi_retriever_tool as multi_retriever_module
from core.tools.utils.dataset_retriever import dataset_retriever_tool as single_retriever_module
from core.tools.utils.dataset_retriever.dataset_multi_retriever_tool import DatasetMultiRetrieverTool
@@ -18,17 +23,88 @@ from core.tools.utils.dataset_retriever.dataset_retriever_tool import DatasetRet
from core.tools.utils.text_processing_utils import remove_leading_symbols
from core.tools.utils.uuid_utils import is_valid_uuid
from core.tools.utils.yaml_utils import _load_yaml_file, load_yaml_file_cached
from models.dataset import Dataset, Document, DocumentSegment
from models.enums import DataSourceType, DocumentCreatedFrom, SegmentStatus
def _retrieve_config() -> DatasetRetrieveConfigEntity:
return DatasetRetrieveConfigEntity(retrieve_strategy=DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE)
def _persist_dataset(
session: Session,
*,
dataset_id: str | None = None,
tenant_id: str | None = None,
name: str = "Knowledge Base",
provider: str = "vendor",
retrieval_model: dict | None = None,
) -> Dataset:
dataset = Dataset(
id=dataset_id or str(uuid.uuid4()),
tenant_id=tenant_id or str(uuid.uuid4()),
name=name,
provider=provider,
data_source_type=DataSourceType.UPLOAD_FILE,
indexing_technique="high_quality",
retrieval_model=retrieval_model,
created_by=str(uuid.uuid4()),
)
session.add(dataset)
session.commit()
return dataset
def _persist_document(
session: Session,
*,
dataset: Dataset,
name: str,
data_source_type: DataSourceType = DataSourceType.UPLOAD_FILE,
doc_metadata: dict | None = None,
) -> Document:
document = Document(
id=str(uuid.uuid4()),
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
position=1,
data_source_type=data_source_type,
batch="batch-1",
name=name,
created_from=DocumentCreatedFrom.WEB,
created_by=dataset.created_by,
doc_metadata=doc_metadata,
)
session.add(document)
session.commit()
return document
class _FakeFlaskApp:
def app_context(self):
return nullcontext()
class _FakeCurrentApp:
def _get_current_object(self) -> _FakeFlaskApp:
return _FakeFlaskApp()
class _DatabaseWithSession:
def __init__(self, session: scoped_session[Session]) -> None:
self.session = session
class _UnusedProviderManager(ProviderManager):
def __init__(self) -> None:
pass
class _UnusedModelInstance(ModelInstance):
def __init__(self) -> None:
pass
class _ImmediateThread:
def __init__(self, target=None, kwargs=None, **_kwargs):
self._target = target
@@ -105,8 +181,13 @@ def test_load_yaml_file_cached_hits(tmp_path):
assert load_yaml_file_cached.cache_info().hits == 1
def test_single_dataset_retriever_from_dataset_builds_name_and_description():
dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1", name="Knowledge", description=None)
def test_single_dataset_retriever_from_dataset_builds_name_and_description(sqlite_session: Session):
dataset = _persist_dataset(
sqlite_session,
dataset_id="dataset-1",
tenant_id="tenant-1",
name="Knowledge",
)
tool = SingleDatasetRetrieverTool.from_dataset(
dataset=dataset,
@@ -120,31 +201,21 @@ def test_single_dataset_retriever_from_dataset_builds_name_and_description():
assert tool.description == "useful for when you want to answer queries about the Knowledge"
def test_single_dataset_retriever_external_run_returns_content_and_resources():
dataset = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
name="Knowledge Base",
provider="external",
indexing_technique="high_quality",
retrieval_model={},
)
def test_single_dataset_retriever_external_run_returns_content_and_resources(sqlite_session: Session):
dataset = _persist_dataset(sqlite_session, provider="external", retrieval_model={})
callback = _TestHitCallback()
dataset_retrieval = Mock()
dataset_retrieval.get_metadata_filter_condition.return_value = (
{"dataset-1": ["doc-a"]},
metadata_filter_result = (
{dataset.id: ["doc-a"]},
{"logical_operator": "and"},
)
session = Mock()
session.scalar.return_value = dataset
external_documents = [
{"content": "first", "metadata": {"document_id": "doc-a"}, "score": 0.9, "title": "Doc A"},
{"content": "second", "metadata": {"document_id": "doc-b"}, "score": 0.8, "title": "Doc B"},
]
tool = SingleDatasetRetrieverTool(
tenant_id="tenant-1",
dataset_id="dataset-1",
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
retrieve_config=_retrieve_config(),
return_resource=True,
retriever_from="dev",
@@ -152,41 +223,35 @@ def test_single_dataset_retriever_external_run_returns_content_and_resources():
inputs={"x": 1},
)
with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval):
with patch.object(
single_retriever_module.DatasetRetrieval,
"get_metadata_filter_condition",
return_value=metadata_filter_result,
):
with patch.object(
single_retriever_module.ExternalDatasetService,
"fetch_external_knowledge_retrieval",
return_value=external_documents,
) as fetch_mock:
result = tool.run(session=session, query="hello")
result = tool.run(session=sqlite_session, query="hello")
assert result == "first\nsecond"
assert callback.queries == [("hello", "dataset-1")]
assert callback.queries == [("hello", dataset.id)]
assert callback.resources is not None
resource_info = callback.resources
assert [item.position for item in resource_info] == [1, 2]
assert resource_info[0].dataset_id == "dataset-1"
assert resource_info[0].dataset_id == dataset.id
fetch_mock.assert_called_once()
assert fetch_mock.call_args.kwargs["session"] is session
assert fetch_mock.call_args.kwargs["session"] is sqlite_session
def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_documents():
dataset = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
name="Knowledge Base",
provider="internal",
indexing_technique="high_quality",
retrieval_model=None,
)
dataset_retrieval = Mock()
dataset_retrieval.get_metadata_filter_condition.return_value = ({"dataset-1": []}, {"logical_operator": "and"})
session = Mock()
session.scalar.return_value = dataset
def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_documents(
sqlite_session: Session,
):
dataset = _persist_dataset(sqlite_session)
tool = SingleDatasetRetrieverTool(
tenant_id="tenant-1",
dataset_id="dataset-1",
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
retrieve_config=_retrieve_config(),
return_resource=False,
retriever_from="prod",
@@ -194,21 +259,21 @@ def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_do
inputs={},
)
with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval):
with patch.object(
single_retriever_module.DatasetRetrieval,
"get_metadata_filter_condition",
return_value=({dataset.id: []}, {"logical_operator": "and"}),
):
with patch.object(single_retriever_module.RetrievalService, "retrieve") as retrieve_mock:
result = tool.run(session=session, query="hello")
result = tool.run(session=sqlite_session, query="hello")
assert result == ""
retrieve_mock.assert_not_called()
def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources():
dataset = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
name="Knowledge Base",
provider="internal",
indexing_technique="high_quality",
def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(sqlite_session: Session):
dataset = _persist_dataset(
sqlite_session,
retrieval_model={
"search_method": "semantic_search",
"score_threshold_enabled": True,
@@ -219,54 +284,64 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources():
"weights": {"vector_setting": {"vector_weight": 0.6}},
},
)
document_low = _persist_document(
sqlite_session,
dataset=dataset,
name="Document Low",
doc_metadata={"lang": "en"},
)
document_high = _persist_document(
sqlite_session,
dataset=dataset,
name="Document High",
data_source_type=DataSourceType.NOTION_IMPORT,
doc_metadata={"lang": "fr"},
)
callback = _TestHitCallback()
dataset_retrieval = Mock()
dataset_retrieval.get_metadata_filter_condition.return_value = (None, None)
low_segment = SimpleNamespace(
id="seg-low",
dataset_id="dataset-1",
document_id="doc-low",
low_segment = DocumentSegment(
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
document_id=document_low.id,
index_node_id="node-low",
content="raw low",
answer="low answer",
hit_count=1,
word_count=10,
position=3,
index_node_hash="hash-low",
get_sign_content=lambda: "signed low",
tokens=10,
created_by=dataset.created_by,
answer="low answer",
status=SegmentStatus.COMPLETED,
completed_at=datetime.now(),
)
high_segment = SimpleNamespace(
id="seg-high",
dataset_id="dataset-1",
document_id="doc-high",
high_segment = DocumentSegment(
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
document_id=document_high.id,
index_node_id="node-high",
content="raw high",
answer=None,
hit_count=9,
word_count=25,
position=1,
index_node_hash="hash-high",
get_sign_content=lambda: "signed high",
tokens=25,
created_by=dataset.created_by,
status=SegmentStatus.COMPLETED,
completed_at=datetime.now(),
)
sqlite_session.add_all([low_segment, high_segment])
sqlite_session.commit()
records = [
SimpleNamespace(segment=low_segment, score=0.2, summary="summary low"),
SimpleNamespace(segment=high_segment, score=0.9, summary=None),
RetrievalSegments(segment=low_segment, score=0.2, summary="summary low"),
RetrievalSegments(segment=high_segment, score=0.9),
]
documents = [
RagDocument(page_content="first", metadata={"doc_id": "node-low", "score": 0.2}),
RagDocument(page_content="second", metadata={"doc_id": "node-high", "score": 0.9}),
]
lookup_doc_low = SimpleNamespace(
id="doc-low", name="Document Low", data_source_type="upload_file", doc_metadata={"lang": "en"}
)
lookup_doc_high = SimpleNamespace(
id="doc-high", name="Document High", data_source_type="notion", doc_metadata={"lang": "fr"}
)
session = Mock()
session.scalar.side_effect = [dataset, lookup_doc_low, lookup_doc_high]
session.get.return_value = dataset
tool = SingleDatasetRetrieverTool(
tenant_id="tenant-1",
dataset_id="dataset-1",
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
retrieve_config=_retrieve_config(),
return_resource=True,
retriever_from="dev",
@@ -275,21 +350,28 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources():
top_k=2,
)
with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval):
with patch.object(single_retriever_module.RetrievalService, "retrieve", return_value=documents):
with patch.object(
single_retriever_module.RetrievalService,
"format_retrieval_documents",
return_value=records,
):
result = tool.run(session=session, query="hello")
with (
patch.object(
single_retriever_module.DatasetRetrieval,
"get_metadata_filter_condition",
return_value=(None, None),
),
patch.object(single_retriever_module.RetrievalService, "retrieve", return_value=documents),
patch.object(
single_retriever_module.RetrievalService,
"format_retrieval_documents",
return_value=records,
),
patch.object(DocumentSegment, "get_sign_content", lambda segment: segment.content.replace("raw", "signed")),
):
result = tool.run(session=sqlite_session, query="hello")
assert result == "signed high\nsummary low\nquestion:signed low answer:low answer"
assert callback.documents == documents
assert callback.resources is not None
resource_info = callback.resources
assert [item.position for item in resource_info] == [1, 2]
assert resource_info[0].segment_id == "seg-high"
assert resource_info[0].segment_id == high_segment.id
assert resource_info[0].hit_count == 9
assert resource_info[1].summary == "summary low"
assert resource_info[1].content == "question:raw low \nanswer:low answer"
@@ -308,11 +390,12 @@ def test_multi_dataset_retriever_from_dataset_sets_tool_name():
assert tool.name == "dataset_tenant_1"
def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing():
def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing(
sqlite_session_factory: sessionmaker[Session],
):
callback = _TestHitCallback()
all_documents: list[RagDocument] = []
db_session = Mock()
db_session.scalar.return_value = None
db_session = scoped_session(sqlite_session_factory)
tool = DatasetMultiRetrieverTool(
tenant_id="tenant-1",
dataset_ids=["dataset-1"],
@@ -322,15 +405,18 @@ def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing
retriever_from="prod",
)
with patch.object(multi_retriever_module, "db", SimpleNamespace(session=db_session)):
with patch.object(multi_retriever_module.RetrievalService, "retrieve") as retrieve_mock:
result = tool._retriever(
flask_app=_FakeFlaskApp(),
dataset_id="dataset-1",
query="hello",
all_documents=all_documents,
hit_callbacks=[callback],
)
try:
with patch.object(multi_retriever_module, "db", _DatabaseWithSession(db_session)):
with patch.object(multi_retriever_module.RetrievalService, "retrieve") as retrieve_mock:
result = tool._retriever(
flask_app=_FakeFlaskApp(),
dataset_id=str(uuid.uuid4()),
query="hello",
all_documents=all_documents,
hit_callbacks=[callback],
)
finally:
db_session.remove()
assert result == []
assert all_documents == []
@@ -338,11 +424,12 @@ def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing
retrieve_mock.assert_not_called()
def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model():
dataset = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
indexing_technique="high_quality",
def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
):
dataset = _persist_dataset(
sqlite_session,
retrieval_model={
"search_method": "semantic_search",
"top_k": 6,
@@ -356,11 +443,10 @@ def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model():
callback = _TestHitCallback()
documents = [RagDocument(page_content="retrieved", metadata={"doc_id": "node-1", "score": 0.4})]
all_documents: list[RagDocument] = []
db_session = Mock()
db_session.scalar.return_value = dataset
db_session = scoped_session(sqlite_session_factory)
tool = DatasetMultiRetrieverTool(
tenant_id="tenant-1",
dataset_ids=["dataset-1"],
tenant_id=dataset.tenant_id,
dataset_ids=[dataset.id],
reranking_provider_name="provider",
reranking_model_name="model",
return_resource=False,
@@ -368,21 +454,26 @@ def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model():
top_k=2,
)
with patch.object(multi_retriever_module, "db", SimpleNamespace(session=db_session)):
with patch.object(multi_retriever_module.RetrievalService, "retrieve", return_value=documents) as retrieve_mock:
tool._retriever(
flask_app=_FakeFlaskApp(),
dataset_id="dataset-1",
query="hello",
all_documents=all_documents,
hit_callbacks=[callback],
)
try:
with patch.object(multi_retriever_module, "db", _DatabaseWithSession(db_session)):
with patch.object(
multi_retriever_module.RetrievalService, "retrieve", return_value=documents
) as retrieve_mock:
tool._retriever(
flask_app=_FakeFlaskApp(),
dataset_id=dataset.id,
query="hello",
all_documents=all_documents,
hit_callbacks=[callback],
)
finally:
db_session.remove()
assert all_documents == documents
assert callback.queries == [("hello", "dataset-1")]
assert callback.queries == [("hello", dataset.id)]
retrieve_mock.assert_called_once_with(
retrieval_method="semantic_search",
dataset_id="dataset-1",
dataset_id=dataset.id,
query="hello",
top_k=6,
score_threshold=0.4,
@@ -392,11 +483,26 @@ def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model():
)
def test_multi_dataset_retriever_run_orders_segments_and_returns_resources():
def test_multi_dataset_retriever_run_orders_segments_and_returns_resources(sqlite_session: Session):
dataset_one = _persist_dataset(sqlite_session, name="Dataset One")
dataset_two = _persist_dataset(sqlite_session, tenant_id=dataset_one.tenant_id, name="Dataset Two")
document_two = _persist_document(
sqlite_session,
dataset=dataset_one,
name="Doc Two",
data_source_type=DataSourceType.NOTION_IMPORT,
doc_metadata={"p": 2},
)
document_one = _persist_document(
sqlite_session,
dataset=dataset_two,
name="Doc One",
doc_metadata={"p": 1},
)
callback = _TestHitCallback()
tool = DatasetMultiRetrieverTool(
tenant_id="tenant-1",
dataset_ids=["dataset-1", "dataset-2"],
tenant_id=dataset_one.tenant_id,
dataset_ids=[dataset_one.id, dataset_two.id],
reranking_provider_name="provider",
reranking_model_name="model",
return_resource=True,
@@ -409,64 +515,63 @@ def test_multi_dataset_retriever_run_orders_segments_and_returns_resources():
second_doc = RagDocument(page_content="second", metadata={"doc_id": "node-1", "score": 0.9})
def fake_retriever(**kwargs):
if kwargs["dataset_id"] == "dataset-1":
if kwargs["dataset_id"] == dataset_one.id:
kwargs["all_documents"].append(first_doc)
else:
kwargs["all_documents"].append(second_doc)
segment_for_node_2 = SimpleNamespace(
id="seg-2",
dataset_id="dataset-1",
document_id="doc-2",
segment_for_node_2 = DocumentSegment(
tenant_id=dataset_one.tenant_id,
dataset_id=dataset_one.id,
document_id=document_two.id,
index_node_id="node-2",
content="raw two",
answer="answer two",
hit_count=2,
word_count=20,
position=2,
index_node_hash="hash-2",
get_sign_content=lambda: "signed two",
tokens=20,
created_by=dataset_one.created_by,
answer="answer two",
status=SegmentStatus.COMPLETED,
completed_at=datetime.now(),
)
segment_for_node_1 = SimpleNamespace(
id="seg-1",
dataset_id="dataset-2",
document_id="doc-1",
segment_for_node_1 = DocumentSegment(
tenant_id=dataset_two.tenant_id,
dataset_id=dataset_two.id,
document_id=document_one.id,
index_node_id="node-1",
content="raw one",
answer=None,
hit_count=7,
word_count=30,
position=1,
index_node_hash="hash-1",
get_sign_content=lambda: "signed one",
tokens=30,
created_by=dataset_two.created_by,
status=SegmentStatus.COMPLETED,
completed_at=datetime.now(),
)
db_session = Mock()
db_session.scalars.return_value.all.return_value = [segment_for_node_2, segment_for_node_1]
db_session.get.side_effect = [
SimpleNamespace(id="dataset-2", name="Dataset Two"),
SimpleNamespace(id="dataset-1", name="Dataset One"),
]
db_session.scalar.side_effect = [
SimpleNamespace(id="doc-1", name="Doc One", data_source_type="upload_file", doc_metadata={"p": 1}),
SimpleNamespace(id="doc-2", name="Doc Two", data_source_type="notion", doc_metadata={"p": 2}),
]
model_manager = Mock()
model_manager.get_model_instance.return_value = Mock()
rerank_runner = Mock()
rerank_runner.run.return_value = [second_doc, first_doc]
fake_current_app = SimpleNamespace(_get_current_object=lambda: _FakeFlaskApp())
sqlite_session.add_all([segment_for_node_2, segment_for_node_1])
sqlite_session.commit()
model_manager = ModelManager(provider_manager=_UnusedProviderManager())
model_instance = _UnusedModelInstance()
rerank_runner = RerankModelRunner(model_instance, session=sqlite_session)
fake_current_app = _FakeCurrentApp()
with patch.object(tool, "_retriever", side_effect=fake_retriever) as retriever_mock:
with patch.object(multi_retriever_module, "current_app", fake_current_app):
with patch.object(multi_retriever_module.threading, "Thread", _ImmediateThread):
with patch.object(multi_retriever_module.ModelManager, "for_tenant", return_value=model_manager):
with patch.object(
multi_retriever_module, "RerankModelRunner", return_value=rerank_runner
) as rerank_runner_class:
result = tool.run(session=db_session, query="hello")
with (
patch.object(DocumentSegment, "get_sign_content", lambda segment: segment.content.replace("raw", "signed")),
patch.object(tool, "_retriever", side_effect=fake_retriever) as retriever_mock,
patch.object(multi_retriever_module, "current_app", fake_current_app),
patch.object(multi_retriever_module.threading, "Thread", _ImmediateThread),
patch.object(multi_retriever_module.ModelManager, "for_tenant", return_value=model_manager),
patch.object(model_manager, "get_model_instance", return_value=model_instance),
patch.object(multi_retriever_module, "RerankModelRunner", return_value=rerank_runner) as rerank_runner_class,
patch.object(rerank_runner, "run", return_value=[second_doc, first_doc]),
):
result = tool.run(session=sqlite_session, query="hello")
assert result == "signed one\nquestion:signed two answer:answer two"
rerank_runner_class.assert_called_once_with(model_manager.get_model_instance.return_value, session=db_session)
rerank_runner_class.assert_called_once_with(model_instance, session=sqlite_session)
assert retriever_mock.call_count == 2
assert callback.documents == [second_doc, first_doc]
assert callback.resources is not None
@@ -89,7 +89,6 @@ def test_calculate_tokens_handles_missing_model():
mock_factory.assert_called_once_with(tenant_id="tenant", user_id=None)
@pytest.mark.parametrize("sqlite_session", [(ToolModelInvoke,)], indirect=True)
def test_invoke_success_and_error_mappings(sqlite_session: Session):
model_instance = _mock_model_instance(schema={ModelPropertyKey.CONTEXT_SIZE: 2048})
model_instance.invoke_llm.return_value = SimpleNamespace(
@@ -158,7 +157,6 @@ def test_invoke_success_and_error_mappings(sqlite_session: Session):
"generic-error",
],
)
@pytest.mark.parametrize("sqlite_session", [(ToolModelInvoke,)], indirect=True)
def test_invoke_error_mappings(exc, expected, sqlite_session: Session):
model_instance = _mock_model_instance(schema={ModelPropertyKey.CONTEXT_SIZE: 2048})
model_instance.invoke_llm.side_effect = exc
@@ -1,9 +1,11 @@
import time
import uuid
from unittest.mock import MagicMock, Mock
from unittest.mock import Mock
import pytest
from pytest_mock import MockerFixture
from sqlalchemy import event
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.rag.index_processor.constant.index_type import IndexTechniqueType
@@ -248,7 +250,6 @@ class TestKnowledgeIndexNode:
def test_run_preview_mode_success(
self,
mocker: MockerFixture,
mock_graph_init_params,
mock_graph_runtime_state,
mock_index_processor,
@@ -283,14 +284,6 @@ class TestKnowledgeIndexNode:
total_segments=2,
)
mock_index_processor.get_preview_output.return_value = mock_preview
session = MagicMock()
session_context = MagicMock()
session_context.__enter__.return_value = session
mocker.patch(
"core.workflow.nodes.knowledge_index.knowledge_index_node.session_factory.create_session",
return_value=session_context,
)
node_id = str(uuid.uuid4())
config = {
"id": node_id,
@@ -310,7 +303,7 @@ class TestKnowledgeIndexNode:
# Assert
assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED
assert result.outputs is not None
assert mock_index_processor.get_preview_output.call_args.kwargs["session"] is session
assert isinstance(mock_index_processor.get_preview_output.call_args.kwargs["session"], Session)
def test_run_production_mode_success(
self,
@@ -548,6 +541,7 @@ class TestKnowledgeIndexNode:
mock_index_processor,
mock_summary_index_service,
sample_node_data,
sqlite_session: Session,
):
# Arrange
dataset_id = str(uuid.uuid4())
@@ -572,7 +566,9 @@ class TestKnowledgeIndexNode:
)
# Act
session = MagicMock()
session = sqlite_session
commits: list[str] = []
event.listen(session, "after_commit", lambda _session: commits.append("commit"))
result = node._invoke_knowledge_index(
session=session,
dataset_id=dataset_id,
@@ -587,7 +583,7 @@ class TestKnowledgeIndexNode:
# Assert
assert mock_summary_index_service.generate_and_vectorize_summary.called
assert mock_index_processor.index_and_clean.called
session.commit.assert_called_once()
assert commits == ["commit"]
assert result == {"status": "indexed"}
def test_version_method(self):
@@ -637,6 +633,7 @@ class TestInvokeKnowledgeIndex:
mock_index_processor,
mock_summary_index_service,
sample_node_data,
sqlite_session: Session,
):
# Arrange
dataset_id = str(uuid.uuid4())
@@ -662,7 +659,9 @@ class TestInvokeKnowledgeIndex:
)
# Act
session = MagicMock()
session = sqlite_session
commits: list[str] = []
event.listen(session, "after_commit", lambda _session: commits.append("commit"))
result = node._invoke_knowledge_index(
session=session,
dataset_id=dataset_id,
@@ -687,5 +686,5 @@ class TestInvokeKnowledgeIndex:
summary_setting,
session=session,
)
session.commit.assert_called_once()
assert commits == ["commit"]
assert result == {"status": "indexed"}