mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 02:43:49 +08:00
Merge branch 'feat/creator-profile-home' into deploy/dev
This commit is contained in:
commit
ff260f9be3
70
.github/workflows/marketplace-performance-e2e.yml
vendored
Normal file
70
.github/workflows/marketplace-performance-e2e.yml
vendored
Normal 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
|
||||
2
.github/workflows/translate-i18n-claude.yml
vendored
2
.github/workflows/translate-i18n-claude.yml
vendored
@ -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 }}
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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."
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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},
|
||||
)
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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](
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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),
|
||||
)
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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."
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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],
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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."
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
*,
|
||||
|
||||
@ -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},
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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.")
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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 ""
|
||||
|
||||
@ -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,
|
||||
),
|
||||
|
||||
@ -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.
|
||||
|
||||
|
||||
@ -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 |
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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:
|
||||
|
||||
189
api/repositories/step_by_step_tour_repository.py
Normal file
189
api/repositories/step_by_step_tour_repository.py
Normal 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
|
||||
230
api/services/account_email_registration_adapters.py
Normal file
230
api/services/account_email_registration_adapters.py
Normal 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,
|
||||
)
|
||||
207
api/services/account_email_registration_service.py
Normal file
207
api/services/account_email_registration_service.py
Normal 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
|
||||
@ -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."""
|
||||
|
||||
|
||||
@ -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: ...
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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)}
|
||||
|
||||
@ -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:
|
||||
"""
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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"
|
||||
|
||||
38
api/services/entities/notification_entities.py
Normal file
38
api/services/entities/notification_entities.py
Normal 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, ...]
|
||||
42
api/services/entities/onboarding_entities.py
Normal file
42
api/services/entities/onboarding_entities.py
Normal 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
|
||||
@ -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."""
|
||||
|
||||
|
||||
48
api/services/notification_gateway.py
Normal file
48
api/services/notification_gateway.py
Normal 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 "",
|
||||
)
|
||||
60
api/services/notification_service.py
Normal file
60
api/services/notification_service.py
Normal 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,
|
||||
)
|
||||
@ -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,
|
||||
)
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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"):
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
26
api/tests/unit_tests/config_override.py
Normal file
26
api/tests/unit_tests/config_override.py
Normal 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
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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']
|
||||
|
||||
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -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"))
|
||||
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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"),
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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),
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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": "",
|
||||
}
|
||||
)
|
||||
@ -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",
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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
Loading…
Reference in New Issue
Block a user