Merge branch 'feat/creator-profile-home' into deploy/dev

This commit is contained in:
CodingOnStar 2026-09-01 17:31:18 +08:00
commit ff260f9be3
618 changed files with 27598 additions and 5325 deletions

View File

@ -0,0 +1,70 @@
name: Marketplace Performance E2E
# Opt-in diagnostic: single-sample timing budgets are too noisy to gate every
# PR, so this lane is only run on demand instead of from the main CI pipeline.
on:
workflow_dispatch:
permissions:
contents: read
jobs:
test:
name: Marketplace Performance E2E
runs-on: depot-ubuntu-24.04-4
timeout-minutes: 60
defaults:
run:
shell: bash
steps:
- name: Checkout code
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
with:
persist-credentials: false
- name: Setup web dependencies
uses: ./.github/actions/setup-web
- name: Setup UV and Python
uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0
with:
enable-cache: true
python-version: '3.12'
cache-dependency-glob: |
api/uv.lock
- name: Install API dependencies
run: uv sync --project api --dev
- name: Install Chromium for marketplace performance E2E
timeout-minutes: 15
working-directory: ./e2e
run: vp run e2e:install:ci:chromium
- name: Run marketplace performance benchmark
working-directory: ./e2e
env:
E2E_ADMIN_EMAIL: e2e-admin@example.com
E2E_ADMIN_NAME: E2E Admin
E2E_ADMIN_PASSWORD: E2eAdmin12345
E2E_FORCE_WEB_BUILD: '1'
E2E_INIT_PASSWORD: E2eInit12345
run: vp run e2e:marketplace-performance
- name: Upload Cucumber report
if: ${{ !cancelled() }}
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
with:
name: cucumber-report-marketplace-performance
path: e2e/cucumber-report
retention-days: 7
- name: Upload E2E logs
if: ${{ !cancelled() }}
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
with:
name: e2e-logs-marketplace-performance
path: e2e/.logs/*.log
include-hidden-files: true
retention-days: 7

View File

@ -162,7 +162,7 @@ jobs:
- name: Run Claude Code for Translation Sync
if: steps.context.outputs.CHANGED_FILES != ''
uses: anthropics/claude-code-action@dcb57747bfceeaa1fa72638cae52295d1d853d4a # v1.0.199
uses: anthropics/claude-code-action@a874e9ecd7bb36efdad65429c6b35815f5a08f10 # v1.0.210
with:
anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }}
github_token: ${{ secrets.GITHUB_TOKEN }}

View File

@ -207,6 +207,7 @@ source_modules =
services.account_avatar_service
services.account_change_email_ports
services.account_change_email_service
services.account_email_registration_service
services.account_deletion_service
services.account_deletion_feedback_service
services.account_education_service

View File

@ -642,13 +642,13 @@ class AppListApi(Resource):
)
permissions = enterprise_rbac_service.RBACService.MyPermissions.get(
str(current_tenant_id),
current_tenant_id,
current_user_id,
session=session,
)
if dify_config.RBAC_ENABLED:
access_filter = resolve_app_access_filter(
str(current_tenant_id),
current_tenant_id,
current_user_id,
session=session,
permissions=permissions,
@ -675,7 +675,7 @@ class AppListApi(Resource):
pagination_model = pagination_model.model_copy(
update={
"data": [
item.model_copy(update={"permission_keys": permission_keys_map.get(str(item.id), [])})
item.model_copy(update={"permission_keys": permission_keys_map.get(item.id, [])})
for item in pagination_model.data
]
}
@ -712,7 +712,7 @@ class AppListApi(Resource):
app_service = AppService()
app = app_service.create_app(current_tenant_id, params, current_user, session=session)
permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get(
str(current_tenant_id),
current_tenant_id,
current_user.id,
[str(app.id)],
session=session,
@ -882,7 +882,7 @@ class AppApi(Resource):
app_model.access_mode = app_setting.access_mode
permissions = enterprise_rbac_service.RBACService.MyPermissions.get(
str(current_tenant_id),
current_tenant_id,
current_user.id,
app_id=str(app_model.id),
session=session,
@ -1020,7 +1020,7 @@ class AppCopyApi(Resource):
raise NotFound("App not found")
permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get(
str(current_tenant_id),
current_tenant_id,
current_user.id,
[str(app.id)],
session=session,
@ -1088,7 +1088,7 @@ class AppPublishToCreatorsPlatformApi(Resource):
# TODO: Move this configuration and OAuth orchestration into the Creators Platform application service
# when that domain is refactored. This controller-level integration is a temporary compatibility bridge.
oauth_code = None
client_id = str(dify_config.CREATORS_PLATFORM_OAUTH_CLIENT_ID or "")
client_id = dify_config.CREATORS_PLATFORM_OAUTH_CLIENT_ID or ""
if client_id:
authorization = application_services().oauth_server.issue_authorization_code(
client_id=client_id,

View File

@ -16,6 +16,7 @@ from controllers.common.schema import register_response_schema_models, register_
from controllers.console import console_ns
from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model
from controllers.console.app.error import (
AgentSessionConfigurationChangedError,
AppUnavailableError,
CompletionRequestError,
ConversationCompletedError,
@ -620,6 +621,10 @@ def _raise_agent_stream_error_before_response(response):
if isinstance(response, _ClosableStream):
response.close()
message = error_payload.get("message")
if error_payload.get("code") == AgentSessionConfigurationChangedError.error_code:
raise AgentSessionConfigurationChangedError(
str(message or AgentSessionConfigurationChangedError.description)
)
raise CompletionRequestError(str(message or "Agent App chat failed."))
return _prepend_stream_chunks(buffered, chunk, iterator)

View File

@ -1,3 +1,7 @@
from core.app.apps.agent_app.errors import (
AGENT_SESSION_CONFIGURATION_CHANGED_ERROR_CODE,
AGENT_SESSION_CONFIGURATION_CHANGED_MESSAGE,
)
from libs.exception import BaseHTTPException
@ -49,6 +53,12 @@ class CompletionRequestError(BaseHTTPException):
code = 400
class AgentSessionConfigurationChangedError(BaseHTTPException):
error_code = AGENT_SESSION_CONFIGURATION_CHANGED_ERROR_CODE
description = AGENT_SESSION_CONFIGURATION_CHANGED_MESSAGE
code = 409
class AppMoreLikeThisDisabledError(BaseHTTPException):
error_code = "app_more_like_this_disabled"
description = "The 'More like this' feature is disabled. Please refresh your page."

View File

@ -412,6 +412,7 @@ class InstructionGenerateApi(Resource):
model_config=req_data.model_config_data,
ideal_output=req_data.ideal_output,
workflow_service=WorkflowService(),
session=session,
)
return {"error": "incompatible parameters"}, 400
except ProviderTokenNotInitError as ex:

View File

@ -353,7 +353,7 @@ class WorkflowResponse(ResponseModel):
return [_serialize_environment_variable(item) for item in value]
class _WorkflowResponseSource:
class WorkflowResponseSource:
def __init__(self, workflow: Workflow, *, session: Session) -> None:
self._workflow = workflow
self._session = session
@ -590,7 +590,8 @@ class DraftWorkflowApi(Resource):
"""
# fetch draft workflow by app_model
workflow_service = WorkflowService()
workflow = workflow_service.get_draft_workflow(app_model=app_model, session=db.session())
session = db.session()
workflow = workflow_service.get_draft_workflow(app_model=app_model, session=session)
if not workflow:
raise DraftWorkflowNotExist()
@ -599,9 +600,11 @@ class DraftWorkflowApi(Resource):
# Return workflow with response-only Agent node job projection so the
# front-end can treat draft graph node data as the editing source.
response = WorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json")
response = WorkflowResponse.model_validate(
WorkflowResponseSource(workflow, session=session), from_attributes=True
).model_dump(mode="json")
response["graph"] = WorkflowAgentPublishService.project_draft_bindings_to_graph(
session=db.session(),
session=session,
draft_workflow=workflow,
)
return response
@ -1283,13 +1286,14 @@ class PublishedWorkflowApi(Resource):
"""
# fetch published workflow by app_model
workflow_service = WorkflowService()
workflow = workflow_service.get_published_workflow(app_model=app_model, session=db.session())
session = db.session()
workflow = workflow_service.get_published_workflow(app_model=app_model, session=session)
# return workflow, if not found, return None
if workflow is None:
return None
return dump_response(WorkflowResponse, workflow)
return dump_response(WorkflowResponse, WorkflowResponseSource(workflow, session=session))
@console_ns.expect(console_ns.models[PublishWorkflowPayload.__name__])
@console_ns.response(200, "Workflow published successfully", console_ns.models[WorkflowPublishResponse.__name__])
@ -1512,7 +1516,7 @@ class PublishedAllWorkflowApi(Resource):
)
return WorkflowPaginationResponse.model_validate(
{
"items": [_WorkflowResponseSource(workflow, session=session) for workflow in workflows],
"items": [WorkflowResponseSource(workflow, session=session) for workflow in workflows],
"page": page,
"limit": limit,
"has_more": has_more,
@ -1606,7 +1610,7 @@ class WorkflowByIdApi(Resource):
if not workflow:
raise NotFound("Workflow not found")
response = dump_response(WorkflowResponse, _WorkflowResponseSource(workflow, session=session))
response = dump_response(WorkflowResponse, WorkflowResponseSource(workflow, session=session))
return response

View File

@ -2,8 +2,6 @@ from flask import request
from flask_restx import Resource
from pydantic import BaseModel, Field, field_validator
from configs import dify_config
from constants.languages import get_valid_language, languages
from controllers.common.fields import SimpleResultDataResponse, VerificationTokenResponse
from controllers.common.schema import register_response_schema_models, register_schema_models
from controllers.console import console_ns
@ -11,31 +9,35 @@ from controllers.console.auth.error import (
EmailAlreadyInUseError,
EmailCodeError,
EmailRegisterLimitError,
EmailRegisterRateLimitExceededError,
InvalidEmailError,
InvalidTokenError,
NormalizedEmailAlreadyInUseError,
PasswordMismatchError,
)
from enums import DeploymentEdition
from extensions.ext_database import db
from controllers.console.flask_admission import console_email_registration_admission
from controllers.console.wraps import model_validate
from extensions.ext_application_services import application_services
from fields.base import ResponseModel
from libs.helper import EmailStr, extract_remote_ip
from libs.helper import EmailStr, dump_response, extract_remote_ip
from libs.helper import timezone as validate_timezone_string
from libs.password import valid_password
from models import Account
from services.account_service import AccountService
from services.billing_service import BillingService
from services.errors.account import (
from services.account_errors import (
AccountEmailAlreadyInUseError,
AccountEmailDomainSuspendedError,
AccountEmailFrozenError,
AccountNormalizedEmailAlreadyInUseError,
AccountRegisterError,
SeatsLimitExceededError,
)
from services.errors.account import (
EmailDomainSuspendedError as EmailDomainSuspendedRegistrationError,
EmailRegistrationPasswordMismatchError,
EmailRegistrationSeatsLimitError,
EmailRegistrationSendIPLimitedError,
EmailRegistrationSendRateLimitError,
EmailRegistrationVerificationLimitError,
InvalidEmailRegistrationAddressError,
InvalidEmailRegistrationCodeError,
InvalidEmailRegistrationTokenError,
)
from ..error import AccountInFreezeError, EmailDomainSuspendedError, EmailSendIpLimitError, SeatsLimitExceeded
from ..wraps import email_password_login_enabled, email_register_enabled, model_validate, setup_required
class EmailRegisterSendPayload(BaseModel):
@ -91,146 +93,91 @@ register_response_schema_models(
@console_ns.route("/email-register/send-email")
class EmailRegisterSendEmailApi(Resource):
@setup_required
@email_password_login_enabled
@email_register_enabled
@console_ns.expect(console_ns.models[EmailRegisterSendPayload.__name__])
@console_ns.response(200, "Success", console_ns.models[SimpleResultDataResponse.__name__])
@console_email_registration_admission
@model_validate(EmailRegisterSendPayload)
def post(self, req_data: EmailRegisterSendPayload):
normalized_email = req_data.email.lower()
ip_address = extract_remote_ip(request)
if AccountService.is_email_send_ip_limit(ip_address):
raise EmailSendIpLimitError()
language = "en-US"
if req_data.language is not None and req_data.language in languages:
language = req_data.language
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
freeze_type = BillingService.get_email_freeze_type(normalized_email)
if freeze_type:
if freeze_type == "email_domain_suspended":
raise EmailDomainSuspendedError()
raise AccountInFreezeError()
account = AccountService.get_account_by_email_with_case_fallback(req_data.email, session=db.session())
token = AccountService.send_email_register_email(email=normalized_email, account=account, language=language)
return {"result": "success", "data": token}
def post(self, args: EmailRegisterSendPayload):
try:
token = application_services().accounts.email_registration.send_code(
remote_ip=extract_remote_ip(request),
requested_email=args.email,
requested_language=args.language,
)
except EmailRegistrationSendIPLimitedError:
raise EmailSendIpLimitError() from None
except EmailRegistrationSendRateLimitError as error:
raise EmailRegisterRateLimitExceededError(error.retry_after_minutes) from None
except AccountEmailDomainSuspendedError:
raise EmailDomainSuspendedError() from None
except AccountEmailFrozenError:
raise AccountInFreezeError() from None
return dump_response(SimpleResultDataResponse, {"result": "success", "data": token})
@console_ns.route("/email-register/validity")
class EmailRegisterCheckApi(Resource):
@setup_required
@email_password_login_enabled
@email_register_enabled
@console_ns.expect(console_ns.models[EmailRegisterValidityPayload.__name__])
@console_ns.response(200, "Success", console_ns.models[VerificationTokenResponse.__name__])
@console_email_registration_admission
@model_validate(EmailRegisterValidityPayload)
def post(self, req_data: EmailRegisterValidityPayload):
user_email = req_data.email.lower()
is_email_register_error_rate_limit = AccountService.is_email_register_error_rate_limit(user_email)
if is_email_register_error_rate_limit:
raise EmailRegisterLimitError()
token_data = AccountService.get_email_register_data(req_data.token)
if token_data is None:
raise InvalidTokenError()
token_email = token_data.get("email")
normalized_token_email = token_email.lower() if isinstance(token_email, str) else token_email
if user_email != normalized_token_email:
raise InvalidEmailError()
if req_data.code != token_data.get("code"):
AccountService.add_email_register_error_rate_limit(user_email)
raise EmailCodeError()
# Verified, revoke the first token
AccountService.revoke_email_register_token(req_data.token)
# Refresh token data by generating a new token
_, new_token = AccountService.generate_email_register_token(
user_email, code=req_data.code, additional_data={"phase": "register"}
def post(self, args: EmailRegisterValidityPayload):
try:
verification = application_services().accounts.email_registration.verify_code(
email=args.email,
code=args.code,
token=args.token,
)
except EmailRegistrationVerificationLimitError:
raise EmailRegisterLimitError() from None
except InvalidEmailRegistrationTokenError:
raise InvalidTokenError() from None
except InvalidEmailRegistrationAddressError:
raise InvalidEmailError() from None
except InvalidEmailRegistrationCodeError:
raise EmailCodeError() from None
return dump_response(
VerificationTokenResponse,
{
"is_valid": True,
"email": verification.email,
"token": verification.token,
},
)
AccountService.reset_email_register_error_rate_limit(user_email)
return {"is_valid": True, "email": normalized_token_email, "token": new_token}
@console_ns.route("/email-register")
class EmailRegisterResetApi(Resource):
@setup_required
@email_password_login_enabled
@email_register_enabled
@console_ns.expect(console_ns.models[EmailRegisterResetPayload.__name__])
@console_ns.response(200, "Success", console_ns.models[EmailRegisterResetResponse.__name__])
@console_email_registration_admission
@model_validate(EmailRegisterResetPayload)
def post(self, req_data: EmailRegisterResetPayload):
# Validate passwords match
if req_data.new_password != req_data.password_confirm:
raise PasswordMismatchError()
# Validate token and get register data
register_data = AccountService.get_email_register_data(req_data.token)
if not register_data:
raise InvalidTokenError()
# Must use token in reset phase
if register_data.get("phase", "") != "register":
raise InvalidTokenError()
# Revoke token to prevent reuse
AccountService.revoke_email_register_token(req_data.token)
email = register_data.get("email", "")
normalized_email = email.lower()
account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session())
if account:
raise EmailAlreadyInUseError()
ip_address = extract_remote_ip(request)
account = self._create_new_account(
email=normalized_email,
password=req_data.password_confirm,
timezone=req_data.timezone,
language=req_data.language,
ip_address=ip_address,
)
token_pair = AccountService.login(account=account, session=db.session(), ip_address=ip_address)
AccountService.reset_login_error_rate_limit(normalized_email)
return {"result": "success", "data": token_pair.model_dump()}
def _create_new_account(
self,
email: str,
password: str,
timezone: str | None = None,
language: str | None = None,
ip_address: str | None = None,
) -> Account:
def post(self, args: EmailRegisterResetPayload):
try:
return AccountService.create_account_and_tenant(
email=email,
name=email,
password=password,
interface_language=get_valid_language(language),
timezone=timezone,
ip_address=ip_address,
check_normalized_email=True,
session=db.session(),
token_pair = application_services().accounts.email_registration.register(
remote_ip=extract_remote_ip(request),
token=args.token,
new_password=args.new_password,
password_confirm=args.password_confirm,
language=args.language,
timezone=args.timezone,
)
except SeatsLimitExceededError:
raise SeatsLimitExceeded()
except EmailDomainSuspendedRegistrationError as exc:
raise EmailDomainSuspendedError() from exc
except AccountNormalizedEmailAlreadyInUseError as exc:
raise NormalizedEmailAlreadyInUseError() from exc
except AccountRegisterError as exc:
raise AccountInFreezeError() from exc
except EmailRegistrationPasswordMismatchError:
raise PasswordMismatchError() from None
except InvalidEmailRegistrationTokenError:
raise InvalidTokenError() from None
except AccountNormalizedEmailAlreadyInUseError:
raise NormalizedEmailAlreadyInUseError() from None
except AccountEmailAlreadyInUseError:
raise EmailAlreadyInUseError() from None
except EmailRegistrationSeatsLimitError:
raise SeatsLimitExceeded() from None
except AccountEmailDomainSuspendedError:
raise EmailDomainSuspendedError() from None
except AccountEmailFrozenError:
raise AccountInFreezeError() from None
return dump_response(
EmailRegisterResetResponse,
{"result": "success", "data": token_pair},
)

View File

@ -55,7 +55,7 @@ class PasswordResetRateLimitExceededError(BaseHTTPException):
code = 429
def __init__(self, minutes: int = 1):
description = self.description.format(minutes=int(minutes)) if self.description else None
description = self.description.format(minutes=minutes) if self.description else None
super().__init__(description=description)
@ -65,7 +65,7 @@ class EmailRegisterRateLimitExceededError(BaseHTTPException):
code = 429
def __init__(self, minutes: int = 1):
description = self.description.format(minutes=int(minutes)) if self.description else None
description = self.description.format(minutes=minutes) if self.description else None
super().__init__(description=description)
@ -75,7 +75,7 @@ class EmailChangeRateLimitExceededError(BaseHTTPException):
code = 429
def __init__(self, minutes: int = 1):
description = self.description.format(minutes=int(minutes)) if self.description else None
description = self.description.format(minutes=minutes) if self.description else None
super().__init__(description=description)
@ -85,7 +85,7 @@ class OwnerTransferRateLimitExceededError(BaseHTTPException):
code = 429
def __init__(self, minutes: int = 1):
description = self.description.format(minutes=int(minutes)) if self.description else None
description = self.description.format(minutes=minutes) if self.description else None
super().__init__(description=description)
@ -137,7 +137,7 @@ class EmailCodeLoginRateLimitExceededError(BaseHTTPException):
code = 429
def __init__(self, minutes: int = 5):
description = self.description.format(minutes=int(minutes)) if self.description else None
description = self.description.format(minutes=minutes) if self.description else None
super().__init__(description=description)
@ -147,7 +147,7 @@ class EmailCodeAccountDeletionRateLimitExceededError(BaseHTTPException):
code = 429
def __init__(self, minutes: int = 5):
description = self.description.format(minutes=int(minutes)) if self.description else None
description = self.description.format(minutes=minutes) if self.description else None
super().__init__(description=description)

View File

@ -26,6 +26,7 @@ from controllers.console.app.workflow import (
DefaultBlockConfigsResponse,
WorkflowPaginationResponse,
WorkflowResponse,
WorkflowResponseSource,
)
from controllers.console.app.wraps import with_session
from controllers.console.datasets.wraps import get_rag_pipeline, load_rag_pipeline
@ -202,14 +203,15 @@ class DraftRagPipelineApi(Resource):
Get draft rag pipeline's workflow
"""
# fetch draft workflow by app_model
rag_pipeline_service = RagPipelineService(db.session())
session = db.session()
rag_pipeline_service = RagPipelineService(session)
workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline)
if not workflow:
raise DraftWorkflowNotExist()
# return workflow, if not found, return 404
return dump_response(WorkflowResponse, workflow)
return dump_response(WorkflowResponse, WorkflowResponseSource(workflow, session=session))
@setup_required
@login_required
@ -548,14 +550,15 @@ class PublishedRagPipelineApi(Resource):
if not pipeline.is_published:
return None
# fetch published workflow by pipeline
rag_pipeline_service = RagPipelineService(db.session())
session = db.session()
rag_pipeline_service = RagPipelineService(session)
workflow = rag_pipeline_service.get_published_workflow(pipeline=pipeline)
# return workflow, if not found, return None
if workflow is None:
return None
return dump_response(WorkflowResponse, workflow)
return dump_response(WorkflowResponse, WorkflowResponseSource(workflow, session=session))
@console_ns.response(200, "Success", console_ns.models[RagPipelineWorkflowPublishResponse.__name__])
@setup_required
@ -684,7 +687,7 @@ class PublishedAllRagPipelineApi(Resource):
return WorkflowPaginationResponse.model_validate(
{
"items": workflows,
"items": [WorkflowResponseSource(workflow, session=session) for workflow in workflows],
"page": page,
"limit": limit,
"has_more": has_more,
@ -763,7 +766,7 @@ class RagPipelineByIdApi(Resource):
if not workflow:
raise NotFound("Workflow not found")
return dump_response(WorkflowResponse, workflow)
return dump_response(WorkflowResponse, WorkflowResponseSource(workflow, session=session))
@console_ns.response(204, "Workflow deleted successfully")
@setup_required

View File

@ -22,6 +22,22 @@ from libs.login import current_account_with_tenant, login_required
from machinery.context import RequestContext
from machinery.errors import AdmissionConfigurationError
from models.account import TenantAccountRole
from services.feature_service import FeatureService
def console_email_registration_admission[T, **P, R](
view: Callable[Concatenate[T, P], R],
) -> Callable[Concatenate[T, P], R | Response]:
"""Apply the complete admission policy for anonymous email registration."""
@wraps(view)
def check_registration_features(self: T, /, *args: P.args, **kwargs: P.kwargs) -> R:
features = FeatureService.get_system_features()
if not features.enable_email_password_login or not features.is_allow_register:
abort(403)
return view(self, *args, **kwargs)
return setup_required(check_registration_features)
def console_account_admission[T, **P, R](

View File

@ -1,56 +1,16 @@
from collections.abc import Mapping
from typing import TypedDict
from flask_restx import Resource
from pydantic import BaseModel, Field
from controllers.common.fields import SimpleResultResponse
from controllers.common.schema import register_response_schema_models, register_schema_models
from controllers.console import console_ns
from controllers.console.wraps import (
account_initialization_required,
model_validate,
only_edition_cloud,
setup_required,
with_current_user,
)
from controllers.console.flask_admission import console_account_admission
from controllers.console.wraps import model_validate
from enums import DeploymentEdition
from extensions.ext_application_services import application_services
from fields.base import ResponseModel
from libs.login import login_required
from models import Account
from services.billing_service import BillingService
# Notification content is stored under three lang tags.
_FALLBACK_LANG = "en-US"
class NotificationLangContent(TypedDict, total=False):
lang: str
title: str
subtitle: str
body: str
titlePicUrl: str
class NotificationItemDict(TypedDict):
notification_id: str | None
frequency: str | None
lang: str
title: str
subtitle: str
body: str
title_pic_url: str
class NotificationResponseDict(TypedDict):
should_show: bool
notifications: list[NotificationItemDict]
def _pick_lang_content(contents: Mapping[str, NotificationLangContent], lang: str) -> NotificationLangContent:
"""Return the single LangContent for *lang*, falling back to English."""
return (
contents.get(lang) or contents.get(_FALLBACK_LANG) or next(iter(contents.values()), NotificationLangContent())
)
from libs.helper import dump_response
from machinery.context import RequestContext
class DismissNotificationPayload(BaseModel):
@ -92,39 +52,10 @@ class NotificationApi(Resource):
},
)
@console_ns.response(200, "Success", console_ns.models[NotificationResponse.__name__])
@setup_required
@login_required
@with_current_user
@account_initialization_required
@only_edition_cloud
def get(self, current_user: Account):
result = BillingService.get_account_notification(str(current_user.id))
# Proto JSON uses camelCase field names (Kratos default marshaling).
response: NotificationResponseDict
if not result.get("shouldShow"):
response = {"should_show": False, "notifications": []}
return response, 200
lang = current_user.interface_language or _FALLBACK_LANG
notifications: list[NotificationItemDict] = []
for notification in result.get("notifications") or []:
contents: Mapping[str, NotificationLangContent] = notification.get("contents") or {}
lang_content = _pick_lang_content(contents, lang)
item: NotificationItemDict = {
"notification_id": notification.get("notificationId"),
"frequency": notification.get("frequency"),
"lang": lang_content.get("lang", lang),
"title": lang_content.get("title", ""),
"subtitle": lang_content.get("subtitle", ""),
"body": lang_content.get("body", ""),
"title_pic_url": lang_content.get("titlePicUrl", ""),
}
notifications.append(item)
response = {"should_show": bool(notifications), "notifications": notifications}
return response, 200
@console_account_admission(editions=frozenset({DeploymentEdition.CLOUD}))
def get(self, request_context: RequestContext):
result = application_services().notifications.get_active(request_context)
return dump_response(NotificationResponse, result), 200
@console_ns.route("/notification/dismiss")
@ -134,17 +65,10 @@ class NotificationDismissApi(Resource):
description="Mark a notification as dismissed for the current user.",
responses={200: "Success", 401: "Unauthorized"},
)
@setup_required
@login_required
@with_current_user
@account_initialization_required
@only_edition_cloud
@console_account_admission(editions=frozenset({DeploymentEdition.CLOUD}))
@console_ns.expect(console_ns.models[DismissNotificationPayload.__name__])
@console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__])
@model_validate(DismissNotificationPayload)
def post(self, payload: DismissNotificationPayload, current_user: Account):
BillingService.dismiss_notification(
notification_id=payload.notification_id,
account_id=str(current_user.id),
)
return {"result": "success"}, 200
def post(self, payload: DismissNotificationPayload, request_context: RequestContext):
application_services().notifications.dismiss(request_context, payload.notification_id)
return dump_response(SimpleResultResponse, {"result": "success"}), 200

View File

@ -7,36 +7,20 @@ action-based so callers do not replace server-side arrays with stale snapshots.
"""
from datetime import datetime
from typing import Literal, cast
from flask_restx import Resource
from pydantic import BaseModel, ConfigDict, Field, model_validator
from controllers.common.schema import register_response_schema_models, register_schema_models
from extensions.ext_database import db
from controllers.console.flask_admission import console_account_admission
from controllers.console.wraps import model_validate
from extensions.ext_application_services import application_services
from fields.base import ResponseModel
from libs.helper import dump_response
from libs.login import login_required
from models import Account
from services.step_by_step_tour_service import StepByStepTourPatch, StepByStepTourService
from machinery.context import RequestContext
from services.entities.onboarding_entities import StepByStepTourAction, StepByStepTourPatch, StepByStepTourTaskId
from . import console_ns
from .wraps import (
account_initialization_required,
model_validate,
setup_required,
with_current_tenant_id,
with_current_user,
)
StepByStepTourAction = Literal[
"skip",
"complete_task",
"uncomplete_task",
"enable_current_workspace",
"disable_current_workspace",
]
StepByStepTourTaskId = Literal["home", "studio", "knowledge", "integration"]
class StepByStepTourStatePatchPayload(BaseModel):
@ -74,39 +58,22 @@ class StepByStepTourStateApi(Resource):
@console_ns.doc("get_step_by_step_tour_state")
@console_ns.doc(description="Get account-level Step-by-step Tour state")
@console_ns.response(200, "Success", console_ns.models[StepByStepTourStateResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@with_current_user
@with_current_tenant_id
def get(self, current_tenant_id: str, current_user: Account):
@console_account_admission()
def get(self, request_context: RequestContext):
return dump_response(
StepByStepTourStateResponse,
StepByStepTourService.get_state(
account=current_user,
current_tenant_id=current_tenant_id,
session=db.session,
),
application_services().step_by_step_tour.get_state(request_context),
)
@console_ns.doc("patch_step_by_step_tour_state")
@console_ns.doc(description="Update account-level Step-by-step Tour state")
@console_ns.expect(console_ns.models[StepByStepTourStatePatchPayload.__name__])
@console_ns.response(200, "Success", console_ns.models[StepByStepTourStateResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@with_current_user
@with_current_tenant_id
@console_account_admission()
@model_validate(StepByStepTourStatePatchPayload)
def patch(self, req_data: StepByStepTourStatePatchPayload, current_tenant_id: str, current_user: Account):
patch = cast(StepByStepTourPatch, req_data.model_dump(exclude_unset=True, exclude_none=True))
def patch(self, req_data: StepByStepTourStatePatchPayload, request_context: RequestContext):
patch = StepByStepTourPatch(action=req_data.action, task_id=req_data.task_id)
return dump_response(
StepByStepTourStateResponse,
StepByStepTourService.patch_state(
account=current_user,
current_tenant_id=current_tenant_id,
patch=patch,
session=db.session,
),
application_services().step_by_step_tour.patch_state(request_context, patch),
)

View File

@ -19,6 +19,7 @@ from controllers.console.app.workflow import (
WorkflowPaginationResponse,
WorkflowPublishResponse,
WorkflowResponse,
WorkflowResponseSource,
WorkflowRestoreResponse,
)
from controllers.console.snippets.payloads import (
@ -179,9 +180,12 @@ class SnippetDraftWorkflowApi(Resource):
raise DraftWorkflowNotExist()
workflow.conversation_variables = []
response = SnippetWorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json")
session = db.session()
response = SnippetWorkflowResponse.model_validate(
WorkflowResponseSource(workflow, session=session), from_attributes=True
).model_dump(mode="json")
response["graph"] = WorkflowAgentPublishService.project_draft_bindings_to_graph(
session=db.session(),
session=session,
draft_workflow=workflow,
)
response["input_fields"] = snippet.input_fields_list
@ -274,7 +278,9 @@ class SnippetPublishedWorkflowApi(Resource):
if not workflow:
return None
response = SnippetWorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json")
response = SnippetWorkflowResponse.model_validate(
WorkflowResponseSource(workflow, session=db.session()), from_attributes=True
).model_dump(mode="json")
response["input_fields"] = snippet.input_fields_list
return response
@ -365,15 +371,15 @@ class SnippetPublishedAllWorkflowApi(Resource):
limit=req_data.limit,
)
response = SnippetWorkflowPaginationResponse.model_validate(
{
"items": workflows,
"page": req_data.page,
"limit": req_data.limit,
"has_more": has_more,
},
from_attributes=True,
).model_dump(mode="json")
response = SnippetWorkflowPaginationResponse.model_validate(
{
"items": [WorkflowResponseSource(workflow, session=session) for workflow in workflows],
"page": req_data.page,
"limit": req_data.limit,
"has_more": has_more,
},
from_attributes=True,
).model_dump(mode="json")
for item in response["items"]:
item["input_fields"] = snippet.input_fields_list
return response
@ -464,9 +470,11 @@ class SnippetWorkflowByIdApi(Resource):
if not workflow:
raise NotFound("Workflow not found")
response = SnippetWorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json")
response["input_fields"] = snippet.input_fields_list
return response
response = SnippetWorkflowResponse.model_validate(
WorkflowResponseSource(workflow, session=session), from_attributes=True
).model_dump(mode="json")
response["input_fields"] = snippet.input_fields_list
return response
@console_ns.doc("delete_snippet_workflow_by_id")
@console_ns.doc(description="Delete a published snippet workflow version")

View File

@ -292,10 +292,8 @@ class AccountInitApi(Resource):
@console_ns.expect(console_ns.models[AccountInitPayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SimpleResultResponse.__name__])
@console_account_admission(require_initialized=False)
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountInitPayload.model_validate(payload)
@model_validate(AccountInitPayload)
def post(self, args: AccountInitPayload, request_context: RequestContext):
try:
application_services().accounts.initialization.initialize(
request_context,
@ -344,9 +342,8 @@ class AccountNameApi(Resource):
@console_ns.expect(console_ns.models[AccountNamePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountNamePayload.model_validate(payload)
@model_validate(AccountNamePayload)
def post(self, args: AccountNamePayload, request_context: RequestContext):
return _update_account_profile(request_context, AccountProfileChanges(name=args.name))
@ -371,9 +368,8 @@ class AccountAvatarApi(Resource):
@console_ns.doc(description="Deprecated. Use PATCH /account/profile instead.")
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountAvatarPayload.model_validate(payload)
@model_validate(AccountAvatarPayload)
def post(self, args: AccountAvatarPayload, request_context: RequestContext):
return _update_account_profile(request_context, AccountProfileChanges(avatar=args.avatar))
@ -387,9 +383,8 @@ class AccountInterfaceLanguageApi(Resource):
@console_ns.expect(console_ns.models[AccountInterfaceLanguagePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountInterfaceLanguagePayload.model_validate(payload)
@model_validate(AccountInterfaceLanguagePayload)
def post(self, args: AccountInterfaceLanguagePayload, request_context: RequestContext):
return _update_account_profile(
request_context,
AccountProfileChanges(interface_language=args.interface_language),
@ -406,9 +401,8 @@ class AccountInterfaceThemeApi(Resource):
@console_ns.expect(console_ns.models[AccountInterfaceThemePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountInterfaceThemePayload.model_validate(payload)
@model_validate(AccountInterfaceThemePayload)
def post(self, args: AccountInterfaceThemePayload, request_context: RequestContext):
return _update_account_profile(
request_context,
AccountProfileChanges(interface_theme=args.interface_theme),
@ -425,9 +419,8 @@ class AccountTimezoneApi(Resource):
@console_ns.expect(console_ns.models[AccountTimezonePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountTimezonePayload.model_validate(payload)
@model_validate(AccountTimezonePayload)
def post(self, args: AccountTimezonePayload, request_context: RequestContext):
return _update_account_profile(request_context, AccountProfileChanges(timezone=args.timezone))
@ -436,10 +429,8 @@ class AccountPasswordApi(Resource):
@console_ns.expect(console_ns.models[AccountPasswordPayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountPasswordPayload.model_validate(payload)
@model_validate(AccountPasswordPayload)
def post(self, args: AccountPasswordPayload, request_context: RequestContext):
try:
assert args.password is not None
account = application_services().accounts.password.change(
@ -498,10 +489,8 @@ class AccountDeleteApi(Resource):
@console_ns.expect(console_ns.models[AccountDeletePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SimpleResultResponse.__name__])
@console_account_admission()
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = AccountDeletePayload.model_validate(payload)
@model_validate(AccountDeletePayload)
def post(self, args: AccountDeletePayload, request_context: RequestContext):
try:
application_services().accounts.deletion.request_deletion(
request_context,
@ -519,10 +508,8 @@ class AccountDeleteUpdateFeedbackApi(Resource):
@console_ns.expect(console_ns.models[AccountDeletionFeedbackPayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SimpleResultResponse.__name__])
@setup_required
def post(self):
payload = console_ns.payload or {}
args = AccountDeletionFeedbackPayload.model_validate(payload)
@model_validate(AccountDeletionFeedbackPayload)
def post(self, args: AccountDeletionFeedbackPayload):
application_services().accounts.deletion_feedback.submit(email=args.email, feedback=args.feedback)
return SimpleResultResponse(result="success").model_dump(mode="json")
@ -547,9 +534,8 @@ class EducationApi(Resource):
@console_ns.expect(console_ns.models[EducationActivatePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[EducationActivateResponse.__name__])
@console_account_admission(editions=frozenset({DeploymentEdition.CLOUD}))
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = EducationActivatePayload.model_validate(payload)
@model_validate(EducationActivatePayload)
def post(self, args: EducationActivatePayload, request_context: RequestContext):
try:
activation = application_services().accounts.education.activate(
request_context,
@ -574,10 +560,8 @@ class EducationAutoCompleteApi(Resource):
@console_ns.doc(params=query_params_from_model(EducationAutocompleteQuery))
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[EducationAutocompleteResponse.__name__])
@console_account_admission(editions=frozenset({DeploymentEdition.CLOUD}))
def get(self, request_context: RequestContext):
payload = request.args.to_dict(flat=True)
args = EducationAutocompleteQuery.model_validate(payload)
@model_validate(EducationAutocompleteQuery)
def get(self, args: EducationAutocompleteQuery, request_context: RequestContext):
return dump_response(
EducationAutocompleteResponse,
application_services().accounts.education.autocomplete(
@ -594,10 +578,8 @@ class ChangeEmailSendEmailApi(Resource):
@console_ns.expect(console_ns.models[ChangeEmailSendPayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SimpleResultDataResponse.__name__])
@console_account_admission(require_change_email_enabled=True)
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = ChangeEmailSendPayload.model_validate(payload)
@model_validate(ChangeEmailSendPayload)
def post(self, args: ChangeEmailSendPayload, request_context: RequestContext):
ip_address = extract_remote_ip(request)
language = "zh-Hans" if args.language == "zh-Hans" else "en-US"
try:
@ -627,10 +609,8 @@ class ChangeEmailCheckApi(Resource):
@console_ns.expect(console_ns.models[ChangeEmailValidityPayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[VerificationTokenResponse.__name__])
@console_account_admission(require_change_email_enabled=True)
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = ChangeEmailValidityPayload.model_validate(payload)
@model_validate(ChangeEmailValidityPayload)
def post(self, args: ChangeEmailValidityPayload, request_context: RequestContext):
try:
verification = application_services().accounts.change_email.verify_code(
request_context,
@ -656,9 +636,8 @@ class ChangeEmailResetApi(Resource):
@console_ns.expect(console_ns.models[ChangeEmailResetPayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountResponse.__name__])
@console_account_admission(require_change_email_enabled=True)
def post(self, request_context: RequestContext):
payload = console_ns.payload or {}
args = ChangeEmailResetPayload.model_validate(payload)
@model_validate(ChangeEmailResetPayload)
def post(self, args: ChangeEmailResetPayload, request_context: RequestContext):
try:
updated_account = application_services().accounts.change_email.reset(
request_context,
@ -684,9 +663,8 @@ class CheckEmailUnique(Resource):
@console_ns.expect(console_ns.models[CheckEmailUniquePayload.__name__])
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SimpleResultResponse.__name__])
@setup_required
def post(self):
payload = console_ns.payload or {}
args = CheckEmailUniquePayload.model_validate(payload)
@model_validate(CheckEmailUniquePayload)
def post(self, args: CheckEmailUniquePayload):
try:
application_services().accounts.change_email.ensure_available(args.email)
except account_errors.AccountEmailDomainSuspendedError:

View File

@ -221,7 +221,7 @@ class CustomizedSnippetDetailApi(Resource):
"""Update customized snippet."""
snippet_service = _snippet_service()
snippet = snippet_service.get_snippet_by_id(
snippet_id=str(snippet_id),
snippet_id=snippet_id,
tenant_id=current_tenant_id,
)
@ -265,7 +265,7 @@ class CustomizedSnippetDetailApi(Resource):
"""Delete customized snippet."""
snippet_service = _snippet_service()
snippet = snippet_service.get_snippet_by_id(
snippet_id=str(snippet_id),
snippet_id=snippet_id,
tenant_id=current_tenant_id,
)
@ -304,7 +304,7 @@ class CustomizedSnippetExportApi(Resource):
"""Export snippet as DSL."""
snippet_service = _snippet_service()
snippet = snippet_service.get_snippet_by_id(
snippet_id=str(snippet_id),
snippet_id=snippet_id,
tenant_id=current_tenant_id,
)
@ -428,7 +428,7 @@ class CustomizedSnippetCheckDependenciesApi(Resource):
"""Check dependencies for a snippet."""
snippet_service = _snippet_service()
snippet = snippet_service.get_snippet_by_id(
snippet_id=str(snippet_id),
snippet_id=snippet_id,
tenant_id=current_tenant_id,
)
@ -458,7 +458,7 @@ class CustomizedSnippetUseCountIncrementApi(Resource):
"""Increment snippet use count when it is inserted into a workflow."""
snippet_service = _snippet_service()
snippet = snippet_service.get_snippet_by_id(
snippet_id=str(snippet_id),
snippet_id=snippet_id,
tenant_id=current_tenant_id,
)

View File

@ -352,19 +352,6 @@ def email_password_login_enabled[**P, R](view: Callable[P, R]) -> Callable[P, R]
return decorated
def email_register_enabled[**P, R](view: Callable[P, R]) -> Callable[P, R]:
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs):
features = FeatureService.get_system_features()
if features.is_allow_register:
return view(*args, **kwargs)
# otherwise, return 403
abort(403)
return decorated
def enable_change_email[**P, R](view: Callable[P, R]) -> Callable[P, R]:
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs):
@ -652,6 +639,23 @@ def with_current_user_id[T, **P, R](
return decorated
def validate_request[M: BaseModel](model: type[M]) -> M:
"""Parse and validate the current request without exposing submitted values."""
if request.method == "GET":
raw = request.args.to_dict(flat=True)
elif request.method == "DELETE":
raw = request.args.to_dict(flat=True) or (request.get_json(silent=True) or {})
else:
raw = request.get_json(silent=True) or {}
try:
return model.model_validate(raw)
except ValidationError as exc:
errors = exc.errors(include_url=False, include_input=False, include_context=False)
raise UnprocessableEntity(json.dumps(errors)) from None
def model_validate[T, M: BaseModel, **P, R](
model: type[M],
) -> Callable[
@ -671,19 +675,7 @@ def model_validate[T, M: BaseModel, **P, R](
) -> Callable[Concatenate[T, P], R]:
@wraps(view)
def wrapper(self: T, *args: P.args, **kwargs: P.kwargs) -> R:
if request.method == "GET":
raw = request.args.to_dict(flat=True)
elif request.method == "DELETE":
raw = request.args.to_dict(flat=True) or (request.get_json(silent=True) or {})
else:
raw = request.get_json(silent=True) or {}
try:
validated = model.model_validate(raw)
except ValidationError as exc:
raise UnprocessableEntity(exc.json())
return view(self, validated, *args, **kwargs)
return view(self, validate_request(model), *args, **kwargs)
return wrapper

View File

@ -67,6 +67,7 @@ class OpenApiErrorCode(StrEnum):
MEMBER_LICENSE_EXCEEDED = "member_license_exceeded"
HUMAN_INPUT_FORM_NOT_FOUND = "form_not_found"
RECIPIENT_SURFACE_MISMATCH = "recipient_surface_mismatch"
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE = "trigger_workflow_service_mode_unavailable"
class ErrorDetail(BaseModel):

View File

@ -35,6 +35,7 @@ from controllers.service_api.app.error import (
ProviderModelCurrentlyNotSupportError,
ProviderNotInitializeError,
ProviderQuotaExceededError,
TriggerWorkflowServiceModeUnavailableError,
)
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
from core.app.apps.base_app_queue_manager import AppQueueManager
@ -57,6 +58,9 @@ from services.errors.app import (
WorkflowIdFormatError,
WorkflowNotFoundError,
)
from services.errors.app import (
TriggerWorkflowServiceModeUnavailableError as TriggerWorkflowServiceModeUnavailableServiceError,
)
from services.errors.llm import InvokeRateLimitError
logger = logging.getLogger(__name__)
@ -70,6 +74,8 @@ def _translate_service_errors() -> Generator[None, None, None]:
raise NotFound(str(ex))
except (IsDraftWorkflowError, WorkflowIdFormatError) as ex:
raise BadRequest(str(ex))
except TriggerWorkflowServiceModeUnavailableServiceError:
raise TriggerWorkflowServiceModeUnavailableError()
except services.errors.conversation.ConversationNotExistsError:
raise NotFound("Conversation Not Exists.")
except services.errors.conversation.ConversationCompletedError:

View File

@ -8,6 +8,7 @@ import services
from controllers.common.controller_schemas import TextToAudioPayload
from controllers.common.fields import AudioBinaryResponse, AudioTranscriptResponse
from controllers.common.schema import register_response_schema_models, register_schema_model
from controllers.console.wraps import model_validate
from controllers.service_api import service_api_ns
from controllers.service_api.app.error import (
AppUnavailableError,
@ -181,14 +182,13 @@ class TextApi(Resource):
# TTS returns provider audio bytes, so the success response is intentionally schema-less.
@service_api_ns.response(200, "Text successfully converted to audio")
@validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON))
def post(self, app_model: App, end_user: EndUser):
@model_validate(TextToAudioPayload)
def post(self, payload: TextToAudioPayload, app_model: App, end_user: EndUser):
"""Convert text to audio using text-to-speech.
Converts the provided text to audio using the specified voice.
"""
try:
payload = TextToAudioPayload.model_validate(service_api_ns.payload or {})
message_id = payload.message_id
text = payload.text
voice = payload.voice

View File

@ -11,6 +11,7 @@ from werkzeug.exceptions import BadRequest, NotFound
import services
from controllers.common.controller_schemas import ConversationRenamePayload
from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models
from controllers.console.wraps import model_validate
from controllers.service_api import service_api_ns
from controllers.service_api.app.error import NotChatAppError
from controllers.service_api.schema import expect_user_json, expect_with_user
@ -293,7 +294,8 @@ class ConversationRenameApi(Resource):
service_api_ns.models[SimpleConversation.__name__],
)
@validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON))
def post(self, app_model: App, end_user: EndUser, conversation_id: UUID):
@model_validate(ConversationRenamePayload)
def post(self, payload: ConversationRenamePayload, app_model: App, end_user: EndUser, conversation_id: UUID):
"""Rename a conversation or auto-generate a name."""
app_mode = AppMode.value_of(app_model.mode)
if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT}:
@ -301,8 +303,6 @@ class ConversationRenameApi(Resource):
conversation_id_str = str(conversation_id)
payload = ConversationRenamePayload.model_validate(service_api_ns.payload or {})
try:
session = db.session()
conversation = ConversationService.rename(
@ -408,7 +408,15 @@ class ConversationVariableDetailApi(Resource):
service_api_ns.models[ConversationVariableResponse.__name__],
)
@validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON))
def put(self, app_model: App, end_user: EndUser, conversation_id: UUID, variable_id: UUID):
@model_validate(ConversationVariableUpdatePayload)
def put(
self,
payload: ConversationVariableUpdatePayload,
app_model: App,
end_user: EndUser,
conversation_id: UUID,
variable_id: UUID,
):
"""Update a conversation variable's value.
Allows updating the value of a specific conversation variable.
@ -421,8 +429,6 @@ class ConversationVariableDetailApi(Resource):
conversation_id_str = str(conversation_id)
variable_id_str = str(variable_id)
payload = ConversationVariableUpdatePayload.model_validate(service_api_ns.payload or {})
try:
variable = ConversationService.update_conversation_variable(
app_model, conversation_id_str, variable_id_str, end_user, payload.value, session=db.session()

View File

@ -1,4 +1,8 @@
from libs.exception import BaseHTTPException
from services.errors.app import (
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_CODE,
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_MESSAGE,
)
class AppUnavailableError(BaseHTTPException):
@ -37,6 +41,12 @@ class WorkflowVersionExecutionNotAllowedError(BaseHTTPException):
code = 403
class TriggerWorkflowServiceModeUnavailableError(BaseHTTPException):
error_code = TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_CODE
description = TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_MESSAGE
code = 403
class ConversationCompletedError(BaseHTTPException):
error_code = "conversation_completed"
description = "The conversation has ended. Please start a new conversation."

View File

@ -28,6 +28,7 @@ from controllers.service_api.app.error import (
ProviderModelCurrentlyNotSupportError,
ProviderNotInitializeError,
ProviderQuotaExceededError,
TriggerWorkflowServiceModeUnavailableError,
WorkflowVersionExecutionNotAllowedError,
)
from controllers.service_api.schema import (
@ -61,7 +62,14 @@ from models.model import App, AppMode, EndUser
from repositories.factory import DifyAPIRepositoryFactory
from services.app_generate_service import AppGenerateService
from services.billing_service import BillingService
from services.errors.app import IsDraftWorkflowError, WorkflowIdFormatError, WorkflowNotFoundError
from services.errors.app import (
IsDraftWorkflowError,
WorkflowIdFormatError,
WorkflowNotFoundError,
)
from services.errors.app import (
TriggerWorkflowServiceModeUnavailableError as TriggerWorkflowServiceModeUnavailableServiceError,
)
from services.errors.llm import InvokeRateLimitError
from services.workflow_app_service import WorkflowAppService
@ -300,6 +308,11 @@ class WorkflowRunApi(Resource):
"- `completion_request_error` : Workflow execution request failed.\n"
"- `invalid_param` : Invalid parameter value."
),
403: (
"- `forbidden` : Token scope, app, or workspace access denied.\n"
"- `trigger_workflow_service_mode_unavailable` : Trigger-entry workflows cannot be invoked through "
"Web App, Service API, OpenAPI, or MCP."
),
429: (
"- `too_many_requests` : Too many concurrent requests for this app.\n"
"- `rate_limit_error` : The upstream model provider rate limit was exceeded."
@ -360,6 +373,8 @@ class WorkflowRunApi(Resource):
# response-contract:ignore compact_generate_response
return helper.compact_generate_response(response)
except TriggerWorkflowServiceModeUnavailableServiceError:
raise TriggerWorkflowServiceModeUnavailableError()
except ProviderTokenNotInitError as ex:
raise ProviderNotInitializeError(ex.description)
except QuotaExceededError:
@ -406,8 +421,11 @@ class WorkflowRunByIdApi(Resource):
"- `invalid_param` : Required parameter missing or invalid."
),
403: (
"`workflow_version_execution_not_allowed` : Workflow version execution is unavailable on the "
"current plan. Upgrade to a paid plan."
"- `forbidden` : Token scope, app, or workspace access denied.\n"
"- `workflow_version_execution_not_allowed` : Workflow version execution is unavailable on the "
"current plan. Upgrade to a paid plan.\n"
"- `trigger_workflow_service_mode_unavailable` : The selected workflow version uses a trigger entry "
"and cannot be invoked through Web App, Service API, OpenAPI, or MCP."
),
404: "`not_found` : Workflow not found.",
429: (
@ -487,6 +505,8 @@ class WorkflowRunByIdApi(Resource):
# response-contract:ignore compact_generate_response
return helper.compact_generate_response(response)
except TriggerWorkflowServiceModeUnavailableServiceError:
raise TriggerWorkflowServiceModeUnavailableError()
except WorkflowNotFoundError as ex:
raise NotFound(str(ex))
except IsDraftWorkflowError as ex:

View File

@ -25,7 +25,7 @@ from controllers.common.schema import (
register_schema_models,
)
from controllers.common.session import with_session
from controllers.console.wraps import edit_permission_required
from controllers.console.wraps import edit_permission_required, model_validate
from controllers.service_api import service_api_ns
from controllers.service_api.dataset.error import DatasetInUseError, DatasetNameDuplicateError, InvalidActionError
from controllers.service_api.wraps import (
@ -669,14 +669,13 @@ class DatasetApi(DatasetApiResource):
)
@cloud_edition_billing_rate_limit_check("knowledge", "dataset")
@with_session
def patch(self, session: Session, _, dataset_id: UUID):
@model_validate(DatasetUpdatePayload)
def patch(self, payload: DatasetUpdatePayload, session: Session, _, dataset_id: UUID):
dataset_id_str = str(dataset_id)
dataset = DatasetService.get_dataset(dataset_id_str, session)
if dataset is None:
raise NotFound("Dataset not found.")
payload_dict = service_api_ns.payload or {}
payload = DatasetUpdatePayload.model_validate(payload_dict)
update_data = payload.model_dump(exclude_unset=True)
if payload.permission is not None:
update_data["permission"] = str(payload.permission)
@ -944,13 +943,13 @@ class DatasetTagsApi(DatasetApiResource):
service_api_ns.models[KnowledgeTagResponse.__name__],
)
@with_session
def post(self, session: Session, _):
@model_validate(TagCreatePayload)
def post(self, payload: TagCreatePayload, session: Session, _):
"""Add a knowledge type tag."""
assert isinstance(current_user, Account)
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
raise Forbidden()
payload = TagCreatePayload.model_validate(service_api_ns.payload or {})
tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=TagType.KNOWLEDGE), session)
response = KnowledgeTagResponse(id=tag.id, name=tag.name, type=tag.type, binding_count="0")
@ -982,12 +981,12 @@ class DatasetTagsApi(DatasetApiResource):
service_api_ns.models[KnowledgeTagResponse.__name__],
)
@with_session
def patch(self, session: Session, _):
@model_validate(TagUpdatePayload)
def patch(self, payload: TagUpdatePayload, session: Session, _):
assert isinstance(current_user, Account)
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
raise Forbidden()
payload = TagUpdatePayload.model_validate(service_api_ns.payload or {})
tag_id = payload.tag_id
tag = TagService.update_tags(
UpdateTagServicePayload(name=payload.name), tag_id, session, tag_type=TagType.KNOWLEDGE
@ -1019,9 +1018,9 @@ class DatasetTagsApi(DatasetApiResource):
)
@edit_permission_required
@with_session
def delete(self, session: Session, _):
@model_validate(TagDeletePayload)
def delete(self, payload: TagDeletePayload, session: Session, _):
"""Delete a knowledge type tag."""
payload = TagDeletePayload.model_validate(service_api_ns.payload or {})
TagService.delete_tag(payload.tag_id, session, tag_type=TagType.KNOWLEDGE)
return "", 204
@ -1049,13 +1048,13 @@ class DatasetTagBindingApi(DatasetApiResource):
}
)
@with_session
def post(self, session: Session, _):
@model_validate(TagBindingPayload)
def post(self, payload: TagBindingPayload, session: Session, _):
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
assert isinstance(current_user, Account)
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
raise Forbidden()
payload = TagBindingPayload.model_validate(service_api_ns.payload or {})
TagService.save_tag_binding(
TagBindingCreatePayload(tag_ids=payload.tag_ids, target_id=payload.target_id, type=TagType.KNOWLEDGE),
session,
@ -1086,13 +1085,13 @@ class DatasetTagUnbindingApi(DatasetApiResource):
}
)
@with_session
def post(self, session: Session, _):
@model_validate(TagUnbindingPayload)
def post(self, payload: TagUnbindingPayload, session: Session, _):
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
assert isinstance(current_user, Account)
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
raise Forbidden()
payload = TagUnbindingPayload.model_validate(service_api_ns.payload or {})
TagService.delete_tag_binding(
TagBindingDeletePayload(tag_ids=payload.tag_ids, target_id=payload.target_id, type=TagType.KNOWLEDGE),
session,

View File

@ -45,6 +45,7 @@ from controllers.common.schema import (
register_schema_models,
)
from controllers.common.session import with_session
from controllers.console.wraps import model_validate
from controllers.service_api import service_api_ns
from controllers.service_api.app.error import ProviderNotInitializeError
from controllers.service_api.dataset.error import (
@ -1069,9 +1070,8 @@ class DocumentBatchDownloadZipApi(DatasetApiResource):
@service_api_ns.response(200, "ZIP archive generated successfully")
@cloud_edition_billing_rate_limit_check("knowledge", "dataset")
@with_session(write=False)
def post(self, session: Session, tenant_id, dataset_id: UUID):
payload = DocumentBatchDownloadZipPayload.model_validate(service_api_ns.payload or {})
@model_validate(DocumentBatchDownloadZipPayload)
def post(self, payload: DocumentBatchDownloadZipPayload, session: Session, tenant_id, dataset_id: UUID):
upload_files, download_name = DocumentService.prepare_document_batch_download_zip(
dataset_id=str(dataset_id),
document_ids=[str(document_id) for document_id in payload.document_ids],

View File

@ -316,15 +316,14 @@ class DocumentMetadataEditServiceApi(DatasetApiResource):
)
@cloud_edition_billing_rate_limit_check("knowledge", "dataset")
@with_session
def post(self, session: Session, tenant_id, dataset_id: UUID):
@model_validate(MetadataOperationData)
def post(self, metadata_args: MetadataOperationData, session: Session, tenant_id, dataset_id: UUID):
"""Update metadata for multiple documents."""
dataset = DatasetService.get_dataset_for_tenant(str(dataset_id), str(tenant_id), session=session)
if dataset is None:
raise NotFound("Dataset not found.")
DatasetService.check_dataset_permission(dataset, current_user, session)
metadata_args = MetadataOperationData.model_validate(service_api_ns.payload or {})
try:
MetadataService.update_documents_metadata(
dataset, metadata_args, cast(Account, current_user), session=session

View File

@ -24,6 +24,7 @@ from controllers.common.schema import (
register_schema_model,
)
from controllers.console.app.wraps import with_session
from controllers.console.wraps import model_validate
from controllers.service_api import service_api_ns
from controllers.service_api.dataset.error import PipelineRunError
from controllers.service_api.schema import event_stream_response, json_or_event_stream_response, multipart_file_params
@ -215,7 +216,8 @@ class DatasourceNodeRunApi(DatasetApiResource):
}
)
@service_api_ns.expect(service_api_ns.models[DatasourceNodeRunPayload.__name__])
def post(self, tenant_id: str, dataset_id: UUID, node_id: str):
@model_validate(DatasourceNodeRunPayload)
def post(self, payload: DatasourceNodeRunPayload, tenant_id: str, dataset_id: UUID, node_id: str):
"""Resource for getting datasource plugins."""
dataset_id_str = str(dataset_id)
# Verify dataset ownership
@ -224,7 +226,6 @@ class DatasourceNodeRunApi(DatasetApiResource):
if not dataset:
raise NotFound("Dataset not found.")
payload = DatasourceNodeRunPayload.model_validate(service_api_ns.payload or {})
assert isinstance(current_user, Account)
rag_pipeline_service: RagPipelineService = RagPipelineService(db.session())
pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str)

View File

@ -6,6 +6,7 @@ from werkzeug.exceptions import InternalServerError
import services
from controllers.common.controller_schemas import TextToAudioPayload as TextToAudioPayloadBase
from controllers.console.wraps import model_validate
from controllers.web import web_ns
from controllers.web.error import (
AppUnavailableError,
@ -131,11 +132,10 @@ class TextApi(WebApiResource):
)
# response-contract:ignore provider audio bytes; TODO: model binary audio response if shape is standardized.
@web_ns.response(200, "Success")
def post(self, app_model: App, end_user: EndUser):
@model_validate(TextToAudioPayload)
def post(self, payload: TextToAudioPayload, app_model: App, end_user: EndUser):
"""Convert text to audio"""
try:
payload = TextToAudioPayload.model_validate(web_ns.payload or {})
message_id = payload.message_id
text = payload.text
voice = payload.voice

View File

@ -8,6 +8,7 @@ from werkzeug.exceptions import NotFound
from controllers.common.controller_schemas import ConversationRenamePayload
from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models
from controllers.console.wraps import model_validate
from controllers.web import web_ns
from controllers.web.error import NotChatAppError
from controllers.web.wraps import WebApiResource
@ -153,15 +154,14 @@ class ConversationRenameApi(WebApiResource):
)
@web_ns.response(200, "Conversation renamed successfully", web_ns.models[SimpleConversation.__name__])
@web_ns.expect(web_ns.models[ConversationRenamePayload.__name__])
def post(self, app_model: App, end_user: EndUser, c_id: UUID):
@model_validate(ConversationRenamePayload)
def post(self, payload: ConversationRenamePayload, app_model: App, end_user: EndUser, c_id: UUID):
app_mode = AppMode.value_of(app_model.mode)
if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT}:
raise NotChatAppError()
conversation_id = str(c_id)
payload = ConversationRenamePayload.model_validate(web_ns.payload or {})
try:
session = db.session()
conversation = ConversationService.rename(

View File

@ -1,4 +1,8 @@
from libs.exception import BaseHTTPException
from services.errors.app import (
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_CODE,
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_MESSAGE,
)
class AppUnavailableError(BaseHTTPException):
@ -31,6 +35,12 @@ class NotWorkflowAppError(BaseHTTPException):
code = 400
class TriggerWorkflowServiceModeUnavailableError(BaseHTTPException):
error_code = TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_CODE
description = TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_MESSAGE
code = 403
class ConversationCompletedError(BaseHTTPException):
error_code = "conversation_completed"
description = "The conversation has ended. Please start a new conversation."

View File

@ -9,6 +9,7 @@ from controllers.common.errors import (
RemoteFileUploadError,
UnsupportedFileTypeError,
)
from controllers.console.wraps import model_validate
from core.file import remote_fetcher
from extensions.ext_database import db
from fields.file_fields import FileWithSignedUrl, RemoteFileInfo
@ -86,7 +87,8 @@ class RemoteFileUploadApi(WebApiResource):
)
@web_ns.response(201, "Remote file uploaded", web_ns.models[FileWithSignedUrl.__name__])
@web_ns.expect(web_ns.models[RemoteFileUploadPayload.__name__])
def post(self, app_model: App, end_user: EndUser):
@model_validate(RemoteFileUploadPayload)
def post(self, payload: RemoteFileUploadPayload, app_model: App, end_user: EndUser):
"""Upload a file from a remote URL.
Downloads a file from the provided remote URL and uploads it
@ -108,7 +110,6 @@ class RemoteFileUploadApi(WebApiResource):
FileTooLargeError: File exceeds size limit
UnsupportedFileTypeError: File type not supported
"""
payload = RemoteFileUploadPayload.model_validate(web_ns.payload or {})
url = str(payload.url)
try:

View File

@ -14,6 +14,7 @@ from controllers.web.error import (
ProviderModelCurrentlyNotSupportError,
ProviderNotInitializeError,
ProviderQuotaExceededError,
TriggerWorkflowServiceModeUnavailableError,
)
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
from controllers.web.wraps import WebApiResource
@ -30,6 +31,9 @@ from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper
from models.model import App, AppMode, EndUser
from services.app_generate_service import AppGenerateService
from services.errors.app import (
TriggerWorkflowServiceModeUnavailableError as TriggerWorkflowServiceModeUnavailableServiceError,
)
from services.errors.llm import InvokeRateLimitError
logger = logging.getLogger(__name__)
@ -78,6 +82,8 @@ class WorkflowRunApi(WebApiResource):
# response-contract:ignore compact_generate_response
return helper.compact_generate_response(response)
except TriggerWorkflowServiceModeUnavailableServiceError:
raise TriggerWorkflowServiceModeUnavailableError()
except ProviderTokenNotInitError as ex:
raise ProviderNotInitializeError(ex.description)
except QuotaExceededError:

View File

@ -31,7 +31,11 @@ from core.agent.publish_visibility import agent_has_workflow_callable_active_sna
from core.app.app_config.easy_ui_based_app.model_config.converter import ModelConfigConverter
from core.app.apps.agent_app.app_config_manager import AgentAppConfigManager
from core.app.apps.agent_app.app_runner import AgentAppRunner
from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError
from core.app.apps.agent_app.errors import (
AgentAppGeneratorError,
AgentAppNotPublishedError,
AgentSessionSnapshotIncompatibleError,
)
from core.app.apps.agent_app.generate_response_converter import AgentAppGenerateResponseConverter
from core.app.apps.agent_app.runtime_request_builder import AgentAppRuntimeRequestBuilder
from core.app.apps.agent_app.session_store import AgentAppWorkspaceStore
@ -531,6 +535,15 @@ class AgentAppGenerator(MessageBasedAppGenerator):
)
except GenerateTaskStoppedError:
pass
except AgentSessionSnapshotIncompatibleError as error:
logger.info(
"Agent App session snapshot no longer matches the current composition",
extra={
"agent_id": application_generate_entity.agent_id,
"conversation_id": conversation_id,
},
)
queue_manager.publish_error(error, PublishFrom.APPLICATION_MANAGER)
except Exception as e:
logger.exception("Unknown Error in Agent App generate worker")
queue_manager.publish_error(e, PublishFrom.APPLICATION_MANAGER)

View File

@ -1,6 +1,24 @@
from core.app.apps.exc import AppGenerateError
AGENT_SESSION_CONFIGURATION_CHANGED_ERROR_CODE = "agent_session_configuration_changed"
AGENT_SESSION_CONFIGURATION_CHANGED_MESSAGE = (
"The Agent configuration changed after this conversation started. Start a new conversation to continue."
)
class AgentAppGeneratorError(ValueError):
"""Raised when an Agent App turn cannot be set up."""
class AgentAppNotPublishedError(AgentAppGeneratorError):
"""Raised when a public Agent App runtime is requested before publish."""
class AgentSessionSnapshotIncompatibleError(AppGenerateError):
"""Raised when a retained session snapshot no longer matches the current composition."""
error_code = AGENT_SESSION_CONFIGURATION_CHANGED_ERROR_CODE
status_code = 409
def __init__(self) -> None:
super().__init__(AGENT_SESSION_CONFIGURATION_CHANGED_MESSAGE)

View File

@ -50,6 +50,8 @@ from models.agent_config_entities import AgentSoulConfig, AgentSoulToolsConfig
from models.provider_ids import ModelProviderID
from services.agent.prompt_mentions import expand_prompt_mentions
from .errors import AgentSessionSnapshotIncompatibleError
class AgentAppRuntimeRequestBuildError(ValueError):
"""Raised when Agent App state cannot be mapped to a valid run request."""
@ -191,6 +193,7 @@ class AgentAppRuntimeRequestBuilder:
metadata=metadata,
)
)
self._validate_session_snapshot_layers(request)
redacted = cast(dict[str, Any], redact_for_agent_backend_log(request))
return AgentAppRuntimeRequest(
request=request,
@ -199,6 +202,24 @@ class AgentAppRuntimeRequestBuilder:
binding_id=context.binding_id,
)
@staticmethod
def _validate_session_snapshot_layers(request: CreateRunRequest) -> None:
"""Reject stale snapshots before they reach the Agent backend.
Draft rows are updated in place, so their IDs cannot prove that a
retained snapshot still belongs to the current composition. Agenton
requires the ordered layer names to match exactly; enforce the same
invariant at the API boundary and return a product-level error.
"""
snapshot = request.session_snapshot
if snapshot is None:
return
snapshot_layer_names = tuple(layer.name for layer in snapshot.layers)
composition_layer_names = tuple(layer.name for layer in request.composition.layers)
if snapshot_layer_names != composition_layer_names:
raise AgentSessionSnapshotIncompatibleError()
def _build_tool_layers(
self,
*,

View File

@ -7,6 +7,7 @@ from dify_agent.protocol import RunFailureType
from pydantic import JsonValue
from clients.agent_backend.errors import AgentBackendError, AgentBackendRunFailedError
from core.app.apps.exc import AppGenerateError
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.task_entities import AppBlockingResponse, AppStreamResponse
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
@ -125,6 +126,13 @@ class AppGenerateResponseConverter[TBlockingResponse: AppBlockingResponse](ABC):
"message": str(e),
}
if isinstance(e, AppGenerateError):
return {
"code": e.error_code,
"status": e.status_code,
"message": str(e),
}
error_responses: dict[type[Exception], dict[str, JsonValue]] = {
ValueError: {"code": "invalid_param", "status": 400},
ProviderTokenNotInitError: {"code": "provider_not_initialize", "status": 400},

View File

@ -1,2 +1,9 @@
class AppGenerateError(ValueError):
"""Base class for application-generation errors with a stable response contract."""
error_code: str
status_code: int
class GenerateTaskStoppedError(Exception):
pass

View File

@ -841,9 +841,8 @@ class LLMGenerator:
model_config: ModelConfig,
ideal_output: str | None,
workflow_service: WorkflowServiceInterface,
session: Session,
):
session = db.session()
app: App | None = session.scalar(select(App).where(App.id == flow_id, App.tenant_id == tenant_id).limit(1))
if not app:
raise ValueError("App not found.")

View File

@ -12,6 +12,7 @@ from core.mcp import types as mcp_types
from graphon.variables.input_entities import VariableEntity, VariableEntityType
from models.model import App, AppMCPServer, AppMode, EndUser
from services.app_generate_service import AppGenerateService
from services.errors.app import TriggerWorkflowServiceModeUnavailableError
logger = logging.getLogger(__name__)
@ -93,11 +94,16 @@ def handle_mcp_request(
result=result_data.model_dump(by_alias=True, mode="json", exclude_none=True),
)
def create_error_response(code: int, message: str) -> mcp_types.JSONRPCError:
def create_error_response(
code: int,
message: str,
*,
data: Mapping[str, Any] | None = None,
) -> mcp_types.JSONRPCError:
"""Create error response with error code and message"""
from core.mcp.types import ErrorData
error_data = ErrorData(code=code, message=message)
error_data = ErrorData(code=code, message=message, data=data)
return mcp_types.JSONRPCError(
jsonrpc="2.0",
id=request_id,
@ -131,6 +137,12 @@ def handle_mcp_request(
case _:
return create_error_response(mcp_types.METHOD_NOT_FOUND, f"Method not found: {request_type.__name__}")
except TriggerWorkflowServiceModeUnavailableError as e:
return create_error_response(
mcp_types.INVALID_REQUEST,
str(e),
data={"code": e.error_code},
)
except ValueError as e:
logger.exception("Invalid params")
return create_error_response(mcp_types.INVALID_PARAMS, str(e))

View File

@ -1193,7 +1193,7 @@ class WorkflowGenerator:
if node.get("node_type") == BuiltinNodeTypes.TOOL and node.get("id")
}
for node in graph.get("nodes") or []:
planned = planned_by_id.get(str(node.get("id") or ""))
planned = planned_by_id.get(node.get("id") or "")
if planned is None:
continue
data = node.get("data")

View File

@ -76,6 +76,10 @@ def _schema_markdown_type(schema: object) -> str:
item_type = _schema_markdown_type(schema.get("items"))
return f"[ {item_type or 'object'} ]"
if isinstance(schema_type, str):
enum_values = schema.get("enum")
if isinstance(enum_values, list) and enum_values:
rendered_values = ", ".join(json.dumps(value, ensure_ascii=False) for value in enum_values)
return f"{schema_type}, <br>**Available values:** {rendered_values}"
return schema_type
return ""

View File

@ -31,6 +31,7 @@ from repositories.factory import DifyAPIRepositoryFactory
from repositories.installation_state_repository import InstallationStateRepository
from repositories.oauth_server_repository import RedisOAuthServerTokenRepository, SQLAlchemyOAuthServerRepository
from repositories.recommended_app_catalog_repository import DatabaseRecommendedAppCatalogRepository
from repositories.step_by_step_tour_repository import SQLAlchemyStepByStepTourStateRepository
from repositories.tag_repository import TagRepository
from repositories.trial_app_query_repository import TrialAppQueryRepository
from repositories.trial_app_usage_repository import TrialAppUsageRepository
@ -70,6 +71,16 @@ from services.account_deletion_adapters import (
from services.account_deletion_feedback_service import AccountDeletionFeedbackService
from services.account_deletion_service import AccountDeletionService
from services.account_education_service import AccountEducationService
from services.account_email_registration_adapters import (
AccountServiceRegistrationGateway,
BillingAccountRegistrationPolicyGateway,
CeleryEmailRegistrationNotificationGateway,
RateLimiterEmailRegistrationSendLimiter,
RedisEmailRegistrationSecurityGateway,
SecureEmailRegistrationCodeGenerator,
TokenManagerEmailRegistrationTokenGateway,
)
from services.account_email_registration_service import AccountEmailRegistrationService
from services.account_initialization_service import AccountInitializationService
from services.account_integration_service import AccountIntegrationService
from services.account_password_hasher import LegacyAccountPasswordHasher
@ -94,6 +105,8 @@ from services.feature_service import FeatureService
from services.feature_service_gateway import FeatureServiceGateway
from services.file_service import FileService
from services.init_validation_service import InitValidationService
from services.notification_gateway import BillingNotificationGateway
from services.notification_service import NotificationService
from services.notion_data_source_gateway import NotionDataSourceGateway
from services.oauth_server_service import OAUTH_ACCESS_TOKEN_EXPIRES_IN, OAuthServerService
from services.partner_tenant_binding_service import PartnerTenantBindingService
@ -112,6 +125,7 @@ from services.retention.workflow_run.archive_log_service import WorkflowRunArchi
from services.schema_definition_service import SchemaDefinitionService
from services.setup_adapters import RedisSetupLock, RegisterServiceAccountProvisioner
from services.setup_service import SetupService
from services.step_by_step_tour_service import StepByStepTourService
from services.tag_application_service import TagApplicationService
from services.trial_app_usage import TrialAppUsageRecorder
from services.web_app_runtime_query_service import WebAppRuntimeQueryService
@ -150,6 +164,7 @@ def _is_user_allowed_to_access_webapp(user_id: str, app_id: str) -> bool:
class AccountServices:
avatar: AccountAvatarService
change_email: AccountChangeEmailService
email_registration: AccountEmailRegistrationService
deletion: AccountDeletionService
deletion_feedback: AccountDeletionFeedbackService
education: AccountEducationService
@ -177,6 +192,8 @@ class ApplicationServices:
feature_queries: FeatureQueryService
oauth_server: OAuthServerService
init_validation: InitValidationService
notifications: NotificationService
step_by_step_tour: StepByStepTourService
partner_tenant_bindings: PartnerTenantBindingService
recommended_app_queries: RecommendedAppQueryService
trial_app_usage: TrialAppUsageRecorder
@ -278,6 +295,29 @@ def build_application_services(
billing_enabled=deployment_edition == DeploymentEdition.CLOUD,
),
),
email_registration=AccountEmailRegistrationService(
accounts=accounts,
tokens=TokenManagerEmailRegistrationTokenGateway(),
codes=SecureEmailRegistrationCodeGenerator(),
notifications=CeleryEmailRegistrationNotificationGateway(),
send_limits=RateLimiterEmailRegistrationSendLimiter(
rate_limiter=RateLimiter(
prefix="email_register_rate_limit",
max_attempts=1,
time_window=60,
redis_client=redis,
)
),
security=RedisEmailRegistrationSecurityGateway(
redis=redis,
verification_failure_limit=5,
verification_lockout_duration=dify_config.EMAIL_REGISTER_LOCKOUT_DURATION,
),
account_policy=BillingAccountRegistrationPolicyGateway(
enabled=deployment_edition == DeploymentEdition.CLOUD,
),
registration=AccountServiceRegistrationGateway(session_factory=database_client),
),
deletion=AccountDeletionService(
accounts=accounts,
memberships=workspace_query_repository,
@ -400,6 +440,16 @@ def build_application_services(
validation_required=(deployment_edition != DeploymentEdition.CLOUD and bool(initialization_password)),
expected_password=initialization_password,
),
notifications=NotificationService(
accounts=accounts,
notifications=BillingNotificationGateway(),
),
step_by_step_tour=StepByStepTourService(
accounts=accounts,
states=SQLAlchemyStepByStepTourStateRepository(session_factory=database_client),
enabled=dify_config.ENABLE_STEP_BY_STEP_TOUR,
rollout_started_at=dify_config.STEP_BY_STEP_TOUR_ROLLOUT_STARTED_AT,
),
partner_tenant_bindings=PartnerTenantBindingService(
sync_bindings=BillingService.sync_partner_tenants_bindings,
),

View File

@ -297,18 +297,16 @@ class Workflow(Base): # bug
workflow.updated_at = workflow.created_at
return workflow
@property
def created_by_account(self) -> Account | None:
return self.get_created_by_account(session=db.session())
def created_by_account(self, session: orm.Session) -> Account | None:
return self.get_created_by_account(session=session)
def get_created_by_account(self, *, session: orm.Session) -> Account | None:
def get_created_by_account(self, session: orm.Session) -> Account | None:
return session.get(Account, self.created_by)
@property
def updated_by_account(self) -> Account | None:
return self.get_updated_by_account(session=db.session())
def updated_by_account(self, session: orm.Session) -> Account | None:
return self.get_updated_by_account(session=session)
def get_updated_by_account(self, *, session: orm.Session) -> Account | None:
def get_updated_by_account(self, session: orm.Session) -> Account | None:
return session.get(Account, self.updated_by) if self.updated_by else None
@property
@ -564,18 +562,17 @@ class Workflow(Base): # bug
return helper.generate_text_hash(json.dumps(entity, sort_keys=True))
@property
@deprecated(
"This property is not accurate for determining if a workflow is published as a tool."
"This method is not accurate for determining if a workflow is published as a tool."
"It only checks if there's a WorkflowToolProvider for the app, "
"not if this specific workflow version is the one being used by the tool."
)
def tool_published(self) -> bool:
return self.get_tool_published(session=db.session())
def tool_published(self, session: orm.Session) -> bool:
return self.get_tool_published(session=session)
def get_tool_published(self, *, session: orm.Session) -> bool:
def get_tool_published(self, session: orm.Session) -> bool:
"""
DEPRECATED: This property is not accurate for determining if a workflow is published as a tool.
DEPRECATED: This method is not accurate for determining if a workflow is published as a tool.
It only checks if there's a WorkflowToolProvider for the app, not if this specific workflow version
is the one being used by the tool.

View File

@ -13501,7 +13501,7 @@ default (the config form sends the full desired feature state on save).
| mode | string, <br>**Available values:** "advanced-chat", "agent", "agent-chat", "all", "channel", "chat", "completion", "workflow", <br>**Default:** all | App mode filter<br>*Enum:* `"advanced-chat"`, `"agent"`, `"agent-chat"`, `"all"`, `"channel"`, `"chat"`, `"completion"`, `"workflow"` | No |
| name | string | Filter by app name | No |
| page | integer, <br>**Default:** 1 | Page number (1-99999) | No |
| publication_status | string | Filter by published or draft Agent configuration status | No |
| publication_status | string, <br>**Available values:** "drafts", "published" | Filter by published or draft Agent configuration status | No |
| sort_by | string, <br>**Available values:** "earliest_created", "last_modified", "recently_created", <br>**Default:** last_modified | Sort apps by last modified, recently created, or earliest created<br>*Enum:* `"earliest_created"`, `"last_modified"`, `"recently_created"` | No |
| tag_ids | [ string ] | Filter by tag IDs | No |
@ -15744,7 +15744,7 @@ AppMCPServer Status Enum
| copyright | string | | No |
| custom_disclaimer | string | | No |
| customize_domain | string | | No |
| customize_token_strategy | string | | No |
| customize_token_strategy | string, <br>**Available values:** "allow", "must", "not_allow" | | No |
| default_language | string | | No |
| description | string | | No |
| icon | string | | No |
@ -16202,7 +16202,7 @@ TEAM: Team collaboration paid plan
| files | [ object ] | | No |
| inputs | object | | Yes |
| query | string | | No |
| response_mode | string | | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | | No |
| retriever_from | string, <br>**Default:** explore_app | | No |
#### CompletionMessagePayload
@ -16223,7 +16223,7 @@ TEAM: Team collaboration paid plan
| files | [ object ] | | No |
| inputs | object | | Yes |
| query | string | | No |
| response_mode | string | | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | | No |
| retriever_from | string, <br>**Default:** explore_app | | No |
#### ComplianceDownloadQuery
@ -18263,9 +18263,9 @@ Flask blueprint initialization.
| ---- | ---- | ----------- | -------- |
| end_date | string | End date (YYYY-MM-DD) | No |
| format | string, <br>**Available values:** "csv", "json", <br>**Default:** csv | Export format<br>*Enum:* `"csv"`, `"json"` | No |
| from_source | string | Filter by feedback source | No |
| from_source | string, <br>**Available values:** "admin", "user" | Filter by feedback source | No |
| has_comment | boolean | Only include feedback with comments | No |
| rating | string | Filter by rating | No |
| rating | string, <br>**Available values:** "dislike", "like" | Filter by rating | No |
| start_date | string | Start date (YYYY-MM-DD) | No |
#### FeedbackStat
@ -18663,7 +18663,7 @@ Icon information model.
| ---- | ---- | ----------- | -------- |
| icon | string | | No |
| icon_background | string | | No |
| icon_type | string | | No |
| icon_type | string, <br>**Available values:** "emoji", "image" | | No |
| icon_url | string | | No |
#### IconType
@ -19245,7 +19245,7 @@ Enum class for large language model mode.
| ---- | ---- | ----------- | -------- |
| content | string | Optional text feedback providing additional detail. | No |
| message_id | string | Message ID | Yes |
| rating | string | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
| rating | string, <br>**Available values:** "dislike", "like" | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
#### MessageFile
@ -19306,7 +19306,7 @@ Metadata Filtering Condition.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| conditions | [ [Condition](#condition) ] | List of metadata conditions to evaluate. | No |
| logical_operator | string | How to combine multiple conditions. | No |
| logical_operator | string, <br>**Available values:** "and", "or" | How to combine multiple conditions. | No |
#### MetadataOperationData
@ -19439,7 +19439,7 @@ Enum class for model property key.
| is_exhausted | boolean | | Yes |
| is_unlimited | boolean | | Yes |
| next_credit_reset_date | integer | | Yes |
| pool_type | string | | Yes |
| pool_type | string, <br>**Available values:** "paid", "trial" | | Yes |
| quota_limit | integer | Credit limit for the effective pool; -1 means unlimited. | Yes |
| quota_used | integer | | Yes |
| remaining_credits | integer | Remaining credits; -1 means unlimited. | Yes |
@ -21434,7 +21434,7 @@ Model class for provider quota configuration.
| ---- | ---- | ----------- | -------- |
| metadata_filtering_conditions | [MetadataFilteringCondition](#metadatafilteringcondition) | Restrict retrieval to chunks whose document metadata matches the given conditions. Conditions are evaluated server-side against document metadata fields. | No |
| reranking_enable | boolean | Whether reranking is enabled. | Yes |
| reranking_mode | string | Reranking mode. Required when `reranking_enable` is `true`. | No |
| reranking_mode | string, <br>**Available values:** "reranking_model", "weighted_score" | Reranking mode. Required when `reranking_enable` is `true`. | No |
| reranking_model | [RerankingModel](#rerankingmodel) | Reranking model configuration. | No |
| score_threshold | number | Minimum similarity score for results. Only effective when score threshold filtering is enabled. | No |
| score_threshold_enabled | boolean | Whether score threshold filtering is enabled. | Yes |
@ -21488,7 +21488,7 @@ Model class for provider quota configuration.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| parent_mode | string | Parent-child segmentation mode. | No |
| parent_mode | string, <br>**Available values:** "full-doc", "paragraph" | Parent-child segmentation mode. | No |
| pre_processing_rules | [ [PreProcessingRule](#preprocessingrule) ] | Pre-processing rules to apply before segmentation. | No |
| segmentation | [Segmentation](#segmentation) | Parent chunk segmentation settings. | No |
| subchunk_segmentation | [Segmentation](#segmentation) | Child chunk segmentation settings. | No |
@ -22477,7 +22477,7 @@ Query parameters for listing snippet published workflows.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| action | string, <br>**Available values:** "complete_task", "disable_current_workspace", "enable_current_workspace", "skip", "uncomplete_task" | State update action<br>*Enum:* `"complete_task"`, `"disable_current_workspace"`, `"enable_current_workspace"`, `"skip"`, `"uncomplete_task"` | Yes |
| task_id | string | Task ID for task actions | No |
| task_id | string, <br>**Available values:** "home", "integration", "knowledge", "studio" | Task ID for task actions | No |
#### StepByStepTourStateResponse
@ -22943,7 +22943,7 @@ Tool label
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| visibility | string | Visibility for the OAuth credential. Defaults to 'only_me'. | No |
| visibility | string, <br>**Available values:** "all_team_members", "only_me" | Visibility for the OAuth credential. Defaults to 'only_me'. | No |
#### ToolOAuthCustomClientPayload
@ -23075,7 +23075,7 @@ removes TOOLS_SELECTOR from PluginParameterType
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| type | string | | No |
| type | string, <br>**Available values:** "api", "builtin", "mcp", "model", "workflow" | | No |
#### ToolProviderListResponse
@ -23693,7 +23693,7 @@ in form definition, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| keyword_setting | [WeightKeywordSetting](#weightkeywordsetting) | Keyword search weight settings. | No |
| vector_setting | [WeightVectorSetting](#weightvectorsetting) | Semantic search weight settings. | No |
| weight_type | string | Strategy for balancing semantic and keyword search weights. | No |
| weight_type | string, <br>**Available values:** "customized", "keyword_first", "semantic_first" | Strategy for balancing semantic and keyword search weights. | No |
#### WeightVectorSetting
@ -24199,7 +24199,7 @@ can reuse its existing handler.
| description | string | | No |
| event | string | | No |
| icon | string | | No |
| mode | string | *Enum:* `"advanced-chat"`, `"workflow"` | Yes |
| mode | string, <br>**Available values:** "advanced-chat", "workflow" | *Enum:* `"advanced-chat"`, `"workflow"` | Yes |
| nodes | [ [WorkflowPlanNodeResponse](#workflowplannoderesponse) ] | | Yes |
| start_inputs | [ [WorkflowPlanStartInputResponse](#workflowplanstartinputresponse) ] | | No |
| title | string | | No |
@ -24214,7 +24214,7 @@ can reuse its existing handler.
| graph | [WorkflowGraph](#workflowgraph) | | Yes |
| icon | string | | No |
| message | string | | No |
| mode | string | | No |
| mode | string, <br>**Available values:** "advanced-chat", "workflow" | | No |
#### WorkflowGenerateResultEventResponse
@ -24227,7 +24227,7 @@ can reuse its existing handler.
| graph | [WorkflowGraph](#workflowgraph) | | Yes |
| icon | string | | No |
| message | string | | No |
| mode | string | | No |
| mode | string, <br>**Available values:** "advanced-chat", "workflow" | | No |
#### WorkflowGenerateStreamEventResponse
@ -24527,9 +24527,9 @@ Lifecycle state for an asynchronous archive download request.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| status | string | Workflow run status filter | No |
| status | string, <br>**Available values:** "failed", "partial-succeeded", "running", "stopped", "succeeded" | Workflow run status filter | No |
| time_range | string | Filter by time range (optional): e.g., 7d (7 days), 4h (4 hours), 30m (30 minutes), 30s (30 seconds). Filters by created_at field. | No |
| triggered_from | string | Filter by trigger source: debugging or app-run. Default: debugging | No |
| triggered_from | string, <br>**Available values:** "app-run", "debugging" | Filter by trigger source: debugging or app-run. Default: debugging | No |
#### WorkflowRunCountResponse
@ -24601,8 +24601,8 @@ Lifecycle state for an asynchronous archive download request.
| ---- | ---- | ----------- | -------- |
| last_id | string | Last run ID for pagination | No |
| limit | integer, <br>**Default:** 20 | Number of items per page (1-100) | No |
| status | string | Workflow run status filter | No |
| triggered_from | string | Filter by trigger source: debugging or app-run. Default: debugging | No |
| status | string, <br>**Available values:** "failed", "partial-succeeded", "running", "stopped", "succeeded" | Workflow run status filter | No |
| triggered_from | string, <br>**Available values:** "app-run", "debugging" | Filter by trigger source: debugging or app-run. Default: debugging | No |
#### WorkflowRunNodeExecutionListResponse
@ -24900,7 +24900,7 @@ Workflow tool configuration
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| language | string | Localized policy label language | No |
| language | string, <br>**Available values:** "en", "ja", "zh" | Localized policy label language | No |
#### _AccessPolicyList
@ -24959,7 +24959,7 @@ Workflow tool configuration
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| language | string | Localized policy label language | No |
| language | string, <br>**Available values:** "en", "ja", "zh" | Localized policy label language | No |
| limit | integer | | No |
| page | integer | | No |
| reverse | boolean | | No |

View File

@ -2211,7 +2211,7 @@ Execute a workflow. Cannot be executed without a published workflow.
| 200 | Successful response. The content type and structure depend on the `response_mode` parameter in the request. - If `response_mode` is `blocking`, returns `application/json` with a `WorkflowBlockingResponse` object. - If `response_mode` is `streaming`, returns `text/event-stream` with a stream of `ChunkWorkflowEvent` objects. | **application/json**: [WorkflowBlockingResponse](#workflowblockingresponse)<br>**text/event-stream**: string<br> |
| 400 | - `not_workflow_app` : App mode does not match the API route. - `provider_not_initialize` : No valid model provider credentials found. - `provider_quota_exceeded` : Model provider quota exhausted. - `model_currently_not_support` : Current model unavailable. - `completion_request_error` : Workflow execution request failed. - `invalid_param` : Invalid parameter value. | |
| 401 | Unauthorized - invalid API token | |
| 403 | Forbidden - token scope, app, dataset, or workspace access denied | |
| 403 | - `forbidden` : Token scope, app, or workspace access denied. - `trigger_workflow_service_mode_unavailable` : Trigger-entry workflows cannot be invoked through Web App, Service API, OpenAPI, or MCP. | |
| 404 | Workflow not found | |
| 429 | - `too_many_requests` : Too many concurrent requests for this app. - `rate_limit_error` : The upstream model provider rate limit was exceeded. | |
| 500 | `internal_server_error` : Internal server error. | |
@ -2287,7 +2287,7 @@ Execute a specific workflow version identified by its ID. Useful for running a p
| 200 | Successful response. The content type and structure depend on the `response_mode` parameter in the request. - If `response_mode` is `blocking`, returns `application/json` with a `WorkflowBlockingResponse` object. - If `response_mode` is `streaming`, returns `text/event-stream` with a stream of `ChunkWorkflowEvent` objects. | **application/json**: [WorkflowBlockingResponse](#workflowblockingresponse)<br>**text/event-stream**: string<br> |
| 400 | - `not_workflow_app` : App mode does not match the API route. - `bad_request` : Workflow is a draft or has an invalid ID format. - `provider_not_initialize` : No valid model provider credentials found. - `provider_quota_exceeded` : Model provider quota exhausted. - `model_currently_not_support` : Current model unavailable. - `completion_request_error` : Workflow execution request failed. - `invalid_param` : Required parameter missing or invalid. | |
| 401 | Unauthorized - invalid API token | |
| 403 | `workflow_version_execution_not_allowed` : Workflow version execution is unavailable on the current plan. Upgrade to a paid plan. | |
| 403 | - `forbidden` : Token scope, app, or workspace access denied. - `workflow_version_execution_not_allowed` : Workflow version execution is unavailable on the current plan. Upgrade to a paid plan. - `trigger_workflow_service_mode_unavailable` : The selected workflow version uses a trigger entry and cannot be invoked through Web App, Service API, OpenAPI, or MCP. | |
| 404 | `not_found` : Workflow not found. | |
| 429 | - `too_many_requests` : Too many concurrent requests for this app. - `rate_limit_error` : The upstream model provider rate limit was exceeded. | |
| 500 | `internal_server_error` : Internal server error. | |
@ -2587,7 +2587,7 @@ Public pause reason emitted by a blocking Chatflow execution.
| files | [ object<br>object<br>object<br>object ] | File list for multimodal understanding, including images, documents, audio, and video. To attach a local file, first upload it via [Upload File](/api-reference/files/upload-file) and use the returned `id` as `upload_file_id` with `transfer_method: local_file`. | No |
| inputs | object | Values for app-defined variables. Refer to the `user_input_form` field in the [Get App Parameters](/api-reference/applications/get-app-parameters) response to discover expected variable names and types. | Yes |
| query | string | User input or question content. | Yes |
| response_mode | string | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. New Agent app mode supports streaming only. When omitted, non-Agent apps run in blocking mode and new Agent apps stream. | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. New Agent app mode supports streaming only. When omitted, non-Agent apps run in blocking mode and new Agent apps stream. | No |
| workflow_id | string | Published workflow version ID to execute for advanced chat. If omitted, the app's current published workflow is used. | No |
#### ChatRequestPayloadWithUser
@ -2599,7 +2599,7 @@ Public pause reason emitted by a blocking Chatflow execution.
| files | [ object<br>object<br>object<br>object ] | File list for multimodal understanding, including images, documents, audio, and video. To attach a local file, first upload it via [Upload File](/api-reference/files/upload-file) and use the returned `id` as `upload_file_id` with `transfer_method: local_file`. | No |
| inputs | object | Values for app-defined variables. Refer to the `user_input_form` field in the [Get App Parameters](/api-reference/applications/get-app-parameters) response to discover expected variable names and types. | Yes |
| query | string | User input or question content. | Yes |
| response_mode | string | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. New Agent app mode supports streaming only. When omitted, non-Agent apps run in blocking mode and new Agent apps stream. | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. New Agent app mode supports streaming only. When omitted, non-Agent apps run in blocking mode and new Agent apps stream. | No |
| user | string | User identifier, unique within the application. This identifier scopes data access; resources created with one `user` value are only visible when queried with the same `user` value. | Yes |
| workflow_id | string | Published workflow version ID to execute for advanced chat. If omitted, the app's current published workflow is used. | No |
@ -2672,7 +2672,7 @@ Public pause reason emitted by a blocking Chatflow execution.
| files | [ object<br>object<br>object<br>object ] | File list for multimodal understanding, including images, documents, audio, and video. To attach a local file, first upload it via [Upload File](/api-reference/files/upload-file) and use the returned `id` as `upload_file_id` with `transfer_method: local_file`. | No |
| inputs | object | Values for app-defined variables. Refer to the `user_input_form` field in the [Get App Parameters](/api-reference/applications/get-app-parameters) response to discover expected variable names and types. | Yes |
| query | string | User input or prompt content. | No |
| response_mode | string | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. When omitted, the request runs in blocking mode. | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. When omitted, the request runs in blocking mode. | No |
#### CompletionRequestPayloadWithUser
@ -2681,7 +2681,7 @@ Public pause reason emitted by a blocking Chatflow execution.
| files | [ object<br>object<br>object<br>object ] | File list for multimodal understanding, including images, documents, audio, and video. To attach a local file, first upload it via [Upload File](/api-reference/files/upload-file) and use the returned `id` as `upload_file_id` with `transfer_method: local_file`. | No |
| inputs | object | Values for app-defined variables. Refer to the `user_input_form` field in the [Get App Parameters](/api-reference/applications/get-app-parameters) response to discover expected variable names and types. | Yes |
| query | string | User input or prompt content. | No |
| response_mode | string | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. When omitted, the request runs in blocking mode. | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode. `streaming` uses Server-Sent Events; `blocking` returns after completion. When omitted, the request runs in blocking mode. | No |
| user | string | User identifier, unique within the application. This identifier scopes data access; resources created with one `user` value are only visible when queried with the same `user` value. | Yes |
#### Condition
@ -2797,7 +2797,7 @@ Enum class for custom configuration status.
| embedding_model_provider | string | Embedding model provider. Use the `provider` field from [Get Available Models](/api-reference/models/get-available-models) with `model_type=text-embedding`. | No |
| external_knowledge_api_id | string | ID of the external knowledge API. | No |
| external_knowledge_id | string | ID of the external knowledge base. | No |
| indexing_technique | string | `high_quality` uses embedding models for precise search; `economy` uses keyword-based indexing. | No |
| indexing_technique | string, <br>**Available values:** "economy", "high_quality" | `high_quality` uses embedding models for precise search; `economy` uses keyword-based indexing. | No |
| name | string | Name of the knowledge base. | Yes |
| permission | [PermissionEnum](#permissionenum) | Controls who can access this knowledge base. `only_me` restricts access to the creator, `all_team_members` grants workspace-wide access, and `partial_members` grants access to specified members. | No |
| provider | string, <br>**Available values:** "external", "vendor", <br>**Default:** vendor | Knowledge base provider: `vendor` for internal knowledge bases, `external` for external ones.<br>*Enum:* `"external"`, `"vendor"` | No |
@ -3039,7 +3039,7 @@ Enum class for custom configuration status.
| external_knowledge_api_id | string | ID of the external knowledge API. | No |
| external_knowledge_id | string | ID of the external knowledge base. | No |
| external_retrieval_model | object | Retrieval settings for external knowledge bases. | No |
| indexing_technique | string | `high_quality` uses embedding models for precise search; `economy` uses keyword-based indexing. | No |
| indexing_technique | string, <br>**Available values:** "economy", "high_quality" | `high_quality` uses embedding models for precise search; `economy` uses keyword-based indexing. | No |
| name | string | Name of the knowledge base. | No |
| partial_member_list | [ object ] | List of team members with access when `permission` is `partial_members`. | No |
| permission | [PermissionEnum](#permissionenum) | Controls who can access this knowledge base. `only_me` restricts access to the creator, `all_team_members` grants workspace-wide access, and `partial_members` grants access to specified members. | No |
@ -3167,7 +3167,7 @@ Request payload for bulk downloading documents as a zip archive.
| keyword | string | Search keyword to filter by document name. | No |
| limit | integer, <br>**Default:** 20 | Number of items per page. Server caps at `100`. | No |
| page | integer, <br>**Default:** 1 | Page number to retrieve. | No |
| status | string | Filter by display status. | No |
| status | string, <br>**Available values:** "archived", "available", "disabled", "error", "indexing", "paused", "queuing" | Filter by display status. | No |
#### DocumentListResponse
@ -3265,7 +3265,7 @@ Request payload for bulk downloading documents as a zip archive.
| doc_language | string, <br>**Default:** English | Language of the document for processing optimization. | No |
| embedding_model | string | Embedding model name. Use the `model` field from [Get Available Models](/api-reference/models/get-available-models) with `model_type=text-embedding`. | No |
| embedding_model_provider | string | Embedding model provider. Use the `provider` field from [Get Available Models](/api-reference/models/get-available-models) with `model_type=text-embedding`. | No |
| indexing_technique | string | `high_quality` uses embedding models for precise search; `economy` uses keyword-based indexing. Required when adding the first document to a knowledge base; subsequent documents inherit the knowledge base's indexing technique if omitted. | No |
| indexing_technique | string, <br>**Available values:** "economy", "high_quality" | `high_quality` uses embedding models for precise search; `economy` uses keyword-based indexing. Required when adding the first document to a knowledge base; subsequent documents inherit the knowledge base's indexing technique if omitted. | No |
| name | string | Document name. | Yes |
| original_document_id | string | Original document ID for replacement. | No |
| process_rule | [ProcessRule](#processrule) | Processing rules for chunking. | No |
@ -3614,14 +3614,14 @@ Model class for i18n object.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| content | string | Optional text feedback providing additional detail. | No |
| rating | string | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
| rating | string, <br>**Available values:** "dislike", "like" | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
#### MessageFeedbackPayloadWithUser
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| content | string | Optional text feedback providing additional detail. | No |
| rating | string | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
| rating | string, <br>**Available values:** "dislike", "like" | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
| user | string | User identifier, unique within the application. This identifier scopes data access; resources created with one `user` value are only visible when queried with the same `user` value. | Yes |
#### MessageFile
@ -3701,7 +3701,7 @@ Metadata Filtering Condition.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| conditions | [ [Condition](#condition) ] | List of metadata conditions to evaluate. | No |
| logical_operator | string | How to combine multiple conditions. | No |
| logical_operator | string, <br>**Available values:** "and", "or" | How to combine multiple conditions. | No |
#### MetadataOperationData
@ -3935,7 +3935,7 @@ Model class for provider with models response.
| ---- | ---- | ----------- | -------- |
| metadata_filtering_conditions | [MetadataFilteringCondition](#metadatafilteringcondition) | Restrict retrieval to chunks whose document metadata matches the given conditions. Conditions are evaluated server-side against document metadata fields. | No |
| reranking_enable | boolean | Whether reranking is enabled. | Yes |
| reranking_mode | string | Reranking mode. Required when `reranking_enable` is `true`. | No |
| reranking_mode | string, <br>**Available values:** "reranking_model", "weighted_score" | Reranking mode. Required when `reranking_enable` is `true`. | No |
| reranking_model | [RerankingModel](#rerankingmodel) | Reranking model configuration. | No |
| score_threshold | number | Minimum similarity score for results. Only effective when score threshold filtering is enabled. | No |
| score_threshold_enabled | boolean | Whether score threshold filtering is enabled. | Yes |
@ -3969,7 +3969,7 @@ Model class for provider with models response.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| parent_mode | string | Parent-child segmentation mode. | No |
| parent_mode | string, <br>**Available values:** "full-doc", "paragraph" | Parent-child segmentation mode. | No |
| pre_processing_rules | [ [PreProcessingRule](#preprocessingrule) ] | Pre-processing rules to apply before segmentation. | No |
| segmentation | [Segmentation](#segmentation) | Parent chunk segmentation settings. | No |
| subchunk_segmentation | [Segmentation](#segmentation) | Child chunk segmentation settings. | No |
@ -4300,7 +4300,7 @@ in form definition, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| keyword_setting | [WeightKeywordSetting](#weightkeywordsetting) | Keyword search weight settings. | No |
| vector_setting | [WeightVectorSetting](#weightvectorsetting) | Semantic search weight settings. | No |
| weight_type | string | Strategy for balancing semantic and keyword search weights. | No |
| weight_type | string, <br>**Available values:** "customized", "keyword_first", "semantic_first" | Strategy for balancing semantic and keyword search weights. | No |
#### WeightVectorSetting
@ -4383,7 +4383,7 @@ Blocking workflow response for a finished or paused execution.
| keyword | string | Keyword to search in logs. | No |
| limit | integer, <br>**Default:** 20 | Number of items per page. | No |
| page | integer, <br>**Default:** 1 | Page number for pagination. | No |
| status | string | Filter by execution status. | No |
| status | string, <br>**Available values:** "failed", "stopped", "succeeded" | Filter by execution status. | No |
#### WorkflowPauseReasonResponse
@ -4452,7 +4452,7 @@ Public pause reason emitted by a blocking Workflow execution.
| ---- | ---- | ----------- | -------- |
| files | [ object<br>object<br>object<br>object ] | File list for workflow system file inputs. Available when file upload is enabled for the workflow. To attach a local file, first upload it via [Upload File](/api-reference/files/upload-file) and use the returned `id` as `upload_file_id` with `transfer_method: local_file`. | No |
| inputs | object | Key-value pairs for workflow input variables. Values for file-type variables should be arrays of file objects with `type`, `transfer_method`, and either `url` or `upload_file_id`. Refer to the `user_input_form` field in the [Get App Parameters](/api-reference/applications/get-app-parameters) response to discover the variable names and types expected by your app. | Yes |
| response_mode | string | Response mode. Use `blocking` for synchronous responses or `streaming` for Server-Sent Events. When omitted, the request runs in blocking mode. | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode. Use `blocking` for synchronous responses or `streaming` for Server-Sent Events. When omitted, the request runs in blocking mode. | No |
#### WorkflowRunPayloadWithUser
@ -4460,7 +4460,7 @@ Public pause reason emitted by a blocking Workflow execution.
| ---- | ---- | ----------- | -------- |
| files | [ object<br>object<br>object<br>object ] | File list for workflow system file inputs. Available when file upload is enabled for the workflow. To attach a local file, first upload it via [Upload File](/api-reference/files/upload-file) and use the returned `id` as `upload_file_id` with `transfer_method: local_file`. | No |
| inputs | object | Key-value pairs for workflow input variables. Values for file-type variables should be arrays of file objects with `type`, `transfer_method`, and either `url` or `upload_file_id`. Refer to the `user_input_form` field in the [Get App Parameters](/api-reference/applications/get-app-parameters) response to discover the variable names and types expected by your app. | Yes |
| response_mode | string | Response mode. Use `blocking` for synchronous responses or `streaming` for Server-Sent Events. When omitted, the request runs in blocking mode. | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode. Use `blocking` for synchronous responses or `streaming` for Server-Sent Events. When omitted, the request runs in blocking mode. | No |
| user | string | User identifier, unique within the application. This identifier scopes data access; resources created with one `user` value are only visible when queried with the same `user` value. | Yes |
#### WorkflowRunResponse

View File

@ -1019,7 +1019,7 @@ Button styles for user actions.
| inputs | object | Input variables for the chat | Yes |
| parent_message_id | string | Parent message ID | No |
| query | string | User query/message | Yes |
| response_mode | string | Response mode: blocking or streaming | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode: blocking or streaming | No |
| retriever_from | string, <br>**Default:** web_app | Source of retriever | No |
#### CompletionMessagePayload
@ -1029,7 +1029,7 @@ Button styles for user actions.
| files | [ object ] | Files to be processed | No |
| inputs | object | Input variables for the completion | Yes |
| query | string | Query text for completion | No |
| response_mode | string | Response mode: blocking or streaming | No |
| response_mode | string, <br>**Available values:** "blocking", "streaming" | Response mode: blocking or streaming | No |
| retriever_from | string, <br>**Default:** web_app | Source of retriever | No |
#### ConversationInfiniteScrollPagination
@ -1322,7 +1322,7 @@ Parsed multipart form fields for HITL uploads.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| content | string | Optional text feedback providing additional detail. | No |
| rating | string | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
| rating | string, <br>**Available values:** "dislike", "like" | Feedback rating. Set to `null` to revoke previously submitted feedback. | No |
#### MessageFile

View File

@ -31,6 +31,14 @@ class SQLAlchemyAccountRepository(AccountRepository):
account = session.get(Account, account_id)
return self._to_snapshot(account) if account is not None else None
@override
def find_by_email(self, email: str) -> AccountSnapshot | None:
with self._session_factory() as session:
account = session.scalar(select(Account).where(Account.email == email).limit(1))
if account is None and email != email.lower():
account = session.scalar(select(Account).where(Account.email == email.lower()).limit(1))
return self._to_snapshot(account) if account is not None else None
@override
def get_credentials(self, account_id: str) -> AccountCredentials | None:
with self._session_factory() as session:

View File

@ -0,0 +1,189 @@
"""SQLAlchemy repository for account Step-by-step Tour state."""
import logging
from collections.abc import Callable
from typing import Protocol, override, runtime_checkable
from sqlalchemy import select, update
from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.orm import Session, sessionmaker
from models.onboarding import AccountStepByStepTourState
from services.entities.onboarding_entities import StepByStepTourState
from services.step_by_step_tour_service import StepByStepTourStateRepository
logger = logging.getLogger(__name__)
_MYSQL_RETRYABLE_LOCK_ERRNOS = frozenset({1205, 1213})
_MAX_LOCK_ATTEMPTS = 3
@runtime_checkable
class _ErrorWithErrno(Protocol):
@property
def errno(self) -> object: ...
class SQLAlchemyStepByStepTourStateRepository(StepByStepTourStateRepository):
def __init__(self, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
@override
def get(self, account_id: str) -> StepByStepTourState | None:
with self._session_factory() as session:
model = self._get_model(account_id, session=session)
return self._to_state(model) if model is not None else None
@override
def initialize(self, account_id: str, first_workspace_id: str) -> StepByStepTourState:
"""Create state with its first workspace, or atomically claim a legacy empty state."""
return self._run_with_lock_retry(
lambda: self._initialize_once(account_id, first_workspace_id),
)
def _initialize_once(self, account_id: str, first_workspace_id: str) -> StepByStepTourState:
with self._session_factory() as session:
model = self._get_model(account_id, session=session)
if model is None:
model = AccountStepByStepTourState(
account_id=account_id,
first_workspace_id=first_workspace_id,
)
session.add(model)
try:
session.commit()
except IntegrityError:
# A concurrent request inserted the account-owned row first.
session.rollback()
model = self._get_model(account_id, session=session)
if model is None:
raise
else:
session.refresh(model)
return self._to_state(model)
if model.first_workspace_id is None:
stmt = (
update(AccountStepByStepTourState)
.where(
AccountStepByStepTourState.account_id == account_id,
AccountStepByStepTourState.first_workspace_id.is_(None),
)
.values(first_workspace_id=first_workspace_id)
.execution_options(synchronize_session=False)
)
session.execute(stmt)
session.commit()
# A competing conditional update may have won while this request waited.
session.refresh(model)
return self._to_state(model)
@override
def mutate(
self,
account_id: str,
mutation: Callable[[StepByStepTourState], StepByStepTourState],
) -> StepByStepTourState:
"""Lock, create if needed, mutate, and persist account state in one transaction."""
return self._run_with_lock_retry(
lambda: self._mutate_once(account_id, mutation),
)
def _mutate_once(
self,
account_id: str,
mutation: Callable[[StepByStepTourState], StepByStepTourState],
) -> StepByStepTourState:
with self._session_factory() as session:
# Probe without a locking read so a missing MySQL unique key does not
# acquire a gap/next-key lock before the insert.
model = self._get_model(account_id, session=session)
if model is None:
model = AccountStepByStepTourState(account_id=account_id)
session.add(model)
try:
session.flush()
except IntegrityError:
# A concurrent mutation created the row. Start a new transaction,
# lock its committed state, and replay the pure mutation on it.
session.rollback()
model = self._get_model(account_id, session=session, lock_for_update=True)
if model is None:
raise
else:
model = self._get_model(account_id, session=session, lock_for_update=True)
if model is None:
raise RuntimeError("Step-by-step Tour state disappeared while acquiring its lock")
state = mutation(self._to_state(model))
if state.account_id != account_id:
raise ValueError("Step-by-step Tour mutation cannot change account ownership")
# first_workspace_id is write-once and owned exclusively by initialize().
model.skipped = state.skipped
model.completed_task_ids = list(state.completed_task_ids)
model.manually_enabled_workspace_ids = list(state.manually_enabled_workspace_ids)
model.manually_disabled_workspace_ids = list(state.manually_disabled_workspace_ids)
session.commit()
session.refresh(model)
return self._to_state(model)
@staticmethod
def _run_with_lock_retry[T](operation: Callable[[], T]) -> T:
for attempt in range(1, _MAX_LOCK_ATTEMPTS):
try:
return operation()
except OperationalError as exc:
if not _is_retryable_mysql_lock_error(exc):
raise
logger.warning(
"Retrying Step-by-step Tour transaction after MySQL lock failure (attempt %s/%s)",
attempt,
_MAX_LOCK_ATTEMPTS,
)
return operation()
@staticmethod
def _get_model(
account_id: str,
*,
session: Session,
lock_for_update: bool = False,
) -> AccountStepByStepTourState | None:
stmt = select(AccountStepByStepTourState).where(AccountStepByStepTourState.account_id == account_id).limit(1)
if lock_for_update:
stmt = stmt.with_for_update().execution_options(populate_existing=True)
return session.execute(stmt).scalar_one_or_none()
@staticmethod
def _to_state(model: AccountStepByStepTourState) -> StepByStepTourState:
return StepByStepTourState(
account_id=model.account_id,
first_workspace_id=model.first_workspace_id,
skipped=model.skipped,
completed_task_ids=tuple(model.completed_task_ids),
manually_enabled_workspace_ids=tuple(model.manually_enabled_workspace_ids),
manually_disabled_workspace_ids=tuple(model.manually_disabled_workspace_ids),
updated_at=model.updated_at,
)
def _is_retryable_mysql_lock_error(exc: OperationalError) -> bool:
orig = exc.orig
if isinstance(orig, _ErrorWithErrno) and _is_retryable_mysql_lock_error_code(orig.errno):
return True
if not isinstance(orig, BaseException) or not orig.args:
return False
return _is_retryable_mysql_lock_error_code(orig.args[0])
def _is_retryable_mysql_lock_error_code(candidate: object) -> bool:
if isinstance(candidate, bool):
return False
if isinstance(candidate, int):
code = candidate
elif isinstance(candidate, str) and candidate.isdecimal():
code = int(candidate)
else:
return False
return code in _MYSQL_RETRYABLE_LOCK_ERRNOS

View File

@ -0,0 +1,230 @@
"""Infrastructure adapters for account email registration."""
import logging
import secrets
from typing import override
from redis import RedisError
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_redis import RedisClientWrapper
from libs.helper import RateLimiter, TokenManager
from models.account import Account
from services.account_email_registration_service import (
AccountRegistrationGateway,
AccountRegistrationPolicyGateway,
EmailRegistrationCodeGenerator,
EmailRegistrationNotificationGateway,
EmailRegistrationSecurityGateway,
EmailRegistrationSendLimiter,
EmailRegistrationTokenGateway,
)
from services.account_errors import (
AccountEmailDomainSuspendedError,
AccountEmailFrozenError,
AccountNormalizedEmailAlreadyInUseError,
EmailRegistrationSeatsLimitError,
)
from services.account_service import AccountService
from services.billing_service import BillingService
from services.entities.account_entities import (
AccountEmailRegistrationPhase,
AccountEmailRegistrationToken,
AccountSessionTokens,
)
from services.errors.account import (
AccountNormalizedEmailAlreadyInUseError as AccountNormalizedEmailAlreadyInUseServiceError,
)
from services.errors.account import AccountRegisterError, EmailDomainSuspendedError, SeatsLimitExceededError
from tasks.mail_register_task import send_email_register_mail_task, send_email_register_mail_task_when_account_exist
logger = logging.getLogger(__name__)
class TokenManagerEmailRegistrationTokenGateway(EmailRegistrationTokenGateway):
@override
def get(self, token: str) -> AccountEmailRegistrationToken | None:
payload = TokenManager.get_token_data(token, "email_register")
if payload is None:
return None
email = payload.get("email")
code = payload.get("code")
phase_value = payload.get("phase")
if not isinstance(email, str) or not isinstance(code, str):
return None
if phase_value is None:
phase = None
else:
try:
phase = AccountEmailRegistrationPhase(phase_value)
except (TypeError, ValueError):
return None
return AccountEmailRegistrationToken(email=email, code=code, phase=phase)
@override
def issue(self, token_data: AccountEmailRegistrationToken) -> str:
additional_data = {"code": token_data.code}
if token_data.phase is not None:
additional_data["phase"] = token_data.phase.value
return TokenManager.generate_token(
email=token_data.email,
token_type="email_register",
additional_data=additional_data,
)
@override
def revoke(self, token: str) -> None:
TokenManager.revoke_token(token, "email_register")
class SecureEmailRegistrationCodeGenerator(EmailRegistrationCodeGenerator):
@override
def generate(self) -> str:
return "".join(str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6))
class CeleryEmailRegistrationNotificationGateway(EmailRegistrationNotificationGateway):
@override
def send_code(self, *, email: str, code: str, language: str) -> None:
send_email_register_mail_task.delay(language=language, to=email, code=code)
@override
def send_account_exists(self, *, email: str, account_name: str, language: str) -> None:
send_email_register_mail_task_when_account_exist.delay(
language=language,
to=email,
account_name=account_name,
)
class RateLimiterEmailRegistrationSendLimiter(EmailRegistrationSendLimiter):
def __init__(self, *, rate_limiter: RateLimiter) -> None:
self._rate_limiter = rate_limiter
@override
def is_limited(self, email: str) -> bool:
return self._rate_limiter.is_rate_limited(email)
@override
def record(self, email: str) -> None:
self._rate_limiter.increment_rate_limit(email)
@property
@override
def retry_after_minutes(self) -> int:
return int(self._rate_limiter.time_window / 60)
class RedisEmailRegistrationSecurityGateway(EmailRegistrationSecurityGateway):
def __init__(
self,
*,
redis: RedisClientWrapper,
verification_failure_limit: int,
verification_lockout_duration: int,
) -> None:
self._redis = redis
self._verification_failure_limit = verification_failure_limit
self._verification_lockout_duration = verification_lockout_duration
@override
def is_ip_limited(self, ip_address: str) -> bool:
return AccountService.is_email_send_ip_limit(ip_address) is True
@override
def is_verification_limited(self, email: str) -> bool:
try:
count = self._redis.get(self._verification_key(email))
return count is not None and int(count) > self._verification_failure_limit
except RedisError:
logger.warning("Failed to read email-registration verification limit", exc_info=True)
return False
@override
def record_verification_failure(self, email: str) -> None:
try:
key = self._verification_key(email)
count = int(self._redis.get(key) or 0) + 1
self._redis.setex(key, self._verification_lockout_duration, count)
except RedisError:
logger.warning("Failed to record email-registration verification failure", exc_info=True)
return None
@override
def reset_verification_failures(self, email: str) -> None:
try:
self._redis.delete(self._verification_key(email))
except RedisError:
logger.warning("Failed to reset email-registration verification failures", exc_info=True)
return None
@override
def reset_login_failures(self, email: str) -> None:
AccountService.reset_login_error_rate_limit(email)
@staticmethod
def _verification_key(email: str) -> str:
return f"email_register_error_rate_limit:{email}"
class BillingAccountRegistrationPolicyGateway(AccountRegistrationPolicyGateway):
def __init__(self, *, enabled: bool) -> None:
self._enabled = enabled
@override
def get_freeze_type(self, email: str) -> str | None:
if not self._enabled:
return None
return BillingService.get_email_freeze_type(email)
class AccountServiceRegistrationGateway(AccountRegistrationGateway):
"""Compatibility adapter around account provisioning and login internals."""
def __init__(self, *, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
@override
def create(
self,
*,
email: str,
password: str,
interface_language: str,
timezone: str | None,
ip_address: str,
) -> str:
with self._session_factory() as session:
try:
account = AccountService.create_account_and_tenant(
email=email,
name=email,
password=password,
interface_language=interface_language,
timezone=timezone,
ip_address=ip_address,
check_normalized_email=True,
session=session,
)
except SeatsLimitExceededError as exc:
raise EmailRegistrationSeatsLimitError from exc
except EmailDomainSuspendedError as exc:
raise AccountEmailDomainSuspendedError from exc
except AccountNormalizedEmailAlreadyInUseServiceError as exc:
raise AccountNormalizedEmailAlreadyInUseError from exc
except AccountRegisterError as exc:
raise AccountEmailFrozenError from exc
return account.id
@override
def login(self, account_id: str, *, ip_address: str) -> AccountSessionTokens:
with self._session_factory() as session:
account = session.get(Account, account_id)
if account is None:
raise RuntimeError("newly registered account no longer exists")
token_pair = AccountService.login(account=account, session=session, ip_address=ip_address)
return AccountSessionTokens(
access_token=token_pair.access_token,
refresh_token=token_pair.refresh_token,
csrf_token=token_pair.csrf_token,
)

View File

@ -0,0 +1,207 @@
"""Application service for the account email-registration use case."""
from typing import Protocol
from constants.languages import get_valid_language, languages
from services.account_errors import (
AccountEmailAlreadyInUseError,
AccountEmailDomainSuspendedError,
AccountEmailFrozenError,
EmailRegistrationPasswordMismatchError,
EmailRegistrationSendIPLimitedError,
EmailRegistrationSendRateLimitError,
EmailRegistrationVerificationLimitError,
InvalidEmailRegistrationAddressError,
InvalidEmailRegistrationCodeError,
InvalidEmailRegistrationTokenError,
)
from services.account_ports import AccountRepository
from services.entities.account_entities import (
AccountEmailRegistrationPhase,
AccountEmailRegistrationToken,
AccountEmailRegistrationVerification,
AccountSessionTokens,
)
class EmailRegistrationTokenGateway(Protocol):
def get(self, token: str) -> AccountEmailRegistrationToken | None: ...
def issue(self, token_data: AccountEmailRegistrationToken) -> str: ...
def revoke(self, token: str) -> None: ...
class EmailRegistrationCodeGenerator(Protocol):
def generate(self) -> str: ...
class EmailRegistrationNotificationGateway(Protocol):
def send_code(self, *, email: str, code: str, language: str) -> None: ...
def send_account_exists(self, *, email: str, account_name: str, language: str) -> None: ...
class EmailRegistrationSendLimiter(Protocol):
def is_limited(self, email: str) -> bool: ...
def record(self, email: str) -> None: ...
@property
def retry_after_minutes(self) -> int: ...
class EmailRegistrationSecurityGateway(Protocol):
def is_ip_limited(self, ip_address: str) -> bool: ...
def is_verification_limited(self, email: str) -> bool: ...
def record_verification_failure(self, email: str) -> None: ...
def reset_verification_failures(self, email: str) -> None: ...
def reset_login_failures(self, email: str) -> None: ...
class AccountRegistrationPolicyGateway(Protocol):
def get_freeze_type(self, email: str) -> str | None: ...
class AccountRegistrationGateway(Protocol):
def create(
self,
*,
email: str,
password: str,
interface_language: str,
timezone: str | None,
ip_address: str,
) -> str: ...
def login(self, account_id: str, *, ip_address: str) -> AccountSessionTokens: ...
class AccountEmailRegistrationService:
def __init__(
self,
*,
accounts: AccountRepository,
tokens: EmailRegistrationTokenGateway,
codes: EmailRegistrationCodeGenerator,
notifications: EmailRegistrationNotificationGateway,
send_limits: EmailRegistrationSendLimiter,
security: EmailRegistrationSecurityGateway,
account_policy: AccountRegistrationPolicyGateway,
registration: AccountRegistrationGateway,
) -> None:
self._accounts = accounts
self._tokens = tokens
self._codes = codes
self._notifications = notifications
self._send_limits = send_limits
self._security = security
self._account_policy = account_policy
self._registration = registration
def send_code(
self,
*,
remote_ip: str,
requested_email: str,
requested_language: str | None,
) -> str:
if self._security.is_ip_limited(remote_ip):
raise EmailRegistrationSendIPLimitedError
normalized_email = requested_email.lower()
self._ensure_email_allowed(normalized_email)
account = self._accounts.find_by_email(requested_email)
delivery_email = account.email if account is not None else normalized_email
if self._send_limits.is_limited(delivery_email):
raise EmailRegistrationSendRateLimitError(self._send_limits.retry_after_minutes)
language = requested_language if requested_language is not None and requested_language in languages else "en-US"
code = self._codes.generate()
token = self._tokens.issue(AccountEmailRegistrationToken(email=delivery_email, code=code))
if account is None:
self._notifications.send_code(email=delivery_email, code=code, language=language)
else:
self._notifications.send_account_exists(
email=delivery_email,
account_name=account.name,
language=language,
)
self._send_limits.record(delivery_email)
return token
def verify_code(
self,
*,
email: str,
code: str,
token: str,
) -> AccountEmailRegistrationVerification:
normalized_email = email.lower()
if self._security.is_verification_limited(normalized_email):
raise EmailRegistrationVerificationLimitError
token_data = self._tokens.get(token)
if token_data is None:
raise InvalidEmailRegistrationTokenError
normalized_token_email = token_data.email.lower()
if normalized_email != normalized_token_email:
raise InvalidEmailRegistrationAddressError
if code != token_data.code:
self._security.record_verification_failure(normalized_email)
raise InvalidEmailRegistrationCodeError
self._tokens.revoke(token)
verified_token = self._tokens.issue(
AccountEmailRegistrationToken(
email=normalized_email,
code=code,
phase=AccountEmailRegistrationPhase.REGISTER,
)
)
self._security.reset_verification_failures(normalized_email)
return AccountEmailRegistrationVerification(email=normalized_token_email, token=verified_token)
def register(
self,
*,
remote_ip: str,
token: str,
new_password: str,
password_confirm: str,
language: str | None,
timezone: str | None,
) -> AccountSessionTokens:
if new_password != password_confirm:
raise EmailRegistrationPasswordMismatchError
token_data = self._tokens.get(token)
if token_data is None or token_data.phase != AccountEmailRegistrationPhase.REGISTER:
raise InvalidEmailRegistrationTokenError
self._tokens.revoke(token)
normalized_email = token_data.email.lower()
if self._accounts.find_by_email(token_data.email) is not None:
raise AccountEmailAlreadyInUseError
account_id = self._registration.create(
email=normalized_email,
password=password_confirm,
interface_language=get_valid_language(language),
timezone=timezone,
ip_address=remote_ip,
)
tokens = self._registration.login(account_id, ip_address=remote_ip)
self._security.reset_login_failures(normalized_email)
return tokens
def _ensure_email_allowed(self, email: str) -> None:
freeze_type = self._account_policy.get_freeze_type(email)
if freeze_type == "email_domain_suspended":
raise AccountEmailDomainSuspendedError
if freeze_type:
raise AccountEmailFrozenError

View File

@ -85,6 +85,46 @@ class AccountEmailAlreadyInUseError(AccountApplicationError):
"""The target email already belongs to an account."""
class AccountNormalizedEmailAlreadyInUseError(AccountEmailAlreadyInUseError):
"""A normalized equivalent of the target email already belongs to an account."""
class EmailRegistrationSendIPLimitedError(AccountApplicationError):
"""The caller IP exceeded the registration-email send policy."""
class EmailRegistrationSendRateLimitError(AccountApplicationError):
"""Too many registration messages were requested for the address."""
def __init__(self, retry_after_minutes: int) -> None:
super().__init__(retry_after_minutes)
self.retry_after_minutes = retry_after_minutes
class EmailRegistrationVerificationLimitError(AccountApplicationError):
"""Too many invalid registration-code attempts were made."""
class InvalidEmailRegistrationTokenError(AccountApplicationError):
"""The registration token is absent, malformed, or in the wrong phase."""
class InvalidEmailRegistrationAddressError(AccountApplicationError):
"""The request address does not match the registration token."""
class InvalidEmailRegistrationCodeError(AccountApplicationError):
"""The verification code does not match the registration token."""
class EmailRegistrationPasswordMismatchError(AccountApplicationError):
"""The registration password confirmation does not match."""
class EmailRegistrationSeatsLimitError(AccountApplicationError):
"""The deployment has no licensed seat available for another account."""
class EducationDiscountPausedError(AccountApplicationError):
"""Education discount activation is temporarily paused."""

View File

@ -19,6 +19,8 @@ from services.entities.account_entities import (
class AccountRepository(Protocol):
def get(self, account_id: str) -> AccountSnapshot | None: ...
def find_by_email(self, email: str) -> AccountSnapshot | None: ...
def get_credentials(self, account_id: str) -> AccountCredentials | None: ...
def update_profile(self, account_id: str, changes: AccountProfileChanges) -> AccountSnapshot | None: ...

View File

@ -93,7 +93,6 @@ from tasks.mail_owner_transfer_task import (
send_old_owner_transfer_notify_email_task,
send_owner_transfer_confirm_task,
)
from tasks.mail_register_task import send_email_register_mail_task, send_email_register_mail_task_when_account_exist
from tasks.mail_reset_password_task import (
send_reset_password_mail_task,
send_reset_password_mail_task_when_account_not_exist,
@ -157,7 +156,6 @@ class AccountService:
CHANGE_EMAIL_PHASE_NEW = ChangeEmailPhase.NEW_EMAIL
reset_password_rate_limiter = RateLimiter(prefix="reset_password_rate_limit", max_attempts=1, time_window=60 * 1)
email_register_rate_limiter = RateLimiter(prefix="email_register_rate_limit", max_attempts=1, time_window=60 * 1)
email_code_login_rate_limiter = RateLimiter(
prefix="email_code_login_rate_limit", max_attempts=3, time_window=300 * 1
)
@ -168,7 +166,16 @@ class AccountService:
FORGOT_PASSWORD_MAX_ERROR_LIMITS = 5
CHANGE_EMAIL_MAX_ERROR_LIMITS = 5
OWNER_TRANSFER_MAX_ERROR_LIMITS = 5
EMAIL_REGISTER_MAX_ERROR_LIMITS = 5
@staticmethod
def _resolve_role_id_by_tag(tenant_id: str, account_id: str, tag: str) -> str:
options = ListOption(page_number=1, results_per_page=100)
roles = RBACService.Roles.list(tenant_id, account_id, options=options).data
for rbac_role in roles:
if rbac_role.is_builtin and rbac_role.category == "global_system_default" and rbac_role.role_tag == tag:
return str(rbac_role.id)
raise ValueError(f"Builtin RBAC role not found for tag {tag!r} in tenant {tenant_id}")
@staticmethod
def _resolve_legacy_role_id(tenant_id: str, account_id: str, role: TenantAccountRole) -> str:
@ -177,9 +184,6 @@ class AccountService:
Looks up the builtin RBAC role whose tag matches the legacy role name
(e.g. ``TenantAccountRole.ADMIN`` builtin role with tag ``"admin"``).
"""
options = ListOption(page_number=1, results_per_page=100)
roles = RBACService.Roles.list(tenant_id, account_id, options=options).data
expected_tag = {
TenantAccountRole.OWNER: "owner",
TenantAccountRole.ADMIN: "admin",
@ -187,15 +191,7 @@ class AccountService:
TenantAccountRole.NORMAL: "normal",
TenantAccountRole.DATASET_OPERATOR: "dataset_operator",
}[role]
for rbac_role in roles:
if (
rbac_role.is_builtin
and rbac_role.category == "global_system_default"
and rbac_role.role_tag == expected_tag
):
return str(rbac_role.id)
raise ValueError(f"Builtin RBAC role not found for {role.value} in tenant {tenant_id}")
return AccountService._resolve_role_id_by_tag(tenant_id, account_id, expected_tag)
@staticmethod
def get_workspace_permission_keys(tenant_id: str, account_id: str, *, session: Session) -> set[str]:
@ -680,40 +676,6 @@ class AccountService:
cls.reset_password_rate_limiter.increment_rate_limit(account_email)
return token
@classmethod
def send_email_register_email(
cls,
account: Account | None = None,
email: str | None = None,
language: str = "en-US",
):
account_email = account.email if account else email
if account_email is None:
raise ValueError("Email must be provided.")
if cls.email_register_rate_limiter.is_rate_limited(account_email):
from controllers.console.auth.error import EmailRegisterRateLimitExceededError
raise EmailRegisterRateLimitExceededError(int(cls.email_register_rate_limiter.time_window / 60))
code, token = cls.generate_email_register_token(account_email)
if account:
send_email_register_mail_task_when_account_exist.delay(
language=language,
to=account_email,
account_name=account.name,
)
else:
send_email_register_mail_task.delay(
language=language,
to=account_email,
code=code,
)
cls.email_register_rate_limiter.increment_rate_limit(account_email)
return token
@classmethod
def send_change_email_email(
cls,
@ -867,19 +829,6 @@ class AccountService:
)
return code, token
@classmethod
def generate_email_register_token(
cls,
email: str,
code: str | None = None,
additional_data: dict[str, Any] = {},
):
if not code:
code = "".join([str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)])
additional_data["code"] = code
token = TokenManager.generate_token(email=email, token_type="email_register", additional_data=additional_data)
return code, token
@classmethod
def generate_change_email_token(
cls,
@ -917,10 +866,6 @@ class AccountService:
def revoke_reset_password_token(cls, token: str):
TokenManager.revoke_token(token, "reset_password")
@classmethod
def revoke_email_register_token(cls, token: str):
TokenManager.revoke_token(token, "email_register")
@classmethod
def revoke_change_email_token(cls, token: str):
TokenManager.revoke_token(token, "change_email")
@ -933,10 +878,6 @@ class AccountService:
def get_reset_password_data(cls, token: str) -> dict[str, Any] | None:
return TokenManager.get_token_data(token, "reset_password")
@classmethod
def get_email_register_data(cls, token: str) -> dict[str, Any] | None:
return TokenManager.get_token_data(token, "email_register")
@classmethod
def get_change_email_data(cls, token: str) -> ChangeEmailTokenData | None:
token_data = TokenManager.get_token_data(token, "change_email")
@ -1067,16 +1008,6 @@ class AccountService:
count = int(count) + 1
redis_client.setex(key, dify_config.FORGOT_PASSWORD_LOCKOUT_DURATION, count)
@staticmethod
@redis_fallback(default_return=None)
def add_email_register_error_rate_limit(email: str) -> None:
key = f"email_register_error_rate_limit:{email}"
count = redis_client.get(key)
if count is None:
count = 0
count = int(count) + 1
redis_client.setex(key, dify_config.EMAIL_REGISTER_LOCKOUT_DURATION, count)
@staticmethod
@redis_fallback(default_return=False)
def is_forgot_password_error_rate_limit(email: str) -> bool:
@ -1096,24 +1027,6 @@ class AccountService:
key = f"forgot_password_error_rate_limit:{email}"
redis_client.delete(key)
@staticmethod
@redis_fallback(default_return=False)
def is_email_register_error_rate_limit(email: str) -> bool:
key = f"email_register_error_rate_limit:{email}"
count = redis_client.get(key)
if count is None:
return False
count = int(count)
if count > AccountService.EMAIL_REGISTER_MAX_ERROR_LIMITS:
return True
return False
@staticmethod
@redis_fallback(default_return=None)
def reset_email_register_error_rate_limit(email: str):
key = f"email_register_error_rate_limit:{email}"
redis_client.delete(key)
@staticmethod
@redis_fallback(default_return=None)
def add_change_email_error_rate_limit(email: str):
@ -1857,28 +1770,39 @@ class TenantService:
raise RoleAlreadyAssignedError("The provided role is already assigned to the member.")
if new_role == "owner":
# Find the current owner and change their role to 'admin'
if dify_config.RBAC_ENABLED:
old_owner_id = AccountService.get_rbac_workspace_owner_account_id(
str(tenant.id), operator.id, session=session
)
owner_role_id = AccountService._resolve_legacy_role_id(
tenant_id=str(tenant.id),
account_id=operator.id,
role=TenantAccountRole.OWNER,
)
no_access_role_id = AccountService._resolve_role_id_by_tag(
tenant_id=str(tenant.id),
account_id=operator.id,
tag="no_access",
)
current_roles = RBACService.MemberRoles.get(
str(tenant.id), operator.id, old_owner_id, session=session
).roles
remaining_role_ids = [str(r.id) for r in current_roles if str(r.id) != owner_role_id]
RBACService.MemberRoles.replace(
tenant_id=str(tenant.id),
account_id=operator.id,
member_account_id=old_owner_id,
role_ids=remaining_role_ids or [no_access_role_id],
session=session,
)
current_owner_join = session.scalar(
select(TenantAccountJoin)
.where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.role == "owner")
.limit(1)
)
if not dify_config.RBAC_ENABLED:
if current_owner_join:
current_owner_join.role = TenantAccountRole.ADMIN
elif current_owner_join:
admin_role_id = AccountService._resolve_legacy_role_id(
tenant_id=str(tenant.id),
account_id=operator.id,
role=TenantAccountRole.ADMIN,
)
RBACService.MemberRoles.replace(
tenant_id=str(tenant.id),
account_id=operator.id,
member_account_id=str(current_owner_join.account_id),
role_ids=[admin_role_id],
session=session,
)
if current_owner_join:
current_owner_join.role = TenantAccountRole.NORMAL
# Update the role of the target member
if dify_config.RBAC_ENABLED:
@ -1894,6 +1818,8 @@ class TenantService:
role_ids=[resolved_role_id],
session=session,
)
if new_tenant_role == TenantAccountRole.OWNER:
target_member_join.role = new_tenant_role
else:
target_member_join.role = new_tenant_role
session.commit()

View File

@ -120,7 +120,7 @@ class AppAnnotationService:
raw_message_id = args.get("message_id")
if raw_message_id:
message_id = str(raw_message_id)
message_id = raw_message_id
message = session.scalar(select(Message).where(Message.id == message_id, Message.app_id == app.id).limit(1))
if not message:
@ -176,19 +176,19 @@ class AppAnnotationService:
@classmethod
def enable_app_annotation(cls, args: EnableAnnotationArgs, app_id: str) -> AnnotationJobStatusDict:
enable_app_annotation_key = f"enable_app_annotation_{str(app_id)}"
enable_app_annotation_key = f"enable_app_annotation_{app_id}"
cache_result = redis_client.get(enable_app_annotation_key)
if cache_result is not None:
return {"job_id": cache_result, "job_status": "processing"}
# async job
job_id = str(uuid.uuid4())
enable_app_annotation_job_key = f"enable_app_annotation_job_{str(job_id)}"
enable_app_annotation_job_key = f"enable_app_annotation_job_{job_id}"
# send batch add segments task
redis_client.setnx(enable_app_annotation_job_key, "waiting")
current_user, current_tenant_id = current_account_with_tenant()
enable_annotation_reply_task.delay(
str(job_id),
job_id,
app_id,
current_user.id,
current_tenant_id,
@ -201,17 +201,17 @@ class AppAnnotationService:
@classmethod
def disable_app_annotation(cls, app_id: str) -> AnnotationJobStatusDict:
_, current_tenant_id = current_account_with_tenant()
disable_app_annotation_key = f"disable_app_annotation_{str(app_id)}"
disable_app_annotation_key = f"disable_app_annotation_{app_id}"
cache_result = redis_client.get(disable_app_annotation_key)
if cache_result is not None:
return {"job_id": cache_result, "job_status": "processing"}
# async job
job_id = str(uuid.uuid4())
disable_app_annotation_job_key = f"disable_app_annotation_job_{str(job_id)}"
disable_app_annotation_job_key = f"disable_app_annotation_job_{job_id}"
# send batch add segments task
redis_client.setnx(disable_app_annotation_job_key, "waiting")
disable_annotation_reply_task.delay(str(job_id), app_id, current_tenant_id)
disable_annotation_reply_task.delay(job_id, app_id, current_tenant_id)
return {"job_id": job_id, "job_status": "waiting"}
@classmethod
@ -539,7 +539,7 @@ class AppAnnotationService:
raise ValueError("The number of annotations exceeds the limit of your subscription.")
# async job
job_id = str(uuid.uuid4())
indexing_cache_key = f"app_annotation_batch_import_{str(job_id)}"
indexing_cache_key = f"app_annotation_batch_import_{job_id}"
# Register job in active tasks list for concurrency tracking
current_time = int(naive_utc_now().timestamp() * 1000)
@ -549,7 +549,7 @@ class AppAnnotationService:
# Set job status
redis_client.setnx(indexing_cache_key, "waiting")
batch_import_annotations_task.delay(str(job_id), result, app_id, current_tenant_id, current_user.id)
batch_import_annotations_task.delay(job_id, result, app_id, current_tenant_id, current_user.id)
except ValueError as e:
return {"error_msg": str(e)}

View File

@ -21,11 +21,17 @@ from core.app.features.rate_limiting import RateLimit
from core.app.features.rate_limiting.rate_limit import rate_limit_context
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig
from core.db import session_factory
from core.trigger.constants import is_trigger_node_type
from enums import DeploymentEdition, QuotaType
from extensions.otel import AppGenerateHandler, trace_span
from models.model import Account, App, AppMode, EndUser
from models.workflow import Workflow, WorkflowRun
from services.errors.app import QuotaExceededError, WorkflowIdFormatError, WorkflowNotFoundError
from services.errors.app import (
QuotaExceededError,
TriggerWorkflowServiceModeUnavailableError,
WorkflowIdFormatError,
WorkflowNotFoundError,
)
from services.errors.llm import InvokeRateLimitError
from services.quota_service import QuotaService, unlimited
from services.workflow_service import WorkflowService
@ -34,6 +40,13 @@ from tasks.app_generate.workflow_execute_task import AppExecutionParams, workflo
logger = logging.getLogger(__name__)
SSE_TASK_START_FALLBACK_MS = 200
_MANUAL_WORKFLOW_INVOKE_SOURCES = frozenset(
{
InvokeFrom.OPENAPI,
InvokeFrom.SERVICE_API,
InvokeFrom.WEB_APP,
}
)
if TYPE_CHECKING:
from controllers.console.app.workflow import LoopNodeRunPayload
@ -290,6 +303,7 @@ class AppGenerateService:
case AppMode.WORKFLOW:
workflow_id = args.get("workflow_id")
workflow = cls._get_workflow(app_model, invoke_from, workflow_id, session=session)
cls._ensure_workflow_service_mode_available(workflow=workflow, invoke_from=invoke_from)
if streaming:
with rate_limit_context(rate_limit, request_id):
payload = AppExecutionParams.new(
@ -343,6 +357,16 @@ class AppGenerateService:
case _:
raise ValueError(f"Invalid app mode {app_model.mode}")
@staticmethod
def _ensure_workflow_service_mode_available(*, workflow: Workflow, invoke_from: InvokeFrom) -> None:
if invoke_from not in _MANUAL_WORKFLOW_INVOKE_SOURCES:
return
for _, node_data in workflow.walk_nodes():
node_type = node_data.get("type")
if isinstance(node_type, str) and is_trigger_node_type(node_type):
raise TriggerWorkflowServiceModeUnavailableError()
@staticmethod
def _get_max_active_requests(app: App) -> int:
"""

View File

@ -1664,7 +1664,7 @@ class DocumentService:
"""Fetch documents for a dataset in a single batch query."""
if not document_ids:
return []
document_id_list: list[str] = [str(document_id) for document_id in document_ids]
document_id_list: list[str] = list(document_ids)
# Fetch all requested documents in one query to avoid N+1 lookups.
documents: Sequence[Document] = session.scalars(
select(Document).where(
@ -1700,7 +1700,7 @@ class DocumentService:
if not document_ids:
return 0
document_id_list: list[str] = [str(document_id) for document_id in document_ids]
document_id_list: list[str] = list(document_ids)
result = session.execute(
update(Document)
@ -1861,7 +1861,7 @@ class DocumentService:
"""
Batch load upload files keyed by document id for ZIP downloads.
"""
document_id_list: list[str] = [str(document_id) for document_id in document_ids]
document_id_list: list[str] = list(document_ids)
documents = DocumentService.get_documents_by_ids(
DatasetRef(tenant_id=tenant_id, dataset_id=dataset_id), document_id_list, session

View File

@ -475,6 +475,7 @@ _LEGACY_WORKSPACE_NORMAL_KEYS: list[str] = [
"plugin.install",
"credential.use",
"app_library.access",
"agent.manage",
]
_LEGACY_WORKSPACE_DATASET_OPERATOR_KEYS: list[str] = [
@ -482,6 +483,7 @@ _LEGACY_WORKSPACE_DATASET_OPERATOR_KEYS: list[str] = [
"plugin.install",
"dataset.create_and_management",
"dataset.external.connect",
"agent.manage",
]
_LEGACY_APP_OWNER_KEYS: list[str] = [
@ -2001,7 +2003,7 @@ class RBACService:
)
)
if current_owner_join and current_owner_join.account_id != member_account_id:
current_owner_join.role = TenantAccountRole.ADMIN
current_owner_join.role = TenantAccountRole.NORMAL
target_member_join.role = tenant_role
session.commit()

View File

@ -116,6 +116,30 @@ class AccountEmailResetResult:
account: AccountSnapshot | None = None
class AccountEmailRegistrationPhase(StrEnum):
REGISTER = "register"
@dataclass(frozen=True, slots=True)
class AccountEmailRegistrationToken:
email: str
code: str
phase: AccountEmailRegistrationPhase | None = None
@dataclass(frozen=True, slots=True)
class AccountEmailRegistrationVerification:
email: str
token: str
@dataclass(frozen=True, slots=True)
class AccountSessionTokens:
access_token: str
refresh_token: str
csrf_token: str
class AccountChangeEmailPhase(StrEnum):
OLD_EMAIL = "old_email"
OLD_EMAIL_VERIFIED = "old_email_verified"

View File

@ -0,0 +1,38 @@
"""Framework-independent notification contracts."""
from collections.abc import Mapping
from typing import NamedTuple
class NotificationContent(NamedTuple):
lang: str
title: str
subtitle: str
body: str
title_pic_url: str
class AccountNotification(NamedTuple):
notification_id: str | None
frequency: str | None
contents: Mapping[str, NotificationContent]
class AccountNotificationBatch(NamedTuple):
should_show: bool
notifications: tuple[AccountNotification, ...]
class NotificationItem(NamedTuple):
notification_id: str | None
frequency: str | None
lang: str
title: str
subtitle: str
body: str
title_pic_url: str
class NotificationResult(NamedTuple):
should_show: bool
notifications: tuple[NotificationItem, ...]

View File

@ -0,0 +1,42 @@
"""Framework-independent Step-by-step Tour contracts."""
from dataclasses import dataclass
from datetime import datetime
from typing import Literal, TypeAlias
# Assignment-form aliases preserve Literal enum values in Pydantic-generated OpenAPI schemas.
StepByStepTourAction: TypeAlias = Literal[ # noqa: UP040
"skip",
"complete_task",
"uncomplete_task",
"enable_current_workspace",
"disable_current_workspace",
]
StepByStepTourTaskId: TypeAlias = Literal["home", "studio", "knowledge", "integration"] # noqa: UP040
@dataclass(frozen=True, slots=True)
class StepByStepTourPatch:
action: StepByStepTourAction
task_id: StepByStepTourTaskId | None = None
@dataclass(frozen=True, slots=True)
class StepByStepTourState:
account_id: str
first_workspace_id: str | None = None
skipped: bool = False
completed_task_ids: tuple[str, ...] = ()
manually_enabled_workspace_ids: tuple[str, ...] = ()
manually_disabled_workspace_ids: tuple[str, ...] = ()
updated_at: datetime | None = None
@dataclass(frozen=True, slots=True)
class StepByStepTourResult:
first_workspace_id: str | None = None
skipped: bool = False
completed_task_ids: tuple[str, ...] = ()
manually_enabled_workspace_ids: tuple[str, ...] = ()
manually_disabled_workspace_ids: tuple[str, ...] = ()
updated_at: datetime | None = None

View File

@ -18,6 +18,21 @@ class WorkflowIdFormatError(Exception):
pass
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_CODE = "trigger_workflow_service_mode_unavailable"
TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_MESSAGE = (
"This workflow uses a trigger entry and cannot be invoked through Web App, Service API, OpenAPI, or MCP."
)
class TriggerWorkflowServiceModeUnavailableError(Exception):
"""Raised when a trigger-entry Workflow is invoked through a manual service surface."""
error_code = TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_CODE
def __init__(self) -> None:
super().__init__(TRIGGER_WORKFLOW_SERVICE_MODE_UNAVAILABLE_MESSAGE)
class QuotaExceededError(ValueError):
"""Raised when billing quota is exceeded for a feature."""

View File

@ -0,0 +1,48 @@
"""Billing-backed notification gateway."""
from collections.abc import Mapping
from typing import Any, override
from services.billing_service import BillingService
from services.entities.notification_entities import (
AccountNotification,
AccountNotificationBatch,
NotificationContent,
)
from services.notification_service import NotificationGateway
class BillingNotificationGateway(NotificationGateway):
@override
def get_active(self, account_id: str) -> AccountNotificationBatch:
payload = BillingService.get_account_notification(account_id)
notifications = tuple(self._map_notification(item) for item in payload.get("notifications") or ())
return AccountNotificationBatch(
should_show=bool(payload.get("shouldShow")),
notifications=notifications,
)
@override
def dismiss(self, notification_id: str, account_id: str) -> None:
BillingService.dismiss_notification(notification_id=notification_id, account_id=account_id)
@classmethod
def _map_notification(cls, payload: Mapping[str, Any]) -> AccountNotification:
raw_contents = payload.get("contents") or {}
contents = {language: cls._map_content(content) for language, content in raw_contents.items() if content}
return AccountNotification(
notification_id=payload.get("notificationId"),
frequency=payload.get("frequency"),
contents=contents,
)
@staticmethod
def _map_content(payload: Mapping[str, Any]) -> NotificationContent:
return NotificationContent(
# The application service owns the requested-language fallback.
lang=payload.get("lang") or "",
title=payload.get("title") or "",
subtitle=payload.get("subtitle") or "",
body=payload.get("body") or "",
title_pic_url=payload.get("titlePicUrl") or "",
)

View File

@ -0,0 +1,60 @@
"""Application service for Console account notifications."""
from typing import Protocol
from machinery.context import RequestContext
from services.account_ports import AccountRepository
from services.entities.notification_entities import (
AccountNotification,
AccountNotificationBatch,
NotificationContent,
NotificationItem,
NotificationResult,
)
_FALLBACK_LANGUAGE = "en-US"
class NotificationGateway(Protocol):
def get_active(self, account_id: str) -> AccountNotificationBatch: ...
def dismiss(self, notification_id: str, account_id: str) -> None: ...
class NotificationService:
def __init__(self, *, accounts: AccountRepository, notifications: NotificationGateway) -> None:
self._accounts = accounts
self._notifications = notifications
def get_active(self, context: RequestContext) -> NotificationResult:
batch = self._notifications.get_active(context.account_id)
if not batch.should_show:
return NotificationResult(should_show=False, notifications=())
account = self._accounts.get(context.account_id)
if account is None:
raise RuntimeError("Console account admission resolved an unknown account")
language = account.interface_language or _FALLBACK_LANGUAGE
notifications = tuple(self._localize(notification, language) for notification in batch.notifications)
return NotificationResult(should_show=bool(notifications), notifications=notifications)
def dismiss(self, context: RequestContext, notification_id: str) -> None:
self._notifications.dismiss(notification_id, context.account_id)
@staticmethod
def _localize(notification: AccountNotification, language: str) -> NotificationItem:
content = (
notification.contents.get(language)
or notification.contents.get(_FALLBACK_LANGUAGE)
or next(iter(notification.contents.values()), NotificationContent(language, "", "", "", ""))
)
return NotificationItem(
notification_id=notification.notification_id,
frequency=notification.frequency,
lang=content.lang or language,
title=content.title,
subtitle=content.subtitle,
body=content.body,
title_pic_url=content.title_pic_url,
)

View File

@ -1,221 +1,161 @@
"""Account-level Step-by-step Tour persistence."""
"""Application service for account-level Step-by-step Tour use cases."""
from collections.abc import Callable
from dataclasses import replace
from datetime import datetime
from typing import NotRequired, TypedDict
from typing import Protocol, get_args
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session, scoped_session
from configs import dify_config
from libs.datetime_utils import ensure_naive_utc
from models.account import Account
from models.onboarding import AccountStepByStepTourState
from machinery.context import RequestContext
from services.account_ports import AccountRepository
from services.entities.onboarding_entities import (
StepByStepTourPatch,
StepByStepTourResult,
StepByStepTourState,
StepByStepTourTaskId,
)
STEP_BY_STEP_TOUR_TASK_IDS = frozenset(("home", "studio", "knowledge", "integration"))
_TASK_IDS: frozenset[str] = frozenset(get_args(StepByStepTourTaskId))
class StepByStepTourStateResponse(TypedDict):
first_workspace_id: str | None
skipped: bool
completed_task_ids: list[str]
manually_enabled_workspace_ids: list[str]
manually_disabled_workspace_ids: list[str]
updated_at: datetime | None
class StepByStepTourStateRepository(Protocol):
def get(self, account_id: str) -> StepByStepTourState | None: ...
def initialize(self, account_id: str, first_workspace_id: str) -> StepByStepTourState: ...
class StepByStepTourPatch(TypedDict):
action: str
task_id: NotRequired[str | None]
def mutate(
self,
account_id: str,
mutation: Callable[[StepByStepTourState], StepByStepTourState],
) -> StepByStepTourState: ...
class StepByStepTourService:
"""Coordinate persisted tour state with account eligibility rules."""
@classmethod
def get_state(
cls,
def __init__(
self,
*,
account: Account,
current_tenant_id: str,
session: Session | scoped_session,
) -> StepByStepTourStateResponse:
eligible = cls.is_eligible(account)
state = cls._get_state(account.id, session=session)
accounts: AccountRepository,
states: StepByStepTourStateRepository,
enabled: bool,
rollout_started_at: datetime | None,
) -> None:
self._accounts = accounts
self._states = states
self._enabled = enabled
self._rollout_started_at = rollout_started_at
if eligible:
state = cls._ensure_state(account.id, session=session, state=state)
if state.first_workspace_id is None:
state.first_workspace_id = current_tenant_id
session.commit()
session.refresh(state)
def get_state(self, context: RequestContext) -> StepByStepTourResult:
workspace_id = self._require_workspace(context)
account = self._accounts.get(context.account_id)
if account is None:
raise RuntimeError("Console account admission resolved an unknown account")
return cls._build_response(state=state)
if not self._is_eligible(account.initialized_at or account.created_at):
return self._to_result(self._states.get(context.account_id))
@classmethod
def patch_state(
cls,
*,
account: Account,
current_tenant_id: str,
patch: StepByStepTourPatch,
session: Session | scoped_session,
) -> StepByStepTourStateResponse:
state = cls._ensure_state(account.id, session=session, state=None)
cls._apply_action(
state=state,
action=patch["action"],
task_id=patch.get("task_id"),
current_tenant_id=current_tenant_id,
return self._to_result(self._states.initialize(context.account_id, workspace_id))
def patch_state(self, context: RequestContext, patch: StepByStepTourPatch) -> StepByStepTourResult:
workspace_id = self._require_workspace(context)
state = self._states.mutate(
context.account_id,
lambda current: self._apply_action(current, patch=patch, workspace_id=workspace_id),
)
return self._to_result(state)
session.commit()
session.refresh(state)
return cls._build_response(state=state)
@classmethod
def is_eligible(cls, account: Account) -> bool:
if not dify_config.ENABLE_STEP_BY_STEP_TOUR:
def _is_eligible(self, account_started_at: datetime) -> bool:
if not self._enabled or self._rollout_started_at is None:
return False
rollout_started_at = dify_config.STEP_BY_STEP_TOUR_ROLLOUT_STARTED_AT
if rollout_started_at is None:
return False
account_started_at = account.initialized_at or account.created_at
if account_started_at is None:
return False
return ensure_naive_utc(account_started_at) >= ensure_naive_utc(rollout_started_at)
@classmethod
def _get_state(
cls,
account_id: str,
*,
session: Session | scoped_session,
) -> AccountStepByStepTourState | None:
stmt = select(AccountStepByStepTourState).where(AccountStepByStepTourState.account_id == account_id).limit(1)
return session.execute(stmt).scalar_one_or_none()
@classmethod
def _ensure_state(
cls,
account_id: str,
*,
session: Session | scoped_session,
state: AccountStepByStepTourState | None,
) -> AccountStepByStepTourState:
if state is None:
state = cls._get_state(account_id, session=session)
if state is not None:
return state
state = AccountStepByStepTourState(account_id=account_id)
session.add(state)
try:
session.flush()
except IntegrityError:
# Another tab/device can create the account row between our read and insert.
session.rollback()
state = cls._get_state(account_id, session=session)
if state is None:
raise
return state
return ensure_naive_utc(account_started_at) >= ensure_naive_utc(self._rollout_started_at)
@classmethod
def _apply_action(
cls,
state: StepByStepTourState,
*,
state: AccountStepByStepTourState,
action: str,
task_id: str | None,
current_tenant_id: str,
) -> None:
match action:
patch: StepByStepTourPatch,
workspace_id: str,
) -> StepByStepTourState:
match patch.action:
case "skip":
state.skipped = True
state.manually_enabled_workspace_ids = cls._remove_id(
state.manually_enabled_workspace_ids,
current_tenant_id,
return replace(
state,
skipped=True,
manually_enabled_workspace_ids=cls._remove_id(
state.manually_enabled_workspace_ids,
workspace_id,
),
)
case "complete_task":
if task_id is None:
raise ValueError("task_id is required")
cls._validate_task_id(task_id)
state.completed_task_ids = cls._add_id(state.completed_task_ids, task_id)
task_id = cls._require_task_id(patch.task_id)
return replace(state, completed_task_ids=cls._add_id(state.completed_task_ids, task_id))
case "uncomplete_task":
if task_id is None:
raise ValueError("task_id is required")
cls._validate_task_id(task_id)
state.completed_task_ids = cls._remove_id(state.completed_task_ids, task_id)
task_id = cls._require_task_id(patch.task_id)
return replace(state, completed_task_ids=cls._remove_id(state.completed_task_ids, task_id))
case "enable_current_workspace":
state.skipped = False
state.manually_enabled_workspace_ids = cls._add_id(
state.manually_enabled_workspace_ids,
current_tenant_id,
)
state.manually_disabled_workspace_ids = cls._remove_id(
state.manually_disabled_workspace_ids,
current_tenant_id,
return replace(
state,
skipped=False,
manually_enabled_workspace_ids=cls._add_id(
state.manually_enabled_workspace_ids,
workspace_id,
),
manually_disabled_workspace_ids=cls._remove_id(
state.manually_disabled_workspace_ids,
workspace_id,
),
)
case "disable_current_workspace":
state.manually_enabled_workspace_ids = cls._remove_id(
state.manually_enabled_workspace_ids,
current_tenant_id,
)
state.manually_disabled_workspace_ids = cls._add_id(
state.manually_disabled_workspace_ids,
current_tenant_id,
return replace(
state,
manually_enabled_workspace_ids=cls._remove_id(
state.manually_enabled_workspace_ids,
workspace_id,
),
manually_disabled_workspace_ids=cls._add_id(
state.manually_disabled_workspace_ids,
workspace_id,
),
)
case _:
raise ValueError(f"Unsupported action: {action}")
@classmethod
def _build_response(
cls,
*,
state: AccountStepByStepTourState | None,
) -> StepByStepTourStateResponse:
if state is None:
return {
"first_workspace_id": None,
"skipped": False,
"completed_task_ids": [],
"manually_enabled_workspace_ids": [],
"manually_disabled_workspace_ids": [],
"updated_at": None,
}
return {
"first_workspace_id": state.first_workspace_id,
"skipped": state.skipped,
"completed_task_ids": cls._normalize_ids(state.completed_task_ids),
"manually_enabled_workspace_ids": cls._normalize_ids(state.manually_enabled_workspace_ids),
"manually_disabled_workspace_ids": cls._normalize_ids(state.manually_disabled_workspace_ids),
"updated_at": state.updated_at,
}
raise ValueError(f"Unsupported action: {patch.action}")
@staticmethod
def _validate_task_id(task_id: str) -> None:
if task_id not in STEP_BY_STEP_TOUR_TASK_IDS:
def _require_workspace(context: RequestContext) -> str:
if context.active_workspace_id is None:
raise RuntimeError("Console account admission did not resolve an active workspace")
return context.active_workspace_id
@staticmethod
def _require_task_id(task_id: str | None) -> str:
if task_id is None:
raise ValueError("task_id is required")
if task_id not in _TASK_IDS:
raise ValueError(f"Unsupported task_id: {task_id}")
return task_id
@classmethod
def _add_id(cls, values: list[str], value: str) -> list[str]:
def _add_id(cls, values: tuple[str, ...], value: str) -> tuple[str, ...]:
normalized = cls._normalize_ids(values)
if value in normalized:
return normalized
return [*normalized, value]
return normalized if value in normalized else (*normalized, value)
@classmethod
def _remove_id(cls, values: list[str], value: str) -> list[str]:
return [item for item in cls._normalize_ids(values) if item != value]
def _remove_id(cls, values: tuple[str, ...], value: str) -> tuple[str, ...]:
return tuple(item for item in cls._normalize_ids(values) if item != value)
@staticmethod
def _normalize_ids(values: list[str]) -> list[str]:
normalized: list[str] = []
for value in values:
if value not in normalized:
normalized.append(value)
return normalized
def _normalize_ids(values: tuple[str, ...]) -> tuple[str, ...]:
return tuple(dict.fromkeys(values))
@staticmethod
def _to_result(state: StepByStepTourState | None) -> StepByStepTourResult:
if state is None:
return StepByStepTourResult()
return StepByStepTourResult(
first_workspace_id=state.first_workspace_id,
skipped=state.skipped,
completed_task_ids=tuple(dict.fromkeys(state.completed_task_ids)),
manually_enabled_workspace_ids=tuple(dict.fromkeys(state.manually_enabled_workspace_ids)),
manually_disabled_workspace_ids=tuple(dict.fromkeys(state.manually_disabled_workspace_ids)),
updated_at=state.updated_at,
)

View File

@ -300,7 +300,7 @@ class TestOwnerTransferApiWithContainers:
)
assert (
factory.get_join(db_session_with_containers, tenant=tenant, account=current_user).role
== TenantAccountRole.ADMIN
== TenantAccountRole.NORMAL
)
mock_new_owner_email.assert_called_once()
mock_old_owner_email.assert_called_once()

View File

@ -102,10 +102,8 @@ class TestConversationRenameApi:
ConversationRenameApi().post(_completion_app(), _end_user(), uuid4())
@patch("controllers.web.conversation.ConversationService.rename")
@patch("controllers.web.conversation.web_ns")
def test_rename_success(self, mock_ns: MagicMock, mock_rename: MagicMock, app: Flask) -> None:
def test_rename_success(self, mock_rename: MagicMock, app: Flask) -> None:
c_id = uuid4()
mock_ns.payload = {"name": "New Name", "auto_generate": False}
conv = SimpleNamespace(
id=str(c_id),
name="New Name",
@ -126,10 +124,8 @@ class TestConversationRenameApi:
"controllers.web.conversation.ConversationService.rename",
side_effect=ConversationNotExistsError(),
)
@patch("controllers.web.conversation.web_ns")
def test_rename_not_found(self, mock_ns: MagicMock, mock_rename: MagicMock, app: Flask) -> None:
def test_rename_not_found(self, mock_rename: MagicMock, app: Flask) -> None:
c_id = uuid4()
mock_ns.payload = {"name": "X", "auto_generate": False}
with app.test_request_context(f"/conversations/{c_id}/name", method="POST", json={"name": "X"}):
with pytest.raises(NotFound, match="Conversation Not Exists"):

View File

@ -1671,7 +1671,7 @@ class TestTenantService:
def test_update_member_role_to_owner(self, db_session_with_containers: Session, mock_external_service_dependencies):
"""
Test updating member role to owner (should change current owner to admin).
Test updating member role to owner (should change current owner to normal).
"""
fake = Faker()
tenant_name = fake.company()
@ -1723,7 +1723,7 @@ class TestTenantService:
.filter_by(tenant_id=tenant.id, account_id=member_account.id)
.first()
)
assert owner_join.role == "admin"
assert owner_join.role == "normal"
assert member_join.role == "owner"
def test_update_member_role_already_assigned(

View File

@ -9,6 +9,7 @@ from clients.agent_backend.factory import create_agent_backend_client, create_ag
from configs import dify_config
from services import agent_app_sandbox_service
from services.agent import home_snapshot_service, workspace_service
from tests.unit_tests.config_override import apply_config_overrides
@pytest.mark.parametrize(
@ -78,9 +79,12 @@ def test_default_agent_backend_clients_forward_authentication(
module: ModuleType,
extra_kwargs: dict[str, float],
) -> None:
monkeypatch.setattr(dify_config, "AGENT_BACKEND_BASE_URL", "http://agent-backend")
monkeypatch.setattr(dify_config, "AGENT_BACKEND_API_TOKEN", "secret-token")
monkeypatch.setattr(dify_config, "AGENT_BACKEND_BINDING_FILE_DOWNLOAD_TIMEOUT_SECONDS", 123.5)
apply_config_overrides(
monkeypatch,
AGENT_BACKEND_BASE_URL="http://agent-backend",
AGENT_BACKEND_API_TOKEN="secret-token",
AGENT_BACKEND_BINDING_FILE_DOWNLOAD_TIMEOUT_SECONDS=123.5,
)
create_client = MagicMock()
monkeypatch.setattr(module, "create_agent_backend_client", create_client)

View File

@ -238,6 +238,42 @@ def test_patch_union_schema_markdown_fills_regular_schema_union_property(tmp_pat
assert "| value | string<br>integer<br>number<br>boolean | | No |" in patched
def test_patch_union_schema_markdown_preserves_nullable_enum_values(tmp_path: Path):
module = _load_generate_swagger_markdown_docs_module()
spec_path = tmp_path / "console-openapi.json"
spec_path.write_text(
json.dumps(
{
"components": {
"schemas": {
"StepByStepTourStatePatchPayload": {
"properties": {
"task_id": {
"anyOf": [
{"enum": ["home", "studio"], "type": "string"},
{"type": "null"},
],
},
},
},
},
}
}
),
encoding="utf-8",
)
markdown = """#### StepByStepTourStatePatchPayload
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| task_id | string | Task ID | No |
"""
patched = module._patch_union_schema_markdown(markdown, spec_path)
assert '| task_id | string, <br>**Available values:** "home", "studio" | Task ID | No |' in patched
def test_patch_union_schema_markdown_fills_array_item_union_property(tmp_path: Path):
module = _load_generate_swagger_markdown_docs_module()
spec_path = tmp_path / "console-openapi.json"

View File

@ -9,6 +9,8 @@ from pathlib import Path
from jsonschema import Draft202012Validator
from tests.unit_tests.config_override import apply_config_overrides
def _walk_values(value):
yield value
@ -162,7 +164,7 @@ def test_apply_runtime_defaults_forces_swagger_routes_on(monkeypatch):
from configs import dify_config
monkeypatch.setenv("SWAGGER_UI_ENABLED", "false")
monkeypatch.setattr(dify_config, "SWAGGER_UI_ENABLED", False)
apply_config_overrides(monkeypatch, SWAGGER_UI_ENABLED=False)
module.apply_runtime_defaults()

View File

@ -21,6 +21,7 @@ from graphon.model_runtime.entities.model_entities import ModelType
from models import Tenant
from models.provider import Provider, ProviderModel, ProviderType
from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider
from tests.unit_tests.config_override import apply_config_overrides
def _invoke_reset() -> int:
@ -88,7 +89,7 @@ def _bind_command_to_sqlite(monkeypatch: pytest.MonkeyPatch, session: Session) -
def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
monkeypatch.setattr(system_commands.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD)
apply_config_overrides(monkeypatch, DEPLOYMENT_EDITION=DeploymentEdition.CLOUD)
exit_code = _invoke_reset()
captured = capsys.readouterr()
@ -107,7 +108,7 @@ def test_reset_purges_provider_and_tool_tables_for_each_tenant(
) -> None:
"""The command must purge LLM provider rows AND every tool provider table
that stores ciphertext encrypted under the tenant key (#35396)."""
monkeypatch.setattr(system_commands.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY)
apply_config_overrides(monkeypatch, DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY)
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)
@ -147,7 +148,7 @@ def test_reset_purges_provider_and_tool_tables_for_each_tenant(
)
def test_reset_iterates_all_tenants(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
"""Multi-tenant deployments must purge every tenant, not just the first."""
monkeypatch.setattr(system_commands.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY)
apply_config_overrides(monkeypatch, DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY)
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)

View File

@ -0,0 +1,26 @@
"""Typed config override support shared by unit-test fixtures and helpers."""
from collections.abc import Generator
from contextlib import contextmanager
import pytest
from configs import dify_config
def apply_config_overrides(monkeypatch: pytest.MonkeyPatch, **values: object) -> None:
"""Override known DifyConfig fields for the lifetime of ``monkeypatch``."""
unknown_fields = values.keys() - type(dify_config).model_fields.keys()
if unknown_fields:
raise ValueError(f"Unknown DifyConfig fields: {sorted(unknown_fields)}")
for name, value in values.items():
monkeypatch.setattr(dify_config, name, value)
@contextmanager
def config_overrides_context(**values: object) -> Generator[None]:
"""Apply validated config overrides as a context manager or decorator."""
with pytest.MonkeyPatch.context() as monkeypatch:
apply_config_overrides(monkeypatch, **values)
yield

View File

@ -41,6 +41,7 @@ import core.db.session_factory as session_factory_module
from extensions import ext_redis
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.base import TypeBase
from tests.unit_tests.config_override import apply_config_overrides
def _patch_redis_clients_on_loaded_modules() -> None:
@ -99,17 +100,9 @@ def reset_redis_mock() -> None:
@pytest.fixture(autouse=True)
def reset_secret_key() -> Iterator[None]:
def reset_secret_key(monkeypatch: pytest.MonkeyPatch) -> None:
"""Ensure SECRET_KEY-dependent logic sees an empty config value by default."""
from configs import dify_config
original = dify_config.SECRET_KEY
dify_config.SECRET_KEY = ""
try:
yield
finally:
dify_config.SECRET_KEY = original
apply_config_overrides(monkeypatch, SECRET_KEY="")
@pytest.fixture
@ -120,14 +113,9 @@ def config_overrides(monkeypatch: pytest.MonkeyPatch) -> Callable[..., None]:
field names keeps tests scoped without replacing that instance with an
unconstrained mock. ``monkeypatch`` restores every value after the test.
"""
from configs import dify_config
def apply(**values: object) -> None:
unknown_fields = values.keys() - type(dify_config).model_fields.keys()
if unknown_fields:
raise ValueError(f"Unknown DifyConfig fields: {sorted(unknown_fields)}")
for name, value in values.items():
monkeypatch.setattr(dify_config, name, value)
apply_config_overrides(monkeypatch, **values)
return apply

View File

@ -1,3 +1,4 @@
from collections.abc import Callable
from datetime import datetime
from inspect import getclosurevars, getsource, unwrap
from types import SimpleNamespace
@ -56,7 +57,7 @@ from controllers.console.agent.roster import (
from controllers.console.app import completion as completion_controller
from controllers.console.app import message as message_controller
from controllers.console.app.completion import AgentBuildChatFinalizeApi, AgentChatMessageApi, AgentChatMessageStopApi
from controllers.console.app.error import CompletionRequestError
from controllers.console.app.error import AgentSessionConfigurationChangedError, CompletionRequestError
from controllers.console.app.message import (
AgentChatMessageListApi,
AgentMessageApi,
@ -75,6 +76,7 @@ from services.entities.agent_entities import (
WorkflowAgentComposerQuery,
WorkflowComposerCopyFromRosterPayload,
)
from tests.unit_tests.config_override import apply_config_overrides
def _rbac_decorators(method: object) -> list[dict[str, object]]:
@ -395,7 +397,10 @@ def account_id() -> str:
def test_agent_app_list_and_create_use_agent_route(
app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, sqlite_session: Session
app: Flask,
monkeypatch: pytest.MonkeyPatch,
account_id: str,
sqlite_session: Session,
) -> None:
captured: dict[str, object] = {}
monkeypatch.setattr(roster_controller.dify_config, "RBAC_ENABLED", True)
@ -484,7 +489,9 @@ def test_agent_app_list_and_create_use_agent_route(
lambda _self, **kwargs: {"agent-list": "debug-conversation-list"},
)
monkeypatch.setattr(
roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 0
roster_controller.AgentRosterService,
"count_agent_app_debug_conversation_messages",
lambda _self, **kwargs: 0,
)
monkeypatch.setattr(
roster_controller.enterprise_rbac_service.RBACService.AgentPermissions,
@ -561,12 +568,22 @@ def test_agent_app_list_and_create_use_agent_route(
assert count_params.agent_is_published is True
with app.test_request_context(
"/console/api/agent",
json={"name": "Iris", "description": "Agent app", "role": "Coordinator", "icon_type": "emoji", "icon": "robot"},
json={
"name": "Iris",
"description": "Agent app",
"role": "Coordinator",
"icon_type": "emoji",
"icon": "robot",
},
):
created, status = unwrap(AgentAppListApi.post)(
AgentAppListApi(),
AgentAppCreatePayload(
name="Iris", description="Agent app", role="Coordinator", icon_type="emoji", icon="robot"
name="Iris",
description="Agent app",
role="Coordinator",
icon_type="emoji",
icon="robot",
),
sqlite_session,
"tenant-1",
@ -596,6 +613,76 @@ def test_agent_app_list_and_create_use_agent_route(
}
def test_agent_app_create_skips_rbac_access_initialization_when_rbac_is_disabled(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
account_id: str,
sqlite_session: Session,
config_overrides: Callable[..., None],
) -> None:
replace_whitelist = MagicMock()
initialize_access = MagicMock()
class FakeAppService:
def get_app(self, app_obj: object, *, session: object) -> object:
return app_obj
def create_app(self, tenant_id: str, params, current_user: object, *, session: object) -> object:
return _app_detail_obj(id="app-created", bound_agent_id="agent-created")
monkeypatch.setattr(roster_controller, "AppService", FakeAppService)
monkeypatch.setattr(
roster_controller.AgentRosterService,
"get_app_backing_agent",
lambda _self, **kwargs: Agent(
id="agent-created",
app_id="app-created",
backing_app_id=None,
role="Created role",
active_config_snapshot_id=None,
),
)
monkeypatch.setattr(
roster_controller.AgentRosterService,
"get_or_create_build_conversation",
lambda _self, **kwargs: "debug-conversation-created",
)
monkeypatch.setattr(
roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 0
)
monkeypatch.setattr(
roster_controller.FeatureService,
"get_system_features",
lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)),
)
config_overrides(RBAC_ENABLED=False)
monkeypatch.setattr(
roster_controller.enterprise_rbac_service.RBACService.AppAccess,
"replace_whitelist",
replace_whitelist,
)
monkeypatch.setattr(roster_controller.initialize_created_app_rbac_access_task, "delay", initialize_access)
with app.test_request_context(
"/console/api/agent",
json={"name": "Iris", "description": "Agent app", "role": "Coordinator", "icon_type": "emoji", "icon": "robot"},
):
created, status = unwrap(AgentAppListApi.post)(
AgentAppListApi(),
AgentAppCreatePayload(
name="Iris", description="Agent app", role="Coordinator", icon_type="emoji", icon="robot"
),
sqlite_session,
"tenant-1",
_account(account_id=account_id),
)
assert status == 201
assert created["id"] == "agent-created"
replace_whitelist.assert_not_called()
initialize_access.assert_not_called()
def test_agent_app_create_payload_allows_optional_role() -> None:
omitted = roster_controller.AgentAppCreatePayload.model_validate(
{"name": "Iris", "description": "Agent app", "icon_type": "emoji", "icon": "robot"}
@ -1009,7 +1096,7 @@ def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata(monkeyp
monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", lambda _session, **kwargs: app_model)
monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, _app: 2)
monkeypatch.setattr(roster_controller, "_agent_app_access_ready", lambda _session, _app: True)
monkeypatch.setattr("models.model.dify_config.SERVICE_API_URL", "https://api.example.test/v1")
apply_config_overrides(monkeypatch, SERVICE_API_URL="https://api.example.test/v1")
response = unwrap(AgentApiAccessApi.get)(AgentApiAccessApi(), MagicMock(), "tenant-1", agent_id)
assert response == {
"access_ready": True,
@ -1777,6 +1864,38 @@ def test_agent_chat_stream_preflight_raises_first_error_event() -> None:
assert stream.closed is True
def test_agent_chat_stream_preflight_preserves_session_configuration_error() -> None:
class ClosableStream:
def __init__(self) -> None:
self.closed = False
self._chunks = iter(
[
"event: ping\n\n",
(
'data: {"event":"error","message":"Start a new conversation to continue.",'
'"code":"agent_session_configuration_changed","status":409}\n\n'
),
]
)
def __iter__(self):
return self
def __next__(self) -> str:
return next(self._chunks)
def close(self) -> None:
self.closed = True
stream = ClosableStream()
with pytest.raises(AgentSessionConfigurationChangedError) as exc_info:
completion_controller._raise_agent_stream_error_before_response(stream)
assert exc_info.value.code == 409
assert exc_info.value.error_code == "agent_session_configuration_changed"
assert "Start a new conversation" in exc_info.value.description
assert stream.closed is True
def test_agent_chat_stream_preflight_preserves_first_normal_event() -> None:
stream = iter(
["event: ping\n\n", 'data: {"event":"message","answer":"hello"}\n\n', 'data: {"event":"message_end"}\n\n']

View File

@ -10,6 +10,7 @@ from core.rbac import RBACPermission, RBACResourceScope
from models import Account
from models.agent import Agent, AgentScope, AgentSource, AgentStatus
from models.model import App, AppMode
from tests.unit_tests.config_override import config_overrides_context
TENANT_ID = "tenant-1"
@ -58,7 +59,7 @@ def _persist_app(
def _patch_guard(account: Account, rbac_enabled: bool):
return (
patch("controllers.console.app.wraps.current_account_with_tenant", return_value=(account, TENANT_ID)),
patch("controllers.console.app.wraps.dify_config.RBAC_ENABLED", rbac_enabled),
config_overrides_context(RBAC_ENABLED=rbac_enabled),
)

View File

@ -77,6 +77,7 @@ from services.app_site_service import (
AppSiteCommandResult,
AppSiteNotFoundError,
)
from tests.unit_tests.config_override import apply_config_overrides
APP_ID = "11111111-1111-1111-1111-111111111111"
TENANT_ID = "22222222-2222-2222-2222-222222222222"
@ -261,8 +262,11 @@ class TestAppEndpoints:
oauth_server.issue_authorization_code.return_value = MagicMock(code="oauth-code-1")
services = MagicMock(oauth_server=oauth_server)
monkeypatch.setattr(app_module.dify_config, "CREATORS_PLATFORM_FEATURES_ENABLED", True)
monkeypatch.setattr(app_module.dify_config, "CREATORS_PLATFORM_OAUTH_CLIENT_ID", "client-1")
apply_config_overrides(
monkeypatch,
CREATORS_PLATFORM_FEATURES_ENABLED=True,
CREATORS_PLATFORM_OAUTH_CLIENT_ID="client-1",
)
monkeypatch.setattr(app_module, "application_services", lambda: services)
monkeypatch.setattr(app_module.AppDslService, "export_dsl", MagicMock(return_value="app: demo"))

View File

@ -21,6 +21,7 @@ from models.model import App, AppMode
from services.app_dsl_service import ImportStatus
from services.entities.dsl_entities import CheckDependenciesResult
from services.entities.feature_entities import SystemFeatureModel, WebAppAuthModel
from tests.unit_tests.config_override import apply_config_overrides
def _unwrap(func):
@ -240,7 +241,7 @@ class TestAppImportApi:
"current_account_with_tenant",
lambda: (_make_account(), "tenant-1"),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
apply_config_overrides(monkeypatch, RBAC_ENABLED=True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
@ -276,7 +277,7 @@ class TestAppImportApi:
"current_account_with_tenant",
lambda: (_make_account(), "tenant-1"),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
apply_config_overrides(monkeypatch, RBAC_ENABLED=True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="import_app",
@ -353,7 +354,7 @@ class TestAppImportConfirmApi:
)
)
monkeypatch.setattr(app_import_module.redis_client, "get", redis_get)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
apply_config_overrides(monkeypatch, RBAC_ENABLED=True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",
@ -397,7 +398,7 @@ class TestAppImportConfirmApi:
b'"name":null,"description":null,"icon_type":null,"icon":null,"icon_background":null}'
),
)
monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True)
apply_config_overrides(monkeypatch, RBAC_ENABLED=True)
app_id = _install_persisting_service_result(
monkeypatch,
method_name="confirm_import",

View File

@ -21,6 +21,7 @@ from models import Account
from models.account import AccountStatus
from models.enums import AppMCPServerStatus
from models.model import App, AppMCPServer, AppMode, IconType
from tests.unit_tests.config_override import config_overrides_context
def _app(
@ -309,7 +310,7 @@ class TestAppMCPServerRefreshController:
current_user = Account(name="Current user", email="user@example.com", status=AccountStatus.ACTIVE)
current_user.id = "account-1"
with (
patch("controllers.common.wraps.dify_config.RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch(
"controllers.common.wraps.current_account_with_tenant",
return_value=(current_user, "tenant-1"),

View File

@ -18,6 +18,7 @@ from libs import login as login_lib
from models import Tenant
from models.account import Account, AccountStatus, TenantAccountRole
from models.model import App, AppMode, IconType
from tests.unit_tests.config_override import apply_config_overrides
def _make_account(role: TenantAccountRole) -> Account:
@ -55,9 +56,12 @@ def _patch_console_guards(
*,
rbac_enabled: bool = False,
) -> None:
monkeypatch.setattr(login_lib.dify_config, "LOGIN_DISABLED", True)
monkeypatch.setattr(login_lib.dify_config, "RBAC_ENABLED", rbac_enabled)
monkeypatch.setattr(console_wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD)
apply_config_overrides(
monkeypatch,
LOGIN_DISABLED=True,
RBAC_ENABLED=rbac_enabled,
DEPLOYMENT_EDITION=DeploymentEdition.CLOUD,
)
monkeypatch.setattr(login_lib, "current_user", account)
monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))

View File

@ -22,6 +22,7 @@ from core.workflow.llm_environment_variable import LLMEnvironmentVariable
from graphon.file import File, FileTransferMethod, FileType
from graphon.variables import SecretVariable, StringVariable
from graphon.variables.variables import RAGPipelineVariable
from tests.unit_tests.config_override import apply_config_overrides
def _make_workflow(**overrides):
@ -77,6 +78,9 @@ def _make_workflow(**overrides):
)
for key, value in overrides.items():
setattr(workflow, key, value)
workflow.get_created_by_account = Mock(return_value=workflow.created_by_account)
workflow.get_updated_by_account = Mock(return_value=workflow.updated_by_account)
workflow.get_tool_published = Mock(return_value=workflow.tool_published)
return workflow
@ -616,6 +620,25 @@ def test_draft_workflow_get_serializes_response_model(monkeypatch: pytest.Monkey
]
def test_published_workflow_get_uses_session_aware_response_source(monkeypatch: pytest.MonkeyPatch) -> None:
workflow = _make_workflow()
session = Mock(spec=Session)
monkeypatch.setattr(workflow_module, "db", SimpleNamespace(session=Mock(return_value=session)))
monkeypatch.setattr(
workflow_module, "WorkflowService", lambda: SimpleNamespace(get_published_workflow=lambda **_kwargs: workflow)
)
api = workflow_module.PublishedWorkflowApi()
handler = inspect.unwrap(api.get)
response = handler(api, app_model=SimpleNamespace(id="app"))
assert response["id"] == "workflow-1"
workflow.get_created_by_account.assert_called_once_with(session=session)
workflow.get_updated_by_account.assert_called_once_with(session=session)
workflow.get_tool_published.assert_called_once_with(session=session)
def test_pipeline_variable_response_accepts_legacy_file_field_names() -> None:
response = workflow_module.PipelineVariableResponse.model_validate(
{
@ -875,7 +898,7 @@ def test_workflow_online_users_filters_inaccessible_workflow(app: Flask, monkeyp
access_filter = SimpleNamespace(is_app_accessible=lambda app_id, _maintainer, _account_id: app_id == app_id_1)
resolve_access = Mock(return_value=access_filter)
monkeypatch.setattr(workflow_module, "resolve_app_access_filter", resolve_access)
monkeypatch.setattr(workflow_module.dify_config, "RBAC_ENABLED", True)
apply_config_overrides(monkeypatch, RBAC_ENABLED=True)
monkeypatch.setattr(workflow_module.file_helpers, "get_signed_file_url", sign_avatar)
short_session = Mock()
monkeypatch.setattr(workflow_module.session_factory, "create_session", lambda: nullcontext(short_session))
@ -966,7 +989,7 @@ def test_workflow_online_users_batches_redis_reads(app: Flask, monkeypatch: pyte
"WorkflowService",
lambda: SimpleNamespace(get_tenant_app_maintainers=lambda app_ids, tenant_id, session: dict.fromkeys(app_ids)),
)
monkeypatch.setattr(workflow_module.dify_config, "RBAC_ENABLED", False)
apply_config_overrides(monkeypatch, RBAC_ENABLED=False)
monkeypatch.setattr(workflow_module.session_factory, "create_session", lambda: nullcontext(Mock()))
first_pipeline = Mock()

View File

@ -18,6 +18,7 @@ from libs import login as login_lib
from models import App, Tenant, WorkflowComment, WorkflowCommentMention, WorkflowCommentReply
from models.account import Account, AccountStatus, TenantAccountRole
from models.model import AppMode, IconType
from tests.unit_tests.config_override import apply_config_overrides
JAN_1_2024_NOON = datetime(2024, 1, 1, 12, 0, 0)
JAN_1_2024_NOON_TS = int(JAN_1_2024_NOON.timestamp())
@ -58,12 +59,11 @@ def _make_app() -> App:
def _patch_console_guards(monkeypatch: pytest.MonkeyPatch, account: Account, app_model: App) -> None:
monkeypatch.setattr(login_lib.dify_config, "LOGIN_DISABLED", True)
apply_config_overrides(monkeypatch, LOGIN_DISABLED=True, DEPLOYMENT_EDITION=DeploymentEdition.CLOUD)
monkeypatch.setattr(login_lib, "current_user", account)
monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(login_lib, "check_csrf_token", lambda *_, **__: None)
monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(console_wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD)
monkeypatch.setattr(app_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(app_wraps, "_load_app_model_from_scoped_session", lambda _app_id: app_model)

View File

@ -15,6 +15,7 @@ from libs import login as login_lib
from models import App, Tenant
from models.account import Account, AccountStatus, TenantAccountRole
from models.model import AppMode, IconType
from tests.unit_tests.config_override import apply_config_overrides
def _make_account() -> Account:
@ -47,14 +48,17 @@ def _make_app(mode: AppMode) -> App:
def _patch_console_guards(monkeypatch: pytest.MonkeyPatch, account: Account, app_model: App) -> None:
# Skip setup and auth guardrails
monkeypatch.setattr("configs.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD)
monkeypatch.setattr(login_lib.dify_config, "LOGIN_DISABLED", True)
apply_config_overrides(
monkeypatch,
DEPLOYMENT_EDITION=DeploymentEdition.CLOUD,
LOGIN_DISABLED=True,
INIT_PASSWORD="",
)
monkeypatch.setattr(login_lib, "current_user", account)
monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(login_lib, "check_csrf_token", lambda *_, **__: None)
monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(app_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
monkeypatch.setattr(console_wraps.dify_config, "INIT_PASSWORD", "")
# Avoid hitting the database when resolving the app model
monkeypatch.setattr(app_wraps, "_load_app_model_from_scoped_session", lambda _app_id: app_model)

View File

@ -20,6 +20,7 @@ from core.workflow.nodes.human_input.pause_reason import HumanInputRequired
from graphon.enums import WorkflowExecutionStatus
from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom
from models.workflow import WorkflowPause, WorkflowRun, WorkflowType
from tests.unit_tests.config_override import apply_config_overrides
@dataclass(frozen=True)
@ -77,7 +78,7 @@ class _PauseEntity:
def test_pause_details_returns_backstage_input_url(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
monkeypatch.setattr(workflow_run_module.dify_config, "APP_WEB_URL", "https://web.example.com")
apply_config_overrides(monkeypatch, APP_WEB_URL="https://web.example.com")
tenant_id = str(uuid4())
run_id = str(uuid4())
@ -137,7 +138,7 @@ def test_pause_details_returns_backstage_input_url(
def test_pause_details_tenant_isolation(app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
monkeypatch.setattr(workflow_run_module.dify_config, "APP_WEB_URL", "https://web.example.com")
apply_config_overrides(monkeypatch, APP_WEB_URL="https://web.example.com")
run_id = str(uuid4())
_persist_run(

View File

@ -12,6 +12,7 @@ from controllers.console.auth.error import AuthenticationFailedError
from controllers.console.auth.login import LoginApi
from enums import DeploymentEdition
from models.account import Account
from tests.unit_tests.config_override import config_overrides_context
def encode_password(password: str) -> str:
@ -35,7 +36,7 @@ class TestAuthenticationSecurity:
@patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit")
@patch("controllers.console.auth.login.AccountService.authenticate")
@patch("controllers.console.auth.login.AccountService.add_login_error_rate_limit")
@patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY)
@config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY)
@patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback")
def test_login_invalid_email_with_registration_allowed(
self, mock_get_invitation, mock_add_rate_limit, mock_authenticate, mock_is_rate_limit, mock_features, mock_db
@ -67,7 +68,7 @@ class TestAuthenticationSecurity:
@patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit")
@patch("controllers.console.auth.login.AccountService.authenticate")
@patch("controllers.console.auth.login.AccountService.add_login_error_rate_limit")
@patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY)
@config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY)
@patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback")
def test_login_wrong_password_returns_error(
self, mock_get_invitation, mock_add_rate_limit, mock_authenticate, mock_is_rate_limit, mock_db
@ -99,7 +100,7 @@ class TestAuthenticationSecurity:
@patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit")
@patch("controllers.console.auth.login.AccountService.authenticate")
@patch("controllers.console.auth.login.AccountService.add_login_error_rate_limit")
@patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY)
@config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY)
@patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback")
def test_login_invalid_email_with_registration_disabled(
self, mock_get_invitation, mock_add_rate_limit, mock_authenticate, mock_is_rate_limit, mock_features, mock_db

View File

@ -19,6 +19,7 @@ from services.data_source_oauth_service import (
InvalidDataSourceOAuthProviderError,
)
from services.entities.data_source_oauth_entities import DataSourceOAuthCallback
from tests.unit_tests.config_override import config_overrides_context
def _request_context() -> RequestContext:
@ -72,10 +73,7 @@ def test_callback_parses_query_and_returns_flask_redirect() -> None:
with (
app.test_request_context("/?code=code-1"),
patch(
"controllers.console.auth.data_source_oauth.dify_config.CONSOLE_WEB_URL",
"https://console.example/root?lang=en#top",
),
config_overrides_context(CONSOLE_WEB_URL="https://console.example/root?lang=en#top"),
patch(
"controllers.console.auth.data_source_oauth.application_services",
return_value=_services(service),

View File

@ -1,30 +1,57 @@
"""Unit tests for email register controller endpoints."""
"""Unit tests for the email-registration Flask adapter."""
from __future__ import annotations
from collections.abc import Callable
from unittest.mock import MagicMock, patch
from collections.abc import Callable, Generator
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from flask import Flask
from pydantic import ValidationError
from controllers.console import bp as console_bp
from controllers.console.auth.email_register import (
EmailRegisterCheckApi,
EmailRegisterResetApi,
EmailRegisterResetPayload,
EmailRegisterSendEmailApi,
)
from controllers.console.auth.error import NormalizedEmailAlreadyInUseError
from controllers.console.error import AccountInFreezeError, EmailDomainSuspendedError
from controllers.console.auth.error import (
EmailAlreadyInUseError,
EmailCodeError,
EmailRegisterLimitError,
EmailRegisterRateLimitExceededError,
InvalidEmailError,
InvalidTokenError,
NormalizedEmailAlreadyInUseError,
PasswordMismatchError,
)
from controllers.console.error import (
AccountInFreezeError,
EmailDomainSuspendedError,
EmailSendIpLimitError,
SeatsLimitExceeded,
)
from enums import DeploymentEdition
from models.account import Account
from services.entities.feature_entities import SystemFeatureModel
from services.errors.account import (
from services.account_email_registration_service import AccountEmailRegistrationService
from services.account_errors import (
AccountEmailAlreadyInUseError,
AccountEmailDomainSuspendedError,
AccountEmailFrozenError,
AccountNormalizedEmailAlreadyInUseError,
AccountRegisterError,
)
from services.errors.account import (
EmailDomainSuspendedError as EmailDomainSuspendedRegistrationError,
EmailRegistrationPasswordMismatchError,
EmailRegistrationSeatsLimitError,
EmailRegistrationSendIPLimitedError,
EmailRegistrationSendRateLimitError,
EmailRegistrationVerificationLimitError,
InvalidEmailRegistrationAddressError,
InvalidEmailRegistrationCodeError,
InvalidEmailRegistrationTokenError,
)
from services.entities.account_entities import AccountEmailRegistrationVerification, AccountSessionTokens
from services.entities.feature_entities import SystemFeatureModel
@pytest.fixture(autouse=True)
@ -32,6 +59,33 @@ def _cloud_edition(config_overrides: Callable[..., None]) -> None:
config_overrides(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD)
@contextmanager
def _request(
app: Flask,
service: Mock,
*,
path: str,
payload: dict[str, str],
) -> Generator[None, None, None]:
services = SimpleNamespace(accounts=SimpleNamespace(email_registration=service))
features = SystemFeatureModel(
deployment_edition=DeploymentEdition.CLOUD,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.auth.email_register.application_services", return_value=services),
patch("controllers.console.flask_admission.FeatureService.get_system_features", return_value=features),
patch("controllers.console.auth.email_register.extract_remote_ip", return_value="127.0.0.1"),
app.test_request_context(path, method="POST", json=payload),
):
yield
def _service() -> Mock:
return Mock(spec=AccountEmailRegistrationService)
def test_normalized_email_conflict_exposes_a_distinct_error_code() -> None:
error = NormalizedEmailAlreadyInUseError()
@ -40,323 +94,210 @@ def test_normalized_email_conflict_exposes_a_distinct_error_code() -> None:
assert error.data["code"] == "normalized_email_already_in_use"
class TestEmailRegisterSendEmailApi:
@patch("controllers.console.auth.email_register.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.email_register.AccountService.send_email_register_email")
@patch("controllers.console.auth.email_register.BillingService.get_email_freeze_type")
@patch("controllers.console.auth.email_register.AccountService.is_email_send_ip_limit", return_value=False)
@patch("controllers.console.auth.email_register.extract_remote_ip", return_value="127.0.0.1")
def test_send_email_normalizes_and_falls_back(
self,
mock_extract_ip,
mock_is_email_send_ip_limit,
mock_is_freeze,
mock_send_mail,
mock_get_account,
app: Flask,
def test_send_email_delegates_with_remote_ip(app: Flask) -> None:
service = _service()
service.send_code.return_value = "token-123"
with _request(
app,
service,
path="/email-register/send-email",
payload={"email": "Invitee@Example.com", "language": "zh-Hans"},
):
mock_send_mail.return_value = "token-123"
mock_is_freeze.return_value = False
account = Account(name="Invitee", email="invitee@example.com")
mock_get_account.return_value = account
response = EmailRegisterSendEmailApi().post()
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
):
with app.test_request_context(
"/email-register/send-email",
method="POST",
json={"email": "Invitee@Example.com", "language": "en-US"},
):
response = EmailRegisterSendEmailApi().post()
assert response == {"result": "success", "data": "token-123"}
assert service.send_code.call_args.kwargs == {
"remote_ip": "127.0.0.1",
"requested_email": "Invitee@Example.com",
"requested_language": "zh-Hans",
}
assert response == {"result": "success", "data": "token-123"}
mock_is_freeze.assert_called_once_with("invitee@example.com")
mock_send_mail.assert_called_once_with(email="invitee@example.com", account=account, language="en-US")
mock_extract_ip.assert_called_once()
mock_is_email_send_ip_limit.assert_called_once_with("127.0.0.1")
@pytest.mark.parametrize(
("freeze_type", "expected_error"),
[
("freeze", AccountInFreezeError),
("email_domain_suspended", EmailDomainSuspendedError),
],
@pytest.mark.parametrize(
("service_error", "http_error"),
[
pytest.param(EmailRegistrationSendIPLimitedError(), EmailSendIpLimitError, id="ip-limit"),
pytest.param(EmailRegistrationSendRateLimitError(1), EmailRegisterRateLimitExceededError, id="send-limit"),
pytest.param(AccountEmailFrozenError(), AccountInFreezeError, id="frozen"),
pytest.param(AccountEmailDomainSuspendedError(), EmailDomainSuspendedError, id="suspended-domain"),
],
)
def test_send_email_translates_application_errors(
app: Flask,
service_error: Exception,
http_error: type[Exception],
) -> None:
service = _service()
service.send_code.side_effect = service_error
with _request(
app,
service,
path="/email-register/send-email",
payload={"email": "invitee@example.com"},
):
with pytest.raises(http_error):
EmailRegisterSendEmailApi().post()
def test_verify_email_code_serializes_application_result(app: Flask) -> None:
service = _service()
service.verify_code.return_value = AccountEmailRegistrationVerification(
email="user@example.com",
token="verified-token",
)
@patch("controllers.console.auth.email_register.BillingService.get_email_freeze_type")
@patch("controllers.console.auth.email_register.AccountService.is_email_send_ip_limit", return_value=False)
@patch("controllers.console.auth.email_register.extract_remote_ip", return_value="127.0.0.1")
def test_send_email_rejects_frozen_email(
self,
mock_extract_ip,
mock_is_email_send_ip_limit,
mock_get_freeze_type,
app: Flask,
freeze_type,
expected_error,
with _request(
app,
service,
path="/email-register/validity",
payload={"email": "User@Example.com", "code": "123456", "token": "pending-token"},
):
mock_get_freeze_type.return_value = freeze_type
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
response = EmailRegisterCheckApi().post()
with (
patch("controllers.console.auth.email_register.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
):
with app.test_request_context(
"/email-register/send-email",
method="POST",
json={"email": "Invitee@Example.com"},
):
with pytest.raises(expected_error):
EmailRegisterSendEmailApi().post()
mock_get_freeze_type.assert_called_once_with("invitee@example.com")
mock_is_email_send_ip_limit.assert_called_once_with("127.0.0.1")
mock_extract_ip.assert_called_once()
class TestEmailRegisterCheckApi:
@patch("controllers.console.auth.email_register.AccountService.reset_email_register_error_rate_limit")
@patch("controllers.console.auth.email_register.AccountService.generate_email_register_token")
@patch("controllers.console.auth.email_register.AccountService.revoke_email_register_token")
@patch("controllers.console.auth.email_register.AccountService.add_email_register_error_rate_limit")
@patch("controllers.console.auth.email_register.AccountService.get_email_register_data")
@patch("controllers.console.auth.email_register.AccountService.is_email_register_error_rate_limit")
def test_validity_normalizes_email_before_checks(
self,
mock_rate_limit_check,
mock_get_data,
mock_add_rate,
mock_revoke,
mock_generate_token,
mock_reset_rate,
app: Flask,
):
mock_rate_limit_check.return_value = False
mock_get_data.return_value = {"email": "User@Example.com", "code": "4321"}
mock_generate_token.return_value = (None, "new-token")
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
):
with app.test_request_context(
"/email-register/validity",
method="POST",
json={"email": "User@Example.com", "code": "4321", "token": "token-123"},
):
response = EmailRegisterCheckApi().post()
assert response == {"is_valid": True, "email": "user@example.com", "token": "new-token"}
mock_rate_limit_check.assert_called_once_with("user@example.com")
mock_generate_token.assert_called_once_with(
"user@example.com", code="4321", additional_data={"phase": "register"}
)
mock_reset_rate.assert_called_once_with("user@example.com")
mock_add_rate.assert_not_called()
mock_revoke.assert_called_once_with("token-123")
class TestEmailRegisterResetApi:
@pytest.mark.parametrize(
("service_error", "expected_error"),
[
(EmailDomainSuspendedRegistrationError(), EmailDomainSuspendedError),
(AccountNormalizedEmailAlreadyInUseError(), NormalizedEmailAlreadyInUseError),
(AccountRegisterError("frozen"), AccountInFreezeError),
],
assert response == {"is_valid": True, "email": "user@example.com", "token": "verified-token"}
service.verify_code.assert_called_once_with(
email="User@Example.com",
code="123456",
token="pending-token",
)
@patch("controllers.console.auth.email_register.AccountService.create_account_and_tenant")
def test_create_new_account_translates_freeze_errors(
self,
mock_create_account,
service_error,
expected_error,
@pytest.mark.parametrize(
("service_error", "http_error"),
[
pytest.param(EmailRegistrationVerificationLimitError(), EmailRegisterLimitError, id="attempt-limit"),
pytest.param(InvalidEmailRegistrationTokenError(), InvalidTokenError, id="token"),
pytest.param(InvalidEmailRegistrationAddressError(), InvalidEmailError, id="email"),
pytest.param(InvalidEmailRegistrationCodeError(), EmailCodeError, id="code"),
],
)
def test_verify_email_code_translates_application_errors(
app: Flask,
service_error: Exception,
http_error: type[Exception],
) -> None:
service = _service()
service.verify_code.side_effect = service_error
with _request(
app,
service,
path="/email-register/validity",
payload={"email": "user@example.com", "code": "wrong", "token": "pending-token"},
):
mock_create_account.side_effect = service_error
with pytest.raises(http_error):
EmailRegisterCheckApi().post()
with pytest.raises(expected_error):
EmailRegisterResetApi()._create_new_account(
email="user@example.com",
password="ValidPass123!",
)
@patch("controllers.console.auth.email_register.AccountService.reset_login_error_rate_limit")
@patch("controllers.console.auth.email_register.AccountService.login")
@patch("controllers.console.auth.email_register.EmailRegisterResetApi._create_new_account")
@patch("controllers.console.auth.email_register.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.email_register.AccountService.revoke_email_register_token")
@patch("controllers.console.auth.email_register.AccountService.get_email_register_data")
@patch("controllers.console.auth.email_register.extract_remote_ip", return_value="127.0.0.1")
def test_reset_creates_account_with_normalized_email(
self,
mock_extract_ip,
mock_get_data,
mock_revoke_token,
mock_get_account,
mock_create_account,
mock_login,
mock_reset_login_rate,
app: Flask,
def test_register_delegates_and_serializes_tokens(app: Flask) -> None:
service = _service()
service.register.return_value = AccountSessionTokens(
access_token="access",
refresh_token="refresh",
csrf_token="csrf",
)
with _request(
app,
service,
path="/email-register",
payload={
"token": "verified-token",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
"language": "zh-Hans",
"timezone": "Asia/Shanghai",
},
):
mock_get_data.return_value = {"phase": "register", "email": "Invitee@Example.com"}
mock_create_account.return_value = Account(name="Invitee", email="invitee@example.com")
token_pair = MagicMock()
token_pair.model_dump.return_value = {"access_token": "a", "refresh_token": "r"}
mock_login.return_value = token_pair
mock_get_account.return_value = None
response = EmailRegisterResetApi().post()
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
):
with app.test_request_context(
"/email-register",
method="POST",
json={"token": "token-123", "new_password": "ValidPass123!", "password_confirm": "ValidPass123!"},
):
response = EmailRegisterResetApi().post()
assert response == {
"result": "success",
"data": {"access_token": "access", "refresh_token": "refresh", "csrf_token": "csrf"},
}
assert service.register.call_args.kwargs == {
"remote_ip": "127.0.0.1",
"token": "verified-token",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
"language": "zh-Hans",
"timezone": "Asia/Shanghai",
}
assert response == {"result": "success", "data": {"access_token": "a", "refresh_token": "r"}}
mock_create_account.assert_called_once_with(
email="invitee@example.com",
password="ValidPass123!",
timezone=None,
language=None,
ip_address="127.0.0.1",
)
mock_reset_login_rate.assert_called_once_with("invitee@example.com")
mock_revoke_token.assert_called_once_with("token-123")
mock_extract_ip.assert_called_once()
@patch("controllers.console.auth.email_register.AccountService.reset_login_error_rate_limit")
@patch("controllers.console.auth.email_register.AccountService.login")
@patch("controllers.console.auth.email_register.EmailRegisterResetApi._create_new_account")
@patch("controllers.console.auth.email_register.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.email_register.AccountService.revoke_email_register_token")
@patch("controllers.console.auth.email_register.AccountService.get_email_register_data")
@patch("controllers.console.auth.email_register.extract_remote_ip", return_value="127.0.0.1")
def test_reset_passes_timezone_to_new_account(
self,
mock_extract_ip,
mock_get_data,
mock_revoke_token,
mock_get_account,
mock_create_account,
mock_login,
mock_reset_login_rate,
app: Flask,
@pytest.mark.parametrize(
("service_error", "http_error"),
[
pytest.param(EmailRegistrationPasswordMismatchError(), PasswordMismatchError, id="password"),
pytest.param(InvalidEmailRegistrationTokenError(), InvalidTokenError, id="token"),
pytest.param(
AccountNormalizedEmailAlreadyInUseError(),
NormalizedEmailAlreadyInUseError,
id="normalized-email-in-use",
),
pytest.param(AccountEmailAlreadyInUseError(), EmailAlreadyInUseError, id="email-in-use"),
pytest.param(EmailRegistrationSeatsLimitError(), SeatsLimitExceeded, id="seat-limit"),
pytest.param(AccountEmailFrozenError(), AccountInFreezeError, id="frozen"),
pytest.param(AccountEmailDomainSuspendedError(), EmailDomainSuspendedError, id="suspended-domain"),
],
)
def test_register_translates_application_errors(
app: Flask,
service_error: Exception,
http_error: type[Exception],
) -> None:
service = _service()
service.register.side_effect = service_error
with _request(
app,
service,
path="/email-register",
payload={
"token": "verified-token",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
},
):
mock_get_data.return_value = {"phase": "register", "email": "Invitee@Example.com"}
mock_create_account.return_value = Account(name="Invitee", email="invitee@example.com")
token_pair = MagicMock()
token_pair.model_dump.return_value = {"access_token": "a", "refresh_token": "r"}
mock_login.return_value = token_pair
mock_get_account.return_value = None
with pytest.raises(http_error):
EmailRegisterResetApi().post()
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
def test_reset_payload_rejects_invalid_timezone() -> None:
with pytest.raises(ValidationError):
EmailRegisterResetPayload.model_validate(
{
"token": "token-123",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
"timezone": "",
}
)
with (
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
):
with app.test_request_context(
"/email-register",
method="POST",
json={
"token": "token-123",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
"timezone": "Asia/Shanghai",
},
):
response = EmailRegisterResetApi().post()
assert response == {"result": "success", "data": {"access_token": "a", "refresh_token": "r"}}
mock_create_account.assert_called_once_with(
email="invitee@example.com",
password="ValidPass123!",
timezone="Asia/Shanghai",
language=None,
ip_address="127.0.0.1",
def test_invalid_password_is_sanitized_by_real_error_handler(caplog: pytest.LogCaptureFixture) -> None:
app = Flask(__name__)
app.config["TESTING"] = True
app.register_blueprint(console_bp)
features = SystemFeatureModel(
deployment_edition=DeploymentEdition.CLOUD,
enable_email_password_login=True,
is_allow_register=True,
)
password_marker = "SecretMarker"
with patch("controllers.console.flask_admission.FeatureService.get_system_features", return_value=features):
response = app.test_client().post(
"/console/api/email-register",
json={
"token": "verified-token",
"new_password": password_marker,
"password_confirm": password_marker,
},
)
mock_reset_login_rate.assert_called_once_with("invitee@example.com")
mock_revoke_token.assert_called_once_with("token-123")
mock_extract_ip.assert_called_once()
@patch("controllers.console.auth.email_register.AccountService.reset_login_error_rate_limit")
@patch("controllers.console.auth.email_register.AccountService.login")
@patch("controllers.console.auth.email_register.EmailRegisterResetApi._create_new_account")
@patch("controllers.console.auth.email_register.AccountService.get_account_by_email_with_case_fallback")
@patch("controllers.console.auth.email_register.AccountService.revoke_email_register_token")
@patch("controllers.console.auth.email_register.AccountService.get_email_register_data")
@patch("controllers.console.auth.email_register.extract_remote_ip", return_value="127.0.0.1")
def test_reset_passes_language_to_new_account(
self,
mock_extract_ip,
mock_get_data,
mock_revoke_token,
mock_get_account,
mock_create_account,
mock_login,
mock_reset_login_rate,
app: Flask,
):
mock_get_data.return_value = {"phase": "register", "email": "Invitee@Example.com"}
mock_create_account.return_value = Account(name="Invitee", email="invitee@example.com")
token_pair = MagicMock()
token_pair.model_dump.return_value = {"access_token": "a", "refresh_token": "r"}
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
):
with app.test_request_context(
"/email-register",
method="POST",
json={
"token": "token-123",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
"language": "zh-Hans",
},
):
response = EmailRegisterResetApi().post()
assert response == {"result": "success", "data": {"access_token": "a", "refresh_token": "r"}}
mock_create_account.assert_called_once_with(
email="invitee@example.com",
password="ValidPass123!",
timezone=None,
language="zh-Hans",
ip_address="127.0.0.1",
)
mock_reset_login_rate.assert_called_once_with("invitee@example.com")
mock_revoke_token.assert_called_once_with("token-123")
mock_extract_ip.assert_called_once()
assert response.status_code == 422
assert password_marker not in response.get_data(as_text=True)
assert password_marker not in caplog.text

View File

@ -1,44 +0,0 @@
from unittest.mock import ANY, patch
import pytest
from pydantic import ValidationError
from controllers.console.auth.email_register import EmailRegisterResetApi, EmailRegisterResetPayload
from models.account import Account
@patch("controllers.console.auth.email_register.AccountService.create_account_and_tenant")
def test_create_new_account_uses_requested_language(mock_create_account):
account = Account(name="Invitee", email="invitee@example.com")
mock_create_account.return_value = account
result = EmailRegisterResetApi()._create_new_account(
"invitee@example.com",
"ValidPass123!",
timezone="Asia/Shanghai",
language="zh-Hans",
)
assert result is account
mock_create_account.assert_called_once_with(
email="invitee@example.com",
name="invitee@example.com",
password="ValidPass123!",
interface_language="zh-Hans",
timezone="Asia/Shanghai",
ip_address=None,
check_normalized_email=True,
session=ANY,
)
def test_reset_payload_rejects_invalid_timezone():
with pytest.raises(ValidationError):
EmailRegisterResetPayload.model_validate(
{
"token": "token-123",
"new_password": "ValidPass123!",
"password_confirm": "ValidPass123!",
"timezone": "",
}
)

View File

@ -461,14 +461,17 @@ class TestEmailCodeLoginApi:
mock_verify_challenge,
mock_db,
app: Flask,
config_overrides: Callable[..., None],
):
config_overrides(
DEPLOYMENT_EDITION=DeploymentEdition.CLOUD,
TURNSTILE_EMAIL_CODE_VERIFY_REQUIRED=True,
)
mock_verify_challenge.return_value = EmailCodeLoginChallengeResult(
status=EmailCodeLoginChallengeStatus.INVALID_TOKEN
)
with (
patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
patch("controllers.console.auth.login.dify_config.TURNSTILE_EMAIL_CODE_VERIFY_REQUIRED", True),
app.test_request_context(
"/email-code-login/validity",
method="POST",
@ -502,10 +505,13 @@ class TestEmailCodeLoginApi:
mock_verify_challenge,
mock_db,
app: Flask,
config_overrides: Callable[..., None],
):
config_overrides(
DEPLOYMENT_EDITION=DeploymentEdition.CLOUD,
TURNSTILE_EMAIL_CODE_VERIFY_REQUIRED=True,
)
with (
patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
patch("controllers.console.auth.login.dify_config.TURNSTILE_EMAIL_CODE_VERIFY_REQUIRED", True),
app.test_request_context(
"/email-code-login/validity",
method="POST",
@ -527,14 +533,17 @@ class TestEmailCodeLoginApi:
mock_verify_challenge,
mock_db,
app: Flask,
config_overrides: Callable[..., None],
):
config_overrides(
DEPLOYMENT_EDITION=DeploymentEdition.CLOUD,
TURNSTILE_EMAIL_CODE_VERIFY_REQUIRED=False,
)
mock_verify_challenge.return_value = EmailCodeLoginChallengeResult(
status=EmailCodeLoginChallengeStatus.INVALID_TOKEN
)
with (
patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
patch("controllers.console.auth.login.dify_config.TURNSTILE_EMAIL_CODE_VERIFY_REQUIRED", False),
app.test_request_context(
"/email-code-login/validity",
method="POST",

View File

@ -17,6 +17,7 @@ from enums import DeploymentEdition
from models.account import Account
from models.engine import db
from services.entities.feature_entities import SystemFeatureModel
from tests.unit_tests.config_override import config_overrides_context
@pytest.fixture
@ -61,7 +62,7 @@ class TestForgotPasswordSendEmailApi:
"controllers.console.auth.forgot_password.FeatureService.get_system_features",
return_value=controller_features,
),
patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
):
with app.test_request_context(
@ -108,7 +109,7 @@ class TestForgotPasswordCheckApi:
enable_email_password_login=True,
)
with (
patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
):
with app.test_request_context(
@ -154,7 +155,7 @@ class TestForgotPasswordResetApi:
enable_email_password_login=True,
)
with (
patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
):
with database_app.test_request_context(

View File

@ -22,6 +22,7 @@ from services.errors.account import AccountRegisterError
from services.errors.account import (
EmailDomainSuspendedError as EmailDomainSuspendedRegistrationError,
)
from tests.unit_tests.config_override import config_overrides_context
@pytest.fixture(autouse=True)
@ -586,7 +587,7 @@ class TestAccountGeneration:
("freeze", AccountRegisterError),
],
)
@patch("controllers.console.auth.oauth.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD)
@config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD)
@patch("controllers.console.auth.oauth.BillingService.get_email_freeze_type")
@patch("controllers.console.auth.oauth._get_account_by_openid_or_email", return_value=None)
@patch("controllers.console.auth.oauth.FeatureService")

View File

@ -7,6 +7,7 @@ from flask import Flask
from controllers.console.auth.oauth import OAuthCallback, OAuthLogin
from libs.oauth import OAuthUserInfo, encode_oauth_state
from models.account import Account, AccountStatus, Tenant
from tests.unit_tests.config_override import config_overrides_context
REDIRECT_URL = "/apps?category=workflow"
CONSOLE_WEB_URL = "https://console.example.com"
@ -74,7 +75,7 @@ def test_oauth_callback_validates_redirect_url_and_appends_new_user_flag(
with (
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL),
config_overrides_context(CONSOLE_WEB_URL=CONSOLE_WEB_URL),
patch("controllers.console.auth.oauth._generate_account", return_value=(account, oauth_new_user)),
patch("controllers.console.auth.oauth.TenantService.create_owner_tenant_if_not_exist"),
patch("controllers.console.auth.oauth.AccountService.login", return_value=token_pair),
@ -109,7 +110,7 @@ def test_oauth_callback_with_invitation_establishes_console_session(app: Flask)
with (
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL),
config_overrides_context(CONSOLE_WEB_URL=CONSOLE_WEB_URL),
patch("controllers.console.auth.oauth.RegisterService") as register_service,
patch("controllers.console.auth.oauth.AccountService.link_account_integrate") as link_account,
patch("controllers.console.auth.oauth.AccountService.login", return_value=token_pair) as login,
@ -155,7 +156,7 @@ def test_oauth_callback_with_invitation_rejects_another_account(app: Flask) -> N
with (
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL),
config_overrides_context(CONSOLE_WEB_URL=CONSOLE_WEB_URL),
patch("controllers.console.auth.oauth.RegisterService") as register_service,
patch("controllers.console.auth.oauth.AccountService.link_account_integrate") as link_account,
patch("controllers.console.auth.oauth.AccountService.login") as login,

View File

@ -26,6 +26,7 @@ from controllers.console.error import AccountNotFound, EmailSendIpLimitError
from enums import DeploymentEdition
from models.account import Account, Tenant, TenantAccountJoin
from services.entities.feature_entities import SystemFeatureModel
from tests.unit_tests.config_override import apply_config_overrides
SQLITE_MODELS = (Account, Tenant, TenantAccountJoin)
@ -46,7 +47,7 @@ def _bind_database_session(session: Session) -> Generator[scoped_session[Session
def enable_password_login_wrappers(monkeypatch: pytest.MonkeyPatch) -> None:
"""Keep endpoint decorators deterministic without requiring the configured app database."""
monkeypatch.setattr("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD)
apply_config_overrides(monkeypatch, DEPLOYMENT_EDITION=DeploymentEdition.CLOUD)
monkeypatch.setattr(
"controllers.console.wraps.FeatureService.get_system_features",
lambda: SystemFeatureModel(

View File

@ -24,6 +24,7 @@ from services.errors.billing import (
BillingUpstreamInvalidResponseError,
BillingUpstreamUnavailableError,
)
from tests.unit_tests.config_override import config_overrides_context
class TestBillingPortal:
@ -188,8 +189,7 @@ class TestPartnerTenants:
console_wraps._is_setup_completed.reset_success()
monkeypatch.setattr(console_wraps.db, "session", sqlite_session)
with (
patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD),
patch("libs.login.dify_config.LOGIN_DISABLED", False),
config_overrides_context(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD, LOGIN_DISABLED=False),
patch("libs.login.check_csrf_token") as mock_csrf,
):
mock_csrf.return_value = None

View File

@ -27,6 +27,7 @@ from models.engine import db
from services.entities.knowledge_entities.rag_pipeline_entities import PipelineTemplateInfoEntity
from services.errors.account import NoPermissionError
from services.errors.rag_pipeline import RagPipelineResourceNotFoundError
from tests.unit_tests.config_override import config_overrides_context
def _template_item() -> dict[str, object]:
@ -376,7 +377,7 @@ class TestPublishCustomizedPipelineTemplateApi:
dataset = object()
with (
patch.object(module.dify_config, "RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch.object(Pipeline, "retrieve_dataset", return_value=dataset),
patch.object(module.DatasetService, "check_dataset_permission") as legacy_acl,
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
@ -409,7 +410,7 @@ class TestPublishCustomizedPipelineTemplateApi:
dataset = object()
with (
patch.object(module.dify_config, "RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch.object(Pipeline, "retrieve_dataset", return_value=dataset),
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
):
@ -425,7 +426,7 @@ class TestPublishCustomizedPipelineTemplateApi:
payload = _payload()
with (
patch.object(module.dify_config, "RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch.object(Pipeline, "retrieve_dataset", return_value=object()),
patch.object(
module.RagPipelineService,
@ -446,7 +447,7 @@ class TestPublishCustomizedPipelineTemplateApi:
payload = _payload()
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(Pipeline, "retrieve_dataset", return_value=dataset),
patch.object(module.DatasetService, "check_dataset_permission") as check_permission,
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
@ -464,7 +465,7 @@ class TestPublishCustomizedPipelineTemplateApi:
account.role = TenantAccountRole.NORMAL
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(Pipeline, "retrieve_dataset", return_value=object()),
patch.object(module.DatasetService, "check_dataset_permission") as check_permission,
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
@ -482,7 +483,7 @@ class TestPublishCustomizedPipelineTemplateApi:
account.role = TenantAccountRole.EDITOR
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(Pipeline, "retrieve_dataset", return_value=object()),
patch.object(
module.DatasetService,
@ -503,7 +504,7 @@ class TestPublishCustomizedPipelineTemplateApi:
account.role = TenantAccountRole.EDITOR
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(Pipeline, "retrieve_dataset", return_value=None),
patch.object(module.DatasetService, "check_dataset_permission") as check_permission,
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,

View File

@ -34,6 +34,7 @@ from models.workflow import Workflow, WorkflowType
from services.errors.llm import InvokeRateLimitError
from services.errors.rag_pipeline import RagPipelineResourceNotFoundError
from services.rag_pipeline.rag_pipeline import RagPipelineService
from tests.unit_tests.config_override import config_overrides_context
DEFAULT_WORKFLOW_TENANT_ID = "00000000-0000-0000-0000-000000000001"
DEFAULT_WORKFLOW_APP_ID = "00000000-0000-0000-0000-000000000002"
@ -259,7 +260,7 @@ def test_rag_pipeline_transform_rejects_read_only_member(sqlite_engine: Engine)
session.add(_dataset())
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
pytest.raises(Forbidden),
):
handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID))
@ -293,7 +294,7 @@ def test_rag_pipeline_transform_enforces_legacy_dataset_permission_before_servic
session.add(_dataset(maintainer="00000000-0000-0000-0000-000000000099"))
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(module.RagPipelineTransformService, "transform_dataset") as transform_dataset,
pytest.raises(Forbidden),
):
@ -315,7 +316,7 @@ def test_rag_pipeline_transform_passes_authorized_dataset_and_account_to_service
session.add(dataset)
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(module.RagPipelineTransformService, "transform_dataset", return_value=expected) as transform,
):
response = handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID))
@ -333,7 +334,7 @@ def test_rag_pipeline_transform_maps_missing_pipeline_to_not_found(sqlite_engine
session.add(_dataset())
with (
patch.object(module.dify_config, "RBAC_ENABLED", False),
config_overrides_context(RBAC_ENABLED=False),
patch.object(
module.RagPipelineTransformService,
"transform_dataset",
@ -355,7 +356,7 @@ def test_rag_pipeline_transform_skips_legacy_acl_when_rbac_is_enabled(sqlite_eng
session.add(_dataset(maintainer="00000000-0000-0000-0000-000000000099"))
with (
patch.object(module.dify_config, "RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch.object(module.RagPipelineTransformService, "transform_dataset", return_value=expected) as transform,
):
response = handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID))

View File

@ -8,6 +8,7 @@ from unittest.mock import ANY, MagicMock, PropertyMock, call, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, Forbidden, NotFound
import services
@ -44,7 +45,7 @@ from core.rag.index_processor.constant.index_type import IndexStructureType
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from extensions.storage.storage_type import StorageType
from models.account import Account, TenantAccountRole
from models.dataset import Dataset, DatasetQuery, Document
from models.dataset import AppDatasetJoin, Dataset, DatasetPermission, DatasetQuery, Document, DocumentSegment
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, IndexingStatus
from models.model import ApiToken, App, AppMode, IconType, UploadFile
from services.dataset_ref_service import DatasetRef
@ -170,7 +171,29 @@ def make_document_status(**overrides) -> Document:
return Document(**base)
class TestDatasetList:
def make_document_segment(*, position: int, completed: bool) -> DocumentSegment:
return DocumentSegment(
tenant_id="tenant-1",
dataset_id="dataset-1",
document_id="doc-1",
position=position,
content=f"segment {position}",
word_count=2,
tokens=2,
created_by="account-1",
completed_at=datetime.datetime(2024, 1, 1, tzinfo=datetime.UTC) if completed else None,
)
class _UsesSQLiteSession:
session: Session
@pytest.fixture(autouse=True)
def _inject_sqlite_session(self, sqlite_session: Session) -> None:
self.session = sqlite_session
class TestDatasetList(_UsesSQLiteSession):
def _mock_user(self):
user = make_account()
return user
@ -185,7 +208,7 @@ class TestDatasetList:
patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
assert status == 200
assert resp["total"] == 1
assert resp["data"][0]["embedding_available"] is True
@ -201,7 +224,7 @@ class TestDatasetList:
method = unwrap(api.get)
current_user = self._mock_user()
dataset = make_dataset()
session = MagicMock()
session = self.session
with app.test_request_context("/datasets"):
with (
patch.object(DatasetService, "get_datasets", return_value=([dataset], 1)),
@ -222,7 +245,7 @@ class TestDatasetList:
patch.object(DatasetService, "get_datasets_by_ids", return_value=(datasets, 2)) as by_ids_mock,
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
by_ids_mock.assert_called_once()
assert status == 200
assert resp["total"] == 2
@ -251,7 +274,7 @@ class TestDatasetList:
return_value=permissions,
) as get_permissions,
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
get_permissions.assert_called_once_with("tenant-1", current_user.id, session=ANY)
assert status == 200
assert resp["data"][0]["permission_keys"] == ["dataset.acl.readonly", "dataset.acl.edit"]
@ -281,7 +304,7 @@ class TestDatasetList:
),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
method(api, MagicMock(), "tenant-1", current_user)
method(api, self.session, "tenant-1", current_user)
assert get_datasets.call_args.kwargs["accessible_dataset_ids"] == []
assert get_datasets.call_args.kwargs["include_own_datasets"] is False
@ -308,7 +331,7 @@ class TestDatasetList:
),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
method(api, MagicMock(), "tenant-1", current_user)
method(api, self.session, "tenant-1", current_user)
assert get_datasets.call_args.kwargs["accessible_dataset_ids"] is None
def test_get_restricted_whitelist_overrides_default_read_permission(
@ -374,7 +397,7 @@ class TestDatasetList:
),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
method(api, MagicMock(), "tenant-1", current_user)
method(api, self.session, "tenant-1", current_user)
assert get_datasets.call_args.kwargs["accessible_dataset_ids"] == [
"dataset-whitelist-only",
]
@ -399,9 +422,9 @@ class TestDatasetList:
),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
method(api, MagicMock(), "tenant-1", current_user)
method(api, self.session, "tenant-1", current_user)
session = get_datasets_by_ids.call_args.kwargs["session"]
assert isinstance(session, MagicMock)
assert session is self.session
assert get_datasets_by_ids.call_args.args == (["dataset-1"], "tenant-1")
assert get_datasets_by_ids.call_args.kwargs == {
"user": current_user,
@ -420,7 +443,7 @@ class TestDatasetList:
patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
assert status == 200
def test_get_allows_legacy_weighted_score_without_weight_type(self, app: Flask):
@ -453,7 +476,7 @@ class TestDatasetList:
patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
assert status == 200
assert resp["data"][0]["retrieval_model_dict"]["weights"]["weight_type"] is None
@ -467,7 +490,7 @@ class TestDatasetList:
patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)),
patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])),
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
assert status == 200
retrieval_model = resp["data"][0]["retrieval_model_dict"]
assert retrieval_model["search_method"] == "semantic_search"
@ -491,7 +514,7 @@ class TestDatasetList:
patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)),
patch.object(ProviderManager, "get_configurations", return_value=config),
):
resp, status = method(api, MagicMock(), "tenant-1", current_user)
resp, status = method(api, self.session, "tenant-1", current_user)
assert resp["data"][0]["embedding_available"] is False
def test_partial_members_permission(self, app: Flask):
@ -499,8 +522,9 @@ class TestDatasetList:
method = unwrap(api.get)
current_user = self._mock_user()
datasets = [make_dataset(permission="partial_members")]
session = MagicMock()
session.execute.return_value.all.return_value = [("ds-1", "u1")]
session = self.session
session.add(DatasetPermission(dataset_id="ds-1", account_id="u1", tenant_id="tenant-1"))
session.flush()
with app.test_request_context("/datasets"):
with (
patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)),
@ -510,7 +534,7 @@ class TestDatasetList:
assert resp["data"][0]["partial_member_list"] == ["u1"]
class TestDatasetListApiPost:
class TestDatasetListApiPost(_UsesSQLiteSession):
def test_post_success(self, app: Flask):
api = DatasetListApi()
method = unwrap(api.post)
@ -522,7 +546,7 @@ class TestDatasetListApiPost:
patch.object(type(console_ns), "payload", payload),
patch.object(DatasetService, "create_empty_dataset", return_value=dataset),
):
_, status = method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", user)
_, status = method(api, DatasetCreatePayload(**payload), self.session, "tenant-1", user)
assert status == 201
def test_post_forbidden(self, app: Flask):
@ -532,7 +556,7 @@ class TestDatasetListApiPost:
user = make_account(TenantAccountRole.NORMAL)
with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload):
with pytest.raises(Forbidden):
method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", user)
method(api, DatasetCreatePayload(**payload), self.session, "tenant-1", user)
def test_post_duplicate_name(self, app: Flask):
api = DatasetListApi()
@ -547,14 +571,14 @@ class TestDatasetListApiPost:
),
):
with pytest.raises(DatasetNameDuplicateError):
method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", user)
method(api, DatasetCreatePayload(**payload), self.session, "tenant-1", user)
def test_post_invalid_payload_missing_name(self, app: Flask):
api = DatasetListApi()
method = unwrap(api.post)
with app.test_request_context("/datasets", json={}), patch.object(type(console_ns), "payload", {}):
with pytest.raises(ValueError):
method(api, DatasetCreatePayload(), MagicMock(), "tenant-1", make_account())
method(api, DatasetCreatePayload(), self.session, "tenant-1", make_account())
def test_post_invalid_indexing_technique(self, app: Flask):
api = DatasetListApi()
@ -562,7 +586,7 @@ class TestDatasetListApiPost:
payload = {"name": "bad", "indexing_technique": "invalid-tech"}
with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload):
with pytest.raises(ValueError, match="Invalid indexing technique"):
method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", make_account())
method(api, DatasetCreatePayload(**payload), self.session, "tenant-1", make_account())
def test_post_invalid_provider(self, app: Flask):
api = DatasetListApi()
@ -570,10 +594,10 @@ class TestDatasetListApiPost:
payload = {"name": "bad", "provider": "unknown"}
with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload):
with pytest.raises(ValueError, match="Invalid provider"):
method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", make_account())
method(api, DatasetCreatePayload(**payload), self.session, "tenant-1", make_account())
class TestDatasetApiGet:
class TestDatasetApiGet(_UsesSQLiteSession):
def test_get_success_basic(self, app: Flask):
api = DatasetApi()
method = unwrap(api.get)
@ -588,7 +612,7 @@ class TestDatasetApiGet:
patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock,
):
provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = []
data, status = method(api, MagicMock(), tenant_id, user, dataset_id)
data, status = method(api, self.session, tenant_id, user, dataset_id)
assert status == 200
assert data["embedding_available"] is True
@ -597,7 +621,7 @@ class TestDatasetApiGet:
api = DatasetApi()
method = unwrap(api.get)
dataset_id = "123e4567-e89b-12d3-a456-426614174000"
user = MagicMock(id="account-1")
user = make_account()
tenant_id = "tenant-1"
dataset = make_dataset(id=dataset_id)
with (
@ -619,7 +643,7 @@ class TestDatasetApiGet:
patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock,
):
provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = []
data, status = method(api, MagicMock(), tenant_id, user, dataset_id)
data, status = method(api, self.session, tenant_id, user, dataset_id)
get_permissions.assert_called_once_with(tenant_id, user.id, dataset_id=dataset_id, session=ANY)
assert status == 200
assert data["permission_keys"] == ["dataset.acl.readonly", "dataset.acl.edit"]
@ -636,7 +660,7 @@ class TestDatasetApiGet:
patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock,
):
provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = []
data, status = method(api, MagicMock(), "tenant", make_account(), dataset_id)
data, status = method(api, self.session, "tenant", make_account(), dataset_id)
assert status == 200
assert data["external_retrieval_model"] == {"top_k": 2, "score_threshold": 0.0, "score_threshold_enabled": None}
@ -649,7 +673,7 @@ class TestDatasetApiGet:
patch.object(DatasetService, "get_dataset", return_value=None),
):
with pytest.raises(NotFound, match="Dataset not found"):
method(api, MagicMock(), "tenant", make_account(), dataset_id)
method(api, self.session, "tenant", make_account(), dataset_id)
def test_get_permission_denied(self, app: Flask):
api = DatasetApi()
@ -666,7 +690,7 @@ class TestDatasetApiGet:
),
):
with pytest.raises(Forbidden, match="no access"):
method(api, MagicMock(), "tenant", make_account(), dataset_id)
method(api, self.session, "tenant", make_account(), dataset_id)
def test_get_high_quality_embedding_unavailable(self, app: Flask):
api = DatasetApi()
@ -687,7 +711,7 @@ class TestDatasetApiGet:
patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock,
):
provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = []
data, _ = method(api, MagicMock(), tenant_id, user, dataset_id)
data, _ = method(api, self.session, tenant_id, user, dataset_id)
assert data["embedding_available"] is False
def test_get_partial_members_permission(self, app: Flask):
@ -704,11 +728,11 @@ class TestDatasetApiGet:
patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock,
):
provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = []
data, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id)
data, _ = method(api, self.session, "tenant", make_account(), dataset_id)
assert data["partial_member_list"] == partial_members
class TestDatasetApiPatch:
class TestDatasetApiPatch(_UsesSQLiteSession):
def test_patch_success_basic(self, app: Flask):
api = DatasetApi()
method = unwrap(api.patch)
@ -725,7 +749,7 @@ class TestDatasetApiPatch:
patch.object(DatasetService, "update_dataset", return_value=dataset),
patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=[]),
):
result, status = method(api, DatasetUpdatePayload(), MagicMock(), tenant_id, user, dataset_id)
result, status = method(api, DatasetUpdatePayload(), self.session, tenant_id, user, dataset_id)
assert status == 200
assert result["partial_member_list"] == []
@ -737,7 +761,7 @@ class TestDatasetApiPatch:
patch.object(DatasetService, "get_dataset", return_value=None),
):
with pytest.raises(NotFound, match="Dataset not found"):
method(api, DatasetUpdatePayload(), MagicMock(), "tenant-1", make_account(), "missing")
method(api, DatasetUpdatePayload(), self.session, "tenant-1", make_account(), "missing")
def test_patch_permission_denied(self, app: Flask):
api = DatasetApi()
@ -752,7 +776,7 @@ class TestDatasetApiPatch:
patch.object(DatasetPermissionService, "check_permission", side_effect=Forbidden("no permission")),
):
with pytest.raises(Forbidden):
method(api, DatasetUpdatePayload(), MagicMock(), "tenant", make_account(), dataset_id)
method(api, DatasetUpdatePayload(), self.session, "tenant", make_account(), dataset_id)
def test_patch_partial_members_update(self, app: Flask):
api = DatasetApi()
@ -769,7 +793,7 @@ class TestDatasetApiPatch:
patch.object(DatasetPermissionService, "update_partial_member_list", return_value=None),
patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["u1", "u2"]),
):
result, _ = method(api, DatasetUpdatePayload(), MagicMock(), "tenant", make_account(), dataset_id)
result, _ = method(api, DatasetUpdatePayload(), self.session, "tenant", make_account(), dataset_id)
assert result["partial_member_list"] == ["u1", "u2"]
def test_patch_clear_partial_members(self, app: Flask):
@ -787,11 +811,11 @@ class TestDatasetApiPatch:
patch.object(DatasetPermissionService, "clear_partial_member_list", return_value=None),
patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=[]),
):
result, _ = method(api, DatasetUpdatePayload(), MagicMock(), "tenant", make_account(), dataset_id)
result, _ = method(api, DatasetUpdatePayload(), self.session, "tenant", make_account(), dataset_id)
assert result["partial_member_list"] == []
class TestDatasetApiDelete:
class TestDatasetApiDelete(_UsesSQLiteSession):
def test_delete_success(self, app: Flask):
api = DatasetApi()
method = unwrap(api.delete)
@ -802,7 +826,7 @@ class TestDatasetApiDelete:
patch.object(DatasetService, "delete_dataset", return_value=True),
patch.object(DatasetPermissionService, "clear_partial_member_list", return_value=None),
):
result, status = method(api, MagicMock(), user, dataset_id)
result, status = method(api, self.session, user, dataset_id)
assert status == 204
assert result == ""
@ -813,7 +837,7 @@ class TestDatasetApiDelete:
user = make_account(TenantAccountRole.NORMAL)
with app.test_request_context(f"/datasets/{dataset_id}"):
with pytest.raises(Forbidden):
method(api, MagicMock(), user, dataset_id)
method(api, self.session, user, dataset_id)
def test_delete_dataset_not_found(self, app: Flask):
api = DatasetApi()
@ -825,7 +849,7 @@ class TestDatasetApiDelete:
patch.object(DatasetService, "delete_dataset", return_value=False),
):
with pytest.raises(NotFound, match="Dataset not found"):
method(api, MagicMock(), user, dataset_id)
method(api, self.session, user, dataset_id)
def test_delete_dataset_in_use(self, app: Flask):
api = DatasetApi()
@ -837,10 +861,10 @@ class TestDatasetApiDelete:
patch.object(DatasetService, "delete_dataset", side_effect=services.errors.dataset.DatasetInUseError()),
):
with pytest.raises(DatasetInUseError):
method(api, MagicMock(), user, dataset_id)
method(api, self.session, user, dataset_id)
class TestDatasetUseCheckApi:
class TestDatasetUseCheckApi(_UsesSQLiteSession):
@pytest.mark.parametrize("is_using", [True, False])
def test_get_use_check(self, app: Flask, is_using: bool):
api = DatasetUseCheckApi()
@ -848,7 +872,7 @@ class TestDatasetUseCheckApi:
dataset_id = "dataset-id"
dataset = make_dataset(id=dataset_id)
current_user = make_account()
session = MagicMock()
session = self.session
with (
app.test_request_context(f"/datasets/{dataset_id}/use-check"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
@ -867,7 +891,7 @@ class TestDatasetUseCheckApi:
api = DatasetUseCheckApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-id")
session = MagicMock()
session = self.session
with (
app.test_request_context("/datasets/dataset-id/use-check"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
@ -884,11 +908,11 @@ class TestDatasetUseCheckApi:
"api_cls",
[DatasetUseCheckApi, DatasetIndexingStatusApi, DatasetErrorDocs, DatasetAutoDisableLogApi],
)
def test_dataset_scoped_read_permission_denied(app: Flask, api_cls):
def test_dataset_scoped_read_permission_denied(app: Flask, api_cls, sqlite_session: Session):
api = api_cls()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
session = MagicMock()
session = sqlite_session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
@ -902,7 +926,7 @@ def test_dataset_scoped_read_permission_denied(app: Flask, api_cls):
method(api, session, "tenant-1", make_account(), "dataset-1")
class TestDatasetQueryApi:
class TestDatasetQueryApi(_UsesSQLiteSession):
def _query_record(self, index: int = 1) -> DatasetQuery:
query = DatasetQuery(
dataset_id="dataset-id",
@ -929,7 +953,7 @@ class TestDatasetQueryApi:
patch.object(DatasetService, "check_dataset_permission", return_value=None),
patch.object(DatasetService, "get_dataset_queries", return_value=(queries, 2)),
):
response, status = method(api, MagicMock(), current_user, dataset_id)
response, status = method(api, self.session, current_user, dataset_id)
assert status == 200
assert response["total"] == 2
assert response["page"] == 1
@ -952,24 +976,30 @@ class TestDatasetQueryApi:
dataset = make_dataset(id="dataset-id")
query = self._query_record()
query.content = json.dumps([{"content_type": "image_query", "content": "file-1"}])
upload_file = SimpleNamespace(
id="file-1",
upload_file = UploadFile(
tenant_id="tenant-1",
storage_type=StorageType.LOCAL,
key="image.png",
name="image.png",
size=10,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-1",
created_at=datetime.datetime(2024, 1, 1, tzinfo=datetime.UTC),
used=False,
)
session = MagicMock()
session.scalar.return_value = upload_file
upload_file.id = "file-1"
session = self.session
session.add(upload_file)
session.flush()
with (
app.test_request_context("/datasets/queries"),
patch.object(DatasetService, "get_dataset", return_value=dataset),
patch.object(DatasetService, "check_dataset_permission", return_value=None),
patch.object(DatasetService, "get_dataset_queries", return_value=([query], 1)),
patch("models.dataset.db") as db_mock,
patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"),
):
db_mock.session.scalar.return_value = upload_file
response, status = method(api, session, make_account(), "dataset-id")
assert status == 200
@ -987,8 +1017,7 @@ class TestDatasetQueryApi:
},
}
]
session.scalar.assert_called_once()
db_mock.session.scalar.assert_not_called()
assert session.get(UploadFile, "file-1") is upload_file
def test_get_queries_dataset_not_found(self, app: Flask):
api = DatasetQueryApi()
@ -1000,7 +1029,7 @@ class TestDatasetQueryApi:
patch.object(DatasetService, "get_dataset", return_value=None),
):
with pytest.raises(NotFound, match="Dataset not found"):
method(api, MagicMock(), current_user, dataset_id)
method(api, self.session, current_user, dataset_id)
def test_get_queries_permission_denied(self, app: Flask):
api = DatasetQueryApi()
@ -1018,7 +1047,7 @@ class TestDatasetQueryApi:
),
):
with pytest.raises(Forbidden):
method(api, MagicMock(), current_user, dataset_id)
method(api, self.session, current_user, dataset_id)
def test_get_queries_pagination_has_more(self, app: Flask):
api = DatasetQueryApi()
@ -1033,13 +1062,13 @@ class TestDatasetQueryApi:
patch.object(DatasetService, "check_dataset_permission", return_value=None),
patch.object(DatasetService, "get_dataset_queries", return_value=(queries, 40)),
):
response, status = method(api, MagicMock(), current_user, dataset_id)
response, status = method(api, self.session, current_user, dataset_id)
assert status == 200
assert response["has_more"] is True
assert len(response["data"]) == 20
class TestDatasetIndexingEstimateApi:
class TestDatasetIndexingEstimateApi(_UsesSQLiteSession):
def _upload_file(self, *, tenant_id: str = "tenant-1", file_id: str = "file-1") -> UploadFile:
upload_file = UploadFile(
tenant_id=tenant_id,
@ -1072,8 +1101,9 @@ class TestDatasetIndexingEstimateApi:
method = unwrap(api.post)
payload = self._base_payload()
mock_file = self._upload_file()
session = MagicMock()
session.scalars.return_value.all.return_value = [mock_file]
session = self.session
session.add(mock_file)
session.flush()
mock_response = IndexingEstimate(total_segments=100, preview=[])
@ -1102,8 +1132,7 @@ class TestDatasetIndexingEstimateApi:
api = DatasetIndexingEstimateApi()
method = unwrap(api.post)
payload = self._base_payload()
session = MagicMock()
session.scalars.return_value.all.return_value = None
session = self.session
with (
app.test_request_context("/"),
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
@ -1122,8 +1151,9 @@ class TestDatasetIndexingEstimateApi:
method = unwrap(api.post)
mock_file = self._upload_file()
payload = self._base_payload()
session = MagicMock()
session.scalars.return_value.all.return_value = [mock_file]
session = self.session
session.add(mock_file)
session.flush()
with (
app.test_request_context("/"),
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
@ -1146,8 +1176,9 @@ class TestDatasetIndexingEstimateApi:
method = unwrap(api.post)
mock_file = self._upload_file()
payload = self._base_payload()
session = MagicMock()
session.scalars.return_value.all.return_value = [mock_file]
session = self.session
session.add(mock_file)
session.flush()
with (
app.test_request_context("/"),
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
@ -1170,8 +1201,9 @@ class TestDatasetIndexingEstimateApi:
method = unwrap(api.post)
mock_file = self._upload_file()
payload = self._base_payload()
session = MagicMock()
session.scalars.return_value.all.return_value = [mock_file]
session = self.session
session.add(mock_file)
session.flush()
with (
app.test_request_context("/"),
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
@ -1189,16 +1221,16 @@ class TestDatasetIndexingEstimateApi:
)
class TestDatasetRelatedAppListApi:
class TestDatasetRelatedAppListApi(_UsesSQLiteSession):
def test_get_success(self, app: Flask):
api = DatasetRelatedAppListApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
app1 = make_related_app(id="app-1", name="App 1")
app2 = make_related_app(id="app-2", name="App 2")
join1 = MagicMock(app_id="app-1")
join2 = MagicMock(app_id="app-2")
session = MagicMock()
join1 = AppDatasetJoin(app_id="app-1", dataset_id="dataset-1")
join2 = AppDatasetJoin(app_id="app-2", dataset_id="dataset-1")
session = self.session
with (
app.test_request_context("/"),
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset),
@ -1251,7 +1283,7 @@ class TestDatasetRelatedAppListApi:
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None),
):
with pytest.raises(NotFound):
method(api, MagicMock(), make_account(), "dataset-1")
method(api, self.session, make_account(), "dataset-1")
def test_get_permission_denied(self, app: Flask):
api = DatasetRelatedAppListApi()
@ -1266,16 +1298,16 @@ class TestDatasetRelatedAppListApi:
),
):
with pytest.raises(Forbidden):
method(api, MagicMock(), make_account(), "dataset-1")
method(api, self.session, make_account(), "dataset-1")
def test_get_filters_none_apps(self, app: Flask):
api = DatasetRelatedAppListApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
app1 = make_related_app()
join1 = MagicMock(app_id="app-1")
join2 = MagicMock(app_id="app-2")
session = MagicMock()
join1 = AppDatasetJoin(app_id="app-1", dataset_id="dataset-1")
join2 = AppDatasetJoin(app_id="app-2", dataset_id="dataset-1")
session = self.session
with (
app.test_request_context("/"),
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset),
@ -1303,26 +1335,17 @@ class TestDatasetRelatedAppListApi:
]
class TestDatasetIndexingStatusApi:
class TestDatasetIndexingStatusApi(_UsesSQLiteSession):
def test_get_success_with_documents(self, app: Flask):
api = DatasetIndexingStatusApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
current_user = make_account()
document = MagicMock()
document.id = "doc-1"
document.indexing_status = "completed"
document.processing_started_at = None
document.parsing_completed_at = None
document.cleaning_completed_at = None
document.splitting_completed_at = None
document.completed_at = None
document.paused_at = None
document.error = None
document.stopped_at = None
session = MagicMock()
session.scalars.return_value.all.return_value = [document]
session.scalar.return_value = 3
document = make_document_status()
session = self.session
session.add(document)
session.add_all([make_document_segment(position=position, completed=True) for position in range(1, 4)])
session.flush()
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
@ -1337,16 +1360,13 @@ class TestDatasetIndexingStatusApi:
assert item["total_segments"] == 3
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
check_permission.assert_called_once_with(dataset, current_user, session)
assert {"dataset-1", "tenant-1"} <= set(session.scalars.call_args.args[0].compile().params.values())
for segment_count_call in session.scalar.call_args_list:
assert {"dataset-1", "tenant-1", "doc-1"} <= set(segment_count_call.args[0].compile().params.values())
assert session.get(Document, "doc-1") is document
def test_get_success_no_documents(self, app: Flask):
api = DatasetIndexingStatusApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
session = MagicMock()
session.scalars.return_value.all.return_value = []
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
@ -1360,20 +1380,11 @@ class TestDatasetIndexingStatusApi:
api = DatasetIndexingStatusApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
document = MagicMock()
document.id = "doc-1"
document.indexing_status = "indexing"
document.processing_started_at = None
document.parsing_completed_at = None
document.cleaning_completed_at = None
document.splitting_completed_at = None
document.completed_at = None
document.paused_at = None
document.error = None
document.stopped_at = None
session = MagicMock()
session.scalars.return_value.all.return_value = [document]
session.scalar.side_effect = [2, 5]
document = make_document_status(indexing_status=IndexingStatus.INDEXING)
session = self.session
session.add(document)
session.add_all([make_document_segment(position=position, completed=position <= 2) for position in range(1, 6)])
session.flush()
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
@ -1386,7 +1397,7 @@ class TestDatasetIndexingStatusApi:
assert item["total_segments"] == 5
class TestDatasetApiKeyApi:
class TestDatasetApiKeyApi(_UsesSQLiteSession):
def test_get_api_keys_success(self, app: Flask):
api = DatasetApiKeyApi()
method = unwrap(api.get)
@ -1404,8 +1415,11 @@ class TestDatasetApiKeyApi:
last_used_at=None,
created_at=None,
)
session = MagicMock()
session.scalars.return_value.all.return_value = [mock_key_1, mock_key_2]
session = self.session
mock_key_1.tenant_id = "tenant-1"
mock_key_2.tenant_id = "tenant-1"
session.add_all([mock_key_1, mock_key_2])
session.flush()
with app.test_request_context("/"):
response = method(api, session, "tenant-1")
assert "data" in response
@ -1418,30 +1432,31 @@ class TestDatasetApiKeyApi:
def test_post_create_api_key_success(self, app: Flask):
api = DatasetApiKeyApi()
method = unwrap(api.post)
mock_token = MagicMock()
mock_token.id = "new-key-id"
mock_token.last_used_at = None
mock_token.created_at = datetime.datetime(2024, 1, 1, 0, 0, 0, tzinfo=datetime.UTC)
mock_api_token_cls = MagicMock()
mock_api_token_cls.return_value = mock_token
mock_api_token_cls.generate_api_key.return_value = "dataset-abc123"
session = MagicMock()
session.scalar.return_value = 3
with app.test_request_context("/"), patch("controllers.console.datasets.datasets.ApiToken", mock_api_token_cls):
session = self.session
with (
app.test_request_context("/"),
patch.object(ApiToken, "generate_api_key", return_value="dataset-abc123") as generate_api_key,
):
response, status = method(api, session, "tenant-1")
assert status == 200
assert isinstance(response, dict)
assert response["id"] == "new-key-id"
assert response["token"] == "dataset-abc123"
assert response["type"] == "dataset"
assert response["created_at"] is not None
mock_api_token_cls.generate_api_key.assert_called_once_with("dataset-", 24, session=session)
generate_api_key.assert_called_once_with("dataset-", 24, session=session)
assert session.get(ApiToken, response["id"]).token == "dataset-abc123"
def test_post_exceed_max_keys(self, app: Flask):
api = DatasetApiKeyApi()
method = unwrap(api.post)
session = MagicMock()
session.scalar.return_value = 10
session = self.session
session.add_all(
[
ApiToken(id=f"key-{index}", tenant_id="tenant-1", type="dataset", token=f"ds-{index}")
for index in range(10)
]
)
session.flush()
with app.test_request_context("/"):
with pytest.raises(BadRequest) as exc_info:
method(api, session, "tenant-1")
@ -1452,36 +1467,42 @@ class TestDatasetApiKeyApi:
}
class TestDatasetApiDeleteApi:
class TestDatasetApiDeleteApi(_UsesSQLiteSession):
def test_delete_success(self, app: Flask):
api = DatasetApiDeleteApi()
method = unwrap(api.delete)
mock_key = MagicMock()
session = MagicMock()
session.scalar.return_value = mock_key
with app.test_request_context("/"):
session = self.session
key = ApiToken(id="api-key-id", tenant_id="tenant-1", type="dataset", token="dataset-secret")
session.add(key)
session.flush()
with (
app.test_request_context("/"),
patch("controllers.console.datasets.datasets.ApiTokenCache.delete") as delete_cache,
):
response, status = method(api, session, "tenant-1", "api-key-id")
assert status == 204
assert response == ""
delete_cache.assert_called_once()
session.flush()
assert session.get(ApiToken, "api-key-id") is None
def test_delete_key_not_found(self, app: Flask):
api = DatasetApiDeleteApi()
method = unwrap(api.delete)
session = MagicMock()
session.scalar.return_value = None
session = self.session
with app.test_request_context("/"):
with pytest.raises(NotFound):
method(api, session, "tenant-1", "api-key-id")
class TestDatasetEnableApiApi:
class TestDatasetEnableApiApi(_UsesSQLiteSession):
@pytest.mark.parametrize(("status_value", "enabled"), [("enable", True), ("disable", False)])
def test_update_api_status(self, app: Flask, status_value: str, enabled: bool):
api = DatasetEnableApiApi()
method = unwrap(api.post)
dataset = make_dataset(id="dataset-1")
current_user = make_account()
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
@ -1499,7 +1520,7 @@ class TestDatasetEnableApiApi:
api = DatasetEnableApiApi()
method = unwrap(api.post)
dataset = make_dataset(id="dataset-1")
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
@ -1582,7 +1603,7 @@ class TestDatasetRetrievalSettingApi:
]
class TestDatasetRetrievalSettingMockApi:
class TestDatasetRetrievalSettingMockApi(_UsesSQLiteSession):
def test_get_success(self, app: Flask):
api = DatasetRetrievalSettingMockApi()
method = unwrap(api.get)
@ -1597,14 +1618,14 @@ class TestDatasetRetrievalSettingMockApi:
assert response["retrieval_method"] == ["semantic"]
class TestDatasetErrorDocs:
class TestDatasetErrorDocs(_UsesSQLiteSession):
def test_get_success(self, app: Flask):
api = DatasetErrorDocs()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
error_doc = make_document_status(id="error-doc", indexing_status=IndexingStatus.ERROR, error="failed")
current_user = make_account()
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
@ -1624,7 +1645,7 @@ class TestDatasetErrorDocs:
def test_get_dataset_not_found(self, app: Flask):
api = DatasetErrorDocs()
method = unwrap(api.get)
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset,
@ -1634,7 +1655,7 @@ class TestDatasetErrorDocs:
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
class TestDatasetPermissionUserListApi:
class TestDatasetPermissionUserListApi(_UsesSQLiteSession):
def test_get_success(self, app: Flask):
api = DatasetPermissionUserListApi()
method = unwrap(api.get)
@ -1649,7 +1670,7 @@ class TestDatasetPermissionUserListApi:
return_value=users,
),
):
response, status = method(api, MagicMock(), make_account(), "dataset-1")
response, status = method(api, self.session, make_account(), "dataset-1")
assert status == 200
assert response["data"] == users
@ -1666,17 +1687,17 @@ class TestDatasetPermissionUserListApi:
),
):
with pytest.raises(Forbidden):
method(api, MagicMock(), make_account(), "dataset-1")
method(api, self.session, make_account(), "dataset-1")
class TestDatasetAutoDisableLogApi:
class TestDatasetAutoDisableLogApi(_UsesSQLiteSession):
def test_get_success(self, app: Flask):
api = DatasetAutoDisableLogApi()
method = unwrap(api.get)
dataset = make_dataset(id="dataset-1")
logs = {"document_ids": ["doc-1"], "count": 1}
current_user = make_account()
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
@ -1693,7 +1714,7 @@ class TestDatasetAutoDisableLogApi:
def test_get_dataset_not_found(self, app: Flask):
api = DatasetAutoDisableLogApi()
method = unwrap(api.get)
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset,

View File

@ -55,6 +55,7 @@ from services.vector_space_admission_service import (
VECTOR_SPACE_ADMISSION_ERROR_CODE,
format_vector_space_admission_error,
)
from tests.unit_tests.config_override import config_overrides_context
def make_serializable_document(**overrides):
@ -504,7 +505,7 @@ class TestDatasetInitApi:
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", payload),
patch("controllers.console.datasets.datasets_document.dify_config.RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch(
"controllers.console.datasets.datasets_document.DocumentService.document_create_args_validate",
return_value=None,
@ -560,7 +561,7 @@ class TestDocumentResource:
api = DocumentResource()
session = MagicMock()
with (
patch("controllers.console.datasets.datasets_document.dify_config.RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch(
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
return_value=dataset,

View File

@ -22,6 +22,7 @@ from services.entities.knowledge_entities.knowledge_entities import MetadataArgs
from services.errors.account import NoPermissionError
from services.errors.metadata import MetadataResourceNotFoundError
from services.metadata_service import MetadataService
from tests.unit_tests.config_override import config_overrides_context
@pytest.fixture
@ -150,7 +151,7 @@ class TestDatasetMetadataGetApi:
method = unwrap(api.get)
with (
app.test_request_context("/"),
patch("controllers.console.datasets.metadata.dify_config.RBAC_ENABLED", True),
config_overrides_context(RBAC_ENABLED=True),
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
patch.object(DatasetService, "check_dataset_permission") as check_permission,
patch.object(

View File

@ -37,6 +37,32 @@ def _snippet(**overrides) -> CustomizedSnippet:
return CustomizedSnippet(**data)
def _workflow(**overrides) -> SimpleNamespace:
data = {
"id": "workflow-1",
"graph_dict": {"nodes": [], "edges": []},
"features_dict": {},
"unique_hash": "hash-1",
"version": "2024-01-01 00:00:00",
"marked_name": "v1",
"marked_comment": "first version",
"created_by_account": None,
"created_at": datetime(2024, 1, 1),
"updated_by_account": None,
"updated_at": datetime(2024, 1, 1),
"tool_published": False,
"environment_variables": [],
"conversation_variables": [],
"rag_pipeline_variables": [],
}
data.update(overrides)
workflow = SimpleNamespace(**data)
workflow.get_created_by_account = Mock(return_value=workflow.created_by_account)
workflow.get_updated_by_account = Mock(return_value=workflow.updated_by_account)
workflow.get_tool_published = Mock(return_value=workflow.tool_published)
return workflow
@pytest.fixture(autouse=True)
def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
snippet_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
@ -114,6 +140,34 @@ def test_draft_workflow_get_raises_when_missing(app: Flask, monkeypatch: pytest.
handler(api, snippet=snippet)
def test_draft_workflow_get_uses_session_aware_response_source(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
workflow = _workflow()
snippet = _snippet()
session = Mock(spec=Session)
monkeypatch.setattr(snippet_workflow_module, "db", SimpleNamespace(session=Mock(return_value=session)))
monkeypatch.setattr(
snippet_workflow_module,
"_snippet_service",
lambda: SimpleNamespace(get_draft_workflow=Mock(return_value=workflow)),
)
monkeypatch.setattr(
snippet_workflow_module.WorkflowAgentPublishService,
"project_draft_bindings_to_graph",
Mock(return_value=workflow.graph_dict),
)
api = snippet_workflow_module.SnippetDraftWorkflowApi()
handler = unwrap(api.get)
with app.test_request_context("/snippets/snippet-1/workflows/draft"):
response = handler(api, snippet=snippet)
assert response["id"] == "workflow-1"
workflow.get_created_by_account.assert_called_once_with(session=session)
workflow.get_updated_by_account.assert_called_once_with(session=session)
workflow.get_tool_published.assert_called_once_with(session=session)
def test_draft_workflow_post_returns_400_for_invalid_graph(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
user = _account("account-1")
snippet = _snippet()
@ -161,6 +215,29 @@ def test_published_workflow_get_returns_none_when_not_published(app) -> None:
assert handler(api, snippet=SimpleNamespace(id="snippet-1", is_published=False)) is None
def test_published_workflow_get_uses_session_aware_response_source(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
workflow = _workflow()
session = Mock(spec=Session)
snippet = SimpleNamespace(id="snippet-1", is_published=True, input_fields_list=[])
monkeypatch.setattr(snippet_workflow_module, "db", SimpleNamespace(session=Mock(return_value=session)))
monkeypatch.setattr(
snippet_workflow_module,
"_snippet_service",
lambda: SimpleNamespace(get_published_workflow=Mock(return_value=workflow)),
)
api = snippet_workflow_module.SnippetPublishedWorkflowApi()
handler = unwrap(api.get)
with app.test_request_context("/snippets/snippet-1/workflows/publish"):
response = handler(api, snippet=snippet)
assert response["id"] == "workflow-1"
workflow.get_created_by_account.assert_called_once_with(session=session)
workflow.get_updated_by_account.assert_called_once_with(session=session)
workflow.get_tool_published.assert_called_once_with(session=session)
@pytest.mark.parametrize("sqlite_session", [(CustomizedSnippet,)], indirect=True)
def test_published_workflow_post_returns_400_when_publish_fails(
app: Flask,
@ -247,23 +324,7 @@ def test_list_published_snippet_workflows_includes_input_fields(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> None:
workflow = SimpleNamespace(
id="workflow-1",
graph_dict={"nodes": [], "edges": []},
features_dict={},
unique_hash="hash-1",
version="2024-01-01 00:00:00",
marked_name="",
marked_comment="",
created_by_account=None,
created_at=datetime(2024, 1, 1),
updated_by_account=None,
updated_at=datetime(2024, 1, 1),
tool_published=False,
environment_variables=[],
conversation_variables=[],
rag_pipeline_variables=[],
)
workflow = _workflow(marked_name="", marked_comment="")
input_fields = [{"variable": "query", "type": "text"}]
snippet = _snippet(input_fields=json.dumps(input_fields))
@ -406,23 +467,7 @@ def test_update_published_snippet_workflow_returns_updated_workflow(
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
) -> None:
workflow = SimpleNamespace(
id="workflow-1",
graph_dict={"nodes": [], "edges": []},
features_dict={},
unique_hash="hash-1",
version="2024-01-01 00:00:00",
marked_name="v1",
marked_comment="first version",
created_by_account=None,
created_at=datetime(2024, 1, 1),
updated_by_account=None,
updated_at=datetime(2024, 1, 1),
tool_published=False,
environment_variables=[],
conversation_variables=[],
rag_pipeline_variables=[],
)
workflow = _workflow()
user = _account("account-1")
input_fields = [{"variable": "query", "type": "text"}]
snippet = _snippet(input_fields=json.dumps(input_fields))

Some files were not shown because too many files have changed in this diff Show More