Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 5 additions & 6 deletions api/controllers/console/explore/banner.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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()
Expand Down
15 changes: 7 additions & 8 deletions api/controllers/console/explore/recommended_app.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
31 changes: 15 additions & 16 deletions api/controllers/console/explore/trial.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
"""
Expand All @@ -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
Expand Down Expand Up @@ -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"):
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
80 changes: 53 additions & 27 deletions api/controllers/console/snippets/snippet_workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
RBACResourceScope,
account_initialization_required,
edit_permission_required,
model_validate,
rbac_permission_required,
setup_required,
with_current_user,
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -575,31 +582,37 @@ 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()
draft_workflow = snippet_service.get_draft_workflow(snippet=snippet)
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(),
)
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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(),
)
Expand Down Expand Up @@ -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(
Expand Down
Loading
Loading