mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: replace manual model_validate with @model_validate in explore/snippets controllers (#40241)
This commit is contained in:
@@ -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"
|
||||
|
||||
+1
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user