Compare commits

...
Author SHA1 Message Date
yyh 5c308c4be1 feat(web): persist detail sidebar state in cookies 2026-07-25 18:35:11 +08:00
EvanandGitHub c56ac3983e refactor(web): migrate vector space billing query to contract client (#39558) 2026-07-25 04:26:14 +00:00
Asuka MinatoandGitHub fa000ae8a7 test: use sqlite3 session in test_human_input_service (#38697) 2026-07-25 03:17:57 +00:00
Asuka MinatoandGitHub 8906a49e56 test: use sqlite3 session in test_datasource_provider_service (#38695) 2026-07-25 03:16:43 +00:00
JingyiandGitHub a2d9aeff37 fix(web): prevent customization horizontal overflow (#39557) 2026-07-25 02:24:07 +00:00
yyhandGitHub aff5c47541 refactor: remove backend-only system feature fields (#39543) 2026-07-25 02:01:12 +00:00
yyhandGitHub 58f83fa7e7 refactor(web): remove obsolete route prefix handler (#39545) 2026-07-24 15:30:52 +00:00
Stephen ZhouGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2d5161a647 feat(dataset): add document detail (#39361)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 15:27:35 +00:00
yyhandGitHub 1227f19c6a fix(web): patch React Server Components DoS (#39539) 2026-07-24 14:00:26 +00:00
EvanandGitHub a0b917425a test: use caplog in ext login tests (#39542) 2026-07-24 13:27:05 +00:00
Stephen ZhouGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
799c7eea3a feat(dataset): add document processing tasks (#39326)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 12:58:15 +00:00
wangxiaoleiandGitHub 5277e4b375 fix: fix when tool execute in progress tool_meta miss tool_provider_type (#39513) 2026-07-24 11:35:12 +00:00
a877e1bd7e fix(agent): block unpublished agents in workflows (#39532)
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-07-24 11:02:43 +00:00
yyhandGitHub a5703762c0 test(e2e): wait for app list before menu interaction (#39538) 2026-07-24 10:37:22 +00:00
yyhandGitHub dcd40d603c fix(agent-v2): link API reference to agent guide (#39534) 2026-07-24 10:14:27 +00:00
JoelandGitHub 8ef002f624 fix: include CSRF token when fetching skill file content (#39527) 2026-07-24 09:35:41 +00:00
yyhandGitHub b1dec5bd42 chore: add web shell code owners (#39526) 2026-07-24 09:17:15 +00:00
yyhandGitHub 4e85acc605 fix(web): hydrate workspace permissions before navigation render (#39524) 2026-07-24 09:06:57 +00:00
Wu TianweiandGitHub a2c67372ed fix(web): allow clipboard writes in embedded apps (#39511) 2026-07-24 08:27:57 +00:00
林玮 (Jade Lin)andGitHub 7bb09ce039 chore(web): improve web accessibility semantics (#39505) 2026-07-24 08:16:04 +00:00
JoelandGitHub e1495c0bae fix: prevent agent preview actions from being clipped (#39510) 2026-07-24 08:02:55 +00:00
JyongandGitHub ededa1f4cd chore: add deploy-knowledge.yml (#39509) 2026-07-24 08:00:11 +00:00
JoelandGitHub b2861b60d8 fix: audio config ui problem (#39504) 2026-07-24 07:42:26 +00:00
3cbe49bf45 fix: don't open the node settings panel when clicking a node in comment mode (#39500)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-24 07:38:01 +00:00
JoelandGitHub dd19b4ff79 fix: download skill files through authenticated API requests (#39499) 2026-07-24 07:37:27 +00:00
yyhGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>GareArc
66a545fc6d refactor: centralize deployment edition in system features (#39454)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: GareArc <garethcxy@dify.ai>
2026-07-24 07:15:02 +00:00
Stephen ZhouGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
eac9f69ad1 feat(dataset): add crawl source selection (#39325)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 07:11:48 +00:00
Yufeng HeandGitHub e4dc29f98e fix(api): make the 10-minute email IP first-strike window take effect (#39479)
Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>
2026-07-24 06:52:30 +00:00
JoelandGitHub 6889974c05 fix: smooth the preview-to-build transition (#39495) 2026-07-24 06:30:07 +00:00
非法操作andGitHub aef1c67721 fix: return disabled auto-upgrade settings when strategy is missing (#39494) 2026-07-24 06:15:02 +00:00
David ParkGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
3ab9e083e0 fix: don't close the active comment panel when a different comment is resolved (#39491)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 06:11:25 +00:00
yyhandGitHub 6e2c21f568 fix(web): simplify agent configure mode state (#39492) 2026-07-24 05:30:45 +00:00
JoelandGitHub abbad7b313 fix: prevent build reset from clearing preview chat (#39490) 2026-07-24 04:59:57 +00:00
yyhandGitHub c4bfea6096 fix(web): adapt workflow agent logs (#39488) 2026-07-24 04:55:30 +00:00
Asuka MinatoandGitHub bbb133cd58 test: use sqlite3 session in test_base_trace_instance (#38744) 2026-07-24 03:55:32 +00:00
JoelGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
34e2205efa chore: support preview mode if not in community version (#39399)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 03:23:01 +00:00
wangxiaoleiandGitHub db29caff2b feat: knowledge add more trace (#38959) 2026-07-24 02:27:16 +00:00
14c50a9f09 feat(web): locate workflow nodes by node_id in the workflow editor (#38187)
Co-authored-by: Crazywoola <100913391+crazywoola@users.noreply.github.com>
2026-07-24 02:16:41 +00:00
Stephen ZhouGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
ba5b26be03 feat(dataset): add website crawl preview (#39324)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 02:10:39 +00:00
samzongGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
922ed4d67d feat(oauth): expose stable account id on provider account endpoint (#39470)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-24 02:09:48 +00:00
zyssyz123andGitHub a8a83a357c fix(agent): show workflow node runs in logs (#39471) 2026-07-24 01:59:57 +00:00
Stephen ZhouGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
302b50c4fb feat(dataset): add source connections (#39319)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-23 13:54:12 +00:00
Byron.wangandGitHub 43254c1ded chore: avoid duplicate token counting during dataset indexing (#39466) 2026-07-23 10:56:03 +00:00
JoelandGitHub 2beafdc457 fix: chat not pass doc if not support vision (#39461) 2026-07-23 10:01:43 +00:00
zyssyz123andGitHub 0198e32447 fix(agent): start new Preview conversations without an ID (#39463) 2026-07-23 09:53:08 +00:00
samzongandGitHub b8f91d0e61 fix(web): echo OAuth state on authorize redirect (#39459)
Signed-off-by: samzong <samzong.lu@gmail.com>
2026-07-23 09:30:52 +00:00
非法操作andGitHub 452dff5c37 fix(api): prevent identity logging deadlock (#39449) 2026-07-23 09:08:04 +00:00
642 changed files with 41698 additions and 5875 deletions
+17
View File
@@ -250,5 +250,22 @@
# Frontend - Workspace
/web/app/components/header/account-dropdown/workplace-selector/ @iamjoel @zxhlyh
# Frontend - App Shell and Console Bootstrap
/web/app/layout.tsx @iamjoel @lyzno1
/web/app/error.tsx @iamjoel @lyzno1
/web/app/(commonLayout)/layout.tsx @iamjoel @lyzno1
/web/app/(commonLayout)/providers.tsx @iamjoel @lyzno1
/web/app/(commonLayout)/hydration-boundary.tsx @iamjoel @lyzno1
/web/app/(commonLayout)/profile-bootstrap-gate.tsx @iamjoel @lyzno1
/web/app/(commonLayout)/error.tsx @iamjoel @lyzno1
/web/app/account/(commonLayout)/layout.tsx @iamjoel @lyzno1
/web/app/components/main-nav/* @iamjoel @lyzno1
/web/app/components/main-nav/components/* @iamjoel @lyzno1
/web/context/query-client.tsx @iamjoel @lyzno1
/web/context/query-client-server.ts @iamjoel @lyzno1
/web/proxy.ts @iamjoel @lyzno1
/web/app/auth/refresh/route.ts @iamjoel @lyzno1
/web/service/server.ts @iamjoel @lyzno1
# Docker
/docker/* @laipz8200
+28
View File
@@ -0,0 +1,28 @@
name: Deploy Knowledge
permissions:
contents: read
on:
workflow_run:
workflows: ["Build and Push API & Web"]
branches:
- "deploy/konwledge"
types:
- completed
jobs:
deploy:
runs-on: depot-ubuntu-24.04
if: |
github.event.workflow_run.conclusion == 'success' &&
github.event.workflow_run.head_branch == 'deploy/konwledge'
steps:
- name: Deploy to server
uses: appleboy/ssh-action@0ff4204d59e8e51228ff73bce53f80d53301dee2 # v1.2.5
with:
host: ${{ secrets.SSH_NEW_RAG_HOST }}
username: ${{ secrets.SSH_USER }}
key: ${{ secrets.SSH_PRIVATE_KEY }}
script: |
${{ vars.SSH_SCRIPT || secrets.SSH_SCRIPT }}
+11
View File
@@ -213,6 +213,17 @@ Before opening a PR / submitting:
DTOs with `register_response_schema_models(...)`, serialize response DTOs with `dump_response(...)`,
and avoid adding new legacy `ns.model(...)`, `@marshal_with(...)`, or GET `@ns.expect(...)` patterns.
### System Features Contract
- Treat the shared Console/Web `/system-features` response as a minimal unauthenticated bootstrap allowlist, not a
general configuration or feature-discovery endpoint. Existing fields do not establish precedent.
- Before adding a field, read `controllers/API_SCHEMA_GUIDE.md#public-system-features-contract` and provide evidence
that both Console and Web have production consumers that require it before authentication.
- Never place backend-only policy, surface-specific configuration, post-authentication state, speculative values, or
large/slow payloads in `SystemFeatureModel`. Use the consumer or domain owner described in the schema guide.
- Agents and reviewers must reject additions whose owner, public exposure, pre-authentication need, or root SSR cost
is not explicit.
### Miscellaneous
- Use `configs.dify_config` for configuration—never read environment variables directly.
+2 -1
View File
@@ -6,6 +6,7 @@ from sqlalchemy import delete, select, update
from sqlalchemy.orm import sessionmaker
from configs import dify_config
from enums.deployment_edition import DeploymentEdition
from events.app_event import app_was_created
from extensions.ext_database import db
from extensions.ext_redis import redis_client
@@ -42,7 +43,7 @@ def reset_encrypt_key_pair():
After the reset, all LLM credentials will become invalid, requiring re-entry.
Only support SELF_HOSTED mode.
"""
if dify_config.EDITION != "SELF_HOSTED":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
click.echo(click.style("This command is only for SELF_HOSTED installations.", fg="red"))
return
with sessionmaker(db.engine, expire_on_commit=False).begin() as session:
+9
View File
@@ -5,6 +5,7 @@ from typing import Any, override
from pydantic.fields import FieldInfo
from pydantic_settings import BaseSettings, PydanticBaseSettingsSource, SettingsConfigDict, TomlConfigSettingsSource
from enums.deployment_edition import DeploymentEdition
from libs.file_utils import search_file_upwards
from .deploy import DeploymentConfig
@@ -116,3 +117,11 @@ class DifyConfig(
),
),
)
@property
def DEPLOYMENT_EDITION(self) -> DeploymentEdition:
if self.EDITION == "CLOUD":
return DeploymentEdition.CLOUD
if self.ENTERPRISE_ENABLED:
return DeploymentEdition.ENTERPRISE
return DeploymentEdition.COMMUNITY
+5
View File
@@ -1497,6 +1497,11 @@ class LoginConfig(BaseSettings):
class AccountConfig(BaseSettings):
ENABLE_CHANGE_EMAIL: bool = Field(
description="whether users can change their email address",
default=True,
)
ACCOUNT_DELETION_TOKEN_EXPIRY_MINUTES: PositiveInt = Field(
description="Duration in minutes for which a account deletion token remains valid",
default=5,
+41
View File
@@ -12,6 +12,47 @@ parameters, response schemas, and Swagger documentation.
- Do not add new Flask-RESTX `fields.*` dictionaries, `Namespace.model(...)` exports, or `@marshal_with(...)` for migrated or new endpoints.
- Do not use `@ns.expect(...)` for GET query parameters. Flask-RESTX documents that as a request body.
## Public System Features Contract
The Console and Web `/system-features` endpoints share `SystemFeatureModel`. They are unauthenticated and may be
requested during root SSR, so treat this response as a minimal public bootstrap allowlist. It is not a general
configuration endpoint, a feature registry, or a mirror of environment and Enterprise settings. Existing fields are
legacy inventory and do not establish precedent for new fields.
A new field is eligible only when all of the following are true:
1. Both Console and Web have named production consumers for the field.
2. Both consumers need the value before authentication and tenant/workspace bootstrap to render initial state or
choose an authentication flow.
3. The value varies at runtime or by deployment and cannot be safely derived from an existing public contract.
4. The value is non-sensitive, safe to disclose without authentication, and has stable public API semantics.
5. Sending the value on every root bootstrap is demonstrably clearer and cheaper than a consumer-owned query.
Do not add:
- Backend-only policy or enforcement inputs, including security decisions, upload limits, or integration toggles.
- Console-only or Web-only configuration.
- Tenant, workspace, account, permission, billing-detail, or other post-authentication state.
- Provider payloads, operational diagnostics, large nested objects, or values without active consumers.
- Speculative fields added for possible future use.
Route excluded values to their actual owner:
- Keep backend enforcement behind a narrow service method or domain policy.
- Serve post-authentication state from an authenticated domain endpoint.
- Serve surface-specific bootstrap state from a Console- or Web-specific endpoint and account for its SSR, caching,
and failure cost explicitly.
- Load large, slow, or page-specific data lazily through a consumer-owned query.
Every pull request that adds a System Features field must:
- Name both production consumer paths and explain why they require the value before authentication.
- Document the root SSR request, payload, caching, and failure-mode impact.
- Update the Pydantic owner, regenerate OpenAPI Markdown and TypeScript/Zod contracts, and update shared fixtures.
- Add Console and Web schema regression coverage. Do not hand-edit generated contracts or add compatibility defaults.
Reviewers should reject a field when its owner, pre-authentication need, or consumers are unclear.
## Naming
- Request body models: use a `Payload` suffix.
+16 -2
View File
@@ -539,9 +539,23 @@ def _parse_observability_time_range(start: str | None, end: str | None, account:
def _query_values(name: str, alias_name: str | None = None) -> list[str]:
values = request.args.getlist(name)
def _get_values(field_name: str) -> list[str]:
values = request.args.getlist(field_name)
indexed_values: list[tuple[int, list[str]]] = []
prefix = f"{field_name}["
for key in request.args:
if not key.startswith(prefix) or not key.endswith("]"):
continue
index = key[len(prefix) : -1]
if index.isdigit():
indexed_values.append((int(index), request.args.getlist(key)))
for _, items in sorted(indexed_values):
values.extend(items)
return values
values = _get_values(name)
if alias_name:
values.extend(request.args.getlist(alias_name))
values.extend(_get_values(alias_name))
return [value.strip() for value in values if value.strip()]
+25 -14
View File
@@ -351,24 +351,31 @@ def _resolve_current_user_agent_debug_conversation_id(
app_model: App,
agent_id: str | None,
draft_type: AgentConfigDraftType,
start_new: bool = False,
) -> str:
"""Resolve the current editor's conversation without crossing draft surfaces."""
"""Resolve or rotate the current editor's conversation within one draft surface.
``start_new`` rotates the scoped mapping through ``AgentRosterService`` so
the old runtime session is retired before the new conversation is used.
Continuations and Build chat keep resolving the existing mapping.
"""
roster_service = AgentRosterService(session)
if agent_id:
return roster_service.get_or_create_agent_app_debug_conversation_id(
tenant_id=current_tenant_id,
agent_id=agent_id,
account_id=current_user.id,
draft_type=draft_type,
)
resolved_agent_id = agent_id
if not resolved_agent_id:
agent = roster_service.get_app_backing_agent(tenant_id=current_tenant_id, app_id=str(app_model.id))
if agent is None:
raise AgentNotFoundError()
resolved_agent_id = agent.id
agent = roster_service.get_app_backing_agent(tenant_id=current_tenant_id, app_id=str(app_model.id))
if agent is None:
raise AgentNotFoundError()
return roster_service.get_or_create_agent_app_debug_conversation_id(
resolve_conversation = (
roster_service.refresh_agent_app_debug_conversation_id
if start_new
else roster_service.get_or_create_agent_app_debug_conversation_id
)
return resolve_conversation(
tenant_id=current_tenant_id,
agent_id=agent.id,
agent_id=resolved_agent_id,
account_id=current_user.id,
draft_type=draft_type,
)
@@ -387,13 +394,17 @@ def _create_chat_message(
args = args_model.model_dump(exclude_none=True, by_alias=True)
if AppMode.value_of(app_model.mode) == AppMode.AGENT:
draft_type = AgentConfigDraftType(args_model.draft_type)
# Preview follows the normal chat contract: an omitted/empty conversation ID starts a new
# conversation. Build chat keeps its stable mapping so build drafts and finalization stay continuous.
debug_conversation_id = _resolve_current_user_agent_debug_conversation_id(
session=session,
current_tenant_id=current_tenant_id or app_model.tenant_id,
current_user=current_user,
app_model=app_model,
agent_id=agent_id,
draft_type=AgentConfigDraftType(args_model.draft_type),
draft_type=draft_type,
start_new=draft_type == AgentConfigDraftType.DRAFT and not args_model.conversation_id,
)
if args_model.conversation_id and args_model.conversation_id != debug_conversation_id:
raise NotFound("Conversation Not Exists.")
@@ -198,6 +198,6 @@ class ForgotPasswordResetApi(Resource):
# Create workspace if needed
if (
not TenantService.get_join_tenants(account, session=db.session())
and FeatureService.get_system_features().is_allow_create_workspace
and FeatureService.is_workspace_creation_allowed()
):
TenantService.create_owner_tenant(account, session=db.session())
+6 -5
View File
@@ -162,9 +162,10 @@ class LoginApi(Resource):
# SELF_HOSTED only have one workspace
tenants = TenantService.get_join_tenants(account, session=db.session())
if len(tenants) == 0:
system_features = FeatureService.get_system_features()
if system_features.is_allow_create_workspace and not system_features.license.workspaces.is_available():
if (
FeatureService.is_workspace_creation_allowed()
and not FeatureService.get_license().workspaces.is_available()
):
raise WorkspacesLimitExceeded()
else:
return SimpleResultOptionalDataResponse(
@@ -310,10 +311,10 @@ class EmailCodeLoginApi(Resource):
if account:
tenants = TenantService.get_join_tenants(account, session=db.session())
if not tenants:
workspaces = FeatureService.get_system_features().license.workspaces
workspaces = FeatureService.get_license().workspaces
if not workspaces.is_available():
raise WorkspacesLimitExceeded()
if not FeatureService.get_system_features().is_allow_create_workspace:
if not FeatureService.is_workspace_creation_allowed():
raise NotAllowedCreateWorkspace()
else:
TenantService.create_owner_tenant(account, session=db.session())
+1 -1
View File
@@ -292,7 +292,7 @@ def _generate_account(
if account:
tenants = TenantService.get_join_tenants(account, session=db.session())
if not tenants:
if not FeatureService.get_system_features().is_allow_create_workspace:
if not FeatureService.is_workspace_creation_allowed():
raise WorkSpaceNotAllowedCreateError()
else:
TenantService.create_owner_tenant(account, session=db.session())
@@ -56,6 +56,7 @@ class OAuthProviderTokenResponse(BaseModel):
class OAuthProviderAccountResponse(BaseModel):
id: str
name: str
email: str
avatar: str | None = None
@@ -251,6 +252,7 @@ class OAuthServerUserAccountApi(Resource):
def post(self, oauth_provider_app: OAuthProviderApp, account: Account):
return jsonable_encoder(
{
"id": account.id,
"name": account.name,
"email": account.email,
"avatar": account.avatar,
+32 -12
View File
@@ -3,10 +3,11 @@ from flask_restx import Resource
from controllers.common.schema import register_response_schema_models
from fields.base import ResponseModel
from libs.helper import dump_response
from libs.login import current_account_with_tenant_optional, login_required
from libs.login import login_required
from services.feature_service import (
FeatureModel,
FeatureService,
LicenseModel,
LimitationModel,
SystemFeatureModel,
)
@@ -32,6 +33,7 @@ register_response_schema_models(
console_ns,
AppDslVersionResponse,
FeatureModel,
LicenseModel,
LimitationModel,
SystemFeatureModel,
TrialModelsResponse,
@@ -121,22 +123,40 @@ class AppDslVersionApi(Resource):
@console_ns.route("/system-features")
class SystemFeatureApi(Resource):
@console_ns.doc("get_system_features")
@console_ns.doc(description="Get system-wide feature configuration")
@console_ns.doc(
description="Get the non-sensitive bootstrap snapshot exposed before Console or Web authentication. "
"This is not a general feature registry."
)
@console_ns.response(
200,
"Success",
console_ns.models[SystemFeatureModel.__name__],
)
def get(self):
"""Get system-wide feature configuration
"""Get the non-sensitive bootstrap snapshot exposed before authentication.
NOTE: This endpoint is unauthenticated by design, as it provides system features
data required for dashboard initialization.
Authentication would create circular dependency (can't login without dashboard loading).
Only non-sensitive configuration data should be returned by this endpoint.
Authentication configuration must be available before the authentication flow can be selected.
Authenticated license detail is served separately by SystemFeatureLicenseApi.
"""
current_user, _ = current_account_with_tenant_optional()
is_authenticated = current_user is not None
return FeatureService.get_system_features(is_authenticated=is_authenticated).model_dump()
return dump_response(SystemFeatureModel, FeatureService.get_system_features())
@console_ns.route("/system-features/license")
class SystemFeatureLicenseApi(Resource):
@console_ns.doc("get_system_license")
@console_ns.doc(description="Get license status and usage detail")
@console_ns.response(
200,
"Success",
console_ns.models[LicenseModel.__name__],
)
@setup_required
@login_required
@account_initialization_required
def get(self):
"""Get full license detail (status, expiry, workspace/seat usage).
Authenticated counterpart to the license *status* exposed on the public
system-features endpoint.
"""
return FeatureService.get_license().model_dump()
+2 -1
View File
@@ -8,6 +8,7 @@ from sqlalchemy.orm import Session
from configs import dify_config
from controllers.fastopenapi import console_router
from enums.deployment_edition import DeploymentEdition
from extensions.ext_database import db
from models.model import DifySetup
from services.account_service import TenantService
@@ -63,7 +64,7 @@ def validate_init_password(payload: InitValidatePayload) -> InitValidateResponse
def get_init_validate_status() -> bool:
if dify_config.EDITION == "SELF_HOSTED":
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
if os.environ.get("INIT_PASSWORD"):
if session.get("is_init_validated"):
return True
+3 -2
View File
@@ -6,6 +6,7 @@ from sqlalchemy import select
from configs import dify_config
from controllers.fastopenapi import console_router
from enums.deployment_edition import DeploymentEdition
from libs.helper import EmailStr, extract_remote_ip
from libs.password import valid_password
from models.model import DifySetup, db
@@ -52,7 +53,7 @@ def get_setup_status_api() -> SetupStatusResponse:
Only bootstrap-safe status information should be returned by this endpoint.
"""
if dify_config.EDITION == "SELF_HOSTED":
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
setup_status = get_setup_status()
if setup_status and not isinstance(setup_status, bool):
return SetupStatusResponse(step="finished", setup_at=setup_status.setup_at.isoformat())
@@ -102,7 +103,7 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse:
def get_setup_status() -> DifySetup | bool | None:
if dify_config.EDITION == "SELF_HOSTED":
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
return db.session.scalar(select(DifySetup).limit(1))
return True
+2 -1
View File
@@ -46,6 +46,7 @@ from controllers.console.wraps import (
with_current_tenant_id,
with_current_user,
)
from enums.deployment_edition import DeploymentEdition
from extensions.ext_database import db
from fields.base import ResponseModel
from fields.member_fields import AccountResponse
@@ -262,7 +263,7 @@ class AccountInitApi(Resource):
payload = console_ns.payload or {}
args = AccountInitPayload.model_validate(payload)
if dify_config.EDITION == "CLOUD":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
if not args.invitation_code:
raise ValueError("invitation_code is required")
+1 -1
View File
@@ -198,7 +198,7 @@ def _check_member_invite_limits(tenant_id: str, new_member_count: int, new_accou
if workspace_members.enabled is True and not workspace_members.is_available(new_member_count):
raise WorkspaceMembersLimitExceeded()
if new_account_count > 0:
seats = FeatureService.get_system_features(is_authenticated=True).license.seats
seats = FeatureService.get_license().seats
if not seats.is_available(new_account_count):
raise SeatsLimitExceeded()
return
+4 -8
View File
@@ -442,12 +442,10 @@ register_enum_models(
)
def _default_auto_upgrade_settings(
tenant_id: str,
category: TenantPluginAutoUpgradeCategory,
) -> AutoUpgradeSettingsResponse:
def _missing_auto_upgrade_settings(tenant_id: str) -> AutoUpgradeSettingsResponse:
"""Represent a missing persisted strategy as effectively disabled."""
return {
"strategy_setting": PluginAutoUpgradeService.default_strategy_setting_for_category(category),
"strategy_setting": TenantPluginAutoUpgradeStrategySetting.DISABLED,
"upgrade_time_of_day": PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id),
"upgrade_mode": TenantPluginAutoUpgradeMode.EXCLUDE,
"exclude_plugins": [],
@@ -1135,9 +1133,7 @@ class PluginFetchAutoUpgradeApi(Resource):
args = ParserAutoUpgradeFetch.model_validate(request.args.to_dict(flat=True))
auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, args.category, session=db.session())
auto_upgrade_dict = (
_auto_upgrade_settings_to_dict(auto_upgrade)
if auto_upgrade
else _default_auto_upgrade_settings(tenant_id, args.category)
_auto_upgrade_settings_to_dict(auto_upgrade) if auto_upgrade else _missing_auto_upgrade_settings(tenant_id)
)
return jsonable_encoder(
@@ -37,6 +37,7 @@ from controllers.console.wraps import (
with_current_user,
)
from enums.cloud_plan import CloudPlan
from enums.deployment_edition import DeploymentEdition
from extensions.ext_database import db
from fields.base import ResponseModel
from libs.helper import dump_response, to_timestamp
@@ -233,7 +234,7 @@ class TenantListApi(Resource):
tenants = [tenant for tenant, _ in tenant_rows]
tenant_dicts = []
is_enterprise_only = dify_config.ENTERPRISE_ENABLED and not dify_config.BILLING_ENABLED
is_saas = dify_config.EDITION == "CLOUD" and dify_config.BILLING_ENABLED
is_saas = dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and dify_config.BILLING_ENABLED
tenant_plans: dict[str, SubscriptionPlan] = {}
if is_saas:
+4 -3
View File
@@ -20,6 +20,7 @@ from controllers.common.wraps import (
from controllers.console.auth.error import AuthenticationFailedError, EmailCodeError
from controllers.console.workspace.error import AccountNotInitializedError
from enums.cloud_plan import CloudPlan
from enums.deployment_edition import DeploymentEdition
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from libs.encryption import FieldEncryption
@@ -129,7 +130,7 @@ def account_initialization_required[R](view: Callable[..., R]) -> Callable[...,
def only_edition_cloud[**P, R](view: Callable[P, R]) -> Callable[P, R]:
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs):
if dify_config.EDITION != "CLOUD":
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
abort(404)
return view(*args, **kwargs)
@@ -151,7 +152,7 @@ def only_edition_enterprise[**P, R](view: Callable[P, R]) -> Callable[P, R]:
def only_edition_self_hosted[**P, R](view: Callable[P, R]) -> Callable[P, R]:
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs):
if dify_config.EDITION != "SELF_HOSTED":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
abort(404)
return view(*args, **kwargs)
@@ -327,7 +328,7 @@ def setup_required[R](view: Callable[..., R]) -> Callable[..., R]:
# The overloads keep Resource methods method-aware for pyrefly while
# preserving support for plain functions used in tests and utilities.
# check setup
if dify_config.EDITION == "SELF_HOSTED" and not _is_setup_completed():
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD and not _is_setup_completed():
if os.environ.get("INIT_PASSWORD"):
raise NotInitValidateError()
raise NotSetupError()
+2 -1
View File
@@ -8,6 +8,7 @@ from werkzeug.exceptions import InternalServerError
from configs import dify_config
from core.rbac import RBACPermission, RBACResourceScope
from enums.deployment_edition import DeploymentEdition
from libs.oauth_bearer import Scope, TokenType
from models.account import Account, Tenant, TenantAccountRole
from models.model import App, EndUser
@@ -26,7 +27,7 @@ class CallerKind(StrEnum):
def current_edition() -> Edition:
if dify_config.EDITION == "CLOUD":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
return Edition.SAAS
if dify_config.ENTERPRISE_ENABLED:
return Edition.EE
+8 -15
View File
@@ -2,6 +2,7 @@ from flask_restx import Resource
from controllers.common.schema import register_response_schema_models
from controllers.web import web_ns
from libs.helper import dump_response
from services.feature_service import FeatureService, SystemFeatureModel
register_response_schema_models(web_ns, SystemFeatureModel)
@@ -10,7 +11,10 @@ register_response_schema_models(web_ns, SystemFeatureModel)
@web_ns.route("/system-features")
class SystemFeatureApi(Resource):
@web_ns.doc("get_system_features")
@web_ns.doc(description="Get system feature flags and configuration")
@web_ns.doc(
description="Get the non-sensitive bootstrap snapshot exposed before Console or Web authentication. "
"This is not a general feature registry."
)
@web_ns.doc(responses={200: "System features retrieved successfully", 500: "Internal server error"})
@web_ns.response(
200,
@@ -18,22 +22,11 @@ class SystemFeatureApi(Resource):
web_ns.models[SystemFeatureModel.__name__],
)
def get(self):
"""Get system feature flags and configuration.
Returns the current system feature flags and configuration
that control various functionalities across the platform.
Returns:
dict: System feature configuration object
"""Get the non-sensitive bootstrap snapshot exposed before authentication.
This endpoint is akin to the `SystemFeatureApi` endpoint in api/controllers/console/feature.py,
except it is intended for use by the web app, instead of the console dashboard.
NOTE: This endpoint is unauthenticated by design, as it provides system features
data required for webapp initialization.
Authentication would create circular dependency (can't authenticate without webapp loading).
Only non-sensitive configuration data should be returned by this endpoint.
Authentication configuration must be available before the authentication flow can be selected.
"""
return FeatureService.get_system_features().model_dump()
return dump_response(SystemFeatureModel, FeatureService.get_system_features())
+4 -1
View File
@@ -8,6 +8,7 @@ from configs import dify_config
from controllers.common.schema import register_response_schema_models
from controllers.web import web_ns
from controllers.web.wraps import WebApiResource
from enums.deployment_edition import DeploymentEdition
from extensions.ext_database import db
from extensions.storage.storage_type import StorageType
from fields.base import ResponseModel
@@ -127,7 +128,9 @@ def _build_site_icon_url(*, site: Site, tenant_id: str) -> str | None:
"""Use direct S3 URLs only in Cloud Mode and preserve preview URLs elsewhere."""
if site.icon_type != IconType.IMAGE or not site.icon:
return None
if dify_config.EDITION == "CLOUD" and StorageType(dify_config.STORAGE_TYPE) == StorageType.S3:
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and (
StorageType(dify_config.STORAGE_TYPE) == StorageType.S3
):
return FileService(db.engine).get_file_presigned_url(file_id=site.icon, tenant_id=tenant_id)
return build_icon_url(site.icon_type, site.icon)
+6
View File
@@ -11,6 +11,7 @@ from werkzeug.exceptions import BadRequest, NotFound, Unauthorized
from constants import HEADER_NAME_APP_CODE
from controllers.web.error import WebAppAuthAccessDeniedError, WebAppAuthRequiredError
from core.logging.context import set_identity_context
from extensions.ext_database import db
from libs.passport import PassportService
from libs.token import extract_webapp_passport
@@ -28,6 +29,11 @@ def validate_jwt_token[**P, R](
@wraps(view)
def decorated(*args: P.args, **kwargs: P.kwargs) -> R:
app_model, end_user = decode_jwt_token()
set_identity_context(
tenant_id=end_user.tenant_id,
user_id=end_user.id,
user_type=end_user.type or "end_user",
)
return view(app_model, end_user, *args, **kwargs)
return decorated
+93
View File
@@ -0,0 +1,93 @@
"""Publication visibility rules for calling roster Agents from Workflows.
``Agent.active_config_is_published`` describes whether the editable shared
draft still matches the active snapshot. It is false both before the first
publish and after a published Agent receives new draft edits, so it must not be
used as a runtime availability flag. App-backed Agents are callable from a
Workflow only when the active snapshot has a revision created by a
publish-visible operation. Direct roster Agents are publish-visible by
construction and only need an active snapshot.
"""
from sqlalchemy import and_, or_, select
from sqlalchemy.orm import Session
from sqlalchemy.sql.elements import ColumnElement
from models.agent import Agent, AgentConfigRevision, AgentConfigRevisionOperation, AgentScope, AgentSource
PUBLISH_VISIBLE_APP_BACKED_REVISION_OPERATIONS = frozenset(
{
AgentConfigRevisionOperation.PUBLISH_DRAFT,
AgentConfigRevisionOperation.SAVE_CURRENT_VERSION,
AgentConfigRevisionOperation.SAVE_NEW_VERSION,
AgentConfigRevisionOperation.SAVE_NEW_AGENT,
AgentConfigRevisionOperation.SAVE_TO_ROSTER,
AgentConfigRevisionOperation.RESTORE_VERSION,
}
)
def workflow_callable_active_snapshot_filter() -> ColumnElement[bool]:
"""Return the SQL predicate for an Agent with a Workflow-callable active snapshot.
The caller remains responsible for tenant, roster scope, lifecycle status,
and model configuration filters. The correlated revision lookup makes the
predicate safe to compose into roster pagination queries.
"""
app_backed_agent = or_(
Agent.source == AgentSource.AGENT_APP,
and_(
Agent.source == AgentSource.IMPORTED,
Agent.scope == AgentScope.ROSTER,
Agent.app_id.is_not(None),
),
)
publish_visible_revision_exists = (
select(AgentConfigRevision.id)
.where(
AgentConfigRevision.tenant_id == Agent.tenant_id,
AgentConfigRevision.agent_id == Agent.id,
AgentConfigRevision.current_snapshot_id == Agent.active_config_snapshot_id,
AgentConfigRevision.operation.in_(PUBLISH_VISIBLE_APP_BACKED_REVISION_OPERATIONS),
)
.correlate(Agent)
.exists()
)
return and_(
Agent.active_config_snapshot_id.is_not(None),
or_(
~app_backed_agent,
publish_visible_revision_exists,
),
)
def agent_has_workflow_callable_active_snapshot(*, session: Session, agent: Agent) -> bool:
"""Return whether ``agent`` has an active snapshot visible to Workflow.
This object-level form is useful after ownership and lifecycle checks have
already loaded an Agent. It intentionally ignores dirty draft state so a
previously published snapshot keeps serving while later edits remain
unpublished.
"""
if not agent.active_config_snapshot_id:
return False
is_app_backed = agent.source == AgentSource.AGENT_APP or (
agent.source == AgentSource.IMPORTED and agent.scope == AgentScope.ROSTER and agent.app_id is not None
)
if not is_app_backed:
return True
return bool(
session.scalar(
select(AgentConfigRevision.id)
.where(
AgentConfigRevision.tenant_id == agent.tenant_id,
AgentConfigRevision.agent_id == agent.id,
AgentConfigRevision.current_snapshot_id == agent.active_config_snapshot_id,
AgentConfigRevision.operation.in_(PUBLISH_VISIBLE_APP_BACKED_REVISION_OPERATIONS),
)
.limit(1)
)
)
+1 -1
View File
@@ -71,7 +71,7 @@ def check_credential_policy_compliance(
)
from services.feature_service import FeatureService
if not FeatureService.get_system_features().plugin_manager.enabled or not credential_id:
if not FeatureService.is_plugin_manager_enabled() or not credential_id:
return
# Check if credential exists in database first (if requested)
+2 -1
View File
@@ -6,6 +6,7 @@ from pydantic import BaseModel
from configs import dify_config
from core.entities import DEFAULT_PLUGIN_ID
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit, RestrictModel
from enums.deployment_edition import DeploymentEdition
from graphon.model_runtime.entities.model_entities import ModelType
@@ -49,7 +50,7 @@ class HostingConfiguration:
self.moderation_config = None
def init_app(self, app: Flask):
if dify_config.EDITION != "CLOUD":
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
return
self.provider_map[f"{DEFAULT_PLUGIN_ID}/azure_openai/azure_openai"] = self.init_azure_openai()
+46 -44
View File
@@ -21,6 +21,7 @@ from core.model_manager import ModelInstance, ModelManager
from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.keyword.keyword_factory import Keyword
from core.rag.docstore.dataset_docstore import DatasetDocumentStore
from core.rag.embedding.token_counter import calculate_segment_token_counts
from core.rag.extractor.entity.datasource_type import DatasourceType
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
@@ -113,17 +114,25 @@ class IndexingRunner:
current_user=current_user,
session=session,
)
token_counts = calculate_segment_token_counts(dataset=dataset, documents=documents)
total_tokens = sum(token_counts)
# save segment
self._load_segments(dataset, requeried_document, documents, session)
self._load_segments(
session=session,
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
token_counts=token_counts,
)
session.commit()
# load
self._load(
index_processor=index_processor,
session=session,
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
session=session,
total_tokens=total_tokens,
)
except DocumentIsPausedError:
raise DocumentIsPausedError(f"Document paused, document id: {document_id}")
@@ -190,17 +199,25 @@ class IndexingRunner:
current_user=current_user,
session=session,
)
token_counts = calculate_segment_token_counts(dataset=dataset, documents=documents)
total_tokens = sum(token_counts)
# save segment
self._load_segments(dataset, requeried_document, documents, session)
self._load_segments(
session=session,
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
token_counts=token_counts,
)
session.commit()
# load
self._load(
index_processor=index_processor,
session=session,
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
session=session,
total_tokens=total_tokens,
)
except DocumentIsPausedError:
raise DocumentIsPausedError(f"Document paused, document id: {document_id}")
@@ -225,7 +242,7 @@ class IndexingRunner:
if not dataset:
raise ValueError("no dataset found")
# get exist document_segment list and delete
# get existing document segments
document_segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
@@ -264,15 +281,15 @@ class IndexingRunner:
child_documents.append(child_document)
document.children = child_documents
documents.append(document)
# Preserve the full document total even when only incomplete segments are re-indexed.
total_tokens = sum(document_segment.tokens for document_segment in document_segments)
# build index
index_type = requeried_document.doc_form
index_processor = IndexProcessorFactory(index_type).init_index_processor()
self._load(
index_processor=index_processor,
session=session,
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
session=session,
total_tokens=total_tokens,
)
except DocumentIsPausedError:
raise DocumentIsPausedError(f"Document paused, document id: {document_id}")
@@ -601,28 +618,16 @@ class IndexingRunner:
def _load(
self,
index_processor: BaseIndexProcessor,
session: Session,
dataset: Dataset,
dataset_document: DatasetDocument,
documents: list[Document],
session: Session,
):
"""
insert index and update document/segment status to completed
"""
total_tokens: int,
) -> None:
"""Build indexes and mark the document complete using the token total computed before hash sharding."""
embedding_model_instance = None
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
embedding_model_instance = self._get_model_manager(dataset.tenant_id).get_model_instance(
tenant_id=dataset.tenant_id,
provider=dataset.embedding_model_provider,
model_type=ModelType.TEXT_EMBEDDING,
model=dataset.embedding_model,
)
# chunk nodes by chunk size
# Build indexes using the existing hash-based worker groups.
indexing_start_at = time.perf_counter()
tokens = 0
create_keyword_thread = None
if (
dataset_document.doc_form != IndexStructureType.PARENT_CHILD_INDEX
@@ -659,12 +664,11 @@ class IndexingRunner:
chunk_documents,
dataset.id,
dataset_document.id,
embedding_model_instance,
)
)
for future in futures:
tokens += future.result()
future.result()
if (
dataset_document.doc_form != IndexStructureType.PARENT_CHILD_INDEX
and dataset.indexing_technique == IndexTechniqueType.ECONOMY
@@ -679,7 +683,7 @@ class IndexingRunner:
document_id=dataset_document.id,
after_indexing_status=IndexingStatus.COMPLETED,
extra_update_params={
DatasetDocument.tokens: tokens,
DatasetDocument.tokens: total_tokens,
DatasetDocument.completed_at: naive_utc_now(),
DatasetDocument.indexing_latency: indexing_end_at - indexing_start_at,
DatasetDocument.error: None,
@@ -720,8 +724,7 @@ class IndexingRunner:
chunk_documents: list[Document],
dataset_id: str,
dataset_document_id: str,
embedding_model_instance: ModelInstance | None,
):
) -> None:
with flask_app.app_context():
with session_factory.create_session() as session:
dataset = session.get(Dataset, dataset_id)
@@ -735,11 +738,6 @@ class IndexingRunner:
# check document is paused
self._check_document_paused_status(dataset_document.id)
tokens = 0
if embedding_model_instance:
page_content_list = [document.page_content for document in chunk_documents]
tokens += sum(embedding_model_instance.get_text_embedding_num_tokens(page_content_list))
multimodal_documents = []
for document in chunk_documents:
if document.attachments and dataset.is_multimodal:
@@ -773,8 +771,6 @@ class IndexingRunner:
session.commit()
return tokens
@staticmethod
def _check_document_paused_status(document_id: str):
indexing_cache_key = f"document_{document_id}_is_paused"
@@ -864,8 +860,14 @@ class IndexingRunner:
return documents
def _load_segments(
self, dataset: Dataset, dataset_document: DatasetDocument, documents: list[Document], session: Session
):
self,
session: Session,
dataset: Dataset,
dataset_document: DatasetDocument,
documents: list[Document],
token_counts: list[int],
) -> None:
"""Persist transformed documents and their precomputed token counts before indexing starts."""
# save node to document segment
doc_store = DatasetDocumentStore(
dataset=dataset, user_id=dataset_document.created_by, document_id=dataset_document.id
@@ -873,9 +875,10 @@ class IndexingRunner:
# add document segments
doc_store.add_documents(
session=session,
docs=documents,
save_child=dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX,
session=session,
token_counts=token_counts,
)
# update document status to indexing
@@ -900,7 +903,6 @@ class IndexingRunner:
DocumentSegment.indexing_at: naive_utc_now(),
},
)
pass
class DocumentIsPausedError(Exception):
+34 -2
View File
@@ -6,9 +6,21 @@ using Python's contextvars for thread-safe and async-safe storage.
import uuid
from contextvars import ContextVar
from typing import NamedTuple
class IdentityContext(NamedTuple):
"""Immutable identity values captured for logging."""
tenant_id: str
user_id: str
user_type: str
_request_id: ContextVar[str] = ContextVar("log_request_id", default="")
_trace_id: ContextVar[str] = ContextVar("log_trace_id", default="")
_EMPTY_IDENTITY_CONTEXT = IdentityContext(tenant_id="", user_id="", user_type="")
_identity: ContextVar[IdentityContext] = ContextVar("log_identity", default=_EMPTY_IDENTITY_CONTEXT)
def get_request_id() -> str:
@@ -21,15 +33,35 @@ def get_trace_id() -> str:
return _trace_id.get()
def get_identity_context() -> IdentityContext:
"""Get the immutable tenant, user, and user-type snapshot for logging."""
return _identity.get()
def set_identity_context(
*, tenant_id: str | None = None, user_id: str | None = None, user_type: str | None = None
) -> None:
"""Set primitive identity values already resolved by an authentication boundary."""
_identity.set(
IdentityContext(
tenant_id=tenant_id or "",
user_id=user_id or "",
user_type=user_type or "",
)
)
def init_request_context() -> None:
"""Initialize request context. Call at start of each request."""
"""Initialize request context and discard identity left by earlier work."""
req_id = uuid.uuid4().hex[:10]
trace_id = uuid.uuid5(uuid.NAMESPACE_DNS, req_id).hex
_request_id.set(req_id)
_trace_id.set(trace_id)
_identity.set(_EMPTY_IDENTITY_CONTEXT)
def clear_request_context() -> None:
"""Clear request context. Call at end of request (optional)."""
"""Clear request context at a request or task lifecycle boundary."""
_request_id.set("")
_trace_id.set("")
_identity.set(_EMPTY_IDENTITY_CONTEXT)
+9 -45
View File
@@ -4,10 +4,7 @@ import contextlib
import logging
from typing import override
import flask
from core.logging.context import get_request_id, get_trace_id
from core.logging.structured_formatter import IdentityDict
from core.logging.context import get_identity_context, get_request_id, get_trace_id
class TraceContextFilter(logging.Filter):
@@ -51,49 +48,16 @@ class TraceContextFilter(logging.Filter):
class IdentityContextFilter(logging.Filter):
"""
Filter that adds user identity context to log records.
Extracts tenant_id, user_id, and user_type from Flask-Login current_user.
"""Add an identity snapshot without invoking authentication or database work.
Logging can run while other libraries hold internal locks, so this filter must
only read primitive ContextVar values populated by authentication boundaries.
"""
@override
def filter(self, record: logging.LogRecord) -> bool:
identity = self._extract_identity()
record.tenant_id = identity.get("tenant_id", "")
record.user_id = identity.get("user_id", "")
record.user_type = identity.get("user_type", "")
identity = get_identity_context()
record.tenant_id = identity.tenant_id
record.user_id = identity.user_id
record.user_type = identity.user_type
return True
def _extract_identity(self) -> IdentityDict:
"""Extract identity from current_user if in request context."""
try:
if not flask.has_request_context():
return {}
from flask_login import current_user
# Check if user is authenticated using the proxy
if not current_user.is_authenticated:
return {}
# Access the underlying user object
user = current_user
from models import Account
from models.model import EndUser
identity: IdentityDict = {}
match user:
case Account():
if user.current_tenant_id:
identity["tenant_id"] = user.current_tenant_id
identity["user_id"] = user.id
identity["user_type"] = "account"
case EndUser():
identity["tenant_id"] = user.tenant_id
identity["user_id"] = user.id
identity["user_type"] = user.type or "end_user"
return identity
except Exception:
return {}
+3 -2
View File
@@ -34,6 +34,7 @@ from core.entities.provider_entities import (
from core.helper import encrypter
from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType
from core.helper.position_helper import is_filtered
from enums.deployment_edition import DeploymentEdition
from extensions import ext_hosting_provider
from extensions.ext_database import db
from extensions.ext_redis import redis_client
@@ -743,7 +744,7 @@ class ProviderManager:
if preferred_provider_type_record:
preferred_provider_type = preferred_provider_type_record.preferred_provider_type
elif dify_config.EDITION == "CLOUD" and system_configuration.enabled:
elif dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and system_configuration.enabled:
preferred_provider_type = ProviderType.SYSTEM
elif custom_configuration.provider or custom_configuration.models:
preferred_provider_type = ProviderType.CUSTOM
@@ -1538,7 +1539,7 @@ class ProviderManager:
quota_type_to_provider_records_dict[provider_record.quota_type] = provider_record # type: ignore[index]
quota_configurations = []
if dify_config.EDITION == "CLOUD":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
from services.credit_pool_service import CreditPoolService
trail_pool = CreditPoolService.get_pool(
@@ -10,6 +10,7 @@ from core.rag.rerank.entity.weight import KeywordSetting, VectorSetting, Weights
from core.rag.rerank.rerank_base import BaseRerankRunner
from core.rag.rerank.rerank_factory import RerankRunnerFactory
from core.rag.rerank.rerank_type import RerankMode
from extensions.otel import trace_span
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
@@ -52,6 +53,7 @@ class DataPostProcessor:
)
self.reorder_runner = self._get_reorder_runner(reorder_enabled)
@trace_span()
def invoke(
self,
query: str,
+8 -24
View File
@@ -1,12 +1,10 @@
import concurrent.futures
import functools
import logging
from collections.abc import Callable, Sequence
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor
from typing import Any, NotRequired, TypedDict
from flask import Flask, current_app
from opentelemetry import context as otel_context
from sqlalchemy import select
from sqlalchemy.orm import Session, load_only
@@ -26,7 +24,7 @@ from core.rag.rerank.rerank_type import RerankMode
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from core.tools.signature import sign_upload_file_preview_url
from extensions.ext_database import db
from extensions.otel import trace_span
from extensions.otel import propagate_context, trace_span
from graphon.model_runtime.entities.model_entities import ModelType
from models.dataset import (
ChildChunk,
@@ -92,20 +90,6 @@ default_retrieval_model: DefaultRetrievalModelDict = {
logger = logging.getLogger(__name__)
def _propagate_otel_context[**P, R](func: Callable[P, R]) -> Callable[P, R]:
captured_context = otel_context.get_current()
@functools.wraps(func)
def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
token = otel_context.attach(captured_context)
try:
return func(*args, **kwargs)
finally:
otel_context.detach(token)
return wrapper
class RetrievalService:
# Cache precompiled regular expressions to avoid repeated compilation
@classmethod
@@ -139,7 +123,7 @@ class RetrievalService:
if query:
futures.append(
executor.submit(
_propagate_otel_context(retrieval_service._retrieve),
propagate_context(retrieval_service._retrieve),
flask_app=current_app._get_current_object(), # type: ignore
retrieval_method=retrieval_method,
dataset=dataset,
@@ -159,7 +143,7 @@ class RetrievalService:
for attachment_id in attachment_ids:
futures.append(
executor.submit(
_propagate_otel_context(retrieval_service._retrieve),
propagate_context(retrieval_service._retrieve),
flask_app=current_app._get_current_object(), # type: ignore
retrieval_method=retrieval_method,
dataset=dataset,
@@ -820,7 +804,7 @@ class RetrievalService:
if retrieval_method == RetrievalMethod.KEYWORD_SEARCH and query:
futures.append(
executor.submit(
_propagate_otel_context(self.keyword_search),
propagate_context(self.keyword_search),
flask_app=current_app._get_current_object(), # type: ignore
dataset_id=dataset.id,
query=query,
@@ -834,7 +818,7 @@ class RetrievalService:
if query:
futures.append(
executor.submit(
_propagate_otel_context(self.embedding_search),
propagate_context(self.embedding_search),
flask_app=current_app._get_current_object(), # type: ignore
dataset_id=dataset.id,
query=query,
@@ -851,7 +835,7 @@ class RetrievalService:
if attachment_id:
futures.append(
executor.submit(
_propagate_otel_context(self.embedding_search),
propagate_context(self.embedding_search),
flask_app=current_app._get_current_object(), # type: ignore
dataset_id=dataset.id,
query=attachment_id,
@@ -868,7 +852,7 @@ class RetrievalService:
if RetrievalMethod.is_support_fulltext_search(retrieval_method) and query:
futures.append(
executor.submit(
_propagate_otel_context(self.full_text_index_search),
propagate_context(self.full_text_index_search),
flask_app=current_app._get_current_object(), # type: ignore
dataset_id=dataset.id,
query=query,
+6 -21
View File
@@ -6,10 +6,7 @@ from typing import Any
from sqlalchemy import delete, func, select
from sqlalchemy.orm import Session
from core.model_manager import ModelManager
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.models.document import AttachmentDocument, Document
from graphon.model_runtime.entities.model_entities import ModelType
from models.dataset import ChildChunk, Dataset, DocumentSegment, SegmentAttachmentBinding
from models.enums import SegmentType
@@ -69,34 +66,22 @@ class DatasetDocumentStore:
def add_documents(
self,
docs: Sequence[Document],
session: Session,
docs: Sequence[Document],
token_counts: list[int],
allow_update: bool = True,
save_child: bool = False,
):
) -> None:
document_token_pairs = list(zip(docs, token_counts, strict=True))
max_position = session.scalar(
select(func.max(DocumentSegment.position)).where(DocumentSegment.document_id == self._document_id)
)
if max_position is None:
max_position = 0
embedding_model = None
if self._dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
model_manager = ModelManager.for_tenant(tenant_id=self._dataset.tenant_id)
embedding_model = model_manager.get_model_instance(
tenant_id=self._dataset.tenant_id,
provider=self._dataset.embedding_model_provider,
model_type=ModelType.TEXT_EMBEDDING,
model=self._dataset.embedding_model,
)
if embedding_model:
page_content_list = [doc.page_content for doc in docs]
tokens_list = embedding_model.get_text_embedding_num_tokens(page_content_list)
else:
tokens_list = [0] * len(docs)
for doc, tokens in zip(docs, tokens_list):
for doc, tokens in document_token_pairs:
if not isinstance(doc, Document):
raise ValueError("doc must be a Document")
+25
View File
@@ -0,0 +1,25 @@
"""Token counting for document segments."""
from core.model_manager import ModelManager
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.models.document import Document
from graphon.model_runtime.entities.model_entities import ModelType
from models.dataset import Dataset
def calculate_segment_token_counts(dataset: Dataset, documents: list[Document]) -> list[int]:
"""Return one token count per document, invoking the embedding model only for high-quality indexes."""
if not documents:
return []
if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY:
return [0] * len(documents)
model_manager = ModelManager.for_tenant(tenant_id=dataset.tenant_id)
embedding_model = model_manager.get_model_instance(
tenant_id=dataset.tenant_id,
provider=dataset.embedding_model_provider,
model_type=ModelType.TEXT_EMBEDDING,
model=dataset.embedding_model,
)
return embedding_model.get_text_embedding_num_tokens([document.page_content for document in documents])
@@ -19,6 +19,7 @@ from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.keyword.keyword_factory import Keyword
from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.docstore.dataset_docstore import DatasetDocumentStore
from core.rag.embedding.token_counter import calculate_segment_token_counts
from core.rag.entities import Rule
from core.rag.extractor.entity.extract_setting import ExtractSetting
from core.rag.extractor.extract_processor import ExtractProcessor
@@ -244,10 +245,16 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
all_multimodal_documents.extend(doc.attachments)
documents.append(doc)
if documents:
token_counts = calculate_segment_token_counts(dataset=dataset, documents=documents)
# save node to document segment
doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id)
# add document segments
doc_store.add_documents(docs=documents, save_child=False, session=session)
doc_store.add_documents(
session=session,
docs=documents,
token_counts=token_counts,
save_child=False,
)
session.commit()
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset, session=session)
@@ -15,6 +15,7 @@ from core.model_manager import ModelInstance
from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.docstore.dataset_docstore import DatasetDocumentStore
from core.rag.embedding.token_counter import calculate_segment_token_counts
from core.rag.entities import ParentMode, Rule
from core.rag.extractor.entity.extract_setting import ExtractSetting
from core.rag.extractor.extract_processor import ExtractProcessor
@@ -304,6 +305,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
doc.attachments = self._get_content_files(doc, current_user=account, session=session)
documents.append(doc)
if documents:
token_counts = calculate_segment_token_counts(dataset=dataset, documents=documents)
# update document parent mode
dataset_process_rule = DatasetProcessRule(
dataset_id=dataset.id,
@@ -321,7 +323,12 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
# save node to document segment
doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id)
# add document segments
doc_store.add_documents(docs=documents, save_child=True, session=session)
doc_store.add_documents(
session=session,
docs=documents,
token_counts=token_counts,
save_child=True,
)
session.commit()
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
all_child_documents = []
@@ -17,6 +17,7 @@ from core.llm_generator.llm_generator import LLMGenerator
from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.docstore.dataset_docstore import DatasetDocumentStore
from core.rag.embedding.token_counter import calculate_segment_token_counts
from core.rag.entities import Rule
from core.rag.extractor.entity.extract_setting import ExtractSetting
from core.rag.extractor.extract_processor import ExtractProcessor
@@ -205,9 +206,15 @@ class QAIndexProcessor(BaseIndexProcessor):
doc = Document(page_content=qa_chunk.question, metadata=metadata)
documents.append(doc)
if documents:
token_counts = calculate_segment_token_counts(dataset=dataset, documents=documents)
# save node to document segment
doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id)
doc_store.add_documents(docs=documents, save_child=False, session=session)
doc_store.add_documents(
session=session,
docs=documents,
token_counts=token_counts,
save_child=False,
)
session.commit()
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset, session=session)
+2
View File
@@ -9,6 +9,7 @@ from core.rag.index_processor.constant.query_type import QueryType
from core.rag.models.document import Document
from core.rag.rerank.rerank_base import BaseRerankRunner
from extensions.ext_storage import storage
from extensions.otel import trace_span
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult
from models.model import UploadFile
@@ -22,6 +23,7 @@ class RerankModelRunner(BaseRerankRunner):
self._session = session
@override
@trace_span()
def run(
self,
query: str,
+101 -24
View File
@@ -65,6 +65,7 @@ from core.workflow.nodes.knowledge_retrieval.retrieval import (
)
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from extensions.otel import propagate_context, trace_span
from graphon.file import File, FileTransferMethod, FileType
from graphon.model_runtime.entities.llm_entities import LLMMode, LLMResult, LLMUsage
from graphon.model_runtime.entities.message_entities import PromptMessage, PromptMessageRole, PromptMessageTool
@@ -116,6 +117,7 @@ class DatasetRetrieval:
else:
self._llm_usage = self._llm_usage.plus(usage)
@trace_span()
def knowledge_retrieval(self, session: Session, request: KnowledgeRetrievalRequest) -> list[Source]:
self._check_knowledge_rate_limit(request.tenant_id)
available_datasets = self._get_available_datasets(request.tenant_id, request.dataset_ids)
@@ -599,6 +601,7 @@ class DatasetRetrieval:
return "\n".join([document_context.content for document_context in document_context_list]), context_files
return "", context_files
@trace_span()
def single_retrieve(
self,
session: Session,
@@ -724,7 +727,7 @@ class DatasetRetrieval:
if results:
thread = threading.Thread(
target=self._on_retrieval_end,
target=propagate_context(self._on_retrieval_end),
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"documents": results,
@@ -737,6 +740,7 @@ class DatasetRetrieval:
return results
return []
@trace_span()
def multiple_retrieve(
self,
app_id: str,
@@ -798,7 +802,7 @@ class DatasetRetrieval:
if query:
query_thread = threading.Thread(
target=self._multiple_retrieve_thread,
target=propagate_context(self._multiple_retrieve_thread_safely),
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"available_datasets": available_datasets,
@@ -824,7 +828,7 @@ class DatasetRetrieval:
if attachment_ids:
for attachment_id in attachment_ids:
attachment_thread = threading.Thread(
target=self._multiple_retrieve_thread,
target=propagate_context(self._multiple_retrieve_thread_safely),
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"available_datasets": available_datasets,
@@ -865,7 +869,7 @@ class DatasetRetrieval:
if all_documents:
# add thread to call _on_retrieval_end
retrieval_end_thread = threading.Thread(
target=self._on_retrieval_end,
target=propagate_context(self._on_retrieval_end),
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"documents": all_documents,
@@ -1161,6 +1165,7 @@ class DatasetRetrieval:
all_documents.extend(documents)
@trace_span()
def _run_retriever_thread(
self,
*,
@@ -1172,27 +1177,51 @@ class DatasetRetrieval:
document_ids_filter: list[str] | None,
metadata_condition: MetadataFilteringCondition | None,
attachment_ids: list[str] | None,
) -> None:
with session_factory.create_session() as session:
self._retriever(
flask_app=flask_app,
session=session,
dataset_id=dataset_id,
query=query or "",
top_k=top_k,
all_documents=all_documents,
document_ids_filter=document_ids_filter,
metadata_condition=metadata_condition,
attachment_ids=attachment_ids,
)
def _run_retriever_thread_safely(
self,
*,
flask_app: Flask,
dataset_id: str,
query: str | None,
top_k: int,
all_documents: list[Document],
document_ids_filter: list[str] | None,
metadata_condition: MetadataFilteringCondition | None,
attachment_ids: list[str] | None,
cancel_event: threading.Event | None,
thread_exceptions: list[Exception] | None,
) -> None:
"""Collect errors only after they pass through the traced retrieval method."""
try:
with session_factory.create_session() as session:
self._retriever(
flask_app=flask_app,
session=session,
dataset_id=dataset_id,
query=query or "",
top_k=top_k,
all_documents=all_documents,
document_ids_filter=document_ids_filter,
metadata_condition=metadata_condition,
attachment_ids=attachment_ids,
)
except Exception as e:
self._run_retriever_thread(
flask_app=flask_app,
dataset_id=dataset_id,
query=query,
top_k=top_k,
all_documents=all_documents,
document_ids_filter=document_ids_filter,
metadata_condition=metadata_condition,
attachment_ids=attachment_ids,
)
except Exception as exc:
if cancel_event:
cancel_event.set()
if thread_exceptions is not None:
thread_exceptions.append(e)
thread_exceptions.append(exc)
def to_dataset_retriever_tool(
self,
@@ -1795,6 +1824,7 @@ class DatasetRetrieval:
return full_text, usage
@trace_span()
def _multiple_retrieve_thread(
self,
flask_app: Flask,
@@ -1813,11 +1843,11 @@ class DatasetRetrieval:
attachment_id: str | None,
dataset_count: int,
cancel_event: threading.Event | None = None,
thread_exceptions: list[Exception] | None = None,
):
) -> None:
try:
with flask_app.app_context():
threads = []
retrieval_thread_exceptions: list[Exception] = []
all_documents_item: list[Document] = []
index_type = None
for dataset in available_datasets:
@@ -1836,7 +1866,7 @@ class DatasetRetrieval:
else:
continue
retrieval_thread = threading.Thread(
target=self._run_retriever_thread,
target=propagate_context(self._run_retriever_thread_safely),
kwargs={
"flask_app": flask_app,
"dataset_id": dataset.id,
@@ -1847,7 +1877,7 @@ class DatasetRetrieval:
"metadata_condition": metadata_condition,
"attachment_ids": [attachment_id] if attachment_id else None,
"cancel_event": cancel_event,
"thread_exceptions": thread_exceptions,
"thread_exceptions": retrieval_thread_exceptions,
},
)
threads.append(retrieval_thread)
@@ -1862,6 +1892,9 @@ class DatasetRetrieval:
if cancel_event and cancel_event.is_set():
break
if retrieval_thread_exceptions:
raise retrieval_thread_exceptions[0]
# Skip second reranking when there is only one dataset
if reranking_enable and dataset_count > 1:
# do rerank for searched documents
@@ -1902,11 +1935,55 @@ class DatasetRetrieval:
all_documents_item = all_documents_item[:top_k] if top_k else all_documents_item
if all_documents_item:
all_documents.extend(all_documents_item)
except Exception as e:
except Exception:
raise
def _multiple_retrieve_thread_safely(
self,
*,
flask_app: Flask,
available_datasets: list[Dataset],
metadata_condition: MetadataFilteringCondition | None,
metadata_filter_document_ids: dict[str, list[str]] | None,
all_documents: list[Document],
tenant_id: str,
reranking_enable: bool,
reranking_mode: str,
reranking_model: RerankingModelDict | None,
weights: WeightsDict | None,
top_k: int,
score_threshold: float,
query: str | None,
attachment_id: str | None,
dataset_count: int,
cancel_event: threading.Event | None = None,
thread_exceptions: list[Exception] | None = None,
) -> None:
"""Collect errors only after they pass through the traced multi-retrieval method."""
try:
self._multiple_retrieve_thread(
flask_app=flask_app,
available_datasets=available_datasets,
metadata_condition=metadata_condition,
metadata_filter_document_ids=metadata_filter_document_ids,
all_documents=all_documents,
tenant_id=tenant_id,
reranking_enable=reranking_enable,
reranking_mode=reranking_mode,
reranking_model=reranking_model,
weights=weights,
top_k=top_k,
score_threshold=score_threshold,
query=query,
attachment_id=attachment_id,
dataset_count=dataset_count,
cancel_event=cancel_event,
)
except Exception as exc:
if cancel_event:
cancel_event.set()
if thread_exceptions is not None:
thread_exceptions.append(e)
thread_exceptions.append(exc)
def _get_available_datasets(self, tenant_id: str, dataset_ids: list[str]) -> list[Dataset]:
with session_factory.create_session() as session:
@@ -4,8 +4,16 @@ from dataclasses import dataclass
from sqlalchemy import select
from core.agent.publish_visibility import workflow_callable_active_snapshot_filter
from core.db.session_factory import session_factory
from models.agent import Agent, AgentConfigSnapshot, AgentStatus, WorkflowAgentBindingType, WorkflowAgentNodeBinding
from models.agent import (
Agent,
AgentConfigSnapshot,
AgentScope,
AgentStatus,
WorkflowAgentBindingType,
WorkflowAgentNodeBinding,
)
class WorkflowAgentBindingError(Exception):
@@ -24,7 +32,7 @@ class WorkflowAgentBindingBundle:
class WorkflowAgentBindingResolver:
"""Resolve the Agent binding owned by the current workflow id and node id."""
"""Resolve an owned binding without allowing unpublished roster snapshots to run."""
def resolve(
self,
@@ -53,18 +61,20 @@ class WorkflowAgentBindingResolver:
if binding.agent_id is None:
raise WorkflowAgentBindingError("agent_not_available", "Workflow Agent binding has no agent.")
agent = session.scalar(
select(Agent)
.where(
Agent.tenant_id == tenant_id,
Agent.id == binding.agent_id,
)
.limit(1)
agent_stmt = select(Agent).where(
Agent.tenant_id == tenant_id,
Agent.id == binding.agent_id,
)
if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT:
agent_stmt = agent_stmt.where(
Agent.scope == AgentScope.ROSTER,
workflow_callable_active_snapshot_filter(),
)
agent = session.scalar(agent_stmt.limit(1))
if agent is None or agent.status == AgentStatus.ARCHIVED:
raise WorkflowAgentBindingError(
"agent_not_available",
f"Agent {binding.agent_id} is not available.",
f"Agent {binding.agent_id} is not available or has not been published.",
)
snapshot_id = (
+26 -11
View File
@@ -6,9 +6,17 @@ from typing import Any
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.agent.publish_visibility import workflow_callable_active_snapshot_filter
from core.workflow.graph_topology import WorkflowGraphTopology
from graphon.enums import BuiltinNodeTypes
from models.agent import Agent, AgentConfigSnapshot, AgentStatus, WorkflowAgentBindingType, WorkflowAgentNodeBinding
from models.agent import (
Agent,
AgentConfigSnapshot,
AgentScope,
AgentStatus,
WorkflowAgentBindingType,
WorkflowAgentNodeBinding,
)
from models.agent_config_entities import (
AgentFileRefConfig,
AgentHumanContactConfig,
@@ -117,21 +125,28 @@ class WorkflowAgentNodeValidator:
binding: WorkflowAgentNodeBinding,
topology: _WorkflowGraphTopology | None = None,
) -> None:
"""Validate binding ownership, publication state, Agent Soul, and node-job references."""
if binding.agent_id is None:
raise WorkflowAgentNodeValidationError(f"Workflow Agent node {binding.node_id} is missing agent binding.")
agent = session.scalar(
select(Agent)
.where(
Agent.tenant_id == binding.tenant_id,
Agent.id == binding.agent_id,
)
.limit(1)
agent_stmt = select(Agent).where(
Agent.tenant_id == binding.tenant_id,
Agent.id == binding.agent_id,
)
if agent is None or agent.status == AgentStatus.ARCHIVED:
raise WorkflowAgentNodeValidationError(
f"Workflow Agent node {binding.node_id} references an unavailable agent."
if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT:
agent_stmt = agent_stmt.where(
Agent.scope == AgentScope.ROSTER,
workflow_callable_active_snapshot_filter(),
)
agent = session.scalar(agent_stmt.limit(1))
if agent is None or agent.status == AgentStatus.ARCHIVED:
availability = (
"an unavailable or unpublished roster agent"
if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT
else "an unavailable agent"
)
raise WorkflowAgentNodeValidationError(f"Workflow Agent node {binding.node_id} references {availability}.")
snapshot_id = (
agent.active_config_snapshot_id
+11
View File
@@ -0,0 +1,11 @@
from enum import StrEnum
class DeploymentEdition(StrEnum):
"""
Enum representing the deployment edition of the platform.
"""
COMMUNITY = "COMMUNITY"
ENTERPRISE = "ENTERPRISE"
CLOUD = "CLOUD"
+19 -4
View File
@@ -1,5 +1,6 @@
import json
from typing import cast, override
import logging
from typing import assert_never, cast, override
import flask_login
from flask import Request, Response, request
@@ -11,6 +12,7 @@ from werkzeug.exceptions import NotFound, Unauthorized
from configs import dify_config
from constants import HEADER_NAME_APP_CODE
from core.db.session_factory import session_factory
from core.logging.context import set_identity_context
from dify_app import DifyApp
from libs.passport import PassportService
from libs.token import extract_access_token, extract_console_cookie_token, extract_webapp_passport
@@ -19,6 +21,8 @@ from models.enums import EndUserType
from models.model import AppMCPServer, EndUser
from services.account_service import AccountService
logger = logging.getLogger(__name__)
type LoginUser = Account | EndUser
@@ -156,13 +160,24 @@ def _load_user_from_request(request_from_flask_login: Request, session: Session)
@user_logged_in.connect
@user_loaded_from_request.connect
def on_user_logged_in(_sender: object, user: LoginUser) -> None:
"""Called when a user logged in.
"""Snapshot authenticated identity into the side-effect-free logging context.
Note: AccountService.load_logged_in_account will populate user.current_tenant_id
through the load_user method, which calls account.set_tenant_id_with_session().
"""
# tenant_id context variable removed - using current_user.current_tenant_id directly
pass
set_identity_context()
try:
match user:
case Account():
set_identity_context(tenant_id=user.current_tenant_id, user_id=user.id, user_type="account")
case EndUser():
set_identity_context(tenant_id=user.tenant_id, user_id=user.id, user_type=user.type or "end_user")
case _ as unreachable:
assert_never(unreachable)
except Exception:
# Logging enrichment must never make authentication fail.
logger.exception("Failed to set logging identity context")
return
@login_manager.unauthorized_handler
+2
View File
@@ -1,3 +1,4 @@
from extensions.otel.context import propagate_context
from extensions.otel.decorators.base import trace_span
from extensions.otel.decorators.handler import SpanHandler
from extensions.otel.decorators.handlers.generate_handler import AppGenerateHandler
@@ -7,5 +8,6 @@ __all__ = [
"AppGenerateHandler",
"SpanHandler",
"WorkflowAppRunnerHandler",
"propagate_context",
"trace_span",
]
+21
View File
@@ -0,0 +1,21 @@
"""Utilities for propagating OpenTelemetry context across execution boundaries."""
import functools
from collections.abc import Callable
from opentelemetry import context as otel_context
def propagate_context[**P, R](func: Callable[P, R]) -> Callable[P, R]:
"""Capture the current context and attach it whenever ``func`` executes."""
captured_context = otel_context.get_current()
@functools.wraps(func)
def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
token = otel_context.attach(captured_context)
try:
return func(*args, **kwargs)
finally:
otel_context.detach(token)
return wrapper
+8 -3
View File
@@ -69,12 +69,17 @@ def on_user_loaded(_sender, user: Union["Account", "EndUser"]):
if user:
try:
current_span = get_current_span()
if not current_span.is_recording():
return
tenant_id = extract_tenant_id(user)
if not tenant_id:
return
if current_span:
current_span.set_attribute(DifySpanAttributes.TENANT_ID, tenant_id)
current_span.set_attribute(GenAIAttributes.USER_ID, user.id)
current_span.set_attributes(
{
DifySpanAttributes.TENANT_ID: tenant_id,
GenAIAttributes.USER_ID: user.id,
}
)
except Exception:
logger.exception("Error setting tenant and user attributes")
pass
+36 -18
View File
@@ -9475,15 +9475,11 @@ Used for frontend component type mapping
| 200 | Success | **application/json**: [SchemaDefinitionsResponse](#schemadefinitionsresponse)<br> |
### [GET] /system-features
**Get system-wide feature configuration**
**Get the non-sensitive bootstrap snapshot exposed before authentication**
Get system-wide feature configuration
NOTE: This endpoint is unauthenticated by design, as it provides system features
data required for dashboard initialization.
Authentication would create circular dependency (can't login without dashboard loading).
Only non-sensitive configuration data should be returned by this endpoint.
Get the non-sensitive bootstrap snapshot exposed before Console or Web authentication. This is not a general feature registry.
Authentication configuration must be available before the authentication flow can be selected.
Authenticated license detail is served separately by SystemFeatureLicenseApi.
#### Responses
@@ -9491,6 +9487,19 @@ Only non-sensitive configuration data should be returned by this endpoint.
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [SystemFeatureModel](#systemfeaturemodel)<br> |
### [GET] /system-features/license
**Get full license detail (status, expiry, workspace/seat usage)**
Get license status and usage detail
Authenticated counterpart to the license *status* exposed on the public
system-features endpoint.
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [LicenseModel](#licensemodel)<br> |
### [POST] /tag-bindings
#### Request Body
@@ -17318,6 +17327,14 @@ Default model entity.
| tool_name | string | | Yes |
| type | string | | Yes |
#### DeploymentEdition
Enum representing the deployment edition of the platform.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| DeploymentEdition | string | Enum representing the deployment edition of the platform. | |
#### DismissNotificationPayload
| Name | Type | Description | Required |
@@ -18815,6 +18832,12 @@ Enum class for large language model mode.
| ---- | ---- | ----------- | -------- |
| LicenseStatus | string | | |
#### LicenseStatusModel
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| status | [LicenseStatus](#licensestatus) | | Yes |
#### LimitationModel
| Name | Type | Description | Required |
@@ -19559,6 +19582,7 @@ Coarse node-level status used by Inspector to pick a banner.
| ---- | ---- | ----------- | -------- |
| avatar | string | | No |
| email | string | | Yes |
| id | string | | Yes |
| interface_language | string | | Yes |
| name | string | | Yes |
| timezone | string | | Yes |
@@ -20492,12 +20516,6 @@ Shared permission levels for resources (datasets, credentials, etc.)
| plugins | [ [PluginEntity](#pluginentity) ] | | Yes |
| total | integer | | Yes |
#### PluginManagerModel
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| enabled | boolean | | Yes |
#### PluginManifestResponse
| Name | Type | Description | Required |
@@ -22023,9 +22041,12 @@ Model class for provider system configuration response.
#### SystemFeatureModel
Non-sensitive bootstrap snapshot exposed before Console or Web authentication.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| branding | [BrandingModel](#brandingmodel) | | Yes |
| deployment_edition | [DeploymentEdition](#deploymentedition) | | Yes |
| enable_app_deploy | boolean | | Yes |
| enable_change_email | boolean, <br>**Default:** true | | Yes |
| enable_collaboration_mode | boolean, <br>**Default:** true | | Yes |
@@ -22038,14 +22059,11 @@ Model class for provider system configuration response.
| enable_social_oauth_login | boolean | | Yes |
| enable_step_by_step_tour | boolean | | Yes |
| enable_trial_app | boolean | | Yes |
| is_allow_create_workspace | boolean | | Yes |
| is_allow_register | boolean | | Yes |
| is_email_setup | boolean | | Yes |
| knowledge_fs_enabled | boolean | | Yes |
| license | [LicenseModel](#licensemodel) | | Yes |
| max_plugin_package_size | integer, <br>**Default:** 15728640 | | Yes |
| license | [LicenseStatusModel](#licensestatusmodel) | | Yes |
| plugin_installation_permission | [PluginInstallationPermissionModel](#plugininstallationpermissionmodel) | | Yes |
| plugin_manager | [PluginManagerModel](#pluginmanagermodel) | | Yes |
| rbac_enabled | boolean | | Yes |
| sso_enforced_for_signin | boolean | | Yes |
| sso_enforced_for_signin_protocol | string | | Yes |
+21 -45
View File
@@ -774,24 +774,13 @@ Retrieve app site information and configuration.
| 500 | Internal Server Error | |
### [GET] /system-features
**Get system feature flags and configuration**
Get system feature flags and configuration
Returns the current system feature flags and configuration
that control various functionalities across the platform.
Returns:
dict: System feature configuration object
**Get the non-sensitive bootstrap snapshot exposed before authentication**
Get the non-sensitive bootstrap snapshot exposed before Console or Web authentication. This is not a general feature registry.
This endpoint is akin to the `SystemFeatureApi` endpoint in api/controllers/console/feature.py,
except it is intended for use by the web app, instead of the console dashboard.
NOTE: This endpoint is unauthenticated by design, as it provides system features
data required for webapp initialization.
Authentication would create circular dependency (can't authenticate without webapp loading).
Only non-sensitive configuration data should be returned by this endpoint.
Authentication configuration must be available before the authentication flow can be selected.
#### Responses
@@ -1058,6 +1047,14 @@ Button styles for user actions.
| auto_generate | boolean | Automatically generate the conversation name. When `true`, the `name` field is ignored. | No |
| name | string | Conversation name. Required when `auto_generate` is `false`. | No |
#### DeploymentEdition
Enum representing the deployment edition of the platform.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| DeploymentEdition | string | Enum representing the deployment edition of the platform. | |
#### EmailCodeLoginSendPayload
| Name | Type | Description | Required |
@@ -1284,33 +1281,18 @@ Parsed multipart form fields for HITL uploads.
| ---- | ---- | ----------- | -------- |
| JsonValue | | | |
#### LicenseLimitationModel
- enabled: whether this limit is enforced
- size: current usage count
- limit: maximum allowed count; 0 means unlimited
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| enabled | boolean | Whether this limit is currently active | Yes |
| limit | integer | Maximum number of resources allowed; 0 means no limit | Yes |
| size | integer | Number of resources already consumed | Yes |
#### LicenseModel
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| expired_at | string | | Yes |
| seats | [LicenseLimitationModel](#licenselimitationmodel) | | Yes |
| status | [LicenseStatus](#licensestatus) | | Yes |
| workspaces | [LicenseLimitationModel](#licenselimitationmodel) | | Yes |
#### LicenseStatus
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| LicenseStatus | string | | |
#### LicenseStatusModel
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| status | [LicenseStatus](#licensestatus) | | Yes |
#### LoginPayload
| Name | Type | Description | Required |
@@ -1419,12 +1401,6 @@ Form input definition.
| ---- | ---- | ----------- | -------- |
| PluginInstallationScope | string | | |
#### PluginManagerModel
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| enabled | boolean | | Yes |
#### RemoteFileInfo
| Name | Type | Description | Required |
@@ -1564,9 +1540,12 @@ Default configuration for form inputs.
#### SystemFeatureModel
Non-sensitive bootstrap snapshot exposed before Console or Web authentication.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| branding | [BrandingModel](#brandingmodel) | | Yes |
| deployment_edition | [DeploymentEdition](#deploymentedition) | | Yes |
| enable_app_deploy | boolean | | Yes |
| enable_change_email | boolean, <br>**Default:** true | | Yes |
| enable_collaboration_mode | boolean, <br>**Default:** true | | Yes |
@@ -1579,14 +1558,11 @@ Default configuration for form inputs.
| enable_social_oauth_login | boolean | | Yes |
| enable_step_by_step_tour | boolean | | Yes |
| enable_trial_app | boolean | | Yes |
| is_allow_create_workspace | boolean | | Yes |
| is_allow_register | boolean | | Yes |
| is_email_setup | boolean | | Yes |
| knowledge_fs_enabled | boolean | | Yes |
| license | [LicenseModel](#licensemodel) | | Yes |
| max_plugin_package_size | integer, <br>**Default:** 15728640 | | Yes |
| license | [LicenseStatusModel](#licensestatusmodel) | | Yes |
| plugin_installation_permission | [PluginInstallationPermissionModel](#plugininstallationpermissionmodel) | | Yes |
| plugin_manager | [PluginManagerModel](#pluginmanagermodel) | | Yes |
| rbac_enabled | boolean | | Yes |
| sso_enforced_for_signin | boolean | | Yes |
| sso_enforced_for_signin_protocol | string | | Yes |
+13 -21
View File
@@ -442,9 +442,9 @@ class AccountService:
# A licensed seat is one Account row, deployment-wide; joining an existing
# account into another workspace does not pass through here and costs no seat.
# is_authenticated=True: server-side enforcement needs the full license payload,
# which the enterprise fill withholds from unauthenticated (browser-facing) calls.
if not FeatureService.get_system_features(is_authenticated=True).license.seats.is_available():
# get_license() carries the full license payload that server-side enforcement needs;
# the public system-features endpoint exposes only license status.
if not FeatureService.get_license().seats.is_available():
raise SeatsLimitExceededError("licensed seats limit exceeded")
if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email):
@@ -1229,12 +1229,12 @@ class AccountService:
if hour_limit_count >= 1:
redis_client.setex(freeze_key, 60 * 60, 1)
return True
else:
redis_client.setex(hour_limit_key, 60 * 10, hour_limit_count + 1) # first time limit 10 minutes
# add hour limit count
redis_client.incr(hour_limit_key)
redis_client.expire(hour_limit_key, 60 * 60)
# First strike claims a 10-minute window atomically; a concurrent
# over-limit request that loses the claim is the second strike and
# freezes the IP for an hour.
if not redis_client.set(hour_limit_key, 1, ex=60 * 10, nx=True):
redis_client.setex(freeze_key, 60 * 60, 1)
return True
@@ -1258,11 +1258,7 @@ class TenantService:
session: Session,
) -> Tenant:
"""Create tenant"""
if (
not FeatureService.get_system_features().is_allow_create_workspace
and not is_setup
and not is_from_dashboard
):
if not FeatureService.is_workspace_creation_allowed() and not is_setup and not is_from_dashboard:
from controllers.console.error import NotAllowedCreateWorkspace
raise NotAllowedCreateWorkspace()
@@ -1325,14 +1321,10 @@ class TenantService:
owner. It persists the legacy membership before creating the matching
RBAC role binding, then makes the workspace current for the account.
"""
if (
not FeatureService.get_system_features().is_allow_create_workspace
and not is_setup
and not is_from_dashboard
):
if not FeatureService.is_workspace_creation_allowed() and not is_setup and not is_from_dashboard:
raise WorkSpaceNotAllowedCreateError()
workspaces = FeatureService.get_system_features().license.workspaces
workspaces = FeatureService.get_license().workspaces
if not workspaces.is_available():
raise WorkspacesLimitExceededError()
@@ -2010,9 +2002,9 @@ class RegisterService:
AccountService.link_account_integrate(provider, open_id, account, session=session)
if (
FeatureService.get_system_features().is_allow_create_workspace
FeatureService.is_workspace_creation_allowed()
and create_workspace_required
and FeatureService.get_system_features().license.workspaces.is_available()
and FeatureService.get_license().workspaces.is_available()
):
try:
TenantService.create_owner_tenant(account, session=session)
+6 -29
View File
@@ -7,6 +7,7 @@ from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from sqlalchemy.sql.elements import ColumnElement
from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot
from libs.helper import to_timestamp
from models import Account
from models.agent import (
@@ -271,6 +272,8 @@ class AgentComposerService:
source_snapshot_id: str | None = None,
idempotency_key: str | None = None,
) -> dict[str, Any]:
"""Copy a callable roster Agent snapshot into a workflow-owned inline Agent."""
workflow = cls._get_draft_workflow(session=session, tenant_id=tenant_id, app_id=app_id)
binding = cls._require_binding(
cls._get_workflow_binding(session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id)
@@ -296,6 +299,8 @@ class AgentComposerService:
source_agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=source_agent_id)
if source_agent.scope != AgentScope.ROSTER or source_agent.status != AgentStatus.ACTIVE:
raise InvalidComposerConfigError("Source agent must be an active roster agent.")
if not agent_has_workflow_callable_active_snapshot(session=session, agent=source_agent):
raise InvalidComposerConfigError("Source agent must have a published config snapshot.")
source_version = cls._require_version(
session=session,
tenant_id=tenant_id,
@@ -529,39 +534,11 @@ class AgentComposerService:
)
if not active_version:
return False
if agent.source in APP_BACKED_AGENT_SOURCES and not cls._has_publish_visible_revision(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
snapshot_id=agent.active_config_snapshot_id,
):
if not agent_has_workflow_callable_active_snapshot(session=session, agent=agent):
return False
return _agent_soul_config_json(agent_soul) == _agent_soul_config_json(active_version.config_snapshot_dict)
@classmethod
def _has_publish_visible_revision(
cls, *, session: Session, tenant_id: str, agent_id: str, snapshot_id: str
) -> bool:
revisions = session.scalars(
select(AgentConfigRevision.operation).where(
AgentConfigRevision.tenant_id == tenant_id,
AgentConfigRevision.agent_id == agent_id,
AgentConfigRevision.current_snapshot_id == snapshot_id,
)
).all()
return any(
operation
in {
AgentConfigRevisionOperation.PUBLISH_DRAFT,
AgentConfigRevisionOperation.SAVE_NEW_VERSION,
AgentConfigRevisionOperation.SAVE_TO_ROSTER,
AgentConfigRevisionOperation.RESTORE_VERSION,
}
for operation in revisions
)
@classmethod
def publish_agent_app_draft(
cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, version_note: str | None = None
+252 -52
View File
@@ -1,8 +1,11 @@
from __future__ import annotations
import json
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime
from decimal import Decimal
from enum import Enum
from typing import Any
import sqlalchemy as sa
@@ -11,6 +14,7 @@ from sqlalchemy.orm import aliased
from configs import dify_config
from core.app.entities.app_invoke_entities import InvokeFrom
from graphon.enums import WorkflowNodeExecutionStatus
from libs.helper import convert_datetime_to_date, escape_like_pattern, to_timestamp
from models.agent import WorkflowAgentNodeBinding
from models.enums import CreatorUserRole, MessageStatus
@@ -194,11 +198,12 @@ class AgentObservabilityService:
self, *, app: App, agent_id: str, conversation_id: str, params: AgentLogQueryParams
) -> dict[str, Any]:
source_filters = self.resolve_source_filters(params.sources)
rows: list[Message] = []
rows: list[dict[str, Any]] = []
for source_filter in source_filters:
if source_filter.kind in {"all", "webapp"}:
rows.extend(
self._list_webapp_messages(
self.serialize_log_message(message)
for message in self._list_webapp_messages(
app=app,
conversation_id=conversation_id,
params=params,
@@ -216,18 +221,18 @@ class AgentObservabilityService:
)
)
deduped = {message.id: message for message in rows}
sort_column = Message.created_at if params.sort_by == "created_at" else Message.updated_at
deduped = {row["id"]: row for row in rows}
sort_key = "created_at" if params.sort_by == "created_at" else "updated_at"
sorted_rows = sorted(
deduped.values(),
key=lambda message: (getattr(message, sort_column.key), message.id),
key=lambda row: (row[sort_key] or 0, row["id"]),
reverse=params.sort_order != "asc",
)
total = len(sorted_rows)
start = (params.page - 1) * params.limit
end = start + params.limit
return {
"data": [self.serialize_log_message(message) for message in sorted_rows[start:end]],
"data": sorted_rows[start:end],
"page": params.page,
"limit": params.limit,
"total": total,
@@ -284,22 +289,20 @@ class AgentObservabilityService:
workflow_app = aliased(App)
stmt = (
select(
Conversation,
WorkflowNodeExecutionModel.id.label("node_execution_id"),
WorkflowNodeExecutionModel.title.label("node_title"),
WorkflowNodeExecutionModel.status.label("node_status"),
WorkflowNodeExecutionModel.created_by_role.label("node_created_by_role"),
WorkflowNodeExecutionModel.created_by.label("node_created_by"),
WorkflowNodeExecutionModel.created_at.label("node_created_at"),
WorkflowNodeExecutionModel.finished_at.label("node_finished_at"),
workflow_app,
WorkflowAgentNodeBinding.workflow_id,
WorkflowAgentNodeBinding.workflow_version,
WorkflowAgentNodeBinding.node_id,
func.count(sa.distinct(Message.id)).label("message_count"),
func.max(Message.created_at).label("created_at"),
func.max(Message.updated_at).label("updated_at"),
func.sum(sa.case((Message.status == MessageStatus.PAUSED, 1), else_=0)).label("paused_count"),
func.sum(
sa.case((or_(Message.error.is_not(None), Message.status == MessageStatus.ERROR), 1), else_=0)
).label("failed_count"),
)
.select_from(Message)
.join(Conversation, Conversation.id == Message.conversation_id)
.join(WorkflowRun, WorkflowRun.id == Message.workflow_run_id)
.select_from(WorkflowNodeExecutionModel)
.join(WorkflowRun, WorkflowRun.id == WorkflowNodeExecutionModel.workflow_run_id)
.join(
WorkflowAgentNodeBinding,
and_(
@@ -310,40 +313,32 @@ class AgentObservabilityService:
WorkflowAgentNodeBinding.workflow_version == WorkflowRun.version,
),
)
.join(
WorkflowNodeExecutionModel,
and_(
WorkflowNodeExecutionModel.workflow_run_id == WorkflowRun.id,
WorkflowNodeExecutionModel.node_id == WorkflowAgentNodeBinding.node_id,
),
)
.join(workflow_app, workflow_app.id == WorkflowAgentNodeBinding.app_id)
.where(Message.workflow_run_id.is_not(None), Conversation.app_id == WorkflowAgentNodeBinding.app_id)
.group_by(
Conversation.id,
workflow_app.id,
WorkflowAgentNodeBinding.workflow_id,
WorkflowAgentNodeBinding.workflow_version,
WorkflowAgentNodeBinding.node_id,
.where(
WorkflowNodeExecutionModel.tenant_id == app.tenant_id,
WorkflowNodeExecutionModel.app_id == WorkflowAgentNodeBinding.app_id,
WorkflowNodeExecutionModel.workflow_id == WorkflowAgentNodeBinding.workflow_id,
WorkflowNodeExecutionModel.node_id == WorkflowAgentNodeBinding.node_id,
)
)
stmt = self._apply_observability_filters(stmt, params=params, source_filter=source_filter)
stmt = self._apply_workflow_node_filters(stmt, params=params, workflow_app=workflow_app)
stmt = self._apply_workflow_source_filter(stmt, source_filter)
rows = list(self._session.execute(stmt).all())
return [
self._serialize_conversation_log(
conversation=row[0],
message_count=row.message_count,
paused_count=row.paused_count,
failed_count=row.failed_count,
self._serialize_workflow_execution_log(
node_execution_id=row.node_execution_id,
title=row.node_title,
status=row.node_status,
created_by_role=row.node_created_by_role,
created_by=row.node_created_by,
created_at=row.node_created_at,
finished_at=row.node_finished_at,
source=self._serialize_workflow_source(
app=row[1],
app=row[7],
workflow_id=row.workflow_id,
workflow_version=row.workflow_version,
node_id=row.node_id,
),
created_at=row.created_at,
updated_at=row.updated_at,
)
for row in rows
]
@@ -363,10 +358,12 @@ class AgentObservabilityService:
conversation_id: str,
params: AgentLogQueryParams,
source_filter: AgentSourceFilter,
) -> list[Message]:
) -> list[dict[str, Any]]:
workflow_app = aliased(App)
stmt = (
select(Message)
.join(WorkflowRun, WorkflowRun.id == Message.workflow_run_id)
select(WorkflowNodeExecutionModel)
.select_from(WorkflowNodeExecutionModel)
.join(WorkflowRun, WorkflowRun.id == WorkflowNodeExecutionModel.workflow_run_id)
.join(
WorkflowAgentNodeBinding,
and_(
@@ -377,18 +374,23 @@ class AgentObservabilityService:
WorkflowAgentNodeBinding.workflow_version == WorkflowRun.version,
),
)
.join(
WorkflowNodeExecutionModel,
and_(
WorkflowNodeExecutionModel.workflow_run_id == WorkflowRun.id,
WorkflowNodeExecutionModel.node_id == WorkflowAgentNodeBinding.node_id,
),
.join(workflow_app, workflow_app.id == WorkflowAgentNodeBinding.app_id)
.where(
WorkflowNodeExecutionModel.id == conversation_id,
WorkflowNodeExecutionModel.tenant_id == app.tenant_id,
WorkflowNodeExecutionModel.app_id == WorkflowAgentNodeBinding.app_id,
WorkflowNodeExecutionModel.workflow_id == WorkflowAgentNodeBinding.workflow_id,
WorkflowNodeExecutionModel.node_id == WorkflowAgentNodeBinding.node_id,
)
.where(Message.conversation_id == conversation_id)
)
stmt = self._apply_message_filters(stmt, params=params, source_filter=source_filter)
stmt = self._apply_workflow_node_filters(stmt, params=params, workflow_app=workflow_app)
stmt = self._apply_workflow_source_filter(stmt, source_filter)
return list(self._session.scalars(stmt.order_by(Message.created_at.desc(), Message.id.desc())).all())
executions = list(
self._session.scalars(
stmt.order_by(WorkflowNodeExecutionModel.created_at.desc(), WorkflowNodeExecutionModel.id.desc())
).all()
)
return [self.serialize_workflow_node_message(execution) for execution in executions]
def _list_workflow_sources(self, *, app: App, agent_id: str) -> list[dict[str, Any]]:
workflow_app = aliased(App)
@@ -443,6 +445,62 @@ class AgentObservabilityService:
stmt = cls._apply_status_filter(stmt, params.statuses)
return stmt
@classmethod
def _apply_workflow_node_filters(cls, stmt, *, params: AgentLogQueryParams, workflow_app):
if params.start:
stmt = stmt.where(WorkflowNodeExecutionModel.created_at >= params.start)
if params.end:
stmt = stmt.where(WorkflowNodeExecutionModel.created_at < params.end)
if params.keyword:
escaped_keyword = escape_like_pattern(params.keyword)
pattern = f"%{escaped_keyword}%"
stmt = stmt.where(
or_(
WorkflowNodeExecutionModel.inputs.ilike(pattern, escape="\\"),
WorkflowNodeExecutionModel.outputs.ilike(pattern, escape="\\"),
WorkflowNodeExecutionModel.error.ilike(pattern, escape="\\"),
WorkflowNodeExecutionModel.title.ilike(pattern, escape="\\"),
workflow_app.name.ilike(pattern, escape="\\"),
)
)
if params.statuses:
stmt = cls._apply_workflow_node_status_filter(stmt, params.statuses)
return stmt
@staticmethod
def _apply_workflow_node_status_filter(stmt, statuses: tuple[str, ...]):
conditions = []
for status in statuses:
normalized = status.strip().lower()
if normalized in {"success", "normal"}:
conditions.append(WorkflowNodeExecutionModel.status == WorkflowNodeExecutionStatus.SUCCEEDED)
elif normalized in {"failed", "error"}:
conditions.append(
WorkflowNodeExecutionModel.status.in_(
(
WorkflowNodeExecutionStatus.FAILED,
WorkflowNodeExecutionStatus.EXCEPTION,
WorkflowNodeExecutionStatus.STOPPED,
)
)
)
elif normalized == "paused":
conditions.append(
WorkflowNodeExecutionModel.status.in_(
(
WorkflowNodeExecutionStatus.PAUSED,
WorkflowNodeExecutionStatus.PENDING,
WorkflowNodeExecutionStatus.RUNNING,
WorkflowNodeExecutionStatus.RETRY,
)
)
)
else:
raise ValueError(f"Unsupported status: {status}")
if not conditions:
return stmt
return stmt.where(or_(*conditions))
@staticmethod
def _apply_workflow_source_filter(stmt, source_filter: AgentSourceFilter):
if source_filter.app_id:
@@ -505,6 +563,148 @@ class AgentObservabilityService:
"updated_at": to_timestamp(updated_at or conversation.updated_at),
}
@classmethod
def _serialize_workflow_execution_log(
cls,
*,
node_execution_id: str,
title: str,
status: object,
created_by_role: object,
created_by: str,
created_at: datetime,
finished_at: datetime | None,
source: dict[str, Any],
) -> dict[str, Any]:
created_by_role_value = cls._enum_value(created_by_role)
return {
"id": node_execution_id,
"conversation_id": node_execution_id,
"title": title,
"end_user_id": created_by if created_by_role_value == CreatorUserRole.END_USER.value else None,
"message_count": 1,
"user_rate": None,
"operation_rate": None,
"unread": False,
"source": source,
"status": cls._workflow_node_status(status),
"created_at": to_timestamp(created_at),
"updated_at": to_timestamp(finished_at or created_at),
}
@classmethod
def serialize_workflow_node_message(cls, node_execution: WorkflowNodeExecutionModel) -> dict[str, Any]:
inputs = cls._json_mapping(node_execution.inputs)
outputs = cls._json_mapping(node_execution.outputs)
metadata = cls._json_mapping(node_execution.execution_metadata)
agent_log = cls._mapping_value(metadata, "agent_log")
agent_backend = cls._mapping_value(agent_log, "agent_backend")
usage = cls._mapping_value(agent_backend, "usage")
prompt_tokens = cls._int_value(usage.get("prompt_tokens"))
completion_tokens = cls._int_value(usage.get("completion_tokens"))
total_tokens = cls._int_value(usage.get("total_tokens") or metadata.get("total_tokens"))
if not total_tokens:
total_tokens = prompt_tokens + completion_tokens
created_by_role = cls._enum_value(node_execution.created_by_role)
return {
"id": node_execution.id,
"message_id": node_execution.id,
"conversation_id": node_execution.id,
"query": cls._workflow_node_query(inputs, fallback=node_execution.title),
"answer": cls._workflow_node_answer(outputs),
"status": cls._workflow_node_status(node_execution.status),
"error": node_execution.error,
"from_end_user_id": (
node_execution.created_by if created_by_role == CreatorUserRole.END_USER.value else None
),
"from_account_id": (
node_execution.created_by if created_by_role == CreatorUserRole.ACCOUNT.value else None
),
"message_tokens": prompt_tokens,
"answer_tokens": completion_tokens,
"total_tokens": total_tokens,
"total_price": str(usage.get("total_price") or metadata.get("total_price") or Decimal(0)),
"currency": str(usage.get("currency") or metadata.get("currency") or ""),
"latency": float(usage.get("latency") or node_execution.elapsed_time or 0),
"created_at": to_timestamp(node_execution.created_at),
"updated_at": to_timestamp(node_execution.finished_at or node_execution.created_at),
}
@staticmethod
def _json_mapping(value: object) -> Mapping[str, Any]:
if isinstance(value, Mapping):
return value
if not isinstance(value, str) or not value:
return {}
try:
parsed = json.loads(value)
except (TypeError, ValueError):
return {}
return parsed if isinstance(parsed, Mapping) else {}
@staticmethod
def _mapping_value(value: Mapping[str, Any], key: str) -> Mapping[str, Any]:
nested = value.get(key)
return nested if isinstance(nested, Mapping) else {}
@staticmethod
def _enum_value(value: object) -> str:
return str(value.value) if isinstance(value, Enum) else str(value)
@staticmethod
def _int_value(value: object) -> int:
if not isinstance(value, (str, int, float, Decimal)):
return 0
try:
return int(value)
except (TypeError, ValueError):
return 0
@classmethod
def _workflow_node_query(cls, inputs: Mapping[str, Any], *, fallback: str) -> str:
request_data = cls._mapping_value(inputs, "agent_backend_request")
composition = cls._mapping_value(request_data, "composition")
layers = composition.get("layers")
prompts: list[str] = []
if isinstance(layers, list):
for layer_name in ("workflow_node_job_prompt", "workflow_user_prompt"):
for layer in layers:
if not isinstance(layer, Mapping) or layer.get("name") != layer_name:
continue
config = cls._mapping_value(layer, "config")
user_prompt = config.get("user")
if isinstance(user_prompt, str) and user_prompt.strip():
prompts.append(user_prompt.strip())
return "\n\n".join(prompts) or fallback
@staticmethod
def _workflow_node_answer(outputs: Mapping[str, Any]) -> str:
for key in ("output", "text", "answer"):
value = outputs.get(key)
if isinstance(value, str):
return value
return json.dumps(outputs, ensure_ascii=False) if outputs else ""
@classmethod
def _workflow_node_status(cls, status: object) -> str:
value = cls._enum_value(status)
if value in {
WorkflowNodeExecutionStatus.FAILED.value,
WorkflowNodeExecutionStatus.EXCEPTION.value,
WorkflowNodeExecutionStatus.STOPPED.value,
}:
return "failed"
if value in {
WorkflowNodeExecutionStatus.PAUSED.value,
WorkflowNodeExecutionStatus.PENDING.value,
WorkflowNodeExecutionStatus.RUNNING.value,
WorkflowNodeExecutionStatus.RETRY.value,
}:
return "paused"
return "success"
@staticmethod
def _conversation_status(*, paused_count: int, failed_count: int) -> str:
if paused_count:
+28 -17
View File
@@ -6,6 +6,7 @@ from sqlalchemy.exc import IntegrityError
from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload
from constants.model_template import default_app_templates
from core.agent.publish_visibility import workflow_callable_active_snapshot_filter
from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore
from core.app.entities.app_invoke_entities import InvokeFrom
from libs.datetime_utils import naive_utc_now
@@ -225,8 +226,11 @@ class AgentRosterService:
def list_invite_options(
self, *, tenant_id: str, page: int = 1, limit: int = 20, keyword: str | None = None, app_id: str | None = None
) -> dict[str, Any]:
"""List active roster Agents whose published snapshot can be called by Workflow."""
stmt = self._build_roster_agents_stmt(tenant_id=tenant_id, keyword=keyword).where(
Agent.active_config_has_model.is_(True)
Agent.active_config_has_model.is_(True),
workflow_callable_active_snapshot_filter(),
)
total = self._session.scalar(select(func.count()).select_from(stmt.subquery())) or 0
agents = list(self._session.scalars(stmt.offset((page - 1) * limit).limit(limit)).all())
@@ -660,15 +664,18 @@ class AgentRosterService:
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
commit: bool = True,
) -> str:
"""Start a new scoped console conversation for the current Agent App editor.
If this account already has a mapping for the requested draft surface, the previous
conversation is abandoned first: any ACTIVE conversation-owned Agent
runtime sessions for that old conversation are sent through best-effort
backend cleanup and then retired locally even when enqueueing fails. The
other draft surface is left untouched.
conversation is abandoned after the replacement mapping is committed: any ACTIVE
conversation-owned Agent runtime sessions for that old conversation are sent through
best-effort backend cleanup and then retired locally even when enqueueing fails. This
order prevents a failed database commit from retiring the still-current runtime session.
The other draft surface is left untouched.
A user and draft surface own one current mapping. If new-conversation requests overlap,
the last committed rotation becomes current and earlier response IDs cannot be continued.
"""
agent = self._session.scalar(
@@ -691,6 +698,7 @@ class AgentRosterService:
app_id=backing_app_id,
account_id=account_id,
)
previous_conversation: tuple[str, str] | None = None
mapping = self._session.scalar(
select(AgentDebugConversation).where(
AgentDebugConversation.tenant_id == tenant_id,
@@ -714,19 +722,22 @@ class AgentRosterService:
previous_app_id = mapping.app_id
previous_conversation_id = mapping.conversation_id
if previous_conversation_id:
self._cleanup_debug_conversation_runtime_sessions(
tenant_id=tenant_id,
agent_id=agent_id,
account_id=account_id,
draft_type=draft_type,
app_id=previous_app_id or backing_app_id,
conversation_id=previous_conversation_id,
)
previous_conversation = (previous_app_id or backing_app_id, previous_conversation_id)
mapping.app_id = backing_app_id
mapping.conversation_id = conversation_id
self._session.flush()
if commit:
self._session.commit()
self._session.commit()
if previous_conversation:
previous_app_id, previous_conversation_id = previous_conversation
self._cleanup_debug_conversation_runtime_sessions(
tenant_id=tenant_id,
agent_id=agent_id,
account_id=account_id,
draft_type=draft_type,
app_id=previous_app_id,
conversation_id=previous_conversation_id,
)
return conversation_id
def _cleanup_debug_conversation_runtime_sessions(
@@ -739,8 +750,8 @@ class AgentRosterService:
app_id: str,
conversation_id: str,
) -> None:
session_store = AgentAppRuntimeSessionStore()
try:
session_store = AgentAppRuntimeSessionStore()
stored_sessions = session_store.list_active_sessions_for_conversation(
tenant_id=tenant_id,
app_id=app_id,
@@ -8,6 +8,7 @@ from pydantic import ValidationError
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.agent.publish_visibility import workflow_callable_active_snapshot_filter
from core.workflow.nodes.agent_v2.validators import WorkflowAgentNodeValidationError, WorkflowAgentNodeValidator
from models.agent import (
Agent,
@@ -447,6 +448,8 @@ class WorkflowAgentPublishService:
node_id: str,
agent_id: str,
) -> tuple[Agent, str]:
"""Resolve an active roster Agent whose published snapshot is callable."""
agent = session.scalar(
select(Agent)
.where(
@@ -454,11 +457,12 @@ class WorkflowAgentPublishService:
Agent.id == agent_id,
Agent.scope == AgentScope.ROSTER,
Agent.status == AgentStatus.ACTIVE,
workflow_callable_active_snapshot_filter(),
)
.limit(1)
)
if agent is None:
raise ValueError(f"Workflow Agent node {node_id} references an unavailable roster agent.")
raise ValueError(f"Workflow Agent node {node_id} references an unavailable or unpublished roster agent.")
if agent.scope != AgentScope.ROSTER:
raise ValueError(f"Workflow Agent node {node_id} roster_agent binding must reference a roster agent.")
if not agent.active_config_snapshot_id:
+15 -3
View File
@@ -9,6 +9,7 @@ import contexts
from core.app.app_config.easy_ui_based_app.agent.manager import AgentConfigManager
from core.plugin.impl.agent import PluginAgentClient
from core.plugin.impl.exc import PluginDaemonClientSideError
from core.tools.entities.tool_entities import EmojiIconDict
from core.tools.tool_manager import ToolManager
from libs.login import current_user
from models import Account
@@ -105,11 +106,22 @@ class AgentService:
tool_output = tool_outputs.get(tool_name, {})
tool_meta_data = tool_meta.get(tool_name, {})
tool_config = tool_meta_data.get("tool_config", {})
if tool_config.get("tool_provider_type", "") != "dataset-retrieval":
tool_provider_type = tool_config.get("tool_provider_type", "")
tool_provider_id = tool_config.get("tool_provider", "")
if not tool_provider_type:
tool_entity = find_agent_tool(tool_name)
if tool_entity:
tool_provider_type = tool_entity.provider_type
tool_provider_id = tool_provider_id or tool_entity.provider_id
tool_icon: str | EmojiIconDict = ""
if tool_provider_type and tool_provider_type != "dataset-retrieval":
tool_icon = ToolManager.get_tool_icon(
tenant_id=app_model.tenant_id,
provider_type=tool_config.get("tool_provider_type", ""),
provider_id=tool_config.get("tool_provider", ""),
provider_type=tool_provider_type,
provider_id=tool_provider_id,
)
if not tool_icon:
tool_entity = find_agent_tool(tool_name)
+67 -35
View File
@@ -5,6 +5,7 @@ from pydantic import BaseModel, ConfigDict, Field
from configs import dify_config
from constants.dsl_version import CURRENT_APP_DSL_VERSION
from enums.cloud_plan import CloudPlan
from enums.deployment_edition import DeploymentEdition
from enums.hosted_provider import HostedTrialProvider
from services.billing_service import BillingInfo, BillingService
from services.enterprise.enterprise_service import EnterpriseService
@@ -75,8 +76,11 @@ class LicenseStatus(StrEnum):
LOST = "lost"
class LicenseModel(FeatureResponseModel):
class LicenseStatusModel(FeatureResponseModel):
status: LicenseStatus = LicenseStatus.NONE
class LicenseModel(LicenseStatusModel):
expired_at: str = ""
workspaces: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0)
seats: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0)
@@ -157,29 +161,25 @@ class KnowledgeRateLimitModel(FeatureResponseModel):
subscription_plan: str = ""
class PluginManagerModel(FeatureResponseModel):
enabled: bool = False
class SystemFeatureModel(FeatureResponseModel):
"""Non-sensitive bootstrap snapshot exposed before Console or Web authentication."""
deployment_edition: DeploymentEdition
enable_app_deploy: bool = False
sso_enforced_for_signin: bool = False
sso_enforced_for_signin_protocol: str = ""
enable_marketplace: bool = False
max_plugin_package_size: int = dify_config.PLUGIN_MAX_PACKAGE_SIZE
enable_email_code_login: bool = False
enable_email_password_login: bool = True
enable_social_oauth_login: bool = False
enable_collaboration_mode: bool = True
is_allow_register: bool = False
is_allow_create_workspace: bool = False
is_email_setup: bool = False
license: LicenseModel = LicenseModel()
license: LicenseStatusModel = LicenseStatusModel()
branding: BrandingModel = BrandingModel()
webapp_auth: WebAppAuthModel = WebAppAuthModel()
plugin_installation_permission: PluginInstallationPermissionModel = PluginInstallationPermissionModel()
enable_change_email: bool = True
plugin_manager: PluginManagerModel = PluginManagerModel()
enable_creators_platform: bool = False
enable_trial_app: bool = False
enable_explore_banner: bool = False
@@ -251,8 +251,8 @@ class FeatureService:
)
@classmethod
def get_system_features(cls, is_authenticated: bool = False) -> SystemFeatureModel:
system_features = SystemFeatureModel()
def get_system_features(cls) -> SystemFeatureModel:
system_features = SystemFeatureModel(deployment_edition=dify_config.DEPLOYMENT_EDITION)
system_features.rbac_enabled = dify_config.RBAC_ENABLED
cls._fulfill_system_params_from_env(system_features)
@@ -261,8 +261,7 @@ class FeatureService:
system_features.branding.enabled = True
system_features.webapp_auth.enabled = True
system_features.enable_change_email = False
system_features.plugin_manager.enabled = True
cls._fulfill_params_from_enterprise(system_features, is_authenticated)
cls._fulfill_params_from_enterprise(system_features)
if dify_config.MARKETPLACE_ENABLED:
system_features.enable_marketplace = True
@@ -272,6 +271,32 @@ class FeatureService:
return system_features
@classmethod
def is_workspace_creation_allowed(cls) -> bool:
"""Resolve the backend workspace-creation policy, including the Enterprise override."""
is_allowed = dify_config.ALLOW_CREATE_WORKSPACE
if not dify_config.ENTERPRISE_ENABLED:
return is_allowed
enterprise_info = EnterpriseService.get_info()
return bool(enterprise_info.get("IsAllowCreateWorkspace", is_allowed))
@classmethod
def is_plugin_manager_enabled(cls) -> bool:
"""Return whether Enterprise plugin credential policies must be enforced."""
return dify_config.ENTERPRISE_ENABLED
@classmethod
def get_license(cls) -> LicenseModel:
"""Return full license detail. Enterprise-only; requires an authenticated caller.
Non-enterprise deployments have no license, so an unconstrained default
(unlimited seats/workspaces) is returned.
"""
if not dify_config.ENTERPRISE_ENABLED:
return LicenseModel()
return cls._build_license(EnterpriseService.get_info())
@classmethod
def get_app_dsl_version(cls) -> str:
return CURRENT_APP_DSL_VERSION
@@ -283,8 +308,8 @@ class FeatureService:
system_features.enable_social_oauth_login = dify_config.ENABLE_SOCIAL_OAUTH_LOGIN
system_features.enable_collaboration_mode = dify_config.ENABLE_COLLABORATION_MODE
system_features.is_allow_register = dify_config.ALLOW_REGISTER
system_features.is_allow_create_workspace = dify_config.ALLOW_CREATE_WORKSPACE
system_features.is_email_setup = dify_config.MAIL_TYPE is not None and dify_config.MAIL_TYPE != ""
system_features.enable_change_email = dify_config.ENABLE_CHANGE_EMAIL
system_features.enable_trial_app = dify_config.ENABLE_TRIAL_APP
system_features.enable_explore_banner = dify_config.ENABLE_EXPLORE_BANNER
system_features.enable_learn_app = dify_config.ENABLE_LEARN_APP
@@ -410,7 +435,27 @@ class FeatureService:
vector_space.limit = billing_info["vector_space"]["limit"]
@classmethod
def _fulfill_params_from_enterprise(cls, features: SystemFeatureModel, is_authenticated: bool = False):
def _build_license(cls, enterprise_info: dict) -> LicenseModel:
license_model = LicenseModel()
if license_info := enterprise_info.get("License"):
license_model.status = LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE))
license_model.expired_at = license_info.get("expiredAt", "")
if workspaces_info := license_info.get("workspaces"):
license_model.workspaces = LicenseLimitationModel(
enabled=workspaces_info.get("enabled", False),
limit=workspaces_info.get("limit", 0),
size=workspaces_info.get("used", 0),
)
if seats_info := license_info.get("licensedSeats"):
license_model.seats = LicenseLimitationModel(
enabled=seats_info.get("enabled", False),
limit=seats_info.get("limit", 0),
size=seats_info.get("used", 0),
)
return license_model
@classmethod
def _fulfill_params_from_enterprise(cls, features: SystemFeatureModel):
enterprise_info = EnterpriseService.get_info()
if "SSOEnforcedForSignin" in enterprise_info:
@@ -428,9 +473,6 @@ class FeatureService:
if "IsAllowRegister" in enterprise_info:
features.is_allow_register = enterprise_info["IsAllowRegister"]
if "IsAllowCreateWorkspace" in enterprise_info:
features.is_allow_create_workspace = enterprise_info["IsAllowCreateWorkspace"]
if "EnableAppDeploy" in enterprise_info:
features.enable_app_deploy = enterprise_info["EnableAppDeploy"]
@@ -450,24 +492,14 @@ class FeatureService:
)
features.webapp_auth.sso_config.protocol = enterprise_info.get("SSOEnforcedForWebProtocol", "")
# SECURITY NOTE: Only license *status* is exposed to unauthenticated callers
# so the login page can detect an expired/inactive license after force-logout.
# All other license details (expiry date, workspace usage) remain auth-gated.
# This behavior reflects prior internal review of information-leakage risks.
# SECURITY NOTE: system-features is unauthenticated, so it exposes only license
# *status* — enough for the login page to detect an expired/inactive license after
# force-logout. Full license detail (expiry, workspace/seat usage) is served
# separately by get_license() behind an authenticated endpoint.
if license_info := enterprise_info.get("License"):
features.license.status = LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE))
if is_authenticated:
features.license.expired_at = license_info.get("expiredAt", "")
if workspaces_info := license_info.get("workspaces"):
features.license.workspaces.enabled = workspaces_info.get("enabled", False)
features.license.workspaces.limit = workspaces_info.get("limit", 0)
features.license.workspaces.size = workspaces_info.get("used", 0)
if seats_info := license_info.get("licensedSeats"):
features.license.seats.enabled = seats_info.get("enabled", False)
features.license.seats.limit = seats_info.get("limit", 0)
features.license.seats.size = seats_info.get("used", 0)
features.license = LicenseStatusModel(
status=LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE))
)
if "PluginInstallationPermission" in enterprise_info:
plugin_installation_info = enterprise_info["PluginInstallationPermission"]
+5 -7
View File
@@ -99,7 +99,7 @@ def _external_source_operation(
*,
request_headers: tuple[str, ...] = ("x-trace-id",),
) -> KnowledgeFSOperation:
"""Declare a source-connection operation restricted to dataset editors."""
"""Declare an external-source operation restricted to dataset editors."""
return _console_operation(
operation_id,
method,
@@ -347,12 +347,10 @@ KNOWLEDGE_FS_CONSOLE_OPERATIONS: Final[tuple[KnowledgeFSOperation, ...]] = (
rbac_permission=RBACPermission.DATASET_READONLY,
legacy_role="reader",
),
_console_operation(
operation_id="putKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy",
method="PUT",
path="knowledge-spaces/{id}/sources/{sourceId}/sync-policy",
rbac_permission=RBACPermission.DATASET_EDIT,
legacy_role="dataset_editor",
_external_source_operation(
"putKnowledgeSpacesByIdSourcesBySourceIdSyncPolicy",
"PUT",
"knowledge-spaces/{id}/sources/{sourceId}/sync-policy",
),
_dataset_read_operation("getKnowledgeSpacesByIdDocuments", "knowledge-spaces/{id}/documents"),
_dataset_edit_operation("postKnowledgeSpacesByIdDocuments", "POST", "knowledge-spaces/{id}/documents"),
+1 -1
View File
@@ -542,7 +542,7 @@ class WorkflowService:
# Validate credentials before publishing, for credential policy check
from services.feature_service import FeatureService
if FeatureService.get_system_features().plugin_manager.enabled:
if FeatureService.is_plugin_manager_enabled():
self._validate_workflow_credentials(draft_workflow, session=session)
# validate graph structure
+2 -1
View File
@@ -4,6 +4,7 @@ from sqlalchemy.orm import Session
from configs import dify_config
from enums.cloud_plan import CloudPlan
from enums.deployment_edition import DeploymentEdition
from models.account import Tenant, TenantAccountJoin, TenantAccountRole
from services.account_service import TenantService
from services.feature_service import FeatureService
@@ -60,7 +61,7 @@ class WorkspaceService:
"remove_webapp_brand": remove_webapp_brand,
"replace_webapp_logo": replace_webapp_logo,
}
if dify_config.EDITION == "CLOUD":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
tenant_info["next_credit_reset_date"] = feature.next_credit_reset_date
from services.credit_pool_service import CreditPoolBalance, CreditPoolService
+2 -1
View File
@@ -1,10 +1,11 @@
from enum import StrEnum
from configs import dify_config
from enums.deployment_edition import DeploymentEdition
from services.workflow.entities import WorkflowScheduleCFSPlanEntity
# Determine queue names based on edition
if dify_config.EDITION == "CLOUD":
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
# Cloud edition: separate queues for different tiers
_professional_queue = "workflow_professional"
_team_queue = "workflow_team"
@@ -318,6 +318,7 @@ def test_oauth_account_successful_retrieval(
assert response.status_code == 200
assert response.get_json() == {
"id": account.id,
"name": "Test User",
"email": account.email,
"avatar": "avatar_url",
@@ -61,7 +61,7 @@ def add_tenant_for_account(
) -> Tenant:
"""Create an additional tenant and join ``account`` to it (real service calls)."""
with patch("services.account_service.FeatureService") as mock_feature_service:
mock_feature_service.get_system_features.return_value.is_allow_create_workspace = True
mock_feature_service.is_workspace_creation_allowed.return_value = True
tenant = TenantService.create_tenant(name=name, session=session)
TenantService.create_tenant_member(tenant, account, session, role=role)
return tenant
@@ -37,8 +37,9 @@ class TestAccountService:
):
# Setup default mock returns
mock_feature_service.get_system_features.return_value.is_allow_register = True
mock_feature_service.get_system_features.return_value.is_allow_create_workspace = True
mock_feature_service.get_system_features.return_value.license.workspaces.is_available.return_value = True
mock_feature_service.is_workspace_creation_allowed.return_value = True
mock_feature_service.get_license.return_value.workspaces.is_available.return_value = True
mock_feature_service.get_license.return_value.seats.is_available.return_value = True
mock_billing_service.is_email_in_freeze.return_value = False
mock_passport_service.return_value.issue.return_value = "mock_jwt_token"
@@ -400,12 +401,10 @@ class TestAccountService:
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
account = AccountService.create_account_and_tenant(
email=email,
@@ -435,9 +434,7 @@ class TestAccountService:
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = False
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = False
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
with pytest.raises(WorkSpaceNotAllowedCreateError):
@@ -461,12 +458,10 @@ class TestAccountService:
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = False
].get_license.return_value.workspaces.is_available.return_value = False
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
with pytest.raises(WorkspacesLimitExceededError):
@@ -492,7 +487,7 @@ class TestAccountService:
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.seats.is_available.return_value = False
].get_license.return_value.seats.is_available.return_value = False
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
with pytest.raises(SeatsLimitExceededError):
@@ -1273,8 +1268,9 @@ class TestTenantService:
patch("services.account_service.BillingService") as mock_billing_service,
):
# Setup default mock returns
mock_feature_service.get_system_features.return_value.is_allow_create_workspace = True
mock_feature_service.get_system_features.return_value.license.workspaces.is_available.return_value = True
mock_feature_service.is_workspace_creation_allowed.return_value = True
mock_feature_service.get_license.return_value.workspaces.is_available.return_value = True
mock_feature_service.get_license.return_value.seats.is_available.return_value = True
mock_billing_service.is_email_in_freeze.return_value = False
yield {
@@ -1289,9 +1285,7 @@ class TestTenantService:
fake = Faker()
tenant_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1310,9 +1304,7 @@ class TestTenantService:
fake = Faker()
tenant_name = fake.company()
# Setup mocks to disable workspace creation
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = False
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = False
with pytest.raises(NotAllowedCreateWorkspace): # NotAllowedCreateWorkspace exception
TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1326,9 +1318,7 @@ class TestTenantService:
fake = Faker()
custom_tenant_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = False
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = False
# Create tenant with setup flag (should bypass workspace creation restriction)
tenant = TenantService.create_tenant(
@@ -1352,9 +1342,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and account
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1388,9 +1376,7 @@ class TestTenantService:
name2 = fake.name()
password2 = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1428,9 +1414,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and account
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1463,9 +1447,7 @@ class TestTenantService:
tenant1_name = fake.company()
tenant2_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create account and tenants
account = AccountService.create_account(
@@ -1502,9 +1484,7 @@ class TestTenantService:
password = generate_valid_password(fake)
tenant_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create account and tenant
account = AccountService.create_account(
@@ -1540,9 +1520,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create account without setting current tenant
account = AccountService.create_account(
@@ -1568,9 +1546,7 @@ class TestTenantService:
tenant1_name = fake.company()
tenant2_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create account and tenants
account = AccountService.create_account(
@@ -1608,9 +1584,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create account
account = AccountService.create_account(
@@ -1637,9 +1611,7 @@ class TestTenantService:
password = generate_valid_password(fake)
tenant_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create account and tenant
account = AccountService.create_account(
@@ -1668,9 +1640,7 @@ class TestTenantService:
admin_name = fake.name()
admin_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1715,9 +1685,7 @@ class TestTenantService:
tenant_name = fake.company()
invalid_role = fake.word()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1736,9 +1704,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and account
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1773,9 +1739,7 @@ class TestTenantService:
member_name = fake.name()
member_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1816,9 +1780,7 @@ class TestTenantService:
password = generate_valid_password(fake)
invalid_action = "invalid_action_that_doesnt_exist"
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and account
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1851,9 +1813,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and account
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1889,9 +1849,7 @@ class TestTenantService:
member_name = fake.name()
member_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -1979,9 +1937,7 @@ class TestTenantService:
name = fake.name()
password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and account
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2015,9 +1971,7 @@ class TestTenantService:
non_member_name = fake.name()
non_member_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2058,9 +2012,7 @@ class TestTenantService:
member_name = fake.name()
member_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2111,9 +2063,7 @@ class TestTenantService:
member_name = fake.name()
member_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2172,9 +2122,7 @@ class TestTenantService:
member_name = fake.name()
member_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2212,9 +2160,7 @@ class TestTenantService:
tenant2_name = fake.company()
tenant3_name = fake.company()
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create multiple tenants
tenant1 = TenantService.create_tenant(name=tenant1_name, session=db_session_with_containers)
@@ -2239,12 +2185,10 @@ class TestTenantService:
password = generate_valid_password(fake)
workspace_name = fake.company()
# Setup mocks
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
# Create account
account = AccountService.create_account(
@@ -2280,12 +2224,10 @@ class TestTenantService:
existing_tenant_name = fake.company()
new_workspace_name = fake.company()
# Setup mocks
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
# Create account and existing tenant
account = AccountService.create_account(
@@ -2323,9 +2265,7 @@ class TestTenantService:
password = generate_valid_password(fake)
workspace_name = fake.company()
# Setup mocks to disable workspace creation
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = False
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = False
# Create account
account = AccountService.create_account(
@@ -2358,9 +2298,7 @@ class TestTenantService:
normal_name = fake.name()
normal_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2427,9 +2365,7 @@ class TestTenantService:
normal_name = fake.name()
normal_password = generate_valid_password(fake)
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant and accounts
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2478,9 +2414,7 @@ class TestTenantService:
theme = fake.random_element(elements=("dark", "light"))
language = fake.random_element(elements=("zh-CN", "en-US"))
# Setup mocks
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
# Create tenant with custom config
tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers)
@@ -2513,8 +2447,9 @@ class TestRegisterService:
):
# Setup default mock returns
mock_feature_service.get_system_features.return_value.is_allow_register = True
mock_feature_service.get_system_features.return_value.is_allow_create_workspace = True
mock_feature_service.get_system_features.return_value.license.workspaces.is_available.return_value = True
mock_feature_service.is_workspace_creation_allowed.return_value = True
mock_feature_service.get_license.return_value.workspaces.is_available.return_value = True
mock_feature_service.get_license.return_value.seats.is_available.return_value = True
mock_billing_service.is_email_in_freeze.return_value = False
mock_passport_service.return_value.issue.return_value = "mock_jwt_token"
@@ -2626,12 +2561,10 @@ class TestRegisterService:
language = fake.random_element(elements=("en-US", "zh-CN"))
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
# Execute registration
@@ -2670,12 +2603,10 @@ class TestRegisterService:
language = fake.random_element(elements=("en-US", "zh-CN"))
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
# Execute registration with OAuth
@@ -2719,12 +2650,10 @@ class TestRegisterService:
language = fake.random_element(elements=("en-US", "zh-CN"))
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
# Execute registration with pending status
@@ -2765,9 +2694,7 @@ class TestRegisterService:
language = fake.random_element(elements=("en-US", "zh-CN"))
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = False
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = False
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
# with pytest.raises(AccountRegisterError, match="Workspace is not allowed to create."):
@@ -2804,12 +2731,10 @@ class TestRegisterService:
language = fake.random_element(elements=("en-US", "zh-CN"))
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = False
].get_license.return_value.workspaces.is_available.return_value = False
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
# with pytest.raises(AccountRegisterError, match="Workspace is not allowed to create."):
@@ -2883,12 +2808,10 @@ class TestRegisterService:
language = fake.random_element(elements=("en-US", "zh-CN"))
# Setup mocks
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.is_allow_create_workspace = True
mock_external_service_dependencies[
"feature_service"
].get_system_features.return_value.license.workspaces.is_available.return_value = True
].get_license.return_value.workspaces.is_available.return_value = True
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
# Create tenant and inviter account
@@ -5,10 +5,12 @@ from faker import Faker
from sqlalchemy.orm import Session
from enums.cloud_plan import CloudPlan
from enums.deployment_edition import DeploymentEdition
from services.feature_service import (
FeatureModel,
FeatureService,
KnowledgeRateLimitModel,
LicenseModel,
LicenseStatus,
SystemFeatureModel,
)
@@ -273,6 +275,7 @@ class TestFeatureService:
tenant_id = self._create_test_tenant_id()
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = True
mock_config.ENABLE_EMAIL_CODE_LOGIN = True
@@ -282,10 +285,9 @@ class TestFeatureService:
mock_config.ALLOW_REGISTER = False
mock_config.ALLOW_CREATE_WORKSPACE = False
mock_config.MAIL_TYPE = "smtp"
mock_config.PLUGIN_MAX_PACKAGE_SIZE = 100
# Act: Execute the method under test
result = FeatureService.get_system_features(is_authenticated=True)
result = FeatureService.get_system_features()
# Assert: Verify the expected outcomes
assert result is not None
@@ -306,7 +308,6 @@ class TestFeatureService:
assert result.enable_email_password_login is False
assert result.enable_collaboration_mode is True
assert result.is_allow_register is False
assert result.is_allow_create_workspace is False
# Verify branding configuration
assert result.branding.application_title == "Test Enterprise"
@@ -320,12 +321,9 @@ class TestFeatureService:
assert result.webapp_auth.allow_email_password_login is False
assert result.webapp_auth.sso_config.protocol == "oidc"
# Verify license configuration
# Verify license status (public system-features exposes status only; detail lives on get_license)
assert result.license.status.value == "active"
assert result.license.expired_at == "2025-12-31"
assert result.license.workspaces.enabled is True
assert result.license.workspaces.limit == 5
assert result.license.workspaces.size == 2
assert not hasattr(result.license, "expired_at")
# Verify plugin installation permission
assert result.plugin_installation_permission.plugin_installation_scope == "official_only"
@@ -350,6 +348,7 @@ class TestFeatureService:
"""
# Arrange: Setup test data with exact same config as success test
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = True
mock_config.ENABLE_EMAIL_CODE_LOGIN = True
@@ -360,20 +359,18 @@ class TestFeatureService:
mock_config.MAIL_TYPE = "smtp"
mock_config.PLUGIN_MAX_PACKAGE_SIZE = 100
# Act: Execute with is_authenticated=False
result = FeatureService.get_system_features(is_authenticated=False)
# Act: Execute the public (unauthenticated) system-features call
result = FeatureService.get_system_features()
# Assert: Basic structure
assert result is not None
assert isinstance(result, SystemFeatureModel)
# --- 1. Verify only license *status* is exposed to unauthenticated clients ---
# Detailed license info (expiry, workspaces) remains auth-gated.
# Detailed license info (expiry, workspaces) is not part of the public model.
assert result.license.status == LicenseStatus.ACTIVE
assert result.license.expired_at == ""
assert result.license.workspaces.enabled is False
assert result.license.workspaces.limit == 0
assert result.license.workspaces.size == 0
assert not hasattr(result.license, "expired_at")
assert not hasattr(result.license, "workspaces")
# --- 2. Verify Public UI Configuration Availability ---
# Ensure that data required for frontend rendering remains accessible.
@@ -394,6 +391,47 @@ class TestFeatureService:
# Marketplace should be visible
assert result.enable_marketplace is True
def test_get_license_returns_full_detail(
self, db_session_with_containers: Session, mock_external_service_dependencies
):
"""
Test the authenticated license accessor.
This test verifies that:
- get_license() returns the full license payload (status, expiry, workspace usage).
- Detail withheld from the public system-features model is present here.
"""
# Arrange
with patch("services.feature_service.dify_config") as mock_config:
mock_config.ENTERPRISE_ENABLED = True
# Act
result = FeatureService.get_license()
# Assert: full license detail is populated
assert isinstance(result, LicenseModel)
assert result.status == LicenseStatus.ACTIVE
assert result.expired_at == "2025-12-31"
assert result.workspaces.enabled is True
assert result.workspaces.limit == 5
assert result.workspaces.size == 2
mock_external_service_dependencies["enterprise_service"].get_info.assert_called_once()
def test_get_license_non_enterprise_is_unconstrained(
self, db_session_with_containers: Session, mock_external_service_dependencies
):
"""Non-enterprise deployments have no license, so limits are unconstrained."""
with patch("services.feature_service.dify_config") as mock_config:
mock_config.ENTERPRISE_ENABLED = False
result = FeatureService.get_license()
assert isinstance(result, LicenseModel)
assert result.status == LicenseStatus.NONE
assert result.workspaces.is_available() is True
assert result.seats.is_available() is True
mock_external_service_dependencies["enterprise_service"].get_info.assert_not_called()
def test_get_system_features_basic_config(
self, db_session_with_containers: Session, mock_external_service_dependencies
):
@@ -408,12 +446,14 @@ class TestFeatureService:
"""
# Arrange: Setup basic config mock (no enterprise)
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY
mock_config.ENTERPRISE_ENABLED = False
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = True
mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True
mock_config.ENABLE_SOCIAL_OAUTH_LOGIN = False
mock_config.ENABLE_COLLABORATION_MODE = False
mock_config.ENABLE_CHANGE_EMAIL = True
mock_config.ALLOW_REGISTER = True
mock_config.ALLOW_CREATE_WORKSPACE = True
mock_config.MAIL_TYPE = "smtp"
@@ -438,15 +478,11 @@ class TestFeatureService:
assert result.enable_social_oauth_login is False
assert result.enable_collaboration_mode is False
assert result.is_allow_register is True
assert result.is_allow_create_workspace is True
assert result.is_email_setup is True
# Verify marketplace configuration
assert result.enable_marketplace is False
# Verify plugin package size (uses default value from dify_config)
assert result.max_plugin_package_size == 15728640
def test_get_features_billing_disabled(
self, db_session_with_containers: Session, mock_external_service_dependencies
):
@@ -611,11 +647,13 @@ class TestFeatureService:
"""
# Arrange: Setup enterprise disabled mock
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY
mock_config.ENTERPRISE_ENABLED = False
mock_config.MARKETPLACE_ENABLED = True
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True
mock_config.ENABLE_SOCIAL_OAUTH_LOGIN = True
mock_config.ENABLE_CHANGE_EMAIL = True
mock_config.ALLOW_REGISTER = False
mock_config.ALLOW_CREATE_WORKSPACE = False
mock_config.MAIL_TYPE = None
@@ -639,19 +677,15 @@ class TestFeatureService:
assert result.enable_email_password_login is True
assert result.enable_social_oauth_login is True
assert result.is_allow_register is False
assert result.is_allow_create_workspace is False
assert result.is_email_setup is False
# Verify marketplace configuration
assert result.enable_marketplace is True
# Verify plugin package size (uses default value from dify_config)
assert result.max_plugin_package_size == 15728640
# Verify default license status
assert result.license.status == "none"
assert result.license.expired_at == ""
assert result.license.workspaces.enabled is False
assert not hasattr(result.license, "expired_at")
assert not hasattr(result.license, "workspaces")
# Verify no enterprise service calls
mock_external_service_dependencies["enterprise_service"].get_info.assert_not_called()
@@ -840,6 +874,7 @@ class TestFeatureService:
"""
# Arrange: Setup edge case webapp auth mock with proper config
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -878,7 +913,6 @@ class TestFeatureService:
assert result.enable_email_code_login is False
assert result.enable_email_password_login is True
assert result.is_allow_register is False
assert result.is_allow_create_workspace is False
# Verify mock interactions
mock_external_service_dependencies["enterprise_service"].get_info.assert_called_once()
@@ -960,6 +994,7 @@ class TestFeatureService:
# Test case 1: Official only scope
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -983,6 +1018,7 @@ class TestFeatureService:
# Test case 2: All plugins scope
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1003,6 +1039,7 @@ class TestFeatureService:
# Test case 3: Specific partners scope
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1026,6 +1063,7 @@ class TestFeatureService:
# Test case 4: None scope
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1118,24 +1156,19 @@ class TestFeatureService:
}
}
# Act: Execute the method under test
result = FeatureService.get_system_features(is_authenticated=True)
# Act: Execute the authenticated license accessor
result = FeatureService.get_license()
# Assert: Verify the expected outcomes
assert result is not None
assert isinstance(result, SystemFeatureModel)
assert isinstance(result, LicenseModel)
# Verify license status
assert result.license.status == "inactive"
assert result.license.expired_at == ""
assert result.license.workspaces.enabled is False
assert result.license.workspaces.size == 0
assert result.license.workspaces.limit == 0
# Verify enterprise features
assert result.branding.enabled is True
assert result.webapp_auth.enabled is True
assert result.enable_change_email is False
# Verify license status and detail
assert result.status == "inactive"
assert result.expired_at == ""
assert result.workspaces.enabled is False
assert result.workspaces.size == 0
assert result.workspaces.limit == 0
# Verify mock interactions
mock_external_service_dependencies["enterprise_service"].get_info.assert_called_once()
@@ -1154,6 +1187,7 @@ class TestFeatureService:
"""
# Arrange: Setup partial enterprise info mock with proper config
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1200,8 +1234,8 @@ class TestFeatureService:
# Verify default license status
assert result.license.status == "none"
assert result.license.expired_at == ""
assert result.license.workspaces.enabled is False
assert not hasattr(result.license, "expired_at")
assert not hasattr(result.license, "workspaces")
# Verify default plugin installation permission
assert result.plugin_installation_permission.plugin_installation_scope == "all"
@@ -1283,6 +1317,7 @@ class TestFeatureService:
"""
# Arrange: Setup edge case protocols mock with proper config
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1493,24 +1528,19 @@ class TestFeatureService:
}
}
# Act: Execute the method under test
result = FeatureService.get_system_features(is_authenticated=True)
# Act: Execute the authenticated license accessor
result = FeatureService.get_license()
# Assert: Verify the expected outcomes
assert result is not None
assert isinstance(result, SystemFeatureModel)
assert isinstance(result, LicenseModel)
# Verify license status
assert result.license.status == "expired"
assert result.license.expired_at == "2023-12-31"
assert result.license.workspaces.enabled is False
assert result.license.workspaces.size == 0
assert result.license.workspaces.limit == 0
# Verify enterprise features
assert result.branding.enabled is True
assert result.webapp_auth.enabled is True
assert result.enable_change_email is False
# Verify license status and detail
assert result.status == "expired"
assert result.expired_at == "2023-12-31"
assert result.workspaces.enabled is False
assert result.workspaces.size == 0
assert result.workspaces.limit == 0
# Verify mock interactions
mock_external_service_dependencies["enterprise_service"].get_info.assert_called_once()
@@ -1587,6 +1617,7 @@ class TestFeatureService:
"""
# Arrange: Setup edge case branding mock with proper config
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1630,7 +1661,6 @@ class TestFeatureService:
assert result.enable_email_code_login is False
assert result.enable_email_password_login is True
assert result.is_allow_register is False
assert result.is_allow_create_workspace is False
# Verify mock interactions
mock_external_service_dependencies["enterprise_service"].get_info.assert_called_once()
@@ -1776,6 +1806,7 @@ class TestFeatureService:
"""
# Arrange: Setup lost license mock with proper config
with patch("services.feature_service.dify_config") as mock_config:
mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
mock_config.ENTERPRISE_ENABLED = True
mock_config.MARKETPLACE_ENABLED = False
mock_config.ENABLE_EMAIL_CODE_LOGIN = False
@@ -1787,7 +1818,7 @@ class TestFeatureService:
mock_config.PLUGIN_MAX_PACKAGE_SIZE = 100
mock_external_service_dependencies["enterprise_service"].get_info.return_value = {
"license": {"status": "lost", "expired_at": None, "plan": None}
"License": {"status": "lost"}
}
# Act: Execute the method under test
@@ -1808,7 +1839,7 @@ class TestFeatureService:
assert result.enable_email_code_login is False
assert result.enable_email_password_login is True
assert result.is_allow_register is False
assert result.is_allow_create_workspace is False
assert result.license.status == LicenseStatus.LOST
# Verify mock interactions
mock_external_service_dependencies["enterprise_service"].get_info.assert_called_once()
@@ -41,7 +41,7 @@ class TestWebhookService:
# Mock feature service
mock_feature_service.get_system_features.return_value.is_allow_register = True
mock_feature_service.get_system_features.return_value.is_allow_create_workspace = True
mock_feature_service.is_workspace_creation_allowed.return_value = True
yield {
"async_service": mock_async_service,
@@ -6,6 +6,7 @@ import pytest
from faker import Faker
from sqlalchemy.orm import Session
from enums.deployment_edition import DeploymentEdition
from models import Account, Tenant, TenantAccountJoin, TenantAccountRole
from services.credit_pool_service import CreditPoolBalance
from services.workspace_service import WorkspaceService
@@ -611,7 +612,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "SELF_HOSTED"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY
mock_external_service_dependencies["feature_service"].get_features.return_value.can_replace_logo = False
mock_external_service_dependencies["tenant_service"].has_roles.return_value = False
@@ -632,7 +633,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -657,7 +658,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -685,7 +686,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -713,7 +714,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -749,7 +750,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -779,7 +780,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -808,7 +809,7 @@ class TestWorkspaceService:
db_session_with_containers, mock_external_service_dependencies
)
mock_external_service_dependencies["dify_config"].EDITION = "CLOUD"
mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
feature = mock_external_service_dependencies["feature_service"].get_features.return_value
feature.can_replace_logo = False
feature.next_credit_reset_date = "2025-02-01"
@@ -112,8 +112,8 @@ def test_publish_blocks_start_and_trigger_coexistence(
monkeypatch.setattr(
feature_service_module.FeatureService,
"get_system_features",
classmethod(lambda _cls: SimpleNamespace(plugin_manager=SimpleNamespace(enabled=False))),
"is_plugin_manager_enabled",
classmethod(lambda _cls: False),
)
monkeypatch.setattr("services.workflow_service.dify_config", SimpleNamespace(BILLING_ENABLED=False))
@@ -110,6 +110,26 @@ def test_generate_specs_writes_unique_operation_ids(tmp_path):
assert len(operation_ids) == len(set(operation_ids))
def test_system_features_specs_exclude_backend_only_fields(tmp_path):
module = _load_generate_swagger_specs_module()
written_paths = module.generate_specs(tmp_path)
excluded_fields = {
"is_allow_create_workspace",
"max_plugin_package_size",
"plugin_manager",
}
for spec_name in ("console-openapi.json", "web-openapi.json"):
spec_path = next(path for path in written_paths if path.name == spec_name)
payload = json.loads(spec_path.read_text(encoding="utf-8"))
schemas = payload["components"]["schemas"]
system_features_schema = schemas["SystemFeatureModel"]
assert excluded_fields.isdisjoint(system_features_schema["properties"])
assert "PluginManagerModel" not in schemas
def test_generate_specs_writes_get_operations_without_request_bodies(tmp_path):
module = _load_generate_swagger_specs_module()
@@ -69,6 +69,20 @@ def _version_response(version_id: str = "version-1") -> dict:
}
def test_query_values_accepts_repeated_and_indexed_arrays() -> None:
app = Flask(__name__)
with app.test_request_context("/?sources=webapp:app-1&sources%5B1%5D=workflow:app-2&sources%5B0%5D=workflow:app-1"):
assert roster_controller._query_values("sources", "source") == [
"webapp:app-1",
"workflow:app-1",
"workflow:app-2",
]
with app.test_request_context("/?source%5B0%5D=workflow:app-3"):
assert roster_controller._query_values("sources", "source") == ["workflow:app-3"]
def _workflow_composer_response(**overrides) -> dict:
response = {
"variant": "workflow",
@@ -1530,8 +1544,35 @@ def test_drain_streaming_generate_response_raises_when_stream_ends_early() -> No
completion_controller._drain_streaming_generate_response(response)
def test_agent_chat_helper_forces_agent_streaming_and_external_trace(
app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str
@pytest.mark.parametrize(
("payload_extra", "expected_draft_type", "expected_start_new"),
[
({}, AgentConfigDraftType.DRAFT, True),
({"conversation_id": None}, AgentConfigDraftType.DRAFT, True),
({"conversation_id": ""}, AgentConfigDraftType.DRAFT, True),
(
{"conversation_id": "00000000-0000-0000-0000-000000000001"},
AgentConfigDraftType.DRAFT,
False,
),
({"draft_type": "debug_build"}, AgentConfigDraftType.DEBUG_BUILD, False),
(
{
"draft_type": "debug_build",
"conversation_id": "00000000-0000-0000-0000-000000000001",
},
AgentConfigDraftType.DEBUG_BUILD,
False,
),
],
)
def test_agent_chat_helper_resolves_scoped_conversation_and_forces_streaming(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
account_id: str,
payload_extra: dict[str, str | None],
expected_draft_type: AgentConfigDraftType,
expected_start_new: bool,
) -> None:
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent")
current_user = SimpleNamespace(id=account_id)
@@ -1543,7 +1584,7 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace(
def resolve_debug_conversation(**kwargs: object) -> str:
captured["resolve_debug_conversation"] = kwargs
return "debug-conversation-1"
return "00000000-0000-0000-0000-000000000001"
monkeypatch.setattr(completion_controller.AppGenerateService, "generate", generate)
monkeypatch.setattr(
@@ -1555,7 +1596,8 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace(
completion_controller.helper, "compact_generate_response", lambda response: {"response": response}
)
with app.test_request_context(
json={"inputs": {}, "query": "hello", "response_mode": "streaming"}, headers={"X-Trace-Id": "trace-1"}
json={"inputs": {}, "query": "hello", "response_mode": "streaming", **payload_extra},
headers={"X-Trace-Id": "trace-1"},
):
result = completion_controller._create_chat_message(
current_user=current_user, app_model=app_model, session=Mock()
@@ -1566,10 +1608,12 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace(
assert captured["streaming"] is True
args = cast(dict[str, object], captured["args"])
assert args["response_mode"] == "streaming"
assert args["conversation_id"] == "debug-conversation-1"
assert args["conversation_id"] == "00000000-0000-0000-0000-000000000001"
assert args["auto_generate_name"] is False
assert args["external_trace_id"] == "trace-1"
assert cast(dict[str, object], captured["resolve_debug_conversation"])["draft_type"] == AgentConfigDraftType.DRAFT
resolve_call = cast(dict[str, object], captured["resolve_debug_conversation"])
assert resolve_call["draft_type"] == expected_draft_type
assert resolve_call["start_new"] is expected_start_new
def test_agent_chat_helper_ignores_private_exit_intent_payload_key(
@@ -1617,14 +1661,28 @@ def test_agent_chat_helper_ignores_private_exit_intent_payload_key(
assert completion_controller.AGENT_RUNTIME_EXIT_INTENT_ARG not in args
def test_agent_chat_helper_rejects_foreign_debug_conversation(
app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str
@pytest.mark.parametrize(
("payload_extra", "expected_draft_type"),
[
({}, AgentConfigDraftType.DRAFT),
({"draft_type": "debug_build"}, AgentConfigDraftType.DEBUG_BUILD),
],
)
def test_agent_chat_helper_rejects_foreign_debug_conversation_before_generation(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
account_id: str,
payload_extra: dict[str, str],
expected_draft_type: AgentConfigDraftType,
) -> None:
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent")
generate = MagicMock()
resolve_debug_conversation = MagicMock(return_value="owned-conversation")
monkeypatch.setattr(completion_controller.AppGenerateService, "generate", generate)
monkeypatch.setattr(
completion_controller,
"_resolve_current_user_agent_debug_conversation_id",
lambda **kwargs: "owned-conversation",
resolve_debug_conversation,
)
with app.test_request_context(
json={
@@ -1632,6 +1690,7 @@ def test_agent_chat_helper_rejects_foreign_debug_conversation(
"query": "hello",
"response_mode": "streaming",
"conversation_id": "00000000-0000-0000-0000-000000000001",
**payload_extra,
}
):
with pytest.raises(NotFound):
@@ -1643,6 +1702,12 @@ def test_agent_chat_helper_rejects_foreign_debug_conversation(
session=Mock(),
)
resolve_debug_conversation.assert_called_once()
resolve_call = resolve_debug_conversation.call_args.kwargs
assert resolve_call["draft_type"] == expected_draft_type
assert resolve_call["start_new"] is False
generate.assert_not_called()
def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app(
monkeypatch: pytest.MonkeyPatch,
@@ -1657,6 +1722,10 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app
calls.append({"get_or_create": kwargs})
return f"debug-{kwargs['agent_id']}"
def refresh_agent_app_debug_conversation_id(self, **kwargs: object) -> str:
calls.append({"refresh": kwargs})
return f"new-{kwargs['agent_id']}"
def get_app_backing_agent(self, **kwargs: object) -> object:
calls.append({"get_app_backing_agent": kwargs})
return SimpleNamespace(id="backing-agent")
@@ -1669,6 +1738,7 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app
app_model=SimpleNamespace(id="app-1"),
agent_id="agent-1",
draft_type=AgentConfigDraftType.DRAFT,
start_new=True,
)
fallback_id = completion_controller._resolve_current_user_agent_debug_conversation_id(
session="session-1", # type: ignore[arg-type]
@@ -1678,10 +1748,20 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app
agent_id=None,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
assert explicit_id == "debug-agent-1"
fallback_preview_id = completion_controller._resolve_current_user_agent_debug_conversation_id(
session="session-1", # type: ignore[arg-type]
current_tenant_id="tenant-1",
current_user=SimpleNamespace(id="account-1"),
app_model=SimpleNamespace(id="app-1"),
agent_id=None,
draft_type=AgentConfigDraftType.DRAFT,
start_new=True,
)
assert explicit_id == "new-agent-1"
assert fallback_id == "debug-backing-agent"
assert fallback_preview_id == "new-backing-agent"
assert calls[1] == {
"get_or_create": {
"refresh": {
"tenant_id": "tenant-1",
"agent_id": "agent-1",
"account_id": "account-1",
@@ -1697,6 +1777,15 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app
"draft_type": AgentConfigDraftType.DEBUG_BUILD,
}
}
assert calls[6] == {"get_app_backing_agent": {"tenant_id": "tenant-1", "app_id": "app-1"}}
assert calls[7] == {
"refresh": {
"tenant_id": "tenant-1",
"agent_id": "backing-agent",
"account_id": "account-1",
"draft_type": AgentConfigDraftType.DRAFT,
}
}
@pytest.mark.parametrize(
@@ -13,6 +13,7 @@ from sqlalchemy import Engine, event
from sqlalchemy.orm import Session
from controllers.console.app import app_import as app_import_module
from enums.deployment_edition import DeploymentEdition
from models.account import Account
from models.base import TypeBase
from models.engine import db
@@ -47,7 +48,10 @@ class _Result:
def _install_features(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
features = SystemFeatureModel(webapp_auth=WebAppAuthModel(enabled=enabled))
features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
webapp_auth=WebAppAuthModel(enabled=enabled),
)
monkeypatch.setattr(app_import_module.FeatureService, "get_system_features", lambda: features)
@@ -11,6 +11,7 @@ from controllers.console.auth.email_register import (
EmailRegisterResetApi,
EmailRegisterSendEmailApi,
)
from enums.deployment_edition import DeploymentEdition
from services.feature_service import SystemFeatureModel
@@ -34,7 +35,11 @@ class TestEmailRegisterSendEmailApi:
mock_account = MagicMock()
mock_get_account.return_value = mock_account
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.auth.email_register.dify_config.BILLING_ENABLED", True),
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
@@ -75,7 +80,11 @@ class TestEmailRegisterCheckApi:
mock_get_data.return_value = {"email": "User@Example.com", "code": "4321"}
mock_generate_token.return_value = (None, "new-token")
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -123,7 +132,11 @@ class TestEmailRegisterResetApi:
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -171,7 +184,11 @@ class TestEmailRegisterResetApi:
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -224,7 +241,11 @@ class TestEmailRegisterResetApi:
mock_login.return_value = token_pair
mock_get_account.return_value = None
feature_flags = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
feature_flags = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags),
@@ -442,10 +442,10 @@ class TestEmailCodeLoginApi:
@patch("controllers.console.auth.login.AccountService.revoke_email_code_login_token")
@patch("controllers.console.auth.login.AccountService.get_user_through_email")
@patch("controllers.console.auth.login.TenantService.get_join_tenants")
@patch("controllers.console.auth.login.FeatureService.get_system_features")
@patch("controllers.console.auth.login.FeatureService.is_workspace_creation_allowed")
def test_email_code_login_creates_workspace_for_user_without_tenant(
self,
mock_get_features,
mock_is_workspace_creation_allowed,
mock_get_tenants,
mock_get_user,
mock_revoke_token,
@@ -465,10 +465,7 @@ class TestEmailCodeLoginApi:
mock_get_data.return_value = {"email": "test@example.com", "code": "123456"}
mock_get_user.return_value = mock_account
mock_get_tenants.return_value = []
mock_features = MagicMock()
mock_features.is_allow_create_workspace = True
mock_features.license.workspaces.is_available.return_value = True
mock_get_features.return_value = mock_features
mock_is_workspace_creation_allowed.return_value = True
# Act & Assert - Should not raise WorkspacesLimitExceeded
with app.test_request_context(
@@ -485,10 +482,12 @@ class TestEmailCodeLoginApi:
@patch("controllers.console.auth.login.AccountService.revoke_email_code_login_token")
@patch("controllers.console.auth.login.AccountService.get_user_through_email")
@patch("controllers.console.auth.login.TenantService.get_join_tenants")
@patch("controllers.console.auth.login.FeatureService.get_system_features")
@patch("controllers.console.auth.login.FeatureService.get_license")
@patch("controllers.console.auth.login.FeatureService.is_workspace_creation_allowed")
def test_email_code_login_workspace_limit_exceeded(
self,
mock_get_features,
mock_is_workspace_creation_allowed,
mock_get_license,
mock_get_tenants,
mock_get_user,
mock_revoke_token,
@@ -507,9 +506,8 @@ class TestEmailCodeLoginApi:
mock_get_data.return_value = {"email": "test@example.com", "code": "123456"}
mock_get_user.return_value = mock_account
mock_get_tenants.return_value = []
mock_features = MagicMock()
mock_features.license.workspaces.is_available.return_value = False
mock_get_features.return_value = mock_features
mock_get_license.return_value.workspaces.is_available.return_value = False
mock_is_workspace_creation_allowed.return_value = True
# Act & Assert
with app.test_request_context(
@@ -526,10 +524,10 @@ class TestEmailCodeLoginApi:
@patch("controllers.console.auth.login.AccountService.revoke_email_code_login_token")
@patch("controllers.console.auth.login.AccountService.get_user_through_email")
@patch("controllers.console.auth.login.TenantService.get_join_tenants")
@patch("controllers.console.auth.login.FeatureService.get_system_features")
@patch("controllers.console.auth.login.FeatureService.is_workspace_creation_allowed")
def test_email_code_login_workspace_creation_not_allowed(
self,
mock_get_features,
mock_is_workspace_creation_allowed,
mock_get_tenants,
mock_get_user,
mock_revoke_token,
@@ -548,9 +546,7 @@ class TestEmailCodeLoginApi:
mock_get_data.return_value = {"email": "test@example.com", "code": "123456"}
mock_get_user.return_value = mock_account
mock_get_tenants.return_value = []
mock_features = MagicMock()
mock_features.is_allow_create_workspace = False
mock_get_features.return_value = mock_features
mock_is_workspace_creation_allowed.return_value = False
# Act & Assert
with app.test_request_context(
@@ -13,6 +13,7 @@ from controllers.console.auth.forgot_password import (
ForgotPasswordResetApi,
ForgotPasswordSendEmailApi,
)
from enums.deployment_edition import DeploymentEdition
from models.account import Account
from models.engine import db
from services.feature_service import SystemFeatureModel
@@ -46,8 +47,15 @@ class TestForgotPasswordSendEmailApi:
mock_get_account.return_value = mock_account
mock_send_email.return_value = "token-123"
wraps_features = SystemFeatureModel(enable_email_password_login=True, is_allow_register=True)
controller_features = SystemFeatureModel(is_allow_register=True)
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
is_allow_register=True,
)
controller_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
is_allow_register=True,
)
with (
patch(
"controllers.console.auth.forgot_password.FeatureService.get_system_features",
@@ -95,7 +103,10 @@ class TestForgotPasswordCheckApi:
mock_get_data.return_value = {"email": "Admin@Example.com", "code": "4321"}
mock_generate_token.return_value = (None, "new-token")
wraps_features = SystemFeatureModel(enable_email_password_login=True)
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
@@ -138,7 +149,10 @@ class TestForgotPasswordResetApi:
db.session.commit()
mock_get_account.return_value = account
wraps_features = SystemFeatureModel(enable_email_password_login=True)
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
)
with (
patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"),
patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features),
@@ -343,10 +343,12 @@ class TestLoginApi:
@patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback")
@patch("controllers.console.auth.login.AccountService.authenticate")
@patch("controllers.console.auth.login.TenantService.get_join_tenants")
@patch("controllers.console.auth.login.FeatureService.get_system_features")
@patch("controllers.console.auth.login.FeatureService.get_license")
@patch("controllers.console.auth.login.FeatureService.is_workspace_creation_allowed")
def test_login_fails_when_no_workspace_and_limit_exceeded(
self,
mock_get_features: MagicMock,
mock_is_workspace_creation_allowed: MagicMock,
mock_get_license: MagicMock,
mock_get_tenants: MagicMock,
mock_authenticate: MagicMock,
mock_get_invitation: MagicMock,
@@ -368,10 +370,8 @@ class TestLoginApi:
mock_authenticate.return_value = mock_account
mock_get_tenants.return_value = [] # No tenants
mock_features = MagicMock()
mock_features.is_allow_create_workspace = True
mock_features.license.workspaces.is_available.return_value = False
mock_get_features.return_value = mock_features
mock_is_workspace_creation_allowed.return_value = True
mock_get_license.return_value.workspaces.is_available.return_value = False
# Act & Assert
with app.test_request_context(
@@ -646,7 +646,7 @@ class TestAccountGeneration:
):
mock_get_account.return_value = mock_account
mock_tenant_service.get_join_tenants.return_value = []
mock_feature_service.get_system_features.return_value.is_allow_create_workspace = True
mock_feature_service.is_workspace_creation_allowed.return_value = True
with app.test_request_context(headers={"Accept-Language": "en-US,en;q=0.9"}):
result, oauth_new_user = _generate_account("github", user_info)
@@ -3,7 +3,7 @@ from __future__ import annotations
from inspect import unwrap
from unittest.mock import patch
from controllers.console.auth.oauth_server import OAuthServerUserAuthorizeApi
from controllers.console.auth.oauth_server import OAuthServerUserAccountApi, OAuthServerUserAuthorizeApi
from models import Account
from models.account import AccountStatus, TenantAccountRole
from models.model import OAuthProviderApp
@@ -45,3 +45,13 @@ def test_oauth_authorize_uses_injected_current_user() -> None:
sign_oauth_authorization_code.assert_called_once_with("client-1", "account-1")
assert response == {"code": "authorization-code"}
def test_oauth_account_returns_stable_account_id() -> None:
api = OAuthServerUserAccountApi()
method = unwrap(api.post)
account = _make_account()
response = method(api, _make_oauth_provider_app(), account)
assert response["id"] == "account-1"
@@ -23,6 +23,7 @@ from controllers.console.auth.forgot_password import (
ForgotPasswordSendEmailApi,
)
from controllers.console.error import AccountNotFound, EmailSendIpLimitError
from enums.deployment_edition import DeploymentEdition
from models.account import Account, Tenant, TenantAccountJoin
from services.feature_service import SystemFeatureModel
@@ -48,7 +49,10 @@ def enable_password_login_wrappers(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("controllers.console.wraps.dify_config.EDITION", "CLOUD")
monkeypatch.setattr(
"controllers.console.wraps.FeatureService.get_system_features",
lambda: SystemFeatureModel(enable_email_password_login=True),
lambda: SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
),
)
@@ -2,14 +2,15 @@ from inspect import unwrap
from pytest_mock import MockerFixture
from models import Account
from services.feature_service import FeatureModel, LimitationModel, SystemFeatureModel
def make_account() -> Account:
account = Account(name="Alice", email="alice@example.com")
account.id = "account-1"
return account
from enums.deployment_edition import DeploymentEdition
from services.feature_service import (
FeatureModel,
LicenseLimitationModel,
LicenseModel,
LicenseStatus,
LimitationModel,
SystemFeatureModel,
)
class TestFeatureApi:
@@ -82,19 +83,16 @@ class TestAppDslVersionApi:
class TestSystemFeatureApi:
def test_get_system_features_authenticated(self, mocker: MockerFixture):
"""
current_user.is_authenticated == True
"""
def test_get_system_features_public(self, mocker: MockerFixture):
"""The public endpoint returns system features without any authentication input."""
from controllers.console.feature import SystemFeatureApi
account = make_account()
current_account = mocker.patch(
"controllers.console.feature.current_account_with_tenant_optional",
return_value=(account, "tenant-123"),
system_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
is_allow_register=True,
enable_learn_app=True,
)
system_features = SystemFeatureModel(is_allow_register=True, enable_learn_app=True)
get_system_features = mocker.patch(
"controllers.console.feature.FeatureService.get_system_features",
return_value=system_features,
@@ -104,30 +102,30 @@ class TestSystemFeatureApi:
result = api.get()
assert result == system_features.model_dump()
assert result["is_allow_register"] is True
assert result["enable_learn_app"] is True
current_account.assert_called_once_with()
get_system_features.assert_called_once_with(is_authenticated=True)
assert result["license"] == {"status": LicenseStatus.NONE}
get_system_features.assert_called_once_with()
def test_get_system_features_unauthenticated(self, mocker: MockerFixture):
"""
current_user.is_authenticated raises Unauthorized
"""
from controllers.console.feature import SystemFeatureApi
class TestSystemFeatureLicenseApi:
def test_get_license_success(self, mocker: MockerFixture):
from controllers.console.feature import SystemFeatureLicenseApi
current_account = mocker.patch(
"controllers.console.feature.current_account_with_tenant_optional",
return_value=(None, None),
license_model = LicenseModel(
status=LicenseStatus.ACTIVE,
expired_at="2025-12-31",
seats=LicenseLimitationModel(enabled=True, limit=5, size=2),
)
system_features = SystemFeatureModel(is_allow_register=False)
get_system_features = mocker.patch(
"controllers.console.feature.FeatureService.get_system_features",
return_value=system_features,
get_license = mocker.patch(
"controllers.console.feature.FeatureService.get_license",
return_value=license_model,
)
api = SystemFeatureApi()
result = api.get()
api = SystemFeatureLicenseApi()
raw_get = unwrap(SystemFeatureLicenseApi.get)
result = raw_get(api)
assert result == system_features.model_dump()
current_account.assert_called_once_with()
get_system_features.assert_called_once_with(is_authenticated=False)
assert result == license_model.model_dump()
assert result["seats"] == {"enabled": True, "limit": 5, "size": 2}
get_license.assert_called_once_with()
@@ -323,8 +323,8 @@ class TestMemberInviteEmailApi:
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = False
license_info = MagicMock()
license_info.seats.is_available.return_value = False
payload = {
"emails": ["a@test.com", "b@test.com"],
@@ -336,9 +336,9 @@ class TestMemberInviteEmailApi:
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 2)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
"controllers.console.workspace.members.FeatureService.get_license",
return_value=license_info,
) as mock_get_license,
patch("controllers.console.workspace.members.RegisterService.invite_new_member") as mock_invite,
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
@@ -346,8 +346,8 @@ class TestMemberInviteEmailApi:
with pytest.raises(SeatsLimitExceeded):
method(api, user)
mock_get_system_features.assert_called_once_with(is_authenticated=True)
system_features.license.seats.is_available.assert_called_once_with(2)
mock_get_license.assert_called_once_with()
license_info.seats.is_available.assert_called_once_with(2)
mock_invite.assert_not_called()
def test_invite_existing_accounts_do_not_consume_seats(self, app: Flask):
@@ -359,8 +359,8 @@ class TestMemberInviteEmailApi:
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = False
license_info = MagicMock()
license_info.seats.is_available.return_value = False
payload = {
"emails": ["a@test.com", "b@test.com"],
@@ -372,9 +372,9 @@ class TestMemberInviteEmailApi:
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 0)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
"controllers.console.workspace.members.FeatureService.get_license",
return_value=license_info,
) as mock_get_license,
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"
) as mock_invite,
@@ -386,8 +386,8 @@ class TestMemberInviteEmailApi:
assert status == 201
assert len(result["invitation_results"]) == 2
mock_get_system_features.assert_not_called()
system_features.license.seats.is_available.assert_not_called()
mock_get_license.assert_not_called()
license_info.seats.is_available.assert_not_called()
assert mock_invite.call_count == 2
def test_invite_mixed_accounts_with_available_seats(self, app: Flask):
@@ -399,8 +399,8 @@ class TestMemberInviteEmailApi:
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = True
license_info = MagicMock()
license_info.seats.is_available.return_value = True
payload = {
"emails": ["a@test.com", "b@test.com"],
@@ -412,9 +412,9 @@ class TestMemberInviteEmailApi:
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 1)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
"controllers.console.workspace.members.FeatureService.get_license",
return_value=license_info,
) as mock_get_license,
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"
) as mock_invite,
@@ -426,8 +426,8 @@ class TestMemberInviteEmailApi:
assert status == 201
assert len(result["invitation_results"]) == 2
mock_get_system_features.assert_called_once_with(is_authenticated=True)
system_features.license.seats.is_available.assert_called_once_with(1)
mock_get_license.assert_called_once_with()
license_info.seats.is_available.assert_called_once_with(1)
assert mock_invite.call_count == 2
def test_invite_skips_seats_limit_when_enterprise_disabled(self, app: Flask):
@@ -439,8 +439,8 @@ class TestMemberInviteEmailApi:
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = False
license_info = MagicMock()
license_info.seats.is_available.return_value = False
payload = {
"emails": ["a@test.com"],
@@ -452,9 +452,9 @@ class TestMemberInviteEmailApi:
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
"controllers.console.workspace.members.FeatureService.get_license",
return_value=license_info,
) as mock_get_license,
patch("controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"),
patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False),
@@ -464,8 +464,8 @@ class TestMemberInviteEmailApi:
assert status == 201
assert result["invitation_results"][0]["status"] == "success"
mock_get_system_features.assert_not_called()
system_features.license.seats.is_available.assert_not_called()
mock_get_license.assert_not_called()
license_info.seats.is_available.assert_not_called()
def test_invite_seats_error_is_reported_as_failed_result(self, app: Flask):
api = MemberInviteEmailApi()
@@ -476,8 +476,8 @@ class TestMemberInviteEmailApi:
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = True
license_info = MagicMock()
license_info.seats.is_available.return_value = True
payload = {
"emails": ["a@test.com"],
@@ -489,8 +489,8 @@ class TestMemberInviteEmailApi:
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
"controllers.console.workspace.members.FeatureService.get_license",
return_value=license_info,
),
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member",
@@ -1343,6 +1343,34 @@ class TestPluginFetchAutoUpgradeApi:
assert result["category"] == TenantPluginAutoUpgradeCategory.TOOL
assert result["auto_upgrade"]["upgrade_time_of_day"] == 1
def test_returns_disabled_settings_when_strategy_is_missing(self, app: Flask):
api = PluginFetchAutoUpgradeApi()
method = unwrap(api.get)
with (
app.test_request_context(f"/?category={TenantPluginAutoUpgradeCategory.MODEL.value}"),
patch(
"controllers.console.workspace.plugin.PluginAutoUpgradeService.get_strategy",
return_value=None,
),
patch(
"controllers.console.workspace.plugin.PluginAutoUpgradeService.default_upgrade_time_of_day",
return_value=78300,
),
):
result = method(api, "t1")
assert result == {
"category": TenantPluginAutoUpgradeCategory.MODEL,
"auto_upgrade": {
"strategy_setting": TenantPluginAutoUpgradeStrategySetting.DISABLED,
"upgrade_time_of_day": 78300,
"upgrade_mode": TenantPluginAutoUpgradeMode.EXCLUDE,
"exclude_plugins": [],
"include_plugins": [],
},
}
class TestPluginAutoUpgradeExcludePluginApi:
def test_success(self, app: Flask):
@@ -11,26 +11,27 @@ from controllers.openapi.auth.data import (
RequestContext,
current_edition,
)
from enums.deployment_edition import DeploymentEdition
from libs.oauth_bearer import Scope, TokenType
def test_current_edition_saas():
with patch("controllers.openapi.auth.data.dify_config") as cfg:
cfg.EDITION = "CLOUD"
cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD
cfg.ENTERPRISE_ENABLED = True
assert current_edition() == Edition.SAAS
def test_current_edition_ee():
with patch("controllers.openapi.auth.data.dify_config") as cfg:
cfg.EDITION = "SELF_HOSTED"
cfg.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE
cfg.ENTERPRISE_ENABLED = True
assert current_edition() == Edition.EE
def test_current_edition_ce():
with patch("controllers.openapi.auth.data.dify_config") as cfg:
cfg.EDITION = "SELF_HOSTED"
cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY
cfg.ENTERPRISE_ENABLED = False
assert current_edition() == Edition.CE
@@ -7,27 +7,26 @@ from unittest.mock import MagicMock, patch
from flask import Flask
from controllers.web.feature import SystemFeatureApi
from enums.deployment_edition import DeploymentEdition
from services.feature_service import SystemFeatureModel
class TestSystemFeatureApi:
@patch("controllers.web.feature.FeatureService.get_system_features")
def test_returns_system_features(self, mock_features: MagicMock, app: Flask) -> None:
mock_model = MagicMock()
mock_model.model_dump.return_value = {"sso_enforced_for_signin": False, "webapp_auth": {"enabled": False}}
mock_features.return_value = mock_model
system_features = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY)
mock_features.return_value = system_features
with app.test_request_context("/system-features"):
result = SystemFeatureApi().get()
assert result == {"sso_enforced_for_signin": False, "webapp_auth": {"enabled": False}}
assert result == system_features.model_dump()
mock_features.assert_called_once()
@patch("controllers.web.feature.FeatureService.get_system_features")
def test_unauthenticated_access(self, mock_features: MagicMock, app: Flask) -> None:
"""SystemFeatureApi is unauthenticated by design — no WebApiResource decorator."""
mock_model = MagicMock()
mock_model.model_dump.return_value = {}
mock_features.return_value = mock_model
mock_features.return_value = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY)
# Verify it's a bare Resource, not WebApiResource
from flask_restx import Resource
@@ -14,6 +14,7 @@ from controllers.web.forgot_password import (
ForgotPasswordResetApi,
ForgotPasswordSendEmailApi,
)
from enums.deployment_edition import DeploymentEdition
from models.account import Account
from models.engine import db
from services.feature_service import SystemFeatureModel
@@ -32,7 +33,10 @@ def database_app() -> Iterator[Flask]:
@pytest.fixture(autouse=True)
def _patch_wraps():
wraps_features = SystemFeatureModel(enable_email_password_login=True)
wraps_features = SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
enable_email_password_login=True,
)
with (
patch("controllers.console.wraps.db") as mock_db,
patch("controllers.console.wraps.dify_config.ENTERPRISE_ENABLED", True),
@@ -10,6 +10,7 @@ from werkzeug.exceptions import Unauthorized
import services.errors.account
from controllers.web.login import EmailCodeLoginApi, EmailCodeLoginSendEmailApi, LoginApi, LoginStatusApi, LogoutApi
from enums.deployment_edition import DeploymentEdition
from services.entities.auth_entities import LoginFailureReason
@@ -34,7 +35,7 @@ def app():
@pytest.fixture(autouse=True)
def _patch_wraps():
wraps_features = SimpleNamespace(enable_email_password_login=True)
console_dify = SimpleNamespace(ENTERPRISE_ENABLED=True, EDITION="CLOUD")
console_dify = SimpleNamespace(ENTERPRISE_ENABLED=True, DEPLOYMENT_EDITION=DeploymentEdition.CLOUD)
web_dify = SimpleNamespace(ENTERPRISE_ENABLED=True)
with (
patch("controllers.console.wraps.db") as mock_db,
@@ -0,0 +1,49 @@
from unittest import mock
import pytest
from werkzeug.exceptions import Unauthorized
from core.logging.context import clear_request_context, get_identity_context
@pytest.fixture(autouse=True)
def _reset_logging_context():
clear_request_context()
yield
clear_request_context()
def test_validate_jwt_token_sets_logging_identity_before_view() -> None:
from controllers.web import wraps
app_model = mock.Mock()
end_user = mock.Mock(id="end-user-id", tenant_id="tenant-id", type=None)
clear_request_context()
@wraps.validate_jwt_token
def protected_view(received_app, received_user):
assert get_identity_context() == ("tenant-id", "end-user-id", "end_user")
return received_app, received_user
with mock.patch.object(wraps, "decode_jwt_token", return_value=(app_model, end_user)):
result = protected_view()
assert result == (app_model, end_user)
def test_validate_jwt_token_does_not_set_identity_when_authentication_fails() -> None:
from controllers.web import wraps
clear_request_context()
@wraps.validate_jwt_token
def protected_view(_app, _user):
raise AssertionError("view must not be called")
with (
mock.patch.object(wraps, "decode_jwt_token", side_effect=Unauthorized()),
pytest.raises(Unauthorized),
):
protected_view()
assert get_identity_context() == ("", "", "")
@@ -0,0 +1,341 @@
from datetime import datetime
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.agent.publish_visibility import (
PUBLISH_VISIBLE_APP_BACKED_REVISION_OPERATIONS,
agent_has_workflow_callable_active_snapshot,
workflow_callable_active_snapshot_filter,
)
from models.agent import (
Agent,
AgentConfigRevision,
AgentConfigRevisionOperation,
AgentConfigSnapshot,
AgentKind,
AgentScope,
AgentSource,
AgentStatus,
WorkflowAgentNodeBinding,
)
from models.agent_config_entities import AgentSoulConfig, AgentSoulModelConfig
from services.agent.roster_service import AgentRosterService
def _agent_soul() -> AgentSoulConfig:
return AgentSoulConfig(
model=AgentSoulModelConfig(
plugin_id="langgenius/openai",
model_provider="openai",
model="gpt-test",
)
)
def _add_agent(
session: Session,
*,
agent_id: str,
snapshot_id: str | None,
name: str,
source: AgentSource,
operation: AgentConfigRevisionOperation | None,
app_id: str | None = None,
scope: AgentScope = AgentScope.ROSTER,
status: AgentStatus = AgentStatus.ACTIVE,
has_model: bool | None = None,
) -> Agent:
agent = Agent(
id=agent_id,
tenant_id="tenant-1",
name=name,
description="",
agent_kind=AgentKind.DIFY_AGENT,
scope=scope,
source=source,
app_id=app_id,
status=status,
active_config_snapshot_id=snapshot_id,
active_config_has_model=snapshot_id is not None if has_model is None else has_model,
# A dirty draft must not hide an already published active snapshot.
active_config_is_published=False,
)
session.add(agent)
if snapshot_id is None:
return agent
session.add(
AgentConfigSnapshot(
id=snapshot_id,
tenant_id="tenant-1",
agent_id=agent_id,
version=1,
config_snapshot=_agent_soul(),
)
)
if operation is not None:
session.add(
AgentConfigRevision(
id=f"revision-{agent_id}",
tenant_id="tenant-1",
agent_id=agent_id,
current_snapshot_id=snapshot_id,
revision=1,
operation=operation,
)
)
return agent
def test_publish_visible_operation_contract() -> None:
assert {
AgentConfigRevisionOperation.PUBLISH_DRAFT,
AgentConfigRevisionOperation.SAVE_CURRENT_VERSION,
AgentConfigRevisionOperation.SAVE_NEW_VERSION,
AgentConfigRevisionOperation.SAVE_NEW_AGENT,
AgentConfigRevisionOperation.SAVE_TO_ROSTER,
AgentConfigRevisionOperation.RESTORE_VERSION,
} == PUBLISH_VISIBLE_APP_BACKED_REVISION_OPERATIONS
@pytest.mark.parametrize(
"sqlite_session",
[(Agent, AgentConfigSnapshot, AgentConfigRevision, WorkflowAgentNodeBinding)],
indirect=True,
)
def test_workflow_callable_filter_distinguishes_never_published_from_dirty_drafts(
sqlite_session: Session,
) -> None:
imported_draft = _add_agent(
sqlite_session,
agent_id="agent-imported-draft",
snapshot_id="snapshot-imported-draft",
name="Imported draft",
source=AgentSource.IMPORTED,
operation=AgentConfigRevisionOperation.IMPORT_PACKAGE,
app_id="app-imported-draft",
)
app_draft = _add_agent(
sqlite_session,
agent_id="agent-app-draft",
snapshot_id="snapshot-app-draft",
name="App draft",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.CREATE_VERSION,
)
published_with_dirty_draft = _add_agent(
sqlite_session,
agent_id="agent-published-dirty",
snapshot_id="snapshot-published-dirty",
name="Published with dirty draft",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.PUBLISH_DRAFT,
)
saved_to_current_version = _add_agent(
sqlite_session,
agent_id="agent-saved-current",
snapshot_id="snapshot-saved-current",
name="Saved to current version",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION,
)
saved_as_new_agent = _add_agent(
sqlite_session,
agent_id="agent-saved-new",
snapshot_id="snapshot-saved-new",
name="Saved as new agent",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.SAVE_NEW_AGENT,
)
saved_as_new_version = _add_agent(
sqlite_session,
agent_id="agent-saved-new-version",
snapshot_id="snapshot-saved-new-version",
name="Saved as new version",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION,
)
saved_to_roster = _add_agent(
sqlite_session,
agent_id="agent-saved-to-roster",
snapshot_id="snapshot-saved-to-roster",
name="Saved to roster",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.SAVE_TO_ROSTER,
)
restored_version = _add_agent(
sqlite_session,
agent_id="agent-restored-version",
snapshot_id="snapshot-restored-version",
name="Restored version",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.RESTORE_VERSION,
)
direct_roster_agent = _add_agent(
sqlite_session,
agent_id="agent-direct-roster",
snapshot_id="snapshot-direct-roster",
name="Direct roster",
source=AgentSource.ROSTER,
operation=AgentConfigRevisionOperation.CREATE_VERSION,
)
direct_imported_roster_agent = _add_agent(
sqlite_session,
agent_id="agent-direct-imported",
snapshot_id="snapshot-direct-imported",
name="Direct imported roster",
source=AgentSource.IMPORTED,
operation=AgentConfigRevisionOperation.IMPORT_PACKAGE,
)
no_snapshot = _add_agent(
sqlite_session,
agent_id="agent-no-snapshot",
snapshot_id=None,
name="No snapshot",
source=AgentSource.ROSTER,
operation=None,
)
published_without_model = _add_agent(
sqlite_session,
agent_id="agent-published-without-model",
snapshot_id="snapshot-published-without-model",
name="Published without model",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.PUBLISH_DRAFT,
has_model=False,
)
archived_agent = _add_agent(
sqlite_session,
agent_id="agent-archived",
snapshot_id="snapshot-archived",
name="Archived",
source=AgentSource.ROSTER,
operation=AgentConfigRevisionOperation.CREATE_VERSION,
status=AgentStatus.ARCHIVED,
)
workflow_only_agent = _add_agent(
sqlite_session,
agent_id="agent-workflow-only",
snapshot_id="snapshot-workflow-only",
name="Workflow only",
source=AgentSource.WORKFLOW,
operation=AgentConfigRevisionOperation.CREATE_VERSION,
scope=AgentScope.WORKFLOW_ONLY,
)
stale_published_snapshot = _add_agent(
sqlite_session,
agent_id="agent-stale-published-snapshot",
snapshot_id="snapshot-current-unpublished",
name="Stale published snapshot",
source=AgentSource.AGENT_APP,
operation=AgentConfigRevisionOperation.CREATE_VERSION,
)
cross_owner_revision = _add_agent(
sqlite_session,
agent_id="agent-cross-owner-revision",
snapshot_id="snapshot-cross-owner-revision",
name="Cross owner revision",
source=AgentSource.AGENT_APP,
operation=None,
)
sqlite_session.add(
AgentConfigSnapshot(
id="snapshot-old-published",
tenant_id="tenant-1",
agent_id=stale_published_snapshot.id,
version=2,
config_snapshot=_agent_soul(),
)
)
sqlite_session.add(
AgentConfigRevision(
id="revision-old-published",
tenant_id="tenant-1",
agent_id=stale_published_snapshot.id,
current_snapshot_id="snapshot-old-published",
revision=2,
operation=AgentConfigRevisionOperation.PUBLISH_DRAFT,
)
)
sqlite_session.add(
AgentConfigRevision(
id="revision-published-dirty-2",
tenant_id="tenant-1",
agent_id=published_with_dirty_draft.id,
current_snapshot_id=published_with_dirty_draft.active_config_snapshot_id,
revision=2,
operation=AgentConfigRevisionOperation.RESTORE_VERSION,
)
)
sqlite_session.add(
AgentConfigRevision(
id="revision-cross-owner",
tenant_id="tenant-other",
agent_id="agent-other",
current_snapshot_id=cross_owner_revision.active_config_snapshot_id,
revision=1,
operation=AgentConfigRevisionOperation.PUBLISH_DRAFT,
)
)
imported_draft.updated_at = datetime(2031, 7, 24, 12, 0, 0)
published_with_dirty_draft.updated_at = datetime(2030, 7, 24, 11, 0, 0)
sqlite_session.commit()
callable_agent_ids = set(
sqlite_session.scalars(select(Agent.id).where(workflow_callable_active_snapshot_filter())).all()
)
assert callable_agent_ids == {
published_with_dirty_draft.id,
saved_to_current_version.id,
saved_as_new_agent.id,
saved_as_new_version.id,
saved_to_roster.id,
restored_version.id,
direct_roster_agent.id,
direct_imported_roster_agent.id,
published_without_model.id,
archived_agent.id,
workflow_only_agent.id,
}
assert agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=published_with_dirty_draft)
assert agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=direct_roster_agent)
assert agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=direct_imported_roster_agent)
assert not agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=imported_draft)
assert not agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=app_draft)
assert not agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=no_snapshot)
assert not agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=stale_published_snapshot)
assert not agent_has_workflow_callable_active_snapshot(session=sqlite_session, agent=cross_owner_revision)
invite_options = AgentRosterService(sqlite_session).list_invite_options(
tenant_id="tenant-1",
page=1,
limit=20,
)
expected_invite_ids = callable_agent_ids - {
published_without_model.id,
archived_agent.id,
workflow_only_agent.id,
}
assert invite_options["total"] == len(expected_invite_ids)
assert {item["id"] for item in invite_options["data"]} == expected_invite_ids
first_page = AgentRosterService(sqlite_session).list_invite_options(
tenant_id="tenant-1",
page=1,
limit=1,
)
assert first_page["total"] == len(expected_invite_ids)
assert first_page["has_more"] is True
assert [item["id"] for item in first_page["data"]] == [published_with_dirty_draft.id]
unpublished_keyword = AgentRosterService(sqlite_session).list_invite_options(
tenant_id="tenant-1",
page=1,
limit=20,
keyword="Imported draft",
)
assert unpublished_keyword["total"] == 0
assert unpublished_keyword["data"] == []
@@ -1,4 +1,3 @@
from types import SimpleNamespace
from typing import cast
import pytest
@@ -51,8 +50,8 @@ def test_check_credential_policy_compliance_returns_when_feature_disabled(
mocker: MockerFixture,
) -> None:
mocker.patch(
"services.feature_service.FeatureService.get_system_features",
return_value=SimpleNamespace(plugin_manager=SimpleNamespace(enabled=False)),
"services.feature_service.FeatureService.is_plugin_manager_enabled",
return_value=False,
)
check_call = mocker.patch(
"services.enterprise.plugin_manager_service.PluginManagerService.check_credential_policy_compliance"
@@ -67,8 +66,8 @@ def test_check_credential_policy_compliance_raises_when_credential_missing(
mocker: MockerFixture,
) -> None:
mocker.patch(
"services.feature_service.FeatureService.get_system_features",
return_value=SimpleNamespace(plugin_manager=SimpleNamespace(enabled=True)),
"services.feature_service.FeatureService.is_plugin_manager_enabled",
return_value=True,
)
mocker.patch("core.helper.credential_utils.is_credential_exists", return_value=False)
@@ -80,8 +79,8 @@ def test_check_credential_policy_compliance_calls_plugin_manager_with_request(
mocker: MockerFixture,
) -> None:
mocker.patch(
"services.feature_service.FeatureService.get_system_features",
return_value=SimpleNamespace(plugin_manager=SimpleNamespace(enabled=True)),
"services.feature_service.FeatureService.is_plugin_manager_enabled",
return_value=True,
)
mocker.patch("core.helper.credential_utils.is_credential_exists", return_value=True)
check_call = mocker.patch(
@@ -101,8 +100,8 @@ def test_check_credential_policy_compliance_skips_existence_check_when_disabled(
mocker: MockerFixture,
) -> None:
mocker.patch(
"services.feature_service.FeatureService.get_system_features",
return_value=SimpleNamespace(plugin_manager=SimpleNamespace(enabled=True)),
"services.feature_service.FeatureService.is_plugin_manager_enabled",
return_value=True,
)
exists_call = mocker.patch("core.helper.credential_utils.is_credential_exists")
check_call = mocker.patch(
@@ -124,8 +123,8 @@ def test_check_credential_policy_compliance_returns_when_credential_id_empty(
mocker: MockerFixture,
) -> None:
mocker.patch(
"services.feature_service.FeatureService.get_system_features",
return_value=SimpleNamespace(plugin_manager=SimpleNamespace(enabled=True)),
"services.feature_service.FeatureService.is_plugin_manager_enabled",
return_value=True,
)
exists_call = mocker.patch("core.helper.credential_utils.is_credential_exists")
check_call = mocker.patch(
@@ -1,15 +1,27 @@
"""Tests for logging context module."""
import uuid
from contextvars import copy_context
import pytest
from core.logging.context import (
clear_request_context,
get_identity_context,
get_request_id,
get_trace_id,
init_request_context,
set_identity_context,
)
@pytest.fixture(autouse=True)
def _reset_logging_context():
clear_request_context()
yield
clear_request_context()
class TestLoggingContext:
"""Tests for the logging context functions."""
@@ -77,3 +89,41 @@ class TestLoggingContext:
# IDs should be different
assert id1 != id2
def test_set_identity_context(self):
set_identity_context(tenant_id="tenant-1", user_id="user-1", user_type="end_user")
identity = get_identity_context()
assert identity.tenant_id == "tenant-1"
assert identity.user_id == "user-1"
assert identity.user_type == "end_user"
def test_set_identity_context_replaces_all_fields(self):
set_identity_context(tenant_id="tenant-1", user_id="user-1", user_type="account")
set_identity_context(user_id="user-2", user_type="end_user")
assert get_identity_context() == ("", "user-2", "end_user")
def test_identity_context_is_copied_as_primitive_values(self):
set_identity_context(tenant_id="tenant-1", user_id="user-1", user_type="end_user")
copied_context = copy_context()
clear_request_context()
assert get_identity_context() == ("", "", "")
assert copied_context.run(get_identity_context) == ("tenant-1", "user-1", "end_user")
def test_init_clears_existing_identity_context(self):
set_identity_context(tenant_id="tenant-1", user_id="user-1", user_type="end_user")
init_request_context()
assert get_identity_context() == ("", "", "")
def test_clear_resets_identity_context(self):
set_identity_context(tenant_id="tenant-1", user_id="user-1", user_type="end_user")
clear_request_context()
assert get_identity_context() == ("", "", "")
+83 -110
View File
@@ -1,5 +1,6 @@
"""Tests for logging filters."""
import io
import logging
from unittest import mock
@@ -19,6 +20,15 @@ def log_record():
)
@pytest.fixture(autouse=True)
def _reset_logging_context():
from core.logging.context import clear_request_context
clear_request_context()
yield
clear_request_context()
class TestTraceContextFilter:
def test_sets_empty_trace_id_without_context(self, log_record):
from core.logging.context import clear_request_context
@@ -147,8 +157,10 @@ class TestTraceContextFilter:
class TestIdentityContextFilter:
def test_sets_empty_identity_without_request_context(self, log_record):
from core.logging.context import clear_request_context
from core.logging.filters import IdentityContextFilter
clear_request_context()
filter = IdentityContextFilter()
result = filter.filter(log_record)
@@ -164,131 +176,92 @@ class TestIdentityContextFilter:
result = filter.filter(log_record)
assert result is True
def test_handles_exception_gracefully(self, log_record):
def test_uses_explicit_identity_context_without_flask_context(self, log_record):
from core.logging.context import set_identity_context
from core.logging.filters import IdentityContextFilter
set_identity_context(tenant_id="tenant_id", user_id="end_user_id", user_type="end_user")
filter = IdentityContextFilter()
filter.filter(log_record)
# Should not raise even if something goes wrong
with mock.patch(
"core.logging.filters.flask.has_request_context", side_effect=Exception("Test error"), autospec=True
):
result = filter.filter(log_record)
assert result is True
assert log_record.tenant_id == ""
assert log_record.tenant_id == "tenant_id"
assert log_record.user_id == "end_user_id"
assert log_record.user_type == "end_user"
def test_sets_empty_identity_unauthenticated(self, log_record):
def test_does_not_trigger_flask_login_request_loader(self, log_record):
from flask import Flask
from flask_login import LoginManager
from core.logging.context import clear_request_context
from core.logging.filters import IdentityContextFilter
mock_user = mock.MagicMock()
mock_user.is_authenticated = False
app = Flask(__name__)
app.secret_key = "test"
login_manager = LoginManager(app)
request_loader = mock.Mock(return_value=None)
login_manager.request_loader(request_loader)
clear_request_context()
with (
mock.patch("flask.has_request_context", return_value=True),
mock.patch("flask_login.current_user", mock_user),
):
filter = IdentityContextFilter()
filter.filter(log_record)
assert log_record.user_id == ""
with app.test_request_context("/"):
from flask import g
def test_sets_identity_for_account(self, log_record):
assert "_login_user" not in g
IdentityContextFilter().filter(log_record)
assert "_login_user" not in g
request_loader.assert_not_called()
assert log_record.tenant_id == ""
assert log_record.user_id == ""
assert log_record.user_type == ""
def test_ended_otel_span_warning_does_not_trigger_request_loader(self):
from flask import Flask, g
from flask_login import LoginManager
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
from core.logging.context import clear_request_context
from core.logging.filters import IdentityContextFilter
class MockAccount:
pass
app = Flask(__name__)
app.secret_key = "test"
login_manager = LoginManager(app)
request_loader = mock.Mock(return_value=None)
login_manager.request_loader(request_loader)
mock_user = MockAccount()
mock_user.id = "account_id"
mock_user.current_tenant_id = "tenant_id"
mock_user.is_authenticated = True
span = TracerProvider().get_tracer(__name__).start_span("ended")
span.end()
with (
mock.patch("flask.has_request_context", return_value=True),
mock.patch("models.Account", MockAccount),
mock.patch("flask_login.current_user", mock_user),
):
filter = IdentityContextFilter()
filter.filter(log_record)
stream = io.StringIO()
handler = logging.StreamHandler(stream)
handler.addFilter(IdentityContextFilter())
handler.setFormatter(logging.Formatter("%(tenant_id)s %(user_id)s %(user_type)s %(message)s"))
assert log_record.tenant_id == "tenant_id"
assert log_record.user_id == "account_id"
assert log_record.user_type == "account"
sdk_logger = logging.getLogger("opentelemetry.sdk.trace")
previous_level = sdk_logger.level
previous_propagate = sdk_logger.propagate
previous_disabled = sdk_logger.disabled
sdk_logger.addHandler(handler)
sdk_logger.setLevel(logging.WARNING)
sdk_logger.propagate = False
sdk_logger.disabled = False
clear_request_context()
def test_sets_identity_for_account_no_tenant(self, log_record):
from core.logging.filters import IdentityContextFilter
try:
with app.test_request_context("/"), trace.use_span(span, end_on_exit=False):
assert "_login_user" not in g
class MockAccount:
pass
span.set_attribute("test.key", "test-value")
mock_user = MockAccount()
mock_user.id = "account_id"
mock_user.current_tenant_id = None
mock_user.is_authenticated = True
assert "_login_user" not in g
finally:
clear_request_context()
sdk_logger.removeHandler(handler)
sdk_logger.setLevel(previous_level)
sdk_logger.propagate = previous_propagate
sdk_logger.disabled = previous_disabled
handler.close()
with (
mock.patch("flask.has_request_context", return_value=True),
mock.patch("models.Account", MockAccount),
mock.patch("flask_login.current_user", mock_user),
):
filter = IdentityContextFilter()
filter.filter(log_record)
assert log_record.tenant_id == ""
assert log_record.user_id == "account_id"
assert log_record.user_type == "account"
def test_sets_identity_for_end_user(self, log_record):
from core.logging.filters import IdentityContextFilter
class MockEndUser:
pass
class AnotherClass:
pass
mock_user = MockEndUser()
mock_user.id = "end_user_id"
mock_user.tenant_id = "tenant_id"
mock_user.type = "custom_type"
mock_user.is_authenticated = True
with (
mock.patch("flask.has_request_context", return_value=True),
mock.patch("models.model.EndUser", MockEndUser),
mock.patch("models.Account", AnotherClass),
mock.patch("flask_login.current_user", mock_user),
):
filter = IdentityContextFilter()
filter.filter(log_record)
assert log_record.tenant_id == "tenant_id"
assert log_record.user_id == "end_user_id"
assert log_record.user_type == "custom_type"
def test_sets_identity_for_end_user_default_type(self, log_record):
from core.logging.filters import IdentityContextFilter
class MockEndUser:
pass
class AnotherClass:
pass
mock_user = MockEndUser()
mock_user.id = "end_user_id"
mock_user.tenant_id = "tenant_id"
mock_user.type = None
mock_user.is_authenticated = True
with (
mock.patch("flask.has_request_context", return_value=True),
mock.patch("models.model.EndUser", MockEndUser),
mock.patch("models.Account", AnotherClass),
mock.patch("flask_login.current_user", mock_user),
):
filter = IdentityContextFilter()
filter.filter(log_record)
assert log_record.tenant_id == "tenant_id"
assert log_record.user_id == "end_user_id"
assert log_record.user_type == "end_user"
request_loader.assert_not_called()
assert "Setting attribute on ended span" in stream.getvalue()
@@ -1,12 +1,18 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from core.ops.base_trace_instance import BaseTraceInstance
from core.ops.entities.config_entity import BaseTracingConfig
from core.ops.entities.trace_entity import BaseTraceInfo
from models import Account, App, TenantAccountJoin
from models import Account, App, Tenant, TenantAccountJoin
from models.account import TenantAccountRole
from models.model import AppMode
TABLES = (App, Account, Tenant, TenantAccountJoin)
class ConcreteTraceInstance(BaseTraceInstance):
@@ -17,91 +23,101 @@ class ConcreteTraceInstance(BaseTraceInstance):
super().trace(trace_info)
@pytest.fixture
def mock_db_session(monkeypatch: pytest.MonkeyPatch):
mock_session = MagicMock(spec=Session)
mock_session.__enter__.return_value = mock_session
mock_session.__exit__.return_value = None
mock_session_class = MagicMock(return_value=mock_session)
monkeypatch.setattr("core.ops.base_trace_instance.Session", mock_session_class)
monkeypatch.setattr("core.ops.base_trace_instance.db", MagicMock())
return mock_session
@pytest.fixture(autouse=True)
def _bind_production_sessions(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None:
"""Bind both service-owned ORM sessions to the test's SQLite engine."""
database = SimpleNamespace(engine=sqlite_engine)
monkeypatch.setattr("core.ops.base_trace_instance.db", database)
monkeypatch.setattr("models.account.db", database)
def test_get_service_account_with_tenant_app_not_found(mock_db_session):
mock_db_session.scalar.return_value = None
def _persist_app(session: Session, *, created_by: str | None) -> App:
app = App(
id="some_app_id",
tenant_id="tenant_id",
name="Trace App",
description="",
mode=AppMode.CHAT,
icon_type=None,
icon=None,
icon_background=None,
app_model_config_id=None,
workflow_id=None,
enable_site=True,
enable_api=True,
max_active_requests=None,
created_by=created_by,
)
session.add(app)
session.commit()
return app
config = MagicMock(spec=BaseTracingConfig)
instance = ConcreteTraceInstance(config)
def _persist_account(session: Session) -> Account:
account = Account(name="Creator", email="creator@example.com")
account.id = "creator_id"
session.add(account)
session.commit()
return account
def _trace_instance() -> ConcreteTraceInstance:
# Tracing configuration is a domain collaborator, not an ORM session.
return ConcreteTraceInstance(MagicMock(spec=BaseTracingConfig))
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_get_service_account_with_tenant_app_not_found(sqlite_session: Session):
with pytest.raises(ValueError, match="App with id some_app_id not found"):
instance.get_service_account_with_tenant("some_app_id")
_trace_instance().get_service_account_with_tenant("some_app_id")
def test_get_service_account_with_tenant_no_creator(mock_db_session):
mock_app = MagicMock(spec=App)
mock_app.id = "some_app_id"
mock_app.created_by = None
mock_db_session.scalar.return_value = mock_app
config = MagicMock(spec=BaseTracingConfig)
instance = ConcreteTraceInstance(config)
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_get_service_account_with_tenant_no_creator(sqlite_session: Session):
_persist_app(sqlite_session, created_by=None)
with pytest.raises(ValueError, match="App with id some_app_id has no creator"):
instance.get_service_account_with_tenant("some_app_id")
_trace_instance().get_service_account_with_tenant("some_app_id")
def test_get_service_account_with_tenant_creator_not_found(mock_db_session):
mock_app = MagicMock(spec=App)
mock_app.id = "some_app_id"
mock_app.created_by = "creator_id"
# First call to scalar returns app, second returns None (for account)
mock_db_session.scalar.side_effect = [mock_app, None]
config = MagicMock(spec=BaseTracingConfig)
instance = ConcreteTraceInstance(config)
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_get_service_account_with_tenant_creator_not_found(sqlite_session: Session):
_persist_app(sqlite_session, created_by="creator_id")
with pytest.raises(ValueError, match="Creator account with id creator_id not found for app some_app_id"):
instance.get_service_account_with_tenant("some_app_id")
_trace_instance().get_service_account_with_tenant("some_app_id")
def test_get_service_account_with_tenant_tenant_not_found(mock_db_session):
mock_app = MagicMock(spec=App)
mock_app.id = "some_app_id"
mock_app.created_by = "creator_id"
mock_account = MagicMock(spec=Account)
mock_account.id = "creator_id"
mock_db_session.scalar.side_effect = [mock_app, mock_account, None]
config = MagicMock(spec=BaseTracingConfig)
instance = ConcreteTraceInstance(config)
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_get_service_account_with_tenant_tenant_not_found(sqlite_session: Session):
_persist_app(sqlite_session, created_by="creator_id")
_persist_account(sqlite_session)
with pytest.raises(ValueError, match="Current tenant not found for account creator_id"):
instance.get_service_account_with_tenant("some_app_id")
_trace_instance().get_service_account_with_tenant("some_app_id")
def test_get_service_account_with_tenant_success(mock_db_session):
mock_app = MagicMock(spec=App)
mock_app.id = "some_app_id"
mock_app.created_by = "creator_id"
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
def test_get_service_account_with_tenant_success(sqlite_session: Session):
_persist_app(sqlite_session, created_by="creator_id")
_persist_account(sqlite_session)
tenant = Tenant(name="Workspace")
tenant.id = "tenant_id"
sqlite_session.add_all(
[
tenant,
TenantAccountJoin(
tenant_id=tenant.id,
account_id="creator_id",
current=True,
role=TenantAccountRole.OWNER,
),
]
)
sqlite_session.commit()
mock_account = MagicMock(spec=Account)
mock_account.id = "creator_id"
result = _trace_instance().get_service_account_with_tenant("some_app_id")
mock_tenant_join = MagicMock(spec=TenantAccountJoin)
mock_tenant_join.tenant_id = "tenant_id"
mock_db_session.scalar.side_effect = [mock_app, mock_account, mock_tenant_join]
config = MagicMock(spec=BaseTracingConfig)
instance = ConcreteTraceInstance(config)
result = instance.get_service_account_with_tenant("some_app_id")
assert result == mock_account
mock_account.set_tenant_id_with_session.assert_called_once_with("tenant_id", session=mock_db_session)
assert result.id == "creator_id"
assert result.current_tenant_id == "tenant_id"
assert result.current_role == TenantAccountRole.OWNER
@@ -8,10 +8,49 @@ which provides document storage and retrieval functionality for datasets in the
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy import func, select
from sqlalchemy.orm import Session
from core.rag.docstore.dataset_docstore import DatasetDocumentStore, DocumentSegment
from core.rag.models.document import AttachmentDocument, Document
from models.dataset import Dataset
from core.rag.models.document import AttachmentDocument, ChildDocument, Document
from models.dataset import ChildChunk, Dataset, SegmentAttachmentBinding
TENANT_ID = "00000000-0000-0000-0000-000000000001"
DATASET_ID = "00000000-0000-0000-0000-000000000002"
DOCUMENT_ID = "00000000-0000-0000-0000-000000000003"
USER_ID = "00000000-0000-0000-0000-000000000004"
def _dataset() -> Dataset:
dataset = MagicMock(spec=Dataset)
dataset.id = DATASET_ID
dataset.tenant_id = TENANT_ID
return dataset
def _persist_segment(
session: Session,
*,
index_node_id: str = "doc-1",
index_node_hash: str = "hash-1",
content: str = "Test content",
tokens: int = 5,
) -> DocumentSegment:
segment = DocumentSegment(
tenant_id=TENANT_ID,
dataset_id=DATASET_ID,
document_id=DOCUMENT_ID,
position=1,
content=content,
word_count=len(content),
tokens=tokens,
created_by=USER_ID,
index_node_id=index_node_id,
index_node_hash=index_node_hash,
)
session.add(segment)
session.flush()
return segment
class TestDatasetDocumentStoreInit:
@@ -132,228 +171,153 @@ class TestDatasetDocumentStoreDocs:
assert result == {}
@pytest.mark.parametrize(
"sqlite_session",
[(DocumentSegment, ChildChunk, SegmentAttachmentBinding)],
indirect=True,
)
class TestDatasetDocumentStoreAddDocuments:
"""Tests for add_documents method."""
def test_add_documents_new_document_with_embedding(self):
"""Test adding new documents with embedding model."""
def test_add_documents_new_document_with_token_count(self, sqlite_session: Session):
"""Test adding a new document with a precomputed token count."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "high_quality"
mock_dataset.embedding_model_provider = "provider"
mock_dataset.embedding_model = "model"
document = Document(
page_content="Test content",
metadata={"doc_id": "doc-1", "doc_hash": "hash-1"},
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Test content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "hash-1",
}
mock_doc.attachments = None
mock_doc.children = None
store.add_documents(session=sqlite_session, docs=[document], token_counts=[10])
sqlite_session.expire_all()
mock_model_instance = MagicMock()
mock_model_instance.get_text_embedding_num_tokens.return_value = [10]
segment = sqlite_session.scalar(
select(DocumentSegment).where(
DocumentSegment.dataset_id == DATASET_ID,
DocumentSegment.index_node_id == "doc-1",
)
)
assert segment is not None
assert segment.content == "Test content"
assert segment.tokens == 10
assert segment.position == 1
with (
patch("core.rag.docstore.dataset_docstore.ModelManager.for_tenant") as mock_manager_class,
):
mock_session = MagicMock()
mock_session.scalar.return_value = None
mock_manager = MagicMock()
mock_manager.get_model_instance.return_value = mock_model_instance
mock_manager_class.return_value = mock_manager
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None):
with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
)
store.add_documents([mock_doc], session=mock_session)
mock_session.add.assert_called()
mock_session.flush.assert_called()
def test_add_documents_update_existing_document(self):
def test_add_documents_update_existing_document(self, sqlite_session: Session):
"""Test updating existing document with allow_update=True."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "economy"
mock_dataset.embedding_model_provider = None
mock_dataset.embedding_model = None
existing_segment = _persist_segment(sqlite_session)
document = Document(
page_content="Updated content",
metadata={"doc_id": "doc-1", "doc_hash": "new-hash"},
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Updated content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "new-hash",
}
mock_doc.attachments = None
mock_doc.children = None
store.add_documents(session=sqlite_session, docs=[document], token_counts=[0])
sqlite_session.expire_all()
mock_existing_segment = MagicMock()
mock_existing_segment.id = "seg-1"
updated_segment = sqlite_session.get(DocumentSegment, existing_segment.id)
assert updated_segment is not None
assert updated_segment.content == "Updated content"
assert updated_segment.index_node_hash == "new-hash"
assert updated_segment.tokens == 0
mock_session = MagicMock()
mock_session.scalar.return_value = 5
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment):
with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
)
store.add_documents([mock_doc], session=mock_session)
mock_session.flush.assert_called()
def test_add_documents_raises_when_not_allowed(self):
def test_add_documents_raises_when_not_allowed(self, sqlite_session: Session):
"""Test that adding existing doc without allow_update raises ValueError."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "economy"
_persist_segment(sqlite_session)
document = Document(
page_content="Test content",
metadata={"doc_id": "doc-1", "doc_hash": "hash-1"},
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Test content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "hash-1",
}
mock_doc.attachments = None
mock_doc.children = None
mock_existing_segment = MagicMock()
mock_session = MagicMock()
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
with pytest.raises(ValueError, match="already exists"):
store.add_documents(
session=sqlite_session,
docs=[document],
token_counts=[0],
allow_update=False,
)
with pytest.raises(ValueError, match="already exists"):
store.add_documents([mock_doc], session=mock_session, allow_update=False)
assert sqlite_session.scalar(select(func.count()).select_from(DocumentSegment)) == 1
def test_add_documents_with_answer_metadata(self):
def test_add_documents_with_answer_metadata(self, sqlite_session: Session):
"""Test adding document with answer in metadata."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "economy"
document = Document(
page_content="Test content",
metadata={
"doc_id": "doc-1",
"doc_hash": "hash-1",
"answer": "Test answer",
},
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Test content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "hash-1",
"answer": "Test answer",
}
mock_doc.attachments = None
mock_doc.children = None
store.add_documents(session=sqlite_session, docs=[document], token_counts=[0])
sqlite_session.expire_all()
mock_session = MagicMock()
mock_session.scalar.return_value = None
segment = sqlite_session.scalar(select(DocumentSegment))
assert segment is not None
assert segment.answer == "Test answer"
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None):
with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
)
store.add_documents([mock_doc], session=mock_session)
mock_session.add.assert_called()
def test_add_documents_with_invalid_document_type(self):
def test_add_documents_with_invalid_document_type(self, sqlite_session: Session):
"""Test that non-Document raises ValueError."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_session = MagicMock()
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
with pytest.raises(ValueError, match="must be a Document"):
store.add_documents(["not a document"], session=mock_session)
store.add_documents(session=sqlite_session, docs=["not a document"], token_counts=[0]) # type: ignore[list-item]
def test_add_documents_with_none_metadata(self):
def test_add_documents_with_none_metadata(self, sqlite_session: Session):
"""Test that document with None metadata raises ValueError."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Test content"
mock_doc.metadata = None
mock_session = MagicMock()
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
)
document = MagicMock(spec=Document)
document.metadata = None
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
with pytest.raises(ValueError, match="metadata must be a dict"):
store.add_documents([mock_doc], session=mock_session)
store.add_documents(session=sqlite_session, docs=[document], token_counts=[0])
def test_add_documents_with_save_child(self):
def test_add_documents_with_save_child(self, sqlite_session: Session):
"""Test adding documents with save_child=True."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "economy"
mock_child = MagicMock(spec=Document)
mock_child.page_content = "Child content"
mock_child.metadata = {
"doc_id": "child-1",
"doc_hash": "child-hash",
}
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Test content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "hash-1",
}
mock_doc.attachments = None
mock_doc.children = [mock_child]
mock_session = MagicMock()
mock_session.scalar.return_value = None
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None):
with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
document = Document(
page_content="Test content",
metadata={"doc_id": "doc-1", "doc_hash": "hash-1"},
children=[
ChildDocument(
page_content="Child content",
metadata={"doc_id": "child-1", "doc_hash": "child-hash"},
)
],
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
store.add_documents([mock_doc], session=mock_session, save_child=True)
store.add_documents(
session=sqlite_session,
docs=[document],
token_counts=[0],
save_child=True,
)
sqlite_session.expire_all()
mock_session.add.assert_called()
child = sqlite_session.scalar(select(ChildChunk))
assert child is not None
assert child.content == "Child content"
assert child.index_node_id == "child-1"
def test_add_documents_rejects_mismatched_token_counts(self, sqlite_session: Session):
document = Document(
page_content="Test content",
metadata={"doc_id": "doc-1", "doc_hash": "hash-1"},
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
with pytest.raises(ValueError):
store.add_documents(session=sqlite_session, docs=[document], token_counts=[])
assert sqlite_session.scalar(select(func.count()).select_from(DocumentSegment)) == 0
class TestDatasetDocumentStoreExists:
@@ -722,88 +686,85 @@ class TestDatasetDocumentStoreMultimodelBinding:
mock_session.add.assert_not_called()
@pytest.mark.parametrize(
"sqlite_session",
[(DocumentSegment, ChildChunk, SegmentAttachmentBinding)],
indirect=True,
)
class TestDatasetDocumentStoreAddDocumentsUpdateChild:
"""Tests for add_documents when updating existing documents with children."""
def test_add_documents_update_existing_with_children(self):
def test_add_documents_update_existing_with_children(self, sqlite_session: Session):
"""Test updating existing document with save_child=True and children."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "economy"
mock_child = MagicMock(spec=Document)
mock_child.page_content = "Updated child content"
mock_child.metadata = {
"doc_id": "child-1",
"doc_hash": "new-child-hash",
}
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Updated content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "new-hash",
}
mock_doc.attachments = None
mock_doc.children = [mock_child]
mock_existing_segment = MagicMock()
mock_existing_segment.id = "seg-1"
mock_session = MagicMock()
mock_session.scalar.return_value = 5
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment):
with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
segment = _persist_segment(sqlite_session)
sqlite_session.add(
ChildChunk(
tenant_id=TENANT_ID,
dataset_id=DATASET_ID,
document_id=DOCUMENT_ID,
segment_id=segment.id,
position=1,
index_node_id="old-child",
index_node_hash="old-child-hash",
content="Old child content",
word_count=len("Old child content"),
created_by=USER_ID,
)
)
sqlite_session.flush()
document = Document(
page_content="Updated content",
metadata={"doc_id": "doc-1", "doc_hash": "new-hash"},
children=[
ChildDocument(
page_content="Updated child content",
metadata={"doc_id": "child-1", "doc_hash": "new-child-hash"},
)
],
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
store.add_documents([mock_doc], session=mock_session, save_child=True)
store.add_documents(
session=sqlite_session,
docs=[document],
token_counts=[0],
save_child=True,
)
sqlite_session.expire_all()
mock_session.execute.assert_called()
mock_session.flush.assert_called()
children = sqlite_session.scalars(select(ChildChunk).order_by(ChildChunk.position)).all()
assert len(children) == 1
assert children[0].content == "Updated child content"
assert children[0].index_node_id == "child-1"
@pytest.mark.parametrize(
"sqlite_session",
[(DocumentSegment, ChildChunk, SegmentAttachmentBinding)],
indirect=True,
)
class TestDatasetDocumentStoreAddDocumentsUpdateAnswer:
"""Tests for add_documents when updating existing documents with answer metadata."""
def test_add_documents_update_existing_with_answer(self):
def test_add_documents_update_existing_with_answer(self, sqlite_session: Session):
"""Test updating existing document with answer in metadata."""
mock_dataset = MagicMock(spec=Dataset)
mock_dataset.id = "test-dataset-id"
mock_dataset.tenant_id = "tenant-1"
mock_dataset.indexing_technique = "economy"
existing_segment = _persist_segment(sqlite_session)
document = Document(
page_content="Updated content",
metadata={
"doc_id": "doc-1",
"doc_hash": "new-hash",
"answer": "Updated answer",
},
)
store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID, document_id=DOCUMENT_ID)
mock_doc = MagicMock(spec=Document)
mock_doc.page_content = "Updated content"
mock_doc.metadata = {
"doc_id": "doc-1",
"doc_hash": "new-hash",
"answer": "Updated answer",
}
mock_doc.attachments = None
mock_doc.children = None
store.add_documents(session=sqlite_session, docs=[document], token_counts=[0])
sqlite_session.expire_all()
mock_existing_segment = MagicMock()
mock_existing_segment.id = "seg-1"
mock_session = MagicMock()
mock_session.scalar.return_value = 5
with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment):
with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"):
store = DatasetDocumentStore(
dataset=mock_dataset,
user_id="test-user-id",
document_id="test-doc-id",
)
store.add_documents([mock_doc], session=mock_session)
mock_session.flush.assert_called()
updated_segment = sqlite_session.get(DocumentSegment, existing_segment.id)
assert updated_segment is not None
assert updated_segment.answer == "Updated answer"
assert updated_segment.tokens == 0
@@ -0,0 +1,56 @@
from unittest.mock import Mock, patch
from core.rag.embedding.token_counter import calculate_segment_token_counts
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.models.document import Document
from models.dataset import Dataset
def test_high_quality_counts_each_document_once() -> None:
dataset = Mock(spec=Dataset)
dataset.tenant_id = "tenant-1"
dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY
dataset.embedding_model_provider = "provider"
dataset.embedding_model = "model"
documents = [
Document(page_content="first", metadata={}),
Document(page_content="second", metadata={}),
Document(page_content="third", metadata={}),
]
with patch("core.rag.embedding.token_counter.ModelManager.for_tenant") as model_manager_factory:
embedding_model = model_manager_factory.return_value.get_model_instance.return_value
embedding_model.get_text_embedding_num_tokens.return_value = [11, 22, 33]
result = calculate_segment_token_counts(dataset=dataset, documents=documents)
assert result == [11, 22, 33]
model_manager_factory.assert_called_once_with(tenant_id=dataset.tenant_id)
model_manager_factory.return_value.get_model_instance.assert_called_once()
embedding_model.get_text_embedding_num_tokens.assert_called_once_with(["first", "second", "third"])
def test_economy_returns_zero_without_loading_model() -> None:
dataset = Mock(spec=Dataset)
dataset.indexing_technique = IndexTechniqueType.ECONOMY
documents = [
Document(page_content="first", metadata={}),
Document(page_content="second", metadata={}),
]
with patch("core.rag.embedding.token_counter.ModelManager.for_tenant") as model_manager_factory:
result = calculate_segment_token_counts(dataset=dataset, documents=documents)
assert result == [0, 0]
model_manager_factory.assert_not_called()
def test_empty_documents_return_without_loading_model() -> None:
dataset = Mock(spec=Dataset)
dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY
with patch("core.rag.embedding.token_counter.ModelManager.for_tenant") as model_manager_factory:
result = calculate_segment_token_counts(dataset=dataset, documents=[])
assert result == []
model_manager_factory.assert_not_called()
@@ -271,14 +271,26 @@ class TestParagraphIndexProcessor:
patch(
"core.rag.index_processor.processor.paragraph_index_processor.DatasetDocumentStore"
) as mock_store_cls,
patch(
"core.rag.index_processor.processor.paragraph_index_processor.calculate_segment_token_counts"
) as mock_token_counter,
patch("core.rag.index_processor.processor.paragraph_index_processor.Vector") as mock_vector_cls,
):
mock_token_counter.side_effect = lambda **_kwargs: phase_events.append("count") or [11, 22]
mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store")
mock_vector_cls.return_value.create.side_effect = lambda _documents: phase_events.append("vector")
processor.index(dataset, dataset_document, ["chunk-1", "chunk-2"], session)
assert phase_events == ["store", "commit", "vector"]
mock_store_cls.return_value.add_documents.assert_called_once()
assert phase_events == ["count", "store", "commit", "vector"]
documents = mock_token_counter.call_args.kwargs["documents"]
assert [document.page_content for document in documents] == ["chunk-1", "chunk-2"]
mock_token_counter.assert_called_once_with(dataset=dataset, documents=documents)
mock_store_cls.return_value.add_documents.assert_called_once_with(
session=session,
docs=documents,
token_counts=[11, 22],
save_child=False,
)
mock_vector_cls.assert_called_once_with(dataset, session=session)
mock_vector_cls.return_value.create.assert_called_once()
mock_vector_cls.return_value.create_multimodal.assert_called_once()
@@ -299,13 +311,18 @@ class TestParagraphIndexProcessor:
patch(
"core.rag.index_processor.processor.paragraph_index_processor.DatasetDocumentStore"
) as mock_store_cls,
patch(
"core.rag.index_processor.processor.paragraph_index_processor.calculate_segment_token_counts"
) as mock_token_counter,
patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls,
):
mock_token_counter.side_effect = lambda **_kwargs: phase_events.append("count") or [0]
mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store")
mock_keyword_cls.return_value.add_texts.side_effect = lambda *_args: phase_events.append("keyword")
processor.index(dataset, dataset_document, ["chunk-3"], session)
assert phase_events == ["store", "commit", "keyword"]
assert phase_events == ["count", "store", "commit", "keyword"]
mock_token_counter.assert_called_once()
mock_keyword_cls.return_value.add_texts.assert_called_once()
def test_index_multimodal_structure_handles_files_and_account_lookup(
@@ -341,6 +358,10 @@ class TestParagraphIndexProcessor:
processor, "_get_content_files", return_value=[AttachmentDocument(page_content="img", metadata={})]
) as mock_files,
patch("core.rag.index_processor.processor.paragraph_index_processor.DatasetDocumentStore"),
patch(
"core.rag.index_processor.processor.paragraph_index_processor.calculate_segment_token_counts",
return_value=[11, 22],
),
patch("core.rag.index_processor.processor.paragraph_index_processor.Vector"),
):
processor.index(dataset, dataset_document, {"general_chunks": []}, session)
@@ -362,17 +362,29 @@ class TestParentChildIndexProcessor:
patch(
"core.rag.index_processor.processor.parent_child_index_processor.DatasetDocumentStore"
) as mock_store_cls,
patch(
"core.rag.index_processor.processor.parent_child_index_processor.calculate_segment_token_counts"
) as mock_token_counter,
patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls,
):
mock_token_counter.side_effect = lambda **_kwargs: phase_events.append("count") or [11]
mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store")
mock_vector_cls.return_value.create.side_effect = lambda _documents: phase_events.append("vector")
processor.index(dataset, dataset_document, {"parent_child_chunks": []}, session)
assert phase_events == ["store", "commit", "vector"]
assert phase_events == ["count", "store", "commit", "vector"]
assert dataset_document.dataset_process_rule_id == "rule-1"
session.add.assert_called_once_with(dataset_rule)
session.flush.assert_called_once()
mock_store_cls.return_value.add_documents.assert_called_once()
documents = mock_token_counter.call_args.kwargs["documents"]
assert [document.page_content for document in documents] == ["parent text"]
mock_token_counter.assert_called_once_with(dataset=dataset, documents=documents)
mock_store_cls.return_value.add_documents.assert_called_once_with(
session=session,
docs=documents,
token_counts=[11],
save_child=True,
)
mock_vector_cls.assert_called_once_with(dataset, session=session)
assert mock_vector_cls.return_value.create.call_count == 1
mock_vector_cls.return_value.create_multimodal.assert_called_once()
@@ -413,6 +425,10 @@ class TestParentChildIndexProcessor:
processor, "_get_content_files", return_value=[AttachmentDocument(page_content="image", metadata={})]
) as mock_files,
patch("core.rag.index_processor.processor.parent_child_index_processor.DatasetDocumentStore"),
patch(
"core.rag.index_processor.processor.parent_child_index_processor.calculate_segment_token_counts",
return_value=[11],
),
patch("core.rag.index_processor.processor.parent_child_index_processor.Vector"),
):
processor.index(dataset, dataset_document, {"parent_child_chunks": []}, session)
@@ -292,14 +292,26 @@ class TestQAIndexProcessor:
"core.rag.index_processor.processor.qa_index_processor.helper.generate_text_hash", return_value="hash"
),
patch("core.rag.index_processor.processor.qa_index_processor.DatasetDocumentStore") as mock_store_cls,
patch(
"core.rag.index_processor.processor.qa_index_processor.calculate_segment_token_counts"
) as mock_token_counter,
patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls,
):
mock_token_counter.side_effect = lambda **_kwargs: phase_events.append("count") or [11, 22]
mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store")
mock_vector_cls.return_value.create.side_effect = lambda _documents: phase_events.append("vector")
processor.index(dataset, dataset_document, {"qa_chunks": []}, session)
assert phase_events == ["store", "commit", "vector"]
mock_store_cls.return_value.add_documents.assert_called_once()
assert phase_events == ["count", "store", "commit", "vector"]
documents = mock_token_counter.call_args.kwargs["documents"]
assert [document.page_content for document in documents] == ["Q1", "Q2"]
mock_token_counter.assert_called_once_with(dataset=dataset, documents=documents)
mock_store_cls.return_value.add_documents.assert_called_once_with(
session=session,
docs=documents,
token_counts=[11, 22],
save_child=False,
)
mock_vector_cls.return_value.create.assert_called_once()
def test_index_requires_high_quality(
@@ -318,6 +330,10 @@ class TestQAIndexProcessor:
"core.rag.index_processor.processor.qa_index_processor.helper.generate_text_hash", return_value="hash"
),
patch("core.rag.index_processor.processor.qa_index_processor.DatasetDocumentStore"),
patch(
"core.rag.index_processor.processor.qa_index_processor.calculate_segment_token_counts",
return_value=[0],
),
):
with pytest.raises(ValueError, match="must be high quality"):
processor.index(dataset, dataset_document, {"qa_chunks": []}, session)
@@ -69,6 +69,7 @@ from graphon.model_runtime.entities.model_entities import ModelType
from libs.datetime_utils import naive_utc_now
from models.dataset import Dataset, DatasetProcessRule, DocumentSegment
from models.dataset import Document as DatasetDocument
from models.enums import SegmentStatus
from models.model import Account
# ============================================================================
@@ -611,7 +612,7 @@ class TestIndexingRunnerLoad:
- Keyword index creation
- Multi-threaded processing
- Document segment status updates
- Token counting
- Precomputed token totals
- Error handling during loading
"""
@@ -677,16 +678,10 @@ class TestIndexingRunnerLoad:
"""Test loading with high quality indexing (vector embeddings)."""
# Arrange
runner = IndexingRunner()
mock_embedding_instance = MagicMock()
mock_embedding_instance.get_text_embedding_num_tokens.return_value = 100
model_manager = mock_dependencies["model_manager"].return_value
model_manager.get_model_instance.return_value = mock_embedding_instance
mock_processor = MagicMock()
# Mock ThreadPoolExecutor
mock_future = MagicMock()
mock_future.result.return_value = 300 # Total tokens
mock_future.result.return_value = None
mock_executor_instance = MagicMock()
mock_executor_instance.__enter__.return_value = mock_executor_instance
mock_executor_instance.__exit__.return_value = None
@@ -694,20 +689,51 @@ class TestIndexingRunnerLoad:
mock_dependencies["executor"].return_value = mock_executor_instance
# Mock update_document_index_status to avoid database calls
with patch.object(runner, "_update_document_index_status"):
with patch.object(runner, "_update_document_index_status") as mock_update_status:
# Act
runner._load(
mock_processor,
sample_dataset,
sample_dataset_document,
sample_documents,
mock_dependencies["session"],
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
total_tokens=300,
)
# Assert
model_manager.get_model_instance.assert_called_once()
mock_dependencies["model_manager"].assert_not_called()
# Verify executor was used for parallel processing
assert mock_executor_instance.submit.called
for submit_call in mock_executor_instance.submit.call_args_list:
assert submit_call.args[0] == runner._process_chunk
assert len(submit_call.args) == 6
mock_future.result.assert_called()
assert mock_update_status.call_args.kwargs["extra_update_params"][DatasetDocument.tokens] == 300
def test_load_propagates_worker_errors(
self, mock_dependencies, sample_dataset, sample_dataset_document, sample_documents
):
runner = IndexingRunner()
mock_future = MagicMock()
mock_future.result.side_effect = RuntimeError("index failed")
mock_executor_instance = MagicMock()
mock_executor_instance.__enter__.return_value = mock_executor_instance
mock_executor_instance.__exit__.return_value = None
mock_executor_instance.submit.return_value = mock_future
mock_dependencies["executor"].return_value = mock_executor_instance
with (
patch.object(runner, "_update_document_index_status") as mock_update_status,
pytest.raises(RuntimeError, match="index failed"),
):
runner._load(
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
total_tokens=300,
)
mock_update_status.assert_not_called()
def test_load_with_economy_indexing(
self, mock_dependencies, sample_dataset, sample_dataset_document, sample_documents
@@ -717,8 +743,6 @@ class TestIndexingRunnerLoad:
runner = IndexingRunner()
sample_dataset.indexing_technique = IndexTechniqueType.ECONOMY
mock_processor = MagicMock()
# Mock thread for keyword indexing
mock_thread_instance = MagicMock()
mock_thread_instance.join = MagicMock()
@@ -728,11 +752,11 @@ class TestIndexingRunnerLoad:
with patch.object(runner, "_update_document_index_status"):
# Act
runner._load(
mock_processor,
sample_dataset,
sample_dataset_document,
sample_documents,
mock_dependencies["session"],
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
total_tokens=0,
)
# Assert
@@ -759,16 +783,9 @@ class TestIndexingRunnerLoad:
)
]
mock_embedding_instance = MagicMock()
mock_embedding_instance.get_text_embedding_num_tokens.return_value = 50
model_manager = mock_dependencies["model_manager"].return_value
model_manager.get_model_instance.return_value = mock_embedding_instance
mock_processor = MagicMock()
# Mock ThreadPoolExecutor
mock_future = MagicMock()
mock_future.result.return_value = 150
mock_future.result.return_value = None
mock_executor_instance = MagicMock()
mock_executor_instance.__enter__.return_value = mock_executor_instance
mock_executor_instance.__exit__.return_value = None
@@ -779,14 +796,15 @@ class TestIndexingRunnerLoad:
with patch.object(runner, "_update_document_index_status"):
# Act
runner._load(
mock_processor,
sample_dataset,
sample_dataset_document,
sample_documents,
mock_dependencies["session"],
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
total_tokens=150,
)
# Assert
mock_dependencies["model_manager"].assert_not_called()
# Verify no keyword thread for parent-child index
mock_dependencies["thread"].assert_not_called()
@@ -850,6 +868,7 @@ class TestIndexingRunnerRun:
segment.index_node_hash = "parent-hash"
segment.document_id = dataset_document.id
segment.dataset_id = dataset_document.dataset_id
segment.tokens = 12
segment.get_child_chunks.return_value = [
SimpleNamespace(content="child", index_node_id="child-node", index_node_hash="child-hash")
]
@@ -862,6 +881,32 @@ class TestIndexingRunnerRun:
segment.get_child_chunks.assert_called_once_with(session=session)
assert load.call_args.kwargs["documents"][0].children[0].page_content == "child"
assert load.call_args.kwargs["total_tokens"] == 12
def test_run_in_indexing_status_uses_tokens_from_all_segments(self, mock_dependencies, sample_dataset_documents):
runner = IndexingRunner()
dataset_document = sample_dataset_documents[0]
dataset = Mock(spec=Dataset)
completed_segment = Mock(spec=DocumentSegment)
completed_segment.status = SegmentStatus.COMPLETED
completed_segment.tokens = 10
incomplete_segment = Mock(spec=DocumentSegment)
incomplete_segment.status = SegmentStatus.WAITING
incomplete_segment.tokens = 20
incomplete_segment.content = "pending"
incomplete_segment.index_node_id = "pending-node"
incomplete_segment.index_node_hash = "pending-hash"
incomplete_segment.document_id = dataset_document.id
incomplete_segment.dataset_id = dataset_document.dataset_id
session = mock_dependencies["session"]
session.get.side_effect = lambda model, _: dataset_document if model is DatasetDocument else dataset
session.scalars.return_value.all.return_value = [completed_segment, incomplete_segment]
with patch.object(runner, "_load") as load:
runner.run_in_indexing_status(dataset_document, session)
assert load.call_args.kwargs["documents"][0].page_content == "pending"
assert load.call_args.kwargs["total_tokens"] == 30
def test_run_success_single_document(self, mock_dependencies, sample_dataset_documents):
"""Test successful run with single document."""
@@ -953,6 +998,98 @@ class TestIndexingRunnerRun:
with pytest.raises(DocumentIsPausedError):
runner.run([doc], mock_dependencies["session"])
def test_run_counts_each_transformed_document_once(self, mock_dependencies, sample_dataset_documents):
runner = IndexingRunner()
dataset_document = sample_dataset_documents[0]
dataset = Mock(spec=Dataset)
dataset.id = dataset_document.dataset_id
dataset.tenant_id = dataset_document.tenant_id
dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY
current_user = Mock(spec=Account)
transformed_documents = [
Document(page_content="first", metadata={"doc_id": "first", "doc_hash": "hash-first"}),
Document(page_content="second", metadata={"doc_id": "second", "doc_hash": "hash-second"}),
]
model_dispatch = {
DatasetDocument: dataset_document,
Dataset: dataset,
Account: current_user,
}
mock_dependencies["session"].get.side_effect = lambda model, _: model_dispatch.get(model)
process_rule = Mock(spec=DatasetProcessRule)
process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}}
mock_dependencies["session"].scalar.return_value = process_rule
with (
patch.object(runner, "_extract", return_value=[Document(page_content="source", metadata={})]),
patch.object(runner, "_transform", return_value=transformed_documents),
patch.object(runner, "_load_segments") as load_segments,
patch.object(runner, "_load") as load,
patch(
"core.indexing_runner.calculate_segment_token_counts",
return_value=[11, 22],
) as calculate_token_counts,
):
runner.run([dataset_document], mock_dependencies["session"])
calculate_token_counts.assert_called_once_with(dataset=dataset, documents=transformed_documents)
load_segments.assert_called_once_with(
session=mock_dependencies["session"],
dataset=dataset,
dataset_document=dataset_document,
documents=transformed_documents,
token_counts=[11, 22],
)
assert load.call_args.kwargs["total_tokens"] == 33
def test_run_in_splitting_status_counts_each_transformed_document_once(
self, mock_dependencies, sample_dataset_documents
):
runner = IndexingRunner()
dataset_document = sample_dataset_documents[0]
dataset_document.created_by = "user-1"
dataset = Mock(spec=Dataset)
dataset.id = dataset_document.dataset_id
dataset.tenant_id = dataset_document.tenant_id
dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY
current_user = Mock(spec=Account)
transformed_documents = [
Document(page_content="first", metadata={"doc_id": "first", "doc_hash": "hash-first"}),
Document(page_content="second", metadata={"doc_id": "second", "doc_hash": "hash-second"}),
]
model_dispatch = {
DatasetDocument: dataset_document,
Dataset: dataset,
Account: current_user,
}
mock_dependencies["session"].get.side_effect = lambda model, _: model_dispatch.get(model)
mock_dependencies["session"].scalars.return_value.all.return_value = []
process_rule = Mock(spec=DatasetProcessRule)
process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}}
mock_dependencies["session"].scalar.return_value = process_rule
with (
patch.object(runner, "_extract", return_value=[Document(page_content="source", metadata={})]),
patch.object(runner, "_transform", return_value=transformed_documents),
patch.object(runner, "_load_segments") as load_segments,
patch.object(runner, "_load") as load,
patch(
"core.indexing_runner.calculate_segment_token_counts",
return_value=[11, 22],
) as calculate_token_counts,
):
runner.run_in_splitting_status(dataset_document, mock_dependencies["session"])
calculate_token_counts.assert_called_once_with(dataset=dataset, documents=transformed_documents)
load_segments.assert_called_once_with(
session=mock_dependencies["session"],
dataset=dataset,
dataset_document=dataset_document,
documents=transformed_documents,
token_counts=[11, 22],
)
assert load.call_args.kwargs["total_tokens"] == 33
def test_run_handles_provider_token_error(self, mock_dependencies, sample_dataset_documents):
"""Test run handles ProviderTokenNotInitError and updates document status."""
# Arrange
@@ -1395,7 +1532,11 @@ class TestIndexingRunnerLoadSegments:
):
# Act
runner._load_segments(
sample_dataset, sample_dataset_document, sample_documents, mock_dependencies["session"]
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
token_counts=[10, 20],
)
# Assert
@@ -1405,7 +1546,10 @@ class TestIndexingRunnerLoadSegments:
document_id=sample_dataset_document.id,
)
mock_docstore_instance.add_documents.assert_called_once_with(
docs=sample_documents, save_child=False, session=mock_dependencies["session"]
session=mock_dependencies["session"],
docs=sample_documents,
save_child=False,
token_counts=[10, 20],
)
def test_load_segments_parent_child_index(
@@ -1435,12 +1579,19 @@ class TestIndexingRunnerLoadSegments:
):
# Act
runner._load_segments(
sample_dataset, sample_dataset_document, sample_documents, mock_dependencies["session"]
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
token_counts=[10, 20],
)
# Assert
mock_docstore_instance.add_documents.assert_called_once_with(
docs=sample_documents, save_child=True, session=mock_dependencies["session"]
session=mock_dependencies["session"],
docs=sample_documents,
save_child=True,
token_counts=[10, 20],
)
def test_load_segments_updates_word_count(
@@ -1462,7 +1613,11 @@ class TestIndexingRunnerLoadSegments:
):
# Act
runner._load_segments(
sample_dataset, sample_dataset_document, sample_documents, mock_dependencies["session"]
session=mock_dependencies["session"],
dataset=sample_dataset,
dataset_document=sample_dataset_document,
documents=sample_documents,
token_counts=[10, 20],
)
# Assert
@@ -1565,7 +1720,6 @@ class TestIndexingRunnerProcessChunk:
"""Unit tests for chunk processing in parallel.
Tests cover:
- Token counting
- Vector index creation
- Segment status updates
- Pause detection during processing
@@ -1590,16 +1744,12 @@ class TestIndexingRunnerProcessChunk:
app.app_context.return_value.__exit__ = MagicMock()
return app
def test_process_chunk_counts_tokens(self, mock_dependencies, mock_flask_app):
"""Test process chunk correctly counts tokens."""
def test_process_chunk_loads_index_and_completes_segments(self, mock_dependencies, mock_flask_app):
"""Test process chunk loads the index and completes segments without counting tokens."""
# Arrange
from core.indexing_runner import IndexingRunner
runner = IndexingRunner()
mock_embedding_instance = MagicMock()
# Mock to return an iterable that sums to 150 tokens
mock_embedding_instance.get_text_embedding_num_tokens.return_value = [75, 75]
mock_processor = MagicMock()
chunk_documents = [
Document(page_content="Chunk 1", metadata={"doc_id": "c1"}),
@@ -1638,18 +1788,19 @@ class TestIndexingRunnerProcessChunk:
mock_factory.return_value.init_index_processor.return_value = mock_processor
# Act - the method creates its own app_context and session
tokens = runner._process_chunk(
result = runner._process_chunk(
mock_flask_app,
IndexStructureType.PARAGRAPH_INDEX,
chunk_documents,
mock_dataset.id,
mock_dataset_document.id,
mock_embedding_instance,
)
# Assert
assert tokens == 150
assert result is None
mock_processor.load.assert_called_once()
mock_dependencies["session"].execute.assert_called_once()
mock_dependencies["session"].commit.assert_called_once()
def test_process_chunk_detects_pause(self, mock_dependencies, mock_flask_app):
"""Test process chunk detects document pause."""
@@ -1657,8 +1808,6 @@ class TestIndexingRunnerProcessChunk:
from core.indexing_runner import IndexingRunner
runner = IndexingRunner()
mock_embedding_instance = MagicMock()
mock_processor = MagicMock()
chunk_documents = [Document(page_content="Chunk", metadata={"doc_id": "c1"})]
mock_dataset = Mock(spec=Dataset)
@@ -1691,5 +1840,4 @@ class TestIndexingRunnerProcessChunk:
chunk_documents,
mock_dataset.id,
mock_dataset_document.id,
mock_embedding_instance,
)
@@ -3752,8 +3752,8 @@ class TestKnowledgeRetrievalRegression:
"""
Repro test for current bug:
reranking runs after `with flask_app.app_context():` exits.
`_multiple_retrieve_thread` catches exceptions and stores them into `thread_exceptions`,
so we must assert from that list (not from an outer try/except).
The outer thread entry point catches exceptions from the traced retrieval method
and stores them in `thread_exceptions`.
"""
dataset_retrieval = DatasetRetrieval()
flask_app = Flask(__name__)
@@ -3806,7 +3806,6 @@ class TestKnowledgeRetrievalRegression:
# output list from _multiple_retrieve_thread
all_documents: list[Document] = []
# IMPORTANT: _multiple_retrieve_thread swallows exceptions and appends them here
thread_exceptions: list[Exception] = []
def target():
@@ -3818,7 +3817,7 @@ class TestKnowledgeRetrievalRegression:
),
_patched_retriever_session(),
):
dataset_retrieval._multiple_retrieve_thread(
dataset_retrieval._multiple_retrieve_thread_safely(
flask_app=flask_app,
available_datasets=[mock_dataset, secondary_dataset],
metadata_condition=None,
@@ -3847,7 +3846,6 @@ class TestKnowledgeRetrievalRegression:
# Ensure reranking branch was actually executed
assert called["init"] >= 1, "DataPostProcessor was never constructed; reranking branch may not have run."
# Current buggy code should record an exception (not raise it)
assert not thread_exceptions, thread_exceptions
def test_run_retriever_thread_provides_session_to_retriever(self):
@@ -3865,14 +3863,12 @@ class TestKnowledgeRetrievalRegression:
document_ids_filter=None,
metadata_condition=None,
attachment_ids=None,
cancel_event=None,
thread_exceptions=[],
)
mock_retriever.assert_called_once()
assert mock_retriever.call_args.kwargs["session"] is session
def test_run_retriever_thread_records_retriever_exception(self):
def test_run_retriever_thread_safely_records_retriever_exception(self):
dataset_retrieval = DatasetRetrieval()
all_documents: list[Document] = []
cancel_event = threading.Event()
@@ -3881,7 +3877,7 @@ class TestKnowledgeRetrievalRegression:
with _patched_retriever_session():
with patch.object(dataset_retrieval, "_retriever", side_effect=expected_error):
dataset_retrieval._run_retriever_thread(
dataset_retrieval._run_retriever_thread_safely(
flask_app=_FakeFlaskApp(),
dataset_id="dataset-1",
query="test query",
@@ -5139,7 +5135,7 @@ class TestSingleAndMultipleRetrieveCoverage:
app = Flask(__name__)
def failing_thread(**kwargs):
kwargs["thread_exceptions"].append(RuntimeError("thread boom"))
raise RuntimeError("thread boom")
with app.app_context():
with (
@@ -1,10 +1,21 @@
import pytest
from sqlalchemy.orm import Session
from core.workflow.nodes.agent_v2.binding_resolver import (
WorkflowAgentBindingError,
WorkflowAgentBindingResolver,
)
from models.agent import Agent, AgentConfigSnapshot, AgentStatus, WorkflowAgentBindingType, WorkflowAgentNodeBinding
from models.agent import (
Agent,
AgentConfigRevision,
AgentConfigRevisionOperation,
AgentConfigSnapshot,
AgentScope,
AgentSource,
AgentStatus,
WorkflowAgentBindingType,
WorkflowAgentNodeBinding,
)
from models.agent_config_entities import AgentSoulConfig, AgentSoulModelConfig, WorkflowNodeJobConfig
@@ -104,6 +115,89 @@ def test_binding_resolver_uses_active_snapshot_for_roster_agent(monkeypatch: pyt
assert bundle.snapshot.id == "active-snapshot"
def test_binding_resolver_rejects_unpublished_roster_agent(monkeypatch: pytest.MonkeyPatch):
binding = _binding()
binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT
monkeypatch.setattr(
"core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session",
lambda: FakeSession([binding, None]),
)
with pytest.raises(WorkflowAgentBindingError) as exc_info:
WorkflowAgentBindingResolver().resolve(**_resolve())
assert exc_info.value.error_code == "agent_not_available"
assert "not been published" in str(exc_info.value)
@pytest.mark.parametrize(
"sqlite_session",
[(Agent, AgentConfigSnapshot, AgentConfigRevision, WorkflowAgentNodeBinding)],
indirect=True,
)
def test_binding_resolver_requires_publish_provenance_for_active_roster_snapshot(
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
) -> None:
binding = _binding()
binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT
binding.workflow_version = "draft"
agent = Agent(
id="agent-1",
tenant_id="tenant-1",
name="Imported Agent",
scope=AgentScope.ROSTER,
source=AgentSource.IMPORTED,
app_id="agent-app-1",
status=AgentStatus.ACTIVE,
active_config_snapshot_id="snapshot-1",
active_config_has_model=True,
# Dirty draft state must not hide a snapshot after it has publish provenance.
active_config_is_published=False,
)
sqlite_session.add_all(
[
binding,
agent,
_snapshot(),
AgentConfigRevision(
id="revision-import",
tenant_id="tenant-1",
agent_id="agent-1",
current_snapshot_id="snapshot-1",
revision=1,
operation=AgentConfigRevisionOperation.IMPORT_PACKAGE,
),
]
)
sqlite_session.commit()
monkeypatch.setattr(
"core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session",
lambda: sqlite_session,
)
with pytest.raises(WorkflowAgentBindingError) as exc_info:
WorkflowAgentBindingResolver().resolve(**_resolve())
assert exc_info.value.error_code == "agent_not_available"
sqlite_session.add(
AgentConfigRevision(
id="revision-publish",
tenant_id="tenant-1",
agent_id="agent-1",
current_snapshot_id="snapshot-1",
revision=2,
operation=AgentConfigRevisionOperation.PUBLISH_DRAFT,
)
)
sqlite_session.commit()
bundle = WorkflowAgentBindingResolver().resolve(**_resolve())
assert bundle.agent.id == agent.id
assert bundle.snapshot.id == "snapshot-1"
def test_binding_resolver_raises_when_binding_missing(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(
"core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session",
@@ -156,6 +156,19 @@ def test_publish_validation_uses_active_snapshot_for_roster_agent():
)
def test_publish_validation_rejects_unpublished_roster_agent():
binding = _binding(WorkflowNodeJobConfig())
binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT
session = Mock()
session.scalar.side_effect = [binding, None]
with pytest.raises(WorkflowAgentNodeValidationError, match="unpublished roster agent"):
WorkflowAgentNodeValidator.validate_published_workflow(
session=session,
workflow=_workflow(_graph([{"source": "start", "target": "agent-node"}])),
)
def test_publish_validation_rejects_non_upstream_previous_output_ref():
node_job = WorkflowNodeJobConfig.model_validate(
{"previous_node_output_refs": [{"node_id": "later-node", "output": "text"}]}

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