refactor: replace manual model_validate with @model_validate in explore/snippets controllers (#40241)

This commit is contained in:
Likalikali
2026-08-09 07:07:08 +00:00
committed by GitHub
parent a1577fc48c
commit 925ca49f21
9 changed files with 437 additions and 129 deletions
@@ -6,6 +6,7 @@ from flask import Flask
from pydantic import ValidationError
import controllers.console.explore.recommended_app as module
from controllers.console.explore.recommended_app import RecommendedAppsQuery
from models import Account
from models.model import AppMode, IconType
@@ -32,7 +33,7 @@ class TestRecommendedAppListApi:
return_value=result_data,
) as service_mock,
):
result = method(api, make_account("fr-FR"))
result = method(api, RecommendedAppsQuery(language="en-US"), make_account("fr-FR"))
service_mock.assert_called_once_with("en-US", session=ANY)
assert result == result_data
@@ -51,7 +52,7 @@ class TestRecommendedAppListApi:
return_value=result_data,
) as service_mock,
):
result = method(api, make_account("fr-FR"))
result = method(api, RecommendedAppsQuery(), make_account("fr-FR"))
service_mock.assert_called_once_with("fr-FR", session=ANY)
assert result == result_data
@@ -70,7 +71,7 @@ class TestRecommendedAppListApi:
return_value=result_data,
) as service_mock,
):
result = method(api, make_account(None))
result = method(api, RecommendedAppsQuery(), make_account(None))
service_mock.assert_called_once_with(module.languages[0], session=ANY)
assert result == result_data
@@ -91,7 +92,7 @@ class TestLearnDifyAppListApi:
return_value=result_data,
) as service_mock,
):
result = method(api, make_account("fr-FR"))
result = method(api, RecommendedAppsQuery(language="en-US"), make_account("fr-FR"))
service_mock.assert_called_once_with("en-US", session=ANY)
assert result == result_data
@@ -110,7 +111,7 @@ class TestLearnDifyAppListApi:
return_value=result_data,
) as service_mock,
):
result = method(api, make_account("fr-FR"))
result = method(api, RecommendedAppsQuery(), make_account("fr-FR"))
service_mock.assert_called_once_with("fr-FR", session=ANY)
assert result == result_data
@@ -8,7 +8,7 @@ from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from flask import Flask, request
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, InternalServerError, NotFound
@@ -28,6 +28,7 @@ from controllers.console.explore.error import (
NotCompletionAppError,
NotWorkflowAppError,
)
from controllers.console.explore.trial import ChatRequest, CompletionRequest, TextToSpeechRequest, WorkflowRunRequest
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
from core.errors.error import (
ModelCurrentlyNotSupportError,
@@ -258,9 +259,15 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
api = module.TrialAppWorkflowRunApi()
method = unwrap(api.post)
with app.test_request_context("/"):
with app.test_request_context("/", json={"inputs": {}}):
with pytest.raises(NotWorkflowAppError):
method(api, self.sqlite_session, account, MagicMock(mode=AppMode.CHAT))
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
MagicMock(mode=AppMode.CHAT),
)
def test_success(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -271,7 +278,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
patch.object(module.AppGenerateService, "generate", return_value=MagicMock()),
patch.object(module.RecommendedAppService, "add_trial_app_record"),
):
result = method(api, self.sqlite_session, account, trial_app_workflow)
result = method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
assert result is not None
@@ -288,7 +301,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderNotInitializeError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
def test_workflow_quota_exceeded(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -303,7 +322,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderQuotaExceededError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
def test_workflow_model_not_support(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -318,7 +343,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderModelCurrentlyNotSupportError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
def test_workflow_invoke_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -333,7 +364,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(CompletionRequestError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
def test_workflow_rate_limit_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -348,7 +385,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(InvokeRateLimitHttpError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
def test_workflow_value_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -363,7 +406,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(ValueError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
def test_workflow_generic_exception(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None:
api = module.TrialAppWorkflowRunApi()
@@ -378,7 +427,13 @@ class TestTrialAppWorkflowRunApi(_UsesSQLiteSession):
),
):
with pytest.raises(InternalServerError):
method(api, self.sqlite_session, account, trial_app_workflow)
method(
api,
WorkflowRunRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_workflow,
)
class TestTrialChatApi(_UsesSQLiteSession):
@@ -388,7 +443,13 @@ class TestTrialChatApi(_UsesSQLiteSession):
with app.test_request_context("/", json={"inputs": {}, "query": "hi"}):
with pytest.raises(NotChatAppError):
method(api, self.sqlite_session, account, MagicMock(mode="completion"))
method(
api,
ChatRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
MagicMock(mode="completion"),
)
def test_success(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -399,7 +460,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
patch.object(module.AppGenerateService, "generate", return_value=MagicMock()),
patch.object(module.RecommendedAppService, "add_trial_app_record"),
):
result = method(api, self.sqlite_session, account, trial_app_chat)
result = method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
assert result is not None
@@ -416,7 +479,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(NotFound):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_conversation_completed(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -431,7 +496,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(ConversationCompletedError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -446,7 +513,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(AppUnavailableError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -461,7 +530,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderNotInitializeError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -476,7 +547,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderQuotaExceededError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -491,7 +564,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderModelCurrentlyNotSupportError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -506,7 +581,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(CompletionRequestError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_rate_limit_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -521,7 +598,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(InvokeRateLimitHttpError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_value_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -536,7 +615,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(ValueError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
def test_chat_generic_exception(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatApi()
@@ -551,7 +632,9 @@ class TestTrialChatApi(_UsesSQLiteSession):
),
):
with pytest.raises(InternalServerError):
method(api, self.sqlite_session, account, trial_app_chat)
method(
api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat
)
class TestTrialCompletionApi(_UsesSQLiteSession):
@@ -561,7 +644,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
with app.test_request_context("/", json={"inputs": {}, "query": ""}):
with pytest.raises(NotCompletionAppError):
method(api, self.sqlite_session, account, MagicMock(mode=AppMode.CHAT))
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
MagicMock(mode=AppMode.CHAT),
)
def test_success(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -572,7 +661,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
patch.object(module.AppGenerateService, "generate", return_value=MagicMock()),
patch.object(module.RecommendedAppService, "add_trial_app_record"),
):
result = method(api, self.sqlite_session, account, trial_app_completion)
result = method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
assert result is not None
@@ -589,7 +684,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(AppUnavailableError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_provider_not_init(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -604,7 +705,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderNotInitializeError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_quota_exceeded(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -619,7 +726,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderQuotaExceededError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_model_not_support(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -634,7 +747,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(ProviderModelCurrentlyNotSupportError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_invoke_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -649,7 +768,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(CompletionRequestError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_rate_limit_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -664,7 +789,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(InternalServerError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_value_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -679,7 +810,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(ValueError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
def test_completion_generic_exception(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None:
api = module.TrialCompletionApi()
@@ -694,7 +831,13 @@ class TestTrialCompletionApi(_UsesSQLiteSession):
),
):
with pytest.raises(InternalServerError):
method(api, self.sqlite_session, account, trial_app_completion)
method(
api,
CompletionRequest.model_validate(request.get_json()),
self.sqlite_session,
account,
trial_app_completion,
)
class TestTrialMessageSuggestedQuestionApi:
@@ -843,7 +986,11 @@ class TestTrialChatAudioApi:
),
):
with pytest.raises(module.AppUnavailableError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_no_audio_uploaded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -862,7 +1009,11 @@ class TestTrialChatAudioApi:
),
):
with pytest.raises(module.NoAudioUploadedError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_missing_file_field_returns_400(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
"""A multipart POST with no `file` field must surface as 400, not 500.
@@ -883,7 +1034,11 @@ class TestTrialChatAudioApi:
patch.object(module.AudioService, "transcript_asr", side_effect=fake_asr),
):
with pytest.raises(module.NoAudioUploadedError) as exc_info:
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
assert exc_info.value.code == 400
@@ -904,7 +1059,11 @@ class TestTrialChatAudioApi:
),
):
with pytest.raises(module.AudioTooLargeError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_unsupported_audio_type(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -923,7 +1082,11 @@ class TestTrialChatAudioApi:
),
):
with pytest.raises(module.UnsupportedAudioTypeError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_provider_not_support_tts(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -942,7 +1105,11 @@ class TestTrialChatAudioApi:
),
):
with pytest.raises(module.ProviderNotSupportSpeechToTextError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_speech_to_text_disabled(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -960,7 +1127,11 @@ class TestTrialChatAudioApi:
),
):
with pytest.raises(SpeechToTextDisabledError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -975,7 +1146,11 @@ class TestTrialChatAudioApi:
patch.object(module.AudioService, "transcript_asr", side_effect=ProviderTokenNotInitError("test")),
):
with pytest.raises(ProviderNotInitializeError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -990,7 +1165,11 @@ class TestTrialChatAudioApi:
patch.object(module.AudioService, "transcript_asr", side_effect=QuotaExceededError()),
):
with pytest.raises(ProviderQuotaExceededError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
class TestTrialChatTextApi:
@@ -1003,7 +1182,9 @@ class TestTrialChatTextApi:
patch.object(module.AudioService, "transcript_tts", return_value={"audio": "base64_data"}),
patch.object(module.RecommendedAppService, "add_trial_app_record"),
):
result = method(api, account, trial_app_chat)
result = method(
api, TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), account, trial_app_chat
)
assert result == {"audio": "base64_data"}
@@ -1018,7 +1199,9 @@ class TestTrialChatTextApi:
patch.object(module.AudioService, "transcript_tts", transcript_tts),
patch.object(module.RecommendedAppService, "add_trial_app_record"),
):
result = method(api, account, trial_app_chat)
result = method(
api, TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), account, trial_app_chat
)
assert result == {"audio": "base64_data"}
assert transcript_tts.call_args.kwargs["message_ref"] == MessageRef(
@@ -1040,7 +1223,12 @@ class TestTrialChatTextApi:
),
):
with pytest.raises(module.AppUnavailableError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_provider_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1055,7 +1243,12 @@ class TestTrialChatTextApi:
),
):
with pytest.raises(module.ProviderNotSupportSpeechToTextError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_audio_too_large(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1070,7 +1263,12 @@ class TestTrialChatTextApi:
),
):
with pytest.raises(module.AudioTooLargeError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_no_audio_uploaded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1085,7 +1283,12 @@ class TestTrialChatTextApi:
),
):
with pytest.raises(module.NoAudioUploadedError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1096,7 +1299,12 @@ class TestTrialChatTextApi:
patch.object(module.AudioService, "transcript_tts", side_effect=ProviderTokenNotInitError("test")),
):
with pytest.raises(ProviderNotInitializeError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1107,7 +1315,12 @@ class TestTrialChatTextApi:
patch.object(module.AudioService, "transcript_tts", side_effect=QuotaExceededError()),
):
with pytest.raises(ProviderQuotaExceededError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1118,7 +1331,12 @@ class TestTrialChatTextApi:
patch.object(module.AudioService, "transcript_tts", side_effect=ModelCurrentlyNotSupportError()),
):
with pytest.raises(ProviderModelCurrentlyNotSupportError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1129,14 +1347,19 @@ class TestTrialChatTextApi:
patch.object(module.AudioService, "transcript_tts", side_effect=InvokeError("test error")),
):
with pytest.raises(CompletionRequestError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
class TestTrialAppWorkflowTaskStopApi:
def test_not_workflow_app(self, app: Flask, trial_app_chat: MagicMock) -> None:
api = module.TrialAppWorkflowTaskStopApi()
with app.test_request_context("/"):
with app.test_request_context("/", json={"inputs": {}}):
with pytest.raises(NotWorkflowAppError):
api.post(trial_app_chat, str(uuid4()))
@@ -1336,7 +1559,11 @@ class TestTrialChatAudioApiExceptionHandlers:
),
):
with pytest.raises(ProviderNotInitializeError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -1355,7 +1582,11 @@ class TestTrialChatAudioApiExceptionHandlers:
),
):
with pytest.raises(ProviderQuotaExceededError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
def test_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatAudioApi()
@@ -1374,7 +1605,11 @@ class TestTrialChatAudioApiExceptionHandlers:
),
):
with pytest.raises(CompletionRequestError):
method(api, account, trial_app_chat)
method(
api,
account,
trial_app_chat,
)
class TestTrialChatTextApiExceptionHandlers:
@@ -1391,7 +1626,12 @@ class TestTrialChatTextApiExceptionHandlers:
),
):
with pytest.raises(module.AppUnavailableError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
def test_unsupported_audio_type(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None:
api = module.TrialChatTextApi()
@@ -1406,4 +1646,9 @@ class TestTrialChatTextApiExceptionHandlers:
),
):
with pytest.raises(module.UnsupportedAudioTypeError):
method(api, account, trial_app_chat)
method(
api,
TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}),
account,
trial_app_chat,
)
@@ -132,7 +132,14 @@ def test_draft_workflow_post_returns_400_for_invalid_graph(app: Flask, monkeypat
method="POST",
json={"graph": {"nodes": [], "edges": []}, "hash": "hash-1"},
):
response, status_code = handler(api, user, snippet)
response, status_code = handler(
api,
snippet_workflow_module.SnippetDraftSyncPayload.model_validate(
{"graph": {"nodes": [], "edges": []}, "hash": "hash-1"}
),
user,
snippet,
)
assert status_code == 400
assert response == {"message": "invalid graph"}
@@ -244,7 +251,11 @@ def test_list_published_snippet_workflows_includes_input_fields(
handler = unwrap(api.get)
with app.test_request_context("/snippets/snippet-1/workflows?page=1&limit=20"):
response = handler(api, snippet=snippet)
response = handler(
api,
snippet_workflow_module.SnippetWorkflowListQuery.model_validate({"page": 1, "limit": 20}),
snippet=snippet,
)
assert response["items"][0]["input_fields"] == input_fields
@@ -411,7 +422,15 @@ def test_update_published_snippet_workflow_returns_updated_workflow(
method="PATCH",
json={"marked_name": "v1", "marked_comment": "first version"},
):
response = handler(api, user, snippet, workflow_id="workflow-1")
response = handler(
api,
snippet_workflow_module.WorkflowUpdatePayload.model_validate(
{"marked_name": "v1", "marked_comment": "first version"}
),
user,
snippet,
workflow_id="workflow-1",
)
update_workflow.assert_called_once()
update_call = update_workflow.call_args.kwargs
@@ -432,7 +451,13 @@ def test_update_published_snippet_workflow_returns_400_when_no_fields(app: Flask
handler = unwrap(api.patch)
with app.test_request_context("/snippets/snippet-1/workflows/workflow-1", method="PATCH", json={}):
response, status_code = handler(api, _account("account-1"), _snippet(), workflow_id="workflow-1")
response, status_code = handler(
api,
snippet_workflow_module.WorkflowUpdatePayload(),
_account("account-1"),
_snippet(),
workflow_id="workflow-1",
)
assert status_code == 400
assert response == {"message": "No valid fields to update"}
@@ -468,7 +493,13 @@ def test_update_published_snippet_workflow_raises_not_found(
json={"marked_name": "v1"},
):
with pytest.raises(NotFound, match="Workflow not found"):
handler(api, user, snippet, workflow_id="missing-workflow")
handler(
api,
snippet_workflow_module.WorkflowUpdatePayload.model_validate({"marked_name": "v1"}),
user,
snippet,
workflow_id="missing-workflow",
)
sqlite_session.refresh(snippet)
assert snippet.name == "Snippet"
@@ -255,6 +255,7 @@ def test_variable_patch_returns_persisted_variable_without_committing_when_no_ch
with app.test_request_context("/", method="PATCH", json={}):
result = handler(
api,
module.WorkflowDraftVariableUpdatePayload(),
_make_account(),
snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"),
variable_id="var-1",