test: use sqlite3 session in test_completion (#38783)

This commit is contained in:
Asuka Minato 2026-07-13 14:25:41 +09:00 committed by GitHub
parent ab521c4cfd
commit 8072b64292
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -12,13 +12,15 @@ Focus on:
"""
import uuid
from decimal import Decimal
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from flask import Flask
from pydantic import ValidationError
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, NotFound
import services
@ -39,7 +41,9 @@ from controllers.service_api.app.error import (
from core.app.apps.agent_app.errors import AgentAppNotPublishedError
from core.errors.error import QuotaExceededError
from graphon.model_runtime.errors.invoke import InvokeError
from models.model import App, AppMode, EndUser
from models.base import TypeBase
from models.enums import ConversationFromSource, EndUserType
from models.model import App, AppMode, Conversation, EndUser, IconType, Message
from services.app_generate_service import AppGenerateService
from services.app_task_service import AppTaskService
from services.conversation_service import ConversationService
@ -48,6 +52,65 @@ from services.errors.conversation import ConversationNotExistsError
from services.errors.llm import InvokeRateLimitError
@pytest.fixture
def orm_session(sqlite_engine: Engine):
models = (App, EndUser, Conversation, Message)
tables = [model.metadata.tables[model.__tablename__] for model in models]
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
with Session(sqlite_engine, expire_on_commit=False) as session:
yield session
def _persist_completion_state(session: Session, mode: AppMode) -> tuple[App, EndUser, Conversation, Message]:
tenant_id = str(uuid.uuid4())
app_model = App(
id=str(uuid.uuid4()),
tenant_id=tenant_id,
name="Completion App",
mode=mode,
icon_type=IconType.EMOJI,
icon="chat",
icon_background="#FFFFFF",
enable_site=False,
enable_api=True,
)
end_user = EndUser(
id=str(uuid.uuid4()),
tenant_id=tenant_id,
app_id=app_model.id,
type=EndUserType.SERVICE_API,
name="API User",
session_id=f"session-{uuid.uuid4()}",
)
conversation = Conversation(
id=str(uuid.uuid4()),
app_id=app_model.id,
mode=mode,
name="Conversation",
_inputs={},
from_source=ConversationFromSource.API,
from_end_user_id=end_user.id,
)
message = Message(
id=str(uuid.uuid4()),
app_id=app_model.id,
conversation_id=conversation.id,
_inputs={},
query="Hello",
message={},
message_unit_price=Decimal(0),
answer="Hi",
answer_unit_price=Decimal(0),
currency="USD",
from_source=ConversationFromSource.API,
from_end_user_id=end_user.id,
app_mode=mode,
)
session.add_all([app_model, end_user, conversation, message])
session.commit()
return app_model, end_user, conversation, message
class TestCompletionRequestPayload:
"""Test suite for CompletionRequestPayload Pydantic model."""
@ -247,64 +310,68 @@ class TestAppGenerateService:
assert callable(AppGenerateService.generate)
@patch.object(AppGenerateService, "generate")
def test_generate_returns_response(self, mock_generate):
def test_generate_returns_response(self, mock_generate, orm_session: Session):
"""Test that generate returns expected response format."""
expected = {"answer": "Hello!"}
mock_generate.return_value = expected
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
result = AppGenerateService.generate(
app_model=Mock(spec=App),
user=Mock(spec=EndUser),
app_model=app_model,
user=end_user,
args={"query": "Hi"},
invoke_from=Mock(),
session=Mock(),
session=orm_session,
streaming=False,
)
assert result == expected
@patch.object(AppGenerateService, "generate")
def test_generate_raises_conversation_not_exists(self, mock_generate):
def test_generate_raises_conversation_not_exists(self, mock_generate, orm_session: Session):
"""Test generate raises ConversationNotExistsError."""
mock_generate.side_effect = services.errors.conversation.ConversationNotExistsError()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with pytest.raises(services.errors.conversation.ConversationNotExistsError):
AppGenerateService.generate(
app_model=Mock(spec=App),
user=Mock(spec=EndUser),
app_model=app_model,
user=end_user,
args={},
invoke_from=Mock(),
session=Mock(),
session=orm_session,
streaming=False,
)
@patch.object(AppGenerateService, "generate")
def test_generate_raises_quota_exceeded(self, mock_generate):
def test_generate_raises_quota_exceeded(self, mock_generate, orm_session: Session):
"""Test generate raises QuotaExceededError."""
mock_generate.side_effect = QuotaExceededError()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with pytest.raises(QuotaExceededError):
AppGenerateService.generate(
app_model=Mock(spec=App),
user=Mock(spec=EndUser),
app_model=app_model,
user=end_user,
args={},
invoke_from=Mock(),
session=Mock(),
session=orm_session,
streaming=False,
)
@patch.object(AppGenerateService, "generate")
def test_generate_raises_invoke_error(self, mock_generate):
def test_generate_raises_invoke_error(self, mock_generate, orm_session: Session):
"""Test generate raises InvokeError."""
mock_generate.side_effect = InvokeError("Model invocation failed")
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with pytest.raises(InvokeError):
AppGenerateService.generate(
app_model=Mock(spec=App),
user=Mock(spec=EndUser),
app_model=app_model,
user=end_user,
args={},
invoke_from=Mock(),
session=Mock(),
session=orm_session,
streaming=False,
)
@ -323,14 +390,14 @@ class TestCompletionControllerLogic:
@patch("controllers.service_api.app.completion.service_api_ns")
@patch("controllers.service_api.app.completion.AppGenerateService")
def test_completion_api_post_success(self, mock_generate_service, mock_service_api_ns, app: Flask):
def test_completion_api_post_success(
self, mock_generate_service, mock_service_api_ns, app: Flask, orm_session: Session
):
"""Test CompletionApi.post success path."""
from controllers.service_api.app.completion import CompletionApi
# Setup mocks
mock_app_model = Mock(spec=App)
mock_app_model.mode = AppMode.COMPLETION
mock_end_user = Mock(spec=EndUser)
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
payload_dict = {"inputs": {"text": "hello"}, "response_mode": "blocking"}
mock_service_api_ns.payload = payload_dict
@ -342,33 +409,29 @@ class TestCompletionControllerLogic:
mock_compact.return_value = {"text": "compacted"}
api = CompletionApi()
response = unwrap(api.post)(api, Mock(), mock_app_model, mock_end_user)
response = unwrap(api.post)(api, orm_session, app_model, end_user)
assert response == {"text": "compacted"}
mock_generate_service.generate.assert_called_once()
@patch("controllers.service_api.app.completion.service_api_ns")
def test_completion_api_post_wrong_app_mode(self, mock_service_api_ns, app: Flask):
def test_completion_api_post_wrong_app_mode(self, mock_service_api_ns, app: Flask, orm_session: Session):
"""Test CompletionApi.post with wrong app mode."""
from controllers.service_api.app.completion import CompletionApi
mock_app_model = Mock(spec=App)
mock_app_model.mode = AppMode.CHAT # Wrong mode
mock_end_user = Mock(spec=EndUser)
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context():
with pytest.raises(AppUnavailableError):
unwrap(CompletionApi().post)(CompletionApi(), Mock(), mock_app_model, mock_end_user)
unwrap(CompletionApi().post)(CompletionApi(), orm_session, app_model, end_user)
@patch("controllers.service_api.app.completion.service_api_ns")
@patch("controllers.service_api.app.completion.AppGenerateService")
def test_chat_api_post_success(self, mock_generate_service, mock_service_api_ns, app: Flask):
def test_chat_api_post_success(self, mock_generate_service, mock_service_api_ns, app: Flask, orm_session: Session):
"""Test ChatApi.post success path."""
from controllers.service_api.app.completion import ChatApi
mock_app_model = Mock(spec=App)
mock_app_model.mode = AppMode.CHAT
mock_end_user = Mock(spec=EndUser)
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
payload_dict = {"inputs": {}, "query": "hello", "response_mode": "blocking"}
mock_service_api_ns.payload = payload_dict
@ -379,52 +442,44 @@ class TestCompletionControllerLogic:
mock_compact.return_value = {"text": "compacted"}
api = ChatApi()
response = unwrap(api.post)(api, Mock(), mock_app_model, mock_end_user)
response = unwrap(api.post)(api, orm_session, app_model, end_user)
assert response == {"text": "compacted"}
@patch("controllers.service_api.app.completion.service_api_ns")
def test_chat_api_post_wrong_app_mode(self, mock_service_api_ns, app: Flask):
def test_chat_api_post_wrong_app_mode(self, mock_service_api_ns, app: Flask, orm_session: Session):
"""Test ChatApi.post with wrong app mode."""
from controllers.service_api.app.completion import ChatApi
mock_app_model = Mock(spec=App)
mock_app_model.mode = AppMode.COMPLETION # Wrong mode
mock_end_user = Mock(spec=EndUser)
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with app.test_request_context():
with pytest.raises(NotChatAppError):
unwrap(ChatApi().post)(ChatApi(), Mock(), mock_app_model, mock_end_user)
unwrap(ChatApi().post)(ChatApi(), orm_session, app_model, end_user)
@patch("controllers.service_api.app.completion.AppTaskService")
def test_completion_stop_api_success(self, mock_task_service, app: Flask):
def test_completion_stop_api_success(self, mock_task_service, app: Flask, orm_session: Session):
"""Test CompletionStopApi.post success."""
from controllers.service_api.app.completion import CompletionStopApi
mock_app_model = Mock(spec=App)
mock_app_model.mode = AppMode.COMPLETION
mock_end_user = Mock(spec=EndUser)
mock_end_user.id = "user_id"
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with app.test_request_context():
api = CompletionStopApi()
response = api.post.__wrapped__(api, mock_app_model, mock_end_user, "task_id")
response = api.post.__wrapped__(api, app_model, end_user, "task_id")
assert response == ({"result": "success"}, 200)
mock_task_service.stop_task.assert_called_once()
@patch("controllers.service_api.app.completion.AppTaskService")
def test_chat_stop_api_success(self, mock_task_service, app: Flask):
def test_chat_stop_api_success(self, mock_task_service, app: Flask, orm_session: Session):
"""Test ChatStopApi.post success."""
from controllers.service_api.app.completion import ChatStopApi
mock_app_model = Mock(spec=App)
mock_app_model.mode = AppMode.CHAT
mock_end_user = Mock(spec=EndUser)
mock_end_user.id = "user_id"
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context():
api = ChatStopApi()
response = api.post.__wrapped__(api, mock_app_model, mock_end_user, "task_id")
response = api.post.__wrapped__(api, app_model, end_user, "task_id")
assert response == ({"result": "success"}, 200)
mock_task_service.stop_task.assert_called_once()
@ -442,52 +497,48 @@ class TestChatRequestPayloadController:
class TestCompletionApiController:
def test_wrong_mode(self, app: Flask) -> None:
def test_wrong_mode(self, app: Flask, orm_session: Session) -> None:
api = CompletionApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context("/completion-messages", method="POST", json={"inputs": {}}):
with pytest.raises(AppUnavailableError):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
def test_conversation_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_conversation_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session) -> None:
monkeypatch.setattr(
AppGenerateService,
"generate",
lambda *_args, **_kwargs: (_ for _ in ()).throw(ConversationNotExistsError()),
)
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
api = CompletionApi()
handler = unwrap(api.post)
with app.test_request_context("/completion-messages", method="POST", json={"inputs": {}}):
with pytest.raises(NotFound):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
class TestCompletionStopApiController:
def test_wrong_mode(self, app: Flask) -> None:
def test_wrong_mode(self, app: Flask, orm_session: Session) -> None:
api = CompletionStopApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
end_user = SimpleNamespace(id="u1")
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context("/completion-messages/1/stop", method="POST"):
with pytest.raises(AppUnavailableError):
handler(api, app_model=app_model, end_user=end_user, task_id="t1")
def test_success(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_success(self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session) -> None:
stop_mock = Mock()
monkeypatch.setattr(AppTaskService, "stop_task", stop_mock)
api = CompletionStopApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
end_user = SimpleNamespace(id="u1")
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with app.test_request_context("/completion-messages/1/stop", method="POST"):
response, status = handler(api, app_model=app_model, end_user=end_user, task_id="t1")
@ -497,17 +548,16 @@ class TestCompletionStopApiController:
class TestChatApiController:
def test_wrong_mode(self, app: Flask) -> None:
def test_wrong_mode(self, app: Flask, orm_session: Session) -> None:
api = ChatApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.COMPLETION.value)
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}):
with pytest.raises(NotChatAppError):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
def test_workflow_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_workflow_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session) -> None:
monkeypatch.setattr(
AppGenerateService,
"generate",
@ -516,14 +566,13 @@ class TestChatApiController:
api = ChatApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}):
with pytest.raises(NotFound):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session) -> None:
monkeypatch.setattr(
AppGenerateService,
"generate",
@ -532,14 +581,15 @@ class TestChatApiController:
api = ChatApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.CHAT.value)
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}):
with pytest.raises(BadRequest):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
def test_agent_not_published_error_mapped(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_agent_not_published_error_mapped(
self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session
) -> None:
monkeypatch.setattr(
AppGenerateService,
"generate",
@ -548,14 +598,15 @@ class TestChatApiController:
api = ChatApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.AGENT.value)
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.AGENT)
with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}):
with pytest.raises(AgentNotPublishedError):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
def test_invalid_conversation_id_fails_fast_as_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_invalid_conversation_id_fails_fast_as_not_found(
self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session
) -> None:
# A well-formed but nonexistent conversation_id must fail fast as 404, before the
# streaming generator is created. Previously the lookup only ran inside the generator,
# so an invalid id surfaced as a hang instead of a clean error.
@ -570,8 +621,7 @@ class TestChatApiController:
api = ChatApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.CHAT.value, id="app-1")
end_user = SimpleNamespace()
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.CHAT)
with app.test_request_context(
"/chat-messages",
@ -579,18 +629,17 @@ class TestChatApiController:
json={"inputs": {}, "query": "hi", "conversation_id": str(uuid.uuid4())},
):
with pytest.raises(NotFound):
handler(api, session=Mock(), app_model=app_model, end_user=end_user)
handler(api, session=orm_session, app_model=app_model, end_user=end_user)
# The lookup must run before generation, so the generator is never started.
generate_mock.assert_not_called()
class TestChatStopApiController:
def test_wrong_mode(self, app: Flask) -> None:
def test_wrong_mode(self, app: Flask, orm_session: Session) -> None:
api = ChatStopApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
end_user = SimpleNamespace(id="u1")
app_model, end_user, _, _ = _persist_completion_state(orm_session, AppMode.COMPLETION)
with app.test_request_context("/chat-messages/1/stop", method="POST"):
with pytest.raises(NotChatAppError):