mirror of
https://github.com/langgenius/dify.git
synced 2026-07-30 16:59:35 +08:00
test: use sqlite3 session in test_completion_completion_app_generator (#38723)
This commit is contained in:
parent
b855d71adb
commit
e12e86a308
@ -283,6 +283,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
|
||||
"""
|
||||
Generate App response.
|
||||
|
||||
:param session: caller-owned database session used for message and historical model-config reads
|
||||
:param app_model: App
|
||||
:param message_id: message ID
|
||||
:param user: account or end user
|
||||
|
||||
@ -1,20 +1,25 @@
|
||||
import contextlib
|
||||
from decimal import Decimal
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, call
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
from pytest_mock import MockerFixture
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import core.app.apps.completion.app_generator as module
|
||||
from core.app.apps.completion.app_generator import CompletionAppGenerator
|
||||
from core.app.apps.exc import GenerateTaskStoppedError
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from graphon.file import FILE_MODEL_IDENTITY
|
||||
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||
from models.enums import ConversationFromSource
|
||||
from models.model import AppMode, AppModelConfig, Conversation, Message
|
||||
from services.errors.app import MoreLikeThisDisabledError
|
||||
from services.errors.message import MessageNotExistsError
|
||||
|
||||
TABLES = (AppModelConfig, Conversation, Message)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def generator(mocker: MockerFixture):
|
||||
@ -53,11 +58,71 @@ def _build_app_model_config():
|
||||
return config
|
||||
|
||||
|
||||
def _persist_message(
|
||||
session: Session,
|
||||
*,
|
||||
message_id: str = "msg",
|
||||
app_id: str = "app1",
|
||||
app_model_config: AppModelConfig | None = None,
|
||||
) -> Message:
|
||||
conversation = Conversation(
|
||||
app_id=app_id,
|
||||
app_model_config_id=app_model_config.id if app_model_config else None,
|
||||
model_provider=None,
|
||||
override_model_configs=None,
|
||||
model_id=None,
|
||||
mode=AppMode.COMPLETION,
|
||||
name="completion conversation",
|
||||
summary=None,
|
||||
inputs={},
|
||||
introduction="",
|
||||
system_instruction="",
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
from_source=ConversationFromSource.CONSOLE,
|
||||
from_end_user_id=None,
|
||||
from_account_id=None,
|
||||
read_at=None,
|
||||
read_account_id=None,
|
||||
)
|
||||
session.add(conversation)
|
||||
session.flush()
|
||||
|
||||
message = Message(
|
||||
id=message_id,
|
||||
app_id=app_id,
|
||||
model_provider=None,
|
||||
model_id=None,
|
||||
override_model_configs=None,
|
||||
conversation_id=conversation.id,
|
||||
inputs={"a": 1},
|
||||
query="q",
|
||||
message={},
|
||||
message_unit_price=Decimal(0),
|
||||
answer="",
|
||||
answer_unit_price=Decimal(0),
|
||||
parent_message_id=None,
|
||||
total_price=None,
|
||||
currency="USD",
|
||||
error=None,
|
||||
message_metadata=None,
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
from_source=ConversationFromSource.CONSOLE,
|
||||
from_end_user_id=None,
|
||||
from_account_id=None,
|
||||
workflow_run_id=None,
|
||||
app_mode=AppMode.COMPLETION,
|
||||
)
|
||||
session.add(message)
|
||||
session.flush()
|
||||
return message
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||
class TestCompletionAppGenerator:
|
||||
def test_generate_invalid_query_type(self, generator):
|
||||
def test_generate_invalid_query_type(self, generator, sqlite_session: Session):
|
||||
with pytest.raises(ValueError):
|
||||
generator.generate(
|
||||
session=MagicMock(),
|
||||
session=sqlite_session,
|
||||
app_model=_build_app_model(),
|
||||
user=_build_user(),
|
||||
args={"query": 123, "inputs": {}, "files": []},
|
||||
@ -65,10 +130,10 @@ class TestCompletionAppGenerator:
|
||||
streaming=True,
|
||||
)
|
||||
|
||||
def test_generate_override_not_debugger(self, generator):
|
||||
def test_generate_override_not_debugger(self, generator, sqlite_session: Session):
|
||||
with pytest.raises(ValueError):
|
||||
generator.generate(
|
||||
session=MagicMock(),
|
||||
session=sqlite_session,
|
||||
app_model=_build_app_model(),
|
||||
user=_build_user(),
|
||||
args={"query": "q", "inputs": {}, "files": [], "model_config": {}},
|
||||
@ -76,7 +141,7 @@ class TestCompletionAppGenerator:
|
||||
streaming=False,
|
||||
)
|
||||
|
||||
def test_generate_success_no_file_config(self, generator, mocker: MockerFixture):
|
||||
def test_generate_success_no_file_config(self, generator, mocker: MockerFixture, sqlite_session: Session):
|
||||
app_model_config = _build_app_model_config()
|
||||
mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config)
|
||||
annotation_reply = {"enabled": False}
|
||||
@ -105,9 +170,8 @@ class TestCompletionAppGenerator:
|
||||
mocker.patch.object(generator, "_handle_response", return_value="response")
|
||||
mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted")
|
||||
|
||||
session = MagicMock()
|
||||
result = generator.generate(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
app_model=_build_app_model(),
|
||||
user=_build_user(),
|
||||
args={"query": "q", "inputs": {"a": 1}, "files": [], "trace_session_id": "session-1"},
|
||||
@ -118,11 +182,11 @@ class TestCompletionAppGenerator:
|
||||
assert result == "converted"
|
||||
assert generator.generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1"
|
||||
module.file_factory.build_from_mappings.assert_not_called()
|
||||
load_annotation_reply_config.assert_called_once_with(session, "app1")
|
||||
load_annotation_reply_config.assert_called_once_with(sqlite_session, "app1")
|
||||
app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply)
|
||||
assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply
|
||||
|
||||
def test_generate_success_with_files(self, generator, mocker: MockerFixture):
|
||||
def test_generate_success_with_files(self, generator, mocker: MockerFixture, sqlite_session: Session):
|
||||
app_model_config = _build_app_model_config()
|
||||
mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config)
|
||||
|
||||
@ -144,7 +208,7 @@ class TestCompletionAppGenerator:
|
||||
mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted")
|
||||
|
||||
result = generator.generate(
|
||||
session=MagicMock(),
|
||||
session=sqlite_session,
|
||||
app_model=_build_app_model(),
|
||||
user=_build_user(),
|
||||
args={"query": "q", "inputs": {"a": 1}, "files": [{"id": "f"}]},
|
||||
@ -155,7 +219,7 @@ class TestCompletionAppGenerator:
|
||||
assert result == "converted"
|
||||
module.file_factory.build_from_mappings.assert_called_once()
|
||||
|
||||
def test_generate_override_model_config_debugger(self, generator, mocker: MockerFixture):
|
||||
def test_generate_override_model_config_debugger(self, generator, mocker: MockerFixture, sqlite_session: Session):
|
||||
app_model_config = _build_app_model_config()
|
||||
mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config)
|
||||
|
||||
@ -180,7 +244,7 @@ class TestCompletionAppGenerator:
|
||||
mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted")
|
||||
|
||||
generator.generate(
|
||||
session=MagicMock(),
|
||||
session=sqlite_session,
|
||||
app_model=_build_app_model(),
|
||||
user=_build_user(),
|
||||
args={"query": "q", "inputs": {}, "files": [], "model_config": override_config},
|
||||
@ -190,118 +254,90 @@ class TestCompletionAppGenerator:
|
||||
|
||||
assert get_app_config.call_args.kwargs["override_config_dict"] == override_config
|
||||
|
||||
def test_generate_more_like_this_message_not_found(self, generator, mocker: MockerFixture):
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.return_value = None
|
||||
def test_generate_more_like_this_message_not_found(self, generator, sqlite_session: Session):
|
||||
_persist_message(sqlite_session, app_id="other-app")
|
||||
|
||||
with pytest.raises(MessageNotExistsError):
|
||||
generator.generate_more_like_this(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
app_model=_build_app_model(),
|
||||
message_id="msg",
|
||||
user=_build_user(),
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
)
|
||||
|
||||
def test_generate_more_like_this_disabled(self, generator, mocker: MockerFixture):
|
||||
def test_generate_more_like_this_disabled(self, generator, sqlite_session: Session):
|
||||
app_model = _build_app_model()
|
||||
current_config = MagicMock(more_like_this=False, more_like_this_dict={"enabled": False})
|
||||
|
||||
message = MagicMock()
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.return_value = message
|
||||
session.get.return_value = current_config
|
||||
current_config = AppModelConfig(app_id=app_model.id, more_like_this='{"enabled": false}')
|
||||
sqlite_session.add(current_config)
|
||||
sqlite_session.flush()
|
||||
app_model.app_model_config_id = current_config.id
|
||||
_persist_message(sqlite_session)
|
||||
|
||||
with pytest.raises(MoreLikeThisDisabledError):
|
||||
generator.generate_more_like_this(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
app_model=app_model,
|
||||
message_id="msg",
|
||||
user=_build_user(),
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
)
|
||||
|
||||
def test_generate_more_like_this_app_model_config_missing(self, generator, mocker: MockerFixture):
|
||||
def test_generate_more_like_this_app_model_config_missing(self, generator, sqlite_session: Session):
|
||||
app_model = _build_app_model()
|
||||
app_model.app_model_config_id = None
|
||||
|
||||
message = MagicMock()
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.return_value = message
|
||||
_persist_message(sqlite_session)
|
||||
|
||||
with pytest.raises(MoreLikeThisDisabledError):
|
||||
generator.generate_more_like_this(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
app_model=app_model,
|
||||
message_id="msg",
|
||||
user=_build_user(),
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
)
|
||||
|
||||
def test_generate_more_like_this_message_config_none(self, generator, mocker: MockerFixture):
|
||||
def test_generate_more_like_this_message_config_none(self, generator, sqlite_session: Session):
|
||||
app_model = _build_app_model()
|
||||
current_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True})
|
||||
|
||||
message = MagicMock(conversation_id="conv-1")
|
||||
conversation = MagicMock(app_model_config_id=None)
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.return_value = message
|
||||
session.get.side_effect = [current_config, conversation]
|
||||
current_config = AppModelConfig(app_id=app_model.id, more_like_this='{"enabled": true}')
|
||||
sqlite_session.add(current_config)
|
||||
sqlite_session.flush()
|
||||
app_model.app_model_config_id = current_config.id
|
||||
_persist_message(sqlite_session)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
generator.generate_more_like_this(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
app_model=app_model,
|
||||
message_id="msg",
|
||||
user=_build_user(),
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
)
|
||||
|
||||
def test_generate_more_like_this_success(self, generator, mocker: MockerFixture):
|
||||
def test_generate_more_like_this_success(self, generator, mocker: MockerFixture, sqlite_session: Session):
|
||||
app_model = _build_app_model()
|
||||
current_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True})
|
||||
|
||||
message = module.Message(id="msg", app_id="app1", conversation_id="conv-1", query="q")
|
||||
message.inputs = {"attachment": {"dify_model_identity": FILE_MODEL_IDENTITY}}
|
||||
message_files = [{"id": "f"}]
|
||||
message_files_with_session = mocker.patch.object(
|
||||
module.Message,
|
||||
"message_files_with_session",
|
||||
return_value=message_files,
|
||||
app_model_config = AppModelConfig(app_id=app_model.id, more_like_this='{"enabled": true}')
|
||||
sqlite_session.add(app_model_config)
|
||||
sqlite_session.flush()
|
||||
app_model.app_model_config_id = app_model_config.id
|
||||
_persist_message(sqlite_session, app_model_config=app_model_config)
|
||||
to_dict = mocker.patch.object(
|
||||
AppModelConfig,
|
||||
"to_dict",
|
||||
return_value={
|
||||
"model": {"completion_params": {"temperature": 0.1}},
|
||||
"file_upload": {"enabled": True},
|
||||
},
|
||||
)
|
||||
|
||||
app_model_config = MagicMock(app_id="app1")
|
||||
app_model_config.to_dict.return_value = {
|
||||
"model": {"completion_params": {"temperature": 0.1}},
|
||||
"file_upload": {"enabled": True},
|
||||
}
|
||||
annotation_reply = {"enabled": False}
|
||||
load_annotation_reply_config = mocker.patch.object(
|
||||
module,
|
||||
"load_annotation_reply_config",
|
||||
return_value=annotation_reply,
|
||||
)
|
||||
conversation = MagicMock(app_model_config_id="cfg-message")
|
||||
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.side_effect = [message, "tenant"]
|
||||
session.get.side_effect = [current_config, conversation, app_model_config]
|
||||
|
||||
global_session = MagicMock()
|
||||
global_session.scalar.side_effect = AssertionError("global session must not be used")
|
||||
global_session.scalars.side_effect = AssertionError("global session must not be used")
|
||||
mocker.patch.object(module.db, "session", global_session)
|
||||
|
||||
def restore_input_file(*, file_mapping, tenant_resolver):
|
||||
assert file_mapping["dify_model_identity"] == FILE_MODEL_IDENTITY
|
||||
assert tenant_resolver() == "tenant"
|
||||
return "input-file"
|
||||
|
||||
mocker.patch("models.model.build_file_from_input_mapping", side_effect=restore_input_file)
|
||||
|
||||
file_extra_config = MagicMock()
|
||||
mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=file_extra_config)
|
||||
build_from_mappings = mocker.patch.object(module.file_factory, "build_from_mappings", return_value=["file1"])
|
||||
mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=None)
|
||||
build_from_mappings = mocker.patch.object(module.file_factory, "build_from_mappings")
|
||||
|
||||
app_config = MagicMock(variables=["v"], to_dict=MagicMock(return_value={}))
|
||||
get_app_config = mocker.patch.object(
|
||||
@ -320,7 +356,7 @@ class TestCompletionAppGenerator:
|
||||
mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted")
|
||||
|
||||
result = generator.generate_more_like_this(
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
app_model=app_model,
|
||||
message_id="msg",
|
||||
user=_build_user(),
|
||||
@ -329,25 +365,12 @@ class TestCompletionAppGenerator:
|
||||
)
|
||||
|
||||
assert result == "converted"
|
||||
assert session.get.call_args_list == [
|
||||
call(module.AppModelConfig, "cfg-current"),
|
||||
call(module.Conversation, "conv-1"),
|
||||
call(module.AppModelConfig, "cfg-message"),
|
||||
]
|
||||
load_annotation_reply_config.assert_called_once_with(session, "app1")
|
||||
app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply)
|
||||
assert session.scalar.call_count == 2
|
||||
message_files_with_session.assert_called_once_with(session=session)
|
||||
build_from_mappings.assert_called_once_with(
|
||||
mappings=message_files,
|
||||
tenant_id="tenant",
|
||||
config=file_extra_config,
|
||||
access_controller=generator._file_access_controller,
|
||||
)
|
||||
assert global_session.mock_calls == []
|
||||
assert generator.generate_entity.call_args.kwargs["inputs"] == {"attachment": "input-file"}
|
||||
load_annotation_reply_config.assert_called_once_with(sqlite_session, app_model.id)
|
||||
to_dict.assert_called_once_with(annotation_reply=annotation_reply)
|
||||
assert generator.generate_entity.call_args.kwargs["inputs"] == {"a": 1}
|
||||
override_dict = get_app_config.call_args.kwargs["override_config_dict"]
|
||||
assert override_dict["model"]["completion_params"]["temperature"] == 0.9
|
||||
build_from_mappings.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("error", "should_publish"),
|
||||
@ -365,13 +388,19 @@ class TestCompletionAppGenerator:
|
||||
(RuntimeError("boom"), True),
|
||||
],
|
||||
)
|
||||
def test_generate_worker_error_handling(self, generator, mocker: MockerFixture, error, should_publish):
|
||||
def test_generate_worker_error_handling(
|
||||
self,
|
||||
generator,
|
||||
mocker: MockerFixture,
|
||||
sqlite_session: Session,
|
||||
error,
|
||||
should_publish,
|
||||
):
|
||||
flask_app = MagicMock()
|
||||
flask_app.app_context.return_value = contextlib.nullcontext()
|
||||
|
||||
session = mocker.MagicMock()
|
||||
session_context = mocker.MagicMock()
|
||||
session_context.__enter__.return_value = session
|
||||
session_context.__enter__.return_value = sqlite_session
|
||||
create_session = mocker.patch.object(module.session_factory, "create_session")
|
||||
create_session.return_value = session_context
|
||||
mocker.patch.object(module.db, "session")
|
||||
|
||||
Loading…
Reference in New Issue
Block a user