mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 10:38:32 +08:00
test: use sqlite3 session in test_completion (#38783)
This commit is contained in:
parent
ab521c4cfd
commit
8072b64292
@ -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):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user