test: use sqlite3 session in test_completion_completion_app_generator (#38723)

This commit is contained in:
Asuka Minato 2026-07-30 11:45:54 +09:00 committed by GitHub
parent b855d71adb
commit e12e86a308
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 128 additions and 98 deletions

View File

@ -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

View File

@ -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")