From 6ee3521430ebbee91476bd5dc9fd77f5fd1a753b Mon Sep 17 00:00:00 2001 From: likalikali <90487424+likalikali@users.noreply.github.com> Date: Sun, 9 Aug 2026 13:16:54 +0800 Subject: [PATCH] refactor: replace manual model_validate with @model_validate in explore/snippets controllers --- api/controllers/console/explore/banner.py | 11 +- .../console/explore/recommended_app.py | 15 +- api/controllers/console/explore/trial.py | 31 +- .../console/snippets/snippet_workflow.py | 80 ++-- .../snippet_workflow_draft_variable.py | 15 +- .../console/explore/test_recommended_app.py | 11 +- .../controllers/console/explore/test_trial.py | 361 +++++++++++++++--- .../console/snippets/test_snippet_workflow.py | 41 +- .../test_snippet_workflow_draft_variable.py | 1 + 9 files changed, 437 insertions(+), 129 deletions(-) diff --git a/api/controllers/console/explore/banner.py b/api/controllers/console/explore/banner.py index 52eadec7f34cd9..b63e025b60901d 100644 --- a/api/controllers/console/explore/banner.py +++ b/api/controllers/console/explore/banner.py @@ -1,13 +1,13 @@ from datetime import datetime from typing import cast -from flask import request from flask_restx import Namespace, Resource from pydantic import BaseModel, Field, RootModel, field_validator from sqlalchemy import select from controllers.common.schema import query_params_from_model, register_response_schema_models from controllers.console import api +from controllers.console.wraps import model_validate from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import dump_response @@ -64,23 +64,22 @@ class BannerApi(Resource): @api.doc(params=query_params_from_model(BannerListQuery)) @api.response(200, "Success", api.models[BannerListResponse.__name__]) - def get(self): + @model_validate(BannerListQuery) + def get(self, req_data: BannerListQuery): """Get banner list.""" if not FeatureService.is_explore_banner_enabled(): return dump_response(BannerListResponse, []) - query = BannerListQuery.model_validate(request.args.to_dict(flat=True)) - # Build base query for enabled banners base_query = select(ExporleBanner).where(ExporleBanner.status == BannerStatus.ENABLED) # Try to get banners in the requested language banners = db.session.scalars( - base_query.where(ExporleBanner.language == query.language).order_by(ExporleBanner.sort) + base_query.where(ExporleBanner.language == req_data.language).order_by(ExporleBanner.sort) ).all() # Fallback to en-US if no banners found and language is not en-US - if not banners and query.language != "en-US": + if not banners and req_data.language != "en-US": banners = db.session.scalars( base_query.where(ExporleBanner.language == "en-US").order_by(ExporleBanner.sort) ).all() diff --git a/api/controllers/console/explore/recommended_app.py b/api/controllers/console/explore/recommended_app.py index af46cdc8c9b4a0..aa91cd2d6c9c99 100644 --- a/api/controllers/console/explore/recommended_app.py +++ b/api/controllers/console/explore/recommended_app.py @@ -1,14 +1,13 @@ from typing import Any from uuid import UUID -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, RootModel, computed_field, field_validator from constants.languages import languages from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models from controllers.console import console_ns -from controllers.console.wraps import account_initialization_required, with_current_user +from controllers.console.wraps import account_initialization_required, model_validate, with_current_user from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import build_icon_url, dump_response @@ -114,10 +113,10 @@ class RecommendedAppListApi(Resource): @login_required @account_initialization_required @with_current_user - def get(self, current_user: Account): + @model_validate(RecommendedAppsQuery) + def get(self, req_data: RecommendedAppsQuery, current_user: Account): # language args - args = RecommendedAppsQuery.model_validate(request.args.to_dict(flat=True)) - language_prefix = _resolve_language(args.language, current_user) + language_prefix = _resolve_language(req_data.language, current_user) return dump_response( RecommendedAppListResponse, @@ -132,9 +131,9 @@ class LearnDifyAppListApi(Resource): @login_required @account_initialization_required @with_current_user - def get(self, current_user: Account): - args = RecommendedAppsQuery.model_validate(request.args.to_dict(flat=True)) - language_prefix = _resolve_language(args.language, current_user) + @model_validate(RecommendedAppsQuery) + def get(self, req_data: RecommendedAppsQuery, current_user: Account): + language_prefix = _resolve_language(req_data.language, current_user) return dump_response( LearnDifyAppListResponse, diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index 7663b4068ef8e9..620f3d3f85e03d 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -49,7 +49,7 @@ from controllers.console.explore.wraps import TrialAppResource from controllers.console.files import FILE_UPLOAD_PARAMS, upload_file_from_request from controllers.console.remote_files import RemoteFileUploadPayload, upload_remote_file_from_request -from controllers.console.wraps import cloud_edition_billing_resource_check, with_current_user +from controllers.console.wraps import cloud_edition_billing_resource_check, model_validate, with_current_user from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.base_app_queue_manager import AppQueueManager @@ -487,7 +487,8 @@ class TrialAppWorkflowRunApi(TrialAppResource): @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, trial_app): + @model_validate(WorkflowRunRequest) + def post(self, req_data: WorkflowRunRequest, session: Session, current_user: Account, trial_app): """ Run workflow """ @@ -498,8 +499,7 @@ def post(self, session: Session, current_user: Account, trial_app): if app_mode != AppMode.WORKFLOW: raise NotWorkflowAppError() - request_data = WorkflowRunRequest.model_validate(console_ns.payload) - args = request_data.model_dump() + args = req_data.model_dump() try: app_id = app_model.id user_id = current_user.id @@ -559,14 +559,14 @@ class TrialChatApi(TrialAppResource): @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, trial_app): + @model_validate(ChatRequest) + def post(self, req_data: ChatRequest, session: Session, current_user: Account, trial_app): app_model = trial_app app_mode = AppMode.value_of(app_model.mode) if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}: raise NotChatAppError() - request_data = ChatRequest.model_validate(console_ns.payload) - args = request_data.model_dump() + args = req_data.model_dump() # Validate UUID values if provided if args.get("conversation_id"): @@ -709,14 +709,13 @@ class TrialChatTextApi(TrialAppResource): @console_ns.expect(console_ns.models[TextToSpeechRequest.__name__]) @console_ns.response(200, "Success", console_ns.models[AudioBinaryResponse.__name__]) @with_current_user - def post(self, current_user: Account, trial_app): + @model_validate(TextToSpeechRequest) + def post(self, req_data: TextToSpeechRequest, current_user: Account, trial_app): app_model = trial_app try: - request_data = TextToSpeechRequest.model_validate(console_ns.payload) - - message_id = request_data.message_id - text = request_data.text - voice = request_data.voice + message_id = req_data.message_id + text = req_data.text + voice = req_data.voice message_ref = None if message_id: app_ref = AppRefService.create_app_ref(app_model) @@ -770,13 +769,13 @@ class TrialCompletionApi(TrialAppResource): @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, trial_app): + @model_validate(CompletionRequest) + def post(self, req_data: CompletionRequest, session: Session, current_user: Account, trial_app): app_model = trial_app if app_model.mode != "completion": raise NotCompletionAppError() - request_data = CompletionRequest.model_validate(console_ns.payload) - args = request_data.model_dump() + args = req_data.model_dump() streaming = args["response_mode"] == "streaming" args["auto_generate_name"] = False diff --git a/api/controllers/console/snippets/snippet_workflow.py b/api/controllers/console/snippets/snippet_workflow.py index 295170d2c0445e..d1399801a06eef 100644 --- a/api/controllers/console/snippets/snippet_workflow.py +++ b/api/controllers/console/snippets/snippet_workflow.py @@ -36,6 +36,7 @@ RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -204,18 +205,18 @@ def get(self, snippet: CustomizedSnippet): @rbac_permission_required( RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False ) - def post(self, current_user: Account, snippet: CustomizedSnippet): + @model_validate(SnippetDraftSyncPayload) + def post(self, req_data: SnippetDraftSyncPayload, current_user: Account, snippet: CustomizedSnippet): """Sync draft workflow for snippet.""" - payload = SnippetDraftSyncPayload.model_validate(console_ns.payload or {}) try: snippet_service = _snippet_service() workflow = snippet_service.sync_draft_workflow( snippet=snippet, - graph=payload.graph, - unique_hash=payload.hash, + graph=req_data.graph, + unique_hash=req_data.hash, account=current_user, - input_fields=payload.input_fields, + input_fields=req_data.input_fields, ) except WorkflowHashNotEqualError: raise DraftWorkflowNotSync() @@ -363,24 +364,24 @@ class SnippetPublishedAllWorkflowApi(Resource): @rbac_permission_required( RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False ) - def get(self, snippet: CustomizedSnippet): + @model_validate(SnippetWorkflowListQuery) + def get(self, req_data: SnippetWorkflowListQuery, snippet: CustomizedSnippet): """Get all published workflow versions for snippet.""" - args = SnippetWorkflowListQuery.model_validate(request.args.to_dict(flat=True)) snippet_service = _snippet_service() with Session(db.engine) as session: workflows, has_more = snippet_service.get_all_published_workflows( session=session, snippet=snippet, - page=args.page, - limit=args.limit, + page=req_data.page, + limit=req_data.limit, ) response = SnippetWorkflowPaginationResponse.model_validate( { "items": workflows, - "page": args.page, - "limit": args.limit, + "page": req_data.page, + "limit": req_data.limit, "has_more": has_more, }, from_attributes=True, @@ -449,10 +450,16 @@ class SnippetWorkflowByIdApi(Resource): @rbac_permission_required( RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False ) - def patch(self, current_user: Account, snippet: CustomizedSnippet, workflow_id: str): + @model_validate(WorkflowUpdatePayload) + def patch( + self, + req_data: WorkflowUpdatePayload, + current_user: Account, + snippet: CustomizedSnippet, + workflow_id: str, + ): """Update a published snippet workflow version's display metadata.""" - payload = WorkflowUpdatePayload.model_validate(console_ns.payload or {}) - update_data = payload.model_dump(exclude_unset=True) + update_data = req_data.model_dump(exclude_unset=True) if not update_data: return {"message": "No valid fields to update"}, 400 @@ -575,16 +582,22 @@ class SnippetDraftNodeRunApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): + @model_validate(SnippetDraftNodeRunPayload) + def post( + self, + req_data: SnippetDraftNodeRunPayload, + current_user: Account, + snippet: CustomizedSnippet, + node_id: str, + ): """ Run a single node in snippet draft workflow. Executes a specific node with provided inputs for single-step debugging. Returns the node execution result including status, outputs, and timing. """ - payload = SnippetDraftNodeRunPayload.model_validate(console_ns.payload or {}) - user_inputs = payload.inputs + user_inputs = req_data.inputs # Get draft workflow for file parsing snippet_service = _snippet_service() @@ -592,14 +605,14 @@ def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): if not draft_workflow: raise NotFound("Draft workflow not found") - files = SnippetGenerateService.parse_files(draft_workflow, payload.files) + files = SnippetGenerateService.parse_files(draft_workflow, req_data.files) workflow_node_execution = SnippetGenerateService.run_draft_node( snippet=snippet, node_id=node_id, user_inputs=user_inputs, account=current_user, - query=payload.query, + query=req_data.query, files=files, session_maker=_snippet_session_maker(), ) @@ -663,14 +676,21 @@ class SnippetDraftRunIterationNodeApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): + @model_validate(SnippetIterationNodeRunPayload) + def post( + self, + req_data: SnippetIterationNodeRunPayload, + current_user: Account, + snippet: CustomizedSnippet, + node_id: str, + ): """ Run a draft workflow iteration node for snippet. Iteration nodes execute their internal sub-graph multiple times over an input list. Returns an SSE event stream with iteration progress and results. """ - args = SnippetIterationNodeRunPayload.model_validate(console_ns.payload or {}).model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = SnippetGenerateService.generate_single_iteration( @@ -708,21 +728,27 @@ class SnippetDraftRunLoopNodeApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): + @model_validate(SnippetLoopNodeRunPayload) + def post( + self, + req_data: SnippetLoopNodeRunPayload, + current_user: Account, + snippet: CustomizedSnippet, + node_id: str, + ): """ Run a draft workflow loop node for snippet. Loop nodes execute their internal sub-graph repeatedly until a condition is met. Returns an SSE event stream with loop progress and results. """ - args = SnippetLoopNodeRunPayload.model_validate(console_ns.payload or {}) try: response = SnippetGenerateService.generate_single_loop( snippet=snippet, user=current_user, node_id=node_id, - args=args, + args=req_data, streaming=True, session_maker=_snippet_session_maker(), ) @@ -751,15 +777,15 @@ class SnippetDraftWorkflowRunApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet): + @model_validate(SnippetDraftRunPayload) + def post(self, req_data: SnippetDraftRunPayload, current_user: Account, snippet: CustomizedSnippet): """ Run draft workflow for snippet. Executes the snippet's draft workflow with the provided inputs and returns an SSE event stream with execution progress and results. """ - payload = SnippetDraftRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = SnippetGenerateService.generate( diff --git a/api/controllers/console/snippets/snippet_workflow_draft_variable.py b/api/controllers/console/snippets/snippet_workflow_draft_variable.py index 20631125831bfe..725c168ea604d1 100644 --- a/api/controllers/console/snippets/snippet_workflow_draft_variable.py +++ b/api/controllers/console/snippets/snippet_workflow_draft_variable.py @@ -36,6 +36,7 @@ from controllers.console.wraps import ( account_initialization_required, edit_permission_required, + model_validate, setup_required, with_current_user, ) @@ -186,9 +187,15 @@ def get(self, current_user: Account, snippet: CustomizedSnippet, variable_id: st @console_ns.response(404, "Variable not found") @_snippet_draft_var_prerequisite @marshal_with(workflow_draft_variable_model) - def patch(self, current_user: Account, snippet: CustomizedSnippet, variable_id: str) -> WorkflowDraftVariable: + @model_validate(WorkflowDraftVariableUpdatePayload) + def patch( + self, + req_data: WorkflowDraftVariableUpdatePayload, + current_user: Account, + snippet: CustomizedSnippet, + variable_id: str, + ) -> WorkflowDraftVariable: draft_var_srv = WorkflowDraftVariableService(session=db.session()) - args_model = WorkflowDraftVariableUpdatePayload.model_validate(console_ns.payload or {}) variable = ensure_variable_access( variable=draft_var_srv.get_variable(variable_id=variable_id), @@ -198,8 +205,8 @@ def patch(self, current_user: Account, snippet: CustomizedSnippet, variable_id: ) _ensure_snippet_draft_variable_row_allowed(variable=variable, variable_id=variable_id) - new_name = args_model.name - raw_value = args_model.value + new_name = req_data.name + raw_value = req_data.value if new_name is None and raw_value is None: return variable diff --git a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py index 9c9338ccc43624..d354e22541d336 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py +++ b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py @@ -6,6 +6,7 @@ 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 @@ def test_get_with_language_param(self, app: Flask): 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 @@ def test_get_fallback_to_user_language(self, app: Flask): 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 @@ def test_get_fallback_to_default_language(self, app: Flask): 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 @@ def test_get_with_language_param(self, app: Flask): 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 @@ def test_get_fallback_to_user_language(self, app: Flask): 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 diff --git a/api/tests/unit_tests/controllers/console/explore/test_trial.py b/api/tests/unit_tests/controllers/console/explore/test_trial.py index 1e4d015b8f2889..9406ad68064a37 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -8,7 +8,7 @@ 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 @@ 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 @@ def test_not_workflow_app(self, app: Flask, account: Account) -> None: 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 @@ def test_success(self, app: Flask, trial_app_workflow: MagicMock, account: Accou 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 @@ def test_workflow_provider_not_init(self, app: Flask, trial_app_workflow: MagicM ), ): 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 @@ def test_workflow_quota_exceeded(self, app: Flask, trial_app_workflow: MagicMock ), ): 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 @@ def test_workflow_model_not_support(self, app: Flask, trial_app_workflow: MagicM ), ): 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 @@ def test_workflow_invoke_error(self, app: Flask, trial_app_workflow: MagicMock, ), ): 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 @@ def test_workflow_rate_limit_error(self, app: Flask, trial_app_workflow: MagicMo ), ): 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 @@ def test_workflow_value_error(self, app: Flask, trial_app_workflow: MagicMock, a ), ): 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 @@ def test_workflow_generic_exception(self, app: Flask, trial_app_workflow: MagicM ), ): 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 @@ def test_not_chat_app(self, app: Flask, account: Account) -> None: 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 @@ def test_success(self, app: Flask, trial_app_chat: MagicMock, account: Account) 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 @@ def test_chat_conversation_not_exists(self, app: Flask, trial_app_chat: MagicMoc ), ): 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 @@ def test_chat_conversation_completed(self, app: Flask, trial_app_chat: MagicMock ), ): 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 @@ def test_chat_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, acc ), ): 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 @@ def test_chat_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, acc ), ): 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 @@ def test_chat_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, accoun ), ): 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 @@ def test_chat_model_not_support(self, app: Flask, trial_app_chat: MagicMock, acc ), ): 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 @@ def test_chat_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_chat_rate_limit_error(self, app: Flask, trial_app_chat: MagicMock, acco ), ): 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 @@ def test_chat_value_error(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_chat_generic_exception(self, app: Flask, trial_app_chat: MagicMock, acc ), ): 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 @@ def test_not_completion_app(self, app: Flask, account: Account) -> None: 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 @@ def test_success(self, app: Flask, trial_app_completion: MagicMock, account: Acc 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 @@ def test_completion_app_config_broken(self, app: Flask, trial_app_completion: Ma ), ): 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 @@ def test_completion_provider_not_init(self, app: Flask, trial_app_completion: Ma ), ): 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 @@ def test_completion_quota_exceeded(self, app: Flask, trial_app_completion: Magic ), ): 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 @@ def test_completion_model_not_support(self, app: Flask, trial_app_completion: Ma ), ): 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 @@ def test_completion_invoke_error(self, app: Flask, trial_app_completion: MagicMo ), ): 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 @@ def test_completion_rate_limit_error(self, app: Flask, trial_app_completion: Mag ), ): 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 @@ def test_completion_value_error(self, app: Flask, trial_app_completion: MagicMoc ), ): 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 @@ def test_completion_generic_exception(self, app: Flask, trial_app_completion: Ma ), ): 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 @@ def test_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_no_audio_uploaded(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def fake_asr(*args, **kwargs): 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 @@ def test_audio_too_large(self, app: Flask, trial_app_chat: MagicMock, account: A ), ): 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 @@ def test_unsupported_audio_type(self, app: Flask, trial_app_chat: MagicMock, acc ), ): 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 @@ def test_provider_not_support_tts(self, app: Flask, trial_app_chat: MagicMock, a ), ): 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 @@ def test_speech_to_text_disabled(self, app: Flask, trial_app_chat: MagicMock, ac ), ): 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 @@ def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: 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 @@ def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Ac 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 @@ def test_success(self, app: Flask, trial_app_chat: MagicMock, account: Account) 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 @@ def test_success_with_message_ref(self, app: Flask, trial_app_chat: MagicMock, a 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 @@ def test_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_provider_not_support(self, app: Flask, trial_app_chat: MagicMock, accou ), ): 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 @@ def test_audio_too_large(self, app: Flask, trial_app_chat: MagicMock, account: A ), ): 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 @@ def test_no_audio_uploaded(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: 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 @@ def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Ac 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 @@ def test_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: 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 @@ def test_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Acco 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 @@ def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Ac ), ): 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 @@ def test_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Acco ), ): with pytest.raises(CompletionRequestError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) class TestTrialChatTextApiExceptionHandlers: @@ -1391,7 +1626,12 @@ def test_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: ), ): 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 @@ def test_unsupported_audio_type(self, app: Flask, trial_app_chat: MagicMock, acc ), ): 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, + ) diff --git a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py index 8a542ef269e38b..b5b2d79411b0f9 100644 --- a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py +++ b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py @@ -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 update_persisted_snippet(*, session: Session, snippet: CustomizedSnippet, ** 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 update_missing_workflow(*, session: Session, snippet: CustomizedSnippet, **_ 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" diff --git a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py index dc8a77d710656c..258c86e53bd763 100644 --- a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py +++ b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py @@ -255,6 +255,7 @@ def record_commit(_session: Session) -> None: 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",