From 444d164aea4dde9acbba8a30e6dfe53337bd8b36 Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Tue, 30 Jun 2026 14:25:20 +0800 Subject: [PATCH 01/54] test(e2e): add agent v2 test infrastructure (#38191) --- e2e/AGENTS.md | 10 + e2e/features/agent-v2/configure-entry.feature | 9 + .../agent-v2/configure.steps.ts | 49 ++++ e2e/features/support/hooks.ts | 9 +- e2e/features/support/world.ts | 6 + e2e/scripts/common.ts | 2 + e2e/support/agent.ts | 229 ++++++++++++++++++ e2e/support/api.ts | 39 ++- e2e/support/process.ts | 22 +- eslint-suppressions.json | 18 -- 10 files changed, 354 insertions(+), 39 deletions(-) create mode 100644 e2e/features/agent-v2/configure-entry.feature create mode 100644 e2e/features/step-definitions/agent-v2/configure.steps.ts create mode 100644 e2e/support/agent.ts diff --git a/e2e/AGENTS.md b/e2e/AGENTS.md index c05b5105be6..017594bbe34 100644 --- a/e2e/AGENTS.md +++ b/e2e/AGENTS.md @@ -294,3 +294,13 @@ Or browse the step definition files directly: - `features/step-definitions/common/` — auth guards and navigation assertions shared by all features - `features/step-definitions//` — domain-specific steps scoped to a single feature area + +## Agent v2 scenarios + +Agent v2 scenarios live under `features/agent-v2/` and use the `@agent-v2` capability tag. + +The E2E web environment enables Agent v2 through `NEXT_PUBLIC_ENABLE_AGENT_V2=true` in `scripts/common.ts`, because `/roster` routes are guarded by that feature flag. + +Use `support/agent.ts` for Agent v2 API fixtures. It owns roster-shaped Agent IDs, configure/access route helpers, composer draft sync, build-draft helpers, publish, API access toggles, and Agent cleanup. Store created roster Agent IDs in `DifyWorld.createdAgentIds`; the shared `After` hook deletes them after each scenario. + +Keep Agent v2 step definitions under `features/step-definitions/agent-v2/`. Prefer API setup for prerequisite state, then use Playwright only for user-observable navigation, editing, and assertions. diff --git a/e2e/features/agent-v2/configure-entry.feature b/e2e/features/agent-v2/configure-entry.feature new file mode 100644 index 00000000000..5c15c84f24d --- /dev/null +++ b/e2e/features/agent-v2/configure-entry.feature @@ -0,0 +1,9 @@ +@agent-v2 @authenticated @infra +Feature: Agent v2 configure entry + Scenario: Open the configure page for an Agent v2 test agent + Given I am signed in as the default E2E admin + And an Agent v2 test agent has been created via API + And a minimal Agent v2 composer draft has been synced + When I open the Agent v2 configure page + Then I should be on the Agent v2 configure page + And I should see the Agent v2 configure workspace diff --git a/e2e/features/step-definitions/agent-v2/configure.steps.ts b/e2e/features/step-definitions/agent-v2/configure.steps.ts new file mode 100644 index 00000000000..cbb49fef685 --- /dev/null +++ b/e2e/features/step-definitions/agent-v2/configure.steps.ts @@ -0,0 +1,49 @@ +import type { DifyWorld } from '../../support/world' +import { Given, Then, When } from '@cucumber/cucumber' +import { expect } from '@playwright/test' +import { + createTestAgent, + getAgentConfigurePath, + saveAgentComposerDraft, +} from '../../../support/agent' + +Given('an Agent v2 test agent has been created via API', async function (this: DifyWorld) { + const agent = await createTestAgent() + this.createdAgentIds.push(agent.id) + this.lastCreatedAgentName = agent.name + this.lastCreatedAgentRole = agent.role +}) + +Given('a minimal Agent v2 composer draft has been synced', async function (this: DifyWorld) { + const agentId = this.createdAgentIds.at(-1) + if (!agentId) + throw new Error('No Agent v2 ID found. Create an Agent v2 test agent first.') + + await saveAgentComposerDraft(agentId) +}) + +When('I open the Agent v2 configure page', async function (this: DifyWorld) { + const agentId = this.createdAgentIds.at(-1) + if (!agentId) + throw new Error('No Agent v2 ID found. Create an Agent v2 test agent first.') + + await this.getPage().goto(getAgentConfigurePath(agentId)) +}) + +Then('I should be on the Agent v2 configure page', async function (this: DifyWorld) { + const agentId = this.createdAgentIds.at(-1) + if (!agentId) + throw new Error('No Agent v2 ID found. Create an Agent v2 test agent first.') + + await expect(this.getPage()).toHaveURL( + new RegExp(`/roster/agent/${agentId}/configure(?:\\?.*)?$`), + ) +}) + +Then('I should see the Agent v2 configure workspace', async function (this: DifyWorld) { + const page = this.getPage() + + await expect(page.getByRole('region', { name: 'Configure' })).toBeVisible({ timeout: 30_000 }) + await expect(page.getByRole('heading', { name: 'Configure' })).toBeVisible() + await expect(page.getByText(this.lastCreatedAgentName!)).toBeVisible() +}) diff --git a/e2e/features/support/hooks.ts b/e2e/features/support/hooks.ts index c1a535ee2c0..e67291a7ad2 100644 --- a/e2e/features/support/hooks.ts +++ b/e2e/features/support/hooks.ts @@ -1,4 +1,5 @@ import type { Browser } from '@playwright/test' +import type { Buffer } from 'node:buffer' import type { DifyWorld } from './world' import { mkdir, writeFile } from 'node:fs/promises' import path from 'node:path' @@ -6,6 +7,7 @@ import { fileURLToPath } from 'node:url' import { After, AfterAll, Before, BeforeAll, setDefaultTimeout, Status } from '@cucumber/cucumber' import { chromium } from '@playwright/test' import { AUTH_BOOTSTRAP_TIMEOUT_MS, ensureAuthenticatedState } from '../../fixtures/auth' +import { deleteTestAgent } from '../../support/agent' import { deleteTestApp } from '../../support/api' import { baseURL, cucumberHeadless, cucumberSlowMo } from '../../test-env' @@ -41,7 +43,7 @@ BeforeAll({ timeout: AUTH_BOOTSTRAP_TIMEOUT_MS }, async () => { slowMo: cucumberSlowMo, }) - console.log(`[e2e] session cache bootstrap against ${baseURL}`) + console.warn(`[e2e] session cache bootstrap against ${baseURL}`) await ensureAuthenticatedState(browser, baseURL) }) @@ -58,7 +60,7 @@ Before(async function (this: DifyWorld, { pickle }) { this.scenarioStartedAt = Date.now() const tags = pickle.tags.map(tag => tag.name).join(' ') - console.log(`[e2e] start ${pickle.name}${tags ? ` ${tags}` : ''}`) + console.warn(`[e2e] start ${pickle.name}${tags ? ` ${tags}` : ''}`) }) After(async function (this: DifyWorld, { pickle, result }) { @@ -85,10 +87,11 @@ After(async function (this: DifyWorld, { pickle, result }) { } const status = result?.status || 'UNKNOWN' - console.log( + console.warn( `[e2e] end ${pickle.name} status=${status}${elapsedMs ? ` durationMs=${elapsedMs}` : ''}`, ) + for (const id of this.createdAgentIds) await deleteTestAgent(id).catch(() => {}) for (const id of this.createdAppIds) await deleteTestApp(id).catch(() => {}) await this.closeSession() diff --git a/e2e/features/support/world.ts b/e2e/features/support/world.ts index b53087171f5..6faf740a7c9 100644 --- a/e2e/features/support/world.ts +++ b/e2e/features/support/world.ts @@ -13,7 +13,10 @@ export class DifyWorld extends World { scenarioStartedAt: number | undefined session: AuthSessionMetadata | undefined lastCreatedAppName: string | undefined + lastCreatedAgentName: string | undefined + lastCreatedAgentRole: string | undefined createdAppIds: string[] = [] + createdAgentIds: string[] = [] capturedDownloads: Download[] = [] shareURL: string | undefined @@ -26,7 +29,10 @@ export class DifyWorld extends World { this.consoleErrors = [] this.pageErrors = [] this.lastCreatedAppName = undefined + this.lastCreatedAgentName = undefined + this.lastCreatedAgentRole = undefined this.createdAppIds = [] + this.createdAgentIds = [] this.capturedDownloads = [] this.shareURL = undefined } diff --git a/e2e/scripts/common.ts b/e2e/scripts/common.ts index 2964892dd0b..20b9ab7fbd8 100644 --- a/e2e/scripts/common.ts +++ b/e2e/scripts/common.ts @@ -1,3 +1,4 @@ +import type { Buffer } from 'node:buffer' import type { ChildProcess } from 'node:child_process' import { spawn } from 'node:child_process' import { createHash } from 'node:crypto' @@ -42,6 +43,7 @@ export const webEnvExampleFile = path.join(webDir, '.env.example') export const apiEnvExampleFile = path.join(apiDir, 'tests', 'integration_tests', '.env.example') export const e2eWebEnvOverrides = { NEXT_PUBLIC_API_PREFIX: 'http://127.0.0.1:5001/console/api', + NEXT_PUBLIC_ENABLE_AGENT_V2: 'true', NEXT_PUBLIC_PUBLIC_API_PREFIX: 'http://127.0.0.1:5001/api', } satisfies Record diff --git a/e2e/support/agent.ts b/e2e/support/agent.ts new file mode 100644 index 00000000000..445c10b5cfc --- /dev/null +++ b/e2e/support/agent.ts @@ -0,0 +1,229 @@ +import { createApiContext, expectApiResponseOK, setAppSiteEnabled } from './api' + +export type AgentSeed = { + app_id?: string + backing_app_id?: string + description?: string + enable_site?: boolean + id: string + name: string + role?: string + site?: { + access_token?: string | null + app_base_url?: string | null + code?: string | null + } | null +} + +export type AgentSoulConfig = Record + +export type AgentComposerResponse = { + agent_soul?: AgentSoulConfig +} + +export type AgentBuildDraftResponse = { + agent_soul: AgentSoulConfig + draft: Record + variant: 'agent_app' +} + +export type AgentApiAccess = { + api_key_count: number + api_reference_url: string + endpoint: string + enabled: boolean + files_upload_endpoint: string +} + +export type AgentApiKey = { + id: string + token?: string +} + +export const defaultAgentSoulConfig: AgentSoulConfig = { + prompt: { + system_prompt: 'You are a Dify Agent E2E test assistant.', + }, +} + +export const getAgentConfigurePath = (agentId: string) => `/roster/agent/${agentId}/configure` +export const getAgentAccessPath = (agentId: string) => `/roster/agent/${agentId}/access` + +export async function createTestAgent({ + description = 'Created by Dify E2E.', + name = `E2E Agent ${Date.now()}`, + role = 'E2E test assistant', +}: { + description?: string + name?: string + role?: string +} = {}): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.post('/console/api/agent', { + data: { + description, + icon: '🤖', + icon_background: '#FFEAD5', + icon_type: 'emoji', + name, + role, + }, + }) + await expectApiResponseOK(response, 'Create Agent v2 test agent') + return (await response.json()) as AgentSeed + } + finally { + await ctx.dispose() + } +} + +export async function getTestAgent(agentId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.get(`/console/api/agent/${agentId}`) + await expectApiResponseOK(response, `Get Agent v2 test agent ${agentId}`) + return (await response.json()) as AgentSeed + } + finally { + await ctx.dispose() + } +} + +export async function deleteTestAgent(agentId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.delete(`/console/api/agent/${agentId}`) + await expectApiResponseOK(response, `Delete Agent v2 test agent ${agentId}`) + } + finally { + await ctx.dispose() + } +} + +export async function saveAgentComposerDraft( + agentId: string, + agentSoul: AgentSoulConfig = defaultAgentSoulConfig, +): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.put(`/console/api/agent/${agentId}/composer`, { + data: { + agent_soul: agentSoul, + save_strategy: 'save_to_current_version', + variant: 'agent_app', + }, + }) + await expectApiResponseOK(response, `Save Agent v2 composer draft for ${agentId}`) + return (await response.json()) as AgentComposerResponse + } + finally { + await ctx.dispose() + } +} + +export async function checkoutAgentBuildDraft(agentId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.post(`/console/api/agent/${agentId}/build-draft/checkout`, { + data: { force: true }, + }) + await expectApiResponseOK(response, `Checkout Agent v2 build draft for ${agentId}`) + return (await response.json()) as AgentBuildDraftResponse + } + finally { + await ctx.dispose() + } +} + +export async function discardAgentBuildDraft(agentId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.delete(`/console/api/agent/${agentId}/build-draft`) + await expectApiResponseOK(response, `Discard Agent v2 build draft for ${agentId}`) + } + finally { + await ctx.dispose() + } +} + +export async function publishAgent(agentId: string, versionNote = 'E2E publish'): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.post(`/console/api/agent/${agentId}/publish`, { + data: { version_note: versionNote }, + }) + await expectApiResponseOK(response, `Publish Agent v2 test agent ${agentId}`) + } + finally { + await ctx.dispose() + } +} + +export async function enableAgentSiteAndGetURL(agentId: string): Promise { + const agent = await getTestAgent(agentId) + const appId = agent.app_id ?? agent.backing_app_id + if (!appId) + throw new Error(`Agent v2 ${agentId} does not expose a backing app ID.`) + + const appDetail = await setAppSiteEnabled(appId, true) + const token = agent.site?.access_token ?? agent.site?.code ?? appDetail.site.access_token + const baseURL = agent.site?.app_base_url ?? appDetail.site.app_base_url + + return `${baseURL.replace(/\/$/, '')}/agent/${token}` +} + +export async function getAgentApiAccess(agentId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.get(`/console/api/agent/${agentId}/api-access`) + await expectApiResponseOK(response, `Get Agent v2 API access for ${agentId}`) + return (await response.json()) as AgentApiAccess + } + finally { + await ctx.dispose() + } +} + +export async function setAgentApiAccess( + agentId: string, + enabled: boolean, +): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.post(`/console/api/agent/${agentId}/api-enable`, { + data: { enable_api: enabled }, + }) + await expectApiResponseOK( + response, + `${enabled ? 'Enable' : 'Disable'} Agent v2 API access for ${agentId}`, + ) + return (await response.json()) as AgentApiAccess + } + finally { + await ctx.dispose() + } +} + +export async function createAgentApiKey(agentId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.post(`/console/api/agent/${agentId}/api-keys`) + await expectApiResponseOK(response, `Create Agent v2 API key for ${agentId}`) + return (await response.json()) as AgentApiKey + } + finally { + await ctx.dispose() + } +} + +export async function deleteAgentApiKey(agentId: string, apiKeyId: string): Promise { + const ctx = await createApiContext() + try { + const response = await ctx.delete(`/console/api/agent/${agentId}/api-keys/${apiKeyId}`) + await expectApiResponseOK(response, `Delete Agent v2 API key ${apiKeyId} for ${agentId}`) + } + finally { + await ctx.dispose() + } +} diff --git a/e2e/support/api.ts b/e2e/support/api.ts index 74c42d3e73f..81773dea595 100644 --- a/e2e/support/api.ts +++ b/e2e/support/api.ts @@ -1,3 +1,4 @@ +import type { APIResponse } from '@playwright/test' import { readFile } from 'node:fs/promises' import { request } from '@playwright/test' import { authStatePath } from '../fixtures/auth' @@ -7,7 +8,7 @@ type StorageState = { cookies: Array<{ name: string, value: string }> } -async function createApiContext() { +export async function createApiContext() { const state = JSON.parse(await readFile(authStatePath, 'utf8')) as StorageState const csrfToken = state.cookies.find(c => c.name.endsWith('csrf_token'))?.value ?? '' @@ -18,6 +19,14 @@ async function createApiContext() { }) } +export async function expectApiResponseOK(response: APIResponse, action: string): Promise { + if (response.ok()) + return + + const body = await response.text().catch(() => '') + throw new Error(`${action} failed with ${response.status()} ${response.statusText()}: ${body}`) +} + export type AppSeed = { id: string name: string @@ -141,20 +150,34 @@ export async function publishWorkflowApp(appId: string): Promise { } } -type AppDetailWithSite = { +export type AppDetailWithSite = { + mode?: string site: { access_token: string, app_base_url: string, enable_site: boolean } } +export function getAppSiteURL({ mode, site }: AppDetailWithSite): string { + const webAppMode = mode === 'completion' || mode === 'workflow' ? mode : 'chat' + return `${site.app_base_url}/${webAppMode}/${site.access_token}` +} + export async function enableAppSiteAndGetURL(appId: string): Promise { + return getAppSiteURL(await setAppSiteEnabled(appId, true)) +} + +export async function setAppSiteEnabled( + appId: string, + enabled: boolean, +): Promise { const ctx = await createApiContext() try { - await ctx.post(`/console/api/apps/${appId}/site-enable`, { - data: { enable_site: true }, + const enableResponse = await ctx.post(`/console/api/apps/${appId}/site-enable`, { + data: { enable_site: enabled }, }) - const res = await ctx.get(`/console/api/apps/${appId}`) - const body = (await res.json()) as AppDetailWithSite - const { app_base_url, access_token } = body.site - return `${app_base_url}/workflow/${access_token}` + await expectApiResponseOK(enableResponse, `${enabled ? 'Enable' : 'Disable'} app site ${appId}`) + + const detailResponse = await ctx.get(`/console/api/apps/${appId}`) + await expectApiResponseOK(detailResponse, `Get app site detail for ${appId}`) + return (await detailResponse.json()) as AppDetailWithSite } finally { await ctx.dispose() diff --git a/e2e/support/process.ts b/e2e/support/process.ts index 4de1161b08d..b54108e6855 100644 --- a/e2e/support/process.ts +++ b/e2e/support/process.ts @@ -122,21 +122,23 @@ const waitForProcessExit = (childProcess: ChildProcess, timeoutMs: number) => return } - const timeout = setTimeout(() => { - cleanup() - resolve() - }, timeoutMs) + let timeout: ReturnType - const onExit = () => { - cleanup() - resolve() - } - - const cleanup = () => { + function cleanup() { clearTimeout(timeout) childProcess.off('exit', onExit) } + function onExit() { + cleanup() + resolve() + } + + timeout = setTimeout(() => { + cleanup() + resolve() + }, timeoutMs) + childProcess.once('exit', onExit) }) diff --git a/eslint-suppressions.json b/eslint-suppressions.json index eb99f899a63..4aeb3f71088 100644 --- a/eslint-suppressions.json +++ b/eslint-suppressions.json @@ -1,22 +1,4 @@ { - "e2e/features/support/hooks.ts": { - "no-console": { - "count": 3 - }, - "node/prefer-global/buffer": { - "count": 1 - } - }, - "e2e/scripts/common.ts": { - "node/prefer-global/buffer": { - "count": 2 - } - }, - "e2e/support/process.ts": { - "ts/no-use-before-define": { - "count": 2 - } - }, "packages/migrate-no-unchecked-indexed-access/src/no-unchecked-indexed-access/migrate.ts": { "no-console": { "count": 11 From 528bf95d1bc6b14694ceee22e7233f836b2c6fe3 Mon Sep 17 00:00:00 2001 From: wangxiaolei Date: Tue, 30 Jun 2026 14:43:46 +0800 Subject: [PATCH 02/54] feat: support dataset permission migrate to rbac (#38166) --- api/commands/__init__.py | 3 +- api/commands/data_migrate.py | 2 + api/commands/rbac.py | 228 +++++++++++++++++- api/controllers/console/workspace/rbac.py | 10 +- api/core/rbac/__init__.py | 4 +- api/core/rbac/entities.py | 8 + api/services/enterprise/rbac_service.py | 16 +- .../test_legacy_model_type_migration.py | 168 +++++++++++++ .../console/workspace/test_rbac.py | 2 +- 9 files changed, 420 insertions(+), 21 deletions(-) diff --git a/api/commands/__init__.py b/api/commands/__init__.py index e4207bea74e..1e40d3f7928 100644 --- a/api/commands/__init__.py +++ b/api/commands/__init__.py @@ -22,7 +22,7 @@ from .plugin import ( setup_system_trigger_oauth_client, transform_datasource_credentials, ) -from .rbac import migrate_member_roles_to_rbac +from .rbac import migrate_dataset_permissions_to_rbac, migrate_member_roles_to_rbac from .retention import ( archive_workflow_runs, archive_workflow_runs_plan, @@ -76,6 +76,7 @@ __all__ = [ "legacy_model_types", "migrate_annotation_vector_database", "migrate_data_for_plugin", + "migrate_dataset_permissions_to_rbac", "migrate_knowledge_vector_database", "migrate_member_roles_to_rbac", "migrate_oss", diff --git a/api/commands/data_migrate.py b/api/commands/data_migrate.py index 4a71d864b26..35c2be999b7 100644 --- a/api/commands/data_migrate.py +++ b/api/commands/data_migrate.py @@ -7,6 +7,7 @@ from typing import cast import click +from commands.rbac import migrate_dataset_permissions_to_rbac from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from services.legacy_model_type_migration import ( @@ -177,3 +178,4 @@ def legacy_model_types( data_migrate.add_command(legacy_model_types) +data_migrate.add_command(migrate_dataset_permissions_to_rbac) diff --git a/api/commands/rbac.py b/api/commands/rbac.py index 630304b67e3..0793d11cbb2 100644 --- a/api/commands/rbac.py +++ b/api/commands/rbac.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from collections.abc import Iterator from concurrent.futures import ThreadPoolExecutor, as_completed @@ -8,8 +9,11 @@ from sqlalchemy import select from configs import dify_config from core.db.session_factory import session_factory -from models import TenantAccountJoin, TenantAccountRole -from services.enterprise.rbac_service import ListOption, RBACService +from core.rbac import RBACResourceWhitelistScope +from models import Dataset, DatasetPermission, DatasetPermissionEnum, TenantAccountJoin, TenantAccountRole +from services.enterprise.rbac_service import ListOption, RBACService, ReplaceMemberBindings, ReplaceUserAccessPolicies + +_RBAC_DEFAULT_ACCESS_POLICY_ID = "default" _LEGACY_ROLE_TO_BUILTIN_TAG = { TenantAccountRole.OWNER.value: "owner", @@ -258,3 +262,223 @@ def migrate_member_roles_to_rbac( fg="green", ) ) + + +def _dataset_permission_enum(permission: DatasetPermissionEnum | str | None) -> DatasetPermissionEnum: + if permission is None: + return DatasetPermissionEnum.ONLY_ME + try: + return DatasetPermissionEnum(permission) + except ValueError as exc: + raise ValueError(f"Unsupported legacy dataset permission: {permission}") from exc + + +def _rbac_dataset_scope_for_legacy_permission(permission: DatasetPermissionEnum) -> RBACResourceWhitelistScope: + if permission is DatasetPermissionEnum.ALL_TEAM: + return RBACResourceWhitelistScope.ALL + if permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.PARTIAL_TEAM}: + return RBACResourceWhitelistScope.SPECIFIC + raise ValueError(f"Unsupported legacy dataset permission: {permission}") + + +def _emit_dataset_permission_migration_event(payload: dict[str, object]) -> None: + click.echo(json.dumps(payload, sort_keys=True)) + + +@click.command( + "rbac-migrate-dataset-permissions", + help=( + "Migrate legacy dataset permission scopes and partial members into RBAC dataset access bindings. " + "Side effect: replacing each dataset whitelist clears existing per-user policy bindings; " + "the command then recreates legacy partial-member default bindings." + ), +) +@click.option("--tenant-id", help="Only migrate datasets in a single workspace.") +@click.option("--dataset-id", help="Only migrate a single dataset.") +@click.option("--batch-size", default=500, show_default=True, type=click.IntRange(min=1)) +@click.option( + "--dry-run/--apply", + default=True, + show_default=True, + help="Preview the migration without writing RBAC bindings. Use --apply to write changes.", +) +def migrate_dataset_permissions_to_rbac( + tenant_id: str | None, + dataset_id: str | None, + batch_size: int, + dry_run: bool, +) -> None: + """Backfill RBAC dataset access config from legacy `Dataset.permission`. + + Legacy mapping: + - all_team_members -> RBAC dataset whitelist scope "all" + - partial_members -> RBAC dataset whitelist scope "specific" plus each partial member gets the + virtual default policy + - only_me -> RBAC dataset whitelist scope "specific" with no member policy bindings + + The command replaces each dataset's RBAC whitelist scope first. RBAC clears + existing per-user policy bindings during that replace, then this command + recreates the legacy partial-member default bindings. Re-running it is + therefore idempotent for a dataset's current legacy configuration. + """ + click.echo(click.style("Starting RBAC dataset permission migration.", fg="green")) + + scanned_count = 0 + scope_migrated_count = 0 + user_policy_migrated_count = 0 + partial_dataset_count = 0 + + last_dataset_id: str | None = None + while True: + with session_factory.create_session() as session: + stmt = ( + select(Dataset.id, Dataset.tenant_id, Dataset.permission, Dataset.created_by) + .order_by(Dataset.id.asc()) + .limit(batch_size) + ) + if tenant_id: + stmt = stmt.where(Dataset.tenant_id == tenant_id) + if dataset_id: + stmt = stmt.where(Dataset.id == dataset_id) + if last_dataset_id: + stmt = stmt.where(Dataset.id > last_dataset_id) + + dataset_rows = list(session.execute(stmt).all()) + if not dataset_rows: + break + + dataset_ids = [str(row.id) for row in dataset_rows] + partial_members_by_dataset_id: dict[str, list[str]] = {item: [] for item in dataset_ids} + permission_rows = session.execute( + select(DatasetPermission.dataset_id, DatasetPermission.account_id).where( + DatasetPermission.dataset_id.in_(dataset_ids) + ) + ).all() + for row in permission_rows: + partial_members_by_dataset_id[str(row.dataset_id)].append(str(row.account_id)) + + for dataset in dataset_rows: + workspace_id = str(dataset.tenant_id) + current_dataset_id = str(dataset.id) + operator_account_id = str(dataset.created_by) + permission_value = _dataset_permission_enum(dataset.permission) + scope = _rbac_dataset_scope_for_legacy_permission(permission_value) + partial_member_ids = sorted(set(partial_members_by_dataset_id[current_dataset_id])) + should_bind_partial_members = permission_value is DatasetPermissionEnum.PARTIAL_TEAM + + click.echo( + f"tenant={workspace_id} dataset={current_dataset_id} " + f"operator={operator_account_id} " + f"legacy_permission={permission_value} -> rbac_scope={scope} " + f"partial_members={len(partial_member_ids) if should_bind_partial_members else 0}" + ) + + scanned_count += 1 + replace_whitelist_payload = ReplaceMemberBindings(scope=scope) + if dry_run: + _emit_dataset_permission_migration_event( + { + "event": "dataset_permission_migration_proposed_change", + "action": "replace_whitelist", + "dry_run": True, + "tenant_id": workspace_id, + "dataset_id": current_dataset_id, + "operator_account_id": operator_account_id, + "before": { + "legacy_dataset_permission": permission_value.value, + "legacy_partial_member_ids": partial_member_ids if should_bind_partial_members else [], + }, + "after": { + "rbac_whitelist_scope": scope.value, + }, + "call": { + "method": "RBACService.DatasetAccess.replace_whitelist", + "kwargs": { + "tenant_id": workspace_id, + "account_id": operator_account_id, + "dataset_id": current_dataset_id, + "payload": replace_whitelist_payload.model_dump(mode="json"), + }, + }, + } + ) + if not dry_run: + RBACService.DatasetAccess.replace_whitelist( + tenant_id=workspace_id, + account_id=operator_account_id, + dataset_id=current_dataset_id, + payload=replace_whitelist_payload, + ) + scope_migrated_count += 1 + + if should_bind_partial_members: + partial_dataset_count += 1 + for member_account_id in partial_member_ids: + replace_user_access_policies_payload = ReplaceUserAccessPolicies( + access_policy_ids=[_RBAC_DEFAULT_ACCESS_POLICY_ID], + ) + if dry_run: + _emit_dataset_permission_migration_event( + { + "event": "dataset_permission_migration_proposed_change", + "action": "replace_user_access_policies", + "dry_run": True, + "tenant_id": workspace_id, + "dataset_id": current_dataset_id, + "operator_account_id": operator_account_id, + "target_account_id": member_account_id, + "before": { + "legacy_dataset_permission": permission_value.value, + "legacy_partial_member_id": member_account_id, + }, + "after": { + "rbac_user_access_policy_ids": [_RBAC_DEFAULT_ACCESS_POLICY_ID], + }, + "call": { + "method": "RBACService.DatasetAccess.replace_user_access_policies", + "kwargs": { + "tenant_id": workspace_id, + "account_id": operator_account_id, + "dataset_id": current_dataset_id, + "target_account_id": member_account_id, + "payload": replace_user_access_policies_payload.model_dump(mode="json"), + }, + }, + } + ) + continue + RBACService.DatasetAccess.replace_user_access_policies( + tenant_id=workspace_id, + account_id=operator_account_id, + dataset_id=current_dataset_id, + target_account_id=member_account_id, + payload=replace_user_access_policies_payload, + ) + user_policy_migrated_count += 1 + + last_dataset_id = dataset_ids[-1] + + if dataset_id: + break + + if scanned_count == 0: + click.echo(click.style("No datasets found for migration.", fg="yellow")) + return + + if dry_run: + click.echo( + click.style( + f"Dry run completed. Scanned {scanned_count} datasets; " + f"{partial_dataset_count} partial-member datasets would be migrated.", + fg="yellow", + ) + ) + else: + click.echo( + click.style( + "RBAC dataset permission migration completed. " + f"Scanned {scanned_count} datasets, migrated {scope_migrated_count} scopes, " + f"wrote {user_policy_migrated_count} user default-policy bindings.", + fg="green", + ) + ) diff --git a/api/controllers/console/workspace/rbac.py b/api/controllers/console/workspace/rbac.py index c3a3420b908..be1783dddd0 100644 --- a/api/controllers/console/workspace/rbac.py +++ b/api/controllers/console/workspace/rbac.py @@ -1,6 +1,5 @@ from __future__ import annotations -from enum import StrEnum from typing import Any from flask import request @@ -14,6 +13,7 @@ from controllers.common.schema import register_response_schema_models from controllers.console import console_ns from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required from core.db.session_factory import session_factory +from core.rbac import RBACResourceWhitelistScope from libs.login import current_account_with_tenant, login_required from models import Account from services.enterprise import rbac_service as svc @@ -511,14 +511,8 @@ class RBACAccessPolicyBindingUnlockApi(Resource): # --------------------------------------------------------------------------- -class _AccessScope(StrEnum): - ALL = "all" - SPECIFIC = "specific" - ONLY_ME = "only_me" - - class _ResourceAccessScopeRequest(BaseModel): - scope: _AccessScope + scope: RBACResourceWhitelistScope class _ReplaceBindingsRequest(BaseModel): diff --git a/api/core/rbac/__init__.py b/api/core/rbac/__init__.py index 495ac90f492..e8cf30e393f 100644 --- a/api/core/rbac/__init__.py +++ b/api/core/rbac/__init__.py @@ -1,3 +1,3 @@ -from core.rbac.entities import RBACPermission, RBACResourceScope +from core.rbac.entities import RBACPermission, RBACResourceScope, RBACResourceWhitelistScope -__all__ = ["RBACPermission", "RBACResourceScope"] +__all__ = ["RBACPermission", "RBACResourceScope", "RBACResourceWhitelistScope"] diff --git a/api/core/rbac/entities.py b/api/core/rbac/entities.py index 16a05f13111..200c33e9424 100644 --- a/api/core/rbac/entities.py +++ b/api/core/rbac/entities.py @@ -13,6 +13,14 @@ class RBACResourceScope(StrEnum): WORKSPACE = "workspace" +class RBACResourceWhitelistScope(StrEnum): + """Whitelist scopes accepted by RBAC app and dataset access config APIs.""" + + ALL = "all" + SPECIFIC = "specific" + ONLY_ME = "only_me" + + class RBACPermission(StrEnum): """Permission points (RBAC scenes) checked by ``rbac_permission_required``. diff --git a/api/services/enterprise/rbac_service.py b/api/services/enterprise/rbac_service.py index ee51fd52df7..16878ae0045 100644 --- a/api/services/enterprise/rbac_service.py +++ b/api/services/enterprise/rbac_service.py @@ -12,6 +12,7 @@ from sqlalchemy.exc import SQLAlchemyError from configs import dify_config from core.db.session_factory import session_factory +from core.rbac import RBACResourceWhitelistScope from models import TenantAccountJoin, TenantAccountRole from services.enterprise.base import EnterpriseRequest @@ -650,17 +651,18 @@ class ReplaceRoleBindings(_RBACModel): class ReplaceMemberBindings(_RBACModel): - scope: str = "specific" + scope: RBACResourceWhitelistScope = RBACResourceWhitelistScope.SPECIFIC @field_validator("scope") @classmethod - def _normalize_scope(cls, value: Any) -> str: + def _normalize_scope(cls, value: Any) -> RBACResourceWhitelistScope: scope = str(value or "").strip().lower() - if scope in {"", "specific"}: - return "specific" - if scope in {"all", "only_me"}: - return scope - raise ValueError(f"invalid scope: {value}") + if scope == "": + return RBACResourceWhitelistScope.SPECIFIC + try: + return RBACResourceWhitelistScope(scope) + except ValueError as exc: + raise ValueError(f"invalid scope: {value}") from exc class DeleteMemberBindings(_RBACModel): diff --git a/api/tests/unit_tests/commands/test_legacy_model_type_migration.py b/api/tests/unit_tests/commands/test_legacy_model_type_migration.py index 7eead948c1c..e14c8ed3243 100644 --- a/api/tests/unit_tests/commands/test_legacy_model_type_migration.py +++ b/api/tests/unit_tests/commands/test_legacy_model_type_migration.py @@ -336,6 +336,174 @@ def _insert_load_balancing_model_config( ) +def test_data_migrate_group_registers_dataset_permission_rbac_migration(command_module) -> None: + command = command_module.data_migrate.commands["rbac-migrate-dataset-permissions"] + + assert command is command_module.migrate_dataset_permissions_to_rbac + assert "operator_account_id" not in {param.name for param in command.params} + + +def test_dataset_permission_rbac_migration_help_mentions_binding_clear_side_effect(command_module) -> None: + result = CliRunner().invoke( + command_module.data_migrate, + ["rbac-migrate-dataset-permissions", "--help"], + ) + + assert result.exit_code == 0 + normalized_output = " ".join(result.output.split()) + assert "clears existing per-user policy bindings" in normalized_output + assert "recreates legacy partial-member default bindings" in normalized_output + + +def test_dataset_permission_rbac_migration_maps_legacy_permissions_to_enum_scopes() -> None: + rbac_module = importlib.import_module("commands.rbac") + + assert ( + rbac_module._rbac_dataset_scope_for_legacy_permission(rbac_module.DatasetPermissionEnum.ALL_TEAM) + is rbac_module.RBACResourceWhitelistScope.ALL + ) + assert ( + rbac_module._rbac_dataset_scope_for_legacy_permission(rbac_module.DatasetPermissionEnum.PARTIAL_TEAM) + is rbac_module.RBACResourceWhitelistScope.SPECIFIC + ) + assert rbac_module._dataset_permission_enum("partial_members") is rbac_module.DatasetPermissionEnum.PARTIAL_TEAM + assert rbac_module._dataset_permission_enum(None) is rbac_module.DatasetPermissionEnum.ONLY_ME + + +def test_dataset_permission_rbac_migration_uses_dataset_creator_as_operator( + command_module, + monkeypatch: pytest.MonkeyPatch, +) -> None: + rbac_module = importlib.import_module("commands.rbac") + dataset_row = SimpleNamespace( + id="dataset-1", + tenant_id="tenant-1", + permission="only_me", + created_by="creator-account-1", + ) + execute_results = [[dataset_row], [], []] + calls: list[dict[str, object]] = [] + session_closed = False + + class FakeExecuteResult: + def __init__(self, rows: list[object]) -> None: + self._rows = rows + + def all(self) -> list[object]: + return self._rows + + class FakeSession: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, traceback) -> None: + nonlocal session_closed + session_closed = True + pass + + def execute(self, stmt): + return FakeExecuteResult(execute_results.pop(0)) + + class FakeSessionFactory: + @staticmethod + def create_session() -> FakeSession: + return FakeSession() + + def fake_replace_whitelist(**kwargs): + assert session_closed is True + calls.append(kwargs) + + monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory) + monkeypatch.setattr(rbac_module.RBACService.DatasetAccess, "replace_whitelist", fake_replace_whitelist) + + command_module.migrate_dataset_permissions_to_rbac.callback( + tenant_id=None, + dataset_id=None, + batch_size=500, + dry_run=False, + ) + + assert calls[0]["tenant_id"] == "tenant-1" + assert calls[0]["account_id"] == "creator-account-1" + assert calls[0]["dataset_id"] == "dataset-1" + assert calls[0]["payload"].scope is rbac_module.RBACResourceWhitelistScope.SPECIFIC + + +def test_dataset_permission_rbac_migration_dry_run_outputs_structured_proposed_changes( + command_module, + monkeypatch: pytest.MonkeyPatch, +) -> None: + rbac_module = importlib.import_module("commands.rbac") + dataset_row = SimpleNamespace( + id="dataset-1", + tenant_id="tenant-1", + permission="partial_members", + created_by="creator-account-1", + ) + permission_row = SimpleNamespace(dataset_id="dataset-1", account_id="member-account-1") + execute_results = [[dataset_row], [permission_row], []] + + class FakeExecuteResult: + def __init__(self, rows: list[object]) -> None: + self._rows = rows + + def all(self) -> list[object]: + return self._rows + + class FakeSession: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, traceback) -> None: + pass + + def execute(self, stmt): + return FakeExecuteResult(execute_results.pop(0)) + + class FakeSessionFactory: + @staticmethod + def create_session() -> FakeSession: + return FakeSession() + + monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory) + monkeypatch.setattr( + rbac_module.RBACService.DatasetAccess, + "replace_whitelist", + lambda **kwargs: pytest.fail("dry-run must not replace whitelist"), + ) + monkeypatch.setattr( + rbac_module.RBACService.DatasetAccess, + "replace_user_access_policies", + lambda **kwargs: pytest.fail("dry-run must not replace user access policies"), + ) + + result = CliRunner().invoke( + command_module.data_migrate, + ["rbac-migrate-dataset-permissions", "--dry-run"], + ) + + assert result.exit_code == 0 + events = [json.loads(line) for line in result.output.splitlines() if line.startswith("{")] + assert [event["action"] for event in events] == ["replace_whitelist", "replace_user_access_policies"] + assert events[0]["before"] == { + "legacy_dataset_permission": "partial_members", + "legacy_partial_member_ids": ["member-account-1"], + } + assert events[0]["after"] == {"rbac_whitelist_scope": "specific"} + assert events[0]["call"] == { + "method": "RBACService.DatasetAccess.replace_whitelist", + "kwargs": { + "tenant_id": "tenant-1", + "account_id": "creator-account-1", + "dataset_id": "dataset-1", + "payload": {"scope": "specific"}, + }, + } + assert events[1]["target_account_id"] == "member-account-1" + assert events[1]["after"] == {"rbac_user_access_policy_ids": ["default"]} + assert events[1]["call"]["kwargs"]["payload"] == {"access_policy_ids": ["default"]} + + def test_data_migrate_command_defaults_output_to_stdout_stream( command_module, monkeypatch: pytest.MonkeyPatch, diff --git a/api/tests/unit_tests/controllers/console/workspace/test_rbac.py b/api/tests/unit_tests/controllers/console/workspace/test_rbac.py index 2960bfef324..92e819f73bf 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_rbac.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_rbac.py @@ -136,7 +136,7 @@ class TestPydanticModels: def test_resource_access_scope_defaults_empty_account_ids(self): parsed = rbac_mod._ResourceAccessScopeRequest.model_validate({"scope": "specific"}) - assert parsed.scope is rbac_mod._AccessScope.SPECIFIC + assert parsed.scope is rbac_mod.RBACResourceWhitelistScope.SPECIFIC def test_resource_access_scope_coerce_null_account_ids(self): rbac_mod._ResourceAccessScopeRequest.model_validate({"scope": "all"}) From 102e1ede6e13f0abeb1bcda7eb3ad6507985b25c Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Tue, 30 Jun 2026 15:55:54 +0900 Subject: [PATCH 03/54] chore: inject more db.session (#38045) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../console/datasets/data_source.py | 131 ++-- api/controllers/console/datasets/datasets.py | 46 +- .../console/datasets/datasets_document.py | 82 ++- .../console/datasets/datasets_segments.py | 72 +- api/controllers/console/datasets/external.py | 2 +- .../console/datasets/hit_testing_base.py | 2 +- api/controllers/console/datasets/metadata.py | 12 +- .../rag_pipeline/rag_pipeline_datasets.py | 16 +- .../service_api/dataset/dataset.py | 26 +- .../service_api/dataset/document.py | 21 +- .../service_api/dataset/metadata.py | 12 +- .../service_api/dataset/segment.py | 66 +- .../easy_ui_based_app/dataset/manager.py | 3 +- .../app/apps/pipeline/pipeline_generator.py | 2 +- .../annotation_reply/annotation_reply.py | 2 +- api/models/dataset.py | 4 +- api/services/dataset_service.py | 683 ++++++++++-------- api/services/metadata_service.py | 10 +- .../rag_pipeline/rag_pipeline_dsl_service.py | 4 +- api/services/summary_index_service.py | 5 +- .../add_annotation_to_index_task.py | 8 +- .../batch_import_annotations_task.py | 2 +- .../delete_annotation_index_task.py | 8 +- .../enable_annotation_reply_task.py | 4 +- .../update_annotation_to_index_task.py | 8 +- .../console/datasets/test_data_source.py | 35 +- .../service_api/dataset/test_dataset.py | 4 +- .../services/dataset_collection_binding.py | 24 +- .../services/dataset_service_update_delete.py | 18 +- .../services/document_service_status.py | 58 +- .../test_dataset_permission_service.py | 105 ++- .../services/test_dataset_service.py | 55 +- ...et_service_batch_update_document_status.py | 34 +- .../test_dataset_service_create_dataset.py | 1 + .../test_dataset_service_delete_dataset.py | 10 +- .../services/test_dataset_service_document.py | 46 +- .../test_dataset_service_permissions.py | 91 ++- .../test_dataset_service_retrieval.py | 12 +- .../test_dataset_service_update_dataset.py | 30 +- .../test_document_service_rename_document.py | 14 +- .../datasets/test_datasets_document.py | 8 +- .../test_datasets_document_download.py | 8 +- .../controllers/console/test_spec.py | 4 +- .../dataset/test_dataset_segment.py | 38 +- .../service_api/dataset/test_document.py | 25 +- .../extensions/test_ext_request_logging.py | 9 +- .../services/test_dataset_service_dataset.py | 183 +++-- .../services/test_dataset_service_document.py | 250 +++++-- .../test_dataset_service_lock_not_owned.py | 7 +- .../services/test_dataset_service_segment.py | 74 +- .../services/test_summary_index_service.py | 2 +- 51 files changed, 1483 insertions(+), 893 deletions(-) diff --git a/api/controllers/console/datasets/data_source.py b/api/controllers/console/datasets/data_source.py index 6dd13b485bc..b2c8bda0581 100644 --- a/api/controllers/console/datasets/data_source.py +++ b/api/controllers/console/datasets/data_source.py @@ -243,72 +243,71 @@ class DataSourceNotionListApi(Resource): if not credential: raise NotFound("Credential not found.") exist_page_ids = [] - with sessionmaker(db.engine).begin() as session: - # import notion in the exist dataset - if query.dataset_id: - dataset = DatasetService.get_dataset(query.dataset_id) - if not dataset: - raise NotFound("Dataset not found.") - if dataset.data_source_type != "notion_import": - raise ValueError("Dataset is not notion type.") + # import notion in the exist dataset + if query.dataset_id: + dataset = DatasetService.get_dataset(query.dataset_id, db.session) + if not dataset: + raise NotFound("Dataset not found.") + if dataset.data_source_type != "notion_import": + raise ValueError("Dataset is not notion type.") - documents = session.scalars( - select(Document).where( - Document.dataset_id == query.dataset_id, - Document.tenant_id == current_tenant_id, - Document.data_source_type == "notion_import", - Document.enabled.is_(True), - ) - ).all() - if documents: - for document in documents: - data_source_info = json.loads(document.data_source_info) - exist_page_ids.append(data_source_info["notion_page_id"]) - # get all authorized pages - from core.datasource.datasource_manager import DatasourceManager - - datasource_runtime = DatasourceManager.get_datasource_runtime( - provider_id="langgenius/notion_datasource/notion_datasource", - datasource_name="notion_datasource", - tenant_id=current_tenant_id, - datasource_type=DatasourceProviderType.ONLINE_DOCUMENT, - ) - datasource_provider_service = DatasourceProviderService() - if credential: - datasource_runtime.runtime.credentials = credential - datasource_runtime = cast(OnlineDocumentDatasourcePlugin, datasource_runtime) - online_document_result: Generator[OnlineDocumentPagesMessage, None, None] = ( - datasource_runtime.get_online_document_pages( - user_id=current_user.id, - datasource_parameters={}, - provider_type=datasource_runtime.datasource_provider_type(), + documents = db.session.scalars( + select(Document).where( + Document.dataset_id == query.dataset_id, + Document.tenant_id == current_tenant_id, + Document.data_source_type == "notion_import", + Document.enabled.is_(True), ) + ).all() + if documents: + for document in documents: + data_source_info = json.loads(document.data_source_info) + exist_page_ids.append(data_source_info["notion_page_id"]) + # get all authorized pages + from core.datasource.datasource_manager import DatasourceManager + + datasource_runtime = DatasourceManager.get_datasource_runtime( + provider_id="langgenius/notion_datasource/notion_datasource", + datasource_name="notion_datasource", + tenant_id=current_tenant_id, + datasource_type=DatasourceProviderType.ONLINE_DOCUMENT, + ) + datasource_provider_service = DatasourceProviderService() + if credential: + datasource_runtime.runtime.credentials = credential + datasource_runtime = cast(OnlineDocumentDatasourcePlugin, datasource_runtime) + online_document_result: Generator[OnlineDocumentPagesMessage, None, None] = ( + datasource_runtime.get_online_document_pages( + user_id=current_user.id, + datasource_parameters={}, + provider_type=datasource_runtime.datasource_provider_type(), ) - try: - pages = [] - workspace_info = {} - for message in online_document_result: - result = message.result - for info in result: - workspace_info = { - "workspace_id": info.workspace_id, - "workspace_name": info.workspace_name, - "workspace_icon": info.workspace_icon, + ) + try: + pages = [] + workspace_info = {} + for message in online_document_result: + result = message.result + for info in result: + workspace_info = { + "workspace_id": info.workspace_id, + "workspace_name": info.workspace_name, + "workspace_icon": info.workspace_icon, + } + for page in info.pages: + page_info = { + "page_id": page.page_id, + "page_name": page.page_name, + "type": page.type, + "parent_id": page.parent_id, + "is_bound": page.page_id in exist_page_ids, + "page_icon": page.page_icon, } - for page in info.pages: - page_info = { - "page_id": page.page_id, - "page_name": page.page_name, - "type": page.type, - "parent_id": page.parent_id, - "is_bound": page.page_id in exist_page_ids, - "page_icon": page.page_icon, - } - pages.append(page_info) - except Exception as e: - raise e - notion_info = [{**workspace_info, "pages": pages}] if workspace_info else [] - return dump_response(NotionIntegrateInfoListResponse, {"notion_info": notion_info}), 200 + pages.append(page_info) + except Exception as e: + raise e + notion_info = [{**workspace_info, "pages": pages}] if workspace_info else [] + return dump_response(NotionIntegrateInfoListResponse, {"notion_info": notion_info}), 200 @console_ns.route("/notion/pages///preview") @@ -401,11 +400,11 @@ class DataSourceNotionDatasetSyncApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) def get(self, dataset_id: UUID) -> tuple[dict[str, str], int]: dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") - documents = DocumentService.get_document_by_dataset_id(dataset_id_str) + documents = DocumentService.get_document_by_dataset_id(dataset_id_str, db.session) for document in documents: document_indexing_sync_task.delay(dataset_id_str, document.id) return {"result": "success"}, 200 @@ -421,11 +420,11 @@ class DataSourceNotionDocumentSyncApi(Resource): def get(self, dataset_id: UUID, document_id: UUID) -> tuple[dict[str, str], int]: dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if document is None: raise NotFound("Document not found.") document_indexing_sync_task.delay(dataset_id_str, document_id_str) diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index 70ce54830c7..dbf546532ee 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -561,6 +561,7 @@ class DatasetListApi(Resource): provider=payload.provider, external_knowledge_api_id=payload.external_knowledge_api_id, external_knowledge_id=payload.external_knowledge_id, + session=db.session, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -598,7 +599,7 @@ class DatasetApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") try: @@ -618,7 +619,7 @@ class DatasetApi(Resource): provider_id = ModelProviderID(dataset.embedding_model_provider) data["embedding_model_provider"] = str(provider_id) if data.get("permission") == "partial_members": - part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str) + part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) data.update({"partial_member_list": part_users_list}) # check embedding setting @@ -661,7 +662,7 @@ class DatasetApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def patch(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -680,10 +681,10 @@ class DatasetApi(Resource): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not dify_config.RBAC_ENABLED: DatasetPermissionService.check_permission( - current_user, dataset, payload.permission, payload.partial_member_list + current_user, dataset, payload.permission, payload.partial_member_list, db.session ) - dataset = DatasetService.update_dataset(dataset_id_str, payload_data, current_user) + dataset = DatasetService.update_dataset(dataset_id_str, payload_data, current_user, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -698,12 +699,14 @@ class DatasetApi(Resource): tenant_id = current_tenant_id if payload.partial_member_list is not None and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: - DatasetPermissionService.update_partial_member_list(tenant_id, dataset_id_str, payload.partial_member_list) + DatasetPermissionService.update_partial_member_list( + tenant_id, dataset_id_str, payload.partial_member_list, db.session + ) # clear partial member list when permission is only_me or all_team_members elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: - DatasetPermissionService.clear_partial_member_list(dataset_id_str) + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) - partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str) + partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) result_data.update({"partial_member_list": partial_member_list}) return result_data, 200 @@ -722,8 +725,8 @@ class DatasetApi(Resource): raise Forbidden() try: - if DatasetService.delete_dataset(dataset_id_str, current_user): - DatasetPermissionService.clear_partial_member_list(dataset_id_str) + if DatasetService.delete_dataset(dataset_id_str, current_user, db.session): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) return "", 204 else: raise NotFound("Dataset not found.") @@ -748,7 +751,7 @@ class DatasetUseCheckApi(Resource): def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset_is_using = DatasetService.dataset_use_check(dataset_id_str) + dataset_is_using = DatasetService.dataset_use_check(dataset_id_str, db.session) return {"is_using": dataset_is_using}, 200 @@ -769,7 +772,7 @@ class DatasetQueryApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -910,7 +913,7 @@ class DatasetRelatedAppListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -919,7 +922,7 @@ class DatasetRelatedAppListApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - app_dataset_joins = DatasetService.get_related_apps(dataset.id) + app_dataset_joins = DatasetService.get_related_apps(dataset.id, db.session) related_apps = [] for app_dataset_join in app_dataset_joins: @@ -1094,7 +1097,7 @@ class DatasetEnableApiApi(Resource): def post(self, dataset_id: UUID, status: str): dataset_id_str = str(dataset_id) - DatasetService.update_dataset_api_status(dataset_id_str, status == "enable") + DatasetService.update_dataset_api_status(dataset_id_str, status == "enable", db.session) return {"result": "success"}, 200 @@ -1163,10 +1166,10 @@ class DatasetErrorDocs(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") - results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str) + results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str, db.session) return dump_response(ErrorDocsResponse, {"data": results, "total": len(results)}), 200 @@ -1190,7 +1193,7 @@ class DatasetPermissionUserListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") try: @@ -1198,7 +1201,7 @@ class DatasetPermissionUserListApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - partial_members_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str) + partial_members_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) return dump_response(PartialMemberListResponse, {"data": partial_members_list}), 200 @@ -1220,7 +1223,8 @@ class DatasetAutoDisableLogApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") - return dump_response(AutoDisableLogsResponse, DatasetService.get_dataset_auto_disable_logs(dataset_id_str)), 200 + auto_disable_logs = DatasetService.get_dataset_auto_disable_logs(dataset_id_str, db.session) + return dump_response(AutoDisableLogsResponse, auto_disable_logs), 200 diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index 499558dc4dd..37f64f54d24 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -181,7 +181,7 @@ class DocumentResource(Resource): def get_document( self, dataset_id: str, document_id: str, current_user: Account, current_tenant_id: str ) -> Document: - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -190,7 +190,7 @@ class DocumentResource(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id, document_id) + document = DocumentService.get_document(dataset_id, document_id, session=db.session) if not document: raise NotFound("Document not found.") @@ -201,7 +201,7 @@ class DocumentResource(Resource): return document def get_batch_documents(self, dataset_id: str, batch: str, current_user: Account) -> Sequence[Document]: - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -210,7 +210,7 @@ class DocumentResource(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - documents = DocumentService.get_batch_documents(dataset_id, batch) + documents = DocumentService.get_batch_documents(dataset_id, batch, db.session) if not documents: raise NotFound("Documents not found.") @@ -241,7 +241,7 @@ class GetProcessRuleApi(Resource): # get the latest process rule document = db.get_or_404(Document, document_id) - dataset = DatasetService.get_dataset(document.dataset_id) + dataset = DatasetService.get_dataset(document.dataset_id, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -317,7 +317,7 @@ class DatasetDocumentListApi(Resource): ) except (ArgumentTypeError, ValueError, Exception): fetch = False - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -421,7 +421,7 @@ class DatasetDocumentListApi(Resource): def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -444,8 +444,10 @@ class DatasetDocumentListApi(Resource): DocumentService.document_create_args_validate(knowledge_config) try: - documents, batch = DocumentService.save_document_with_dataset_id(dataset, knowledge_config, current_user) - dataset = DatasetService.get_dataset(dataset_id_str) + documents, batch = DocumentService.save_document_with_dataset_id( + dataset, knowledge_config, current_user, session=db.session + ) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -464,7 +466,7 @@ class DatasetDocumentListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def delete(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -472,7 +474,7 @@ class DatasetDocumentListApi(Resource): try: document_ids = request.args.getlist("document_id") - DocumentService.delete_documents(dataset, document_ids) + DocumentService.delete_documents(dataset, document_ids, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -531,6 +533,7 @@ class DatasetInitApi(Resource): tenant_id=current_tenant_id, knowledge_config=knowledge_config, account=current_user, + session=db.session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -867,7 +870,7 @@ class DocumentApi(DocumentResource): if metadata == "only": response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} elif metadata == "without": - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} response = { "id": document.id, @@ -901,7 +904,7 @@ class DocumentApi(DocumentResource): "need_summary": document.need_summary if document.need_summary is not None else False, } else: - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} response = { "id": document.id, @@ -950,7 +953,7 @@ class DocumentApi(DocumentResource): def delete(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -959,7 +962,7 @@ class DocumentApi(DocumentResource): document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) try: - DocumentService.delete_document(document) + DocumentService.delete_document(document, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -983,7 +986,7 @@ class DocumentDownloadApi(DocumentResource): def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID) -> dict[str, Any]: # Reuse the shared permission/tenant checks implemented in DocumentResource. document = self.get_document(str(dataset_id), str(document_id), current_user, current_tenant_id) - return {"url": DocumentService.get_document_download_url(document)} + return {"url": DocumentService.get_document_download_url(document, db.session)} @console_ns.route("/datasets//documents/download-zip") @@ -1013,6 +1016,7 @@ class DocumentBatchDownloadZipApi(DocumentResource): document_ids=document_ids, tenant_id=current_tenant_id, current_user=current_user, + session=db.session, ) # Delegate ZIP packing to FileService, but keep Flask response+cleanup in the route. @@ -1161,7 +1165,7 @@ class DocumentStatusApi(DocumentResource): self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable", "archive", "un_archive"] ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -1178,7 +1182,7 @@ class DocumentStatusApi(DocumentResource): document_ids = request.args.getlist("document_id") try: - DocumentService.batch_update_document_status(dataset, document_ids, action, current_user) + DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -1202,11 +1206,11 @@ class DocumentPauseApi(DocumentResource): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) # 404 if document not found if document is None: @@ -1218,7 +1222,7 @@ class DocumentPauseApi(DocumentResource): try: # pause document - DocumentService.pause_document(document) + DocumentService.pause_document(document, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot pause completed document.") @@ -1237,10 +1241,10 @@ class DocumentRecoverApi(DocumentResource): """recover document.""" dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) # 404 if document not found if document is None: @@ -1251,7 +1255,7 @@ class DocumentRecoverApi(DocumentResource): raise ArchivedDocumentImmutableError() try: # pause document - DocumentService.recover_document(document) + DocumentService.recover_document(document, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Document is not in paused status.") @@ -1271,13 +1275,13 @@ class DocumentRetryApi(DocumentResource): """retry document.""" payload = DocumentRetryPayload.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) retry_documents = [] if not dataset: raise NotFound("Dataset not found.") for document_id in payload.document_ids: try: - document = DocumentService.get_document(dataset.id, document_id) + document = DocumentService.get_document(dataset.id, document_id, session=db.session) # 404 if document not found if document is None: @@ -1295,7 +1299,7 @@ class DocumentRetryApi(DocumentResource): logger.exception("Failed to retry document, document id: %s", document_id) continue # retry document - DocumentService.retry_document(dataset_id_str, retry_documents) + DocumentService.retry_document(dataset_id_str, retry_documents, db.session) return "", 204 @@ -1313,14 +1317,14 @@ class DocumentRenameApi(DocumentResource): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not current_user.is_dataset_editor: raise Forbidden() - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, db.session) if not dataset: raise NotFound("Dataset not found.") - DatasetService.check_dataset_operator_permission(current_user, dataset) + DatasetService.check_dataset_operator_permission(current_user, dataset, session=db.session) payload = DocumentRenamePayload.model_validate(console_ns.payload or {}) try: - document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name) + document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -1338,11 +1342,11 @@ class WebsiteDocumentSyncApi(DocumentResource): def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID): """sync website document.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") document_id_str = str(document_id) - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") if document.tenant_id != current_tenant_id: @@ -1353,7 +1357,7 @@ class WebsiteDocumentSyncApi(DocumentResource): if DocumentService.check_archived(document): raise ArchivedDocumentImmutableError() # sync document - DocumentService.sync_website_document(dataset_id_str, document) + DocumentService.sync_website_document(dataset_id_str, document, db.session) return {"result": "success"}, 200 @@ -1373,10 +1377,10 @@ class DocumentPipelineExecutionLogApi(DocumentResource): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") log = db.session.scalar( @@ -1431,7 +1435,7 @@ class DocumentGenerateSummaryApi(Resource): dataset_id_str = str(dataset_id) # Get dataset - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -1465,7 +1469,7 @@ class DocumentGenerateSummaryApi(Resource): raise ValueError("Summary index is not enabled for this dataset. Please enable it in the dataset settings.") # Verify all documents exist and belong to the dataset - documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list) + documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list, db.session) if len(documents) != len(document_list): found_ids = {doc.id for doc in documents} @@ -1481,6 +1485,7 @@ class DocumentGenerateSummaryApi(Resource): DocumentService.update_documents_need_summary( dataset_id=dataset_id_str, document_ids=document_ids_to_update, + session=db.session, need_summary=True, ) @@ -1531,7 +1536,7 @@ class DocumentSummaryStatusApi(DocumentResource): document_id_str = str(document_id) # Get dataset - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -1547,6 +1552,7 @@ class DocumentSummaryStatusApi(DocumentResource): result = SummaryIndexService.get_document_summary_status_detail( document_id=document_id_str, dataset_id=dataset_id_str, + session=db.session, ) return result, 200 diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 5ba115ff491..1a6f4c7a712 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -176,7 +176,7 @@ class DatasetDocumentSegmentListApi(Resource): def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") @@ -185,7 +185,7 @@ class DatasetDocumentSegmentListApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") @@ -286,14 +286,14 @@ class DatasetDocumentSegmentListApi(Resource): def delete(self, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_ids = request.args.getlist("segment_id") @@ -305,7 +305,7 @@ class DatasetDocumentSegmentListApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - SegmentService.delete_segments(segment_ids, document, dataset) + SegmentService.delete_segments(segment_ids, document, dataset, db.session) return "", 204 @@ -331,11 +331,11 @@ class DatasetDocumentSegmentApi(Resource): action: Literal["enable", "disable"], ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check user's model setting @@ -371,7 +371,7 @@ class DatasetDocumentSegmentApi(Resource): if cache_result is not None: raise InvalidActionError("Document is being indexed, please try again later") try: - SegmentService.update_segments_status(segment_ids, action, dataset, document) + SegmentService.update_segments_status(segment_ids, action, dataset, document, db.session) except Exception as e: raise InvalidActionError(str(e)) return dump_response(SimpleResultResponse, {"result": "success"}), 200 @@ -394,12 +394,12 @@ class DatasetDocumentSegmentAddApi(Resource): def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") if not current_user.is_dataset_editor: @@ -428,7 +428,7 @@ class DatasetDocumentSegmentAddApi(Resource): payload = SegmentCreatePayload.model_validate(console_ns.payload or {}) payload_dict = payload.model_dump(exclude_none=True) SegmentService.segment_create_args_validate(payload_dict, document) - segment = type_cast(DocumentSegment, SegmentService.create_segment(payload_dict, document, dataset)) + segment = type_cast(DocumentSegment, SegmentService.create_segment(payload_dict, document, dataset, db.session)) summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), @@ -455,14 +455,14 @@ class DatasetDocumentSegmentUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -504,7 +504,11 @@ class DatasetDocumentSegmentUpdateApi(Resource): # Update segment (summary update with change detection is handled in SegmentService.update_segment) segment = SegmentService.update_segment( - SegmentUpdateArgs.model_validate(payload.model_dump(exclude_none=True)), segment, document, dataset + SegmentUpdateArgs.model_validate(payload.model_dump(exclude_none=True)), + segment, + document, + dataset, + db.session, ) summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) response = { @@ -527,14 +531,14 @@ class DatasetDocumentSegmentUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check segment @@ -553,7 +557,7 @@ class DatasetDocumentSegmentUpdateApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - SegmentService.delete_segment(segment, document, dataset) + SegmentService.delete_segment(segment, document, dataset, db.session) return "", 204 @@ -576,12 +580,12 @@ class DatasetDocumentSegmentBatchImportApi(Resource): def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") @@ -651,12 +655,12 @@ class ChildChunkAddApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check segment @@ -693,7 +697,7 @@ class ChildChunkAddApi(Resource): # validate args try: payload = ChildChunkCreatePayload.model_validate(console_ns.payload or {}) - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset) + child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkDetailResponse, {"data": child_chunk}), 200 @@ -709,14 +713,14 @@ class ChildChunkAddApi(Resource): def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check segment @@ -766,14 +770,14 @@ class ChildChunkAddApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check segment @@ -795,7 +799,7 @@ class ChildChunkAddApi(Resource): # validate args payload = ChildChunkBatchUpdatePayload.model_validate(console_ns.payload or {}) try: - child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset) + child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset, db.session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkBatchUpdateResponse, {"data": child_chunks}), 200 @@ -825,14 +829,14 @@ class ChildChunkUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check segment @@ -866,7 +870,7 @@ class ChildChunkUpdateApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) try: - SegmentService.delete_child_chunk(child_chunk, dataset) + SegmentService.delete_child_chunk(child_chunk, dataset, db.session) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) return "", 204 @@ -893,14 +897,14 @@ class ChildChunkUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check segment @@ -936,7 +940,9 @@ class ChildChunkUpdateApi(Resource): # validate args try: payload = ChildChunkUpdatePayload.model_validate(console_ns.payload or {}) - child_chunk = SegmentService.update_child_chunk(payload.content, child_chunk, segment, document, dataset) + child_chunk = SegmentService.update_child_chunk( + payload.content, child_chunk, segment, document, dataset, db.session + ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkDetailResponse, {"data": child_chunk}), 200 diff --git a/api/controllers/console/datasets/external.py b/api/controllers/console/datasets/external.py index 99a61807a4d..7a3c746b80f 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -377,7 +377,7 @@ class ExternalKnowledgeHitTestingApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_PIPELINE_TEST) def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") diff --git a/api/controllers/console/datasets/hit_testing_base.py b/api/controllers/console/datasets/hit_testing_base.py index 82c30fc7ffb..6464e435fc2 100644 --- a/api/controllers/console/datasets/hit_testing_base.py +++ b/api/controllers/console/datasets/hit_testing_base.py @@ -85,7 +85,7 @@ class DatasetsHitTestingBase: dataset_id: str, current_user: Account | None = None, current_tenant_id: str | None = None ) -> Dataset: current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, db.session) if dataset is None: raise NotFound("Dataset not found.") diff --git a/api/controllers/console/datasets/metadata.py b/api/controllers/console/datasets/metadata.py index 90ce263dfe5..8802fcf2814 100644 --- a/api/controllers/console/datasets/metadata.py +++ b/api/controllers/console/datasets/metadata.py @@ -61,7 +61,7 @@ class DatasetMetadataCreateApi(Resource): metadata_args = MetadataArgs.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -81,7 +81,7 @@ class DatasetMetadataCreateApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") metadata = MetadataService.get_dataset_metadatas(db.session(), dataset) @@ -105,7 +105,7 @@ class DatasetMetadataApi(Resource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -125,7 +125,7 @@ class DatasetMetadataApi(Resource): def delete(self, current_user: Account, dataset_id: UUID, metadata_id: UUID): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -162,7 +162,7 @@ class DatasetMetadataBuiltInFieldActionApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def post(self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable"]): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -191,7 +191,7 @@ class DocumentMetadataEditApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py index 0f277b6a4cc..a373c8b1a41 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py @@ -1,6 +1,5 @@ from flask_restx import Resource from pydantic import BaseModel -from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden import services @@ -66,19 +65,19 @@ class CreateRagPipelineDatasetApi(Resource): yaml_content=payload.yaml_content, ) try: - with Session(db.engine, expire_on_commit=False) as session: - rag_pipeline_dsl_service = RagPipelineDslService(session) - import_info = rag_pipeline_dsl_service.create_rag_pipeline_dataset( - tenant_id=current_tenant_id, - rag_pipeline_dataset_create_entity=rag_pipeline_dataset_create_entity, - ) - session.commit() + rag_pipeline_dsl_service = RagPipelineDslService(db.session) + import_info = rag_pipeline_dsl_service.create_rag_pipeline_dataset( + tenant_id=current_tenant_id, + rag_pipeline_dataset_create_entity=rag_pipeline_dataset_create_entity, + ) if rag_pipeline_dataset_create_entity.permission == "partial_members": DatasetPermissionService.update_partial_member_list( current_tenant_id, import_info["dataset_id"], rag_pipeline_dataset_create_entity.partial_member_list, + db.session, ) + db.session.commit() except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -111,5 +110,6 @@ class CreateEmptyRagPipelineDatasetApi(Resource): permission=DatasetPermissionEnum.ONLY_ME, partial_member_list=None, ), + session=db.session, ) return dump_response(DatasetDetailResponse, dataset), 201 diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index bfb7a045082..e903e92e7a6 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -519,6 +519,7 @@ class DatasetListApi(DatasetApiResource): embedding_model_name=payload.embedding_model, retrieval_model=payload.retrieval_model, summary_index_setting=payload.summary_index_setting, + session=db.session, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -561,7 +562,7 @@ class DatasetApi(DatasetApiResource): ) def get(self, _, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") try: @@ -597,7 +598,7 @@ class DatasetApi(DatasetApiResource): retrieval_model_dict["search_method"] = "keyword_search" if data.get("permission") == "partial_members": - part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str) + part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) data.update({"partial_member_list": part_users_list}) return _dump_service_dataset_with_partial_members(data), 200 @@ -635,7 +636,7 @@ class DatasetApi(DatasetApiResource): @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def patch(self, _, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -676,9 +677,10 @@ class DatasetApi(DatasetApiResource): dataset, str(payload.permission) if payload.permission else None, payload.partial_member_list, + db.session, ) - dataset = DatasetService.update_dataset(dataset_id_str, update_data, current_user) + dataset = DatasetService.update_dataset(dataset_id_str, update_data, current_user, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -688,12 +690,14 @@ class DatasetApi(DatasetApiResource): tenant_id = current_user.current_tenant_id if payload.partial_member_list and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: - DatasetPermissionService.update_partial_member_list(tenant_id, dataset_id_str, payload.partial_member_list) + DatasetPermissionService.update_partial_member_list( + tenant_id, dataset_id_str, payload.partial_member_list, db.session + ) # clear partial member list when permission is only_me or all_team_members elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: - DatasetPermissionService.clear_partial_member_list(dataset_id_str) + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) - partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str) + partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) result_data.update({"partial_member_list": partial_member_list}) return _dump_service_dataset_with_partial_members(result_data), 200 @@ -746,8 +750,8 @@ class DatasetApi(DatasetApiResource): dataset_id_str = str(dataset_id) try: - if DatasetService.delete_dataset(dataset_id_str, current_user): - DatasetPermissionService.clear_partial_member_list(dataset_id_str) + if DatasetService.delete_dataset(dataset_id_str, current_user, db.session): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) return "", 204 else: raise NotFound("Dataset not found.") @@ -812,7 +816,7 @@ class DocumentStatusApi(DatasetApiResource): InvalidActionError: If the action is invalid or cannot be performed. """ dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") @@ -831,7 +835,7 @@ class DocumentStatusApi(DatasetApiResource): document_ids = data.get("document_ids", []) try: - DocumentService.batch_update_document_status(dataset, document_ids, action, current_user) + DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: diff --git a/api/controllers/service_api/dataset/document.py b/api/controllers/service_api/dataset/document.py index 9bae862814a..49ccb1bd55c 100644 --- a/api/controllers/service_api/dataset/document.py +++ b/api/controllers/service_api/dataset/document.py @@ -400,6 +400,7 @@ def _create_document_by_text(tenant_id: str, dataset_id: UUID) -> tuple[Mapping[ account=current_user, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", + session=db.session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -459,6 +460,7 @@ def _update_document_by_text(tenant_id: str, dataset_id: UUID, document_id: UUID account=current_user, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", + session=db.session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -756,6 +758,7 @@ class DocumentAddByFileApi(DatasetApiResource): account=dataset.created_by_account, dataset_process_rule=dataset_process_rule, created_from="api", + session=db.session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -832,6 +835,7 @@ def _update_document_by_file(tenant_id: str, dataset_id: UUID, document_id: UUID account=dataset.created_by_account, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", + session=db.session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -1002,6 +1006,7 @@ class DocumentBatchDownloadZipApi(DatasetApiResource): document_ids=[str(document_id) for document_id in payload.document_ids], tenant_id=str(tenant_id), current_user=current_user, + session=db.session, ) with ExitStack() as stack: @@ -1058,7 +1063,7 @@ class DocumentIndexingStatusApi(DatasetApiResource): if not dataset: raise NotFound("Dataset not found.") # get documents - documents = DocumentService.get_batch_documents(dataset_id_str, batch) + documents = DocumentService.get_batch_documents(dataset_id_str, batch, db.session) if not documents: raise NotFound("Documents not found.") documents_status = [] @@ -1134,7 +1139,7 @@ class DocumentDownloadApi(DatasetApiResource): @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def get(self, tenant_id, dataset_id: UUID, document_id: UUID): dataset = self.get_dataset(str(dataset_id), str(tenant_id)) - document = DocumentService.get_document(dataset.id, str(document_id)) + document = DocumentService.get_document(dataset.id, str(document_id), session=db.session) if not document: raise NotFound("Document not found.") @@ -1142,7 +1147,7 @@ class DocumentDownloadApi(DatasetApiResource): if document.tenant_id != str(tenant_id): raise Forbidden("No permission.") - return {"url": DocumentService.get_document_download_url(document)} + return {"url": DocumentService.get_document_download_url(document, db.session)} @service_api_ns.route("/datasets//documents/") @@ -1190,7 +1195,7 @@ class DocumentApi(DatasetApiResource): dataset = self.get_dataset(dataset_id_str, tenant_id) - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") @@ -1215,7 +1220,7 @@ class DocumentApi(DatasetApiResource): if metadata == "only": response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} elif metadata == "without": - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} data_source_info = document.data_source_detail_dict response = { @@ -1250,7 +1255,7 @@ class DocumentApi(DatasetApiResource): "need_summary": document.need_summary if document.need_summary is not None else False, } else: - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} data_source_info = document.data_source_detail_dict response = { @@ -1345,7 +1350,7 @@ class DocumentApi(DatasetApiResource): if not dataset: raise ValueError("Dataset does not exist.") - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) # 404 if document not found if document is None: @@ -1357,7 +1362,7 @@ class DocumentApi(DatasetApiResource): try: # delete document - DocumentService.delete_document(document) + DocumentService.delete_document(document, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") diff --git a/api/controllers/service_api/dataset/metadata.py b/api/controllers/service_api/dataset/metadata.py index 3bb39f0cd4f..aec3b06a91e 100644 --- a/api/controllers/service_api/dataset/metadata.py +++ b/api/controllers/service_api/dataset/metadata.py @@ -81,7 +81,7 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): metadata_args = MetadataArgs.model_validate(service_api_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -116,7 +116,7 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): def get(self, tenant_id, dataset_id: UUID): """Get all metadata for a dataset.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") metadata = MetadataService.get_dataset_metadatas(db.session(), dataset) @@ -154,7 +154,7 @@ class DatasetMetadataServiceApi(DatasetApiResource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -189,7 +189,7 @@ class DatasetMetadataServiceApi(DatasetApiResource): """Delete metadata.""" dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -257,7 +257,7 @@ class DatasetMetadataBuiltInFieldActionServiceApi(DatasetApiResource): def post(self, tenant_id, dataset_id: UUID, action: Literal["enable", "disable"]): """Enable or disable built-in metadata field.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) @@ -303,7 +303,7 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): def post(self, tenant_id, dataset_id: UUID): """Update metadata for multiple documents.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str) + dataset = DatasetService.get_dataset(dataset_id_str, db.session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, db.session) diff --git a/api/controllers/service_api/dataset/segment.py b/api/controllers/service_api/dataset/segment.py index dd8d7c76632..7b0e31952c7 100644 --- a/api/controllers/service_api/dataset/segment.py +++ b/api/controllers/service_api/dataset/segment.py @@ -175,7 +175,7 @@ class SegmentApi(DatasetApiResource): raise NotFound("Dataset not found.") document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") if document.indexing_status != "completed": @@ -210,7 +210,9 @@ class SegmentApi(DatasetApiResource): for args_item in segment_items: SegmentService.segment_create_args_validate(args_item, document) - segments = cast(list[DocumentSegment], SegmentService.multi_create_segment(segment_items, document, dataset)) + segments = cast( + list[DocumentSegment], SegmentService.multi_create_segment(segment_items, document, dataset, db.session) + ) segment_ids = [segment.id for segment in segments] summaries: dict[str, str | None] = {} if segment_ids: @@ -267,7 +269,7 @@ class SegmentApi(DatasetApiResource): raise NotFound("Dataset not found.") document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") # check embedding model setting @@ -349,15 +351,17 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) # check segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") - SegmentService.delete_segment(segment, document, dataset) + SegmentService.delete_segment(segment, document, dataset, db.session) return "", 204 @service_api_ns.doc( @@ -395,7 +399,7 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -416,13 +420,15 @@ class DatasetSegmentApi(DatasetApiResource): raise ProviderNotInitializeError(ex.description) segment_id_str = str(segment_id) # check segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") payload = SegmentUpdatePayload.model_validate(service_api_ns.payload or {}) - updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset) + updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, db.session) summary = SummaryIndexService.get_segment_summary(segment_id=updated_segment.id, dataset_id=dataset_id_str) response = { "data": segment_response_with_summary(updated_segment, summary.summary_content if summary else None), @@ -469,12 +475,14 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) # check segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") @@ -533,13 +541,15 @@ class ChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) # check segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") @@ -564,7 +574,7 @@ class ChildChunkApi(DatasetApiResource): payload = ChildChunkCreatePayload.model_validate(service_api_ns.payload or {}) try: - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset) + child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) @@ -607,13 +617,15 @@ class ChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) # check segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") @@ -677,13 +689,15 @@ class DatasetChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) # check segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") @@ -694,7 +708,7 @@ class DatasetChildChunkApi(DatasetApiResource): child_chunk_id_str = str(child_chunk_id) # check child chunk child_chunk = SegmentService.get_child_chunk_by_id( - child_chunk_id=child_chunk_id_str, tenant_id=current_tenant_id + child_chunk_id=child_chunk_id_str, tenant_id=current_tenant_id, session=db.session ) if not child_chunk: raise NotFound("Child chunk not found.") @@ -704,7 +718,7 @@ class DatasetChildChunkApi(DatasetApiResource): raise NotFound("Child chunk not found.") try: - SegmentService.delete_child_chunk(child_chunk, dataset) + SegmentService.delete_child_chunk(child_chunk, dataset, db.session) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) @@ -751,13 +765,15 @@ class DatasetChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # get document - document = DocumentService.get_document(dataset_id_str, document_id_str) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) # get segment - segment = SegmentService.get_segment_by_id(segment_id=segment_id_str, tenant_id=current_tenant_id) + segment = SegmentService.get_segment_by_id( + segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session + ) if not segment: raise NotFound("Segment not found.") @@ -768,7 +784,7 @@ class DatasetChildChunkApi(DatasetApiResource): child_chunk_id_str = str(child_chunk_id) # get child chunk child_chunk = SegmentService.get_child_chunk_by_id( - child_chunk_id=child_chunk_id_str, tenant_id=current_tenant_id + child_chunk_id=child_chunk_id_str, tenant_id=current_tenant_id, session=db.session ) if not child_chunk: raise NotFound("Child chunk not found.") @@ -781,7 +797,9 @@ class DatasetChildChunkApi(DatasetApiResource): payload = ChildChunkUpdatePayload.model_validate(service_api_ns.payload or {}) try: - child_chunk = SegmentService.update_child_chunk(payload.content, child_chunk, segment, document, dataset) + child_chunk = SegmentService.update_child_chunk( + payload.content, child_chunk, segment, document, dataset, db.session + ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/core/app/app_config/easy_ui_based_app/dataset/manager.py b/api/core/app/app_config/easy_ui_based_app/dataset/manager.py index be538455afb..140d4e6a2a6 100644 --- a/api/core/app/app_config/easy_ui_based_app/dataset/manager.py +++ b/api/core/app/app_config/easy_ui_based_app/dataset/manager.py @@ -9,6 +9,7 @@ from core.app.app_config.entities import ( ) from core.entities.agent_entities import PlanningStrategy from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict +from extensions.ext_database import db from models.model import AppMode, AppModelConfigDict from services.dataset_service import DatasetService @@ -256,7 +257,7 @@ class DatasetConfigManager: @classmethod def is_dataset_exists(cls, tenant_id: str, dataset_id: str) -> bool: # verify if the dataset ID exists - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, db.session) if not dataset: return False diff --git a/api/core/app/apps/pipeline/pipeline_generator.py b/api/core/app/apps/pipeline/pipeline_generator.py index 255740b86a1..cafc95d035a 100644 --- a/api/core/app/apps/pipeline/pipeline_generator.py +++ b/api/core/app/apps/pipeline/pipeline_generator.py @@ -144,7 +144,7 @@ class PipelineGenerator(BaseAppGenerator): DocumentService.check_document_creation_limits(len(datasource_info_list), features) for datasource_info in datasource_info_list: - position = DocumentService.get_documents_position(dataset.id) + position = DocumentService.get_documents_position(dataset.id, session) document = self._build_document( tenant_id=pipeline.tenant_id, dataset_id=dataset.id, diff --git a/api/core/app/features/annotation_reply/annotation_reply.py b/api/core/app/features/annotation_reply/annotation_reply.py index 0bd904811a0..520ba7b85b3 100644 --- a/api/core/app/features/annotation_reply/annotation_reply.py +++ b/api/core/app/features/annotation_reply/annotation_reply.py @@ -45,7 +45,7 @@ class AnnotationReplyFeature: embedding_model_name = collection_binding_detail.model_name dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( - embedding_provider_name, embedding_model_name, CollectionBindingType.ANNOTATION + embedding_provider_name, embedding_model_name, db.session, CollectionBindingType.ANNOTATION ) dataset = Dataset( diff --git a/api/models/dataset.py b/api/models/dataset.py index 998bc02ee85..58fec01d675 100644 --- a/api/models/dataset.py +++ b/api/models/dataset.py @@ -14,7 +14,7 @@ from uuid import uuid4 import sqlalchemy as sa from sqlalchemy import DateTime, String, func, select -from sqlalchemy.orm import Mapped, Session, mapped_column +from sqlalchemy.orm import Mapped, Session, mapped_column, scoped_session from configs import dify_config from core.rag.entities import ParentMode, Rule @@ -1670,7 +1670,7 @@ class Pipeline(TypeBase): init=False, ) - def retrieve_dataset(self, session: Session): + def retrieve_dataset(self, session: Session | scoped_session): return session.scalar(select(Dataset).where(Dataset.pipeline_id == self.id)) diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index 2dd8c533828..4dbcf372bb0 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -13,11 +13,10 @@ import sqlalchemy as sa from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator from redis.exceptions import LockNotOwnedError from sqlalchemy import ColumnElement, delete, exists, func, select, update -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, scoped_session from werkzeug.exceptions import Forbidden, NotFound from configs import dify_config -from core.db.session_factory import session_factory from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.helper.name_generator import generate_incremental_name from core.model_manager import ModelManager @@ -109,6 +108,13 @@ from tasks.sync_website_document_indexing_task import sync_website_document_inde logger = logging.getLogger(__name__) +def _session_for_helpers(session: scoped_session | Session) -> Session: + """Return a concrete SQLAlchemy session for helpers that do not accept scoped_session.""" + if isinstance(session, scoped_session): + return session() + return session + + class ProcessRulesDict(TypedDict): mode: ProcessRuleMode rules: dict[str, Any] @@ -353,9 +359,9 @@ class DatasetService: return datasets.items, datasets.total @staticmethod - def get_process_rules(dataset_id) -> ProcessRulesDict: + def get_process_rules(dataset_id, session: scoped_session | Session) -> ProcessRulesDict: # get the latest process rule - dataset_process_rule = db.session.execute( + dataset_process_rule = session.execute( select(DatasetProcessRule) .where(DatasetProcessRule.dataset_id == dataset_id) .order_by(DatasetProcessRule.created_at.desc()) @@ -411,9 +417,11 @@ class DatasetService: embedding_model_name: str | None = None, retrieval_model: RetrievalModel | None = None, summary_index_setting: dict[str, Any] | None = None, + *, + session: scoped_session | Session, ): # check if dataset name already exists - if db.session.scalar(select(Dataset).where(Dataset.name == name, Dataset.tenant_id == tenant_id).limit(1)): + if session.scalar(select(Dataset).where(Dataset.name == name, Dataset.tenant_id == tenant_id).limit(1)): raise DatasetNameDuplicateError(f"Dataset with name {name} already exists.") embedding_model = None if indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -459,8 +467,8 @@ class DatasetService: dataset.provider = provider if summary_index_setting is not None: dataset.summary_index_setting = summary_index_setting - db.session.add(dataset) - db.session.flush() + session.add(dataset) + session.flush() if provider == "external" and external_knowledge_api_id: external_knowledge_api = ExternalDatasetService.get_external_knowledge_api( @@ -477,9 +485,9 @@ class DatasetService: external_knowledge_id=external_knowledge_id, created_by=account.id, ) - db.session.add(external_knowledge_binding) + session.add(external_knowledge_binding) - db.session.commit() + session.commit() enterprise_rbac_service.try_sync_creator_access_policy_member_bindings( tenant_id, account.id, @@ -492,10 +500,11 @@ class DatasetService: def create_empty_rag_pipeline_dataset( tenant_id: str, rag_pipeline_dataset_create_entity: RagPipelineDatasetCreateEntity, + session: scoped_session | Session, ): if rag_pipeline_dataset_create_entity.name: # check if dataset name already exists - if db.session.scalar( + if session.scalar( select(Dataset) .where(Dataset.name == rag_pipeline_dataset_create_entity.name, Dataset.tenant_id == tenant_id) .limit(1) @@ -505,7 +514,7 @@ class DatasetService: ) else: # generate a random name as Untitled 1 2 3 ... - datasets = db.session.scalars(select(Dataset).where(Dataset.tenant_id == tenant_id)).all() + datasets = session.scalars(select(Dataset).where(Dataset.tenant_id == tenant_id)).all() names = [dataset.name for dataset in datasets] rag_pipeline_dataset_create_entity.name = generate_incremental_name( names, @@ -519,8 +528,8 @@ class DatasetService: description=rag_pipeline_dataset_create_entity.description, created_by=current_user.id, ) - db.session.add(pipeline) - db.session.flush() + session.add(pipeline) + session.flush() dataset = Dataset( tenant_id=tenant_id, @@ -534,13 +543,13 @@ class DatasetService: maintainer=current_user.id, pipeline_id=pipeline.id, ) - db.session.add(dataset) - db.session.commit() + session.add(dataset) + session.commit() return dataset @staticmethod - def get_dataset(dataset_id) -> Dataset | None: - dataset: Dataset | None = db.session.get(Dataset, dataset_id) + def get_dataset(dataset_id, session: scoped_session | Session) -> Dataset | None: + dataset: Dataset | None = session.get(Dataset, dataset_id) return dataset @staticmethod @@ -622,7 +631,7 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def update_dataset(dataset_id, data, user): + def update_dataset(dataset_id, data, user, session: scoped_session | Session): """ Update dataset configuration and settings. @@ -639,7 +648,7 @@ class DatasetService: NoPermissionError: If user lacks permission to update the dataset """ # Retrieve and validate dataset existence - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, session) if not dataset: raise ValueError("Dataset not found") # check if dataset name is exists @@ -648,21 +657,22 @@ class DatasetService: tenant_id=dataset.tenant_id, dataset_id=dataset_id, name=data.get("name", dataset.name), + session=session, ): raise ValueError("Dataset name already exists") # Verify user has permission to update this dataset - DatasetService.check_dataset_permission(dataset, user, db.session) + DatasetService.check_dataset_permission(dataset, user, session) # Handle external dataset updates if dataset.provider == "external": - return DatasetService._update_external_dataset(dataset, data, user) + return DatasetService._update_external_dataset(dataset, data, user, session) else: - return DatasetService._update_internal_dataset(dataset, data, user) + return DatasetService._update_internal_dataset(dataset, data, user, session) @staticmethod - def _has_dataset_same_name(tenant_id: str, dataset_id: str, name: str): - dataset = db.session.scalar( + def _has_dataset_same_name(tenant_id: str, dataset_id: str, name: str, session: scoped_session | Session): + dataset = session.scalar( select(Dataset) .where( Dataset.id != dataset_id, @@ -674,7 +684,7 @@ class DatasetService: return dataset is not None @staticmethod - def _update_external_dataset(dataset, data, user): + def _update_external_dataset(dataset, data, user, session: scoped_session | Session): """ Update external dataset configuration. @@ -718,18 +728,22 @@ class DatasetService: # Update metadata fields dataset.updated_by = user.id if user else None dataset.updated_at = naive_utc_now() - db.session.add(dataset) + session.add(dataset) # Update external knowledge binding - DatasetService._update_external_knowledge_binding(dataset.id, external_knowledge_id, external_knowledge_api_id) + DatasetService._update_external_knowledge_binding( + dataset.id, external_knowledge_id, external_knowledge_api_id, session + ) # Commit changes to database - db.session.commit() + session.commit() return dataset @staticmethod - def _update_external_knowledge_binding(dataset_id, external_knowledge_id, external_knowledge_api_id): + def _update_external_knowledge_binding( + dataset_id, external_knowledge_id, external_knowledge_api_id, session: scoped_session | Session + ): """ Update external knowledge binding configuration. @@ -738,25 +752,24 @@ class DatasetService: external_knowledge_id: External knowledge identifier external_knowledge_api_id: External knowledge API identifier """ - with sessionmaker(db.engine).begin() as session: - external_knowledge_binding = session.scalar( - select(ExternalKnowledgeBindings).where(ExternalKnowledgeBindings.dataset_id == dataset_id).limit(1) - ) + external_knowledge_binding = session.scalar( + select(ExternalKnowledgeBindings).where(ExternalKnowledgeBindings.dataset_id == dataset_id).limit(1) + ) - if not external_knowledge_binding: - raise ValueError("External knowledge binding not found.") + if not external_knowledge_binding: + raise ValueError("External knowledge binding not found.") - # Update binding if values have changed - if ( - external_knowledge_binding.external_knowledge_id != external_knowledge_id - or external_knowledge_binding.external_knowledge_api_id != external_knowledge_api_id - ): - external_knowledge_binding.external_knowledge_id = external_knowledge_id - external_knowledge_binding.external_knowledge_api_id = external_knowledge_api_id - session.add(external_knowledge_binding) + # Update binding if values have changed + if ( + external_knowledge_binding.external_knowledge_id != external_knowledge_id + or external_knowledge_binding.external_knowledge_api_id != external_knowledge_api_id + ): + external_knowledge_binding.external_knowledge_id = external_knowledge_id + external_knowledge_binding.external_knowledge_api_id = external_knowledge_api_id + session.add(external_knowledge_binding) @staticmethod - def _update_internal_dataset(dataset, data, user): + def _update_internal_dataset(dataset, data, user, session: scoped_session | Session): """ Update internal dataset configuration. @@ -778,7 +791,7 @@ class DatasetService: filtered_data = {k: v for k, v in data.items() if v is not None or k == "description"} # Handle indexing technique changes and embedding model updates - action = DatasetService._handle_indexing_technique_change(dataset, data, filtered_data) + action = DatasetService._handle_indexing_technique_change(dataset, data, filtered_data, session) # Add metadata fields filtered_data["updated_by"] = user.id @@ -794,14 +807,14 @@ class DatasetService: filtered_data["icon_info"] = data.get("icon_info") # Update dataset in database - db.session.execute(update(Dataset).where(Dataset.id == dataset.id).values(**filtered_data)) - db.session.commit() + session.execute(update(Dataset).where(Dataset.id == dataset.id).values(**filtered_data)) + session.commit() # Reload dataset to get updated values - db.session.refresh(dataset) + session.refresh(dataset) # update pipeline knowledge base node data - DatasetService._update_pipeline_knowledge_base_node_data(dataset, user.id) + DatasetService._update_pipeline_knowledge_base_node_data(dataset, user.id, session) # Trigger vector index task if indexing technique changed if action: @@ -822,14 +835,16 @@ class DatasetService: return dataset @staticmethod - def _update_pipeline_knowledge_base_node_data(dataset: Dataset, updata_user_id: str): + def _update_pipeline_knowledge_base_node_data( + dataset: Dataset, updata_user_id: str, session: scoped_session | Session + ): """ Update pipeline knowledge base node data. """ if dataset.runtime_mode != DatasetRuntimeMode.RAG_PIPELINE: return - pipeline = db.session.get(Pipeline, dataset.pipeline_id) + pipeline = session.get(Pipeline, dataset.pipeline_id) if not pipeline: return @@ -887,25 +902,25 @@ class DatasetService: marked_name="", marked_comment="", ) - db.session.add(workflow) + session.add(workflow) # Update draft workflow if draft_workflow: updated_graph = update_knowledge_nodes(draft_workflow.graph) if updated_graph != draft_workflow.graph: draft_workflow.graph = updated_graph - db.session.add(draft_workflow) + session.add(draft_workflow) # Commit all changes in one transaction - db.session.commit() + session.commit() except Exception: logging.exception("Failed to update pipeline knowledge base node data") - db.session.rollback() + session.rollback() raise @staticmethod - def _handle_indexing_technique_change(dataset, data, filtered_data): + def _handle_indexing_technique_change(dataset, data, filtered_data, session: scoped_session | Session): """ Handle changes in indexing technique and configure embedding models accordingly. @@ -913,6 +928,7 @@ class DatasetService: dataset: Current dataset object data: Update data dictionary filtered_data: Filtered update data + session: SQLAlchemy session used for embedding collection binding lookups Returns: str: Action to perform ('add', 'remove', 'update', or None) @@ -928,21 +944,24 @@ class DatasetService: return "remove" elif data["indexing_technique"] == IndexTechniqueType.HIGH_QUALITY: # Configure embedding model for high quality mode - DatasetService._configure_embedding_model_for_high_quality(data, filtered_data) + DatasetService._configure_embedding_model_for_high_quality(data, filtered_data, session) return "add" else: # Handle embedding model updates when indexing technique remains the same - return DatasetService._handle_embedding_model_update_when_technique_unchanged(dataset, data, filtered_data) + return DatasetService._handle_embedding_model_update_when_technique_unchanged( + dataset, data, filtered_data, session + ) return None @staticmethod - def _configure_embedding_model_for_high_quality(data, filtered_data): + def _configure_embedding_model_for_high_quality(data, filtered_data, session: scoped_session | Session): """ Configure embedding model settings for high quality indexing. Args: data: Update data dictionary filtered_data: Filtered update data to modify + session: SQLAlchemy session used for embedding collection binding lookups """ # assert isinstance(current_user, Account) and current_user.current_tenant_id is not None try: @@ -961,6 +980,7 @@ class DatasetService: dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( embedding_model.provider, embedding_model_name, + session, ) filtered_data["collection_binding_id"] = dataset_collection_binding.id except LLMBadRequestError: @@ -971,7 +991,9 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def _handle_embedding_model_update_when_technique_unchanged(dataset, data, filtered_data): + def _handle_embedding_model_update_when_technique_unchanged( + dataset, data, filtered_data, session: scoped_session | Session + ): """ Handle embedding model updates when indexing technique remains the same. @@ -979,6 +1001,7 @@ class DatasetService: dataset: Current dataset object data: Update data dictionary filtered_data: Filtered update data to modify + session: SQLAlchemy session used for embedding collection binding lookups Returns: str: Action to perform ('update' or None) @@ -993,7 +1016,7 @@ class DatasetService: DatasetService._preserve_existing_embedding_settings(dataset, filtered_data) return None else: - return DatasetService._update_embedding_model_settings(dataset, data, filtered_data) + return DatasetService._update_embedding_model_settings(dataset, data, filtered_data, session) @staticmethod def _preserve_existing_embedding_settings(dataset, filtered_data): @@ -1019,7 +1042,7 @@ class DatasetService: del filtered_data["embedding_model"] @staticmethod - def _update_embedding_model_settings(dataset, data, filtered_data): + def _update_embedding_model_settings(dataset, data, filtered_data, session: scoped_session | Session): """ Update embedding model settings with new values. @@ -1027,6 +1050,7 @@ class DatasetService: dataset: Current dataset object data: Update data dictionary filtered_data: Filtered update data to modify + session: SQLAlchemy session used for embedding collection binding lookups Returns: str: Action to perform ('update' or None) @@ -1042,7 +1066,7 @@ class DatasetService: # Only update if values are different if current_provider_str != new_provider_str or data["embedding_model"] != dataset.embedding_model: - DatasetService._apply_new_embedding_settings(dataset, data, filtered_data) + DatasetService._apply_new_embedding_settings(dataset, data, filtered_data, session) return "update" except LLMBadRequestError: raise ValueError( @@ -1053,7 +1077,7 @@ class DatasetService: return None @staticmethod - def _apply_new_embedding_settings(dataset, data, filtered_data): + def _apply_new_embedding_settings(dataset, data, filtered_data, session: scoped_session | Session): """ Apply new embedding model settings to the dataset. @@ -1061,6 +1085,7 @@ class DatasetService: dataset: Current dataset object data: Update data dictionary filtered_data: Filtered update data to modify + session: SQLAlchemy session used for embedding collection binding lookups """ # assert isinstance(current_user, Account) and current_user.current_tenant_id is not None @@ -1096,6 +1121,7 @@ class DatasetService: dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( embedding_model.provider, embedding_model_name, + session, ) filtered_data["collection_binding_id"] = dataset_collection_binding.id @@ -1177,6 +1203,7 @@ class DatasetService: dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( embedding_model.provider, embedding_model_name, + session, ) dataset.collection_binding_id = dataset_collection_binding.id elif knowledge_configuration.indexing_technique == IndexTechniqueType.ECONOMY: @@ -1213,6 +1240,7 @@ class DatasetService: dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( embedding_model.provider, embedding_model_name, + session, ) is_multimodal = DatasetService.check_is_multimodal_model( current_user.current_tenant_id, @@ -1276,6 +1304,7 @@ class DatasetService: DatasetCollectionBindingService.get_dataset_collection_binding( embedding_model.provider, embedding_model_name, + session, ) ) dataset.collection_binding_id = dataset_collection_binding.id @@ -1305,24 +1334,24 @@ class DatasetService: deal_dataset_index_update_task.delay(dataset.id, action) @staticmethod - def delete_dataset(dataset_id, user): - dataset = DatasetService.get_dataset(dataset_id) + def delete_dataset(dataset_id, user, session: scoped_session | Session): + dataset = DatasetService.get_dataset(dataset_id, session) if dataset is None: return False - DatasetService.check_dataset_permission(dataset, user, db.session) + DatasetService.check_dataset_permission(dataset, user, session) dataset_was_deleted.send(dataset) - db.session.delete(dataset) - db.session.commit() + session.delete(dataset) + session.commit() return True @staticmethod - def dataset_use_check(dataset_id) -> bool: + def dataset_use_check(dataset_id, session: scoped_session | Session) -> bool: stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_id)) - return db.session.execute(stmt).scalar_one() + return session.execute(stmt).scalar_one() @staticmethod def check_dataset_permission(dataset, user, session: scoped_session | Session): @@ -1347,7 +1376,9 @@ class DatasetService: raise NoPermissionError("You do not have permission to access this dataset.") @staticmethod - def check_dataset_operator_permission(user: Account | None = None, dataset: Dataset | None = None): + def check_dataset_operator_permission( + user: Account | None = None, dataset: Dataset | None = None, *, session: scoped_session | Session + ): if not dataset: raise ValueError("Dataset not found") @@ -1362,7 +1393,7 @@ class DatasetService: elif dataset.permission == DatasetPermissionEnum.PARTIAL_TEAM: if not any( dp.dataset_id == dataset.id - for dp in db.session.scalars( + for dp in session.scalars( select(DatasetPermission).where(DatasetPermission.account_id == user.id) ).all() ): @@ -1377,16 +1408,16 @@ class DatasetService: return dataset_queries.items, dataset_queries.total @staticmethod - def get_related_apps(dataset_id: str): - return db.session.scalars( + def get_related_apps(dataset_id: str, session: scoped_session | Session): + return session.scalars( select(AppDatasetJoin) .where(AppDatasetJoin.dataset_id == dataset_id) .order_by(AppDatasetJoin.created_at.desc()) ).all() @staticmethod - def update_dataset_api_status(dataset_id: str, status: bool): - dataset = DatasetService.get_dataset(dataset_id) + def update_dataset_api_status(dataset_id: str, status: bool, session: scoped_session | Session): + dataset = DatasetService.get_dataset(dataset_id, session) if dataset is None: raise NotFound("Dataset not found.") dataset.enable_api = status @@ -1394,10 +1425,10 @@ class DatasetService: raise ValueError("Current user or current user id not found") dataset.updated_by = current_user.id dataset.updated_at = naive_utc_now() - db.session.commit() + session.commit() @staticmethod - def get_dataset_auto_disable_logs(dataset_id: str) -> AutoDisableLogsDict: + def get_dataset_auto_disable_logs(dataset_id: str, session: scoped_session | Session) -> AutoDisableLogsDict: assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None features = FeatureService.get_features(current_user.current_tenant_id, exclude_vector_space=True) @@ -1408,7 +1439,7 @@ class DatasetService: } # get recent 30 days auto disable logs start_date = datetime.datetime.now() - datetime.timedelta(days=30) - dataset_auto_disable_logs = db.session.scalars( + dataset_auto_disable_logs = session.scalars( select(DatasetAutoDisableLog).where( DatasetAutoDisableLog.dataset_id == dataset_id, DatasetAutoDisableLog.created_at >= start_date, @@ -1596,9 +1627,12 @@ class DocumentService: } @staticmethod - def get_document(dataset_id: str, document_id: str | None = None) -> Document | None: + def get_document( + dataset_id: str, document_id: str | None = None, *, session: scoped_session | Session + ) -> Document | None: + """Fetch a document by id within a dataset using the caller-provided session.""" if document_id: - document = db.session.scalar( + document = session.scalar( select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) ) return document @@ -1606,13 +1640,15 @@ class DocumentService: return None @staticmethod - def get_documents_by_ids(dataset_id: str, document_ids: Sequence[str]) -> Sequence[Document]: + def get_documents_by_ids( + dataset_id: str, document_ids: Sequence[str], session: scoped_session | Session + ) -> Sequence[Document]: """Fetch documents for a dataset in a single batch query.""" if not document_ids: return [] document_id_list: list[str] = [str(document_id) for document_id in document_ids] # Fetch all requested documents in one query to avoid N+1 lookups. - documents: Sequence[Document] = db.session.scalars( + documents: Sequence[Document] = session.scalars( select(Document).where( Document.dataset_id == dataset_id, Document.id.in_(document_id_list), @@ -1621,7 +1657,12 @@ class DocumentService: return documents @staticmethod - def update_documents_need_summary(dataset_id: str, document_ids: Sequence[str], need_summary: bool = True) -> int: + def update_documents_need_summary( + dataset_id: str, + document_ids: Sequence[str], + session: scoped_session | Session, + need_summary: bool = True, + ) -> int: """ Update need_summary field for multiple documents. @@ -1631,6 +1672,7 @@ class DocumentService: Args: dataset_id: Dataset ID document_ids: List of document IDs to update + session: SQLAlchemy session used for the update need_summary: Value to set for need_summary field (default: True) Returns: @@ -1641,33 +1683,32 @@ class DocumentService: document_id_list: list[str] = [str(document_id) for document_id in document_ids] - with session_factory.create_session() as session: - result = session.execute( - update(Document) - .where( - Document.id.in_(document_id_list), - Document.dataset_id == dataset_id, - Document.doc_form != IndexStructureType.QA_INDEX, # Skip qa_model documents - ) - .values(need_summary=need_summary) - .execution_options(synchronize_session=False) + result = session.execute( + update(Document) + .where( + Document.id.in_(document_id_list), + Document.dataset_id == dataset_id, + Document.doc_form != IndexStructureType.QA_INDEX, # Skip qa_model documents ) - updated_count = result.rowcount # type: ignore[union-attr,attr-defined] - session.commit() - logger.info( - "Updated need_summary to %s for %d documents in dataset %s", - need_summary, - updated_count, - dataset_id, - ) - return updated_count + .values(need_summary=need_summary) + .execution_options(synchronize_session=False) + ) + updated_count = result.rowcount # type: ignore[union-attr,attr-defined] + session.commit() + logger.info( + "Updated need_summary to %s for %d documents in dataset %s", + need_summary, + updated_count, + dataset_id, + ) + return updated_count @staticmethod - def get_document_download_url(document: Document) -> str: + def get_document_download_url(document: Document, session: scoped_session | Session) -> str: """ Return a signed download URL for an upload-file document. """ - upload_file = DocumentService._get_upload_file_for_upload_file_document(document) + upload_file = DocumentService._get_upload_file_for_upload_file_document(document, session) return file_helpers.get_signed_file_url(upload_file_id=upload_file.id, as_attachment=True) @staticmethod @@ -1721,15 +1762,16 @@ class DocumentService: document_ids: Sequence[str], tenant_id: str, current_user: Account, + session: scoped_session | Session, ) -> tuple[list[UploadFile], str]: """ Resolve upload files for batch ZIP downloads and generate a client-visible filename. """ - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, session) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, session) except NoPermissionError as e: raise Forbidden(str(e)) @@ -1737,6 +1779,7 @@ class DocumentService: dataset_id=dataset_id, document_ids=document_ids, tenant_id=tenant_id, + session=session, ) upload_files = [upload_files_by_document_id[document_id] for document_id in document_ids] download_name = DocumentService._generate_document_batch_download_zip_filename() @@ -1770,7 +1813,7 @@ class DocumentService: return str(upload_file_id) @staticmethod - def _get_upload_file_for_upload_file_document(document: Document) -> UploadFile: + def _get_upload_file_for_upload_file_document(document: Document, session: scoped_session | Session) -> UploadFile: """ Load the `UploadFile` row for an upload-file document. """ @@ -1779,7 +1822,9 @@ class DocumentService: invalid_source_message="Document does not have an uploaded file to download.", missing_file_message="Uploaded file not found.", ) - upload_files_by_id = FileService.get_upload_files_by_ids(db.session(), document.tenant_id, [upload_file_id]) + upload_files_by_id = FileService.get_upload_files_by_ids( + _session_for_helpers(session), document.tenant_id, [upload_file_id] + ) upload_file = upload_files_by_id.get(upload_file_id) if not upload_file: raise NotFound("Uploaded file not found.") @@ -1791,13 +1836,14 @@ class DocumentService: dataset_id: str, document_ids: Sequence[str], tenant_id: str, + session: scoped_session | Session, ) -> dict[str, UploadFile]: """ Batch load upload files keyed by document id for ZIP downloads. """ document_id_list: list[str] = [str(document_id) for document_id in document_ids] - documents = DocumentService.get_documents_by_ids(dataset_id, document_id_list) + documents = DocumentService.get_documents_by_ids(dataset_id, document_id_list, session) documents_by_id: dict[str, Document] = {str(document.id): document for document in documents} missing_document_ids: set[str] = set(document_id_list) - set(documents_by_id.keys()) @@ -1818,7 +1864,9 @@ class DocumentService: upload_file_ids.append(upload_file_id) upload_file_ids_by_document_id[document_id] = upload_file_id - upload_files_by_id = FileService.get_upload_files_by_ids(db.session(), tenant_id, upload_file_ids) + upload_files_by_id = FileService.get_upload_files_by_ids( + _session_for_helpers(session), tenant_id, upload_file_ids + ) missing_upload_file_ids: set[str] = set(upload_file_ids) - set(upload_files_by_id.keys()) if missing_upload_file_ids: raise NotFound("Only uploaded-file documents can be downloaded as ZIP.") @@ -1829,14 +1877,14 @@ class DocumentService: } @staticmethod - def get_document_by_id(document_id: str) -> Document | None: - document = db.session.get(Document, document_id) + def get_document_by_id(document_id: str, session: scoped_session | Session) -> Document | None: + document = session.get(Document, document_id) return document @staticmethod - def get_document_by_ids(document_ids: list[str]) -> Sequence[Document]: - documents = db.session.scalars( + def get_document_by_ids(document_ids: list[str], session: scoped_session | Session) -> Sequence[Document]: + documents = session.scalars( select(Document).where( Document.id.in_(document_ids), Document.enabled == True, @@ -1847,8 +1895,8 @@ class DocumentService: return documents @staticmethod - def get_document_by_dataset_id(dataset_id: str) -> Sequence[Document]: - documents = db.session.scalars( + def get_document_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, Document.enabled == True, @@ -1858,8 +1906,8 @@ class DocumentService: return documents @staticmethod - def get_working_documents_by_dataset_id(dataset_id: str) -> Sequence[Document]: - documents = db.session.scalars( + def get_working_documents_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, Document.enabled == True, @@ -1871,8 +1919,8 @@ class DocumentService: return documents @staticmethod - def get_error_documents_by_dataset_id(dataset_id: str) -> Sequence[Document]: - documents = db.session.scalars( + def get_error_documents_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, Document.indexing_status.in_([IndexingStatus.ERROR, IndexingStatus.PAUSED]), @@ -1881,9 +1929,9 @@ class DocumentService: return documents @staticmethod - def get_batch_documents(dataset_id: str, batch: str) -> Sequence[Document]: + def get_batch_documents(dataset_id: str, batch: str, session: scoped_session | Session) -> Sequence[Document]: assert isinstance(current_user, Account) - documents = db.session.scalars( + documents = session.scalars( select(Document).where( Document.batch == batch, Document.dataset_id == dataset_id, @@ -1894,8 +1942,8 @@ class DocumentService: return documents @staticmethod - def get_document_file_detail(file_id: str): - file_detail = db.session.get(UploadFile, file_id) + def get_document_file_detail(file_id: str, session: scoped_session | Session): + file_detail = session.get(UploadFile, file_id) return file_detail @staticmethod @@ -1906,7 +1954,7 @@ class DocumentService: return False @staticmethod - def delete_document(document): + def delete_document(document, session: scoped_session | Session): # trigger document_was_deleted signal file_id = None if document.data_source_type == DataSourceType.UPLOAD_FILE: @@ -1918,15 +1966,15 @@ class DocumentService: document.id, dataset_id=document.dataset_id, doc_form=document.doc_form, file_id=file_id ) - db.session.delete(document) - db.session.commit() + session.delete(document) + session.commit() @staticmethod - def delete_documents(dataset: Dataset, document_ids: list[str]): + def delete_documents(dataset: Dataset, document_ids: list[str], session: scoped_session | Session): # Check if document_ids is not empty to avoid WHERE false condition if not document_ids or len(document_ids) == 0: return - documents = db.session.scalars(select(Document).where(Document.id.in_(document_ids))).all() + documents = session.scalars(select(Document).where(Document.id.in_(document_ids))).all() file_ids = [ document.data_source_info_dict.get("upload_file_id", "") for document in documents @@ -1936,8 +1984,8 @@ class DocumentService: # Delete documents first, then dispatch cleanup task after commit # to avoid deadlock between main transaction and async task for document in documents: - db.session.delete(document) - db.session.commit() + session.delete(document) + session.commit() # Dispatch cleanup task after commit to avoid lock contention # Task cleans up segments, files, and vector indexes @@ -1945,14 +1993,14 @@ class DocumentService: batch_clean_document_task.delay(document_ids, dataset.id, dataset.doc_form, file_ids) @staticmethod - def rename_document(dataset_id: str, document_id: str, name: str) -> Document: + def rename_document(dataset_id: str, document_id: str, name: str, session: scoped_session | Session) -> Document: assert isinstance(current_user, Account) - dataset = DatasetService.get_dataset(dataset_id) + dataset = DatasetService.get_dataset(dataset_id, session) if not dataset: raise ValueError("Dataset not found.") - document = DocumentService.get_document(dataset_id, document_id) + document = DocumentService.get_document(dataset_id, document_id, session=session) if not document: raise ValueError("Document not found.") @@ -1967,20 +2015,20 @@ class DocumentService: document.doc_metadata = doc_metadata document.name = name - db.session.add(document) + session.add(document) if document.data_source_info_dict and "upload_file_id" in document.data_source_info_dict: - db.session.execute( + session.execute( update(UploadFile) .where(UploadFile.id == document.data_source_info_dict["upload_file_id"]) .values(name=name) ) - db.session.commit() + session.commit() return document @staticmethod - def pause_document(document): + def pause_document(document, session: scoped_session | Session): if document.indexing_status not in { IndexingStatus.WAITING, IndexingStatus.PARSING, @@ -1995,14 +2043,14 @@ class DocumentService: document.paused_by = current_user.id document.paused_at = naive_utc_now() - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() # set document paused flag indexing_cache_key = f"document_{document.id}_is_paused" redis_client.setnx(indexing_cache_key, "True") @staticmethod - def recover_document(document): + def recover_document(document, session: scoped_session | Session): if not document.is_paused: raise DocumentIndexingError() # update document to be recover @@ -2010,8 +2058,8 @@ class DocumentService: document.paused_by = None document.paused_at = None - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() # delete paused flag indexing_cache_key = f"document_{document.id}_is_paused" redis_client.delete(indexing_cache_key) @@ -2019,7 +2067,7 @@ class DocumentService: recover_document_indexing_task.delay(document.dataset_id, document.id) @staticmethod - def retry_document(dataset_id: str, documents: list[Document]): + def retry_document(dataset_id: str, documents: list[Document], session: scoped_session | Session): for document in documents: # add retry flag retry_indexing_cache_key = f"document_{document.id}_is_retried" @@ -2028,8 +2076,8 @@ class DocumentService: raise ValueError("Document is being retried, please try again later") # retry document indexing document.indexing_status = IndexingStatus.WAITING - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() redis_client.setex(retry_indexing_cache_key, 600, 1) # trigger async task @@ -2039,7 +2087,7 @@ class DocumentService: retry_document_indexing_task.delay(dataset_id, document_ids, current_user.id) @staticmethod - def sync_website_document(dataset_id: str, document: Document): + def sync_website_document(dataset_id: str, document: Document, session: scoped_session | Session): # add sync flag sync_indexing_cache_key = f"document_{document.id}_is_sync" cache_result = redis_client.get(sync_indexing_cache_key) @@ -2051,16 +2099,16 @@ class DocumentService: if data_source_info: data_source_info["mode"] = "scrape" document.data_source_info = json.dumps(data_source_info, ensure_ascii=False) - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() redis_client.setex(sync_indexing_cache_key, 600, 1) sync_website_document_indexing_task.delay(dataset_id, document.id) @staticmethod - def get_documents_position(dataset_id): - document = db.session.scalar( + def get_documents_position(dataset_id, session: scoped_session | Session): + document = session.scalar( select(Document).where(Document.dataset_id == dataset_id).order_by(Document.position.desc()).limit(1) ) if document: @@ -2075,6 +2123,8 @@ class DocumentService: account: Account | Any, dataset_process_rule: DatasetProcessRule | None = None, created_from: str = DocumentCreatedFrom.WEB, + *, + session: scoped_session | Session, ) -> tuple[list[Document], str]: # check doc_form DatasetService.check_doc_form(dataset, knowledge_config.doc_form) @@ -2126,7 +2176,7 @@ class DocumentService: dataset.embedding_model = dataset_embedding_model dataset.embedding_model_provider = dataset_embedding_model_provider dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( - dataset_embedding_model_provider, dataset_embedding_model + dataset_embedding_model_provider, dataset_embedding_model, session ) dataset.collection_binding_id = dataset_collection_binding.id if not dataset.retrieval_model: @@ -2146,7 +2196,9 @@ class DocumentService: documents = [] if knowledge_config.original_document_id: - document = DocumentService.update_document_with_dataset_id(dataset, knowledge_config, account) + document = DocumentService.update_document_with_dataset_id( + dataset, knowledge_config, account, session=session + ) documents.append(document) batch = document.batch else: @@ -2184,8 +2236,8 @@ class DocumentService: process_rule.mode, ) return [], "" - db.session.add(dataset_process_rule) - db.session.flush() + session.add(dataset_process_rule) + session.flush() else: # Fallback when no process_rule provided in knowledge_config: # 1) reuse dataset.latest_process_rule if present @@ -2198,13 +2250,13 @@ class DocumentService: rules=json.dumps(DatasetProcessRule.AUTOMATIC_RULES), created_by=account.id, ) - db.session.add(dataset_process_rule) - db.session.flush() + session.add(dataset_process_rule) + session.flush() lock_name = f"add_document_lock_dataset_id_{dataset.id}" try: with redis_client.lock(lock_name, timeout=600): assert dataset_process_rule - position = DocumentService.get_documents_position(dataset.id) + position = DocumentService.get_documents_position(dataset.id, session) document_ids = [] duplicate_document_ids = [] if knowledge_config.data_source.info_list.data_source_type == "upload_file": @@ -2212,7 +2264,7 @@ class DocumentService: raise ValueError("File source info is required") upload_file_list = knowledge_config.data_source.info_list.file_info_list.file_ids files = list( - db.session.scalars( + session.scalars( select(UploadFile).where( UploadFile.tenant_id == dataset.tenant_id, UploadFile.id.in_(upload_file_list), @@ -2224,7 +2276,7 @@ class DocumentService: file_names = [file.name for file in files] db_documents = list( - db.session.scalars( + session.scalars( select(Document).where( Document.dataset_id == dataset.id, Document.tenant_id == current_user.current_tenant_id, @@ -2249,7 +2301,7 @@ class DocumentService: document.data_source_info = json.dumps(data_source_info) document.batch = batch document.indexing_status = IndexingStatus.WAITING - db.session.add(document) + session.add(document) documents.append(document) duplicate_document_ids.append(document.id) continue @@ -2267,8 +2319,8 @@ class DocumentService: file.name, batch, ) - db.session.add(document) - db.session.flush() + session.add(document) + session.flush() document_ids.append(document.id) documents.append(document) position += 1 @@ -2279,7 +2331,7 @@ class DocumentService: exist_page_ids = [] exist_document = {} documents = list( - db.session.scalars( + session.scalars( select(Document).where( Document.dataset_id == dataset.id, Document.tenant_id == current_user.current_tenant_id, @@ -2319,8 +2371,8 @@ class DocumentService: truncated_page_name, batch, ) - db.session.add(document) - db.session.flush() + session.add(document) + session.flush() document_ids.append(document.id) documents.append(document) position += 1 @@ -2359,12 +2411,12 @@ class DocumentService: document_name, batch, ) - db.session.add(document) - db.session.flush() + session.add(document) + session.flush() document_ids.append(document.id) documents.append(document) position += 1 - db.session.commit() + session.commit() # trigger async task if document_ids: @@ -2486,8 +2538,8 @@ class DocumentService: # f"Invalid process rule mode: {process_rule.mode}, can not find dataset process rule" # ) # return - # db.session.add(dataset_process_rule) - # db.session.commit() + # session.add(dataset_process_rule) + # session.commit() # lock_name = "add_document_lock_dataset_id_{}".format(dataset.id) # with redis_client.lock(lock_name, timeout=600): # position = DocumentService.get_documents_position(dataset.id) @@ -2497,7 +2549,7 @@ class DocumentService: # upload_file_list = knowledge_config.data_source.info_list.file_info_list.file_ids # for file_id in upload_file_list: # file = ( - # db.session.query(UploadFile) + # session.query(UploadFile) # .filter(UploadFile.tenant_id == dataset.tenant_id, UploadFile.id == file_id) # .first() # ) @@ -2528,7 +2580,7 @@ class DocumentService: # document.data_source_info = json.dumps(data_source_info) # document.batch = batch # document.indexing_status = "waiting" - # db.session.add(document) + # session.add(document) # documents.append(document) # duplicate_document_ids.append(document.id) # continue @@ -2545,8 +2597,8 @@ class DocumentService: # file_name, # batch, # ) - # db.session.add(document) - # db.session.flush() + # session.add(document) + # session.flush() # document_ids.append(document.id) # documents.append(document) # position += 1 @@ -2602,8 +2654,8 @@ class DocumentService: # truncated_page_name, # batch, # ) - # db.session.add(document) - # db.session.flush() + # session.add(document) + # session.flush() # document_ids.append(document.id) # documents.append(document) # position += 1 @@ -2642,12 +2694,12 @@ class DocumentService: # document_name, # batch, # ) - # db.session.add(document) - # db.session.flush() + # session.add(document) + # session.flush() # document_ids.append(document.id) # documents.append(document) # position += 1 - # db.session.commit() + # session.commit() # # trigger async task # if document_ids: @@ -2728,11 +2780,11 @@ class DocumentService: return document @staticmethod - def get_tenant_documents_count(): + def get_tenant_documents_count(session: scoped_session | Session): assert isinstance(current_user, Account) documents_count = ( - db.session.scalar( + session.scalar( select(func.count(Document.id)).where( Document.completed_at.isnot(None), Document.enabled == True, @@ -2751,11 +2803,13 @@ class DocumentService: account: Account, dataset_process_rule: DatasetProcessRule | None = None, created_from: str = DocumentCreatedFrom.WEB, + *, + session: scoped_session | Session, ): assert isinstance(current_user, Account) DatasetService.check_dataset_model_setting(dataset) - document = DocumentService.get_document(dataset.id, document_data.original_document_id) + document = DocumentService.get_document(dataset.id, document_data.original_document_id, session=session) if document is None: raise NotFound("Document not found") if document.display_status != "available": @@ -2778,8 +2832,8 @@ class DocumentService: created_by=account.id, ) if dataset_process_rule is not None: - db.session.add(dataset_process_rule) - db.session.commit() + session.add(dataset_process_rule) + session.commit() document.dataset_process_rule_id = dataset_process_rule.id # update document data source if document_data.data_source: @@ -2790,7 +2844,7 @@ class DocumentService: raise ValueError("No file info list found.") upload_file_list = document_data.data_source.info_list.file_info_list.file_ids for file_id in upload_file_list: - file = db.session.scalar( + file = session.scalar( select(UploadFile) .where(UploadFile.tenant_id == dataset.tenant_id, UploadFile.id == file_id) .limit(1) @@ -2810,7 +2864,7 @@ class DocumentService: notion_info_list = document_data.data_source.info_list.notion_info_list for notion_info in notion_info_list: workspace_id = notion_info.workspace_id - data_source_binding = db.session.scalar( + data_source_binding = session.scalar( select(DataSourceOauthBinding) .where( sa.and_( @@ -2861,22 +2915,24 @@ class DocumentService: document.updated_at = naive_utc_now() document.created_from = created_from document.doc_form = IndexStructureType(document_data.doc_form) - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() # update document segment - db.session.execute( + session.execute( update(DocumentSegment) .where(DocumentSegment.document_id == document.id) .values(status=SegmentStatus.RE_SEGMENT) ) - db.session.commit() + session.commit() # trigger async task document_indexing_update_task.delay(document.dataset_id, document.id) return document @staticmethod - def save_document_without_dataset_id(tenant_id: str, knowledge_config: KnowledgeConfig, account: Account): + def save_document_without_dataset_id( + tenant_id: str, knowledge_config: KnowledgeConfig, account: Account, session: scoped_session | Session + ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None assert knowledge_config.data_source @@ -2911,6 +2967,7 @@ class DocumentService: dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( knowledge_config.embedding_model_provider, knowledge_config.embedding_model, + session, ) dataset_collection_binding_id = dataset_collection_binding.id if knowledge_config.retrieval_model: @@ -2939,16 +2996,18 @@ class DocumentService: is_multimodal=knowledge_config.is_multimodal, ) - db.session.add(dataset) - db.session.flush() + session.add(dataset) + session.flush() - documents, batch = DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account) + documents, batch = DocumentService.save_document_with_dataset_id( + dataset, knowledge_config, account, session=session + ) cut_length = 18 cut_name = documents[0].name[:cut_length] dataset.name = cut_name + "..." dataset.description = "useful for when you want to answer queries about the " + documents[0].name - db.session.commit() + session.commit() return dataset, documents, batch @@ -3057,7 +3116,11 @@ class DocumentService: @staticmethod def batch_update_document_status( - dataset: Dataset, document_ids: list[str], action: Literal["enable", "disable", "archive", "un_archive"], user + dataset: Dataset, + document_ids: list[str], + action: Literal["enable", "disable", "archive", "un_archive"], + user, + session: scoped_session | Session, ): """ Batch update document status. @@ -3084,7 +3147,7 @@ class DocumentService: # First pass: validate all documents and prepare updates for document_id in document_ids: - document = DocumentService.get_document(dataset.id, document_id) + document = DocumentService.get_document(dataset.id, document_id, session=session) if not document: continue @@ -3110,13 +3173,13 @@ class DocumentService: for field, value in updates.items(): setattr(document, field, value) - db.session.add(document) + session.add(document) # Batch commit all changes - db.session.commit() + session.commit() except Exception as e: # Rollback on any error - db.session.rollback() + session.rollback() raise e # Execute async tasks and set Redis cache after successful commit # propagation_error is used to capture any errors for submitting async task execution @@ -3264,7 +3327,9 @@ class SegmentService: raise ValueError(f"Exceeded maximum attachment limit of {single_chunk_attachment_limit}") @classmethod - def create_segment(cls, args: dict[str, Any], document: Document, dataset: Dataset): + def create_segment( + cls, args: dict[str, Any], document: Document, dataset: Dataset, session: scoped_session | Session + ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3285,7 +3350,7 @@ class SegmentService: lock_name = f"add_segment_lock_document_id_{document.id}" try: with redis_client.lock(lock_name, timeout=600): - max_position = db.session.scalar( + max_position = session.scalar( select(func.max(DocumentSegment.position)).where(DocumentSegment.document_id == document.id) ) segment_document = DocumentSegment( @@ -3307,12 +3372,12 @@ class SegmentService: segment_document.word_count += len(args["answer"]) segment_document.answer = args["answer"] - db.session.add(segment_document) + session.add(segment_document) # update document word count assert document.word_count is not None document.word_count += segment_document.word_count - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() if args["attachment_ids"]: for attachment_id in args["attachment_ids"]: @@ -3323,8 +3388,8 @@ class SegmentService: segment_id=segment_document.id, attachment_id=attachment_id, ) - db.session.add(binding) - db.session.commit() + session.add(binding) + session.commit() # save vector index try: @@ -3337,14 +3402,16 @@ class SegmentService: segment_document.disabled_at = naive_utc_now() segment_document.status = SegmentStatus.ERROR segment_document.error = str(e) - db.session.commit() - segment = db.session.get(DocumentSegment, segment_document.id) + session.commit() + segment = session.get(DocumentSegment, segment_document.id) return segment except LockNotOwnedError: pass @classmethod - def multi_create_segment(cls, segments: list, document: Document, dataset: Dataset): + def multi_create_segment( + cls, segments: list, document: Document, dataset: Dataset, session: scoped_session | Session + ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3361,7 +3428,7 @@ class SegmentService: model_type=ModelType.TEXT_EMBEDDING, model=dataset.embedding_model, ) - max_position = db.session.scalar( + max_position = session.scalar( select(func.max(DocumentSegment.position)).where(DocumentSegment.document_id == document.id) ) pre_segment_data_list = [] @@ -3402,7 +3469,7 @@ class SegmentService: segment_document.answer = segment_item["answer"] segment_document.word_count += len(segment_item["answer"]) increment_word_count += segment_document.word_count - db.session.add(segment_document) + session.add(segment_document) segment_data_list.append(segment_document) position += 1 @@ -3414,7 +3481,7 @@ class SegmentService: # update document word count assert document.word_count is not None document.word_count += increment_word_count - db.session.add(document) + session.add(document) try: # save vector index VectorService.create_segments_vector( @@ -3427,13 +3494,20 @@ class SegmentService: segment_document.disabled_at = naive_utc_now() segment_document.status = SegmentStatus.ERROR segment_document.error = str(e) - db.session.commit() + session.commit() return segment_data_list except LockNotOwnedError: pass @classmethod - def update_segment(cls, args: SegmentUpdateArgs, segment: DocumentSegment, document: Document, dataset: Dataset): + def update_segment( + cls, + args: SegmentUpdateArgs, + segment: DocumentSegment, + document: Document, + dataset: Dataset, + session: scoped_session | Session, + ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3448,8 +3522,8 @@ class SegmentService: segment.enabled = action segment.disabled_at = naive_utc_now() segment.disabled_by = current_user.id - db.session.add(segment) - db.session.commit() + session.add(segment) + session.commit() # Set cache to prevent indexing the same segment multiple times redis_client.setex(indexing_cache_key, 600, 1) disable_segment_from_index_task.delay(segment.id) @@ -3477,13 +3551,13 @@ class SegmentService: segment.enabled = True segment.disabled_at = None segment.disabled_by = None - db.session.add(segment) - db.session.commit() + session.add(segment) + session.commit() # update document word count if word_count_change != 0: assert document.word_count is not None document.word_count = max(0, document.word_count + word_count_change) - db.session.add(document) + session.add(document) # update segment index task if document.doc_form == IndexStructureType.PARENT_CHILD_INDEX and args.regenerate_child_chunks: # regenerate child chunks @@ -3507,7 +3581,7 @@ class SegmentService: else: raise ValueError("The knowledge base index technique is not high quality!") # get the process rule - processing_rule = db.session.get(DatasetProcessRule, document.dataset_process_rule_id) + processing_rule = session.get(DatasetProcessRule, document.dataset_process_rule_id) if processing_rule: VectorService.generate_child_chunks( segment, document, dataset, embedding_model_instance, processing_rule, True @@ -3525,7 +3599,7 @@ class SegmentService: # Query existing summary from database from models.dataset import DocumentSegmentSummary - existing_summary = db.session.scalar( + existing_summary = session.scalar( select(DocumentSegmentSummary) .where( DocumentSegmentSummary.chunk_id == segment.id, @@ -3583,9 +3657,9 @@ class SegmentService: if word_count_change != 0: assert document.word_count is not None document.word_count = max(0, document.word_count + word_count_change) - db.session.add(document) - db.session.add(segment) - db.session.commit() + session.add(document) + session.add(segment) + session.commit() if document.doc_form == IndexStructureType.PARENT_CHILD_INDEX and args.regenerate_child_chunks: # get embedding model instance if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -3607,7 +3681,7 @@ class SegmentService: else: raise ValueError("The knowledge base index technique is not high quality!") # get the process rule - processing_rule = db.session.get(DatasetProcessRule, document.dataset_process_rule_id) + processing_rule = session.get(DatasetProcessRule, document.dataset_process_rule_id) if processing_rule: VectorService.generate_child_chunks( segment, document, dataset, embedding_model_instance, processing_rule, True @@ -3619,7 +3693,7 @@ class SegmentService: if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: from models.dataset import DocumentSegmentSummary - existing_summary = db.session.scalar( + existing_summary = session.scalar( select(DocumentSegmentSummary) .where( DocumentSegmentSummary.chunk_id == segment.id, @@ -3690,14 +3764,16 @@ class SegmentService: segment.disabled_at = naive_utc_now() segment.status = SegmentStatus.ERROR segment.error = str(e) - db.session.commit() - new_segment = db.session.get(DocumentSegment, segment.id) + session.commit() + new_segment = session.get(DocumentSegment, segment.id) if not new_segment: raise ValueError("new_segment is not found") return new_segment @classmethod - def delete_segment(cls, segment: DocumentSegment, document: Document, dataset: Dataset): + def delete_segment( + cls, segment: DocumentSegment, document: Document, dataset: Dataset, session: scoped_session | Session + ): indexing_cache_key = f"segment_{segment.id}_delete_indexing" cache_result = redis_client.get(indexing_cache_key) if cache_result is not None: @@ -3712,7 +3788,7 @@ class SegmentService: child_node_ids = [] if segment.index_node_id: child_node_ids = list( - db.session.scalars( + session.scalars( select(ChildChunk.index_node_id).where( ChildChunk.segment_id == segment.id, ChildChunk.dataset_id == dataset.id, @@ -3724,20 +3800,22 @@ class SegmentService: [segment.index_node_id], dataset.id, document.id, [segment.id], child_node_ids ) - db.session.delete(segment) + session.delete(segment) # update document word count assert document.word_count is not None document.word_count -= segment.word_count - db.session.add(document) - db.session.commit() + session.add(document) + session.commit() @classmethod - def delete_segments(cls, segment_ids: list, document: Document, dataset: Dataset): + def delete_segments( + cls, segment_ids: list, document: Document, dataset: Dataset, session: scoped_session | Session + ): assert current_user is not None # Check if segment_ids is not empty to avoid WHERE false condition if not segment_ids or len(segment_ids) == 0: return - segments_info = db.session.execute( + segments_info = session.execute( select(DocumentSegment.index_node_id, DocumentSegment.id, DocumentSegment.word_count).where( DocumentSegment.id.in_(segment_ids), DocumentSegment.dataset_id == dataset.id, @@ -3758,7 +3836,7 @@ class SegmentService: if index_node_ids: child_node_ids = [ nid - for nid in db.session.scalars( + for nid in session.scalars( select(ChildChunk.index_node_id).where( ChildChunk.segment_id.in_(segment_db_ids), ChildChunk.dataset_id == dataset.id, @@ -3778,15 +3856,20 @@ class SegmentService: else: document.word_count = max(0, document.word_count - total_words) - db.session.add(document) + session.add(document) # Delete database records - db.session.execute(delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids))) - db.session.commit() + session.execute(delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids))) + session.commit() @classmethod def update_segments_status( - cls, segment_ids: list, action: Literal["enable", "disable"], dataset: Dataset, document: Document + cls, + segment_ids: list, + action: Literal["enable", "disable"], + dataset: Dataset, + document: Document, + session: scoped_session | Session, ): assert current_user is not None @@ -3795,7 +3878,7 @@ class SegmentService: return match action: case "enable": - segments = db.session.scalars( + segments = session.scalars( select(DocumentSegment).where( DocumentSegment.id.in_(segment_ids), DocumentSegment.dataset_id == dataset.id, @@ -3814,13 +3897,13 @@ class SegmentService: segment.enabled = True segment.disabled_at = None segment.disabled_by = None - db.session.add(segment) + session.add(segment) real_deal_segment_ids.append(segment.id) - db.session.commit() + session.commit() enable_segments_to_index_task.delay(real_deal_segment_ids, dataset.id, document.id) case "disable": - segments = db.session.scalars( + segments = session.scalars( select(DocumentSegment).where( DocumentSegment.id.in_(segment_ids), DocumentSegment.dataset_id == dataset.id, @@ -3839,15 +3922,20 @@ class SegmentService: segment.enabled = False segment.disabled_at = naive_utc_now() segment.disabled_by = current_user.id - db.session.add(segment) + session.add(segment) real_deal_segment_ids.append(segment.id) - db.session.commit() + session.commit() disable_segments_from_index_task.delay(real_deal_segment_ids, dataset.id, document.id) @classmethod def create_child_chunk( - cls, content: str, segment: DocumentSegment, document: Document, dataset: Dataset + cls, + content: str, + segment: DocumentSegment, + document: Document, + dataset: Dataset, + session: scoped_session | Session, ) -> ChildChunk: assert isinstance(current_user, Account) @@ -3855,7 +3943,7 @@ class SegmentService: with redis_client.lock(lock_name, timeout=20): index_node_id = str(uuid.uuid4()) index_node_hash = helper.generate_text_hash(content) - max_position = db.session.scalar( + max_position = session.scalar( select(func.max(ChildChunk.position)).where( ChildChunk.tenant_id == current_user.current_tenant_id, ChildChunk.dataset_id == dataset.id, @@ -3877,15 +3965,15 @@ class SegmentService: type=SegmentType.CUSTOMIZED, created_by=current_user.id, ) - db.session.add(child_chunk) + session.add(child_chunk) # save vector index try: VectorService.create_child_chunk_vector(child_chunk, dataset) except Exception as e: logger.exception("create child chunk index failed") - db.session.rollback() + session.rollback() raise ChildChunkIndexingError(str(e)) - db.session.commit() + session.commit() return child_chunk @@ -3896,9 +3984,10 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, + session: scoped_session | Session, ) -> list[ChildChunk]: assert isinstance(current_user, Account) - child_chunks = db.session.scalars( + child_chunks = session.scalars( select(ChildChunk).where( ChildChunk.dataset_id == dataset.id, ChildChunk.document_id == document.id, @@ -3926,11 +4015,11 @@ class SegmentService: delete_child_chunks = list(child_chunks_map.values()) try: if update_child_chunks: - db.session.bulk_save_objects(update_child_chunks) + session.bulk_save_objects(update_child_chunks) if delete_child_chunks: for child_chunk in delete_child_chunks: - db.session.delete(child_chunk) + session.delete(child_chunk) if new_child_chunks_args: child_chunk_count = len(child_chunks) for position, args in enumerate(new_child_chunks_args, start=child_chunk_count + 1): @@ -3951,14 +4040,14 @@ class SegmentService: created_by=current_user.id, ) - db.session.add(child_chunk) - db.session.flush() + session.add(child_chunk) + session.flush() new_child_chunks.append(child_chunk) VectorService.update_child_chunk_vector(new_child_chunks, update_child_chunks, delete_child_chunks, dataset) - db.session.commit() + session.commit() except Exception as e: logger.exception("update child chunk index failed") - db.session.rollback() + session.rollback() raise ChildChunkIndexingError(str(e)) return sorted(new_child_chunks + update_child_chunks, key=lambda x: x.position) @@ -3970,6 +4059,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, + session: scoped_session | Session, ) -> ChildChunk: assert current_user is not None @@ -3979,25 +4069,25 @@ class SegmentService: child_chunk.updated_by = current_user.id child_chunk.updated_at = naive_utc_now() child_chunk.type = SegmentType.CUSTOMIZED - db.session.add(child_chunk) + session.add(child_chunk) VectorService.update_child_chunk_vector([], [child_chunk], [], dataset) - db.session.commit() + session.commit() except Exception as e: logger.exception("update child chunk index failed") - db.session.rollback() + session.rollback() raise ChildChunkIndexingError(str(e)) return child_chunk @classmethod - def delete_child_chunk(cls, child_chunk: ChildChunk, dataset: Dataset): - db.session.delete(child_chunk) + def delete_child_chunk(cls, child_chunk: ChildChunk, dataset: Dataset, session: scoped_session | Session): + session.delete(child_chunk) try: VectorService.delete_child_chunk_vector(child_chunk, dataset) except Exception as e: logger.exception("delete child chunk index failed") - db.session.rollback() + session.rollback() raise ChildChunkDeleteIndexError(str(e)) - db.session.commit() + session.commit() @classmethod def get_child_chunks( @@ -4021,9 +4111,11 @@ class SegmentService: return db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False) @classmethod - def get_child_chunk_by_id(cls, child_chunk_id: str, tenant_id: str) -> ChildChunk | None: + def get_child_chunk_by_id( + cls, child_chunk_id: str, tenant_id: str, session: scoped_session | Session + ) -> ChildChunk | None: """Get a child chunk by its ID.""" - result = db.session.scalar( + result = session.scalar( select(ChildChunk).where(ChildChunk.id == child_chunk_id, ChildChunk.tenant_id == tenant_id).limit(1) ) return result if isinstance(result, ChildChunk) else None @@ -4057,9 +4149,11 @@ class SegmentService: return paginated_segments.items, paginated_segments.total @classmethod - def get_segment_by_id(cls, segment_id: str, tenant_id: str) -> DocumentSegment | None: + def get_segment_by_id( + cls, segment_id: str, tenant_id: str, session: scoped_session | Session + ) -> DocumentSegment | None: """Get a segment by its ID.""" - result = db.session.scalar( + result = session.scalar( select(DocumentSegment) .where(DocumentSegment.id == segment_id, DocumentSegment.tenant_id == tenant_id) .limit(1) @@ -4071,6 +4165,7 @@ class SegmentService: cls, document_id: str, dataset_id: str, + session: scoped_session | Session, status: str | None = None, enabled: bool | None = None, ) -> Sequence[DocumentSegment]: @@ -4097,15 +4192,15 @@ class SegmentService: if enabled is not None: query = query.where(DocumentSegment.enabled == enabled) - return db.session.scalars(query).all() + return session.scalars(query).all() class DatasetCollectionBindingService: @classmethod def get_dataset_collection_binding( - cls, provider_name: str, model_name: str, collection_type: str = "dataset" + cls, provider_name: str, model_name: str, session: scoped_session | Session, collection_type: str = "dataset" ) -> DatasetCollectionBinding: - dataset_collection_binding = db.session.scalar( + dataset_collection_binding = session.scalar( select(DatasetCollectionBinding) .where( DatasetCollectionBinding.provider_name == provider_name, @@ -4123,15 +4218,15 @@ class DatasetCollectionBindingService: collection_name=Dataset.gen_collection_name_by_id(str(uuid.uuid4())), type=collection_type, ) - db.session.add(dataset_collection_binding) - db.session.commit() + session.add(dataset_collection_binding) + session.commit() return dataset_collection_binding @classmethod def get_dataset_collection_binding_by_id_and_type( - cls, collection_binding_id: str, collection_type: str = "dataset" + cls, collection_binding_id: str, session: scoped_session | Session, collection_type: str = "dataset" ) -> DatasetCollectionBinding: - dataset_collection_binding = db.session.scalar( + dataset_collection_binding = session.scalar( select(DatasetCollectionBinding) .where( DatasetCollectionBinding.id == collection_binding_id, DatasetCollectionBinding.type == collection_type @@ -4147,8 +4242,8 @@ class DatasetCollectionBindingService: class DatasetPermissionService: @classmethod - def get_dataset_partial_member_list(cls, dataset_id): - user_list_query = db.session.scalars( + def get_dataset_partial_member_list(cls, dataset_id, session: scoped_session | Session): + user_list_query = session.scalars( select( DatasetPermission.account_id, ).where(DatasetPermission.dataset_id == dataset_id) @@ -4157,9 +4252,9 @@ class DatasetPermissionService: return user_list_query @classmethod - def update_partial_member_list(cls, tenant_id, dataset_id, user_list): + def update_partial_member_list(cls, tenant_id, dataset_id, user_list, session: scoped_session | Session): try: - db.session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) + session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) permissions = [] for user in user_list: permission = DatasetPermission( @@ -4169,14 +4264,16 @@ class DatasetPermissionService: ) permissions.append(permission) - db.session.add_all(permissions) - db.session.commit() + session.add_all(permissions) + session.commit() except Exception as e: - db.session.rollback() + session.rollback() raise e @classmethod - def check_permission(cls, user, dataset, requested_permission, requested_partial_member_list): + def check_permission( + cls, user, dataset, requested_permission, requested_partial_member_list, session: scoped_session | Session + ): if not user.is_dataset_editor: raise NoPermissionError("User does not have permission to edit this dataset.") @@ -4187,16 +4284,16 @@ class DatasetPermissionService: if not requested_partial_member_list: raise ValueError("Partial member list is required when setting to partial members.") - local_member_list = cls.get_dataset_partial_member_list(dataset.id) + local_member_list = cls.get_dataset_partial_member_list(dataset.id, session) request_member_list = [user["user_id"] for user in requested_partial_member_list] if set(local_member_list) != set(request_member_list): raise ValueError("Dataset operators cannot change the dataset permissions.") @classmethod - def clear_partial_member_list(cls, dataset_id): + def clear_partial_member_list(cls, dataset_id, session: scoped_session | Session): try: - db.session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) - db.session.commit() + session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) + session.commit() except Exception as e: - db.session.rollback() + session.rollback() raise e diff --git a/api/services/metadata_service.py b/api/services/metadata_service.py index d9cd65b2b39..4e83858ea0e 100644 --- a/api/services/metadata_service.py +++ b/api/services/metadata_service.py @@ -107,7 +107,7 @@ class MetadataService: ).all() if dataset_metadata_bindings: document_ids = [binding.document_id for binding in dataset_metadata_bindings] - documents = DocumentService.get_document_by_ids(document_ids) + documents = DocumentService.get_document_by_ids(document_ids, session) for document in documents: if not document.doc_metadata: doc_metadata = {} @@ -145,7 +145,7 @@ class MetadataService: ).all() if dataset_metadata_bindings: document_ids = [binding.document_id for binding in dataset_metadata_bindings] - documents = DocumentService.get_document_by_ids(document_ids) + documents = DocumentService.get_document_by_ids(document_ids, session) for document in documents: if not document.doc_metadata: doc_metadata = {} @@ -179,7 +179,7 @@ class MetadataService: try: MetadataService.knowledge_base_metadata_lock_check(dataset.id, None) session.add(dataset) - documents = DocumentService.get_working_documents_by_dataset_id(dataset.id) + documents = DocumentService.get_working_documents_by_dataset_id(dataset.id, session) if documents: for document in documents: if not document.doc_metadata: @@ -208,7 +208,7 @@ class MetadataService: try: MetadataService.knowledge_base_metadata_lock_check(dataset.id, None) session.add(dataset) - documents = DocumentService.get_working_documents_by_dataset_id(dataset.id) + documents = DocumentService.get_working_documents_by_dataset_id(dataset.id, session) document_ids = [] if documents: for document in documents: @@ -246,7 +246,7 @@ class MetadataService: lock_key = f"document_metadata_lock_{operation.document_id}" try: MetadataService.knowledge_base_metadata_lock_check(None, operation.document_id) - document = DocumentService.get_document(dataset.id, operation.document_id) + document = DocumentService.get_document(dataset.id, operation.document_id, session=session) if document is None: raise ValueError("Document not found.") if operation.partial_update: diff --git a/api/services/rag_pipeline/rag_pipeline_dsl_service.py b/api/services/rag_pipeline/rag_pipeline_dsl_service.py index bdbf3e080e9..4f9cde37a7b 100644 --- a/api/services/rag_pipeline/rag_pipeline_dsl_service.py +++ b/api/services/rag_pipeline/rag_pipeline_dsl_service.py @@ -15,7 +15,7 @@ from Crypto.Util.Padding import pad, unpad from flask_login import current_user from pydantic import BaseModel from sqlalchemy import select -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, scoped_session from core.file import remote_fetcher from core.helper.name_generator import generate_incremental_name @@ -83,7 +83,7 @@ class RagPipelineDslService: when generated IDs are needed mid-operation; they never commit or rollback. """ - def __init__(self, session: Session): + def __init__(self, session: Session | scoped_session): self._session = session def import_rag_pipeline( diff --git a/api/services/summary_index_service.py b/api/services/summary_index_service.py index 1657cfdd256..3e065653bdf 100644 --- a/api/services/summary_index_service.py +++ b/api/services/summary_index_service.py @@ -7,7 +7,7 @@ from datetime import UTC, datetime from typing import TypedDict, cast from sqlalchemy import select -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, scoped_session from core.db.session_factory import session_factory from core.model_manager import ModelManager @@ -1407,6 +1407,7 @@ class SummaryIndexService: def get_document_summary_status_detail( document_id: str, dataset_id: str, + session: Session | scoped_session, ) -> DocumentSummaryStatusDetailDict: """ Get detailed summary status for a document. @@ -1414,6 +1415,7 @@ class SummaryIndexService: Args: document_id: Document ID dataset_id: Dataset ID + session: SQLAlchemy session used for segment lookup Returns: Dictionary containing: @@ -1431,6 +1433,7 @@ class SummaryIndexService: segments = SegmentService.get_segments_by_document_and_dataset( document_id=document_id, dataset_id=dataset_id, + session=session, status="completed", enabled=True, ) diff --git a/api/tasks/annotation/add_annotation_to_index_task.py b/api/tasks/annotation/add_annotation_to_index_task.py index dafa36cc343..81f57d7d6f3 100644 --- a/api/tasks/annotation/add_annotation_to_index_task.py +++ b/api/tasks/annotation/add_annotation_to_index_task.py @@ -4,6 +4,7 @@ import time import click from celery import shared_task +from core.db.session_factory import session_factory from core.rag.datasource.vdb.vector_factory import Vector from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.models.document import Document @@ -31,9 +32,10 @@ def add_annotation_to_index_task( start_at = time.perf_counter() try: - dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - collection_binding_id, "annotation" - ) + with session_factory.create_session() as session: + dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + collection_binding_id, session, "annotation" + ) dataset = Dataset( id=app_id, tenant_id=tenant_id, diff --git a/api/tasks/annotation/batch_import_annotations_task.py b/api/tasks/annotation/batch_import_annotations_task.py index 89844ef44b4..343b8ee2fa3 100644 --- a/api/tasks/annotation/batch_import_annotations_task.py +++ b/api/tasks/annotation/batch_import_annotations_task.py @@ -63,7 +63,7 @@ def batch_import_annotations_task(job_id: str, content_list: list[dict], app_id: if app_annotation_setting: dataset_collection_binding = ( DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - app_annotation_setting.collection_binding_id, "annotation" + app_annotation_setting.collection_binding_id, session, "annotation" ) ) if not dataset_collection_binding: diff --git a/api/tasks/annotation/delete_annotation_index_task.py b/api/tasks/annotation/delete_annotation_index_task.py index c9aa8fadb78..79a8db2548b 100644 --- a/api/tasks/annotation/delete_annotation_index_task.py +++ b/api/tasks/annotation/delete_annotation_index_task.py @@ -4,6 +4,7 @@ import time import click from celery import shared_task +from core.db.session_factory import session_factory from core.rag.datasource.vdb.vector_factory import Vector from core.rag.index_processor.constant.index_type import IndexTechniqueType from models.dataset import Dataset @@ -20,9 +21,10 @@ def delete_annotation_index_task(annotation_id: str, app_id: str, tenant_id: str logger.info(click.style(f"Start delete app annotation index: {app_id}", fg="green")) start_at = time.perf_counter() try: - dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - collection_binding_id, "annotation" - ) + with session_factory.create_session() as session: + dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + collection_binding_id, session, "annotation" + ) dataset = Dataset( id=app_id, diff --git a/api/tasks/annotation/enable_annotation_reply_task.py b/api/tasks/annotation/enable_annotation_reply_task.py index 4cbca13a92e..32c010eaef1 100644 --- a/api/tasks/annotation/enable_annotation_reply_task.py +++ b/api/tasks/annotation/enable_annotation_reply_task.py @@ -51,7 +51,7 @@ def enable_annotation_reply_task( try: documents = [] dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( - embedding_provider_name, embedding_model_name, CollectionBindingType.ANNOTATION + embedding_provider_name, embedding_model_name, session, CollectionBindingType.ANNOTATION ) annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) @@ -60,7 +60,7 @@ def enable_annotation_reply_task( if dataset_collection_binding.id != annotation_setting.collection_binding_id: old_dataset_collection_binding = ( DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - annotation_setting.collection_binding_id, CollectionBindingType.ANNOTATION + annotation_setting.collection_binding_id, session, CollectionBindingType.ANNOTATION ) ) if old_dataset_collection_binding and annotations: diff --git a/api/tasks/annotation/update_annotation_to_index_task.py b/api/tasks/annotation/update_annotation_to_index_task.py index f41da1d373e..eecc1f6fc7b 100644 --- a/api/tasks/annotation/update_annotation_to_index_task.py +++ b/api/tasks/annotation/update_annotation_to_index_task.py @@ -4,6 +4,7 @@ import time import click from celery import shared_task +from core.db.session_factory import session_factory from core.rag.datasource.vdb.vector_factory import Vector from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.models.document import Document @@ -31,9 +32,10 @@ def update_annotation_to_index_task( start_at = time.perf_counter() try: - dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - collection_binding_id, "annotation" - ) + with session_factory.create_session() as session: + dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + collection_binding_id, session, "annotation" + ) dataset = Dataset( id=app_id, diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py index c644281190c..aab79538f3c 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py @@ -6,9 +6,11 @@ import inspect from collections.abc import Iterator from datetime import UTC, datetime from unittest.mock import MagicMock, PropertyMock, patch +from uuid import uuid4 import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.console.datasets import data_source @@ -22,6 +24,8 @@ from controllers.console.datasets.data_source import ( ) from core.rag.index_processor.constant.index_type import IndexStructureType from models import Account, DataSourceOauthBinding +from models.dataset import Document +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus @pytest.fixture @@ -277,9 +281,13 @@ class TestDataSourceNotionListApi: assert status == 200 - def test_get_success_with_dataset_id(self, app: Flask, current_user: Account, mock_engine: None) -> None: + def test_get_success_with_dataset_id( + self, app: Flask, current_user: Account, mock_engine: None, db_session_with_containers: Session + ) -> None: api = DataSourceNotionListApi() method = inspect.unwrap(api.get) + tenant_id = str(uuid4()) + dataset_id = str(uuid4()) page = MagicMock( page_id="p1", @@ -301,10 +309,24 @@ class TestDataSourceNotionListApi: ) dataset = MagicMock(data_source_type="notion_import") - document = MagicMock(data_source_info='{"notion_page_id": "p1"}') + document = Document( + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.NOTION_IMPORT, + data_source_info='{"notion_page_id": "p1"}', + batch=f"batch-{uuid4()}", + name="Notion Page", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid4()), + indexing_status=IndexingStatus.COMPLETED, + enabled=True, + ) + db_session_with_containers.add(document) + db_session_with_containers.commit() with ( - app.test_request_context("/?credential_id=c1&dataset_id=ds1"), + app.test_request_context(f"/?credential_id=c1&dataset_id={dataset_id}"), patch( "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", return_value={"token": "t"}, @@ -313,7 +335,6 @@ class TestDataSourceNotionListApi: "controllers.console.datasets.data_source.DatasetService.get_dataset", return_value=dataset, ), - patch("controllers.console.datasets.data_source.sessionmaker") as mock_session_class, patch( "core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", return_value=MagicMock( @@ -322,11 +343,7 @@ class TestDataSourceNotionListApi: ), ), ): - mock_session = MagicMock() - mock_session_class.return_value.begin.return_value.__enter__.return_value = mock_session - mock_session.scalars.return_value.all.return_value = [document] - - response, status = method(api, "tenant-1", current_user) + response, status = method(api, tenant_id, current_user) assert status == 200 diff --git a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py index 1717fea789f..6e8af5ca43e 100644 --- a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py +++ b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py @@ -729,13 +729,15 @@ class TestDatasetApiPatch: assert response["name"] == "Updated Dataset" assert response["partial_member_list"] == ["user-1"] mock_dataset_svc.update_dataset.assert_called_once() - _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args + _, update_data, _, session = mock_dataset_svc.update_dataset.call_args.args + assert isinstance(session, (Session, scoped_session)) assert update_data["name"] == "Updated Dataset" assert update_data["permission"] == "partial_members" mock_perm_svc.update_partial_member_list.assert_called_once_with( mock_dataset.tenant_id, mock_dataset.id, [{"user_id": "user-1", "role": "editor"}], + SessionMatcher(), ) diff --git a/api/tests/test_containers_integration_tests/services/dataset_collection_binding.py b/api/tests/test_containers_integration_tests/services/dataset_collection_binding.py index 638a61c8151..d180da9555e 100644 --- a/api/tests/test_containers_integration_tests/services/dataset_collection_binding.py +++ b/api/tests/test_containers_integration_tests/services/dataset_collection_binding.py @@ -88,7 +88,7 @@ class TestDatasetCollectionBindingServiceGetBinding: # Act result = DatasetCollectionBindingService.get_dataset_collection_binding( - provider_name, model_name, collection_type + provider_name, model_name, session=db_session_with_containers, collection_type=collection_type ) # Assert @@ -109,7 +109,7 @@ class TestDatasetCollectionBindingServiceGetBinding: # Act result = DatasetCollectionBindingService.get_dataset_collection_binding( - provider_name, model_name, collection_type + provider_name, model_name, session=db_session_with_containers, collection_type=collection_type ) # Assert @@ -128,7 +128,7 @@ class TestDatasetCollectionBindingServiceGetBinding: # Act result = DatasetCollectionBindingService.get_dataset_collection_binding( - provider_name, model_name, collection_type + provider_name, model_name, session=db_session_with_containers, collection_type=collection_type ) # Assert @@ -143,7 +143,9 @@ class TestDatasetCollectionBindingServiceGetBinding: model_name = "text-embedding-ada-002" # Act - result = DatasetCollectionBindingService.get_dataset_collection_binding(provider_name, model_name) + result = DatasetCollectionBindingService.get_dataset_collection_binding( + provider_name, model_name, session=db_session_with_containers + ) # Assert assert result.type == CollectionBindingType.DATASET @@ -192,7 +194,7 @@ class TestDatasetCollectionBindingServiceGetBindingByIdAndType: # Act result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - binding.id, CollectionBindingType.DATASET + binding.id, session=db_session_with_containers, collection_type=CollectionBindingType.DATASET ) # Assert @@ -210,7 +212,7 @@ class TestDatasetCollectionBindingServiceGetBindingByIdAndType: # Act & Assert with pytest.raises(ValueError, match="Dataset collection binding not found"): DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - non_existent_id, CollectionBindingType.DATASET + non_existent_id, session=db_session_with_containers, collection_type=CollectionBindingType.DATASET ) def test_get_dataset_collection_binding_by_id_and_type_different_collection_type( @@ -228,7 +230,7 @@ class TestDatasetCollectionBindingServiceGetBindingByIdAndType: # Act result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( - binding.id, "custom_type" + binding.id, session=db_session_with_containers, collection_type="custom_type" ) # Assert @@ -249,7 +251,9 @@ class TestDatasetCollectionBindingServiceGetBindingByIdAndType: ) # Act - result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type(binding.id) + result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + binding.id, session=db_session_with_containers + ) # Assert assert result.id == binding.id @@ -268,4 +272,6 @@ class TestDatasetCollectionBindingServiceGetBindingByIdAndType: # Act & Assert with pytest.raises(ValueError, match="Dataset collection binding not found"): - DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type(binding.id, "wrong_type") + DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + binding.id, session=db_session_with_containers, collection_type="wrong_type" + ) diff --git a/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py b/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py index 1cffc43658a..7d4cd63267a 100644 --- a/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py +++ b/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py @@ -141,7 +141,7 @@ class TestDatasetServiceDeleteDataset: # Act with patch("services.dataset_service.dataset_was_deleted") as mock_dataset_was_deleted: - result = DatasetService.delete_dataset(dataset.id, owner) + result = DatasetService.delete_dataset(dataset.id, owner, session=db_session_with_containers) # Assert assert result is True @@ -168,7 +168,7 @@ class TestDatasetServiceDeleteDataset: dataset_id = str(uuid4()) # Act - result = DatasetService.delete_dataset(dataset_id, owner) + result = DatasetService.delete_dataset(dataset_id, owner, session=db_session_with_containers) # Assert assert result is False @@ -198,7 +198,7 @@ class TestDatasetServiceDeleteDataset: # Act & Assert with pytest.raises(NoPermissionError): - DatasetService.delete_dataset(dataset.id, normal_user) + DatasetService.delete_dataset(dataset.id, normal_user, session=db_session_with_containers) # Verify no deletion was attempted assert db_session_with_containers.get(Dataset, dataset.id) is not None @@ -230,7 +230,7 @@ class TestDatasetServiceDatasetUseCheck: DatasetUpdateDeleteTestDataFactory.create_app_dataset_join(db_session_with_containers, app.id, dataset.id) # Act - result = DatasetService.dataset_use_check(dataset.id) + result = DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) # Assert assert result is True @@ -254,7 +254,7 @@ class TestDatasetServiceDatasetUseCheck: dataset = DatasetUpdateDeleteTestDataFactory.create_dataset(db_session_with_containers, tenant.id, owner.id) # Act - result = DatasetService.dataset_use_check(dataset.id) + result = DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) # Assert assert result is False @@ -292,7 +292,7 @@ class TestDatasetServiceUpdateDatasetApiStatus: patch("services.dataset_service.current_user", owner), patch("services.dataset_service.naive_utc_now", return_value=current_time), ): - DatasetService.update_dataset_api_status(dataset.id, True) + DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) @@ -327,7 +327,7 @@ class TestDatasetServiceUpdateDatasetApiStatus: patch("services.dataset_service.current_user", owner), patch("services.dataset_service.naive_utc_now", return_value=current_time), ): - DatasetService.update_dataset_api_status(dataset.id, False) + DatasetService.update_dataset_api_status(dataset.id, False, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) @@ -351,7 +351,7 @@ class TestDatasetServiceUpdateDatasetApiStatus: # Act & Assert with pytest.raises(NotFound, match="Dataset not found"): - DatasetService.update_dataset_api_status(dataset_id, True) + DatasetService.update_dataset_api_status(dataset_id, True, session=db_session_with_containers) def test_update_dataset_api_status_missing_current_user_error(self, db_session_with_containers: Session): """ @@ -378,7 +378,7 @@ class TestDatasetServiceUpdateDatasetApiStatus: patch("services.dataset_service.current_user", None), pytest.raises(ValueError, match="Current user or current user id not found"), ): - DatasetService.update_dataset_api_status(dataset.id, True) + DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) # Verify no commit was attempted db_session_with_containers.rollback() diff --git a/api/tests/test_containers_integration_tests/services/document_service_status.py b/api/tests/test_containers_integration_tests/services/document_service_status.py index 327f14ddfe7..7e78cef1db3 100644 --- a/api/tests/test_containers_integration_tests/services/document_service_status.py +++ b/api/tests/test_containers_integration_tests/services/document_service_status.py @@ -301,7 +301,7 @@ class TestDocumentServicePauseDocument: ) # Act - DocumentService.pause_document(document) + DocumentService.pause_document(document, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -336,7 +336,7 @@ class TestDocumentServicePauseDocument: ) # Act - DocumentService.pause_document(document) + DocumentService.pause_document(document, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -366,7 +366,7 @@ class TestDocumentServicePauseDocument: ) # Act - DocumentService.pause_document(document) + DocumentService.pause_document(document, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -398,7 +398,7 @@ class TestDocumentServicePauseDocument: # Act & Assert with pytest.raises(DocumentIndexingError): - DocumentService.pause_document(document) + DocumentService.pause_document(document, session=db_session_with_containers) db_session_with_containers.refresh(document) assert document.is_paused is False @@ -429,7 +429,7 @@ class TestDocumentServicePauseDocument: # Act & Assert with pytest.raises(DocumentIndexingError): - DocumentService.pause_document(document) + DocumentService.pause_document(document, session=db_session_with_containers) db_session_with_containers.refresh(document) assert document.is_paused is False @@ -507,7 +507,7 @@ class TestDocumentServiceRecoverDocument: ) # Act - DocumentService.recover_document(document) + DocumentService.recover_document(document, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -547,7 +547,7 @@ class TestDocumentServiceRecoverDocument: # Act & Assert with pytest.raises(DocumentIndexingError): - DocumentService.recover_document(document) + DocumentService.recover_document(document, session=db_session_with_containers) db_session_with_containers.refresh(document) assert document.is_paused is False @@ -632,7 +632,7 @@ class TestDocumentServiceRetryDocument: mock_document_service_dependencies["redis_client"].get.return_value = None # Act - DocumentService.retry_document(dataset.id, [document]) + DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -679,7 +679,7 @@ class TestDocumentServiceRetryDocument: mock_document_service_dependencies["redis_client"].get.return_value = None # Act - DocumentService.retry_document(dataset.id, [document1, document2]) + DocumentService.retry_document(dataset.id, [document1, document2], session=db_session_with_containers) # Assert db_session_with_containers.refresh(document1) @@ -719,7 +719,7 @@ class TestDocumentServiceRetryDocument: # Act & Assert with pytest.raises(ValueError, match="Document is being retried, please try again later"): - DocumentService.retry_document(dataset.id, [document]) + DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers) db_session_with_containers.refresh(document) assert document.indexing_status == IndexingStatus.ERROR @@ -753,7 +753,7 @@ class TestDocumentServiceRetryDocument: # Act & Assert with pytest.raises(ValueError, match="Current user or current user id not found"): - DocumentService.retry_document(dataset.id, [document]) + DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers) class TestDocumentServiceBatchUpdateDocumentStatus: @@ -851,7 +851,9 @@ class TestDocumentServiceBatchUpdateDocumentStatus: mock_document_service_dependencies["redis_client"].get.return_value = None # Act - DocumentService.batch_update_document_status(dataset, document_ids, "enable", user) + DocumentService.batch_update_document_status( + dataset, document_ids, "enable", user, session=db_session_with_containers + ) # Assert db_session_with_containers.refresh(document1) @@ -893,7 +895,9 @@ class TestDocumentServiceBatchUpdateDocumentStatus: mock_document_service_dependencies["redis_client"].get.return_value = None # Act - DocumentService.batch_update_document_status(dataset, document_ids, "disable", user) + DocumentService.batch_update_document_status( + dataset, document_ids, "disable", user, session=db_session_with_containers + ) # Assert db_session_with_containers.refresh(document) @@ -935,7 +939,9 @@ class TestDocumentServiceBatchUpdateDocumentStatus: mock_document_service_dependencies["redis_client"].get.return_value = None # Act - DocumentService.batch_update_document_status(dataset, document_ids, "archive", user) + DocumentService.batch_update_document_status( + dataset, document_ids, "archive", user, session=db_session_with_containers + ) # Assert db_session_with_containers.refresh(document) @@ -977,7 +983,9 @@ class TestDocumentServiceBatchUpdateDocumentStatus: mock_document_service_dependencies["redis_client"].get.return_value = None # Act - DocumentService.batch_update_document_status(dataset, document_ids, "un_archive", user) + DocumentService.batch_update_document_status( + dataset, document_ids, "un_archive", user, session=db_session_with_containers + ) # Assert db_session_with_containers.refresh(document) @@ -1006,7 +1014,9 @@ class TestDocumentServiceBatchUpdateDocumentStatus: document_ids = [] # Act - DocumentService.batch_update_document_status(dataset, document_ids, "enable", user) + DocumentService.batch_update_document_status( + dataset, document_ids, "enable", user, session=db_session_with_containers + ) # Assert mock_document_service_dependencies["add_task"].delay.assert_not_called() @@ -1042,7 +1052,9 @@ class TestDocumentServiceBatchUpdateDocumentStatus: # Act & Assert with pytest.raises(DocumentIndexingError, match="is being indexed"): - DocumentService.batch_update_document_status(dataset, document_ids, "enable", user) + DocumentService.batch_update_document_status( + dataset, document_ids, "enable", user, session=db_session_with_containers + ) class TestDocumentServiceRenameDocument: @@ -1121,7 +1133,7 @@ class TestDocumentServiceRenameDocument: ) # Act - result = DocumentService.rename_document(dataset.id, document.id, new_name) + result = DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -1164,7 +1176,7 @@ class TestDocumentServiceRenameDocument: ) # Act - DocumentService.rename_document(dataset.id, document.id, new_name) + DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -1214,7 +1226,7 @@ class TestDocumentServiceRenameDocument: ) # Act - DocumentService.rename_document(dataset.id, document.id, new_name) + DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -1243,7 +1255,7 @@ class TestDocumentServiceRenameDocument: # Act & Assert with pytest.raises(ValueError, match="Dataset not found"): - DocumentService.rename_document(dataset_id, document_id, new_name) + DocumentService.rename_document(dataset_id, document_id, new_name, session=db_session_with_containers) def test_rename_document_not_found_error( self, db_session_with_containers: Session, mock_document_service_dependencies @@ -1272,7 +1284,7 @@ class TestDocumentServiceRenameDocument: # Act & Assert with pytest.raises(ValueError, match="Document not found"): - DocumentService.rename_document(dataset.id, document_id, new_name) + DocumentService.rename_document(dataset.id, document_id, new_name, session=db_session_with_containers) def test_rename_document_permission_error( self, db_session_with_containers: Session, mock_document_service_dependencies @@ -1309,4 +1321,4 @@ class TestDocumentServiceRenameDocument: # Act & Assert with pytest.raises(ValueError, match="No permission"): - DocumentService.rename_document(dataset.id, document.id, new_name) + DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py index b88204b2a6a..0df4ece5d2d 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py @@ -134,7 +134,9 @@ class TestDatasetPermissionServiceGetPartialMemberList: DatasetPermissionTestDataFactory.create_dataset_permission(dataset.id, account_id, tenant.id) # Act - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) # Assert assert set(result) == set(expected_account_ids) @@ -156,7 +158,9 @@ class TestDatasetPermissionServiceGetPartialMemberList: DatasetPermissionTestDataFactory.create_dataset_permission(dataset.id, user.id, tenant.id) # Act - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) # Assert assert set(result) == set(expected_account_ids) @@ -171,7 +175,9 @@ class TestDatasetPermissionServiceGetPartialMemberList: dataset = DatasetPermissionTestDataFactory.create_dataset(tenant.id, owner.id) # Act - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) # Assert assert result == [] @@ -199,10 +205,14 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: user_list = DatasetPermissionTestDataFactory.build_user_list_payload([member_1.id, member_2.id]) # Act - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, user_list) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, user_list, session=db_session_with_containers + ) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert set(result) == {member_1.id, member_2.id} def test_update_partial_member_list_replace_existing(self, db_session_with_containers: Session): @@ -230,15 +240,21 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: dataset = DatasetPermissionTestDataFactory.create_dataset(tenant.id, owner.id) old_users = DatasetPermissionTestDataFactory.build_user_list_payload([old_member_1.id, old_member_2.id]) - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, old_users) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, old_users, session=db_session_with_containers + ) new_users = DatasetPermissionTestDataFactory.build_user_list_payload([new_member_1.id, new_member_2.id]) # Act - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, new_users) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, new_users, session=db_session_with_containers + ) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert set(result) == {new_member_1.id, new_member_2.id} def test_update_partial_member_list_empty_list(self, db_session_with_containers: Session): @@ -257,13 +273,19 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: ) dataset = DatasetPermissionTestDataFactory.create_dataset(tenant.id, owner.id) users = DatasetPermissionTestDataFactory.build_user_list_payload([member_1.id, member_2.id]) - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, users) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, users, session=db_session_with_containers + ) # Act - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, []) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, [], session=db_session_with_containers + ) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert result == [] def test_update_partial_member_list_database_error_rollback(self, db_session_with_containers: Session): @@ -285,10 +307,11 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: tenant.id, dataset.id, DatasetPermissionTestDataFactory.build_user_list_payload([existing_member.id]), + session=db_session_with_containers, ) user_list = DatasetPermissionTestDataFactory.build_user_list_payload([replacement_member.id]) rollback_called = {"count": 0} - original_rollback = db.session.rollback + original_rollback = db_session_with_containers.rollback # Act / Assert with pytest.MonkeyPatch.context() as mp: @@ -300,13 +323,17 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: rollback_called["count"] += 1 original_rollback() - mp.setattr("services.dataset_service.db.session.commit", _raise_commit) - mp.setattr("services.dataset_service.db.session.rollback", _rollback_and_mark) + mp.setattr(db_session_with_containers, "commit", _raise_commit) + mp.setattr(db_session_with_containers, "rollback", _rollback_and_mark) with pytest.raises(Exception, match="Database connection error"): - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, user_list) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, user_list, session=db_session_with_containers + ) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert rollback_called["count"] == 1 assert result == [existing_member.id] assert db_session_with_containers.query(DatasetPermission).filter_by(dataset_id=dataset.id).count() == 1 @@ -331,13 +358,17 @@ class TestDatasetPermissionServiceClearPartialMemberList: ) dataset = DatasetPermissionTestDataFactory.create_dataset(tenant.id, owner.id) users = DatasetPermissionTestDataFactory.build_user_list_payload([member_1.id, member_2.id]) - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, users) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, users, session=db_session_with_containers + ) # Act - DatasetPermissionService.clear_partial_member_list(dataset.id) + DatasetPermissionService.clear_partial_member_list(dataset.id, session=db_session_with_containers) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert result == [] def test_clear_partial_member_list_empty_list(self, db_session_with_containers: Session): @@ -349,10 +380,12 @@ class TestDatasetPermissionServiceClearPartialMemberList: dataset = DatasetPermissionTestDataFactory.create_dataset(tenant.id, owner.id) # Act - DatasetPermissionService.clear_partial_member_list(dataset.id) + DatasetPermissionService.clear_partial_member_list(dataset.id, session=db_session_with_containers) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert result == [] def test_clear_partial_member_list_database_error_rollback(self, db_session_with_containers: Session): @@ -371,9 +404,11 @@ class TestDatasetPermissionServiceClearPartialMemberList: ) dataset = DatasetPermissionTestDataFactory.create_dataset(tenant.id, owner.id) users = DatasetPermissionTestDataFactory.build_user_list_payload([member_1.id, member_2.id]) - DatasetPermissionService.update_partial_member_list(tenant.id, dataset.id, users) + DatasetPermissionService.update_partial_member_list( + tenant.id, dataset.id, users, session=db_session_with_containers + ) rollback_called = {"count": 0} - original_rollback = db.session.rollback + original_rollback = db_session_with_containers.rollback # Act / Assert with pytest.MonkeyPatch.context() as mp: @@ -385,13 +420,15 @@ class TestDatasetPermissionServiceClearPartialMemberList: rollback_called["count"] += 1 original_rollback() - mp.setattr("services.dataset_service.db.session.commit", _raise_commit) - mp.setattr("services.dataset_service.db.session.rollback", _rollback_and_mark) + mp.setattr(db_session_with_containers, "commit", _raise_commit) + mp.setattr(db_session_with_containers, "rollback", _rollback_and_mark) with pytest.raises(Exception, match="Database connection error"): - DatasetPermissionService.clear_partial_member_list(dataset.id) + DatasetPermissionService.clear_partial_member_list(dataset.id, session=db_session_with_containers) # Assert - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert rollback_called["count"] == 1 assert set(result) == {member_1.id, member_2.id} assert db_session_with_containers.query(DatasetPermission).filter_by(dataset_id=dataset.id).count() == 2 @@ -486,7 +523,9 @@ class TestDatasetServiceCheckDatasetPermission: DatasetService.check_dataset_permission(dataset, user, db_session_with_containers) # Assert - permissions = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + permissions = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert user.id in permissions def test_check_dataset_permission_partial_members_without_permission_error( @@ -547,10 +586,12 @@ class TestDatasetServiceCheckDatasetOperatorPermission: DatasetPermissionTestDataFactory.create_dataset_permission(dataset.id, user.id, tenant.id) # Act (should not raise) - DatasetService.check_dataset_operator_permission(user=user, dataset=dataset) + DatasetService.check_dataset_operator_permission(user=user, dataset=dataset, session=db_session_with_containers) # Assert - permissions = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + permissions = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert user.id in permissions def test_check_dataset_operator_permission_partial_members_without_permission_error( @@ -574,4 +615,6 @@ class TestDatasetServiceCheckDatasetOperatorPermission: # Act & Assert with pytest.raises(NoPermissionError, match="You do not have permission to access this dataset"): - DatasetService.check_dataset_operator_permission(user=user, dataset=dataset) + DatasetService.check_dataset_operator_permission( + user=user, dataset=dataset, session=db_session_with_containers + ) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_service.py index 201b65b30d0..912e00b0b7d 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service.py @@ -137,6 +137,7 @@ class TestDatasetServiceCreateDataset: description="Test description", indexing_technique=None, account=account, + session=db_session_with_containers, ) # Assert @@ -159,6 +160,7 @@ class TestDatasetServiceCreateDataset: description=None, indexing_technique=IndexTechniqueType.ECONOMY, account=account, + session=db_session_with_containers, ) # Assert @@ -183,6 +185,7 @@ class TestDatasetServiceCreateDataset: description=None, indexing_technique=IndexTechniqueType.HIGH_QUALITY, account=account, + session=db_session_with_containers, ) # Assert @@ -215,6 +218,7 @@ class TestDatasetServiceCreateDataset: description=None, indexing_technique=None, account=account, + session=db_session_with_containers, ) def test_create_external_dataset_success(self, db_session_with_containers: Session): @@ -236,6 +240,7 @@ class TestDatasetServiceCreateDataset: provider="external", external_knowledge_api_id=external_knowledge_api_id, external_knowledge_id=external_knowledge_id, + session=db_session_with_containers, ) # Assert @@ -276,6 +281,7 @@ class TestDatasetServiceCreateDataset: indexing_technique=IndexTechniqueType.HIGH_QUALITY, account=account, retrieval_model=retrieval_model, + session=db_session_with_containers, ) # Assert @@ -310,6 +316,7 @@ class TestDatasetServiceCreateDataset: account=account, embedding_model_provider=embedding_provider, embedding_model_name=embedding_model_name, + session=db_session_with_containers, ) # Assert @@ -345,6 +352,7 @@ class TestDatasetServiceCreateDataset: indexing_technique=None, account=account, retrieval_model=retrieval_model, + session=db_session_with_containers, ) # Assert @@ -364,6 +372,7 @@ class TestDatasetServiceCreateDataset: indexing_technique=None, account=account, permission=DatasetPermissionEnum.ALL_TEAM, + session=db_session_with_containers, ) # Assert @@ -389,6 +398,7 @@ class TestDatasetServiceCreateDataset: provider="external", external_knowledge_api_id=external_knowledge_api_id, external_knowledge_id="knowledge-123", + session=db_session_with_containers, ) def test_create_external_dataset_missing_knowledge_id_error(self, db_session_with_containers: Session): @@ -410,6 +420,7 @@ class TestDatasetServiceCreateDataset: provider="external", external_knowledge_api_id=external_knowledge_api_id, external_knowledge_id=None, + session=db_session_with_containers, ) @@ -431,7 +442,9 @@ class TestDatasetServiceCreateRagPipelineDataset: # Act with patch("services.dataset_service.current_user", account): result = DatasetService.create_empty_rag_pipeline_dataset( - tenant_id=tenant.id, rag_pipeline_dataset_create_entity=entity + tenant_id=tenant.id, + rag_pipeline_dataset_create_entity=entity, + session=db_session_with_containers, ) # Assert @@ -467,7 +480,9 @@ class TestDatasetServiceCreateRagPipelineDataset: ): mock_generate_name.return_value = generated_name result = DatasetService.create_empty_rag_pipeline_dataset( - tenant_id=tenant.id, rag_pipeline_dataset_create_entity=entity + tenant_id=tenant.id, + rag_pipeline_dataset_create_entity=entity, + session=db_session_with_containers, ) # Assert @@ -505,7 +520,9 @@ class TestDatasetServiceCreateRagPipelineDataset: pytest.raises(DatasetNameDuplicateError, match=f"Dataset with name {duplicate_name} already exists"), ): DatasetService.create_empty_rag_pipeline_dataset( - tenant_id=tenant.id, rag_pipeline_dataset_create_entity=entity + tenant_id=tenant.id, + rag_pipeline_dataset_create_entity=entity, + session=db_session_with_containers, ) def test_create_rag_pipeline_dataset_with_custom_permission(self, db_session_with_containers: Session): @@ -523,7 +540,9 @@ class TestDatasetServiceCreateRagPipelineDataset: # Act with patch("services.dataset_service.current_user", account): result = DatasetService.create_empty_rag_pipeline_dataset( - tenant_id=tenant.id, rag_pipeline_dataset_create_entity=entity + tenant_id=tenant.id, + rag_pipeline_dataset_create_entity=entity, + session=db_session_with_containers, ) # Assert @@ -550,7 +569,9 @@ class TestDatasetServiceCreateRagPipelineDataset: # Act with patch("services.dataset_service.current_user", account): result = DatasetService.create_empty_rag_pipeline_dataset( - tenant_id=tenant.id, rag_pipeline_dataset_create_entity=entity + tenant_id=tenant.id, + rag_pipeline_dataset_create_entity=entity, + session=db_session_with_containers, ) # Assert @@ -580,7 +601,9 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act / Assert with pytest.raises(ValueError, match="Dataset name already exists"): - DatasetService.update_dataset(source_dataset.id, {"name": "Existing Dataset"}, account) + DatasetService.update_dataset( + source_dataset.id, {"name": "Existing Dataset"}, account, session=db_session_with_containers + ) def test_delete_dataset_with_documents_success(self, db_session_with_containers: Session): """Delete a dataset that already has documents.""" @@ -599,7 +622,7 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act with patch("services.dataset_service.dataset_was_deleted") as dataset_deleted_signal: - result = DatasetService.delete_dataset(dataset.id, account) + result = DatasetService.delete_dataset(dataset.id, account, session=db_session_with_containers) # Assert assert result is True @@ -620,7 +643,7 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act with patch("services.dataset_service.dataset_was_deleted") as dataset_deleted_signal: - result = DatasetService.delete_dataset(dataset.id, account) + result = DatasetService.delete_dataset(dataset.id, account, session=db_session_with_containers) # Assert assert result is True @@ -641,7 +664,7 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act with patch("services.dataset_service.dataset_was_deleted") as dataset_deleted_signal: - result = DatasetService.delete_dataset(dataset.id, account) + result = DatasetService.delete_dataset(dataset.id, account, session=db_session_with_containers) # Assert assert result is True @@ -670,7 +693,7 @@ class TestDatasetServiceRetrievalConfiguration: ) # Act - result = DatasetService.get_dataset(dataset.id) + result = DatasetService.get_dataset(dataset.id, session=db_session_with_containers) # Assert assert result is not None @@ -702,7 +725,7 @@ class TestDatasetServiceRetrievalConfiguration: } # Act - result = DatasetService.update_dataset(dataset.id, update_data, account) + result = DatasetService.update_dataset(dataset.id, update_data, account, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) @@ -730,7 +753,7 @@ class TestDocumentServicePauseRecoverRetry: with patch("services.dataset_service.current_user") as mock_user: mock_user.id = account.id - DocumentService.pause_document(doc) + DocumentService.pause_document(doc, session=db_session_with_containers) db_session_with_containers.refresh(doc) assert doc.is_paused is True @@ -750,7 +773,7 @@ class TestDocumentServicePauseRecoverRetry: with patch("services.dataset_service.current_user") as mock_user: mock_user.id = account.id with pytest.raises(DocumentIndexingError): - DocumentService.pause_document(doc) + DocumentService.pause_document(doc, session=db_session_with_containers) def test_recover_document_success(self, db_session_with_containers: Session): from extensions.ext_redis import redis_client @@ -761,11 +784,11 @@ class TestDocumentServicePauseRecoverRetry: # Pause first with patch("services.dataset_service.current_user") as mock_user: mock_user.id = account.id - DocumentService.pause_document(doc) + DocumentService.pause_document(doc, session=db_session_with_containers) # Recover with patch("services.dataset_service.recover_document_indexing_task") as recover_task: - DocumentService.recover_document(doc) + DocumentService.recover_document(doc, session=db_session_with_containers) db_session_with_containers.refresh(doc) assert doc.is_paused is False @@ -795,7 +818,7 @@ class TestDocumentServicePauseRecoverRetry: patch("services.dataset_service.retry_document_indexing_task") as retry_task, ): mock_user.id = account.id - DocumentService.retry_document(dataset.id, [doc1, doc2]) + DocumentService.retry_document(dataset.id, [doc1, doc2], session=db_session_with_containers) db_session_with_containers.refresh(doc1) db_session_with_containers.refresh(doc2) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_batch_update_document_status.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_batch_update_document_status.py index c1d088755c1..2f1022a47ca 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_batch_update_document_status.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_batch_update_document_status.py @@ -196,7 +196,11 @@ class TestDatasetServiceBatchUpdateDocumentStatus: # Act DocumentService.batch_update_document_status( - dataset=dataset, document_ids=document_ids, action="enable", user=user + dataset=dataset, + document_ids=document_ids, + action="enable", + user=user, + session=db_session_with_containers, ) # Assert @@ -228,6 +232,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="enable", user=user, + session=db_session_with_containers, ) # Assert @@ -256,6 +261,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=document_ids, action="disable", user=user, + session=db_session_with_containers, ) # Assert @@ -291,6 +297,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[disabled_doc.id], action="disable", user=user, + session=db_session_with_containers, ) # Assert @@ -321,6 +328,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[non_completed_doc.id], action="disable", user=user, + session=db_session_with_containers, ) def test_batch_update_archive_documents_success(self, db_session_with_containers: Session, patched_dependencies): @@ -338,6 +346,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="archive", user=user, + session=db_session_with_containers, ) # Assert @@ -364,6 +373,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="archive", user=user, + session=db_session_with_containers, ) # Assert @@ -389,6 +399,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="archive", user=user, + session=db_session_with_containers, ) # Assert @@ -412,6 +423,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="un_archive", user=user, + session=db_session_with_containers, ) # Assert @@ -439,6 +451,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="un_archive", user=user, + session=db_session_with_containers, ) # Assert @@ -464,6 +477,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="un_archive", user=user, + session=db_session_with_containers, ) # Assert @@ -495,6 +509,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="enable", user=user, + session=db_session_with_containers, ) assert "test_document.pdf" in str(exc_info.value) @@ -517,6 +532,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[document.id], action="enable", user=user, + session=db_session_with_containers, ) db_session_with_containers.refresh(document) @@ -531,7 +547,11 @@ class TestDatasetServiceBatchUpdateDocumentStatus: # Act result = DocumentService.batch_update_document_status( - dataset=dataset, document_ids=[], action="enable", user=user + dataset=dataset, + document_ids=[], + action="enable", + user=user, + session=db_session_with_containers, ) # Assert @@ -552,6 +572,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=[missing_document_id], action="enable", user=user, + session=db_session_with_containers, ) # Assert @@ -590,6 +611,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=document_ids, action="enable", user=user, + session=db_session_with_containers, ) # Assert @@ -628,6 +650,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=document_ids, action="enable", user=user, + session=db_session_with_containers, ) # Assert @@ -679,6 +702,7 @@ class TestDatasetServiceBatchUpdateDocumentStatus: document_ids=document_ids, action="enable", user=user, + session=db_session_with_containers, ) # Assert @@ -709,5 +733,9 @@ class TestDatasetServiceBatchUpdateDocumentStatus: with pytest.raises(ValueError, match="Invalid action"): DocumentService.batch_update_document_status( - dataset=dataset, document_ids=[doc.id], action="invalid_action", user=user + dataset=dataset, + document_ids=[doc.id], + action="invalid_action", + user=user, + session=db_session_with_containers, ) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_create_dataset.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_create_dataset.py index 08de79f4b7e..292ae46190a 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_create_dataset.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_create_dataset.py @@ -58,4 +58,5 @@ class TestDatasetServiceCreateRagPipelineDataset: DatasetService.create_empty_rag_pipeline_dataset( tenant_id=tenant.id, rag_pipeline_dataset_create_entity=self._build_entity(), + session=db_session_with_containers, ) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_delete_dataset.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_delete_dataset.py index c43a5d59789..1c5ec6835bd 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_delete_dataset.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_delete_dataset.py @@ -130,7 +130,7 @@ class TestDatasetServiceDeleteDataset: "events.event_handlers.clean_when_dataset_deleted.clean_dataset_task.delay", autospec=True, ) as clean_dataset_delay: - result = DatasetService.delete_dataset(dataset.id, owner) + result = DatasetService.delete_dataset(dataset.id, owner, session=db_session_with_containers) # Assert db_session_with_containers.expire_all() @@ -166,7 +166,7 @@ class TestDatasetServiceDeleteDataset: "events.event_handlers.clean_when_dataset_deleted.clean_dataset_task.delay", autospec=True, ) as clean_dataset_delay: - result = DatasetService.delete_dataset(dataset.id, owner) + result = DatasetService.delete_dataset(dataset.id, owner, session=db_session_with_containers) # Assert db_session_with_containers.expire_all() @@ -194,7 +194,7 @@ class TestDatasetServiceDeleteDataset: "events.event_handlers.clean_when_dataset_deleted.clean_dataset_task.delay", autospec=True, ) as clean_dataset_delay: - result = DatasetService.delete_dataset(dataset.id, owner) + result = DatasetService.delete_dataset(dataset.id, owner, session=db_session_with_containers) # Assert db_session_with_containers.expire_all() @@ -222,7 +222,7 @@ class TestDatasetServiceDeleteDataset: "events.event_handlers.clean_when_dataset_deleted.clean_dataset_task.delay", autospec=True, ) as clean_dataset_delay: - result = DatasetService.delete_dataset(dataset.id, owner) + result = DatasetService.delete_dataset(dataset.id, owner, session=db_session_with_containers) # Assert db_session_with_containers.expire_all() @@ -241,7 +241,7 @@ class TestDatasetServiceDeleteDataset: "events.event_handlers.clean_when_dataset_deleted.clean_dataset_task.delay", autospec=True, ) as clean_dataset_delay: - result = DatasetService.delete_dataset(missing_dataset_id, owner) + result = DatasetService.delete_dataset(missing_dataset_id, owner, session=db_session_with_containers) # Assert assert result is False diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py index 946ac661940..ab5ba173792 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py @@ -125,14 +125,14 @@ def current_user_mock(): def test_get_document_returns_none_when_document_id_is_missing(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) - assert DocumentService.get_document(dataset.id, None) is None + assert DocumentService.get_document(dataset.id, None, session=db_session_with_containers) is None def test_get_document_queries_by_dataset_and_document_id(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) document = DocumentServiceIntegrationFactory.create_document(db_session_with_containers, dataset=dataset) - result = DocumentService.get_document(dataset.id, document.id) + result = DocumentService.get_document(dataset.id, document.id, session=db_session_with_containers) assert result is not None assert result.id == document.id @@ -141,7 +141,7 @@ def test_get_document_queries_by_dataset_and_document_id(db_session_with_contain def test_get_documents_by_ids_returns_empty_for_empty_input(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) - result = DocumentService.get_documents_by_ids(dataset.id, []) + result = DocumentService.get_documents_by_ids(dataset.id, [], session=db_session_with_containers) assert result == [] @@ -156,7 +156,7 @@ def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers position=2, ) - result = DocumentService.get_documents_by_ids(dataset.id, [doc_a.id, doc_b.id]) + result = DocumentService.get_documents_by_ids(dataset.id, [doc_a.id, doc_b.id], db_session_with_containers) assert {document.id for document in result} == {doc_a.id, doc_b.id} @@ -164,7 +164,7 @@ def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers def test_update_documents_need_summary_returns_zero_for_empty_input(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) - assert DocumentService.update_documents_need_summary(dataset.id, []) == 0 + assert DocumentService.update_documents_need_summary(dataset.id, [], db_session_with_containers) == 0 def test_update_documents_need_summary_updates_matching_non_qa_documents(db_session_with_containers: Session): @@ -185,6 +185,7 @@ def test_update_documents_need_summary_updates_matching_non_qa_documents(db_sess updated_count = DocumentService.update_documents_need_summary( dataset.id, [paragraph_doc.id, qa_doc.id], + db_session_with_containers, need_summary=False, ) @@ -212,7 +213,7 @@ def test_get_document_download_url_uses_signed_url_helper(db_session_with_contai ) with patch("services.dataset_service.file_helpers.get_signed_file_url", return_value="signed-url") as get_url: - result = DocumentService.get_document_download_url(document) + result = DocumentService.get_document_download_url(document, session=db_session_with_containers) assert result == "signed-url" get_url.assert_called_once_with(upload_file_id=upload_file.id, as_attachment=True) @@ -282,7 +283,7 @@ def test_get_upload_file_for_upload_file_document_raises_when_file_service_retur with patch("services.dataset_service.FileService.get_upload_files_by_ids", return_value={}): with pytest.raises(NotFound, match="Uploaded file not found"): - DocumentService._get_upload_file_for_upload_file_document(document) + DocumentService._get_upload_file_for_upload_file_document(document, session=db_session_with_containers) def test_get_upload_file_for_upload_file_document_returns_upload_file(db_session_with_containers: Session): @@ -298,7 +299,7 @@ def test_get_upload_file_for_upload_file_document_returns_upload_file(db_session data_source_info={"upload_file_id": upload_file.id}, ) - result = DocumentService._get_upload_file_for_upload_file_document(document) + result = DocumentService._get_upload_file_for_upload_file_document(document, session=db_session_with_containers) assert result.id == upload_file.id @@ -313,6 +314,7 @@ def test_get_upload_files_by_document_id_for_zip_download_raises_for_missing_doc dataset_id=dataset.id, document_ids=[str(uuid4())], tenant_id=dataset.tenant_id, + session=db_session_with_containers, ) @@ -337,6 +339,7 @@ def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_a dataset_id=dataset.id, document_ids=[document.id], tenant_id=dataset.tenant_id, + session=db_session_with_containers, ) @@ -355,6 +358,7 @@ def test_get_upload_files_by_document_id_for_zip_download_rejects_missing_upload dataset_id=dataset.id, document_ids=[document.id], tenant_id=dataset.tenant_id, + session=db_session_with_containers, ) @@ -390,6 +394,7 @@ def test_get_upload_files_by_document_id_for_zip_download_returns_document_keyed dataset_id=dataset.id, document_ids=[document_a.id, document_b.id], tenant_id=dataset.tenant_id, + session=db_session_with_containers, ) assert mapping[document_a.id].id == upload_file_a.id @@ -397,7 +402,7 @@ def test_get_upload_files_by_document_id_for_zip_download_returns_document_keyed def test_prepare_document_batch_download_zip_raises_not_found_for_missing_dataset( - current_user_mock, flask_app_with_containers + current_user_mock, flask_app_with_containers, db_session_with_containers: Session ): with flask_app_with_containers.app_context(): with pytest.raises(NotFound, match="Dataset not found"): @@ -406,6 +411,7 @@ def test_prepare_document_batch_download_zip_raises_not_found_for_missing_datase document_ids=[str(uuid4())], tenant_id=current_user_mock.current_tenant_id, current_user=current_user_mock, + session=db_session_with_containers, ) @@ -429,6 +435,7 @@ def test_prepare_document_batch_download_zip_translates_permission_error_to_forb document_ids=[], tenant_id=current_user_mock.current_tenant_id, current_user=current_user_mock, + session=db_session_with_containers, ) @@ -470,6 +477,7 @@ def test_prepare_document_batch_download_zip_returns_upload_files_in_requested_o document_ids=[document_b.id, document_a.id], tenant_id=current_user_mock.current_tenant_id, current_user=current_user_mock, + session=db_session_with_containers, ) assert [upload_file.id for upload_file in upload_files] == [upload_file_b.id, upload_file_a.id] @@ -490,7 +498,7 @@ def test_get_document_by_dataset_id_returns_enabled_documents(db_session_with_co enabled=False, ) - result = DocumentService.get_document_by_dataset_id(dataset.id) + result = DocumentService.get_document_by_dataset_id(dataset.id, session=db_session_with_containers) assert [document.id for document in result] == [enabled_document.id] @@ -513,7 +521,7 @@ def test_get_working_documents_by_dataset_id_returns_completed_enabled_unarchive indexing_status=IndexingStatus.ERROR, ) - result = DocumentService.get_working_documents_by_dataset_id(dataset.id) + result = DocumentService.get_working_documents_by_dataset_id(dataset.id, session=db_session_with_containers) assert [document.id for document in result] == [available_document.id] @@ -538,7 +546,7 @@ def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db indexing_status=IndexingStatus.COMPLETED, ) - result = DocumentService.get_error_documents_by_dataset_id(dataset.id) + result = DocumentService.get_error_documents_by_dataset_id(dataset.id, session=db_session_with_containers) assert {document.id for document in result} == {error_document.id, paused_document.id} @@ -561,7 +569,7 @@ def test_get_batch_documents_filters_by_current_user_tenant(db_session_with_cont with patch("services.dataset_service.current_user", create_autospec(Account, instance=True)) as current_user: current_user.current_tenant_id = dataset.tenant_id - result = DocumentService.get_batch_documents(dataset.id, batch) + result = DocumentService.get_batch_documents(dataset.id, batch, session=db_session_with_containers) assert [document.id for document in result] == [matching_document.id] @@ -574,7 +582,7 @@ def test_get_document_file_detail_returns_upload_file(db_session_with_containers created_by=dataset.created_by, ) - result = DocumentService.get_document_file_detail(upload_file.id) + result = DocumentService.get_document_file_detail(upload_file.id, session=db_session_with_containers) assert result is not None assert result.id == upload_file.id @@ -594,7 +602,7 @@ def test_delete_document_emits_signal_and_commits(db_session_with_containers: Se ) with patch("services.dataset_service.document_was_deleted.send") as signal_send: - DocumentService.delete_document(document) + DocumentService.delete_document(document, session=db_session_with_containers) assert db_session_with_containers.get(Document, document.id) is None signal_send.assert_called_once_with( @@ -609,7 +617,7 @@ def test_delete_documents_ignores_empty_input(db_session_with_containers: Sessio dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) with patch("services.dataset_service.batch_clean_document_task.delay") as delay: - DocumentService.delete_documents(dataset, []) + DocumentService.delete_documents(dataset, [], session=db_session_with_containers) delay.assert_not_called() @@ -643,7 +651,7 @@ def test_delete_documents_deletes_rows_and_dispatches_cleanup_task(db_session_wi ) with patch("services.dataset_service.batch_clean_document_task.delay") as delay: - DocumentService.delete_documents(dataset, [document_a.id, document_b.id]) + DocumentService.delete_documents(dataset, [document_a.id, document_b.id], session=db_session_with_containers) assert db_session_with_containers.get(Document, document_a.id) is None assert db_session_with_containers.get(Document, document_b.id) is None @@ -658,10 +666,10 @@ def test_get_documents_position_returns_next_position_when_documents_exist(db_se dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) DocumentServiceIntegrationFactory.create_document(db_session_with_containers, dataset=dataset, position=3) - assert DocumentService.get_documents_position(dataset.id) == 4 + assert DocumentService.get_documents_position(dataset.id, session=db_session_with_containers) == 4 def test_get_documents_position_defaults_to_one_when_dataset_is_empty(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) - assert DocumentService.get_documents_position(dataset.id) == 1 + assert DocumentService.get_documents_position(dataset.id, session=db_session_with_containers) == 1 diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py index ba5883f408d..6b32273624b 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py @@ -182,7 +182,7 @@ class TestDatasetServicePermissionsAndLifecycle: def test_delete_dataset_returns_false_when_dataset_is_missing(self, db_session_with_containers: Session): owner, _tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) - result = DatasetService.delete_dataset(str(uuid4()), user=owner) + result = DatasetService.delete_dataset(str(uuid4()), user=owner, session=db_session_with_containers) assert result is False @@ -195,7 +195,7 @@ class TestDatasetServicePermissionsAndLifecycle: ) with patch("services.dataset_service.dataset_was_deleted.send") as send_deleted_signal: - result = DatasetService.delete_dataset(dataset.id, user=owner) + result = DatasetService.delete_dataset(dataset.id, user=owner, session=db_session_with_containers) assert result is True assert db_session_with_containers.get(Dataset, dataset.id) is None @@ -213,7 +213,7 @@ class TestDatasetServicePermissionsAndLifecycle: dataset_id=dataset.id, ) - assert DatasetService.dataset_use_check(dataset.id) is True + assert DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) is True def test_dataset_use_check_returns_false_when_join_missing(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -223,7 +223,7 @@ class TestDatasetServicePermissionsAndLifecycle: created_by=owner.id, ) - assert DatasetService.dataset_use_check(dataset.id) is False + assert DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) is False def test_check_dataset_permission_rejects_cross_tenant_access(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -320,7 +320,9 @@ class TestDatasetServicePermissionsAndLifecycle: ) with pytest.raises(NoPermissionError, match="do not have permission"): - DatasetService.check_dataset_operator_permission(user=operator, dataset=dataset) + DatasetService.check_dataset_operator_permission( + user=operator, dataset=dataset, session=db_session_with_containers + ) def test_check_dataset_operator_permission_rejects_partial_team_without_binding( self, db_session_with_containers: Session @@ -339,7 +341,9 @@ class TestDatasetServicePermissionsAndLifecycle: ) with pytest.raises(NoPermissionError, match="do not have permission"): - DatasetService.check_dataset_operator_permission(user=operator, dataset=dataset) + DatasetService.check_dataset_operator_permission( + user=operator, dataset=dataset, session=db_session_with_containers + ) def test_check_dataset_operator_permission_allows_partial_team_with_binding( self, db_session_with_containers: Session @@ -363,12 +367,16 @@ class TestDatasetServicePermissionsAndLifecycle: account_id=operator.id, ) - DatasetService.check_dataset_operator_permission(user=operator, dataset=dataset) + DatasetService.check_dataset_operator_permission( + user=operator, dataset=dataset, session=db_session_with_containers + ) - def test_update_dataset_api_status_raises_not_found_for_missing_dataset(self, flask_app_with_containers: Flask): + def test_update_dataset_api_status_raises_not_found_for_missing_dataset( + self, flask_app_with_containers: Flask, db_session_with_containers: Session + ): with flask_app_with_containers.app_context(): with pytest.raises(NotFound, match="Dataset not found"): - DatasetService.update_dataset_api_status(str(uuid4()), True) + DatasetService.update_dataset_api_status(str(uuid4()), True, session=db_session_with_containers) def test_update_dataset_api_status_requires_current_user_id(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -381,7 +389,7 @@ class TestDatasetServicePermissionsAndLifecycle: with patch("services.dataset_service.current_user", SimpleNamespace(id=None)): with pytest.raises(ValueError, match="Current user or current user id not found"): - DatasetService.update_dataset_api_status(dataset.id, True) + DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) def test_update_dataset_api_status_updates_fields_and_commits(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -397,7 +405,7 @@ class TestDatasetServicePermissionsAndLifecycle: patch("services.dataset_service.current_user", owner), patch("services.dataset_service.naive_utc_now", return_value=now), ): - DatasetService.update_dataset_api_status(dataset.id, True) + DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.enable_api is True @@ -416,7 +424,7 @@ class TestDatasetServicePermissionsAndLifecycle: patch("services.dataset_service.current_user", owner), patch("services.dataset_service.FeatureService.get_features", return_value=features), ): - result = DatasetService.get_dataset_auto_disable_logs(str(uuid4())) + result = DatasetService.get_dataset_auto_disable_logs(str(uuid4()), session=db_session_with_containers) assert result == {"document_ids": [], "count": 0} @@ -447,7 +455,7 @@ class TestDatasetServicePermissionsAndLifecycle: patch("services.dataset_service.current_user", owner), patch("services.dataset_service.FeatureService.get_features", return_value=features), ): - result = DatasetService.get_dataset_auto_disable_logs(dataset.id) + result = DatasetService.get_dataset_auto_disable_logs(dataset.id, session=db_session_with_containers) assert result["count"] == 2 assert len(result["document_ids"]) == 2 @@ -461,12 +469,16 @@ class TestDatasetCollectionBindingServiceIntegration: model_name="model", ) - result = DatasetCollectionBindingService.get_dataset_collection_binding("provider", "model") + result = DatasetCollectionBindingService.get_dataset_collection_binding( + "provider", "model", session=db_session_with_containers + ) assert result.id == binding.id def test_get_dataset_collection_binding_creates_binding_when_missing(self, db_session_with_containers: Session): - result = DatasetCollectionBindingService.get_dataset_collection_binding("provider", "missing-model") + result = DatasetCollectionBindingService.get_dataset_collection_binding( + "provider", "missing-model", session=db_session_with_containers + ) persisted = db_session_with_containers.get(DatasetCollectionBinding, result.id) assert persisted is not None @@ -475,10 +487,14 @@ class TestDatasetCollectionBindingServiceIntegration: assert persisted.type == "dataset" assert persisted.collection_name - def test_get_dataset_collection_binding_by_id_and_type_raises_when_missing(self, flask_app_with_containers: Flask): + def test_get_dataset_collection_binding_by_id_and_type_raises_when_missing( + self, flask_app_with_containers: Flask, db_session_with_containers: Session + ): with flask_app_with_containers.app_context(): with pytest.raises(ValueError, match="Dataset collection binding not found"): - DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type(str(uuid4())) + DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + str(uuid4()), session=db_session_with_containers + ) def test_get_dataset_collection_binding_by_id_and_type_returns_binding(self, db_session_with_containers: Session): binding = DatasetPermissionIntegrationFactory.create_collection_binding( @@ -487,7 +503,9 @@ class TestDatasetCollectionBindingServiceIntegration: model_name="model", ) - result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type(binding.id) + result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type( + binding.id, session=db_session_with_containers + ) assert result.id == binding.id @@ -516,7 +534,9 @@ class TestDatasetPermissionServiceIntegration: account_id=member_b.id, ) - result = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) + result = DatasetPermissionService.get_dataset_partial_member_list( + dataset.id, session=db_session_with_containers + ) assert set(result) == {member_a.id, member_b.id} @@ -542,33 +562,44 @@ class TestDatasetPermissionServiceIntegration: tenant.id, dataset.id, [{"user_id": member_a.id}, {"user_id": member_b.id}], + session=db_session_with_containers, ) permissions = db_session_with_containers.query(DatasetPermission).filter_by(dataset_id=dataset.id).all() assert {permission.account_id for permission in permissions} == {member_a.id, member_b.id} - def test_check_permission_requires_dataset_editor(self): + def test_check_permission_requires_dataset_editor(self, db_session_with_containers: Session): user = SimpleNamespace(is_dataset_editor=False, is_dataset_operator=False) dataset = SimpleNamespace(id="dataset-1", permission=DatasetPermissionEnum.ALL_TEAM) with pytest.raises(NoPermissionError, match="does not have permission"): - DatasetPermissionService.check_permission(user, dataset, DatasetPermissionEnum.ALL_TEAM, []) + DatasetPermissionService.check_permission( + user, dataset, DatasetPermissionEnum.ALL_TEAM, [], session=db_session_with_containers + ) - def test_check_permission_prevents_dataset_operator_from_changing_permission_mode(self): + def test_check_permission_prevents_dataset_operator_from_changing_permission_mode( + self, db_session_with_containers: Session + ): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = SimpleNamespace(id="dataset-1", permission=DatasetPermissionEnum.ALL_TEAM) with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): - DatasetPermissionService.check_permission(user, dataset, DatasetPermissionEnum.ONLY_ME, []) + DatasetPermissionService.check_permission( + user, dataset, DatasetPermissionEnum.ONLY_ME, [], session=db_session_with_containers + ) - def test_check_permission_requires_partial_member_list_for_partial_members_mode(self): + def test_check_permission_requires_partial_member_list_for_partial_members_mode( + self, db_session_with_containers: Session + ): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = SimpleNamespace(id="dataset-1", permission=DatasetPermissionEnum.PARTIAL_TEAM) with pytest.raises(ValueError, match="Partial member list is required"): - DatasetPermissionService.check_permission(user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, []) + DatasetPermissionService.check_permission( + user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [], session=db_session_with_containers + ) - def test_check_permission_rejects_dataset_operator_member_list_changes(self): + def test_check_permission_rejects_dataset_operator_member_list_changes(self, db_session_with_containers: Session): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = SimpleNamespace(id="dataset-1", permission=DatasetPermissionEnum.PARTIAL_TEAM) @@ -579,9 +610,12 @@ class TestDatasetPermissionServiceIntegration: dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-2"}], + session=db_session_with_containers, ) - def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged(self): + def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged( + self, db_session_with_containers: Session + ): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = SimpleNamespace(id="dataset-1", permission=DatasetPermissionEnum.PARTIAL_TEAM) @@ -591,6 +625,7 @@ class TestDatasetPermissionServiceIntegration: dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-1"}], + session=db_session_with_containers, ) def test_clear_partial_member_list_deletes_permissions_and_commits(self, db_session_with_containers: Session): @@ -609,7 +644,7 @@ class TestDatasetPermissionServiceIntegration: account_id=member.id, ) - DatasetPermissionService.clear_partial_member_list(dataset.id) + DatasetPermissionService.clear_partial_member_list(dataset.id, session=db_session_with_containers) remaining = db_session_with_containers.query(DatasetPermission).filter_by(dataset_id=dataset.id).all() assert remaining == [] diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py index 05632b1ec2a..b6768d2ca26 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py @@ -548,7 +548,7 @@ class TestDatasetServiceGetDataset: ) # Act - result = DatasetService.get_dataset(dataset.id) + result = DatasetService.get_dataset(dataset.id, session=db_session_with_containers) # Assert assert result is not None @@ -560,7 +560,7 @@ class TestDatasetServiceGetDataset: dataset_id = str(uuid4()) # Act - result = DatasetService.get_dataset(dataset_id) + result = DatasetService.get_dataset(dataset_id, session=db_session_with_containers) # Assert assert result is None @@ -639,7 +639,7 @@ class TestDatasetServiceGetProcessRules: ) # Act - result = DatasetService.get_process_rules(dataset.id) + result = DatasetService.get_process_rules(dataset.id, session=db_session_with_containers) # Assert assert result["mode"] == "custom" @@ -654,7 +654,7 @@ class TestDatasetServiceGetProcessRules: ) # Act - result = DatasetService.get_process_rules(dataset.id) + result = DatasetService.get_process_rules(dataset.id, session=db_session_with_containers) # Assert assert result["mode"] == DocumentService.DEFAULT_RULES["mode"] @@ -724,7 +724,7 @@ class TestDatasetServiceGetRelatedApps: DatasetRetrievalTestDataFactory.create_app_dataset_join(db_session_with_containers, dataset.id) # Act - result = DatasetService.get_related_apps(dataset.id) + result = DatasetService.get_related_apps(dataset.id, session=db_session_with_containers) # Assert assert len(result) == 2 @@ -739,7 +739,7 @@ class TestDatasetServiceGetRelatedApps: ) # Act - result = DatasetService.get_related_apps(dataset.id) + result = DatasetService.get_related_apps(dataset.id, session=db_session_with_containers) # Assert assert result == [] diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py index ac0483a45d7..d9fb23e8e33 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py @@ -189,7 +189,7 @@ class TestDatasetServiceUpdateDataset: "external_knowledge_api_id": external_api.id, } - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) updated_binding = db_session_with_containers.query(ExternalKnowledgeBindings).filter_by(id=binding_id).first() @@ -221,7 +221,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_api_id": str(uuid4())} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge id is required" in str(context.value) db_session_with_containers.rollback() @@ -245,7 +245,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_id": "knowledge_id"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge api id is required" in str(context.value) db_session_with_containers.rollback() @@ -272,7 +272,7 @@ class TestDatasetServiceUpdateDataset: } with pytest.raises(ValueError) as context: - DatasetService.update_dataset(dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge binding not found" in str(context.value) db_session_with_containers.rollback() @@ -303,7 +303,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": "text-embedding-ada-002", } - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -338,7 +338,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": None, } - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -371,7 +371,7 @@ class TestDatasetServiceUpdateDataset: } with patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task: - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_task.delay.assert_called_once_with(dataset.id, "remove") db_session_with_containers.refresh(dataset) @@ -418,7 +418,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -426,7 +426,7 @@ class TestDatasetServiceUpdateDataset: model_type=ModelType.TEXT_EMBEDDING, model="text-embedding-ada-002", ) - mock_get_binding.assert_called_once_with("openai", "text-embedding-ada-002") + mock_get_binding.assert_called_once_with("openai", "text-embedding-ada-002", db_session_with_containers) mock_task.delay.assert_called_once_with(dataset.id, "add") db_session_with_containers.refresh(dataset) @@ -462,7 +462,7 @@ class TestDatasetServiceUpdateDataset: "retrieval_model": "new_model", } - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -514,7 +514,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -522,7 +522,7 @@ class TestDatasetServiceUpdateDataset: model_type=ModelType.TEXT_EMBEDDING, model="text-embedding-3-small", ) - mock_get_binding.assert_called_once_with("openai", "text-embedding-3-small") + mock_get_binding.assert_called_once_with("openai", "text-embedding-3-small", db_session_with_containers) mock_task.delay.assert_called_once_with(dataset.id, "update") mock_regenerate_task.delay.assert_called_once_with( dataset.id, @@ -545,7 +545,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(str(uuid4()), update_data, user) + DatasetService.update_dataset(str(uuid4()), update_data, user, session=db_session_with_containers) assert "Dataset not found" in str(context.value) @@ -568,7 +568,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(NoPermissionError): - DatasetService.update_dataset(dataset.id, update_data, outsider) + DatasetService.update_dataset(dataset.id, update_data, outsider, session=db_session_with_containers) def test_update_internal_dataset_embedding_model_error(self, db_session_with_containers: Session): """Test error when embedding model is not available.""" @@ -595,6 +595,6 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.side_effect = Exception("No Embedding Model available") with pytest.raises(Exception) as context: - DatasetService.update_dataset(dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "No Embedding Model available".lower() in str(context.value).lower() diff --git a/api/tests/test_containers_integration_tests/services/test_document_service_rename_document.py b/api/tests/test_containers_integration_tests/services/test_document_service_rename_document.py index 34532ed7f81..6a0ea17560a 100644 --- a/api/tests/test_containers_integration_tests/services/test_document_service_rename_document.py +++ b/api/tests/test_containers_integration_tests/services/test_document_service_rename_document.py @@ -118,7 +118,7 @@ def test_rename_document_success(db_session_with_containers, mock_env): ) # Act - result = DocumentService.rename_document(dataset.id, document_id, new_name) + result = DocumentService.rename_document(dataset.id, document_id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -147,7 +147,7 @@ def test_rename_document_with_built_in_fields(db_session_with_containers, mock_e ) # Act - DocumentService.rename_document(dataset.id, document.id, new_name) + DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -179,7 +179,7 @@ def test_rename_document_updates_upload_file_when_present(db_session_with_contai ) # Act - DocumentService.rename_document(dataset.id, document.id, new_name) + DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -210,7 +210,7 @@ def test_rename_document_does_not_update_upload_file_when_missing_id(db_session_ ) # Act - DocumentService.rename_document(dataset.id, document.id, new_name) + DocumentService.rename_document(dataset.id, document.id, new_name, session=db_session_with_containers) # Assert db_session_with_containers.refresh(document) @@ -226,7 +226,7 @@ def test_rename_document_dataset_not_found(db_session_with_containers, mock_env) # Act / Assert with pytest.raises(ValueError, match="Dataset not found"): - DocumentService.rename_document(missing_dataset_id, str(uuid4()), "x") + DocumentService.rename_document(missing_dataset_id, str(uuid4()), "x", session=db_session_with_containers) def test_rename_document_not_found(db_session_with_containers, mock_env): @@ -236,7 +236,7 @@ def test_rename_document_not_found(db_session_with_containers, mock_env): # Act / Assert with pytest.raises(ValueError, match="Document not found"): - DocumentService.rename_document(dataset.id, str(uuid4()), "x") + DocumentService.rename_document(dataset.id, str(uuid4()), "x", session=db_session_with_containers) def test_rename_document_permission_denied_when_tenant_mismatch(db_session_with_containers, mock_env): @@ -251,4 +251,4 @@ def test_rename_document_permission_denied_when_tenant_mismatch(db_session_with_ # Act / Assert with pytest.raises(ValueError, match="No permission"): - DocumentService.rename_document(dataset.id, document.id, "x") + DocumentService.rename_document(dataset.id, document.id, "x", session=db_session_with_containers) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py index 0b4ce39bafb..e8bb6f88674 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py @@ -1,5 +1,5 @@ import inspect -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -742,7 +742,7 @@ class TestDocumentRetryApi: resp, status = method(api, "ds-1") assert status == 204 - retry_mock.assert_called_once_with("ds-1", []) + retry_mock.assert_called_once_with("ds-1", [], ANY) def test_retry_success(self, app: Flask, patch_tenant, patch_dataset): api = DocumentRetryApi() @@ -771,7 +771,7 @@ class TestDocumentRetryApi: response, status = method(api, "ds-1") assert status == 204 - retry_mock.assert_called_once_with("ds-1", [document]) + retry_mock.assert_called_once_with("ds-1", [document], ANY) def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_dataset): api = DocumentRetryApi() @@ -796,7 +796,7 @@ class TestDocumentRetryApi: response, status = method(api, "ds-1") assert status == 204 - retry_mock.assert_called_once_with("ds-1", []) + retry_mock.assert_called_once_with("ds-1", [], ANY) class TestDocumentPipelineExecutionLogApi: diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py index 3b9a1dcc5ea..6288fe363f5 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py @@ -107,7 +107,7 @@ def _wire_common_success_mocks( import services.dataset_service as dataset_service_module # Return a dataset object and allow permission checks to pass. - monkeypatch.setattr(module.DatasetService, "get_dataset", lambda _dataset_id: SimpleNamespace(id="ds-1")) + monkeypatch.setattr(module.DatasetService, "get_dataset", lambda *_args, **_kwargs: SimpleNamespace(id="ds-1")) monkeypatch.setattr(module.DatasetService, "check_dataset_permission", lambda *_args, **_kwargs: None) # Return a document that will be validated inside DocumentResource.get_document. @@ -150,7 +150,7 @@ def test_batch_download_zip_returns_send_file( """Ensure batch ZIP download returns a zip attachment via `send_file`.""" monkeypatch.setattr( - datasets_document_module.DatasetService, "get_dataset", lambda _dataset_id: SimpleNamespace(id="ds-1") + datasets_document_module.DatasetService, "get_dataset", lambda *_args, **_kwargs: SimpleNamespace(id="ds-1") ) monkeypatch.setattr( datasets_document_module.DatasetService, "check_dataset_permission", lambda *_args, **_kwargs: None @@ -218,7 +218,7 @@ def test_batch_download_zip_response_is_openable_zip( # Arrange: same controller mocks as the lightweight send_file test, but we keep the real `send_file`. monkeypatch.setattr( - datasets_document_module.DatasetService, "get_dataset", lambda _dataset_id: SimpleNamespace(id="ds-1") + datasets_document_module.DatasetService, "get_dataset", lambda *_args, **_kwargs: SimpleNamespace(id="ds-1") ) monkeypatch.setattr( datasets_document_module.DatasetService, "check_dataset_permission", lambda *_args, **_kwargs: None @@ -284,7 +284,7 @@ def test_batch_download_zip_rejects_non_upload_file_document( """Ensure batch ZIP download rejects non upload-file documents.""" monkeypatch.setattr( - datasets_document_module.DatasetService, "get_dataset", lambda _dataset_id: SimpleNamespace(id="ds-1") + datasets_document_module.DatasetService, "get_dataset", lambda *_args, **_kwargs: SimpleNamespace(id="ds-1") ) monkeypatch.setattr( datasets_document_module.DatasetService, "check_dataset_permission", lambda *_args, **_kwargs: None diff --git a/api/tests/unit_tests/controllers/console/test_spec.py b/api/tests/unit_tests/controllers/console/test_spec.py index ed02923caf6..58d7027751b 100644 --- a/api/tests/unit_tests/controllers/console/test_spec.py +++ b/api/tests/unit_tests/controllers/console/test_spec.py @@ -1,6 +1,8 @@ from inspect import unwrap from unittest.mock import patch +import pytest + import controllers.console.spec as spec_module @@ -22,7 +24,7 @@ class TestSpecSchemaDefinitionsApi: assert status == 200 assert resp == schema_definitions - def test_get_exception_returns_empty_list(self, caplog): + def test_get_exception_returns_empty_list(self, caplog: pytest.LogCaptureFixture): api = spec_module.SpecSchemaDefinitionsApi() method = unwrap(api.get) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py index 5eb76e309c1..9170f38df2a 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py @@ -356,9 +356,13 @@ class TestSegmentServiceMockedBehavior: """Test segment creation returns list of segments.""" mock_segments = [Mock(spec=DocumentSegment), Mock(spec=DocumentSegment)] mock_create.return_value = mock_segments + session = Mock() result = SegmentService.multi_create_segment( - segments=[{"content": "Test"}, {"content": "Test 2"}], document=mock_document, dataset=mock_dataset + segments=[{"content": "Test"}, {"content": "Test 2"}], + document=mock_document, + dataset=mock_dataset, + session=session, ) assert result is not None @@ -385,8 +389,13 @@ class TestSegmentServiceMockedBehavior: def test_get_segment_by_id_returns_segment(self, mock_get, mock_segment): """Test get_segment_by_id returns segment.""" mock_get.return_value = mock_segment + session = Mock() - result = SegmentService.get_segment_by_id(segment_id=mock_segment.id, tenant_id=mock_segment.tenant_id) + result = SegmentService.get_segment_by_id( + segment_id=mock_segment.id, + tenant_id=mock_segment.tenant_id, + session=session, + ) assert result == mock_segment @@ -394,16 +403,22 @@ class TestSegmentServiceMockedBehavior: def test_get_segment_by_id_returns_none_when_not_found(self, mock_get): """Test get_segment_by_id returns None when not found.""" mock_get.return_value = None + session = Mock() - result = SegmentService.get_segment_by_id(segment_id=str(uuid.uuid4()), tenant_id=str(uuid.uuid4())) + result = SegmentService.get_segment_by_id( + segment_id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + session=session, + ) assert result is None @patch.object(SegmentService, "delete_segment") def test_delete_segment_called(self, mock_delete, mock_segment, mock_document, mock_dataset): """Test segment deletion is called.""" - SegmentService.delete_segment(mock_segment, mock_document, mock_dataset) - mock_delete.assert_called_once_with(mock_segment, mock_document, mock_dataset) + session = Mock() + SegmentService.delete_segment(mock_segment, mock_document, mock_dataset, session) + mock_delete.assert_called_once_with(mock_segment, mock_document, mock_dataset, session) class TestChildChunkServiceMockedBehavior: @@ -431,7 +446,11 @@ class TestChildChunkServiceMockedBehavior: mock_create.return_value = mock_child_chunk result = SegmentService.create_child_chunk( - content="New chunk content", segment=mock_segment, document=Mock(spec=Document), dataset=Mock(spec=Dataset) + content="New chunk content", + segment=mock_segment, + document=Mock(spec=Document), + dataset=Mock(spec=Dataset), + session=Mock(), ) assert result == mock_child_chunk @@ -462,7 +481,9 @@ class TestChildChunkServiceMockedBehavior: mock_get.return_value = mock_child_chunk result = SegmentService.get_child_chunk_by_id( - child_chunk_id=mock_child_chunk.id, tenant_id=mock_child_chunk.tenant_id + child_chunk_id=mock_child_chunk.id, + tenant_id=mock_child_chunk.tenant_id, + session=Mock(), ) assert result == mock_child_chunk @@ -480,6 +501,7 @@ class TestChildChunkServiceMockedBehavior: segment=Mock(spec=DocumentSegment), document=Mock(spec=Document), dataset=Mock(spec=Dataset), + session=Mock(), ) assert result.content == "Updated content" @@ -1156,7 +1178,7 @@ class TestDatasetSegmentApiDelete: # Assert assert response == ("", 204) - mock_seg_svc.delete_segment.assert_called_once_with(mock_segment, mock_doc, mock_dataset) + mock_seg_svc.delete_segment.assert_called_once_with(mock_segment, mock_doc, mock_dataset, mock_db.session) @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py index 16b54acd8c6..e68eb647063 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py @@ -269,7 +269,7 @@ class TestDocumentService: mock_doc.indexing_status = "completed" mock_get.return_value = mock_doc - result = DocumentService.get_document(dataset_id="dataset_id", document_id="doc_id") + result = DocumentService.get_document(dataset_id="dataset_id", document_id="doc_id", session=Mock()) assert result is not None assert result.name == "Test Document" assert result.indexing_status == "completed" @@ -278,8 +278,9 @@ class TestDocumentService: def test_delete_document_called(self, mock_delete): """Test delete_document is called with document.""" mock_doc = Mock() - DocumentService.delete_document(document=mock_doc) - mock_delete.assert_called_once_with(document=mock_doc) + session = Mock() + DocumentService.delete_document(document=mock_doc, session=session) + mock_delete.assert_called_once_with(document=mock_doc, session=session) class TestDocumentIndexingStatus: @@ -454,24 +455,24 @@ class TestDocumentDisplayStatusLogic: class TestDocumentServiceBatchMethods: """Test DocumentService batch operations.""" - @patch("services.dataset_service.db.session.scalars") - def test_get_documents_by_ids(self, mock_scalars): + def test_get_documents_by_ids(self): """Test batch retrieval of documents by IDs.""" dataset_id = str(uuid.uuid4()) doc_ids = [str(uuid.uuid4()), str(uuid.uuid4())] mock_result = Mock() mock_result.all.return_value = [Mock(id=doc_ids[0]), Mock(id=doc_ids[1])] - mock_scalars.return_value = mock_result + session = Mock() + session.scalars.return_value = mock_result - documents = DocumentService.get_documents_by_ids(dataset_id, doc_ids) + documents = DocumentService.get_documents_by_ids(dataset_id, doc_ids, session) assert len(documents) == 2 - mock_scalars.assert_called_once() + session.scalars.assert_called_once() def test_get_documents_by_ids_empty(self): """Test batch retrieval with empty list returns empty.""" - assert DocumentService.get_documents_by_ids("ds_id", []) == [] + assert DocumentService.get_documents_by_ids("ds_id", [], Mock()) == [] class TestDocumentServiceFileOperations: @@ -487,7 +488,7 @@ class TestDocumentServiceFileOperations: mock_get_file.return_value = mock_file mock_signed_url.return_value = "https://example.com/download" - url = DocumentService.get_document_download_url(mock_doc) + url = DocumentService.get_document_download_url(mock_doc, Mock()) assert url == "https://example.com/download" mock_signed_url.assert_called_with(upload_file_id="file_id", as_attachment=True) @@ -516,7 +517,7 @@ class TestDocumentServiceSaveValidation: # Skip actual logic by mocking dependent calls or raising error to stop early with pytest.raises(TestStopError): # We just want to check check_doc_form is called early - DocumentService.save_document_with_dataset_id(dataset, config, Mock()) + DocumentService.save_document_with_dataset_id(dataset, config, Mock(), session=Mock()) # This will fail if we raise exception before check_doc_form, # but check_doc_form is the first thing called. @@ -782,7 +783,7 @@ class TestDocumentApiDelete: # Assert assert response == ("", 204) - mock_doc_svc.delete_document.assert_called_once_with(mock_document) + mock_doc_svc.delete_document.assert_called_once_with(mock_document, mock_db.session) @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") diff --git a/api/tests/unit_tests/extensions/test_ext_request_logging.py b/api/tests/unit_tests/extensions/test_ext_request_logging.py index 70e80707882..664de8cbd8b 100644 --- a/api/tests/unit_tests/extensions/test_ext_request_logging.py +++ b/api/tests/unit_tests/extensions/test_ext_request_logging.py @@ -1,6 +1,7 @@ import json import logging from unittest import mock +from unittest.mock import MagicMock import pytest from flask import Flask, Response @@ -73,8 +74,8 @@ class TestRequestLoggingExtension: def test_receiver_should_not_be_invoked_if_configuration_is_disabled( self, monkeypatch: pytest.MonkeyPatch, - mock_request_receiver, - mock_response_receiver, + mock_request_receiver: MagicMock, + mock_response_receiver: MagicMock, ): monkeypatch.setattr(dify_config, "ENABLE_REQUEST_LOGGING", False) @@ -90,8 +91,8 @@ class TestRequestLoggingExtension: def test_receiver_should_be_called_if_enabled( self, enable_request_logging, - mock_request_receiver, - mock_response_receiver, + mock_request_receiver: MagicMock, + mock_response_receiver: MagicMock, ): """ Test the request logging extension with JSON data. diff --git a/api/tests/unit_tests/services/test_dataset_service_dataset.py b/api/tests/unit_tests/services/test_dataset_service_dataset.py index 044e0e5ab40..6560e5a1c0b 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -18,7 +18,6 @@ from .dataset_service_test_helpers import ( TenantAccountRole, _make_knowledge_configuration, _make_retrieval_model, - _make_session_context, json, patch, pytest, @@ -345,7 +344,9 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session.scalar.return_value = object() with pytest.raises(DatasetNameDuplicateError, match="Dataset with name Dataset already exists"): - DatasetService.create_empty_dataset("tenant-1", "Dataset", None, "economy", account) + DatasetService.create_empty_dataset( + "tenant-1", "Dataset", None, "economy", account, session=mock_db.session + ) def test_create_empty_dataset_uses_default_embedding_model_for_high_quality_dataset(self): account = SimpleNamespace(id="user-1") @@ -370,6 +371,7 @@ class TestDatasetServiceCreationAndUpdate: description="Description", indexing_technique="high_quality", account=account, + session=mock_db.session, ) assert dataset.embedding_model_provider == "provider" @@ -421,6 +423,7 @@ class TestDatasetServiceCreationAndUpdate: embedding_model_name="embedding-model", retrieval_model=retrieval_model, summary_index_setting={"enable": True}, + session=mock_db.session, ) assert dataset.embedding_model_provider == "provider" @@ -451,7 +454,7 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session.scalar.return_value = object() with pytest.raises(DatasetNameDuplicateError, match="Existing Dataset already exists"): - DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity) + DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, mock_db.session) def test_create_empty_rag_pipeline_dataset_generates_name_and_creates_dataset(self): entity = RagPipelineDatasetCreateEntity( @@ -482,7 +485,7 @@ class TestDatasetServiceCreationAndUpdate: SimpleNamespace(name="Untitled 1"), ] - dataset = DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity) + dataset = DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, mock_db.session) assert entity.name == "Untitled 2" assert dataset.pipeline_id == "pipeline-1" @@ -505,12 +508,13 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session.scalar.return_value = None with pytest.raises(ValueError, match="Current user or current user id not found"): - DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity) + DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, mock_db.session) def test_update_dataset_raises_when_dataset_is_missing(self): + session = MagicMock() with patch.object(DatasetService, "get_dataset", return_value=None): with pytest.raises(ValueError, match="Dataset not found"): - DatasetService.update_dataset("dataset-1", {}, SimpleNamespace(id="user-1")) + DatasetService.update_dataset("dataset-1", {}, SimpleNamespace(id="user-1"), session) def test_update_dataset_raises_when_new_name_conflicts(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1", tenant_id="tenant-1") @@ -521,7 +525,12 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_has_dataset_same_name", return_value=True), ): with pytest.raises(ValueError, match="Dataset name already exists"): - DatasetService.update_dataset("dataset-1", {"name": "New Dataset"}, SimpleNamespace(id="user-1")) + DatasetService.update_dataset( + "dataset-1", + {"name": "New Dataset"}, + SimpleNamespace(id="user-1"), + MagicMock(), + ) def test_update_dataset_routes_external_datasets_to_external_helper(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1", tenant_id="tenant-1") @@ -533,13 +542,14 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "check_dataset_permission") as check_permission, patch.object(DatasetService, "_update_external_dataset", return_value="updated") as update_external, ): - result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user) + session = MagicMock() + result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session) assert result == "updated" check_permission.assert_called_once() assert check_permission.call_args.args[:2] == (dataset, user) assert len(check_permission.call_args.args) == 3 - update_external.assert_called_once_with(dataset, {"name": dataset.name}, user) + update_external.assert_called_once_with(dataset, {"name": dataset.name}, user, session) def test_update_dataset_routes_internal_datasets_to_internal_helper(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1", tenant_id="tenant-1") @@ -551,19 +561,20 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "check_dataset_permission") as check_permission, patch.object(DatasetService, "_update_internal_dataset", return_value="updated") as update_internal, ): - result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user) + session = MagicMock() + result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session) assert result == "updated" check_permission.assert_called_once() assert check_permission.call_args.args[:2] == (dataset, user) assert len(check_permission.call_args.args) == 3 - update_internal.assert_called_once_with(dataset, {"name": dataset.name}, user) + update_internal.assert_called_once_with(dataset, {"name": dataset.name}, user, session) def test_has_dataset_same_name_returns_true_when_query_matches(self): with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = object() - result = DatasetService._has_dataset_same_name("tenant-1", "dataset-1", "Dataset") + result = DatasetService._has_dataset_same_name("tenant-1", "dataset-1", "Dataset", mock_db.session) assert result is True @@ -592,6 +603,7 @@ class TestDatasetServiceCreationAndUpdate: "external_knowledge_api_id": "api-1", }, user, + mock_db.session, ) assert result is dataset @@ -603,7 +615,7 @@ class TestDatasetServiceCreationAndUpdate: assert dataset.updated_by == "user-1" assert dataset.updated_at is now get_external_knowledge_api.assert_called_once_with("api-1", dataset.tenant_id) - update_binding.assert_called_once_with("dataset-1", "knowledge-1", "api-1") + update_binding.assert_called_once_with("dataset-1", "knowledge-1", "api-1", mock_db.session) mock_db.session.add.assert_called_once_with(dataset) mock_db.session.commit.assert_called_once() @@ -618,7 +630,7 @@ class TestDatasetServiceCreationAndUpdate: dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") with pytest.raises(ValueError, match=message): - DatasetService._update_external_dataset(dataset, payload, SimpleNamespace(id="user-1")) + DatasetService._update_external_dataset(dataset, payload, SimpleNamespace(id="user-1"), MagicMock()) def test_update_external_dataset_rejects_cross_tenant_external_api_id(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") @@ -639,6 +651,7 @@ class TestDatasetServiceCreationAndUpdate: "external_knowledge_api_id": "foreign-api", }, SimpleNamespace(id="user-1"), + mock_db.session, ) get_external_knowledge_api.assert_called_once_with("foreign-api", dataset.tenant_id) @@ -650,16 +663,7 @@ class TestDatasetServiceCreationAndUpdate: session = MagicMock() session.scalar.return_value = binding session.add = MagicMock() - session_context = _make_session_context(session) - - mock_sessionmaker = MagicMock() - mock_sessionmaker.return_value.begin.return_value = session_context - - with ( - patch("services.dataset_service.db") as mock_db, - patch("services.dataset_service.sessionmaker", mock_sessionmaker), - ): - DatasetService._update_external_knowledge_binding("dataset-1", "new-knowledge", "new-api") + DatasetService._update_external_knowledge_binding("dataset-1", "new-knowledge", "new-api", session) assert binding.external_knowledge_id == "new-knowledge" assert binding.external_knowledge_api_id == "new-api" @@ -668,17 +672,8 @@ class TestDatasetServiceCreationAndUpdate: def test_update_external_knowledge_binding_raises_for_missing_binding(self): session = MagicMock() session.scalar.return_value = None - session_context = _make_session_context(session) - - mock_sessionmaker = MagicMock() - mock_sessionmaker.return_value.begin.return_value = session_context - - with ( - patch("services.dataset_service.db"), - patch("services.dataset_service.sessionmaker", mock_sessionmaker), - ): - with pytest.raises(ValueError, match="External knowledge binding not found"): - DatasetService._update_external_knowledge_binding("dataset-1", "knowledge-1", "api-1") + with pytest.raises(ValueError, match="External knowledge binding not found"): + DatasetService._update_external_knowledge_binding("dataset-1", "knowledge-1", "api-1", session) def test_update_internal_dataset_updates_fields_and_dispatches_regeneration_tasks(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") @@ -704,7 +699,7 @@ class TestDatasetServiceCreationAndUpdate: patch("services.dataset_service.deal_dataset_vector_index_task") as vector_task, patch("services.dataset_service.regenerate_summary_index_task") as regenerate_task, ): - result = DatasetService._update_internal_dataset(dataset, update_payload.copy(), user) + result = DatasetService._update_internal_dataset(dataset, update_payload.copy(), user, mock_db.session) assert result is dataset updated_values = mock_db.session.execute.call_args.args[0].compile().params @@ -721,7 +716,7 @@ class TestDatasetServiceCreationAndUpdate: assert "external_retrieval_model" not in updated_values mock_db.session.commit.assert_called_once() mock_db.session.refresh.assert_called_once_with(dataset) - update_pipeline.assert_called_once_with(dataset, "user-1") + update_pipeline.assert_called_once_with(dataset, "user-1", mock_db.session) vector_task.delay.assert_called_once_with("dataset-1", "update") regenerate_task.delay.assert_called_once_with( "dataset-1", @@ -733,7 +728,7 @@ class TestDatasetServiceCreationAndUpdate: dataset = SimpleNamespace(runtime_mode="workflow", pipeline_id="pipeline-1") with patch("services.dataset_service.db") as mock_db: - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1") + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) mock_db.session.get.assert_not_called() @@ -743,7 +738,7 @@ class TestDatasetServiceCreationAndUpdate: with patch("services.dataset_service.db") as mock_db: mock_db.session.get.return_value = None - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1") + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) mock_db.session.commit.assert_not_called() @@ -782,7 +777,7 @@ class TestDatasetServiceCreationAndUpdate: ): mock_db.session.get.return_value = pipeline - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1") + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) published_graph = json.loads(workflow_new.call_args.kwargs["graph"]) assert published_graph["nodes"][0]["data"]["embedding_model"] == "embedding-model" @@ -805,15 +800,16 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session.get.return_value = pipeline with pytest.raises(RuntimeError, match="boom"): - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1") + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) mock_db.session.rollback.assert_called_once() def test_handle_indexing_technique_change_returns_none_without_indexing_technique(self): filtered_data: dict[str, object] = {} dataset = SimpleNamespace(indexing_technique="economy") + session = MagicMock() - result = DatasetService._handle_indexing_technique_change(dataset, {}, filtered_data) + result = DatasetService._handle_indexing_technique_change(dataset, {}, filtered_data, session) assert result is None assert filtered_data == {} @@ -821,11 +817,13 @@ class TestDatasetServiceCreationAndUpdate: def test_handle_indexing_technique_change_switches_to_economy(self): filtered_data: dict[str, object] = {} dataset = SimpleNamespace(indexing_technique="high_quality") + session = MagicMock() result = DatasetService._handle_indexing_technique_change( dataset, {"indexing_technique": "economy"}, filtered_data, + session, ) assert result == "remove" @@ -838,20 +836,23 @@ class TestDatasetServiceCreationAndUpdate: def test_handle_indexing_technique_change_switches_to_high_quality(self): filtered_data: dict[str, object] = {} dataset = SimpleNamespace(indexing_technique="economy") + session = MagicMock() with patch.object(DatasetService, "_configure_embedding_model_for_high_quality") as configure_embedding: result = DatasetService._handle_indexing_technique_change( dataset, {"indexing_technique": "high_quality"}, filtered_data, + session, ) assert result == "add" - configure_embedding.assert_called_once_with({"indexing_technique": "high_quality"}, filtered_data) + configure_embedding.assert_called_once_with({"indexing_technique": "high_quality"}, filtered_data, session) def test_handle_indexing_technique_change_delegates_when_technique_is_unchanged(self): filtered_data: dict[str, object] = {} dataset = SimpleNamespace(indexing_technique="high_quality") + session = MagicMock() with patch.object( DatasetService, @@ -862,10 +863,16 @@ class TestDatasetServiceCreationAndUpdate: dataset, {"indexing_technique": "high_quality"}, filtered_data, + session, ) assert result == "update" - update_embedding.assert_called_once_with(dataset, {"indexing_technique": "high_quality"}, filtered_data) + update_embedding.assert_called_once_with( + dataset, + {"indexing_technique": "high_quality"}, + filtered_data, + session, + ) def test_configure_embedding_model_for_high_quality_updates_filtered_data(self): class FakeAccount: @@ -875,6 +882,7 @@ class TestDatasetServiceCreationAndUpdate: current_user.current_tenant_id = "tenant-1" embedding_model = SimpleNamespace(provider="provider", model_name="embedding-model") filtered_data: dict[str, object] = {} + session = MagicMock() with ( patch("services.dataset_service.Account", FakeAccount), @@ -890,6 +898,7 @@ class TestDatasetServiceCreationAndUpdate: DatasetService._configure_embedding_model_for_high_quality( {"embedding_model_provider": "provider", "embedding_model": "embedding-model"}, filtered_data, + session, ) assert filtered_data == { @@ -911,6 +920,7 @@ class TestDatasetServiceCreationAndUpdate: current_user = FakeAccount() current_user.current_tenant_id = "tenant-1" + session = MagicMock() with ( patch("services.dataset_service.Account", FakeAccount), @@ -923,6 +933,7 @@ class TestDatasetServiceCreationAndUpdate: DatasetService._configure_embedding_model_for_high_quality( {"embedding_model_provider": "provider", "embedding_model": "embedding-model"}, {}, + session, ) def test_handle_embedding_model_update_when_technique_unchanged_preserves_existing_settings(self): @@ -931,12 +942,14 @@ class TestDatasetServiceCreationAndUpdate: embedding_model="embedding-model", ) filtered_data: dict[str, object] = {} + session = MagicMock() with patch.object(DatasetService, "_preserve_existing_embedding_settings") as preserve_settings: result = DatasetService._handle_embedding_model_update_when_technique_unchanged( dataset, {}, filtered_data, + session, ) assert result is None @@ -947,16 +960,23 @@ class TestDatasetServiceCreationAndUpdate: embedding_model_provider="provider", embedding_model="embedding-model", ) + session = MagicMock() with patch.object(DatasetService, "_update_embedding_model_settings", return_value="update") as update_settings: result = DatasetService._handle_embedding_model_update_when_technique_unchanged( dataset, {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, {}, + session, ) assert result == "update" - update_settings.assert_called_once() + update_settings.assert_called_once_with( + dataset, + {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, + {}, + session, + ) def test_preserve_existing_embedding_settings_keeps_current_binding(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock( @@ -991,27 +1011,36 @@ class TestDatasetServiceCreationAndUpdate: embedding_model_provider="provider", embedding_model="embedding-model", ) + session = MagicMock() with patch.object(DatasetService, "_apply_new_embedding_settings") as apply_settings: result = DatasetService._update_embedding_model_settings( dataset, {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, {}, + session, ) assert result == "update" - apply_settings.assert_called_once() + apply_settings.assert_called_once_with( + dataset, + {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, + {}, + session, + ) def test_update_embedding_model_settings_returns_none_for_unchanged_values(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock( embedding_model_provider="provider", embedding_model="embedding-model", ) + session = MagicMock() result = DatasetService._update_embedding_model_settings( dataset, {"embedding_model_provider": "provider", "embedding_model": "embedding-model"}, {}, + session, ) assert result is None @@ -1021,6 +1050,7 @@ class TestDatasetServiceCreationAndUpdate: embedding_model_provider="provider", embedding_model="embedding-model", ) + session = MagicMock() with patch.object(DatasetService, "_apply_new_embedding_settings", side_effect=LLMBadRequestError()): with pytest.raises(ValueError, match="No Embedding Model available"): @@ -1028,6 +1058,7 @@ class TestDatasetServiceCreationAndUpdate: dataset, {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, {}, + session, ) def test_apply_new_embedding_settings_updates_binding_for_new_model(self): @@ -1038,6 +1069,7 @@ class TestDatasetServiceCreationAndUpdate: current_user.current_tenant_id = "tenant-1" dataset = DatasetServiceUnitDataFactory.create_dataset_mock(collection_binding_id="binding-1") filtered_data: dict[str, object] = {} + session = MagicMock() with ( patch("services.dataset_service.Account", FakeAccount), @@ -1057,6 +1089,7 @@ class TestDatasetServiceCreationAndUpdate: dataset, {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, filtered_data, + session, ) assert filtered_data == { @@ -1077,6 +1110,7 @@ class TestDatasetServiceCreationAndUpdate: collection_binding_id="binding-1", ) filtered_data: dict[str, object] = {} + session = MagicMock() with ( patch("services.dataset_service.Account", FakeAccount), @@ -1091,6 +1125,7 @@ class TestDatasetServiceCreationAndUpdate: dataset, {"embedding_model_provider": "provider-two", "embedding_model": "embedding-model-two"}, filtered_data, + session, ) assert filtered_data == { @@ -1380,11 +1415,21 @@ class TestDatasetServicePermissionsAndLifecycle: """Unit tests for dataset permissions, deletion, and metadata helpers.""" def test_check_dataset_operator_permission_validates_required_arguments(self): + session = MagicMock() + with pytest.raises(ValueError, match="Dataset not found"): - DatasetService.check_dataset_operator_permission(user=SimpleNamespace(id="user-1"), dataset=None) + DatasetService.check_dataset_operator_permission( + user=SimpleNamespace(id="user-1"), + dataset=None, + session=session, + ) with pytest.raises(ValueError, match="User not found"): - DatasetService.check_dataset_operator_permission(user=None, dataset=SimpleNamespace(id="dataset-1")) + DatasetService.check_dataset_operator_permission( + user=None, + dataset=SimpleNamespace(id="dataset-1"), + session=session, + ) class TestDatasetCollectionBindingService: @@ -1395,44 +1440,49 @@ class TestDatasetPermissionService: """Unit tests for dataset partial-member management helpers.""" def test_update_partial_member_list_rolls_back_on_exception(self): - with patch("services.dataset_service.db") as mock_db: - mock_db.session.add_all.side_effect = RuntimeError("boom") + session = MagicMock() + session.add_all.side_effect = RuntimeError("boom") - with pytest.raises(RuntimeError, match="boom"): - DatasetPermissionService.update_partial_member_list( - "tenant-1", - "dataset-1", - [{"user_id": "user-1"}], - ) + with pytest.raises(RuntimeError, match="boom"): + DatasetPermissionService.update_partial_member_list( + "tenant-1", + "dataset-1", + [{"user_id": "user-1"}], + session, + ) - mock_db.session.rollback.assert_called_once() + session.rollback.assert_called_once() def test_check_permission_requires_dataset_editor(self): user = SimpleNamespace(is_dataset_editor=False, is_dataset_operator=False) dataset = DatasetServiceUnitDataFactory.create_dataset_mock() + session = MagicMock() with pytest.raises(NoPermissionError, match="does not have permission"): - DatasetPermissionService.check_permission(user, dataset, "all_team", []) + DatasetPermissionService.check_permission(user, dataset, "all_team", [], session) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = DatasetServiceUnitDataFactory.create_dataset_mock(permission="all_team") + session = MagicMock() with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): - DatasetPermissionService.check_permission(user, dataset, "only_me", []) + DatasetPermissionService.check_permission(user, dataset, "only_me", [], session) def test_check_permission_requires_partial_member_list_for_partial_members_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = DatasetServiceUnitDataFactory.create_dataset_mock(permission="partial_members") + session = MagicMock() with pytest.raises(ValueError, match="Partial member list is required"): - DatasetPermissionService.check_permission(user, dataset, "partial_members", []) + DatasetPermissionService.check_permission(user, dataset, "partial_members", [], session) def test_check_permission_rejects_dataset_operator_member_list_changes(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) dataset = DatasetServiceUnitDataFactory.create_dataset_mock( dataset_id="dataset-1", permission="partial_members" ) + session = MagicMock() with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): @@ -1441,6 +1491,7 @@ class TestDatasetPermissionService: dataset, "partial_members", [{"user_id": "user-2"}], + session, ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged(self): @@ -1448,6 +1499,7 @@ class TestDatasetPermissionService: dataset = DatasetServiceUnitDataFactory.create_dataset_mock( dataset_id="dataset-1", permission="partial_members" ) + session = MagicMock() with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( @@ -1455,13 +1507,14 @@ class TestDatasetPermissionService: dataset, "partial_members", [{"user_id": "user-1"}], + session, ) def test_clear_partial_member_list_rolls_back_on_exception(self): - with patch("services.dataset_service.db") as mock_db: - mock_db.session.execute.side_effect = RuntimeError("boom") + session = MagicMock() + session.execute.side_effect = RuntimeError("boom") - with pytest.raises(RuntimeError, match="boom"): - DatasetPermissionService.clear_partial_member_list("dataset-1") + with pytest.raises(RuntimeError, match="boom"): + DatasetPermissionService.clear_partial_member_list("dataset-1", session) - mock_db.session.rollback.assert_called_once() + session.rollback.assert_called_once() diff --git a/api/tests/unit_tests/services/test_dataset_service_document.py b/api/tests/unit_tests/services/test_dataset_service_document.py index 9a8243936b3..c108b06ac6b 100644 --- a/api/tests/unit_tests/services/test_dataset_service_document.py +++ b/api/tests/unit_tests/services/test_dataset_service_document.py @@ -104,30 +104,34 @@ class TestDocumentServiceMutations: assert DocumentService.check_archived(document) is expected def test_rename_document_raises_when_dataset_is_missing(self, rename_account_context): + session = MagicMock() + with patch.object(DatasetService, "get_dataset", return_value=None): with pytest.raises(ValueError, match="Dataset not found"): - DocumentService.rename_document("dataset-1", "doc-1", "New Name") + DocumentService.rename_document("dataset-1", "doc-1", "New Name", session) def test_rename_document_raises_when_document_is_missing(self, rename_account_context): dataset = DatasetServiceUnitDataFactory.create_dataset_mock() + session = MagicMock() with ( patch.object(DatasetService, "get_dataset", return_value=dataset), patch.object(DocumentService, "get_document", return_value=None), ): with pytest.raises(ValueError, match="Document not found"): - DocumentService.rename_document(dataset.id, "doc-1", "New Name") + DocumentService.rename_document(dataset.id, "doc-1", "New Name", session) def test_rename_document_rejects_cross_tenant_access(self, rename_account_context): dataset = DatasetServiceUnitDataFactory.create_dataset_mock() document = DatasetServiceUnitDataFactory.create_document_mock(tenant_id="tenant-other") + session = MagicMock() with ( patch.object(DatasetService, "get_dataset", return_value=dataset), patch.object(DocumentService, "get_document", return_value=document), ): with pytest.raises(ValueError, match="No permission"): - DocumentService.rename_document(dataset.id, document.id, "New Name") + DocumentService.rename_document(dataset.id, document.id, "New Name", session) def test_rename_document_updates_document_metadata_and_upload_file_name(self, rename_account_context): dataset = DatasetServiceUnitDataFactory.create_dataset_mock( @@ -146,7 +150,7 @@ class TestDocumentServiceMutations: patch.object(DocumentService, "get_document", return_value=document), patch("services.dataset_service.db") as mock_db, ): - result = DocumentService.rename_document(dataset.id, document.id, "New Name") + result = DocumentService.rename_document(dataset.id, document.id, "New Name", mock_db.session) assert result is document assert document.name == "New Name" @@ -157,27 +161,30 @@ class TestDocumentServiceMutations: def test_recover_document_raises_when_document_is_not_paused(self): document = DatasetServiceUnitDataFactory.create_document_mock(is_paused=False) + session = MagicMock() with pytest.raises(DocumentIndexingError): - DocumentService.recover_document(document) + DocumentService.recover_document(document, session) def test_retry_document_raises_when_retry_flag_is_already_set(self): document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1") + session = MagicMock() with patch("services.dataset_service.redis_client") as mock_redis: mock_redis.get.return_value = "1" with pytest.raises(ValueError, match="being retried"): - DocumentService.retry_document("dataset-1", [document]) + DocumentService.retry_document("dataset-1", [document], session) def test_sync_website_document_raises_when_sync_flag_exists(self): document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1") + session = MagicMock() with patch("services.dataset_service.redis_client") as mock_redis: mock_redis.get.return_value = "1" with pytest.raises(ValueError, match="being synced"): - DocumentService.sync_website_document("dataset-1", document) + DocumentService.sync_website_document("dataset-1", document, session) def test_sync_website_document_updates_status_sets_cache_and_dispatches_task(self): document = DatasetServiceUnitDataFactory.create_document_mock( @@ -193,7 +200,7 @@ class TestDocumentServiceMutations: ): mock_redis.get.return_value = None - DocumentService.sync_website_document("dataset-1", document) + DocumentService.sync_website_document("dataset-1", document, mock_db.session) assert document.indexing_status == "waiting" assert '"mode": "scrape"' in document.data_source_info @@ -258,6 +265,7 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: tenant_id="tenant-1", knowledge_config=knowledge_config, account=account_context, + session=mock_db.session, ) assert dataset is created_dataset @@ -274,7 +282,12 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: == "useful for when you want to answer queries about the VeryLongDocumentNameForDataset.txt" ) dataset_cls.assert_called_once() - save_document.assert_called_once_with(created_dataset, knowledge_config, account_context) + save_document.assert_called_once_with( + created_dataset, + knowledge_config, + account_context, + session=mock_db.session, + ) assert mock_db.session.commit.call_count == 1 def test_save_document_without_dataset_id_uses_provided_retrieval_model(self, account_context): @@ -312,9 +325,14 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: "save_document_with_dataset_id", return_value=([SimpleNamespace(name="Doc")], "batch-1"), ), - patch("services.dataset_service.db"), + patch("services.dataset_service.db") as mock_db, ): - DocumentService.save_document_without_dataset_id("tenant-1", knowledge_config, account_context) + DocumentService.save_document_without_dataset_id( + "tenant-1", + knowledge_config, + account_context, + mock_db.session, + ) assert created_dataset.retrieval_model == retrieval_model.model_dump() assert created_dataset.collection_binding_id is None @@ -337,8 +355,9 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: ), patch.object(DocumentService, "check_documents_upload_quota") as check_quota, ): + session = MagicMock() with pytest.raises(ValueError, match="does not support batch upload"): - DocumentService.save_document_without_dataset_id("tenant-1", knowledge_config, account_context) + DocumentService.save_document_without_dataset_id("tenant-1", knowledge_config, account_context, session) check_quota.assert_not_called() @@ -367,13 +386,19 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: ) ), ) + session = MagicMock() with ( patch.object(DocumentService, "get_document", return_value=None), patch.object(DatasetService, "check_dataset_model_setting") as check_model_setting, ): with pytest.raises(NotFound, match="Document not found"): - DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=session, + ) check_model_setting.assert_called_once_with(dataset) @@ -390,13 +415,19 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: ) ), ) + session = MagicMock() with ( patch.object(DocumentService, "get_document", return_value=document), patch.object(DatasetService, "check_dataset_model_setting"), ): with pytest.raises(ValueError, match="Document is not available"): - DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=session, + ) def test_update_document_with_dataset_id_upload_file_process_rule_and_name_override(self, account_context): dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") @@ -433,7 +464,12 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: ): mock_db.session.scalar.return_value = SimpleNamespace(id="file-1", name="upload.txt") - result = DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + result = DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=mock_db.session, + ) assert result is document assert document.dataset_process_rule_id == "rule-2" @@ -481,7 +517,12 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: mock_db.session.scalar.return_value = None with pytest.raises(ValueError, match="Data source binding not found"): - DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=mock_db.session, + ) def test_update_document_with_dataset_id_website_crawl_updates_segments_and_dispatches_task(self, account_context): dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") @@ -510,7 +551,12 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.document_indexing_update_task") as update_task, ): - result = DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + result = DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=mock_db.session, + ) assert result is document assert document.data_source_type == "website_crawl" @@ -681,8 +727,14 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config = _make_upload_knowledge_config(file_ids=None) with patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=True)): + session = MagicMock() with pytest.raises(ValueError, match="File source info is required"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) def test_save_document_with_dataset_id_blocks_batch_upload_for_sandbox_plan(self, account_context): dataset = _make_dataset() @@ -695,8 +747,14 @@ class TestDocumentServiceSaveDocumentWithDatasetId: ), patch.object(DocumentService, "check_documents_upload_quota") as check_quota, ): + session = MagicMock() with pytest.raises(ValueError, match="does not support batch upload"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) check_quota.assert_not_called() @@ -709,8 +767,14 @@ class TestDocumentServiceSaveDocumentWithDatasetId: patch("services.dataset_service.dify_config.BATCH_UPLOAD_LIMIT", 1), patch.object(DocumentService, "check_documents_upload_quota") as check_quota, ): + session = MagicMock() with pytest.raises(ValueError, match="batch upload limit of 1"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) check_quota.assert_not_called() @@ -725,20 +789,32 @@ class TestDocumentServiceSaveDocumentWithDatasetId: DocumentService, "update_document_with_dataset_id", return_value=updated_document ) as update_document, ): - documents, batch = DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + session = MagicMock() + documents, batch = DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) assert dataset.data_source_type == "upload_file" assert documents == [updated_document] assert batch == "batch-existing" - update_document.assert_called_once_with(dataset, knowledge_config, account_context) + update_document.assert_called_once_with(dataset, knowledge_config, account_context, session=session) def test_save_document_with_dataset_id_requires_data_source_for_new_documents(self, account_context): dataset = _make_dataset() knowledge_config = _make_upload_knowledge_config(data_source=None) with patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)): + session = MagicMock() with pytest.raises(ValueError, match="Data source is required when creating new documents"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) def test_save_document_with_dataset_id_requires_existing_process_rule_for_custom_mode(self, account_context): dataset = _make_dataset(latest_process_rule=None) @@ -748,8 +824,14 @@ class TestDocumentServiceSaveDocumentWithDatasetId: ) with patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)): + session = MagicMock() with pytest.raises(ValueError, match="No process rule found"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) def test_save_document_with_dataset_id_rejects_invalid_indexing_technique(self, account_context): dataset = _make_dataset(indexing_technique=None) @@ -761,8 +843,14 @@ class TestDocumentServiceSaveDocumentWithDatasetId: ) with patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)): + session = MagicMock() with pytest.raises(ValueError, match="Indexing technique is invalid"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) def test_save_document_with_dataset_id_returns_empty_for_invalid_process_rule_mode(self, account_context): dataset = _make_dataset() @@ -770,7 +858,12 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config.process_rule = SimpleNamespace(mode="unsupported-mode", rules=None) with patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)): - documents, batch = DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + documents, batch = DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=MagicMock(), + ) assert documents == [] assert batch == "" @@ -807,6 +900,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config, account_context, dataset_process_rule=dataset_process_rule, + session=mock_db.session, ) assert documents == [duplicate_document, created_document] @@ -887,6 +981,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config, account_context, dataset_process_rule=dataset_process_rule, + session=mock_db.session, ) assert created_document in documents @@ -938,6 +1033,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config, account_context, dataset_process_rule=dataset_process_rule, + session=mock_db.session, ) assert documents == [first_document, second_document] @@ -990,7 +1086,7 @@ class TestDocumentServiceBatchUpdateStatus: with pytest.raises(DocumentIndexingError, match="Busy document is being indexed"): DocumentService.batch_update_document_status( - dataset, [document.id], "archive", SimpleNamespace(id="user-1") + dataset, [document.id], "archive", SimpleNamespace(id="user-1"), mock_db.session ) mock_db.session.commit.assert_not_called() @@ -1009,7 +1105,7 @@ class TestDocumentServiceBatchUpdateStatus: with pytest.raises(RuntimeError, match="commit failed"): DocumentService.batch_update_document_status( - dataset, [document.id], "enable", SimpleNamespace(id="user-1") + dataset, [document.id], "enable", SimpleNamespace(id="user-1"), mock_db.session ) mock_db.session.rollback.assert_called_once() @@ -1029,7 +1125,7 @@ class TestDocumentServiceBatchUpdateStatus: with pytest.raises(RuntimeError, match="task failed"): DocumentService.batch_update_document_status( - dataset, [document.id], "enable", SimpleNamespace(id="user-1") + dataset, [document.id], "enable", SimpleNamespace(id="user-1"), mock_db.session ) mock_db.session.commit.assert_called_once() @@ -1052,7 +1148,7 @@ class TestDocumentServiceTenantAndUpdateEdges: with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = 12 - result = DocumentService.get_tenant_documents_count() + result = DocumentService.get_tenant_documents_count(mock_db.session) assert result == 12 @@ -1091,7 +1187,12 @@ class TestDocumentServiceTenantAndUpdateEdges: process_rule_cls.return_value = created_process_rule mock_db.session.scalar.return_value = SimpleNamespace(id="file-1", name="upload.txt") - result = DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + result = DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=mock_db.session, + ) assert result is document assert document.dataset_process_rule_id == "rule-2" @@ -1117,8 +1218,14 @@ class TestDocumentServiceTenantAndUpdateEdges: patch.object(DocumentService, "get_document", return_value=_make_document()), patch.object(DatasetService, "check_dataset_model_setting"), ): + session = MagicMock() with pytest.raises(ValueError, match="No file info list found"): - DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=session, + ) def test_update_document_with_dataset_id_raises_when_upload_file_is_missing(self, account_context): dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") @@ -1141,7 +1248,12 @@ class TestDocumentServiceTenantAndUpdateEdges: mock_db.session.scalar.return_value = None with pytest.raises(FileNotExistsError): - DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=mock_db.session, + ) def test_update_document_with_dataset_id_requires_notion_info_list(self, account_context): dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") @@ -1155,8 +1267,14 @@ class TestDocumentServiceTenantAndUpdateEdges: patch.object(DocumentService, "get_document", return_value=_make_document()), patch.object(DatasetService, "check_dataset_model_setting"), ): + session = MagicMock() with pytest.raises(ValueError, match="No notion info list found"): - DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=session, + ) def test_update_document_with_dataset_id_notion_import_updates_page_info(self, account_context): dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") @@ -1191,7 +1309,12 @@ class TestDocumentServiceTenantAndUpdateEdges: ): mock_db.session.scalar.return_value = SimpleNamespace(id="binding-1") - result = DocumentService.update_document_with_dataset_id(dataset, document_data, account_context) + result = DocumentService.update_document_with_dataset_id( + dataset, + document_data, + account_context, + session=mock_db.session, + ) assert result is document assert document.data_source_type == "notion_import" @@ -1260,9 +1383,14 @@ class TestDocumentServiceSaveWithoutDatasetBilling: "save_document_with_dataset_id", return_value=([SimpleNamespace(name="Doc")], "batch-1"), ), - patch("services.dataset_service.db"), + patch("services.dataset_service.db") as mock_db, ): - DocumentService.save_document_without_dataset_id("tenant-1", knowledge_config, account_context) + DocumentService.save_document_without_dataset_id( + "tenant-1", + knowledge_config, + account_context, + mock_db.session, + ) check_quota.assert_called_once_with(3, features) @@ -1287,8 +1415,9 @@ class TestDocumentServiceSaveWithoutDatasetBilling: patch("services.dataset_service.dify_config.BATCH_UPLOAD_LIMIT", "1"), patch.object(DocumentService, "check_documents_upload_quota") as check_quota, ): + session = MagicMock() with pytest.raises(ValueError, match="batch upload limit of 1"): - DocumentService.save_document_without_dataset_id("tenant-1", knowledge_config, account_context) + DocumentService.save_document_without_dataset_id("tenant-1", knowledge_config, account_context, session) check_quota.assert_not_called() @@ -1458,7 +1587,13 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: provider="default-provider", ) - documents, batch = DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + session = MagicMock() + documents, batch = DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=session, + ) assert documents == [updated_document] assert batch == "batch-existing" @@ -1474,7 +1609,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: "top_k": 4, "score_threshold_enabled": False, } - get_binding.assert_called_once_with("default-provider", "default-embedding") + get_binding.assert_called_once_with("default-provider", "default-embedding", session) def test_save_document_with_dataset_id_uses_explicit_embedding_and_retrieval_model(self, account_context): dataset = _make_dataset(indexing_technique=None) @@ -1503,10 +1638,11 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: ) as get_binding, patch.object(DocumentService, "update_document_with_dataset_id", return_value=_make_document()), ): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + session = MagicMock() + DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context, session=session) model_manager_cls.for_tenant.return_value.get_default_model_instance.assert_not_called() - get_binding.assert_called_once_with("explicit-provider", "explicit-model") + get_binding.assert_called_once_with("explicit-provider", "explicit-model", session) assert dataset.embedding_model == "explicit-model" assert dataset.embedding_model_provider == "explicit-provider" assert dataset.retrieval_model == knowledge_config.retrieval_model.model_dump() @@ -1541,7 +1677,12 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: process_rule_cls.return_value = created_process_rule mock_db.session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] - documents, batch = DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + documents, batch = DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=mock_db.session, + ) assert documents == [created_document] assert batch == "20260101010101100023" @@ -1581,7 +1722,12 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: process_rule_cls.return_value = created_process_rule mock_db.session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=mock_db.session, + ) assert process_rule_cls.call_args.kwargs == { "dataset_id": "dataset-1", @@ -1615,7 +1761,12 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: process_rule_cls.return_value = created_process_rule mock_db.session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=mock_db.session, + ) assert process_rule_cls.call_args.kwargs == { "dataset_id": "dataset-1", @@ -1640,7 +1791,12 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: mock_db.session.scalars.return_value.all.return_value = [SimpleNamespace(id="file-1", name="file.txt")] with pytest.raises(FileNotExistsError, match="One or more files not found"): - DocumentService.save_document_with_dataset_id(dataset, knowledge_config, account_context) + DocumentService.save_document_with_dataset_id( + dataset, + knowledge_config, + account_context, + session=mock_db.session, + ) def test_save_document_with_dataset_id_requires_notion_info_list_for_notion_import(self, account_context): dataset = _make_dataset() @@ -1663,6 +1819,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: knowledge_config, account_context, dataset_process_rule=SimpleNamespace(id="rule-1"), + session=MagicMock(), ) def test_save_document_with_dataset_id_requires_website_info_list_for_website_crawl(self, account_context): @@ -1686,4 +1843,5 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: knowledge_config, account_context, dataset_process_rule=SimpleNamespace(id="rule-1"), + session=MagicMock(), ) diff --git a/api/tests/unit_tests/services/test_dataset_service_lock_not_owned.py b/api/tests/unit_tests/services/test_dataset_service_lock_not_owned.py index 352a765de28..5e5d406edb8 100644 --- a/api/tests/unit_tests/services/test_dataset_service_lock_not_owned.py +++ b/api/tests/unit_tests/services/test_dataset_service_lock_not_owned.py @@ -94,7 +94,9 @@ def test_save_document_with_dataset_id_ignores_lock_not_owned( # Avoid touching real doc_form logic monkeypatch.setattr("services.dataset_service.DatasetService.check_doc_form", lambda *a, **k: None) # Avoid real DB interactions - monkeypatch.setattr("services.dataset_service.db", Mock()) + db_mock = Mock() + db_mock.session = Mock() + monkeypatch.setattr("services.dataset_service.db", db_mock) # Act: this would hit the redis lock, whose __enter__ raises LockNotOwnedError. # Our implementation should catch it and still return (documents, batch). @@ -102,6 +104,7 @@ def test_save_document_with_dataset_id_ignores_lock_not_owned( dataset=dataset, knowledge_config=knowledge_config, account=account, + session=db_mock.session, ) # Assert @@ -148,7 +151,7 @@ def test_add_segment_ignores_lock_not_owned( monkeypatch.setattr("services.dataset_service.VectorService", Mock()) # Act - result = SegmentService.create_segment(args=args, document=document, dataset=dataset) + result = SegmentService.create_segment(args=args, document=document, dataset=dataset, session=db_mock.session) # Assert # Under LockNotOwnedError except, add_segment should swallow the error and return None. diff --git a/api/tests/unit_tests/services/test_dataset_service_segment.py b/api/tests/unit_tests/services/test_dataset_service_segment.py index 1f8586e32f3..a625f17ef37 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -51,7 +51,13 @@ class TestSegmentServiceChildChunks: mock_redis.lock.return_value = _make_lock_context() mock_db.session.scalar.return_value = 2 - child_chunk = SegmentService.create_child_chunk("child content", segment, document, dataset) + child_chunk = SegmentService.create_child_chunk( + "child content", + segment, + document, + dataset, + mock_db.session, + ) assert isinstance(child_chunk, ChildChunk) assert child_chunk.position == 3 @@ -79,7 +85,7 @@ class TestSegmentServiceChildChunks: vector_service.create_child_chunk_vector.side_effect = RuntimeError("vector failed") with pytest.raises(ChildChunkIndexingError, match="vector failed"): - SegmentService.create_child_chunk("child content", segment, document, dataset) + SegmentService.create_child_chunk("child content", segment, document, dataset, mock_db.session) mock_db.session.rollback.assert_called_once() mock_db.session.commit.assert_not_called() @@ -127,6 +133,7 @@ class TestSegmentServiceChildChunks: segment, document, dataset, + mock_db.session, ) assert [chunk.position for chunk in result] == [1, 3] @@ -164,6 +171,7 @@ class TestSegmentServiceChildChunks: segment, document, dataset, + mock_db.session, ) mock_db.session.rollback.assert_called_once() @@ -179,7 +187,7 @@ class TestSegmentServiceChildChunks: patch("services.dataset_service.VectorService") as vector_service, ): result = SegmentService.update_child_chunk( - "new content", child_chunk, _make_segment(), _make_document(), dataset + "new content", child_chunk, _make_segment(), _make_document(), dataset, mock_db.session ) assert result is child_chunk @@ -202,7 +210,7 @@ class TestSegmentServiceChildChunks: vector_service.delete_child_chunk_vector.side_effect = RuntimeError("delete failed") with pytest.raises(ChildChunkDeleteIndexError, match="delete failed"): - SegmentService.delete_child_chunk(child_chunk, dataset) + SegmentService.delete_child_chunk(child_chunk, dataset, mock_db.session) mock_db.session.delete.assert_called_once_with(child_chunk) mock_db.session.rollback.assert_called_once() @@ -247,13 +255,13 @@ class TestSegmentServiceQueries: with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = child_chunk - result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1") + result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1", mock_db.session) assert result is child_chunk with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = SimpleNamespace() - result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1") + result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1", mock_db.session) assert result is None @@ -294,13 +302,13 @@ class TestSegmentServiceQueries: segment.id = "segment-1" with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = segment - result = SegmentService.get_segment_by_id("segment-1", "tenant-1") + result = SegmentService.get_segment_by_id("segment-1", "tenant-1", mock_db.session) assert result is segment with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = SimpleNamespace() - result = SegmentService.get_segment_by_id("segment-1", "tenant-1") + result = SegmentService.get_segment_by_id("segment-1", "tenant-1", mock_db.session) assert result is None @@ -323,6 +331,7 @@ class TestSegmentServiceQueries: result = SegmentService.get_segments_by_document_and_dataset( document_id="doc-1", dataset_id="dataset-1", + session=mock_db.session, status="completed", enabled=True, ) @@ -409,7 +418,12 @@ class TestSegmentServiceMutations: mock_db.session.add.side_effect = add_side_effect vector_service.create_segments_vector.side_effect = RuntimeError("vector failed") - result = SegmentService.create_segment(args=args, document=document, dataset=dataset) + result = SegmentService.create_segment( + args=args, + document=document, + dataset=dataset, + session=mock_db.session, + ) created_segment = vector_service.create_segments_vector.call_args.args[1][0] attachment_bindings = [ @@ -459,7 +473,7 @@ class TestSegmentServiceMutations: mock_db.session.scalar.return_value = 1 vector_service.create_segments_vector.side_effect = RuntimeError("vector failed") - result = SegmentService.multi_create_segment(segments, document, dataset) + result = SegmentService.multi_create_segment(segments, document, dataset, mock_db.session) assert result assert len(result) == 2 @@ -488,7 +502,7 @@ class TestSegmentServiceMutations: ): mock_redis.get.return_value = None - result = SegmentService.update_segment(args, segment, document, dataset) + result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is segment assert segment.enabled is False @@ -508,7 +522,9 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = None with pytest.raises(ValueError, match="Can't update disabled segment"): - SegmentService.update_segment(SegmentUpdateArgs(content="new content"), segment, document, dataset) + SegmentService.update_segment( + SegmentUpdateArgs(content="new content"), segment, document, dataset, MagicMock() + ) def test_update_segment_rejects_when_indexing_cache_exists(self, account_context): segment = _make_segment(enabled=True) @@ -519,7 +535,9 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = "1" with pytest.raises(ValueError, match="Segment is indexing"): - SegmentService.update_segment(SegmentUpdateArgs(content="new content"), segment, document, dataset) + SegmentService.update_segment( + SegmentUpdateArgs(content="new content"), segment, document, dataset, MagicMock() + ) def test_update_segment_updates_keywords_for_same_content_segment(self, account_context): segment = _make_segment(content="same content", keywords=["old"]) @@ -536,7 +554,7 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = None mock_db.session.get.return_value = refreshed_segment - result = SegmentService.update_segment(args, segment, document, dataset) + result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is refreshed_segment assert segment.keywords == ["new"] @@ -575,7 +593,7 @@ class TestSegmentServiceMutations: # scalar call: existing_summary mock_db.session.scalar.return_value = existing_summary - result = SegmentService.update_segment(args, segment, document, dataset) + result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is refreshed_segment vector_service.generate_child_chunks.assert_called_once_with( @@ -617,7 +635,7 @@ class TestSegmentServiceMutations: mock_db.session.scalar.return_value = existing_summary mock_db.session.get.return_value = refreshed_segment - result = SegmentService.update_segment(args, segment, document, dataset) + result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is refreshed_segment assert segment.content == "new content" @@ -657,7 +675,7 @@ class TestSegmentServiceMutations: mock_db.session.scalar.return_value = existing_summary mock_db.session.get.return_value = refreshed_segment - result = SegmentService.update_segment(args, segment, document, dataset) + result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is refreshed_segment generate_summary.assert_called_once_with(segment, dataset, {"enable": True}) @@ -677,7 +695,7 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = None mock_db.session.scalars.return_value.all.return_value = ["child-1", "child-2"] - SegmentService.delete_segment(segment, document, dataset) + SegmentService.delete_segment(segment, document, dataset, mock_db.session) assert document.word_count == 6 mock_redis.setex.assert_called_once_with(f"segment_{segment.id}_delete_indexing", 600, 1) @@ -701,7 +719,7 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = "1" with pytest.raises(ValueError, match="Segment is deleting"): - SegmentService.delete_segment(segment, document, dataset) + SegmentService.delete_segment(segment, document, dataset, MagicMock()) def test_delete_segments_removes_records_and_clamps_document_word_count(self): dataset = _make_dataset() @@ -723,7 +741,7 @@ class TestSegmentServiceMutations: # scalars() for child_node_ids mock_db.session.scalars.return_value.all.return_value = ["child-1"] - SegmentService.delete_segments(["segment-1", "segment-2"], document, dataset) + SegmentService.delete_segments(["segment-1", "segment-2"], document, dataset, mock_db.session) assert document.word_count == 0 mock_db.session.add.assert_called_once_with(document) @@ -753,7 +771,9 @@ class TestSegmentServiceMutations: mock_db.session.scalars.return_value.all.return_value = [segment_a, segment_b] mock_redis.get.side_effect = [None, "1"] - SegmentService.update_segments_status(["segment-a", "segment-b"], "enable", dataset, document) + SegmentService.update_segments_status( + ["segment-a", "segment-b"], "enable", dataset, document, mock_db.session + ) assert segment_a.enabled is True assert segment_a.disabled_at is None @@ -780,7 +800,9 @@ class TestSegmentServiceMutations: mock_db.session.scalars.return_value.all.return_value = [segment_a, segment_b] mock_redis.get.side_effect = [None, "1"] - SegmentService.update_segments_status(["segment-a", "segment-b"], "disable", dataset, document) + SegmentService.update_segments_status( + ["segment-a", "segment-b"], "disable", dataset, document, mock_db.session + ) assert segment_a.enabled is False assert segment_a.disabled_at == "now" @@ -808,7 +830,7 @@ class TestSegmentServiceChildChunkTailHelpers: with pytest.raises(ChildChunkIndexingError, match="vector failed"): SegmentService.update_child_chunk( - "new content", child_chunk, SimpleNamespace(), SimpleNamespace(), dataset + "new content", child_chunk, SimpleNamespace(), SimpleNamespace(), dataset, mock_db.session ) mock_db.session.rollback.assert_called_once() @@ -822,7 +844,7 @@ class TestSegmentServiceChildChunkTailHelpers: patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.VectorService") as vector_service, ): - SegmentService.delete_child_chunk(child_chunk, dataset) + SegmentService.delete_child_chunk(child_chunk, dataset, mock_db.session) mock_db.session.delete.assert_called_once_with(child_chunk) vector_service.delete_child_chunk_vector.assert_called_once_with(child_chunk, dataset) @@ -860,6 +882,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: segment, document, dataset, + mock_db.session, ) assert result is refreshed_segment @@ -895,6 +918,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: segment, document, dataset, + mock_db.session, ) assert result is refreshed_segment @@ -943,6 +967,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: segment, document, dataset, + mock_db.session, ) assert result is refreshed_segment @@ -986,6 +1011,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: segment, document, dataset, + mock_db.session, ) assert result is refreshed_segment diff --git a/api/tests/unit_tests/services/test_summary_index_service.py b/api/tests/unit_tests/services/test_summary_index_service.py index cef11c0038d..19418c43926 100644 --- a/api/tests/unit_tests/services/test_summary_index_service.py +++ b/api/tests/unit_tests/services/test_summary_index_service.py @@ -1169,7 +1169,7 @@ def test_get_document_summary_status_detail_counts_and_previews(monkeypatch: pyt monkeypatch.setattr(SummaryIndexService, "get_document_summaries", MagicMock(return_value=[summary1])) - detail = SummaryIndexService.get_document_summary_status_detail("doc-1", "dataset-1") + detail = SummaryIndexService.get_document_summary_status_detail("doc-1", "dataset-1", MagicMock()) assert detail["total_segments"] == 2 assert detail["summary_status"]["completed"] == 1 assert detail["summary_status"]["not_started"] == 1 From 03d59dba47b504638db862131f4976c0e6823b5b Mon Sep 17 00:00:00 2001 From: Bond Zhu <37842169+MRZHUH@users.noreply.github.com> Date: Tue, 30 Jun 2026 15:19:48 +0800 Subject: [PATCH 04/54] fix(workflow): guard on_tool_execution stdout traces behind DEBUG (#38200) --- .../workflow_tool_callback_handler.py | 10 ++-- .../test_workflow_tool_callback_handler.py | 49 ++++++++++++++++--- 2 files changed, 49 insertions(+), 10 deletions(-) diff --git a/api/core/callback_handler/workflow_tool_callback_handler.py b/api/core/callback_handler/workflow_tool_callback_handler.py index 23aabd99708..8c48a62d93f 100644 --- a/api/core/callback_handler/workflow_tool_callback_handler.py +++ b/api/core/callback_handler/workflow_tool_callback_handler.py @@ -1,6 +1,7 @@ from collections.abc import Generator, Iterable, Mapping from typing import Any +from configs import dify_config from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler, print_text from core.ops.ops_trace_manager import TraceQueueManager from core.tools.entities.tool_entities import ToolInvokeMessage @@ -19,8 +20,9 @@ class DifyWorkflowCallbackHandler(DifyAgentCallbackHandler): trace_manager: TraceQueueManager | None = None, ) -> Generator[ToolInvokeMessage, None, None]: for tool_output in tool_outputs: - print_text("\n[on_tool_execution]\n", color=self.color) - print_text("Tool: " + tool_name + "\n", color=self.color) - print_text("Outputs: " + tool_output.model_dump_json()[:1000] + "\n", color=self.color) - print_text("\n") + if dify_config.DEBUG: + print_text("\n[on_tool_execution]\n", color=self.color) + print_text("Tool: " + tool_name + "\n", color=self.color) + print_text("Outputs: " + tool_output.model_dump_json()[:1000] + "\n", color=self.color) + print_text("\n") yield tool_output diff --git a/api/tests/unit_tests/core/callback_handler/test_workflow_tool_callback_handler.py b/api/tests/unit_tests/core/callback_handler/test_workflow_tool_callback_handler.py index 5b53c5965ce..81ac48b2036 100644 --- a/api/tests/unit_tests/core/callback_handler/test_workflow_tool_callback_handler.py +++ b/api/tests/unit_tests/core/callback_handler/test_workflow_tool_callback_handler.py @@ -32,8 +32,16 @@ def mock_print_text(mocker: MockerFixture): return mocker.patch("core.callback_handler.workflow_tool_callback_handler.print_text") +@pytest.fixture +def enable_debug(mocker: MockerFixture): + """Force DEBUG on so the handler emits its verbose stdout traces.""" + mocker.patch("core.callback_handler.workflow_tool_callback_handler.dify_config.DEBUG", True) + + class TestDifyWorkflowCallbackHandler: - def test_on_tool_execution_single_output_success(self, handler: DifyWorkflowCallbackHandler, mock_print_text): + def test_on_tool_execution_single_output_success( + self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug + ): # Arrange tool_name = "test_tool" tool_inputs = {"a": 1} @@ -63,7 +71,9 @@ class TestDifyWorkflowCallbackHandler: ] ) - def test_on_tool_execution_multiple_outputs(self, handler: DifyWorkflowCallbackHandler, mock_print_text): + def test_on_tool_execution_multiple_outputs( + self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug + ): # Arrange tool_name = "multi_tool" outputs = [ @@ -101,6 +111,29 @@ class TestDifyWorkflowCallbackHandler: assert results == [] mock_print_text.assert_not_called() + def test_on_tool_execution_skips_print_when_debug_disabled( + self, handler: DifyWorkflowCallbackHandler, mock_print_text, mocker: MockerFixture + ): + """When DEBUG is off, outputs are still yielded but nothing is printed + and model_dump_json() is never invoked.""" + # Arrange + mocker.patch("core.callback_handler.workflow_tool_callback_handler.dify_config.DEBUG", False) + message = MagicMock() + + # Act + results = list( + handler.on_tool_execution( + tool_name="quiet_tool", + tool_inputs={}, + tool_outputs=[message], + ) + ) + + # Assert + assert results == [message] + mock_print_text.assert_not_called() + message.model_dump_json.assert_not_called() + @pytest.mark.parametrize( ("invalid_outputs", "expected_exception"), [ @@ -110,7 +143,7 @@ class TestDifyWorkflowCallbackHandler: ], ) def test_on_tool_execution_invalid_outputs_type( - self, handler: DifyWorkflowCallbackHandler, invalid_outputs, expected_exception + self, handler: DifyWorkflowCallbackHandler, invalid_outputs, expected_exception, enable_debug ): # Arrange tool_name = "invalid_tool" @@ -125,7 +158,9 @@ class TestDifyWorkflowCallbackHandler: ) ) - def test_on_tool_execution_long_json_truncation(self, handler: DifyWorkflowCallbackHandler, mock_print_text): + def test_on_tool_execution_long_json_truncation( + self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug + ): # Arrange tool_name = "long_json_tool" long_json = "x" * 1500 @@ -147,7 +182,9 @@ class TestDifyWorkflowCallbackHandler: color="blue", ) - def test_on_tool_execution_model_dump_json_exception(self, handler: DifyWorkflowCallbackHandler, mock_print_text): + def test_on_tool_execution_model_dump_json_exception( + self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug + ): # Arrange tool_name = "exception_tool" bad_message = MagicMock() @@ -167,7 +204,7 @@ class TestDifyWorkflowCallbackHandler: assert mock_print_text.call_count >= 2 def test_on_tool_execution_none_message_id_and_trace_manager( - self, handler: DifyWorkflowCallbackHandler, mock_print_text + self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug ): # Arrange tool_name = "optional_params_tool" From 4303103304a3cc0b4ec1e0f07f6abf4d1e58984c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=9E=E6=B3=95=E6=93=8D=E4=BD=9C?= Date: Tue, 30 Jun 2026 15:29:20 +0800 Subject: [PATCH 05/54] fix: debug plugin permission setting not work (#38197) --- .../__tests__/use-reference-setting.spec.ts | 31 ++++++++++++++++--- .../plugin-page/use-reference-setting.ts | 4 ++- 2 files changed, 30 insertions(+), 5 deletions(-) diff --git a/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts b/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts index 2ab99dfe878..9d5958e2718 100644 --- a/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts +++ b/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts @@ -128,7 +128,7 @@ describe('useReferenceSetting Hook', () => { expect(result.current.canDebugger).toBe(true) }) - it('should ignore legacy admin permission for managers without plugin keys', () => { + it('should allow debug for managers with legacy admin permission when RBAC is disabled', () => { vi.mocked(useAppContext).mockReturnValue({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, @@ -146,10 +146,10 @@ describe('useReferenceSetting Hook', () => { const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool)) expect(result.current.canManagement).toBe(false) - expect(result.current.canDebugger).toBe(false) + expect(result.current.canDebugger).toBe(true) }) - it('should ignore legacy admin permission for owners without plugin keys', () => { + it('should allow debug for owners with legacy admin permission when RBAC is disabled', () => { vi.mocked(useAppContext).mockReturnValue({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: true, @@ -167,7 +167,30 @@ describe('useReferenceSetting Hook', () => { const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool)) expect(result.current.canManagement).toBe(false) - expect(result.current.canDebugger).toBe(false) + expect(result.current.canDebugger).toBe(true) + }) + + it('should allow debug for normal users when legacy debug permission is everyone and RBAC is disabled', () => { + vi.mocked(useAppContext).mockReturnValue({ + isCurrentWorkspaceManager: false, + isCurrentWorkspaceOwner: false, + langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, + workspacePermissionKeys: ['plugin.install'], + } as ReturnType) + + vi.mocked(usePluginPermissionSettings).mockReturnValue({ + data: { + install_permission: PermissionType.everyone, + debug_permission: PermissionType.everyone, + }, + } as ReturnType) + + const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool), { + systemFeatures: { rbac_enabled: false }, + }) + + expect(result.current.canDebugPlugin).toBe(true) + expect(result.current.canDebugger).toBe(true) }) it('should use plugin keys even when legacy admin permission is configured and RBAC is enabled', () => { diff --git a/web/app/components/plugins/plugin-page/use-reference-setting.ts b/web/app/components/plugins/plugin-page/use-reference-setting.ts index 20c306658fc..b2bce30f75c 100644 --- a/web/app/components/plugins/plugin-page/use-reference-setting.ts +++ b/web/app/components/plugins/plugin-page/use-reference-setting.ts @@ -52,7 +52,9 @@ export const usePluginSettingsAccess = () => { const canInstallPlugin = hasPermission(workspacePermissionKeys, 'plugin.install') && legacyCanInstallPlugin const canUpdatePlugin = hasPermission(workspacePermissionKeys, 'plugin.install') && legacyCanInstallPlugin const canDeletePlugin = hasPermission(workspacePermissionKeys, 'plugin.delete') && legacyCanInstallPlugin - const canDebugPlugin = hasPermission(workspacePermissionKeys, 'plugin.debug') && legacyCanDebugPlugin + const canDebugPlugin = rbacEnabled + ? hasPermission(workspacePermissionKeys, 'plugin.debug') + : legacyCanDebugPlugin return { permission: permissions, From b4be4d90a500ce08cb3e2367817d1bb15426ff02 Mon Sep 17 00:00:00 2001 From: Blackoutta <37723456+Blackoutta@users.noreply.github.com> Date: Tue, 30 Jun 2026 15:29:42 +0800 Subject: [PATCH 06/54] fix: stress test setup process and report structure workflow for Dify 1.15.0+ (#38194) --- scripts/stress-test/README.md | 54 +++++--- scripts/stress-test/common/config_helper.py | 36 +++++- scripts/stress-test/run_locust_stress_test.sh | 58 +++++++-- .../setup/configure_openai_plugin.py | 17 ++- scripts/stress-test/setup/create_api_key.py | 20 ++- .../stress-test/setup/import_workflow_app.py | 66 ++++++---- .../setup/install_openai_plugin.py | 34 +++-- scripts/stress-test/setup/login_admin.py | 33 +++-- scripts/stress-test/setup/publish_workflow.py | 19 ++- scripts/stress-test/setup/setup_admin.py | 31 +++-- scripts/stress-test/setup_all.py | 73 +++++++++-- scripts/stress-test/test_setup_scripts.py | 120 ++++++++++++++++++ 12 files changed, 444 insertions(+), 117 deletions(-) create mode 100644 scripts/stress-test/test_setup_scripts.py diff --git a/scripts/stress-test/README.md b/scripts/stress-test/README.md index 15f21cd5329..f3a5827ed11 100644 --- a/scripts/stress-test/README.md +++ b/scripts/stress-test/README.md @@ -84,10 +84,10 @@ The stress test tests a single endpoint with comprehensive SSE metrics tracking: ## Prerequisites -1. **Dependencies are automatically installed** when running setup: +1. **Dependencies**: - - Locust (load testing framework) - - sseclient-py (SSE client library) + - Locust runs through `uvx --from locust`, outside the API project environment. + - `sseclient-py` is included in the API project dependencies. 1. **Complete Dify setup**: @@ -96,6 +96,25 @@ The stress test tests a single endpoint with comprehensive SSE metrics tracking: python scripts/stress-test/setup_all.py ``` + For a brand-new Dify instance, the setup script creates the first admin account. Override the defaults if needed: + + ```bash + STRESS_TEST_ADMIN_EMAIL='your-admin@example.com' \ + STRESS_TEST_ADMIN_USERNAME='dify' \ + STRESS_TEST_ADMIN_PASSWORD='your-password' \ + python scripts/stress-test/setup_all.py + ``` + + For an already-initialized Dify instance with an admin account, provide the existing admin login: + + ```bash + STRESS_TEST_ADMIN_EMAIL='your-admin@example.com' \ + STRESS_TEST_ADMIN_PASSWORD='your-password' \ + python scripts/stress-test/setup_all.py + ``` + + `STRESS_TEST_ADMIN_USERNAME` is only used in the brand-new instance case, when `/console/api/setup` creates the first admin account. + 1. **Ensure services are running**: **IMPORTANT**: For accurate stress testing, run the API server with Gunicorn in production mode: @@ -141,11 +160,11 @@ The stress test tests a single endpoint with comprehensive SSE metrics tracking: # Run with default configuration (headless mode) ./scripts/stress-test/run_locust_stress_test.sh -# Or run directly with uv -uv run --project api python -m locust -f scripts/stress-test/sse_benchmark.py --host http://localhost:5001 +# Or run directly with uvx +uvx --from locust locust -f scripts/stress-test/sse_benchmark.py --host http://localhost:5001 # Run with Web UI (access at http://localhost:8089) -uv run --project api python -m locust -f scripts/stress-test/sse_benchmark.py --host http://localhost:5001 --web-port 8089 +uvx --from locust locust -f scripts/stress-test/sse_benchmark.py --host http://localhost:5001 --web-port 8089 ``` The script will: @@ -182,12 +201,13 @@ self.questions = [ ### Report Structure -After running the stress test, you'll find these files in the `reports/` directory: +After running the stress test, you'll find one directory per run under `reports/`: -- `locust_summary_YYYYMMDD_HHMMSS.txt` - Complete console output with metrics -- `locust_report_YYYYMMDD_HHMMSS.html` - Interactive HTML report with charts -- `locust_YYYYMMDD_HHMMSS_stats.csv` - CSV with detailed statistics -- `locust_YYYYMMDD_HHMMSS_stats_history.csv` - Time-series data +- `YYYYMMDD_HHMMSS/locust_summary.txt` - Complete console output with metrics +- `YYYYMMDD_HHMMSS/locust_report.html` - Interactive HTML report with charts +- `YYYYMMDD_HHMMSS/locust_stats.csv` - CSV with detailed statistics +- `YYYYMMDD_HHMMSS/locust_stats_history.csv` - Time-series data +- `YYYYMMDD_HHMMSS/sse_metrics_YYYYMMDD_HHMMSS.json` - Custom SSE metrics ### Key Metrics @@ -399,8 +419,8 @@ docker compose -f docker/docker-compose.middleware.yaml up -d db 1. **"ModuleNotFoundError: No module named 'locust'"**: ```bash - # Dependencies are installed automatically, but if needed: - uv --project api add --dev locust sseclient-py + # Locust is intentionally run outside the api project environment: + uvx --from locust locust --version ``` 1. **"API key configuration not found"**: @@ -453,15 +473,15 @@ Run Locust directly with custom options: ```bash # With specific user count and spawn rate -uv run --project api python -m locust -f scripts/stress-test/sse_benchmark.py \ +uvx --from locust locust -f scripts/stress-test/sse_benchmark.py \ --host http://localhost:5001 --users 50 --spawn-rate 5 # Generate CSV reports -uv run --project api python -m locust -f scripts/stress-test/sse_benchmark.py \ +uvx --from locust locust -f scripts/stress-test/sse_benchmark.py \ --host http://localhost:5001 --csv reports/results # Run for specific duration -uv run --project api python -m locust -f scripts/stress-test/sse_benchmark.py \ +uvx --from locust locust -f scripts/stress-test/sse_benchmark.py \ --host http://localhost:5001 --run-time 5m --headless ``` @@ -469,7 +489,7 @@ uv run --project api python -m locust -f scripts/stress-test/sse_benchmark.py \ ```bash # Compare multiple stress test runs -ls -la reports/stress_test_*.txt | tail -5 +ls -la scripts/stress-test/reports/*/locust_summary.txt | tail -5 ``` ## Interpreting Performance Issues diff --git a/scripts/stress-test/common/config_helper.py b/scripts/stress-test/common/config_helper.py index fffb5e00d80..4038259c736 100644 --- a/scripts/stress-test/common/config_helper.py +++ b/scripts/stress-test/common/config_helper.py @@ -8,9 +8,9 @@ from typing import NotRequired, TypedDict class AdminConfig(TypedDict): """Configuration for admin section.""" + email: str username: str password: str - base_url: str class AuthConfig(TypedDict): @@ -18,6 +18,7 @@ class AuthConfig(TypedDict): access_token: str refresh_token: NotRequired[str] + csrf_token: NotRequired[str] expires_at: NotRequired[int] @@ -253,18 +254,25 @@ class ConfigHelper: Returns: Access token string or None if not found """ - auth = self.get_state_section[AuthConfig]("auth") + auth = self.get_state_section("auth") if auth: return auth.get("access_token") return None + def get_csrf_token(self) -> str | None: + """Get the CSRF token from auth section.""" + auth = self.get_state_section("auth") + if auth: + return auth.get("csrf_token") + return None + def get_app_id(self) -> str | None: """Get the app ID from app section. Returns: App ID string or None if not found """ - app = self.get_state_section[AppConfig]("app") + app = self.get_state_section("app") if app: return app.get("app_id") return None @@ -275,11 +283,31 @@ class ConfigHelper: Returns: API key token string or None if not found """ - api_key = self.get_state_section[ApiKeyConfig]("api_key") + api_key = self.get_state_section("api_key") if api_key: return api_key.get("token") return None + def console_auth_headers(self) -> dict[str, str]: + access_token = self.get_token() + csrf_token = self.get_csrf_token() + headers: dict[str, str] = {} + if access_token: + headers["authorization"] = f"Bearer {access_token}" + if csrf_token: + headers["X-CSRF-Token"] = csrf_token + return headers + + def console_auth_cookies(self) -> dict[str, str]: + access_token = self.get_token() + csrf_token = self.get_csrf_token() + cookies = {"locale": "en-US"} + if access_token: + cookies["access_token"] = access_token + if csrf_token: + cookies["csrf_token"] = csrf_token + return cookies + # Create a default instance for convenience config_helper = ConfigHelper() diff --git a/scripts/stress-test/run_locust_stress_test.sh b/scripts/stress-test/run_locust_stress_test.sh index 665cb68754c..a861ada701c 100755 --- a/scripts/stress-test/run_locust_stress_test.sh +++ b/scripts/stress-test/run_locust_stress_test.sh @@ -21,10 +21,12 @@ NC='\033[0m' # No Color # Configuration TIMESTAMP=$(date +"%Y%m%d_%H%M%S") -REPORT_DIR="${STRESS_TEST_DIR}/reports" -CSV_PREFIX="${REPORT_DIR}/locust_${TIMESTAMP}" -HTML_REPORT="${REPORT_DIR}/locust_report_${TIMESTAMP}.html" -SUMMARY_REPORT="${REPORT_DIR}/locust_summary_${TIMESTAMP}.txt" +START_EPOCH=$(date +%s) +REPORT_ROOT="${STRESS_TEST_DIR}/reports" +REPORT_DIR="${REPORT_ROOT}/${TIMESTAMP}" +CSV_PREFIX="${REPORT_DIR}/locust" +HTML_REPORT="${REPORT_DIR}/locust_report.html" +SUMMARY_REPORT="${REPORT_DIR}/locust_summary.txt" # Create reports directory if it doesn't exist mkdir -p "${REPORT_DIR}" @@ -111,6 +113,7 @@ echo # Use SSE stress test script LOCUST_SCRIPT="${STRESS_TEST_DIR}/sse_benchmark.py" +LOCUST_RUN=(uvx --from locust locust) # Prepare Locust command if [ "$choice" = "2" ]; then @@ -119,7 +122,7 @@ if [ "$choice" = "2" ]; then echo # Run with web UI - uv --project api run locust \ + "${LOCUST_RUN[@]}" \ -f ${LOCUST_SCRIPT} \ --host http://localhost:5001 \ --web-port 8089 @@ -128,7 +131,7 @@ else echo # Run in headless mode with CSV output - uv --project api run locust \ + "${LOCUST_RUN[@]}" \ -f ${LOCUST_SCRIPT} \ --host http://localhost:5001 \ --users $USERS \ @@ -139,6 +142,22 @@ else --csv=$CSV_PREFIX \ --html=$HTML_REPORT \ 2>&1 | tee $SUMMARY_REPORT + SSE_METRICS_REPORT=$(python3 - <= ${START_EPOCH} +] +if reports: + source = max(reports, key=lambda p: p.stat().st_mtime) + target = Path("${REPORT_DIR}") / source.name + if source != target: + shutil.move(str(source), str(target)) + print(target) +EOF +) echo echo -e "${GREEN}═══════════════════════════════════════════════════════════════${NC}" @@ -150,6 +169,11 @@ else echo -e " ${YELLOW}HTML Report:${NC} $HTML_REPORT" echo -e " ${YELLOW}CSV Stats:${NC} ${CSV_PREFIX}_stats.csv" echo -e " ${YELLOW}CSV History:${NC} ${CSV_PREFIX}_stats_history.csv" + if [ -n "$SSE_METRICS_REPORT" ]; then + echo -e " ${YELLOW}SSE Metrics:${NC} $SSE_METRICS_REPORT" + else + echo -e " ${YELLOW}SSE Metrics:${NC} not found" + fi echo echo -e "${CYAN}View HTML report:${NC}" echo " open $HTML_REPORT # macOS" @@ -184,14 +208,22 @@ try: print(f" RPS: {row.get('Requests/s', 'N/A')}") break - # Show SSE-specific metrics print() print("SSE Streaming Metrics:") - for row in rows: - if 'Time to First Event' in row.get('Name', ''): - print(f" Time to First Event: {row.get('Median Response Time', 'N/A')} ms (median)") - elif 'Stream Duration' in row.get('Name', ''): - print(f" Stream Duration: {row.get('Median Response Time', 'N/A')} ms (median)") + import json + sse_metrics_report = "${SSE_METRICS_REPORT}" + if sse_metrics_report: + with open(sse_metrics_report, 'r') as metrics_file: + metrics = json.load(metrics_file).get("metrics", {}) + print(f" Total Connections: {metrics.get('total_connections', 'N/A')}") + print(f" Total Events: {metrics.get('total_events', 'N/A')}") + print(f" Connection Rate: {metrics.get('overall_conn_rate', 0):.2f} conn/s") + print(f" Event Throughput: {metrics.get('overall_event_rate', 0):.2f} events/s") + print(f" TTFE: {metrics.get('ttfe_p50', 0):.1f} ms p50 / {metrics.get('ttfe_p95', 0):.1f} ms p95") + print(f" Stream Duration: {metrics.get('stream_duration_p50', 0):.1f} ms p50 / {metrics.get('stream_duration_p95', 0):.1f} ms p95") + print(f" Inter-event Latency: {metrics.get('inter_event_latency_p50', 0):.1f} ms p50 / {metrics.get('inter_event_latency_p95', 0):.1f} ms p95") + else: + print(" No SSE metrics JSON report found for this run") except Exception as e: print(f"Could not parse metrics: {e}") @@ -199,4 +231,4 @@ EOF fi echo -e "${CYAN}═══════════════════════════════════════════════════════════════${NC}" -fi \ No newline at end of file +fi diff --git a/scripts/stress-test/setup/configure_openai_plugin.py b/scripts/stress-test/setup/configure_openai_plugin.py index 110417fa737..38a79ac7cd9 100755 --- a/scripts/stress-test/setup/configure_openai_plugin.py +++ b/scripts/stress-test/setup/configure_openai_plugin.py @@ -9,7 +9,7 @@ import httpx from common import Logger, config_helper -def configure_openai_plugin() -> None: +def configure_openai_plugin() -> bool: """Configure OpenAI plugin with mock server credentials.""" log = Logger("ConfigPlugin") @@ -20,7 +20,7 @@ def configure_openai_plugin() -> None: if not access_token: log.error("No access token found in config") log.info("Please run login_admin.py first to get access token") - return + return False log.step("Configuring OpenAI plugin with mock server...") @@ -50,14 +50,14 @@ def configure_openai_plugin() -> None: "Sec-Fetch-Mode": "cors", "Sec-Fetch-Site": "same-site", "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/139.0.0.0 Safari/537.36", - "authorization": f"Bearer {access_token}", + **config_helper.console_auth_headers(), "content-type": "application/json", "sec-ch-ua": '"Not;A=Brand";v="99", "Google Chrome";v="139", "Chromium";v="139"', "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": '"macOS"', } - cookies = {"locale": "en-US"} + cookies = config_helper.console_auth_cookies() try: # Make the configuration request @@ -73,25 +73,32 @@ def configure_openai_plugin() -> None: log.success("OpenAI plugin configured successfully!") log.key_value("API Base", config_payload["credentials"]["openai_api_base"]) log.key_value("API Key", config_payload["credentials"]["openai_api_key"]) + return True elif response.status_code == 201: log.success("OpenAI plugin credentials created successfully!") log.key_value("API Base", config_payload["credentials"]["openai_api_base"]) log.key_value("API Key", config_payload["credentials"]["openai_api_key"]) + return True elif response.status_code == 401: log.error("Configuration failed: Unauthorized") log.info("Token may have expired. Please run login_admin.py again") + return False else: log.error(f"Configuration failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - configure_openai_plugin() + if not configure_openai_plugin(): + sys.exit(1) diff --git a/scripts/stress-test/setup/create_api_key.py b/scripts/stress-test/setup/create_api_key.py index cd04fe57eb6..1bc42062c02 100755 --- a/scripts/stress-test/setup/create_api_key.py +++ b/scripts/stress-test/setup/create_api_key.py @@ -11,7 +11,7 @@ import httpx from common import Logger, config_helper -def create_api_key() -> None: +def create_api_key() -> bool: """Create API key for the imported app.""" log = Logger("CreateAPIKey") @@ -21,14 +21,14 @@ def create_api_key() -> None: access_token = config_helper.get_token() if not access_token: log.error("No access token found in config") - return + return False # Read app_id from config app_id = config_helper.get_app_id() if not app_id: log.error("No app_id found in config") log.info("Please run import_workflow_app.py first to import the app") - return + return False log.step(f"Creating API key for app: {app_id}") @@ -50,14 +50,14 @@ def create_api_key() -> None: "Sec-Fetch-Mode": "cors", "Sec-Fetch-Site": "same-site", "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/139.0.0.0 Safari/537.36", - "authorization": f"Bearer {access_token}", + **config_helper.console_auth_headers(), "content-type": "application/json", "sec-ch-ua": '"Not;A=Brand";v="99", "Google Chrome";v="139", "Chromium";v="139"', "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": '"macOS"', } - cookies = {"locale": "en-US"} + cookies = config_helper.console_auth_cookies() try: # Make the API key creation request @@ -91,23 +91,31 @@ def create_api_key() -> None: if config_helper.write_config("api_key_config", api_key_config): log.info(f"API key saved to: {config_helper.get_config_path('benchmark_state')}") + return True + return False else: log.error("No API token received") log.debug(f"Response: {json.dumps(response_data, indent=2)}") + return False elif response.status_code == 401: log.error("API key creation failed: Unauthorized") log.info("Token may have expired. Please run login_admin.py again") + return False else: log.error(f"API key creation failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - create_api_key() + if not create_api_key(): + sys.exit(1) diff --git a/scripts/stress-test/setup/import_workflow_app.py b/scripts/stress-test/setup/import_workflow_app.py index 41a76bd29be..2da40f806ac 100755 --- a/scripts/stress-test/setup/import_workflow_app.py +++ b/scripts/stress-test/setup/import_workflow_app.py @@ -11,7 +11,11 @@ import httpx from common import Logger, config_helper # type: ignore[import] -def import_workflow_app() -> None: +def is_successful_import_response(response_data: dict[str, object]) -> bool: + return response_data.get("status") != "failed" and bool(response_data.get("app_id")) + + +def import_workflow_app() -> bool: """Import workflow app from DSL file and save app_id.""" log = Logger("ImportApp") @@ -22,14 +26,14 @@ def import_workflow_app() -> None: if not access_token: log.error("No access token found in config") log.info("Please run login_admin.py first to get access token") - return + return False # Read workflow DSL file dsl_path = Path(__file__).parent / "dsl" / "workflow_llm.yml" if not dsl_path.exists(): log.error(f"DSL file not found: {dsl_path}") - return + return False with open(dsl_path) as f: yaml_content = f.read() @@ -57,14 +61,14 @@ def import_workflow_app() -> None: "Sec-Fetch-Mode": "cors", "Sec-Fetch-Site": "same-site", "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/139.0.0.0 Safari/537.36", - "authorization": f"Bearer {access_token}", + **config_helper.console_auth_headers(), "content-type": "application/json", "sec-ch-ua": '"Not;A=Brand";v="99", "Google Chrome";v="139", "Chromium";v="139"', "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": '"macOS"', } - cookies = {"locale": "en-US"} + cookies = config_helper.console_auth_cookies() try: # Make the import request @@ -79,50 +83,56 @@ def import_workflow_app() -> None: if response.status_code == 200: response_data = response.json() - # Check import status - if response_data.get("status") == "completed": + if is_successful_import_response(response_data): app_id = response_data.get("app_id") - - if app_id: - log.success("Workflow app imported successfully!") - log.key_value("App ID", app_id) - log.key_value("App Mode", response_data.get("app_mode")) - log.key_value("DSL Version", response_data.get("imported_dsl_version")) - - # Save app_id to config - app_config = { - "app_id": app_id, - "app_mode": response_data.get("app_mode"), - "app_name": "workflow_llm", - "dsl_version": response_data.get("imported_dsl_version"), - } - - if config_helper.write_config("app_config", app_config): - log.info(f"App config saved to: {config_helper.get_config_path('benchmark_state')}") - else: - log.error("Import completed but no app_id received") + if response_data.get("status") != "completed": + log.warning(f"Import status: {response_data.get('status')}") log.debug(f"Response: {json.dumps(response_data, indent=2)}") + log.success("Workflow app imported successfully!") + log.key_value("App ID", app_id) + log.key_value("App Mode", response_data.get("app_mode")) + log.key_value("DSL Version", response_data.get("imported_dsl_version")) + + # Save app_id to config + app_config = { + "app_id": app_id, + "app_mode": response_data.get("app_mode"), + "app_name": "workflow_llm", + "dsl_version": response_data.get("imported_dsl_version"), + } + + if config_helper.write_config("app_config", app_config): + log.info(f"App config saved to: {config_helper.get_config_path('benchmark_state')}") + return True + return False elif response_data.get("status") == "failed": log.error("Import failed") log.error(f"Error: {response_data.get('error')}") + return False else: - log.warning(f"Import status: {response_data.get('status')}") + log.error("Import response did not include app_id") log.debug(f"Response: {json.dumps(response_data, indent=2)}") + return False elif response.status_code == 401: log.error("Import failed: Unauthorized") log.info("Token may have expired. Please run login_admin.py again") + return False else: log.error(f"Import failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - import_workflow_app() + if not import_workflow_app(): + sys.exit(1) diff --git a/scripts/stress-test/setup/install_openai_plugin.py b/scripts/stress-test/setup/install_openai_plugin.py index 055e5661f81..8e8d5a634de 100755 --- a/scripts/stress-test/setup/install_openai_plugin.py +++ b/scripts/stress-test/setup/install_openai_plugin.py @@ -11,7 +11,11 @@ import httpx from common import Logger, config_helper -def install_openai_plugin() -> None: +def is_non_blocking_install_response(response_data: dict[str, object]) -> bool: + return not response_data.get("code") + + +def install_openai_plugin() -> bool: """Install OpenAI plugin using saved access token.""" log = Logger("InstallPlugin") @@ -22,7 +26,7 @@ def install_openai_plugin() -> None: if not access_token: log.error("No access token found in config") log.info("Please run login_admin.py first to get access token") - return + return False log.step("Installing OpenAI plugin...") @@ -50,14 +54,14 @@ def install_openai_plugin() -> None: "Sec-Fetch-Mode": "cors", "Sec-Fetch-Site": "same-site", "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/139.0.0.0 Safari/537.36", - "authorization": f"Bearer {access_token}", + **config_helper.console_auth_headers(), "content-type": "application/json", "sec-ch-ua": '"Not;A=Brand";v="99", "Google Chrome";v="139", "Chromium";v="139"', "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": '"macOS"', } - cookies = {"locale": "en-US"} + cookies = config_helper.console_auth_cookies() try: # Make the installation request @@ -74,8 +78,13 @@ def install_openai_plugin() -> None: task_id = response_data.get("task_id") if not task_id: + if is_non_blocking_install_response(response_data): + log.warning("No installation task returned; plugin may already be installed") + log.debug(f"Response: {response.text}") + return True log.error("No task ID received from installation request") - return + log.debug(f"Response: {response.text}") + return False log.progress(f"Installation task created: {task_id}") log.info("Polling for task completion...") @@ -103,7 +112,7 @@ def install_openai_plugin() -> None: success=False, message=f"Failed to get task status: {task_response.status_code}", ) - return + return False task_data = task_response.json() task_info = task_data.get("task", {}) @@ -119,7 +128,7 @@ def install_openai_plugin() -> None: plugin_info = plugins[0] log.key_value("Plugin ID", plugin_info.get("plugin_id")) log.key_value("Message", plugin_info.get("message")) - break + return True elif status == "failed": log.spinner_stop(success=False, message="Installation failed") @@ -128,30 +137,37 @@ def install_openai_plugin() -> None: if plugins: for plugin in plugins: log.list_item(f"{plugin.get('plugin_id')}: {plugin.get('message')}") - break + return False # Continue polling if status is "pending" or other else: log.spinner_stop(success=False, message="Installation timed out") log.error("Installation timed out after 60 seconds") + return False elif response.status_code == 401: log.error("Installation failed: Unauthorized") log.info("Token may have expired. Please run login_admin.py again") + return False elif response.status_code == 409: log.warning("Plugin may already be installed") log.debug(f"Response: {response.text}") + return True else: log.error(f"Installation failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - install_openai_plugin() + if not install_openai_plugin(): + sys.exit(1) diff --git a/scripts/stress-test/setup/login_admin.py b/scripts/stress-test/setup/login_admin.py index 572b8fb6500..751d850a45e 100755 --- a/scripts/stress-test/setup/login_admin.py +++ b/scripts/stress-test/setup/login_admin.py @@ -5,13 +5,19 @@ from pathlib import Path sys.path.append(str(Path(__file__).parent.parent)) +import base64 import json import httpx from common import Logger, config_helper -def login_admin() -> None: +def encode_sensitive_field(value: str) -> str: + """Encode fields the same way the web client does before login.""" + return base64.b64encode(value.encode("utf-8")).decode() + + +def login_admin() -> bool: """Login with admin account and save access token.""" log = Logger("Login") @@ -23,7 +29,7 @@ def login_admin() -> None: if not admin_config: log.error("Admin config not found") log.info("Please run setup_admin.py first to create the admin account") - return + return False log.info(f"Logging in with email: {admin_config['email']}") @@ -34,7 +40,7 @@ def login_admin() -> None: # Prepare login payload login_payload = { "email": admin_config["email"], - "password": admin_config["password"], + "password": encode_sensitive_field(admin_config["password"]), "remember_me": True, } @@ -50,29 +56,28 @@ def login_admin() -> None: if response.status_code == 200: log.success("Login successful!") - # Extract token from response response_data = response.json() # Check if login was successful if response_data.get("result") != "success": log.error(f"Login failed: {response_data}") - return + return False - # Extract tokens from data field - token_data = response_data.get("data", {}) - access_token = token_data.get("access_token", "") - refresh_token = token_data.get("refresh_token", "") + access_token = response.cookies.get("access_token", "") + refresh_token = response.cookies.get("refresh_token", "") + csrf_token = response.cookies.get("csrf_token", "") if not access_token: log.error("No access token found in response") log.debug(f"Full response: {json.dumps(response_data, indent=2)}") - return + return False # Save token to config file token_config = { "email": admin_config["email"], "access_token": access_token, "refresh_token": refresh_token, + "csrf_token": csrf_token, } # Save token config @@ -82,20 +87,26 @@ def login_admin() -> None: # Show truncated token for verification token_display = f"{access_token[:20]}..." if len(access_token) > 20 else "Token saved" log.key_value("Access token", token_display) + return True elif response.status_code == 401: log.error("Login failed: Invalid credentials") log.debug(f"Response: {response.text}") + return False else: log.error(f"Login failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - login_admin() + if not login_admin(): + sys.exit(1) diff --git a/scripts/stress-test/setup/publish_workflow.py b/scripts/stress-test/setup/publish_workflow.py index b772eccebdd..3e565ff7624 100755 --- a/scripts/stress-test/setup/publish_workflow.py +++ b/scripts/stress-test/setup/publish_workflow.py @@ -11,7 +11,7 @@ import httpx from common import Logger, config_helper -def publish_workflow() -> None: +def publish_workflow() -> bool: """Publish the imported workflow app.""" log = Logger("PublishWorkflow") @@ -21,13 +21,13 @@ def publish_workflow() -> None: access_token = config_helper.get_token() if not access_token: log.error("No access token found in config") - return + return False # Read app_id from config app_id = config_helper.get_app_id() if not app_id: log.error("No app_id found in config") - return + return False log.step(f"Publishing workflow for app: {app_id}") @@ -51,14 +51,14 @@ def publish_workflow() -> None: "Sec-Fetch-Mode": "cors", "Sec-Fetch-Site": "same-site", "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/139.0.0.0 Safari/537.36", - "authorization": f"Bearer {access_token}", + **config_helper.console_auth_headers(), "content-type": "application/json", "sec-ch-ua": '"Not;A=Brand";v="99", "Google Chrome";v="139", "Chromium";v="139"', "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": '"macOS"', } - cookies = {"locale": "en-US"} + cookies = config_helper.console_auth_cookies() try: # Make the publish request @@ -83,23 +83,30 @@ def publish_workflow() -> None: except json.JSONDecodeError: # Response might be empty or non-JSON pass + return True elif response.status_code == 401: log.error("Workflow publish failed: Unauthorized") log.info("Token may have expired. Please run login_admin.py again") + return False elif response.status_code == 404: log.error("Workflow publish failed: App not found") log.info("Make sure the app was imported successfully") + return False else: log.error(f"Workflow publish failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - publish_workflow() + if not publish_workflow(): + sys.exit(1) diff --git a/scripts/stress-test/setup/setup_admin.py b/scripts/stress-test/setup/setup_admin.py index a5e9161210c..6e5dbd373f9 100755 --- a/scripts/stress-test/setup/setup_admin.py +++ b/scripts/stress-test/setup/setup_admin.py @@ -5,22 +5,27 @@ from pathlib import Path sys.path.append(str(Path(__file__).parent.parent)) +import os + import httpx from common import Logger, config_helper -def setup_admin_account() -> None: +def build_admin_config() -> dict[str, str]: + return { + "email": os.getenv("STRESS_TEST_ADMIN_EMAIL", "test@dify.ai"), + "username": os.getenv("STRESS_TEST_ADMIN_USERNAME", "dify"), + "password": os.getenv("STRESS_TEST_ADMIN_PASSWORD", "password123"), + } + + +def setup_admin_account() -> bool: """Setup Dify API with an admin account.""" log = Logger("SetupAdmin") log.header("Setting up Admin Account") - # Admin account credentials - admin_config = { - "email": "test@dify.ai", - "username": "dify", - "password": "password123", - } + admin_config = build_admin_config() # Save credentials to config file if config_helper.write_config("admin_config", admin_config): @@ -52,20 +57,26 @@ def setup_admin_account() -> None: log.success("Admin account created successfully!") log.key_value("Email", admin_config["email"]) log.key_value("Username", admin_config["username"]) + return True - elif response.status_code == 400: - log.warning("Setup may have already been completed or invalid data provided") + elif response.status_code in {400, 403}: + log.warning("Setup may have already been completed") log.debug(f"Response: {response.text}") + return True else: log.error(f"Setup failed with status code: {response.status_code}") log.debug(f"Response: {response.text}") + return False except httpx.ConnectError: log.error("Could not connect to Dify API at http://localhost:5001") log.info("Make sure the API server is running with: ./dev/start-api") + return False except Exception as e: log.error(f"An error occurred: {e}") + return False if __name__ == "__main__": - setup_admin_account() + if not setup_admin_account(): + sys.exit(1) diff --git a/scripts/stress-test/setup_all.py b/scripts/stress-test/setup_all.py index ece420f9257..47a789c4317 100755 --- a/scripts/stress-test/setup_all.py +++ b/scripts/stress-test/setup_all.py @@ -1,12 +1,56 @@ #!/usr/bin/env python3 +import os import socket import subprocess import sys import time from pathlib import Path -from common import Logger, ProgressLogger +import httpx +from common import Logger, ProgressLogger, config_helper + + +def build_admin_config() -> dict[str, str]: + return { + "email": os.getenv("STRESS_TEST_ADMIN_EMAIL", "test@dify.ai"), + "username": os.getenv("STRESS_TEST_ADMIN_USERNAME", "dify"), + "password": os.getenv("STRESS_TEST_ADMIN_PASSWORD", "password123"), + } + + +def get_setup_step(base_url: str = "http://localhost:5001") -> str | None: + try: + response = httpx.get(f"{base_url}/console/api/setup", timeout=5) + if response.status_code == 200: + return response.json().get("step") + except (httpx.HTTPError, ValueError): + return None + return None + + +def confirm(prompt: str) -> bool: + answer = input(f"\n{prompt} [Y/n]: ").strip().lower() + return answer in ("", "y", "yes") + + +def confirm_admin_credentials(log: Logger) -> bool: + admin_config = build_admin_config() + setup_step = get_setup_step() + config_helper.write_config("admin_config", admin_config) + + if setup_step == "finished": + log.warning("Dify is already initialized; setup will use the existing admin account to log in.") + log.key_value("Admin email", admin_config["email"]) + log.info("Set STRESS_TEST_ADMIN_EMAIL and STRESS_TEST_ADMIN_PASSWORD if this is not the right account.") + return confirm("Continue with this admin login?") + + log.info("Dify is not initialized; setup will create the first admin account with:") + log.key_value("Admin email", admin_config["email"]) + log.key_value("Admin username", admin_config["username"]) + log.key_value("Admin password", admin_config["password"]) + log.info("Set STRESS_TEST_ADMIN_EMAIL, STRESS_TEST_ADMIN_USERNAME, or STRESS_TEST_ADMIN_PASSWORD to override.") + return confirm("Create/use this admin account?") def run_script(script_name: str, description: str) -> bool: @@ -89,19 +133,20 @@ def main() -> None: if not dify_running or not mock_running: print("\n⚠️ Both services must be running before proceeding.") - retry = input("\nWould you like to check again? (yes/no): ") - if retry.lower() in ["yes", "y"]: + if confirm("Would you like to check again?"): return main() # Recursively call main to check again else: print("❌ Setup cancelled. Please start the required services and try again.") sys.exit(1) log.success("All required services are running!") - input("\nPress Enter to continue with setup...") + if not confirm_admin_credentials(log): + print("❌ Setup cancelled. Please set the admin environment variables and try again.") + sys.exit(1) # Define setup steps + setup_step = get_setup_step() setup_steps = [ - ("setup_admin.py", "Creating admin account"), ("login_admin.py", "Logging in and getting access token"), ("install_openai_plugin.py", "Installing OpenAI plugin"), ("configure_openai_plugin.py", "Configuring OpenAI plugin with mock server"), @@ -109,6 +154,8 @@ def main() -> None: ("create_api_key.py", "Creating API key for the app"), ("publish_workflow.py", "Publishing the workflow"), ] + if setup_step != "finished": + setup_steps.insert(0, ("setup_admin.py", "Creating admin account")) # Create progress logger progress = ProgressLogger(len(setup_steps), log) @@ -131,6 +178,18 @@ def main() -> None: log.error(f"Setup failed at: {failed_step}") log.separator() log.info("Troubleshooting:") + if failed_step == "login_admin.py": + if get_setup_step() == "finished": + log.list_item( + "Dify is already initialized; set STRESS_TEST_ADMIN_EMAIL and " + "STRESS_TEST_ADMIN_PASSWORD to an existing admin account." + ) + else: + admin_config = build_admin_config() + log.list_item( + "Dify is not initialized; setup creates the first admin with " + f"{admin_config['email']} / {admin_config['username']} unless overridden by environment variables." + ) log.list_item("Check if the Dify API server is running (./dev/start-api)") log.list_item("Check if the mock OpenAI server is running (port 5004)") log.list_item("Review the error messages above") @@ -151,9 +210,7 @@ def main() -> None: # Optionally run a test log.separator() - test_input = input("Would you like to run a test workflow now? (yes/no): ") - - if test_input.lower() in ["yes", "y"]: + if confirm("Would you like to run a test workflow now?"): log.step("Running test workflow...") run_script("run_workflow.py", "Testing workflow with default question") diff --git a/scripts/stress-test/test_setup_scripts.py b/scripts/stress-test/test_setup_scripts.py new file mode 100644 index 00000000000..01ddd2b8fb9 --- /dev/null +++ b/scripts/stress-test/test_setup_scripts.py @@ -0,0 +1,120 @@ +#!/usr/bin/env python3 + +import base64 +import importlib.util +import sys +from pathlib import Path + +from common.config_helper import ConfigHelper + + +def _load_setup_module(name: str): + module_path = Path(__file__).parent / "setup" / f"{name}.py" + spec = importlib.util.spec_from_file_location(f"stress_test_{name}", module_path) + assert spec and spec.loader + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def _load_stress_test_module(name: str): + module_path = Path(__file__).parent / f"{name}.py" + spec = importlib.util.spec_from_file_location(f"stress_test_{name}", module_path) + assert spec and spec.loader + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def test_config_helper_getters_read_state_sections(tmp_path): + helper = ConfigHelper(base_dir=tmp_path) + helper.write_state( + { + "auth": {"access_token": "console-token", "csrf_token": "csrf-token"}, + "app": {"app_id": "app-id"}, + "api_key": {"token": "app-token"}, + } + ) + + assert helper.get_token() == "console-token" + assert helper.get_csrf_token() == "csrf-token" + assert helper.get_app_id() == "app-id" + assert helper.get_api_key() == "app-token" + assert helper.console_auth_headers() == { + "authorization": "Bearer console-token", + "X-CSRF-Token": "csrf-token", + } + assert helper.console_auth_cookies() == { + "locale": "en-US", + "access_token": "console-token", + "csrf_token": "csrf-token", + } + + +def test_login_admin_encodes_password_like_web_client(): + login_admin = _load_setup_module("login_admin") + + encoded = login_admin.encode_sensitive_field("password123") + + assert encoded == base64.b64encode(b"password123").decode() + + +def test_setup_admin_reads_credentials_from_environment(monkeypatch): + setup_admin = _load_setup_module("setup_admin") + monkeypatch.setenv("STRESS_TEST_ADMIN_EMAIL", "real-admin@example.com") + monkeypatch.setenv("STRESS_TEST_ADMIN_USERNAME", "real-admin") + monkeypatch.setenv("STRESS_TEST_ADMIN_PASSWORD", "secret") + + assert setup_admin.build_admin_config() == { + "email": "real-admin@example.com", + "username": "real-admin", + "password": "secret", + } + + +def test_setup_all_reads_credentials_from_environment(monkeypatch): + setup_all = _load_stress_test_module("setup_all") + monkeypatch.setenv("STRESS_TEST_ADMIN_EMAIL", "real-admin@example.com") + monkeypatch.setenv("STRESS_TEST_ADMIN_USERNAME", "real-admin") + monkeypatch.setenv("STRESS_TEST_ADMIN_PASSWORD", "secret") + + assert setup_all.build_admin_config() == { + "email": "real-admin@example.com", + "username": "real-admin", + "password": "secret", + } + + +def test_setup_all_confirm_defaults_to_yes(monkeypatch): + setup_all = _load_stress_test_module("setup_all") + monkeypatch.setattr("builtins.input", lambda prompt: "") + + assert setup_all.confirm("Continue?") is True + + +def test_create_api_key_fails_without_app_id(tmp_path): + create_api_key = _load_setup_module("create_api_key") + create_api_key.config_helper.base_dir = tmp_path + create_api_key.config_helper.write_state({"auth": {"access_token": "token"}}) + + assert create_api_key.create_api_key() is False + + +def test_plugin_install_response_without_task_is_non_blocking(): + install_openai_plugin = _load_setup_module("install_openai_plugin") + + assert install_openai_plugin.is_non_blocking_install_response({"ok": True}) is True + assert install_openai_plugin.is_non_blocking_install_response({}) is True + assert install_openai_plugin.is_non_blocking_install_response({"code": "plugin_error"}) is False + + +def test_import_response_with_warnings_and_app_id_is_success(): + import_workflow_app = _load_setup_module("import_workflow_app") + + assert import_workflow_app.is_successful_import_response( + {"status": "completed-with-warnings", "app_id": "app-id"} + ) + assert not import_workflow_app.is_successful_import_response({"status": "failed", "app_id": "app-id"}) + assert not import_workflow_app.is_successful_import_response({"status": "completed"}) From f8afa49ab987aa865830f0eccce84deb289b85c4 Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Tue, 30 Jun 2026 15:54:13 +0800 Subject: [PATCH 07/54] chore(dify-ui): update theme tokens (#38189) --- packages/dify-ui/src/themes/dark.css | 977 +++++++++++++------------ packages/dify-ui/src/themes/light.css | 979 ++++++++++++++------------ packages/dify-ui/src/themes/theme.css | 945 +++++++++++++------------ 3 files changed, 1515 insertions(+), 1386 deletions(-) diff --git a/packages/dify-ui/src/themes/dark.css b/packages/dify-ui/src/themes/dark.css index 3f4a163a725..fb969b5fa01 100644 --- a/packages/dify-ui/src/themes/dark.css +++ b/packages/dify-ui/src/themes/dark.css @@ -1,33 +1,367 @@ /* Attention: Generate by code. Don't update by hand!!! */ html[data-theme="dark"] { + --color-text-primary: #fbfbfc; + --color-text-secondary: #d9d9de; + --color-text-tertiary: rgb(200 206 218 / 0.6); + --color-text-quaternary: rgb(200 206 218 / 0.4); + --color-text-empty-state-icon: rgb(200 206 218 / 0.3); + --color-text-primary-on-surface: rgb(255 255 255 / 0.95); + --color-text-secondary-on-surface: rgb(255 255 255 / 0.9); + --color-text-destructive: #f97066; + --color-text-destructive-secondary: #f97066; + --color-text-success: #17b26a; + --color-text-success-secondary: #47cd89; + --color-text-warning: #f79009; + --color-text-warning-secondary: #fdb022; + --color-text-accent: #6694ff; + --color-text-accent-secondary: #3072ff; + --color-text-accent-light-mode-only: #d9d9de; + --color-text-placeholder: rgb(200 206 218 / 0.3); + --color-text-disabled: rgb(200 206 218 / 0.3); + --color-text-text-selected: rgb(8 90 252 / 0.3); + --color-text-logo-text: #e9e9ec; + --color-text-inverted: #ffffff; + --color-text-inverted-dimmed: rgb(255 255 255 / 0.8); + + --color-background-body: #1d1d20; + --color-background-body-transparent: rgb(29 29 32 / 0); + --color-background-default: #222225; + --color-background-default-hover: #27272b; + --color-background-default-hover-alpha-0: rgb(39 39 43 / 0); + --color-background-default-subtle: #222225; + --color-background-default-dodge: #3a3a40; + --color-background-default-burn: #1d1d20; + --color-background-default-dimmed: #27272b; + --color-background-default-lighter: rgb(200 206 218 / 0.04); + --color-background-section: rgb(24 24 27 / 0.4); + --color-background-section-burn: rgb(24 24 27 / 0.6); + --color-background-section-burn-inverted: #27272b; + --color-background-soft: rgb(24 24 27 / 0.25); + --color-background-neutral-subtle: #1d1d20; + --color-background-surface-white: rgb(255 255 255 / 0.9); + --color-background-sidenav-bg: rgb(39 39 42 / 0.92); + --color-background-overlay-fullscreen: rgb(39 39 42 / 0.97); + --color-background-overlay-backdrop: rgb(24 24 27 / 0.95); + --color-background-overlay: rgb(24 24 27 / 0.8); + --color-background-overlay-alt: rgb(24 24 27 / 0.4); + --color-background-overlay-destructive: rgb(240 68 56 / 0.3); + --color-background-interaction-from-bg-1: rgb(24 24 27 / 0.4); + --color-background-interaction-from-bg-2: rgb(24 24 27 / 0.14); + --color-background-gradient-bg-fill-chat-bg-1: #222225; + --color-background-gradient-bg-fill-chat-bg-2: #1d1d20; + --color-background-gradient-bg-fill-chat-bubble-bg-1: rgb(200 206 218 / 0.08); + --color-background-gradient-bg-fill-chat-bubble-bg-2: rgb(200 206 218 / 0.02); + --color-background-gradient-bg-fill-debug-bg-1: rgb(200 206 218 / 0.08); + --color-background-gradient-bg-fill-debug-bg-2: rgb(24 24 27 / 0.04); + + --color-background-gradient-mask-gray: rgb(24 24 27 / 0.08); + --color-background-gradient-mask-transparent: rgb(0 0 0 / 0); + --color-background-gradient-mask-transparent-dark: rgb(0 0 0 / 0); + --color-background-gradient-mask-side-panel-1: rgb(24 24 27 / 0.04); + --color-background-gradient-mask-side-panel-2: rgb(24 24 27 / 0.9); + --color-background-gradient-mask-input-clear-1: #393a3e; + --color-background-gradient-mask-input-clear-2: rgb(57 58 62 / 0); + + --color-state-base-hover-subtle: rgb(200 206 218 / 0.04); + --color-state-base-hover: rgb(200 206 218 / 0.08); + --color-state-base-hover-alt: rgb(200 206 218 / 0.14); + --color-state-base-active: rgb(200 206 218 / 0.2); + --color-state-base-handle: rgb(200 206 218 / 0.3); + --color-state-base-handle-hover: rgb(200 206 218 / 0.5); + + --color-state-accent-hover: rgb(8 90 252 / 0.14); + --color-state-accent-hover-alt: rgb(8 90 252 / 0.25); + --color-state-accent-active: rgb(8 90 252 / 0.14); + --color-state-accent-active-alt: rgb(8 90 252 / 0.2); + --color-state-accent-solid: #3072ff; + + --color-state-destructive-hover: rgb(240 68 56 / 0.14); + --color-state-destructive-hover-transparent: rgb(240 68 56 / 0); + --color-state-destructive-hover-alt: rgb(240 68 56 / 0.25); + --color-state-destructive-active: rgb(240 68 56 / 0.3); + --color-state-destructive-solid: #f97066; + --color-state-destructive-border: #f97066; + + --color-state-warning-hover: rgb(247 144 9 / 0.14); + --color-state-warning-hover-transparent: rgb(247 144 9 / 0); + --color-state-warning-hover-alt: rgb(247 144 9 / 0.25); + --color-state-warning-active: rgb(247 144 9 / 0.3); + --color-state-warning-solid: #f79009; + + --color-state-success-hover: rgb(23 178 106 / 0.14); + --color-state-success-hover-alt: rgb(23 178 106 / 0.25); + --color-state-success-active: rgb(23 178 106 / 0.3); + --color-state-success-solid: #47cd89; + + --color-workflow-workflow-progress-bg-1: rgb(24 24 27 / 0.25); + --color-workflow-workflow-progress-bg-2: rgb(24 24 27 / 0.04); + + --color-workflow-block-bg: #27272b; + --color-workflow-block-bg-transparent: rgb(39 39 43 / 0.96); + --color-workflow-block-border: rgb(255 255 255 / 0.08); + --color-workflow-block-border-highlight: rgb(200 206 218 / 0.2); + --color-workflow-block-parma-bg: rgb(255 255 255 / 0.05); + --color-workflow-block-wrapper-bg-1: #27272b; + --color-workflow-block-wrapper-bg-2: rgb(39 39 43 / 0.2); + + --color-workflow-link-line-normal: #676f83; + --color-workflow-link-line-normal-transparent: rgb(103 111 131 / 0.2); + --color-workflow-link-line-active: #3072ff; + --color-workflow-link-line-handle: #3072ff; + --color-workflow-link-line-failure-active: #fdb022; + --color-workflow-link-line-failure-handle: #fdb022; + --color-workflow-link-line-failure-button-bg: #f79009; + --color-workflow-link-line-failure-button-hover: #dc6803; + + --color-workflow-link-line-success-active: #47cd89; + --color-workflow-link-line-success-handle: #47cd89; + + --color-workflow-link-line-error-active: #f97066; + --color-workflow-link-line-error-handle: #f97066; + + --color-workflow-minimap-block: rgb(200 206 218 / 0.08); + --color-workflow-minimap-bg: #27272b; + + --color-workflow-display-glass-1: rgb(255 255 255 / 0.03); + --color-workflow-display-glass-2: rgb(255 255 255 / 0.05); + --color-workflow-display-highlight: rgb(255 255 255 / 0.12); + --color-workflow-display-outline: rgb(24 24 27 / 0.95); + --color-workflow-display-vignette-dark: rgb(0 0 0 / 0.4); + --color-workflow-display-success-bg: rgb(23 178 106 / 0.2); + --color-workflow-display-success-bg-line-pattern: rgb(24 24 27 / 0.8); + --color-workflow-display-success-border-1: rgb(23 178 106 / 0.9); + --color-workflow-display-success-border-2: rgb(23 178 106 / 0.8); + --color-workflow-display-success-vignette-color: rgb(23 178 106 / 0.25); + + --color-workflow-display-error-bg: rgb(240 68 56 / 0.2); + --color-workflow-display-error-bg-line-pattern: rgb(24 24 27 / 0.8); + --color-workflow-display-error-border-1: rgb(240 68 56 / 0.9); + --color-workflow-display-error-border-2: rgb(240 68 56 / 0.8); + --color-workflow-display-error-vignette-color: rgb(240 68 56 / 0.25); + + --color-workflow-display-warning-bg: rgb(247 144 9 / 0.2); + --color-workflow-display-warning-bg-line-pattern: rgb(24 24 27 / 0.8); + --color-workflow-display-warning-border-1: rgb(247 144 9 / 0.9); + --color-workflow-display-warning-border-2: rgb(247 144 9 / 0.8); + --color-workflow-display-warning-vignette-color: rgb(247 144 9 / 0.25); + + --color-workflow-display-normal-bg: rgb(11 165 236 / 0.2); + --color-workflow-display-normal-bg-line-pattern: rgb(24 24 27 / 0.8); + --color-workflow-display-normal-border-1: rgb(11 165 236 / 0.9); + --color-workflow-display-normal-border-2: rgb(11 165 236 / 0.8); + --color-workflow-display-normal-vignette-color: rgb(11 165 236 / 0.25); + + --color-workflow-display-disabled-bg: rgb(200 206 218 / 0.2); + --color-workflow-display-disabled-bg-line-pattern: rgb(24 24 27 / 0.8); + --color-workflow-display-disabled-border-1: rgb(200 206 218 / 0.6); + --color-workflow-display-disabled-border-2: rgb(200 206 218 / 0.25); + --color-workflow-display-disabled-vignette-color: rgb(200 206 218 / 0.25); + --color-workflow-display-disabled-outline: rgb(24 24 27 / 0.95); + + --color-workflow-canvas-workflow-dot-color: rgb(133 133 173 / 0.11); + --color-workflow-canvas-workflow-bg: #1d1d20; + --color-workflow-canvas-workflow-top-bar-1: rgb(29 29 32 / 0.9); + --color-workflow-canvas-workflow-top-bar-2: rgb(29 29 32 / 0.08); + --color-workflow-canvas-canvas-overlay: rgb(29 29 32 / 0.8); + + --color-workflow-debug-run-status-bg: rgb(230 46 5 / 0.4); + --color-workflow-debug-run-status-bg-alt: rgb(255 46 0 / 0.5); + --color-workflow-debug-breakpoint: #ff692e; + --color-workflow-debug-text: #ff9c66; + --color-workflow-debug-text-disabled: rgb(255 68 5 / 0.2); + + --color-workflow-test-run-run-status-bg: rgb(8 90 252 / 0.5); + --color-workflow-test-run-paused-bg: rgb(247 144 9 / 0.3); + --color-workflow-test-run-paused-text: #fdb022; + --color-workflow-test-run-run-status-bg-alt: rgb(45 90 190 / 0.9); + --color-workflow-test-run-text: #d0dfff; + + --color-components-icon-bg-red-solid: #d92d20; + --color-components-icon-bg-rose-solid: #e31b54; + --color-components-icon-bg-pink-solid: #dd2590; + --color-components-icon-bg-orange-dark-solid: #ff4405; + --color-components-icon-bg-orange-solid: #f79009; + --color-components-icon-bg-yellow-solid: #eaaa08; + --color-components-icon-bg-green-solid: #4ca30d; + --color-components-icon-bg-teal-solid: #0e9384; + --color-components-icon-bg-blue-light-solid: #0ba5ec; + --color-components-icon-bg-blue-solid: #085afc; + --color-components-icon-bg-indigo-solid: #444ce7; + --color-components-icon-bg-violet-solid: #7839ee; + --color-components-icon-bg-midnight-solid: #5d698d; + --color-components-icon-bg-red-soft: rgb(240 68 56 / 0.2); + --color-components-icon-bg-rose-soft: rgb(246 61 104 / 0.2); + --color-components-icon-bg-pink-soft: rgb(238 70 188 / 0.2); + --color-components-icon-bg-orange-dark-soft: rgb(255 68 5 / 0.2); + --color-components-icon-bg-orange-soft: rgb(247 144 9 / 0.2); + --color-components-icon-bg-yellow-soft: rgb(234 170 8 / 0.2); + --color-components-icon-bg-green-soft: rgb(102 198 28 / 0.2); + --color-components-icon-bg-teal-soft: rgb(21 183 158 / 0.2); + --color-components-icon-bg-blue-light-soft: rgb(11 165 236 / 0.2); + --color-components-icon-bg-blue-soft: rgb(0 51 255 / 0.2); + --color-components-icon-bg-indigo-soft: rgb(97 114 243 / 0.2); + --color-components-icon-bg-violet-soft: rgb(135 91 247 / 0.2); + --color-components-icon-bg-midnight-soft: rgb(130 141 173 / 0.2); + + --color-components-avatar-default-avatar-bg: #222225; + --color-components-avatar-mask-darkmode-dimmed: rgb(0 0 0 / 0.12); + --color-components-avatar-shape-fill-stop-0: rgb(255 255 255 / 0.95); + --color-components-avatar-shape-fill-stop-100: rgb(255 255 255 / 0.8); + + --color-components-avatar-bg-mask-stop-0: rgb(255 255 255 / 0.2); + --color-components-avatar-bg-mask-stop-100: rgb(255 255 255 / 0.03); + + --color-components-chat-input-audio-bg: rgb(0 51 255 / 0.2); + --color-components-chat-input-audio-bg-alt: rgb(24 24 27 / 0.9); + --color-components-chat-input-audio-wave-default: rgb(200 206 218 / 0.14); + --color-components-chat-input-audio-wave-active: #6694ff; + --color-components-chat-input-bg-mask-1: rgb(24 24 27 / 0.04); + --color-components-chat-input-bg-mask-2: rgb(24 24 27 / 0.6); + --color-components-chat-input-border: rgb(200 206 218 / 0.2); + + --color-components-label-gray: rgb(200 206 218 / 0.14); + + --color-components-premium-badge-highlight-stop-0: rgb(255 255 255 / 0.12); + --color-components-premium-badge-highlight-stop-100: rgb(255 255 255 / 0.2); + --color-components-premium-badge-orange-bg-stop-0: #ff692e; + --color-components-premium-badge-orange-bg-stop-100: #e04f16; + --color-components-premium-badge-orange-stroke-stop-0: rgb(255 255 255 / 0.2); + --color-components-premium-badge-orange-stroke-stop-100: #ff4405; + --color-components-premium-badge-orange-text-stop-0: #fef6ee; + --color-components-premium-badge-orange-text-stop-100: #f9dbaf; + --color-components-premium-badge-orange-glow: #b93815; + --color-components-premium-badge-orange-glow-hover: #fdead7; + --color-components-premium-badge-orange-bg-stop-0-hover: #ff692e; + --color-components-premium-badge-orange-bg-stop-100-hover: #b93815; + --color-components-premium-badge-orange-stroke-stop-0-hover: rgb(255 255 255 / 0.5); + --color-components-premium-badge-orange-stroke-stop-100-hover: #ff4405; + + --color-components-premium-badge-blue-bg-stop-0: #3072ff; + --color-components-premium-badge-blue-bg-stop-100: #085afc; + --color-components-premium-badge-blue-stroke-stop-0: rgb(255 255 255 / 0.2); + --color-components-premium-badge-blue-stroke-stop-100: #085afc; + --color-components-premium-badge-blue-text-stop-0: #e9f0ff; + --color-components-premium-badge-blue-text-stop-100: #a0bdff; + --color-components-premium-badge-blue-glow: #002cde; + --color-components-premium-badge-blue-glow-hover: #d0dfff; + --color-components-premium-badge-blue-bg-stop-0-hover: #6694ff; + --color-components-premium-badge-blue-bg-stop-100-hover: #002cde; + --color-components-premium-badge-blue-stroke-stop-0-hover: rgb(255 255 255 / 0.5); + --color-components-premium-badge-blue-stroke-stop-100-hover: #085afc; + + --color-components-premium-badge-indigo-bg-stop-0: #6172f3; + --color-components-premium-badge-indigo-bg-stop-100: #3538cd; + --color-components-premium-badge-indigo-stroke-stop-0: rgb(255 255 255 / 0.2); + --color-components-premium-badge-indigo-stroke-stop-100: #444ce7; + --color-components-premium-badge-indigo-text-stop-0: #eef4ff; + --color-components-premium-badge-indigo-text-stop-100: #c7d7fe; + --color-components-premium-badge-indigo-glow: #3538cd; + --color-components-premium-badge-indigo-glow-hover: #e0eaff; + --color-components-premium-badge-indigo-bg-stop-0-hover: #a4bcfd; + --color-components-premium-badge-indigo-bg-stop-100-hover: #3538cd; + --color-components-premium-badge-indigo-stroke-stop-0-hover: rgb(255 255 255 / 0.5); + --color-components-premium-badge-indigo-stroke-stop-100-hover: #444ce7; + + --color-components-premium-badge-grey-bg-stop-0: #676f83; + --color-components-premium-badge-grey-bg-stop-100: #495464; + --color-components-premium-badge-grey-stroke-stop-0: rgb(255 255 255 / 0.12); + --color-components-premium-badge-grey-stroke-stop-100: #495464; + --color-components-premium-badge-grey-text-stop-0: #f9fafb; + --color-components-premium-badge-grey-text-stop-100: #e9ebf0; + --color-components-premium-badge-grey-glow: #354052; + --color-components-premium-badge-grey-glow-hover: #f2f4f7; + --color-components-premium-badge-grey-bg-stop-0-hover: #98a2b2; + --color-components-premium-badge-grey-bg-stop-100-hover: #354052; + --color-components-premium-badge-grey-stroke-stop-0-hover: rgb(255 255 255 / 0.5); + --color-components-premium-badge-grey-stroke-stop-100-hover: #676f83; + + --color-components-dropzone-bg: rgb(24 24 27 / 0.4); + --color-components-dropzone-bg-alt: rgb(24 24 27 / 0.8); + --color-components-dropzone-bg-accent: rgb(0 51 255 / 0.2); + --color-components-dropzone-border: rgb(200 206 218 / 0.14); + --color-components-dropzone-border-alt: rgb(200 206 218 / 0.2); + --color-components-dropzone-border-accent: #6694ff; + + --color-components-panel-bg: #222225; + --color-components-panel-bg-transparent: rgb(34 34 37 / 0); + --color-components-panel-bg-alt: #222225; + --color-components-panel-bg-blur: rgb(44 44 48 / 0.95); + --color-components-panel-bg-blur-burn: rgb(31 31 35 / 0.9); + --color-components-panel-border: rgb(200 206 218 / 0.14); + --color-components-panel-border-subtle: rgb(200 206 218 / 0.08); + --color-components-panel-gradient-1: #27272b; + --color-components-panel-gradient-2: #222225; + --color-components-panel-on-panel-item-bg: #27272b; + --color-components-panel-on-panel-item-bg-transparent: rgb(44 44 48 / 0.95); + --color-components-panel-on-panel-item-bg-hover: #3a3a40; + --color-components-panel-on-panel-item-bg-hover-transparent: rgb(58 58 64 / 0); + --color-components-panel-on-panel-item-bg-destructive-hover-transparent: rgb(255 251 250 / 0); + --color-components-panel-on-panel-item-bg-alt: #3a3a40; + + --color-components-marketplace-header-bg: rgb(31 31 35 / 0.9); + + --color-components-card-bg: #222225; + --color-components-card-bg-alt: #27272b; + --color-components-card-bg-transparent: rgb(34 34 37 / 0); + --color-components-card-bg-alt-transparent: rgb(39 39 43 / 0); + --color-components-card-border: rgb(255 255 255 / 0.03); + + --color-components-actionbar-border: rgb(200 206 218 / 0.08); + --color-components-actionbar-bg: #222225; + --color-components-actionbar-bg-accent: #27272b; + --color-components-actionbar-border-accent: #3072ff; + --color-components-input-bg-normal: rgb(255 255 255 / 0.08); - --color-components-input-text-placeholder: rgb(200 206 218 / 0.3); --color-components-input-bg-hover: rgb(255 255 255 / 0.03); --color-components-input-bg-active: rgb(255 255 255 / 0.05); - --color-components-input-border-active: #747481; - --color-components-input-border-destructive: #f97066; - --color-components-input-text-filled: #f4f4f5; - --color-components-input-bg-destructive: rgb(255 255 255 / 0.01); --color-components-input-bg-disabled: rgb(255 255 255 / 0.03); + --color-components-input-bg-destructive: rgb(255 255 255 / 0.01); + --color-components-input-text-filled: #f4f4f5; + --color-components-input-text-placeholder: rgb(200 206 218 / 0.3); --color-components-input-text-disabled: rgb(200 206 218 / 0.3); --color-components-input-text-filled-disabled: rgb(200 206 218 / 0.6); + --color-components-input-text-filled-blue: #6694ff; --color-components-input-border-hover: #3a3a40; + --color-components-input-border-active: #747481; --color-components-input-border-active-prompt-1: #36bffa; - --color-components-input-border-active-prompt-2: #296dff; + --color-components-input-border-active-prompt-2: #085afc; + --color-components-input-border-destructive: #f97066; - --color-components-kbd-bg-gray: rgb(255 255 255 / 0.03); - --color-components-kbd-bg-white: rgb(255 255 255 / 0.12); + --color-components-progress-brand-progress: #3072ff; + --color-components-progress-brand-border: #3072ff; + --color-components-progress-brand-bg: rgb(8 90 252 / 0.04); - --color-components-tooltip-bg: rgb(24 24 27 / 0.95); + --color-components-progress-white-progress: #ffffff; + --color-components-progress-white-border: rgb(255 255 255 / 0.95); + --color-components-progress-white-bg: rgb(255 255 255 / 0.01); + --color-components-progress-gray-progress: #98a2b2; + --color-components-progress-gray-border: #98a2b2; + --color-components-progress-gray-bg: rgb(200 206 218 / 0.02); + + --color-components-progress-warning-progress: #fdb022; + --color-components-progress-warning-border: #fdb022; + --color-components-progress-warning-bg: rgb(247 144 9 / 0.04); + + --color-components-progress-error-progress: #f97066; + --color-components-progress-error-border: #f97066; + --color-components-progress-error-bg: rgb(240 68 56 / 0.04); + + --color-components-progress-bar-progress: rgb(200 206 218 / 0.14); + --color-components-progress-bar-progress-highlight: rgb(200 206 218 / 0.2); + --color-components-progress-bar-progress-solid: rgb(255 255 255 / 0.95); + --color-components-progress-bar-border: rgb(255 255 255 / 0.03); + --color-components-progress-bar-bg: rgb(200 206 218 / 0.08); + + --color-components-button-button-seam: rgb(0 0 0 / 0.15); --color-components-button-primary-text: rgb(255 255 255 / 0.95); - --color-components-button-primary-bg: #155aef; - --color-components-button-primary-border: rgb(255 255 255 / 0.12); - --color-components-button-primary-bg-hover: #296dff; - --color-components-button-primary-border-hover: rgb(255 255 255 / 0.2); - --color-components-button-primary-bg-disabled: rgb(255 255 255 / 0.03); - --color-components-button-primary-border-disabled: rgb(255 255 255 / 0.08); --color-components-button-primary-text-disabled: rgb(255 255 255 / 0.2); + --color-components-button-primary-bg: #085afc; + --color-components-button-primary-bg-hover: #3072ff; + --color-components-button-primary-bg-disabled: rgb(255 255 255 / 0.03); + --color-components-button-primary-border: rgb(255 255 255 / 0.12); + --color-components-button-primary-border-hover: rgb(255 255 255 / 0.2); + --color-components-button-primary-border-disabled: rgb(255 255 255 / 0.08); --color-components-button-secondary-text: rgb(255 255 255 / 0.8); --color-components-button-secondary-text-disabled: rgb(255 255 255 / 0.2); @@ -38,6 +372,15 @@ html[data-theme="dark"] { --color-components-button-secondary-border-hover: rgb(255 255 255 / 0.12); --color-components-button-secondary-border-disabled: rgb(255 255 255 / 0.05); + --color-components-button-secondary-accent-text: rgb(255 255 255 / 0.8); + --color-components-button-secondary-accent-text-disabled: rgb(255 255 255 / 0.2); + --color-components-button-secondary-accent-bg: rgb(255 255 255 / 0.05); + --color-components-button-secondary-accent-bg-hover: rgb(255 255 255 / 0.08); + --color-components-button-secondary-accent-bg-disabled: rgb(255 255 255 / 0.03); + --color-components-button-secondary-accent-border: rgb(255 255 255 / 0.08); + --color-components-button-secondary-accent-border-hover: rgb(255 255 255 / 0.12); + --color-components-button-secondary-accent-border-disabled: rgb(255 255 255 / 0.05); + --color-components-button-tertiary-text: #d9d9de; --color-components-button-tertiary-text-disabled: rgb(255 255 255 / 0.2); --color-components-button-tertiary-bg: rgb(255 255 255 / 0.08); @@ -76,133 +419,99 @@ html[data-theme="dark"] { --color-components-button-destructive-ghost-text-disabled: rgb(240 68 56 / 0.2); --color-components-button-destructive-ghost-bg-hover: rgb(240 68 56 / 0.14); - --color-components-button-secondary-accent-text: rgb(255 255 255 / 0.8); - --color-components-button-secondary-accent-text-disabled: rgb(255 255 255 / 0.2); - --color-components-button-secondary-accent-bg: rgb(255 255 255 / 0.05); - --color-components-button-secondary-accent-bg-hover: rgb(255 255 255 / 0.08); - --color-components-button-secondary-accent-bg-disabled: rgb(255 255 255 / 0.03); - --color-components-button-secondary-accent-border: rgb(255 255 255 / 0.08); - --color-components-button-secondary-accent-border-hover: rgb(255 255 255 / 0.12); - --color-components-button-secondary-accent-border-disabled: rgb(255 255 255 / 0.05); - --color-components-button-indigo-bg: #444ce7; --color-components-button-indigo-bg-hover: #6172f3; --color-components-button-indigo-bg-disabled: rgb(255 255 255 / 0.03); + --color-components-button-debug-text: rgb(255 255 255 / 0.95); + --color-components-button-debug-text-disabled: rgb(255 255 255 / 0.2); + --color-components-button-debug-bg: #ff4405; + --color-components-button-debug-bg-hover: #ff692e; + --color-components-button-debug-bg-disabled: rgb(255 68 5 / 0.08); + --color-components-button-debug-border: rgb(255 255 255 / 0.12); + --color-components-button-debug-border-hover: rgb(255 255 255 / 0.2); + --color-components-button-debug-border-disabled: rgb(255 255 255 / 0.08); + --color-components-checkbox-icon: rgb(255 255 255 / 0.95); --color-components-checkbox-icon-disabled: rgb(255 255 255 / 0.2); - --color-components-checkbox-bg: #296dff; - --color-components-checkbox-bg-hover: #5289ff; + --color-components-checkbox-bg: #085afc; + --color-components-checkbox-bg-hover: #3072ff; --color-components-checkbox-bg-disabled: rgb(255 255 255 / 0.03); + --color-components-checkbox-bg-disabled-checked: rgb(8 90 252 / 0.2); + --color-components-checkbox-bg-unchecked: rgb(255 255 255 / 0.03); + --color-components-checkbox-bg-unchecked-hover: rgb(255 255 255 / 0.05); --color-components-checkbox-border: rgb(255 255 255 / 0.4); --color-components-checkbox-border-hover: rgb(255 255 255 / 0.6); --color-components-checkbox-border-disabled: rgb(255 255 255 / 0.01); - --color-components-checkbox-bg-unchecked: rgb(255 255 255 / 0.03); - --color-components-checkbox-bg-unchecked-hover: rgb(255 255 255 / 0.05); - --color-components-checkbox-bg-disabled-checked: rgb(21 90 239 / 0.2); - --color-components-radio-border-checked: #296dff; - --color-components-radio-border-checked-hover: #5289ff; - --color-components-radio-border-checked-disabled: rgb(21 90 239 / 0.2); + --color-components-radio-bg: rgb(255 255 255 / 0); + --color-components-radio-bg-hover: rgb(255 255 255 / 0.05); --color-components-radio-bg-disabled: rgb(255 255 255 / 0.03); + --color-components-radio-border-checked: #085afc; + --color-components-radio-border-checked-hover: #3072ff; + --color-components-radio-border-checked-disabled: rgb(8 90 252 / 0.2); --color-components-radio-border: rgb(255 255 255 / 0.4); --color-components-radio-border-hover: rgb(255 255 255 / 0.6); --color-components-radio-border-disabled: rgb(255 255 255 / 0.01); - --color-components-radio-bg: rgb(255 255 255 / 0); - --color-components-radio-bg-hover: rgb(255 255 255 / 0.05); --color-components-toggle-knob: #f4f4f5; + --color-components-toggle-knob-hover: #fefefe; --color-components-toggle-knob-disabled: rgb(255 255 255 / 0.2); - --color-components-toggle-bg: #296dff; - --color-components-toggle-bg-hover: #5289ff; + --color-components-toggle-bg: #085afc; + --color-components-toggle-bg-hover: #3072ff; --color-components-toggle-bg-disabled: rgb(255 255 255 / 0.08); --color-components-toggle-bg-unchecked: rgb(255 255 255 / 0.2); --color-components-toggle-bg-unchecked-hover: rgb(255 255 255 / 0.3); --color-components-toggle-bg-unchecked-disabled: rgb(255 255 255 / 0.08); - --color-components-toggle-knob-hover: #fefefe; - - --color-components-card-bg: #222225; - --color-components-card-border: rgb(255 255 255 / 0.03); - --color-components-card-bg-alt: #27272b; - --color-components-card-bg-transparent: rgb(34 34 37 / 0); - --color-components-card-bg-alt-transparent: rgb(39 39 43 / 0); - - --color-components-menu-item-text: rgb(200 206 218 / 0.6); - --color-components-menu-item-text-active: rgb(255 255 255 / 0.95); - --color-components-menu-item-text-hover: rgb(200 206 218 / 0.8); - --color-components-menu-item-text-active-accent: rgb(255 255 255 / 0.95); - --color-components-menu-item-bg-active: rgb(200 206 218 / 0.14); - --color-components-menu-item-bg-hover: rgb(200 206 218 / 0.08); - - --color-components-panel-bg: #222225; - --color-components-panel-bg-blur: rgb(44 44 48 / 0.95); - --color-components-panel-border: rgb(200 206 218 / 0.14); - --color-components-panel-border-subtle: rgb(200 206 218 / 0.08); - --color-components-panel-gradient-2: #222225; - --color-components-panel-gradient-1: #27272b; - --color-components-panel-bg-alt: #222225; - --color-components-panel-on-panel-item-bg: #27272b; - --color-components-panel-on-panel-item-bg-hover: #3a3a40; - --color-components-panel-on-panel-item-bg-alt: #3a3a40; - --color-components-panel-on-panel-item-bg-transparent: rgb(44 44 48 / 0.95); - --color-components-panel-on-panel-item-bg-hover-transparent: rgb(58 58 64 / 0); - --color-components-panel-on-panel-item-bg-destructive-hover-transparent: rgb(255 251 250 / 0); - - --color-components-panel-bg-transparent: rgb(34 34 37 / 0); - - --color-components-main-nav-nav-button-text: rgb(200 206 218 / 0.6); - --color-components-main-nav-nav-button-text-active: #f4f4f5; - --color-components-main-nav-nav-button-bg: rgb(255 255 255 / 0); - --color-components-main-nav-nav-button-bg-active: rgb(200 206 218 / 0.14); - --color-components-main-nav-nav-button-border: rgb(255 255 255 / 0.08); - --color-components-main-nav-nav-button-bg-hover: rgb(200 206 218 / 0.04); - --color-components-main-nav-glass-text-glow: #3146ff2e; - --color-components-main-nav-glass-surface-first: #0033ff14; - --color-components-main-nav-glass-surface-middle-1: #0033ff1f; - --color-components-main-nav-glass-surface-middle-2: #0033ff1a; - --color-components-main-nav-glass-surface-end: #0033ff14; - --color-components-main-nav-glass-edge-highlight-first: #fffffffa; - --color-components-main-nav-glass-edge-highlight-middle: #ffffff00; - --color-components-main-nav-glass-edge-highlight-end: #ffffff6b; - --color-components-main-nav-glass-edge-reflection-first: #0033ff00; - --color-components-main-nav-glass-edge-reflection-middle: #0033ff99; - --color-components-main-nav-glass-edge-reflection-end: #0033ff00; - --color-components-main-nav-glass-inner-glow: #ffffff4d; - --color-components-main-nav-glass-shadow-reflection: #0033ff0a; - --color-components-main-nav-glass-shadow-reflection-glow: #ffffff00; - - --color-components-main-nav-nav-user-border: rgb(255 255 255 / 0.05); - - --color-components-slider-knob: #f4f4f5; - --color-components-slider-knob-hover: #fefefe; - --color-components-slider-knob-disabled: rgb(255 255 255 / 0.2); - --color-components-slider-range: #296dff; - --color-components-slider-track: rgb(255 255 255 / 0.2); - --color-components-slider-knob-border-hover: rgb(16 24 40 / 0.3); - --color-components-slider-knob-border: rgb(16 24 40 / 0.2); - - --color-components-segmented-control-item-active-bg: rgb(255 255 255 / 0.08); - --color-components-segmented-control-item-active-border: rgb(200 206 218 / 0.08); - --color-components-segmented-control-bg-normal: rgb(24 24 27 / 0.7); - --color-components-segmented-control-item-active-accent-bg: rgb(21 90 239 / 0.2); - --color-components-segmented-control-item-active-accent-border: rgb(21 90 239 / 0.3); --color-components-option-card-option-bg: rgb(200 206 218 / 0.04); - --color-components-option-card-option-selected-bg: rgb(255 255 255 / 0.05); - --color-components-option-card-option-selected-border: #5289ff; --color-components-option-card-option-border: rgb(200 206 218 / 0.2); --color-components-option-card-option-bg-hover: rgb(200 206 218 / 0.14); --color-components-option-card-option-border-hover: rgb(200 206 218 / 0.3); + --color-components-option-card-option-selected-bg: rgb(255 255 255 / 0.05); + --color-components-option-card-option-selected-border: #3072ff; - --color-components-tab-active: #296dff; + --color-components-slider-knob: #f4f4f5; + --color-components-slider-knob-border: rgb(16 24 40 / 0.2); + --color-components-slider-knob-border-hover: rgb(16 24 40 / 0.3); + --color-components-slider-knob-hover: #fefefe; + --color-components-slider-knob-disabled: rgb(255 255 255 / 0.2); + --color-components-slider-range: #085afc; + --color-components-slider-track: rgb(255 255 255 / 0.2); + + --color-components-tooltip-bg: rgb(24 24 27 / 0.95); + + --color-components-kbd-bg-gray: rgb(255 255 255 / 0.03); + --color-components-kbd-bg-white: rgb(255 255 255 / 0.12); + + --color-components-menu-item-text: rgb(200 206 218 / 0.6); + --color-components-menu-item-text-hover: rgb(200 206 218 / 0.8); + --color-components-menu-item-text-active: rgb(255 255 255 / 0.95); + --color-components-menu-item-text-active-accent: rgb(255 255 255 / 0.95); + --color-components-menu-item-bg-hover: rgb(200 206 218 / 0.08); + --color-components-menu-item-bg-active: rgb(200 206 218 / 0.14); + + --color-components-segmented-control-bg-normal: rgb(24 24 27 / 0.7); + --color-components-segmented-control-item-active-bg: rgb(255 255 255 / 0.08); + --color-components-segmented-control-item-active-border: rgb(200 206 218 / 0.08); + --color-components-segmented-control-item-active-accent-bg: rgb(8 90 252 / 0.2); + --color-components-segmented-control-item-active-accent-border: rgb(8 90 252 / 0.3); --color-components-badge-white-to-dark: rgb(24 24 27 / 0.8); + --color-components-badge-white-to-dark-alpha: rgb(24 24 27 / 0); + --color-components-badge-bg-green-soft: rgb(23 178 106 / 0.14); + --color-components-badge-bg-orange-soft: rgb(247 144 9 / 0.14); + --color-components-badge-bg-red-soft: rgb(240 68 56 / 0.14); + --color-components-badge-bg-blue-light-soft: rgb(11 165 236 / 0.14); + --color-components-badge-bg-gray-soft: rgb(200 206 218 / 0.08); + --color-components-badge-bg-dimm: rgb(255 255 255 / 0.03); + + --color-components-badge-status-light-border-outer: #222225; + --color-components-badge-status-light-high-light: rgb(255 255 255 / 0.3); --color-components-badge-status-light-success-bg: #17b26a; --color-components-badge-status-light-success-border-inner: #47cd89; --color-components-badge-status-light-success-halo: rgb(23 178 106 / 0.3); - --color-components-badge-status-light-border-outer: #222225; - --color-components-badge-status-light-high-light: rgb(255 255 255 / 0.3); --color-components-badge-status-light-warning-bg: #f79009; --color-components-badge-status-light-warning-border-inner: #fdb022; --color-components-badge-status-light-warning-halo: rgb(247 144 9 / 0.3); @@ -219,221 +528,60 @@ html[data-theme="dark"] { --color-components-badge-status-light-disabled-border-inner: #98a2b2; --color-components-badge-status-light-disabled-halo: rgb(200 206 218 / 0.08); - --color-components-badge-bg-green-soft: rgb(23 178 106 / 0.14); - --color-components-badge-bg-orange-soft: rgb(247 144 9 / 0.14); - --color-components-badge-bg-red-soft: rgb(240 68 56 / 0.14); - --color-components-badge-bg-blue-light-soft: rgb(11 165 236 / 0.14); - --color-components-badge-bg-gray-soft: rgb(200 206 218 / 0.08); - --color-components-badge-bg-dimm: rgb(255 255 255 / 0.03); + --color-components-tab-active: #3072ff; + + --color-components-main-nav-text: #a8a8b3; + --color-components-main-nav-text-active: #ffffff; + --color-components-main-nav-nav-button-border: rgb(255 255 255 / 0.08); + --color-components-main-nav-nav-button-text: rgb(200 206 218 / 0.6); + --color-components-main-nav-nav-button-text-active: #f4f4f5; + --color-components-main-nav-nav-button-bg: rgb(255 255 255 / 0); + --color-components-main-nav-nav-button-bg-hover: rgb(200 206 218 / 0.04); + --color-components-main-nav-nav-button-bg-active: rgb(200 206 218 / 0.14); + + --color-components-main-nav-nav-user-border: rgb(255 255 255 / 0.05); + + --color-components-main-nav-glass-inner-glow: rgb(210 219 255 / 0.05); + --color-components-main-nav-glass-shadow-reflection: rgb(210 219 255 / 0.04); + --color-components-main-nav-glass-shadow-reflection-glow: rgb(255 255 255 / 0.02); + --color-components-main-nav-glass-text-glow: rgb(245 246 255 / 0.27); + --color-components-main-nav-glass-surface-first: rgb(196 207 255 / 0.08); + --color-components-main-nav-glass-surface-middle-1: rgb(210 219 255 / 0.12); + --color-components-main-nav-glass-surface-middle-2: rgb(210 219 255 / 0.1); + --color-components-main-nav-glass-surface-end: rgb(196 207 255 / 0.08); + + --color-components-main-nav-glass-edge-reflection-first: rgb(92 124 255 / 0); + --color-components-main-nav-glass-edge-reflection-middle: rgb(210 219 255 / 0.8); + --color-components-main-nav-glass-edge-reflection-end: rgb(92 124 255 / 0); + + --color-components-main-nav-glass-edge-highlight-first: rgb(196 207 255 / 0.15); + --color-components-main-nav-glass-edge-highlight-middle: rgb(72 108 255 / 0); + --color-components-main-nav-glass-edge-highlight-end: rgb(196 207 255 / 0.05); - --color-components-chart-line: #5289ff; - --color-components-chart-area-1: rgb(21 90 239 / 0.2); - --color-components-chart-area-2: rgb(21 90 239 / 0.04); - --color-components-chart-current-1: #5289ff; - --color-components-chart-current-2: rgb(21 90 239 / 0.3); --color-components-chart-bg: rgb(24 24 27 / 0.95); + --color-components-chart-line: #3072ff; + --color-components-chart-current-1: #3072ff; + --color-components-chart-current-2: rgb(8 90 252 / 0.3); + --color-components-chart-area-1: rgb(8 90 252 / 0.2); + --color-components-chart-area-2: rgb(8 90 252 / 0.04); - --color-components-actionbar-bg: #222225; - --color-components-actionbar-border: rgb(200 206 218 / 0.08); - --color-components-actionbar-bg-accent: #27272b; - --color-components-actionbar-border-accent: #5289ff; + --color-divider-subtle: rgb(200 206 218 / 0.08); + --color-divider-regular: rgb(200 206 218 / 0.14); + --color-divider-deep: rgb(200 206 218 / 0.2); + --color-divider-intense: rgb(200 206 218 / 0.4); + --color-divider-burn: rgb(24 24 27 / 0.95); + --color-divider-solid: #3a3a40; + --color-divider-solid-alt: #747481; + --color-divider-accent: rgb(200 206 218 / 0.14); - --color-components-dropzone-bg-alt: rgb(24 24 27 / 0.8); - --color-components-dropzone-bg: rgb(24 24 27 / 0.4); - --color-components-dropzone-bg-accent: rgb(21 90 239 / 0.2); - --color-components-dropzone-border: rgb(200 206 218 / 0.14); - --color-components-dropzone-border-alt: rgb(200 206 218 / 0.2); - --color-components-dropzone-border-accent: #84abff; - - --color-components-progress-brand-progress: #5289ff; - --color-components-progress-brand-border: #5289ff; - --color-components-progress-brand-bg: rgb(21 90 239 / 0.04); - - --color-components-progress-white-progress: #ffffff; - --color-components-progress-white-border: rgb(255 255 255 / 0.95); - --color-components-progress-white-bg: rgb(255 255 255 / 0.01); - - --color-components-progress-gray-progress: #98a2b2; - --color-components-progress-gray-border: #98a2b2; - --color-components-progress-gray-bg: rgb(200 206 218 / 0.02); - - --color-components-progress-warning-progress: #fdb022; - --color-components-progress-warning-border: #fdb022; - --color-components-progress-warning-bg: rgb(247 144 9 / 0.04); - - --color-components-progress-error-progress: #f97066; - --color-components-progress-error-border: #f97066; - --color-components-progress-error-bg: rgb(240 68 56 / 0.04); - - --color-components-chat-input-audio-bg: rgb(21 90 239 / 0.2); - --color-components-chat-input-audio-wave-default: rgb(200 206 218 / 0.14); - --color-components-chat-input-bg-mask-1: rgb(24 24 27 / 0.04); - --color-components-chat-input-bg-mask-2: rgb(24 24 27 / 0.6); - --color-components-chat-input-border: rgb(200 206 218 / 0.2); - --color-components-chat-input-audio-wave-active: #84abff; - --color-components-chat-input-audio-bg-alt: rgb(24 24 27 / 0.9); - - --color-components-avatar-shape-fill-stop-0: rgb(255 255 255 / 0.95); - --color-components-avatar-shape-fill-stop-100: rgb(255 255 255 / 0.8); - - --color-components-avatar-bg-mask-stop-0: rgb(255 255 255 / 0.2); - --color-components-avatar-bg-mask-stop-100: rgb(255 255 255 / 0.03); - - --color-components-avatar-default-avatar-bg: #222225; - --color-components-avatar-mask-darkmode-dimmed: rgb(0 0 0 / 0.12); - - --color-components-label-gray: rgb(200 206 218 / 0.14); - - --color-components-premium-badge-blue-bg-stop-0: #5289ff; - --color-components-premium-badge-blue-bg-stop-100: #296dff; - --color-components-premium-badge-blue-stroke-stop-0: rgb(255 255 255 / 0.2); - --color-components-premium-badge-blue-stroke-stop-100: #296dff; - --color-components-premium-badge-blue-text-stop-0: #eff4ff; - --color-components-premium-badge-blue-text-stop-100: #b2caff; - --color-components-premium-badge-blue-glow: #004aeb; - --color-components-premium-badge-blue-bg-stop-0-hover: #84abff; - --color-components-premium-badge-blue-bg-stop-100-hover: #004aeb; - --color-components-premium-badge-blue-glow-hover: #d1e0ff; - --color-components-premium-badge-blue-stroke-stop-0-hover: rgb(255 255 255 / 0.5); - --color-components-premium-badge-blue-stroke-stop-100-hover: #296dff; - - --color-components-premium-badge-highlight-stop-0: rgb(255 255 255 / 0.12); - --color-components-premium-badge-highlight-stop-100: rgb(255 255 255 / 0.2); - --color-components-premium-badge-indigo-bg-stop-0: #6172f3; - --color-components-premium-badge-indigo-bg-stop-100: #3538cd; - --color-components-premium-badge-indigo-stroke-stop-0: rgb(255 255 255 / 0.2); - --color-components-premium-badge-indigo-stroke-stop-100: #444ce7; - --color-components-premium-badge-indigo-text-stop-0: #eef4ff; - --color-components-premium-badge-indigo-text-stop-100: #c7d7fe; - --color-components-premium-badge-indigo-glow: #3538cd; - --color-components-premium-badge-indigo-glow-hover: #e0eaff; - --color-components-premium-badge-indigo-bg-stop-0-hover: #a4bcfd; - --color-components-premium-badge-indigo-bg-stop-100-hover: #3538cd; - --color-components-premium-badge-indigo-stroke-stop-0-hover: rgb(255 255 255 / 0.5); - --color-components-premium-badge-indigo-stroke-stop-100-hover: #444ce7; - - --color-components-premium-badge-grey-bg-stop-0: #676f83; - --color-components-premium-badge-grey-bg-stop-100: #495464; - --color-components-premium-badge-grey-stroke-stop-0: rgb(255 255 255 / 0.12); - --color-components-premium-badge-grey-stroke-stop-100: #495464; - --color-components-premium-badge-grey-text-stop-0: #f9fafb; - --color-components-premium-badge-grey-text-stop-100: #e9ebf0; - --color-components-premium-badge-grey-glow: #354052; - --color-components-premium-badge-grey-glow-hover: #f2f4f7; - --color-components-premium-badge-grey-bg-stop-0-hover: #98a2b2; - --color-components-premium-badge-grey-bg-stop-100-hover: #354052; - --color-components-premium-badge-grey-stroke-stop-0-hover: rgb(255 255 255 / 0.5); - --color-components-premium-badge-grey-stroke-stop-100-hover: #676f83; - - --color-components-premium-badge-orange-bg-stop-0: #ff692e; - --color-components-premium-badge-orange-bg-stop-100: #e04f16; - --color-components-premium-badge-orange-stroke-stop-0: rgb(255 255 255 / 0.2); - --color-components-premium-badge-orange-stroke-stop-100: #ff4405; - --color-components-premium-badge-orange-text-stop-0: #fef6ee; - --color-components-premium-badge-orange-text-stop-100: #f9dbaf; - --color-components-premium-badge-orange-glow: #b93815; - --color-components-premium-badge-orange-glow-hover: #fdead7; - --color-components-premium-badge-orange-bg-stop-0-hover: #ff692e; - --color-components-premium-badge-orange-bg-stop-100-hover: #b93815; - --color-components-premium-badge-orange-stroke-stop-0-hover: rgb(255 255 255 / 0.5); - --color-components-premium-badge-orange-stroke-stop-100-hover: #ff4405; - - --color-components-progress-bar-bg: rgb(200 206 218 / 0.08); - --color-components-progress-bar-progress: rgb(200 206 218 / 0.14); - --color-components-progress-bar-border: rgb(255 255 255 / 0.03); - --color-components-progress-bar-progress-solid: rgb(255 255 255 / 0.95); - --color-components-progress-bar-progress-highlight: rgb(200 206 218 / 0.2); - - --color-components-icon-bg-red-solid: #d92d20; - --color-components-icon-bg-rose-solid: #e31b54; - --color-components-icon-bg-pink-solid: #dd2590; - --color-components-icon-bg-orange-dark-solid: #ff4405; - --color-components-icon-bg-yellow-solid: #eaaa08; - --color-components-icon-bg-green-solid: #4ca30d; - --color-components-icon-bg-teal-solid: #0e9384; - --color-components-icon-bg-blue-light-solid: #0ba5ec; - --color-components-icon-bg-blue-solid: #155aef; - --color-components-icon-bg-indigo-solid: #444ce7; - --color-components-icon-bg-violet-solid: #7839ee; - --color-components-icon-bg-midnight-solid: #5d698d; - --color-components-icon-bg-rose-soft: rgb(246 61 104 / 0.2); - --color-components-icon-bg-pink-soft: rgb(238 70 188 / 0.2); - --color-components-icon-bg-orange-dark-soft: rgb(255 68 5 / 0.2); - --color-components-icon-bg-yellow-soft: rgb(234 170 8 / 0.2); - --color-components-icon-bg-green-soft: rgb(102 198 28 / 0.2); - --color-components-icon-bg-teal-soft: rgb(21 183 158 / 0.2); - --color-components-icon-bg-blue-light-soft: rgb(11 165 236 / 0.2); - --color-components-icon-bg-blue-soft: rgb(21 90 239 / 0.2); - --color-components-icon-bg-indigo-soft: rgb(97 114 243 / 0.2); - --color-components-icon-bg-violet-soft: rgb(135 91 247 / 0.2); - --color-components-icon-bg-midnight-soft: rgb(130 141 173 / 0.2); - --color-components-icon-bg-red-soft: rgb(240 68 56 / 0.2); - --color-components-icon-bg-orange-solid: #f79009; - --color-components-icon-bg-orange-soft: rgb(247 144 9 / 0.2); - - --color-text-primary: #fbfbfc; - --color-text-secondary: #d9d9de; - --color-text-tertiary: rgb(200 206 218 / 0.6); - --color-text-quaternary: rgb(200 206 218 / 0.4); - --color-text-destructive: #f97066; - --color-text-success: #17b26a; - --color-text-warning: #f79009; - --color-text-destructive-secondary: #f97066; - --color-text-success-secondary: #47cd89; - --color-text-warning-secondary: #fdb022; - --color-text-accent: #5289ff; - --color-text-primary-on-surface: rgb(255 255 255 / 0.95); - --color-text-placeholder: rgb(200 206 218 / 0.3); - --color-text-disabled: rgb(200 206 218 / 0.3); - --color-text-accent-secondary: #84abff; - --color-text-accent-light-mode-only: #d9d9de; - --color-text-text-selected: rgb(21 90 239 / 0.3); - --color-text-secondary-on-surface: rgb(255 255 255 / 0.9); - --color-text-logo-text: #e9e9ec; - --color-text-empty-state-icon: rgb(200 206 218 / 0.3); - --color-text-inverted: #ffffff; - --color-text-inverted-dimmed: rgb(255 255 255 / 0.8); - - --color-background-body: #1d1d20; - --color-background-default-subtle: #222225; - --color-background-neutral-subtle: #1d1d20; - --color-background-sidenav-bg: rgb(39 39 42 / 0.92); - --color-background-default: #222225; - --color-background-soft: rgb(24 24 27 / 0.25); - --color-background-gradient-bg-fill-chat-bg-1: #222225; - --color-background-gradient-bg-fill-chat-bg-2: #1d1d20; - --color-background-gradient-bg-fill-chat-bubble-bg-1: rgb(200 206 218 / 0.08); - --color-background-gradient-bg-fill-chat-bubble-bg-2: rgb(200 206 218 / 0.02); - --color-background-gradient-bg-fill-debug-bg-1: rgb(200 206 218 / 0.08); - --color-background-gradient-bg-fill-debug-bg-2: rgb(24 24 27 / 0.04); - - --color-background-gradient-mask-gray: rgb(24 24 27 / 0.08); - --color-background-gradient-mask-transparent: rgb(0 0 0 / 0); - --color-background-gradient-mask-input-clear-2: rgb(57 58 62 / 0); - --color-background-gradient-mask-input-clear-1: #393a3e; - --color-background-gradient-mask-transparent-dark: rgb(0 0 0 / 0); - --color-background-gradient-mask-side-panel-2: rgb(24 24 27 / 0.9); - --color-background-gradient-mask-side-panel-1: rgb(24 24 27 / 0.04); - - --color-background-default-burn: #1d1d20; - --color-background-overlay-fullscreen: rgb(39 39 42 / 0.97); - --color-background-default-lighter: rgb(200 206 218 / 0.04); - --color-background-section: rgb(24 24 27 / 0.4); - --color-background-interaction-from-bg-1: rgb(24 24 27 / 0.4); - --color-background-interaction-from-bg-2: rgb(24 24 27 / 0.14); - --color-background-section-burn: rgb(24 24 27 / 0.6); - --color-background-default-dodge: #3a3a40; - --color-background-overlay: rgb(24 24 27 / 0.8); - --color-background-default-dimmed: #27272b; - --color-background-default-hover: #27272b; - --color-background-overlay-alt: rgb(24 24 27 / 0.4); - --color-background-surface-white: rgb(255 255 255 / 0.9); - --color-background-overlay-destructive: rgb(240 68 56 / 0.3); - --color-background-overlay-backdrop: rgb(24 24 27 / 0.95); - --color-background-body-transparent: rgb(29 29 32 / 0); - --color-background-section-burn-inverted: #27272b; + --color-effects-highlight: rgb(200 206 218 / 0.08); + --color-effects-highlight-subtle: rgb(200 206 218 / 0.04); + --color-effects-highlight-lightmode-off: rgb(200 206 218 / 0.08); + --color-effects-image-frame: #ffffff; + --color-effects-icon-border: rgb(255 255 255 / 0.15); --color-shadow-shadow-1: rgb(0 0 0 / 0.05); + --color-shadow-shadow-2: rgb(0 0 0 / 0.08); --color-shadow-shadow-3: rgb(0 0 0 / 0.1); --color-shadow-shadow-4: rgb(0 0 0 / 0.12); --color-shadow-shadow-5: rgb(0 0 0 / 0.16); @@ -441,124 +589,18 @@ html[data-theme="dark"] { --color-shadow-shadow-7: rgb(0 0 0 / 0.24); --color-shadow-shadow-8: rgb(0 0 0 / 0.28); --color-shadow-shadow-9: rgb(0 0 0 / 0.36); - --color-shadow-shadow-2: rgb(0 0 0 / 0.08); --color-shadow-shadow-10: rgb(0 0 0 / 0.4); - --color-workflow-block-border: rgb(255 255 255 / 0.08); - --color-workflow-block-parma-bg: rgb(255 255 255 / 0.05); - --color-workflow-block-bg: #27272b; - --color-workflow-block-bg-transparent: rgb(39 39 43 / 0.96); - --color-workflow-block-border-highlight: rgb(200 206 218 / 0.2); - --color-workflow-block-wrapper-bg-1: #27272b; - --color-workflow-block-wrapper-bg-2: rgb(39 39 43 / 0.2); - - --color-workflow-canvas-workflow-dot-color: rgb(133 133 173 / 0.11); - --color-workflow-canvas-workflow-bg: #1d1d20; - --color-workflow-canvas-workflow-top-bar-1: rgb(29 29 32 / 0.9); - --color-workflow-canvas-workflow-top-bar-2: rgb(29 29 32 / 0.08); - --color-workflow-canvas-canvas-overlay: rgb(29 29 32 / 0.8); - - --color-workflow-link-line-active: #5289ff; - --color-workflow-link-line-normal: #676f83; - --color-workflow-link-line-handle: #5289ff; - --color-workflow-link-line-normal-transparent: rgb(103 111 131 / 0.2); - --color-workflow-link-line-failure-active: #fdb022; - --color-workflow-link-line-failure-handle: #fdb022; - --color-workflow-link-line-failure-button-bg: #f79009; - --color-workflow-link-line-failure-button-hover: #dc6803; - - --color-workflow-link-line-success-active: #47cd89; - --color-workflow-link-line-success-handle: #47cd89; - - --color-workflow-link-line-error-active: #f97066; - --color-workflow-link-line-error-handle: #f97066; - - --color-workflow-minimap-bg: #27272b; - --color-workflow-minimap-block: rgb(200 206 218 / 0.08); - - --color-workflow-display-success-bg: rgb(23 178 106 / 0.2); - --color-workflow-display-success-border-1: rgb(23 178 106 / 0.9); - --color-workflow-display-success-border-2: rgb(23 178 106 / 0.8); - --color-workflow-display-success-vignette-color: rgb(23 178 106 / 0.25); - --color-workflow-display-success-bg-line-pattern: rgb(24 24 27 / 0.8); - - --color-workflow-display-glass-1: rgb(255 255 255 / 0.03); - --color-workflow-display-glass-2: rgb(255 255 255 / 0.05); - --color-workflow-display-vignette-dark: rgb(0 0 0 / 0.4); - --color-workflow-display-highlight: rgb(255 255 255 / 0.12); - --color-workflow-display-outline: rgb(24 24 27 / 0.95); - --color-workflow-display-error-bg: rgb(240 68 56 / 0.2); - --color-workflow-display-error-bg-line-pattern: rgb(24 24 27 / 0.8); - --color-workflow-display-error-border-1: rgb(240 68 56 / 0.9); - --color-workflow-display-error-border-2: rgb(240 68 56 / 0.8); - --color-workflow-display-error-vignette-color: rgb(240 68 56 / 0.25); - - --color-workflow-display-warning-bg: rgb(247 144 9 / 0.2); - --color-workflow-display-warning-bg-line-pattern: rgb(24 24 27 / 0.8); - --color-workflow-display-warning-border-1: rgb(247 144 9 / 0.9); - --color-workflow-display-warning-border-2: rgb(247 144 9 / 0.8); - --color-workflow-display-warning-vignette-color: rgb(247 144 9 / 0.25); - - --color-workflow-display-normal-bg: rgb(11 165 236 / 0.2); - --color-workflow-display-normal-bg-line-pattern: rgb(24 24 27 / 0.8); - --color-workflow-display-normal-border-1: rgb(11 165 236 / 0.9); - --color-workflow-display-normal-border-2: rgb(11 165 236 / 0.8); - --color-workflow-display-normal-vignette-color: rgb(11 165 236 / 0.25); - - --color-workflow-display-disabled-bg: rgb(200 206 218 / 0.2); - --color-workflow-display-disabled-bg-line-pattern: rgb(24 24 27 / 0.8); - --color-workflow-display-disabled-border-1: rgb(200 206 218 / 0.6); - --color-workflow-display-disabled-border-2: rgb(200 206 218 / 0.25); - --color-workflow-display-disabled-vignette-color: rgb(200 206 218 / 0.25); - --color-workflow-display-disabled-outline: rgb(24 24 27 / 0.95); - - --color-workflow-workflow-progress-bg-1: rgb(24 24 27 / 0.25); - --color-workflow-workflow-progress-bg-2: rgb(24 24 27 / 0.04); - - --color-divider-subtle: rgb(200 206 218 / 0.08); - --color-divider-regular: rgb(200 206 218 / 0.14); - --color-divider-deep: rgb(200 206 218 / 0.2); - --color-divider-burn: rgb(24 24 27 / 0.95); - --color-divider-intense: rgb(200 206 218 / 0.4); - --color-divider-solid: #3a3a40; - --color-divider-solid-alt: #747481; - --color-divider-accent: rgb(200 206 218 / 0.14); - - --color-state-base-hover: rgb(200 206 218 / 0.08); - --color-state-base-active: rgb(200 206 218 / 0.2); - --color-state-base-hover-alt: rgb(200 206 218 / 0.14); - --color-state-base-handle: rgb(200 206 218 / 0.3); - --color-state-base-handle-hover: rgb(200 206 218 / 0.5); - --color-state-base-hover-subtle: rgb(200 206 218 / 0.04); - - --color-state-accent-hover: rgb(21 90 239 / 0.14); - --color-state-accent-active: rgb(21 90 239 / 0.14); - --color-state-accent-hover-alt: rgb(21 90 239 / 0.25); - --color-state-accent-solid: #5289ff; - --color-state-accent-active-alt: rgb(21 90 239 / 0.2); - - --color-state-destructive-hover: rgb(240 68 56 / 0.14); - --color-state-destructive-hover-alt: rgb(240 68 56 / 0.25); - --color-state-destructive-active: rgb(240 68 56 / 0.3); - --color-state-destructive-solid: #f97066; - --color-state-destructive-border: #f97066; - --color-state-destructive-hover-transparent: rgb(240 68 56 / 0); - - --color-state-success-hover: rgb(23 178 106 / 0.14); - --color-state-success-hover-alt: rgb(23 178 106 / 0.25); - --color-state-success-active: rgb(23 178 106 / 0.3); - --color-state-success-solid: #47cd89; - - --color-state-warning-hover: rgb(247 144 9 / 0.14); - --color-state-warning-hover-alt: rgb(247 144 9 / 0.25); - --color-state-warning-active: rgb(247 144 9 / 0.3); - --color-state-warning-solid: #f79009; - --color-state-warning-hover-transparent: rgb(247 144 9 / 0); - - --color-effects-highlight: rgb(200 206 218 / 0.08); - --color-effects-highlight-lightmode-off: rgb(200 206 218 / 0.08); - --color-effects-image-frame: #ffffff; - --color-effects-icon-border: rgb(255 255 255 / 0.15); + --color-third-party-LangChain: #ffffff; + --color-third-party-Langfuse: #ffffff; + --color-third-party-Github: #ffffff; + --color-third-party-Github-tertiary: rgb(200 206 218 / 0.6); + --color-third-party-Github-secondary: #d9d9de; + --color-third-party-aws: #141f2e; + --color-third-party-aws-alt: #192639; + --color-third-party-model-bg-openai: #121212; + --color-third-party-model-bg-anthropic: #1d1917; + --color-third-party-model-bg-default: #1d1d20; --color-util-colors-orange-dark-orange-dark-50: #57130a; --color-util-colors-orange-dark-orange-dark-100: #771a0d; @@ -633,6 +675,15 @@ html[data-theme="dark"] { --color-util-colors-blue-light-blue-light-600: #36bffa; --color-util-colors-blue-light-blue-light-700: #7cd4fd; + --color-util-colors-blue-brand-blue-brand-50: #001782; + --color-util-colors-blue-brand-blue-brand-100: #001ea0; + --color-util-colors-blue-brand-blue-brand-200: #0025be; + --color-util-colors-blue-brand-blue-brand-300: #002cde; + --color-util-colors-blue-brand-blue-brand-400: #0033ff; + --color-util-colors-blue-brand-blue-brand-500: #085afc; + --color-util-colors-blue-brand-blue-brand-600: #3072ff; + --color-util-colors-blue-brand-blue-brand-700: #6694ff; + --color-util-colors-gray-blue-gray-blue-50: #0d0f1c; --color-util-colors-gray-blue-gray-blue-100: #101323; --color-util-colors-gray-blue-gray-blue-200: #293056; @@ -642,15 +693,6 @@ html[data-theme="dark"] { --color-util-colors-gray-blue-gray-blue-600: #717bbc; --color-util-colors-gray-blue-gray-blue-700: #b3b8db; - --color-util-colors-blue-brand-blue-brand-50: #002066; - --color-util-colors-blue-brand-blue-brand-100: #00329e; - --color-util-colors-blue-brand-blue-brand-200: #003dc1; - --color-util-colors-blue-brand-blue-brand-300: #004aeb; - --color-util-colors-blue-brand-blue-brand-400: #155aef; - --color-util-colors-blue-brand-blue-brand-500: #296dff; - --color-util-colors-blue-brand-blue-brand-600: #5289ff; - --color-util-colors-blue-brand-blue-brand-700: #84abff; - --color-util-colors-red-red-50: #55160c; --color-util-colors-red-red-100: #7a271a; --color-util-colors-red-red-200: #912018; @@ -669,6 +711,15 @@ html[data-theme="dark"] { --color-util-colors-green-green-600: #47cd89; --color-util-colors-green-green-700: #75e0a7; + --color-util-colors-green-light-green-light-50: #15290a; + --color-util-colors-green-light-green-light-100: #2b5314; + --color-util-colors-green-light-green-light-200: #326212; + --color-util-colors-green-light-green-light-300: #3b7c0f; + --color-util-colors-green-light-green-light-400: #4ca30d; + --color-util-colors-green-light-green-light-500: #66c61c; + --color-util-colors-green-light-green-light-600: #85e13a; + --color-util-colors-green-light-green-light-700: #a6ef67; + --color-util-colors-warning-warning-50: #4e1d09; --color-util-colors-warning-warning-100: #7a2e0e; --color-util-colors-warning-warning-200: #93370d; @@ -723,15 +774,6 @@ html[data-theme="dark"] { --color-util-colors-gray-gray-600: #98a2b2; --color-util-colors-gray-gray-700: #d0d5dc; - --color-util-colors-green-light-green-light-50: #15290a; - --color-util-colors-green-light-green-light-100: #2b5314; - --color-util-colors-green-light-green-light-200: #326212; - --color-util-colors-green-light-green-light-300: #3b7c0f; - --color-util-colors-green-light-green-light-500: #66c61c; - --color-util-colors-green-light-green-light-400: #4ca30d; - --color-util-colors-green-light-green-light-600: #85e13a; - --color-util-colors-green-light-green-light-700: #a6ef67; - --color-util-colors-rose-rose-50: #510b24; --color-util-colors-rose-rose-100: #89123e; --color-util-colors-rose-rose-200: #a11043; @@ -750,30 +792,31 @@ html[data-theme="dark"] { --color-util-colors-midnight-midnight-600: #a7aec5; --color-util-colors-midnight-midnight-700: #c6cbd9; - --color-third-party-LangChain: #ffffff; - --color-third-party-Langfuse: #ffffff; - --color-third-party-Github: #ffffff; - --color-third-party-Github-tertiary: rgb(200 206 218 / 0.6); - --color-third-party-Github-secondary: #d9d9de; - --color-third-party-model-bg-openai: #121212; - --color-third-party-model-bg-anthropic: #1d1917; - --color-third-party-model-bg-default: #1d1d20; - - --color-third-party-aws: #141f2e; - --color-third-party-aws-alt: #192639; - --color-saas-background: #0b0b0e; + --color-saas-background-inverted: rgb(255 255 255 / 0.9); + --color-saas-background-inverted-hover: #ffffff; --color-saas-pricing-grid-bg: rgb(200 206 218 / 0.2); --color-saas-dify-blue-static: #0033ff; - --color-saas-dify-blue-static-hover: #002cd6; + --color-saas-dify-blue-static-hover: #085afc; --color-saas-dify-blue-accessible: #0a68ff; --color-saas-dify-blue-inverted: #ffffff; --color-saas-dify-blue-inverted-dimmed: rgb(255 255 255 / 0.88); - --color-saas-background-inverted: rgb(255 255 255 / 0.9); - --color-saas-background-inverted-hover: #ffffff; - --color-dify-logo-blue: #e8e8e8; --color-dify-logo-black: #e8e8e8; + --color-dify-logo-outline-1: #ffffff; + --color-dify-logo-outline-2: #e8e8e8; + + --color-brand-color-opacity-50: #020821; + --color-brand-color-opacity-100: #031042; + --color-brand-color-opacity-200: #051969; + --color-brand-color-opacity-300: #07228e; + --color-brand-color-opacity-400: #082bb4; + --color-brand-color-opacity-500: #0033ff; + --color-brand-color-opacity-600: #4f6de5; + --color-brand-color-opacity-700: #899eee; + --color-brand-color-opacity-800: #bac6f5; + --color-brand-color-opacity-900: #e2e7fb; + --color-brand-color-opacity-1000: #f3f5fd; } diff --git a/packages/dify-ui/src/themes/light.css b/packages/dify-ui/src/themes/light.css index dd3252f3614..059e1e207e1 100644 --- a/packages/dify-ui/src/themes/light.css +++ b/packages/dify-ui/src/themes/light.css @@ -1,33 +1,367 @@ /* Attention: Generate by code. Don't update by hand!!! */ html[data-theme="light"] { + --color-text-primary: #101828; + --color-text-secondary: #354052; + --color-text-tertiary: #676f83; + --color-text-quaternary: rgb(16 24 40 / 0.3); + --color-text-empty-state-icon: #d0d5dc; + --color-text-primary-on-surface: #ffffff; + --color-text-secondary-on-surface: rgb(255 255 255 / 0.9); + --color-text-destructive: #d92d20; + --color-text-destructive-secondary: #f04438; + --color-text-success: #079455; + --color-text-success-secondary: #17b26a; + --color-text-warning: #dc6803; + --color-text-warning-secondary: #f79009; + --color-text-accent: #0033ff; + --color-text-accent-secondary: #085afc; + --color-text-accent-light-mode-only: #0033ff; + --color-text-placeholder: #98a2b2; + --color-text-disabled: #d0d5dc; + --color-text-text-selected: rgb(0 51 255 / 0.14); + --color-text-logo-text: #18222f; + --color-text-inverted: #000000; + --color-text-inverted-dimmed: rgb(0 0 0 / 0.95); + + --color-background-body: #f2f4f7; + --color-background-body-transparent: rgb(242 244 247 / 0); + --color-background-default: #ffffff; + --color-background-default-hover: #f9fafb; + --color-background-default-hover-alpha-0: rgb(249 250 251 / 0); + --color-background-default-subtle: #fcfcfd; + --color-background-default-dodge: #ffffff; + --color-background-default-burn: #e9ebf0; + --color-background-default-dimmed: #e9ebf0; + --color-background-default-lighter: rgb(255 255 255 / 0.5); + --color-background-section: #f9fafb; + --color-background-section-burn: #f2f4f7; + --color-background-section-burn-inverted: #f2f4f7; + --color-background-soft: #f9fafb; + --color-background-neutral-subtle: #f9fafb; + --color-background-surface-white: rgb(255 255 255 / 0.95); + --color-background-sidenav-bg: rgb(255 255 255 / 0.8); + --color-background-overlay-fullscreen: rgb(249 250 251 / 0.95); + --color-background-overlay-backdrop: rgb(242 244 247 / 0.95); + --color-background-overlay: rgb(16 24 40 / 0.6); + --color-background-overlay-alt: rgb(16 24 40 / 0.4); + --color-background-overlay-destructive: rgb(240 68 56 / 0.3); + --color-background-interaction-from-bg-1: rgb(200 206 218 / 0.2); + --color-background-interaction-from-bg-2: rgb(200 206 218 / 0.14); + --color-background-gradient-bg-fill-chat-bg-1: #f9fafb; + --color-background-gradient-bg-fill-chat-bg-2: #f2f4f7; + --color-background-gradient-bg-fill-chat-bubble-bg-1: #ffffff; + --color-background-gradient-bg-fill-chat-bubble-bg-2: rgb(255 255 255 / 0.6); + --color-background-gradient-bg-fill-debug-bg-1: rgb(255 255 255 / 0); + --color-background-gradient-bg-fill-debug-bg-2: rgb(200 206 218 / 0.14); + + --color-background-gradient-mask-gray: rgb(200 206 218 / 0.2); + --color-background-gradient-mask-transparent: rgb(255 255 255 / 0); + --color-background-gradient-mask-transparent-dark: rgb(0 0 0 / 0); + --color-background-gradient-mask-side-panel-1: rgb(16 24 40 / 0.02); + --color-background-gradient-mask-side-panel-2: rgb(16 24 40 / 0.3); + --color-background-gradient-mask-input-clear-1: #e9ebf0; + --color-background-gradient-mask-input-clear-2: rgb(233 235 240 / 0); + + --color-state-base-hover-subtle: rgb(200 206 218 / 0.08); + --color-state-base-hover: rgb(200 206 218 / 0.2); + --color-state-base-hover-alt: rgb(200 206 218 / 0.4); + --color-state-base-active: rgb(200 206 218 / 0.4); + --color-state-base-handle: rgb(16 24 40 / 0.2); + --color-state-base-handle-hover: rgb(16 24 40 / 0.3); + + --color-state-accent-hover: #e9f0ff; + --color-state-accent-hover-alt: #d0dfff; + --color-state-accent-active: rgb(0 51 255 / 0.08); + --color-state-accent-active-alt: rgb(0 51 255 / 0.14); + --color-state-accent-solid: #0033ff; + + --color-state-destructive-hover: #fef3f2; + --color-state-destructive-hover-transparent: rgb(254 243 242 / 0); + --color-state-destructive-hover-alt: #fee4e2; + --color-state-destructive-active: #fecdca; + --color-state-destructive-solid: #f04438; + --color-state-destructive-border: #fda29b; + + --color-state-warning-hover: #fffaeb; + --color-state-warning-hover-transparent: rgb(255 250 235 / 0); + --color-state-warning-hover-alt: #fef0c7; + --color-state-warning-active: #fedf89; + --color-state-warning-solid: #f79009; + + --color-state-success-hover: #ecfdf3; + --color-state-success-hover-alt: #dcfae6; + --color-state-success-active: #abefc6; + --color-state-success-solid: #17b26a; + + --color-workflow-workflow-progress-bg-1: rgb(200 206 218 / 0.2); + --color-workflow-workflow-progress-bg-2: rgb(200 206 218 / 0.04); + + --color-workflow-block-bg: #fcfcfd; + --color-workflow-block-bg-transparent: rgb(252 252 253 / 0.9); + --color-workflow-block-border: #ffffff; + --color-workflow-block-border-highlight: rgb(0 51 255 / 0.14); + --color-workflow-block-parma-bg: #f2f4f7; + --color-workflow-block-wrapper-bg-1: #e9ebf0; + --color-workflow-block-wrapper-bg-2: rgb(233 235 240 / 0.2); + + --color-workflow-link-line-normal: #d0d5dc; + --color-workflow-link-line-normal-transparent: rgb(208 213 220 / 0.2); + --color-workflow-link-line-active: #085afc; + --color-workflow-link-line-handle: #085afc; + --color-workflow-link-line-failure-active: #f79009; + --color-workflow-link-line-failure-handle: #f79009; + --color-workflow-link-line-failure-button-bg: #dc6803; + --color-workflow-link-line-failure-button-hover: #b54708; + + --color-workflow-link-line-success-active: #17b26a; + --color-workflow-link-line-success-handle: #17b26a; + + --color-workflow-link-line-error-active: #f04438; + --color-workflow-link-line-error-handle: #f04438; + + --color-workflow-minimap-block: rgb(200 206 218 / 0.3); + --color-workflow-minimap-bg: #e9ebf0; + + --color-workflow-display-glass-1: rgb(255 255 255 / 0.12); + --color-workflow-display-glass-2: rgb(255 255 255 / 0.5); + --color-workflow-display-highlight: rgb(255 255 255 / 0.5); + --color-workflow-display-outline: rgb(0 0 0 / 0.05); + --color-workflow-display-vignette-dark: rgb(0 0 0 / 0.12); + --color-workflow-display-success-bg: #ecfdf3; + --color-workflow-display-success-bg-line-pattern: rgb(23 178 106 / 0.3); + --color-workflow-display-success-border-1: rgb(23 178 106 / 0.8); + --color-workflow-display-success-border-2: rgb(23 178 106 / 0.5); + --color-workflow-display-success-vignette-color: rgb(23 178 106 / 0.2); + + --color-workflow-display-error-bg: #fef3f2; + --color-workflow-display-error-bg-line-pattern: rgb(240 68 56 / 0.3); + --color-workflow-display-error-border-1: rgb(240 68 56 / 0.8); + --color-workflow-display-error-border-2: rgb(240 68 56 / 0.5); + --color-workflow-display-error-vignette-color: rgb(240 68 56 / 0.2); + + --color-workflow-display-warning-bg: #fffaeb; + --color-workflow-display-warning-bg-line-pattern: rgb(247 144 9 / 0.3); + --color-workflow-display-warning-border-1: rgb(247 144 9 / 0.8); + --color-workflow-display-warning-border-2: rgb(247 144 9 / 0.5); + --color-workflow-display-warning-vignette-color: rgb(247 144 9 / 0.2); + + --color-workflow-display-normal-bg: #f0f9ff; + --color-workflow-display-normal-bg-line-pattern: rgb(11 165 236 / 0.3); + --color-workflow-display-normal-border-1: rgb(11 165 236 / 0.8); + --color-workflow-display-normal-border-2: rgb(11 165 236 / 0.5); + --color-workflow-display-normal-vignette-color: rgb(11 165 236 / 0.2); + + --color-workflow-display-disabled-bg: #f9fafb; + --color-workflow-display-disabled-bg-line-pattern: rgb(200 206 218 / 0.3); + --color-workflow-display-disabled-border-1: rgb(200 206 218 / 0.6); + --color-workflow-display-disabled-border-2: rgb(200 206 218 / 0.4); + --color-workflow-display-disabled-vignette-color: rgb(200 206 218 / 0.4); + --color-workflow-display-disabled-outline: rgb(0 0 0 / 0); + + --color-workflow-canvas-workflow-dot-color: rgb(133 133 173 / 0.15); + --color-workflow-canvas-workflow-bg: #f2f4f7; + --color-workflow-canvas-workflow-top-bar-1: rgb(242 244 247 / 0.9); + --color-workflow-canvas-workflow-top-bar-2: rgb(242 244 247 / 0.05); + --color-workflow-canvas-canvas-overlay: rgb(242 244 247 / 0.8); + + --color-workflow-debug-run-status-bg: rgb(255 68 5 / 0.08); + --color-workflow-debug-run-status-bg-alt: rgb(255 68 5 / 0.14); + --color-workflow-debug-breakpoint: #e62e05; + --color-workflow-debug-text: #e62e05; + --color-workflow-debug-text-disabled: rgb(255 68 5 / 0.2); + + --color-workflow-test-run-run-status-bg: rgb(0 51 255 / 0.08); + --color-workflow-test-run-paused-bg: rgb(247 144 9 / 0.14); + --color-workflow-test-run-paused-text: #dc6803; + --color-workflow-test-run-run-status-bg-alt: rgb(0 51 255 / 0.14); + --color-workflow-test-run-text: #002cde; + + --color-components-icon-bg-red-solid: #d92d20; + --color-components-icon-bg-rose-solid: #e31b54; + --color-components-icon-bg-pink-solid: #dd2590; + --color-components-icon-bg-orange-dark-solid: #ff4405; + --color-components-icon-bg-orange-solid: #f79009; + --color-components-icon-bg-yellow-solid: #eaaa08; + --color-components-icon-bg-green-solid: #4ca30d; + --color-components-icon-bg-teal-solid: #0e9384; + --color-components-icon-bg-blue-light-solid: #0ba5ec; + --color-components-icon-bg-blue-solid: #0033ff; + --color-components-icon-bg-indigo-solid: #444ce7; + --color-components-icon-bg-violet-solid: #7839ee; + --color-components-icon-bg-midnight-solid: #828dad; + --color-components-icon-bg-red-soft: #fef3f2; + --color-components-icon-bg-rose-soft: #fff1f3; + --color-components-icon-bg-pink-soft: #fdf2fa; + --color-components-icon-bg-orange-dark-soft: #fff4ed; + --color-components-icon-bg-orange-soft: #fffaeb; + --color-components-icon-bg-yellow-soft: #fefbe8; + --color-components-icon-bg-green-soft: #f3fee7; + --color-components-icon-bg-teal-soft: #f0fdf9; + --color-components-icon-bg-blue-light-soft: #f0f9ff; + --color-components-icon-bg-blue-soft: #e9f0ff; + --color-components-icon-bg-indigo-soft: #eef4ff; + --color-components-icon-bg-violet-soft: #f5f3ff; + --color-components-icon-bg-midnight-soft: #f0f2f5; + + --color-components-avatar-default-avatar-bg: #d0d5dc; + --color-components-avatar-mask-darkmode-dimmed: rgb(255 255 255 / 0); + --color-components-avatar-shape-fill-stop-0: #ffffff; + --color-components-avatar-shape-fill-stop-100: rgb(255 255 255 / 0.9); + + --color-components-avatar-bg-mask-stop-0: rgb(255 255 255 / 0.12); + --color-components-avatar-bg-mask-stop-100: rgb(255 255 255 / 0.08); + + --color-components-chat-input-audio-bg: #e9f0ff; + --color-components-chat-input-audio-bg-alt: #fcfcfd; + --color-components-chat-input-audio-wave-default: rgb(0 51 255 / 0.2); + --color-components-chat-input-audio-wave-active: #0033ff; + --color-components-chat-input-bg-mask-1: rgb(255 255 255 / 0.01); + --color-components-chat-input-bg-mask-2: #f2f4f7; + --color-components-chat-input-border: #ffffff; + + --color-components-label-gray: #f2f4f7; + + --color-components-premium-badge-highlight-stop-0: rgb(255 255 255 / 0.12); + --color-components-premium-badge-highlight-stop-100: rgb(255 255 255 / 0.3); + --color-components-premium-badge-orange-bg-stop-0: #ff692e; + --color-components-premium-badge-orange-bg-stop-100: #e04f16; + --color-components-premium-badge-orange-stroke-stop-0: rgb(255 255 255 / 0.95); + --color-components-premium-badge-orange-stroke-stop-100: #e62e05; + --color-components-premium-badge-orange-text-stop-0: #fefaf5; + --color-components-premium-badge-orange-text-stop-100: #fdead7; + --color-components-premium-badge-orange-glow: #772917; + --color-components-premium-badge-orange-glow-hover: #f7b27a; + --color-components-premium-badge-orange-bg-stop-0-hover: #ff4405; + --color-components-premium-badge-orange-bg-stop-100-hover: #b93815; + --color-components-premium-badge-orange-stroke-stop-0-hover: rgb(255 255 255 / 0.95); + --color-components-premium-badge-orange-stroke-stop-100-hover: #bc1b06; + + --color-components-premium-badge-blue-bg-stop-0: #3072ff; + --color-components-premium-badge-blue-bg-stop-100: #0033ff; + --color-components-premium-badge-blue-stroke-stop-0: rgb(255 255 255 / 0.95); + --color-components-premium-badge-blue-stroke-stop-100: #0033ff; + --color-components-premium-badge-blue-text-stop-0: #f5f8ff; + --color-components-premium-badge-blue-text-stop-100: #d0dfff; + --color-components-premium-badge-blue-glow: #001ea0; + --color-components-premium-badge-blue-glow-hover: #6694ff; + --color-components-premium-badge-blue-bg-stop-0-hover: #085afc; + --color-components-premium-badge-blue-bg-stop-100-hover: #002cde; + --color-components-premium-badge-blue-stroke-stop-0-hover: rgb(255 255 255 / 0.95); + --color-components-premium-badge-blue-stroke-stop-100-hover: #001ea0; + + --color-components-premium-badge-indigo-bg-stop-0: #8098f9; + --color-components-premium-badge-indigo-bg-stop-100: #444ce7; + --color-components-premium-badge-indigo-stroke-stop-0: rgb(255 255 255 / 0.95); + --color-components-premium-badge-indigo-stroke-stop-100: #6172f3; + --color-components-premium-badge-indigo-text-stop-0: #f5f8ff; + --color-components-premium-badge-indigo-text-stop-100: #e0eaff; + --color-components-premium-badge-indigo-glow: #2d3282; + --color-components-premium-badge-indigo-glow-hover: #a4bcfd; + --color-components-premium-badge-indigo-bg-stop-0-hover: #6172f3; + --color-components-premium-badge-indigo-bg-stop-100-hover: #2d31a6; + --color-components-premium-badge-indigo-stroke-stop-0-hover: rgb(255 255 255 / 0.95); + --color-components-premium-badge-indigo-stroke-stop-100-hover: #2d31a6; + + --color-components-premium-badge-grey-bg-stop-0: #98a2b2; + --color-components-premium-badge-grey-bg-stop-100: #676f83; + --color-components-premium-badge-grey-stroke-stop-0: rgb(255 255 255 / 0.95); + --color-components-premium-badge-grey-stroke-stop-100: #676f83; + --color-components-premium-badge-grey-text-stop-0: #fcfcfd; + --color-components-premium-badge-grey-text-stop-100: #f2f4f7; + --color-components-premium-badge-grey-glow: #101828; + --color-components-premium-badge-grey-glow-hover: #d0d5dc; + --color-components-premium-badge-grey-bg-stop-0-hover: #676f83; + --color-components-premium-badge-grey-bg-stop-100-hover: #354052; + --color-components-premium-badge-grey-stroke-stop-0-hover: rgb(255 255 255 / 0.95); + --color-components-premium-badge-grey-stroke-stop-100-hover: #354052; + + --color-components-dropzone-bg: #f9fafb; + --color-components-dropzone-bg-alt: #f2f4f7; + --color-components-dropzone-bg-accent: rgb(0 51 255 / 0.14); + --color-components-dropzone-border: rgb(16 24 40 / 0.08); + --color-components-dropzone-border-alt: rgb(16 24 40 / 0.2); + --color-components-dropzone-border-accent: #6694ff; + + --color-components-panel-bg: #ffffff; + --color-components-panel-bg-transparent: rgb(255 255 255 / 0); + --color-components-panel-bg-alt: #f9fafb; + --color-components-panel-bg-blur: rgb(255 255 255 / 0.95); + --color-components-panel-bg-blur-burn: rgb(255 255 255 / 0.9); + --color-components-panel-border: rgb(16 24 40 / 0.08); + --color-components-panel-border-subtle: rgb(16 24 40 / 0.08); + --color-components-panel-gradient-1: #ffffff; + --color-components-panel-gradient-2: #f9fafb; + --color-components-panel-on-panel-item-bg: #ffffff; + --color-components-panel-on-panel-item-bg-transparent: rgb(255 255 255 / 0.95); + --color-components-panel-on-panel-item-bg-hover: #f9fafb; + --color-components-panel-on-panel-item-bg-hover-transparent: rgb(249 250 251 / 0); + --color-components-panel-on-panel-item-bg-destructive-hover-transparent: rgb(254 243 242 / 0); + --color-components-panel-on-panel-item-bg-alt: #f9fafb; + + --color-components-marketplace-header-bg: rgb(255 255 255 / 0.98); + + --color-components-card-bg: #fcfcfd; + --color-components-card-bg-alt: #ffffff; + --color-components-card-bg-transparent: rgb(252 252 253 / 0); + --color-components-card-bg-alt-transparent: rgb(255 255 255 / 0); + --color-components-card-border: #ffffff; + + --color-components-actionbar-border: rgb(16 24 40 / 0.04); + --color-components-actionbar-bg: rgb(255 255 255 / 0.95); + --color-components-actionbar-bg-accent: #f5f8ff; + --color-components-actionbar-border-accent: #a0bdff; + --color-components-input-bg-normal: rgb(200 206 218 / 0.25); - --color-components-input-text-placeholder: #98a2b2; --color-components-input-bg-hover: rgb(200 206 218 / 0.14); --color-components-input-bg-active: #f9fafb; - --color-components-input-border-active: #d0d5dc; - --color-components-input-border-destructive: #fda29b; - --color-components-input-text-filled: #101828; - --color-components-input-bg-destructive: #ffffff; --color-components-input-bg-disabled: rgb(200 206 218 / 0.14); + --color-components-input-bg-destructive: #ffffff; + --color-components-input-text-filled: #101828; + --color-components-input-text-placeholder: #98a2b2; --color-components-input-text-disabled: #d0d5dc; --color-components-input-text-filled-disabled: #676f83; + --color-components-input-text-filled-blue: #0033ff; --color-components-input-border-hover: #d0d5dc; + --color-components-input-border-active: #d0d5dc; --color-components-input-border-active-prompt-1: #0ba5ec; - --color-components-input-border-active-prompt-2: #155aef; + --color-components-input-border-active-prompt-2: #0033ff; + --color-components-input-border-destructive: #fda29b; - --color-components-kbd-bg-gray: rgb(16 24 40 / 0.04); - --color-components-kbd-bg-white: rgb(255 255 255 / 0.12); + --color-components-progress-brand-progress: #0033ff; + --color-components-progress-brand-border: #0033ff; + --color-components-progress-brand-bg: rgb(0 51 255 / 0.04); - --color-components-tooltip-bg: rgb(255 255 255 / 0.95); + --color-components-progress-white-progress: #ffffff; + --color-components-progress-white-border: rgb(255 255 255 / 0.95); + --color-components-progress-white-bg: rgb(255 255 255 / 0.01); + --color-components-progress-gray-progress: #98a2b2; + --color-components-progress-gray-border: #98a2b2; + --color-components-progress-gray-bg: rgb(200 206 218 / 0.02); + + --color-components-progress-warning-progress: #f79009; + --color-components-progress-warning-border: #f79009; + --color-components-progress-warning-bg: rgb(247 144 9 / 0.04); + + --color-components-progress-error-progress: #f04438; + --color-components-progress-error-border: #f04438; + --color-components-progress-error-bg: rgb(240 68 56 / 0.04); + + --color-components-progress-bar-progress: rgb(0 51 255 / 0.14); + --color-components-progress-bar-progress-highlight: rgb(0 51 255 / 0.2); + --color-components-progress-bar-progress-solid: #0033ff; + --color-components-progress-bar-border: rgb(16 24 40 / 0.04); + --color-components-progress-bar-bg: rgb(0 51 255 / 0.04); + + --color-components-button-button-seam: rgb(0 0 0 / 0.03); --color-components-button-primary-text: #ffffff; - --color-components-button-primary-bg: #155aef; - --color-components-button-primary-border: rgb(16 24 40 / 0.04); - --color-components-button-primary-bg-hover: #004aeb; - --color-components-button-primary-border-hover: rgb(16 24 40 / 0.08); - --color-components-button-primary-bg-disabled: rgb(21 90 239 / 0.14); - --color-components-button-primary-border-disabled: rgb(255 255 255 / 0); --color-components-button-primary-text-disabled: rgb(255 255 255 / 0.6); + --color-components-button-primary-bg: #0033ff; + --color-components-button-primary-bg-hover: #002cde; + --color-components-button-primary-bg-disabled: rgb(0 51 255 / 0.14); + --color-components-button-primary-border: rgb(16 24 40 / 0.04); + --color-components-button-primary-border-hover: rgb(16 24 40 / 0.08); + --color-components-button-primary-border-disabled: rgb(255 255 255 / 0); --color-components-button-secondary-text: #354052; --color-components-button-secondary-text-disabled: rgb(16 24 40 / 0.25); @@ -38,6 +372,15 @@ html[data-theme="light"] { --color-components-button-secondary-border-hover: rgb(16 24 40 / 0.2); --color-components-button-secondary-border-disabled: rgb(16 24 40 / 0.04); + --color-components-button-secondary-accent-text: #0033ff; + --color-components-button-secondary-accent-text-disabled: #a0bdff; + --color-components-button-secondary-accent-bg: #ffffff; + --color-components-button-secondary-accent-bg-hover: #f2f4f7; + --color-components-button-secondary-accent-bg-disabled: #f9fafb; + --color-components-button-secondary-accent-border: rgb(16 24 40 / 0.14); + --color-components-button-secondary-accent-border-hover: rgb(16 24 40 / 0.14); + --color-components-button-secondary-accent-border-disabled: rgb(16 24 40 / 0.04); + --color-components-button-tertiary-text: #354052; --color-components-button-tertiary-text-disabled: rgb(16 24 40 / 0.25); --color-components-button-tertiary-bg: #f2f4f7; @@ -76,133 +419,99 @@ html[data-theme="light"] { --color-components-button-destructive-ghost-text-disabled: rgb(240 68 56 / 0.2); --color-components-button-destructive-ghost-bg-hover: #fee4e2; - --color-components-button-secondary-accent-text: #155aef; - --color-components-button-secondary-accent-text-disabled: #b2caff; - --color-components-button-secondary-accent-bg: #ffffff; - --color-components-button-secondary-accent-bg-hover: #f2f4f7; - --color-components-button-secondary-accent-bg-disabled: #f9fafb; - --color-components-button-secondary-accent-border: rgb(16 24 40 / 0.14); - --color-components-button-secondary-accent-border-hover: rgb(16 24 40 / 0.14); - --color-components-button-secondary-accent-border-disabled: rgb(16 24 40 / 0.04); - --color-components-button-indigo-bg: #444ce7; --color-components-button-indigo-bg-hover: #3538cd; --color-components-button-indigo-bg-disabled: rgb(97 114 243 / 0.14); + --color-components-button-debug-text: #ffffff; + --color-components-button-debug-text-disabled: rgb(255 255 255 / 0.6); + --color-components-button-debug-bg: #ff4405; + --color-components-button-debug-bg-hover: #e62e05; + --color-components-button-debug-bg-disabled: rgb(255 68 5 / 0.2); + --color-components-button-debug-border: rgb(16 24 40 / 0.04); + --color-components-button-debug-border-hover: rgb(16 24 40 / 0.08); + --color-components-button-debug-border-disabled: rgb(255 255 255 / 0); + --color-components-checkbox-icon: #ffffff; --color-components-checkbox-icon-disabled: rgb(255 255 255 / 0.5); - --color-components-checkbox-bg: #155aef; - --color-components-checkbox-bg-hover: #004aeb; + --color-components-checkbox-bg: #0033ff; + --color-components-checkbox-bg-hover: #002cde; --color-components-checkbox-bg-disabled: #f2f4f7; + --color-components-checkbox-bg-disabled-checked: #a0bdff; + --color-components-checkbox-bg-unchecked: #ffffff; + --color-components-checkbox-bg-unchecked-hover: #ffffff; --color-components-checkbox-border: #d0d5dc; --color-components-checkbox-border-hover: #98a2b2; --color-components-checkbox-border-disabled: rgb(24 24 27 / 0.04); - --color-components-checkbox-bg-unchecked: #ffffff; - --color-components-checkbox-bg-unchecked-hover: #ffffff; - --color-components-checkbox-bg-disabled-checked: #b2caff; - --color-components-radio-border-checked: #155aef; - --color-components-radio-border-checked-hover: #004aeb; - --color-components-radio-border-checked-disabled: #b2caff; + --color-components-radio-bg: rgb(255 255 255 / 0); + --color-components-radio-bg-hover: rgb(255 255 255 / 0); --color-components-radio-bg-disabled: rgb(255 255 255 / 0); + --color-components-radio-border-checked: #0033ff; + --color-components-radio-border-checked-hover: #002cde; + --color-components-radio-border-checked-disabled: #a0bdff; --color-components-radio-border: #d0d5dc; --color-components-radio-border-hover: #98a2b2; --color-components-radio-border-disabled: rgb(24 24 27 / 0.04); - --color-components-radio-bg: rgb(255 255 255 / 0); - --color-components-radio-bg-hover: rgb(255 255 255 / 0); --color-components-toggle-knob: #ffffff; + --color-components-toggle-knob-hover: #ffffff; --color-components-toggle-knob-disabled: rgb(255 255 255 / 0.95); - --color-components-toggle-bg: #155aef; - --color-components-toggle-bg-hover: #004aeb; - --color-components-toggle-bg-disabled: #d1e0ff; + --color-components-toggle-bg: #0033ff; + --color-components-toggle-bg-hover: #002cde; + --color-components-toggle-bg-disabled: #d0dfff; --color-components-toggle-bg-unchecked: #e9ebf0; --color-components-toggle-bg-unchecked-hover: #d0d5dc; --color-components-toggle-bg-unchecked-disabled: #f2f4f7; - --color-components-toggle-knob-hover: #ffffff; - - --color-components-card-bg: #fcfcfd; - --color-components-card-border: #ffffff; - --color-components-card-bg-alt: #ffffff; - --color-components-card-bg-transparent: rgb(252 252 253 / 0); - --color-components-card-bg-alt-transparent: rgb(255 255 255 / 0); - - --color-components-menu-item-text: #495464; - --color-components-menu-item-text-active: #18222f; - --color-components-menu-item-text-hover: #354052; - --color-components-menu-item-text-active-accent: #18222f; - --color-components-menu-item-bg-active: rgb(21 90 239 / 0.08); - --color-components-menu-item-bg-hover: rgb(200 206 218 / 0.2); - - --color-components-panel-bg: #ffffff; - --color-components-panel-bg-blur: rgb(255 255 255 / 0.95); - --color-components-panel-border: rgb(16 24 40 / 0.08); - --color-components-panel-border-subtle: rgb(16 24 40 / 0.08); - --color-components-panel-gradient-2: #f9fafb; - --color-components-panel-gradient-1: #ffffff; - --color-components-panel-bg-alt: #f9fafb; - --color-components-panel-on-panel-item-bg: #ffffff; - --color-components-panel-on-panel-item-bg-hover: #f9fafb; - --color-components-panel-on-panel-item-bg-alt: #f9fafb; - --color-components-panel-on-panel-item-bg-transparent: rgb(255 255 255 / 0.95); - --color-components-panel-on-panel-item-bg-hover-transparent: rgb(249 250 251 / 0); - --color-components-panel-on-panel-item-bg-destructive-hover-transparent: rgb(254 243 242 / 0); - - --color-components-panel-bg-transparent: rgb(255 255 255 / 0); - - --color-components-main-nav-nav-button-text: #495464; - --color-components-main-nav-nav-button-text-active: #155aef; - --color-components-main-nav-nav-button-bg: rgb(255 255 255 / 0); - --color-components-main-nav-nav-button-bg-active: #fcfcfd; - --color-components-main-nav-nav-button-border: rgb(255 255 255 / 0.95); - --color-components-main-nav-nav-button-bg-hover: rgb(16 24 40 / 0.04); - --color-components-main-nav-glass-text-glow: #3146ff2e; - --color-components-main-nav-glass-surface-first: #0033ff14; - --color-components-main-nav-glass-surface-middle-1: #0033ff1f; - --color-components-main-nav-glass-surface-middle-2: #0033ff1a; - --color-components-main-nav-glass-surface-end: #0033ff14; - --color-components-main-nav-glass-edge-highlight-first: #fffffffa; - --color-components-main-nav-glass-edge-highlight-middle: #ffffff00; - --color-components-main-nav-glass-edge-highlight-end: #ffffff6b; - --color-components-main-nav-glass-edge-reflection-first: #0033ff00; - --color-components-main-nav-glass-edge-reflection-middle: #0033ff99; - --color-components-main-nav-glass-edge-reflection-end: #0033ff00; - --color-components-main-nav-glass-inner-glow: #ffffff4d; - --color-components-main-nav-glass-shadow-reflection: #0033ff0a; - --color-components-main-nav-glass-shadow-reflection-glow: #ffffff00; - - --color-components-main-nav-nav-user-border: #ffffff; - - --color-components-slider-knob: #ffffff; - --color-components-slider-knob-hover: #ffffff; - --color-components-slider-knob-disabled: rgb(255 255 255 / 0.95); - --color-components-slider-range: #296dff; - --color-components-slider-track: #e9ebf0; - --color-components-slider-knob-border-hover: rgb(16 24 40 / 0.2); - --color-components-slider-knob-border: rgb(16 24 40 / 0.14); - - --color-components-segmented-control-item-active-bg: #ffffff; - --color-components-segmented-control-item-active-border: #ffffff; - --color-components-segmented-control-bg-normal: rgb(200 206 218 / 0.2); - --color-components-segmented-control-item-active-accent-bg: #ffffff; - --color-components-segmented-control-item-active-accent-border: #ffffff; --color-components-option-card-option-bg: #fcfcfd; - --color-components-option-card-option-selected-bg: #ffffff; - --color-components-option-card-option-selected-border: #296dff; --color-components-option-card-option-border: #e9ebf0; --color-components-option-card-option-bg-hover: #ffffff; --color-components-option-card-option-border-hover: #d0d5dc; + --color-components-option-card-option-selected-bg: #ffffff; + --color-components-option-card-option-selected-border: #0033ff; - --color-components-tab-active: #155aef; + --color-components-slider-knob: #ffffff; + --color-components-slider-knob-border: rgb(16 24 40 / 0.14); + --color-components-slider-knob-border-hover: rgb(16 24 40 / 0.2); + --color-components-slider-knob-hover: #ffffff; + --color-components-slider-knob-disabled: rgb(255 255 255 / 0.95); + --color-components-slider-range: #0033ff; + --color-components-slider-track: #e9ebf0; + + --color-components-tooltip-bg: rgb(255 255 255 / 0.95); + + --color-components-kbd-bg-gray: rgb(16 24 40 / 0.04); + --color-components-kbd-bg-white: rgb(255 255 255 / 0.12); + + --color-components-menu-item-text: #495464; + --color-components-menu-item-text-hover: #354052; + --color-components-menu-item-text-active: #18222f; + --color-components-menu-item-text-active-accent: #18222f; + --color-components-menu-item-bg-hover: rgb(200 206 218 / 0.2); + --color-components-menu-item-bg-active: rgb(0 51 255 / 0.08); + + --color-components-segmented-control-bg-normal: rgb(200 206 218 / 0.2); + --color-components-segmented-control-item-active-bg: #ffffff; + --color-components-segmented-control-item-active-border: #ffffff; + --color-components-segmented-control-item-active-accent-bg: #ffffff; + --color-components-segmented-control-item-active-accent-border: #ffffff; --color-components-badge-white-to-dark: #ffffff; + --color-components-badge-white-to-dark-alpha: rgb(255 255 255 / 0); + --color-components-badge-bg-green-soft: rgb(23 178 106 / 0.08); + --color-components-badge-bg-orange-soft: rgb(247 144 9 / 0.08); + --color-components-badge-bg-red-soft: rgb(240 68 56 / 0.08); + --color-components-badge-bg-blue-light-soft: rgb(11 165 236 / 0.08); + --color-components-badge-bg-gray-soft: rgb(16 24 40 / 0.04); + --color-components-badge-bg-dimm: rgb(255 255 255 / 0.05); + + --color-components-badge-status-light-border-outer: #ffffff; + --color-components-badge-status-light-high-light: rgb(255 255 255 / 0.3); --color-components-badge-status-light-success-bg: #47cd89; --color-components-badge-status-light-success-border-inner: #17b26a; --color-components-badge-status-light-success-halo: rgb(23 178 106 / 0.25); - --color-components-badge-status-light-border-outer: #ffffff; - --color-components-badge-status-light-high-light: rgb(255 255 255 / 0.3); --color-components-badge-status-light-warning-bg: #fdb022; --color-components-badge-status-light-warning-border-inner: #f79009; --color-components-badge-status-light-warning-halo: rgb(247 144 9 / 0.25); @@ -219,221 +528,60 @@ html[data-theme="light"] { --color-components-badge-status-light-disabled-border-inner: #676f83; --color-components-badge-status-light-disabled-halo: rgb(16 24 40 / 0.04); - --color-components-badge-bg-green-soft: rgb(23 178 106 / 0.08); - --color-components-badge-bg-orange-soft: rgb(247 144 9 / 0.08); - --color-components-badge-bg-red-soft: rgb(240 68 56 / 0.08); - --color-components-badge-bg-blue-light-soft: rgb(11 165 236 / 0.08); - --color-components-badge-bg-gray-soft: rgb(16 24 40 / 0.04); - --color-components-badge-bg-dimm: rgb(255 255 255 / 0.05); + --color-components-tab-active: #0033ff; + + --color-components-main-nav-text: #495464; + --color-components-main-nav-text-active: #0033ff; + --color-components-main-nav-nav-button-border: rgb(255 255 255 / 0.95); + --color-components-main-nav-nav-button-text: #495464; + --color-components-main-nav-nav-button-text-active: #0033ff; + --color-components-main-nav-nav-button-bg: rgb(255 255 255 / 0); + --color-components-main-nav-nav-button-bg-hover: rgb(16 24 40 / 0.04); + --color-components-main-nav-nav-button-bg-active: #fcfcfd; + + --color-components-main-nav-nav-user-border: #ffffff; + + --color-components-main-nav-glass-inner-glow: rgb(255 255 255 / 0.3); + --color-components-main-nav-glass-shadow-reflection: rgb(0 51 255 / 0.04); + --color-components-main-nav-glass-shadow-reflection-glow: rgb(255 255 255 / 0); + --color-components-main-nav-glass-text-glow: rgb(49 70 255 / 0.18); + --color-components-main-nav-glass-surface-first: rgb(0 51 255 / 0.08); + --color-components-main-nav-glass-surface-middle-1: rgb(0 51 255 / 0.12); + --color-components-main-nav-glass-surface-middle-2: rgb(0 51 255 / 0.1); + --color-components-main-nav-glass-surface-end: rgb(0 51 255 / 0.08); + + --color-components-main-nav-glass-edge-reflection-first: rgb(0 51 255 / 0); + --color-components-main-nav-glass-edge-reflection-middle: rgb(0 51 255 / 0.6); + --color-components-main-nav-glass-edge-reflection-end: rgb(0 51 255 / 0); + + --color-components-main-nav-glass-edge-highlight-first: rgb(255 255 255 / 0.98); + --color-components-main-nav-glass-edge-highlight-middle: rgb(255 255 255 / 0); + --color-components-main-nav-glass-edge-highlight-end: rgb(255 255 255 / 0.42); - --color-components-chart-line: #296dff; - --color-components-chart-area-1: rgb(21 90 239 / 0.14); - --color-components-chart-area-2: rgb(21 90 239 / 0.04); - --color-components-chart-current-1: #155aef; - --color-components-chart-current-2: #d1e0ff; --color-components-chart-bg: #ffffff; + --color-components-chart-line: #0033ff; + --color-components-chart-current-1: #0033ff; + --color-components-chart-current-2: #d0dfff; + --color-components-chart-area-1: rgb(0 51 255 / 0.14); + --color-components-chart-area-2: rgb(0 51 255 / 0.04); - --color-components-actionbar-bg: rgb(255 255 255 / 0.95); - --color-components-actionbar-border: rgb(16 24 40 / 0.04); - --color-components-actionbar-bg-accent: #f5f7ff; - --color-components-actionbar-border-accent: #b2caff; + --color-divider-subtle: rgb(16 24 40 / 0.04); + --color-divider-regular: rgb(16 24 40 / 0.08); + --color-divider-deep: rgb(16 24 40 / 0.14); + --color-divider-intense: rgb(16 24 40 / 0.3); + --color-divider-burn: rgb(16 24 40 / 0.04); + --color-divider-solid: #d0d5dc; + --color-divider-solid-alt: #98a2b2; + --color-divider-accent: #e5eaff; - --color-components-dropzone-bg-alt: #f2f4f7; - --color-components-dropzone-bg: #f9fafb; - --color-components-dropzone-bg-accent: rgb(21 90 239 / 0.14); - --color-components-dropzone-border: rgb(16 24 40 / 0.08); - --color-components-dropzone-border-alt: rgb(16 24 40 / 0.2); - --color-components-dropzone-border-accent: #84abff; - - --color-components-progress-brand-progress: #296dff; - --color-components-progress-brand-border: #296dff; - --color-components-progress-brand-bg: rgb(21 90 239 / 0.04); - - --color-components-progress-white-progress: #ffffff; - --color-components-progress-white-border: rgb(255 255 255 / 0.95); - --color-components-progress-white-bg: rgb(255 255 255 / 0.01); - - --color-components-progress-gray-progress: #98a2b2; - --color-components-progress-gray-border: #98a2b2; - --color-components-progress-gray-bg: rgb(200 206 218 / 0.02); - - --color-components-progress-warning-progress: #f79009; - --color-components-progress-warning-border: #f79009; - --color-components-progress-warning-bg: rgb(247 144 9 / 0.04); - - --color-components-progress-error-progress: #f04438; - --color-components-progress-error-border: #f04438; - --color-components-progress-error-bg: rgb(240 68 56 / 0.04); - - --color-components-chat-input-audio-bg: #eff4ff; - --color-components-chat-input-audio-wave-default: rgb(21 90 239 / 0.2); - --color-components-chat-input-bg-mask-1: rgb(255 255 255 / 0.01); - --color-components-chat-input-bg-mask-2: #f2f4f7; - --color-components-chat-input-border: #ffffff; - --color-components-chat-input-audio-wave-active: #296dff; - --color-components-chat-input-audio-bg-alt: #fcfcfd; - - --color-components-avatar-shape-fill-stop-0: #ffffff; - --color-components-avatar-shape-fill-stop-100: rgb(255 255 255 / 0.9); - - --color-components-avatar-bg-mask-stop-0: rgb(255 255 255 / 0.12); - --color-components-avatar-bg-mask-stop-100: rgb(255 255 255 / 0.08); - - --color-components-avatar-default-avatar-bg: #d0d5dc; - --color-components-avatar-mask-darkmode-dimmed: rgb(255 255 255 / 0); - - --color-components-label-gray: #f2f4f7; - - --color-components-premium-badge-blue-bg-stop-0: #5289ff; - --color-components-premium-badge-blue-bg-stop-100: #155aef; - --color-components-premium-badge-blue-stroke-stop-0: rgb(255 255 255 / 0.95); - --color-components-premium-badge-blue-stroke-stop-100: #155aef; - --color-components-premium-badge-blue-text-stop-0: #f5f7ff; - --color-components-premium-badge-blue-text-stop-100: #d1e0ff; - --color-components-premium-badge-blue-glow: #00329e; - --color-components-premium-badge-blue-bg-stop-0-hover: #296dff; - --color-components-premium-badge-blue-bg-stop-100-hover: #004aeb; - --color-components-premium-badge-blue-glow-hover: #84abff; - --color-components-premium-badge-blue-stroke-stop-0-hover: rgb(255 255 255 / 0.95); - --color-components-premium-badge-blue-stroke-stop-100-hover: #00329e; - - --color-components-premium-badge-highlight-stop-0: rgb(255 255 255 / 0.12); - --color-components-premium-badge-highlight-stop-100: rgb(255 255 255 / 0.3); - --color-components-premium-badge-indigo-bg-stop-0: #8098f9; - --color-components-premium-badge-indigo-bg-stop-100: #444ce7; - --color-components-premium-badge-indigo-stroke-stop-0: rgb(255 255 255 / 0.95); - --color-components-premium-badge-indigo-stroke-stop-100: #6172f3; - --color-components-premium-badge-indigo-text-stop-0: #f5f8ff; - --color-components-premium-badge-indigo-text-stop-100: #e0eaff; - --color-components-premium-badge-indigo-glow: #2d3282; - --color-components-premium-badge-indigo-glow-hover: #a4bcfd; - --color-components-premium-badge-indigo-bg-stop-0-hover: #6172f3; - --color-components-premium-badge-indigo-bg-stop-100-hover: #2d31a6; - --color-components-premium-badge-indigo-stroke-stop-0-hover: rgb(255 255 255 / 0.95); - --color-components-premium-badge-indigo-stroke-stop-100-hover: #2d31a6; - - --color-components-premium-badge-grey-bg-stop-0: #98a2b2; - --color-components-premium-badge-grey-bg-stop-100: #676f83; - --color-components-premium-badge-grey-stroke-stop-0: rgb(255 255 255 / 0.95); - --color-components-premium-badge-grey-stroke-stop-100: #676f83; - --color-components-premium-badge-grey-text-stop-0: #fcfcfd; - --color-components-premium-badge-grey-text-stop-100: #f2f4f7; - --color-components-premium-badge-grey-glow: #101828; - --color-components-premium-badge-grey-glow-hover: #d0d5dc; - --color-components-premium-badge-grey-bg-stop-0-hover: #676f83; - --color-components-premium-badge-grey-bg-stop-100-hover: #354052; - --color-components-premium-badge-grey-stroke-stop-0-hover: rgb(255 255 255 / 0.95); - --color-components-premium-badge-grey-stroke-stop-100-hover: #354052; - - --color-components-premium-badge-orange-bg-stop-0: #ff692e; - --color-components-premium-badge-orange-bg-stop-100: #e04f16; - --color-components-premium-badge-orange-stroke-stop-0: rgb(255 255 255 / 0.95); - --color-components-premium-badge-orange-stroke-stop-100: #e62e05; - --color-components-premium-badge-orange-text-stop-0: #fefaf5; - --color-components-premium-badge-orange-text-stop-100: #fdead7; - --color-components-premium-badge-orange-glow: #772917; - --color-components-premium-badge-orange-glow-hover: #f7b27a; - --color-components-premium-badge-orange-bg-stop-0-hover: #ff4405; - --color-components-premium-badge-orange-bg-stop-100-hover: #b93815; - --color-components-premium-badge-orange-stroke-stop-0-hover: rgb(255 255 255 / 0.95); - --color-components-premium-badge-orange-stroke-stop-100-hover: #bc1b06; - - --color-components-progress-bar-bg: rgb(21 90 239 / 0.04); - --color-components-progress-bar-progress: rgb(21 90 239 / 0.14); - --color-components-progress-bar-border: rgb(16 24 40 / 0.04); - --color-components-progress-bar-progress-solid: #296dff; - --color-components-progress-bar-progress-highlight: rgb(21 90 239 / 0.2); - - --color-components-icon-bg-red-solid: #d92d20; - --color-components-icon-bg-rose-solid: #e31b54; - --color-components-icon-bg-pink-solid: #dd2590; - --color-components-icon-bg-orange-dark-solid: #ff4405; - --color-components-icon-bg-yellow-solid: #eaaa08; - --color-components-icon-bg-green-solid: #4ca30d; - --color-components-icon-bg-teal-solid: #0e9384; - --color-components-icon-bg-blue-light-solid: #0ba5ec; - --color-components-icon-bg-blue-solid: #155aef; - --color-components-icon-bg-indigo-solid: #444ce7; - --color-components-icon-bg-violet-solid: #7839ee; - --color-components-icon-bg-midnight-solid: #828dad; - --color-components-icon-bg-rose-soft: #fff1f3; - --color-components-icon-bg-pink-soft: #fdf2fa; - --color-components-icon-bg-orange-dark-soft: #fff4ed; - --color-components-icon-bg-yellow-soft: #fefbe8; - --color-components-icon-bg-green-soft: #f3fee7; - --color-components-icon-bg-teal-soft: #f0fdf9; - --color-components-icon-bg-blue-light-soft: #f0f9ff; - --color-components-icon-bg-blue-soft: #eff4ff; - --color-components-icon-bg-indigo-soft: #eef4ff; - --color-components-icon-bg-violet-soft: #f5f3ff; - --color-components-icon-bg-midnight-soft: #f0f2f5; - --color-components-icon-bg-red-soft: #fef3f2; - --color-components-icon-bg-orange-solid: #f79009; - --color-components-icon-bg-orange-soft: #fffaeb; - - --color-text-primary: #101828; - --color-text-secondary: #354052; - --color-text-tertiary: #676f83; - --color-text-quaternary: rgb(16 24 40 / 0.3); - --color-text-destructive: #d92d20; - --color-text-success: #079455; - --color-text-warning: #dc6803; - --color-text-destructive-secondary: #f04438; - --color-text-success-secondary: #17b26a; - --color-text-warning-secondary: #f79009; - --color-text-accent: #155aef; - --color-text-primary-on-surface: #ffffff; - --color-text-placeholder: #98a2b2; - --color-text-disabled: #d0d5dc; - --color-text-accent-secondary: #296dff; - --color-text-accent-light-mode-only: #155aef; - --color-text-text-selected: rgb(21 90 239 / 0.14); - --color-text-secondary-on-surface: rgb(255 255 255 / 0.9); - --color-text-logo-text: #18222f; - --color-text-empty-state-icon: #d0d5dc; - --color-text-inverted: #000000; - --color-text-inverted-dimmed: rgb(0 0 0 / 0.95); - - --color-background-body: #f2f4f7; - --color-background-default-subtle: #fcfcfd; - --color-background-neutral-subtle: #f9fafb; - --color-background-sidenav-bg: rgb(255 255 255 / 0.8); - --color-background-default: #ffffff; - --color-background-soft: #f9fafb; - --color-background-gradient-bg-fill-chat-bg-1: #f9fafb; - --color-background-gradient-bg-fill-chat-bg-2: #f2f4f7; - --color-background-gradient-bg-fill-chat-bubble-bg-1: #ffffff; - --color-background-gradient-bg-fill-chat-bubble-bg-2: rgb(255 255 255 / 0.6); - --color-background-gradient-bg-fill-debug-bg-1: rgb(255 255 255 / 0); - --color-background-gradient-bg-fill-debug-bg-2: rgb(200 206 218 / 0.14); - - --color-background-gradient-mask-gray: rgb(200 206 218 / 0.2); - --color-background-gradient-mask-transparent: rgb(255 255 255 / 0); - --color-background-gradient-mask-input-clear-2: rgb(233 235 240 / 0); - --color-background-gradient-mask-input-clear-1: #e9ebf0; - --color-background-gradient-mask-transparent-dark: rgb(0 0 0 / 0); - --color-background-gradient-mask-side-panel-2: rgb(16 24 40 / 0.3); - --color-background-gradient-mask-side-panel-1: rgb(16 24 40 / 0.02); - - --color-background-default-burn: #e9ebf0; - --color-background-overlay-fullscreen: rgb(249 250 251 / 0.95); - --color-background-default-lighter: rgb(255 255 255 / 0.5); - --color-background-section: #f9fafb; - --color-background-interaction-from-bg-1: rgb(200 206 218 / 0.2); - --color-background-interaction-from-bg-2: rgb(200 206 218 / 0.14); - --color-background-section-burn: #f2f4f7; - --color-background-default-dodge: #ffffff; - --color-background-overlay: rgb(16 24 40 / 0.6); - --color-background-default-dimmed: #e9ebf0; - --color-background-default-hover: #f9fafb; - --color-background-overlay-alt: rgb(16 24 40 / 0.4); - --color-background-surface-white: rgb(255 255 255 / 0.95); - --color-background-overlay-destructive: rgb(240 68 56 / 0.3); - --color-background-overlay-backdrop: rgb(242 244 247 / 0.95); - --color-background-body-transparent: rgb(242 244 247 / 0); - --color-background-section-burn-inverted: #f2f4f7; + --color-effects-highlight: #ffffff; + --color-effects-highlight-subtle: rgb(255 255 255 / 0.5); + --color-effects-highlight-lightmode-off: rgb(255 255 255 / 0); + --color-effects-image-frame: #ffffff; + --color-effects-icon-border: rgb(16 24 40 / 0.08); --color-shadow-shadow-1: rgb(9 9 11 / 0.03); + --color-shadow-shadow-2: rgb(9 9 11 / 0.04); --color-shadow-shadow-3: rgb(9 9 11 / 0.05); --color-shadow-shadow-4: rgb(9 9 11 / 0.06); --color-shadow-shadow-5: rgb(9 9 11 / 0.08); @@ -441,124 +589,18 @@ html[data-theme="light"] { --color-shadow-shadow-7: rgb(9 9 11 / 0.12); --color-shadow-shadow-8: rgb(9 9 11 / 0.14); --color-shadow-shadow-9: rgb(9 9 11 / 0.18); - --color-shadow-shadow-2: rgb(9 9 11 / 0.04); --color-shadow-shadow-10: rgb(9 9 11 / 0.05); - --color-workflow-block-border: #ffffff; - --color-workflow-block-parma-bg: #f2f4f7; - --color-workflow-block-bg: #fcfcfd; - --color-workflow-block-bg-transparent: rgb(252 252 253 / 0.9); - --color-workflow-block-border-highlight: rgb(21 90 239 / 0.14); - --color-workflow-block-wrapper-bg-1: #e9ebf0; - --color-workflow-block-wrapper-bg-2: rgb(233 235 240 / 0.2); - - --color-workflow-canvas-workflow-dot-color: rgb(133 133 173 / 0.15); - --color-workflow-canvas-workflow-bg: #f2f4f7; - --color-workflow-canvas-workflow-top-bar-1: rgb(242 244 247 / 0.9); - --color-workflow-canvas-workflow-top-bar-2: rgb(242 244 247 / 0.05); - --color-workflow-canvas-canvas-overlay: rgb(242 244 247 / 0.8); - - --color-workflow-link-line-active: #296dff; - --color-workflow-link-line-normal: #d0d5dc; - --color-workflow-link-line-handle: #296dff; - --color-workflow-link-line-normal-transparent: rgb(208 213 220 / 0.2); - --color-workflow-link-line-failure-active: #f79009; - --color-workflow-link-line-failure-handle: #f79009; - --color-workflow-link-line-failure-button-bg: #dc6803; - --color-workflow-link-line-failure-button-hover: #b54708; - - --color-workflow-link-line-success-active: #17b26a; - --color-workflow-link-line-success-handle: #17b26a; - - --color-workflow-link-line-error-active: #f04438; - --color-workflow-link-line-error-handle: #f04438; - - --color-workflow-minimap-bg: #e9ebf0; - --color-workflow-minimap-block: rgb(200 206 218 / 0.3); - - --color-workflow-display-success-bg: #ecfdf3; - --color-workflow-display-success-border-1: rgb(23 178 106 / 0.8); - --color-workflow-display-success-border-2: rgb(23 178 106 / 0.5); - --color-workflow-display-success-vignette-color: rgb(23 178 106 / 0.2); - --color-workflow-display-success-bg-line-pattern: rgb(23 178 106 / 0.3); - - --color-workflow-display-glass-1: rgb(255 255 255 / 0.12); - --color-workflow-display-glass-2: rgb(255 255 255 / 0.5); - --color-workflow-display-vignette-dark: rgb(0 0 0 / 0.12); - --color-workflow-display-highlight: rgb(255 255 255 / 0.5); - --color-workflow-display-outline: rgb(0 0 0 / 0.05); - --color-workflow-display-error-bg: #fef3f2; - --color-workflow-display-error-bg-line-pattern: rgb(240 68 56 / 0.3); - --color-workflow-display-error-border-1: rgb(240 68 56 / 0.8); - --color-workflow-display-error-border-2: rgb(240 68 56 / 0.5); - --color-workflow-display-error-vignette-color: rgb(240 68 56 / 0.2); - - --color-workflow-display-warning-bg: #fffaeb; - --color-workflow-display-warning-bg-line-pattern: rgb(247 144 9 / 0.3); - --color-workflow-display-warning-border-1: rgb(247 144 9 / 0.8); - --color-workflow-display-warning-border-2: rgb(247 144 9 / 0.5); - --color-workflow-display-warning-vignette-color: rgb(247 144 9 / 0.2); - - --color-workflow-display-normal-bg: #f0f9ff; - --color-workflow-display-normal-bg-line-pattern: rgb(11 165 236 / 0.3); - --color-workflow-display-normal-border-1: rgb(11 165 236 / 0.8); - --color-workflow-display-normal-border-2: rgb(11 165 236 / 0.5); - --color-workflow-display-normal-vignette-color: rgb(11 165 236 / 0.2); - - --color-workflow-display-disabled-bg: #f9fafb; - --color-workflow-display-disabled-bg-line-pattern: rgb(200 206 218 / 0.3); - --color-workflow-display-disabled-border-1: rgb(200 206 218 / 0.6); - --color-workflow-display-disabled-border-2: rgb(200 206 218 / 0.4); - --color-workflow-display-disabled-vignette-color: rgb(200 206 218 / 0.4); - --color-workflow-display-disabled-outline: rgb(0 0 0 / 0); - - --color-workflow-workflow-progress-bg-1: rgb(200 206 218 / 0.2); - --color-workflow-workflow-progress-bg-2: rgb(200 206 218 / 0.04); - - --color-divider-subtle: rgb(16 24 40 / 0.04); - --color-divider-regular: rgb(16 24 40 / 0.08); - --color-divider-deep: rgb(16 24 40 / 0.14); - --color-divider-burn: rgb(16 24 40 / 0.04); - --color-divider-intense: rgb(16 24 40 / 0.3); - --color-divider-solid: #d0d5dc; - --color-divider-solid-alt: #98a2b2; - --color-divider-accent: #e5eaff; - - --color-state-base-hover: rgb(200 206 218 / 0.2); - --color-state-base-active: rgb(200 206 218 / 0.4); - --color-state-base-hover-alt: rgb(200 206 218 / 0.4); - --color-state-base-handle: rgb(16 24 40 / 0.2); - --color-state-base-handle-hover: rgb(16 24 40 / 0.3); - --color-state-base-hover-subtle: rgb(200 206 218 / 0.08); - - --color-state-accent-hover: #eff4ff; - --color-state-accent-active: rgb(21 90 239 / 0.08); - --color-state-accent-hover-alt: #d1e0ff; - --color-state-accent-solid: #296dff; - --color-state-accent-active-alt: rgb(21 90 239 / 0.14); - - --color-state-destructive-hover: #fef3f2; - --color-state-destructive-hover-alt: #fee4e2; - --color-state-destructive-active: #fecdca; - --color-state-destructive-solid: #f04438; - --color-state-destructive-border: #fda29b; - --color-state-destructive-hover-transparent: rgb(254 243 242 / 0); - - --color-state-success-hover: #ecfdf3; - --color-state-success-hover-alt: #dcfae6; - --color-state-success-active: #abefc6; - --color-state-success-solid: #17b26a; - - --color-state-warning-hover: #fffaeb; - --color-state-warning-hover-alt: #fef0c7; - --color-state-warning-active: #fedf89; - --color-state-warning-solid: #f79009; - --color-state-warning-hover-transparent: rgb(255 250 235 / 0); - - --color-effects-highlight: #ffffff; - --color-effects-highlight-lightmode-off: rgb(255 255 255 / 0); - --color-effects-image-frame: #ffffff; - --color-effects-icon-border: rgb(16 24 40 / 0.08); + --color-third-party-LangChain: #1c3c3c; + --color-third-party-Langfuse: #000000; + --color-third-party-Github: #1b1f24; + --color-third-party-Github-tertiary: #1b1f24; + --color-third-party-Github-secondary: #1b1f24; + --color-third-party-aws: #141f2e; + --color-third-party-aws-alt: #0f1824; + --color-third-party-model-bg-openai: #e3e5e8; + --color-third-party-model-bg-anthropic: #eeede7; + --color-third-party-model-bg-default: #f9fafb; --color-util-colors-orange-dark-orange-dark-50: #fff4ed; --color-util-colors-orange-dark-orange-dark-100: #ffe6d5; @@ -633,6 +675,15 @@ html[data-theme="light"] { --color-util-colors-blue-light-blue-light-600: #0086c9; --color-util-colors-blue-light-blue-light-700: #026aa2; + --color-util-colors-blue-brand-blue-brand-50: #f5f8ff; + --color-util-colors-blue-brand-blue-brand-100: #d0dfff; + --color-util-colors-blue-brand-blue-brand-200: #a0bdff; + --color-util-colors-blue-brand-blue-brand-300: #6694ff; + --color-util-colors-blue-brand-blue-brand-400: #3072ff; + --color-util-colors-blue-brand-blue-brand-500: #085afc; + --color-util-colors-blue-brand-blue-brand-600: #0033ff; + --color-util-colors-blue-brand-blue-brand-700: #002cde; + --color-util-colors-gray-blue-gray-blue-50: #f8f9fc; --color-util-colors-gray-blue-gray-blue-100: #eaecf5; --color-util-colors-gray-blue-gray-blue-200: #d5d9eb; @@ -642,15 +693,6 @@ html[data-theme="light"] { --color-util-colors-gray-blue-gray-blue-600: #3e4784; --color-util-colors-gray-blue-gray-blue-700: #363f72; - --color-util-colors-blue-brand-blue-brand-50: #f5f7ff; - --color-util-colors-blue-brand-blue-brand-100: #d1e0ff; - --color-util-colors-blue-brand-blue-brand-200: #b2caff; - --color-util-colors-blue-brand-blue-brand-300: #84abff; - --color-util-colors-blue-brand-blue-brand-400: #5289ff; - --color-util-colors-blue-brand-blue-brand-500: #296dff; - --color-util-colors-blue-brand-blue-brand-600: #155aef; - --color-util-colors-blue-brand-blue-brand-700: #004aeb; - --color-util-colors-red-red-50: #fef3f2; --color-util-colors-red-red-100: #fee4e2; --color-util-colors-red-red-200: #fecdca; @@ -669,6 +711,15 @@ html[data-theme="light"] { --color-util-colors-green-green-600: #079455; --color-util-colors-green-green-700: #067647; + --color-util-colors-green-light-green-light-50: #f3fee7; + --color-util-colors-green-light-green-light-100: #e3fbcc; + --color-util-colors-green-light-green-light-200: #d0f8ab; + --color-util-colors-green-light-green-light-300: #a6ef67; + --color-util-colors-green-light-green-light-400: #85e13a; + --color-util-colors-green-light-green-light-500: #66c61c; + --color-util-colors-green-light-green-light-600: #4ca30d; + --color-util-colors-green-light-green-light-700: #3b7c0f; + --color-util-colors-warning-warning-50: #fffaeb; --color-util-colors-warning-warning-100: #fef0c7; --color-util-colors-warning-warning-200: #fedf89; @@ -723,15 +774,6 @@ html[data-theme="light"] { --color-util-colors-gray-gray-600: #495464; --color-util-colors-gray-gray-700: #354052; - --color-util-colors-green-light-green-light-50: #f3fee7; - --color-util-colors-green-light-green-light-100: #e3fbcc; - --color-util-colors-green-light-green-light-200: #d0f8ab; - --color-util-colors-green-light-green-light-300: #a6ef67; - --color-util-colors-green-light-green-light-500: #66c61c; - --color-util-colors-green-light-green-light-400: #85e13a; - --color-util-colors-green-light-green-light-600: #4ca30d; - --color-util-colors-green-light-green-light-700: #3b7c0f; - --color-util-colors-rose-rose-50: #fff1f3; --color-util-colors-rose-rose-100: #ffe4e8; --color-util-colors-rose-rose-200: #fecdd6; @@ -750,30 +792,31 @@ html[data-theme="light"] { --color-util-colors-midnight-midnight-600: #5d698d; --color-util-colors-midnight-midnight-700: #3e465e; - --color-third-party-LangChain: #1c3c3c; - --color-third-party-Langfuse: #000000; - --color-third-party-Github: #1b1f24; - --color-third-party-Github-tertiary: #1b1f24; - --color-third-party-Github-secondary: #1b1f24; - --color-third-party-model-bg-openai: #e3e5e8; - --color-third-party-model-bg-anthropic: #eeede7; - --color-third-party-model-bg-default: #f9fafb; - - --color-third-party-aws: #141f2e; - --color-third-party-aws-alt: #0f1824; - --color-saas-background: #ffffff; + --color-saas-background-inverted: #0b0b0e; + --color-saas-background-inverted-hover: #222225; --color-saas-pricing-grid-bg: rgb(200 206 218 / 0.5); --color-saas-dify-blue-static: #0033ff; - --color-saas-dify-blue-static-hover: #002cd6; + --color-saas-dify-blue-static-hover: #002cde; --color-saas-dify-blue-accessible: #0033ff; --color-saas-dify-blue-inverted: #0033ff; --color-saas-dify-blue-inverted-dimmed: #0033ff; - --color-saas-background-inverted: #0b0b0e; - --color-saas-background-inverted-hover: #222225; - --color-dify-logo-blue: #0033ff; --color-dify-logo-black: #000000; + --color-dify-logo-outline-1: rgb(0 0 0 / 0); + --color-dify-logo-outline-2: rgb(0 0 0 / 0); + + --color-brand-color-opacity-50: #f3f5fd; + --color-brand-color-opacity-100: #e2e7fb; + --color-brand-color-opacity-200: #bac6f5; + --color-brand-color-opacity-300: #899eee; + --color-brand-color-opacity-400: #4f6de5; + --color-brand-color-opacity-500: #0033ff; + --color-brand-color-opacity-600: #082bb4; + --color-brand-color-opacity-700: #07228e; + --color-brand-color-opacity-800: #051969; + --color-brand-color-opacity-900: #031042; + --color-brand-color-opacity-1000: #020821; } diff --git a/packages/dify-ui/src/themes/theme.css b/packages/dify-ui/src/themes/theme.css index c14e54ea549..c6ed246f9b1 100644 --- a/packages/dify-ui/src/themes/theme.css +++ b/packages/dify-ui/src/themes/theme.css @@ -7,34 +7,368 @@ */ @theme inline { + --color-text-primary: var(--color-text-primary); + --color-text-secondary: var(--color-text-secondary); + --color-text-tertiary: var(--color-text-tertiary); + --color-text-quaternary: var(--color-text-quaternary); + --color-text-empty-state-icon: var(--color-text-empty-state-icon); + --color-text-primary-on-surface: var(--color-text-primary-on-surface); + --color-text-secondary-on-surface: var(--color-text-secondary-on-surface); + --color-text-destructive: var(--color-text-destructive); + --color-text-destructive-secondary: var(--color-text-destructive-secondary); + --color-text-success: var(--color-text-success); + --color-text-success-secondary: var(--color-text-success-secondary); + --color-text-warning: var(--color-text-warning); + --color-text-warning-secondary: var(--color-text-warning-secondary); + --color-text-accent: var(--color-text-accent); + --color-text-accent-secondary: var(--color-text-accent-secondary); + --color-text-accent-light-mode-only: var(--color-text-accent-light-mode-only); + --color-text-placeholder: var(--color-text-placeholder); + --color-text-disabled: var(--color-text-disabled); + --color-text-text-selected: var(--color-text-text-selected); + --color-text-logo-text: var(--color-text-logo-text); + --color-text-inverted: var(--color-text-inverted); + --color-text-inverted-dimmed: var(--color-text-inverted-dimmed); + + --color-background-body: var(--color-background-body); + --color-background-body-transparent: var(--color-background-body-transparent); + --color-background-default: var(--color-background-default); + --color-background-default-hover: var(--color-background-default-hover); + --color-background-default-hover-alpha-0: var(--color-background-default-hover-alpha-0); + --color-background-default-subtle: var(--color-background-default-subtle); + --color-background-default-dodge: var(--color-background-default-dodge); + --color-background-default-burn: var(--color-background-default-burn); + --color-background-default-dimmed: var(--color-background-default-dimmed); + --color-background-default-lighter: var(--color-background-default-lighter); + --color-background-section: var(--color-background-section); + --color-background-section-burn: var(--color-background-section-burn); + --color-background-section-burn-inverted: var(--color-background-section-burn-inverted); + --color-background-soft: var(--color-background-soft); + --color-background-neutral-subtle: var(--color-background-neutral-subtle); + --color-background-surface-white: var(--color-background-surface-white); + --color-background-sidenav-bg: var(--color-background-sidenav-bg); + --color-background-overlay-fullscreen: var(--color-background-overlay-fullscreen); + --color-background-overlay-backdrop: var(--color-background-overlay-backdrop); + --color-background-overlay: var(--color-background-overlay); + --color-background-overlay-alt: var(--color-background-overlay-alt); + --color-background-overlay-destructive: var(--color-background-overlay-destructive); + --color-background-interaction-from-bg-1: var(--color-background-interaction-from-bg-1); + --color-background-interaction-from-bg-2: var(--color-background-interaction-from-bg-2); + --color-background-gradient-bg-fill-chat-bg-1: var(--color-background-gradient-bg-fill-chat-bg-1); + --color-background-gradient-bg-fill-chat-bg-2: var(--color-background-gradient-bg-fill-chat-bg-2); + --color-background-gradient-bg-fill-chat-bubble-bg-1: var(--color-background-gradient-bg-fill-chat-bubble-bg-1); + --color-background-gradient-bg-fill-chat-bubble-bg-2: var(--color-background-gradient-bg-fill-chat-bubble-bg-2); + --color-background-gradient-bg-fill-debug-bg-1: var(--color-background-gradient-bg-fill-debug-bg-1); + --color-background-gradient-bg-fill-debug-bg-2: var(--color-background-gradient-bg-fill-debug-bg-2); + + --color-background-gradient-mask-gray: var(--color-background-gradient-mask-gray); + --color-background-gradient-mask-transparent: var(--color-background-gradient-mask-transparent); + --color-background-gradient-mask-transparent-dark: var(--color-background-gradient-mask-transparent-dark); + --color-background-gradient-mask-side-panel-1: var(--color-background-gradient-mask-side-panel-1); + --color-background-gradient-mask-side-panel-2: var(--color-background-gradient-mask-side-panel-2); + --color-background-gradient-mask-input-clear-1: var(--color-background-gradient-mask-input-clear-1); + --color-background-gradient-mask-input-clear-2: var(--color-background-gradient-mask-input-clear-2); + + --color-state-base-hover-subtle: var(--color-state-base-hover-subtle); + --color-state-base-hover: var(--color-state-base-hover); + --color-state-base-hover-alt: var(--color-state-base-hover-alt); + --color-state-base-active: var(--color-state-base-active); + --color-state-base-handle: var(--color-state-base-handle); + --color-state-base-handle-hover: var(--color-state-base-handle-hover); + + --color-state-accent-hover: var(--color-state-accent-hover); + --color-state-accent-hover-alt: var(--color-state-accent-hover-alt); + --color-state-accent-active: var(--color-state-accent-active); + --color-state-accent-active-alt: var(--color-state-accent-active-alt); + --color-state-accent-solid: var(--color-state-accent-solid); + + --color-state-destructive-hover: var(--color-state-destructive-hover); + --color-state-destructive-hover-transparent: var(--color-state-destructive-hover-transparent); + --color-state-destructive-hover-alt: var(--color-state-destructive-hover-alt); + --color-state-destructive-active: var(--color-state-destructive-active); + --color-state-destructive-solid: var(--color-state-destructive-solid); + --color-state-destructive-border: var(--color-state-destructive-border); + + --color-state-warning-hover: var(--color-state-warning-hover); + --color-state-warning-hover-transparent: var(--color-state-warning-hover-transparent); + --color-state-warning-hover-alt: var(--color-state-warning-hover-alt); + --color-state-warning-active: var(--color-state-warning-active); + --color-state-warning-solid: var(--color-state-warning-solid); + + --color-state-success-hover: var(--color-state-success-hover); + --color-state-success-hover-alt: var(--color-state-success-hover-alt); + --color-state-success-active: var(--color-state-success-active); + --color-state-success-solid: var(--color-state-success-solid); + + --color-workflow-workflow-progress-bg-1: var(--color-workflow-workflow-progress-bg-1); + --color-workflow-workflow-progress-bg-2: var(--color-workflow-workflow-progress-bg-2); + + --color-workflow-block-bg: var(--color-workflow-block-bg); + --color-workflow-block-bg-transparent: var(--color-workflow-block-bg-transparent); + --color-workflow-block-border: var(--color-workflow-block-border); + --color-workflow-block-border-highlight: var(--color-workflow-block-border-highlight); + --color-workflow-block-parma-bg: var(--color-workflow-block-parma-bg); + --color-workflow-block-wrapper-bg-1: var(--color-workflow-block-wrapper-bg-1); + --color-workflow-block-wrapper-bg-2: var(--color-workflow-block-wrapper-bg-2); + + --color-workflow-link-line-normal: var(--color-workflow-link-line-normal); + --color-workflow-link-line-normal-transparent: var(--color-workflow-link-line-normal-transparent); + --color-workflow-link-line-active: var(--color-workflow-link-line-active); + --color-workflow-link-line-handle: var(--color-workflow-link-line-handle); + --color-workflow-link-line-failure-active: var(--color-workflow-link-line-failure-active); + --color-workflow-link-line-failure-handle: var(--color-workflow-link-line-failure-handle); + --color-workflow-link-line-failure-button-bg: var(--color-workflow-link-line-failure-button-bg); + --color-workflow-link-line-failure-button-hover: var(--color-workflow-link-line-failure-button-hover); + + --color-workflow-link-line-success-active: var(--color-workflow-link-line-success-active); + --color-workflow-link-line-success-handle: var(--color-workflow-link-line-success-handle); + + --color-workflow-link-line-error-active: var(--color-workflow-link-line-error-active); + --color-workflow-link-line-error-handle: var(--color-workflow-link-line-error-handle); + + --color-workflow-minimap-block: var(--color-workflow-minimap-block); + --color-workflow-minimap-bg: var(--color-workflow-minimap-bg); + + --color-workflow-display-glass-1: var(--color-workflow-display-glass-1); + --color-workflow-display-glass-2: var(--color-workflow-display-glass-2); + --color-workflow-display-highlight: var(--color-workflow-display-highlight); + --color-workflow-display-outline: var(--color-workflow-display-outline); + --color-workflow-display-vignette-dark: var(--color-workflow-display-vignette-dark); + --color-workflow-display-success-bg: var(--color-workflow-display-success-bg); + --color-workflow-display-success-bg-line-pattern: var(--color-workflow-display-success-bg-line-pattern); + --color-workflow-display-success-border-1: var(--color-workflow-display-success-border-1); + --color-workflow-display-success-border-2: var(--color-workflow-display-success-border-2); + --color-workflow-display-success-vignette-color: var(--color-workflow-display-success-vignette-color); + + --color-workflow-display-error-bg: var(--color-workflow-display-error-bg); + --color-workflow-display-error-bg-line-pattern: var(--color-workflow-display-error-bg-line-pattern); + --color-workflow-display-error-border-1: var(--color-workflow-display-error-border-1); + --color-workflow-display-error-border-2: var(--color-workflow-display-error-border-2); + --color-workflow-display-error-vignette-color: var(--color-workflow-display-error-vignette-color); + + --color-workflow-display-warning-bg: var(--color-workflow-display-warning-bg); + --color-workflow-display-warning-bg-line-pattern: var(--color-workflow-display-warning-bg-line-pattern); + --color-workflow-display-warning-border-1: var(--color-workflow-display-warning-border-1); + --color-workflow-display-warning-border-2: var(--color-workflow-display-warning-border-2); + --color-workflow-display-warning-vignette-color: var(--color-workflow-display-warning-vignette-color); + + --color-workflow-display-normal-bg: var(--color-workflow-display-normal-bg); + --color-workflow-display-normal-bg-line-pattern: var(--color-workflow-display-normal-bg-line-pattern); + --color-workflow-display-normal-border-1: var(--color-workflow-display-normal-border-1); + --color-workflow-display-normal-border-2: var(--color-workflow-display-normal-border-2); + --color-workflow-display-normal-vignette-color: var(--color-workflow-display-normal-vignette-color); + + --color-workflow-display-disabled-bg: var(--color-workflow-display-disabled-bg); + --color-workflow-display-disabled-bg-line-pattern: var(--color-workflow-display-disabled-bg-line-pattern); + --color-workflow-display-disabled-border-1: var(--color-workflow-display-disabled-border-1); + --color-workflow-display-disabled-border-2: var(--color-workflow-display-disabled-border-2); + --color-workflow-display-disabled-vignette-color: var(--color-workflow-display-disabled-vignette-color); + --color-workflow-display-disabled-outline: var(--color-workflow-display-disabled-outline); + + --color-workflow-canvas-workflow-dot-color: var(--color-workflow-canvas-workflow-dot-color); + --color-workflow-canvas-workflow-bg: var(--color-workflow-canvas-workflow-bg); + --color-workflow-canvas-workflow-top-bar-1: var(--color-workflow-canvas-workflow-top-bar-1); + --color-workflow-canvas-workflow-top-bar-2: var(--color-workflow-canvas-workflow-top-bar-2); + --color-workflow-canvas-canvas-overlay: var(--color-workflow-canvas-canvas-overlay); + + --color-workflow-debug-run-status-bg: var(--color-workflow-debug-run-status-bg); + --color-workflow-debug-run-status-bg-alt: var(--color-workflow-debug-run-status-bg-alt); + --color-workflow-debug-breakpoint: var(--color-workflow-debug-breakpoint); + --color-workflow-debug-text: var(--color-workflow-debug-text); + --color-workflow-debug-text-disabled: var(--color-workflow-debug-text-disabled); + + --color-workflow-test-run-run-status-bg: var(--color-workflow-test-run-run-status-bg); + --color-workflow-test-run-paused-bg: var(--color-workflow-test-run-paused-bg); + --color-workflow-test-run-paused-text: var(--color-workflow-test-run-paused-text); + --color-workflow-test-run-run-status-bg-alt: var(--color-workflow-test-run-run-status-bg-alt); + --color-workflow-test-run-text: var(--color-workflow-test-run-text); + + --color-components-icon-bg-red-solid: var(--color-components-icon-bg-red-solid); + --color-components-icon-bg-rose-solid: var(--color-components-icon-bg-rose-solid); + --color-components-icon-bg-pink-solid: var(--color-components-icon-bg-pink-solid); + --color-components-icon-bg-orange-dark-solid: var(--color-components-icon-bg-orange-dark-solid); + --color-components-icon-bg-orange-solid: var(--color-components-icon-bg-orange-solid); + --color-components-icon-bg-yellow-solid: var(--color-components-icon-bg-yellow-solid); + --color-components-icon-bg-green-solid: var(--color-components-icon-bg-green-solid); + --color-components-icon-bg-teal-solid: var(--color-components-icon-bg-teal-solid); + --color-components-icon-bg-blue-light-solid: var(--color-components-icon-bg-blue-light-solid); + --color-components-icon-bg-blue-solid: var(--color-components-icon-bg-blue-solid); + --color-components-icon-bg-indigo-solid: var(--color-components-icon-bg-indigo-solid); + --color-components-icon-bg-violet-solid: var(--color-components-icon-bg-violet-solid); + --color-components-icon-bg-midnight-solid: var(--color-components-icon-bg-midnight-solid); + --color-components-icon-bg-red-soft: var(--color-components-icon-bg-red-soft); + --color-components-icon-bg-rose-soft: var(--color-components-icon-bg-rose-soft); + --color-components-icon-bg-pink-soft: var(--color-components-icon-bg-pink-soft); + --color-components-icon-bg-orange-dark-soft: var(--color-components-icon-bg-orange-dark-soft); + --color-components-icon-bg-orange-soft: var(--color-components-icon-bg-orange-soft); + --color-components-icon-bg-yellow-soft: var(--color-components-icon-bg-yellow-soft); + --color-components-icon-bg-green-soft: var(--color-components-icon-bg-green-soft); + --color-components-icon-bg-teal-soft: var(--color-components-icon-bg-teal-soft); + --color-components-icon-bg-blue-light-soft: var(--color-components-icon-bg-blue-light-soft); + --color-components-icon-bg-blue-soft: var(--color-components-icon-bg-blue-soft); + --color-components-icon-bg-indigo-soft: var(--color-components-icon-bg-indigo-soft); + --color-components-icon-bg-violet-soft: var(--color-components-icon-bg-violet-soft); + --color-components-icon-bg-midnight-soft: var(--color-components-icon-bg-midnight-soft); + + --color-components-avatar-default-avatar-bg: var(--color-components-avatar-default-avatar-bg); + --color-components-avatar-mask-darkmode-dimmed: var(--color-components-avatar-mask-darkmode-dimmed); + --color-components-avatar-shape-fill-stop-0: var(--color-components-avatar-shape-fill-stop-0); + --color-components-avatar-shape-fill-stop-100: var(--color-components-avatar-shape-fill-stop-100); + + --color-components-avatar-bg-mask-stop-0: var(--color-components-avatar-bg-mask-stop-0); + --color-components-avatar-bg-mask-stop-100: var(--color-components-avatar-bg-mask-stop-100); + + --color-components-chat-input-audio-bg: var(--color-components-chat-input-audio-bg); + --color-components-chat-input-audio-bg-alt: var(--color-components-chat-input-audio-bg-alt); + --color-components-chat-input-audio-wave-default: var(--color-components-chat-input-audio-wave-default); + --color-components-chat-input-audio-wave-active: var(--color-components-chat-input-audio-wave-active); + --color-components-chat-input-bg-mask-1: var(--color-components-chat-input-bg-mask-1); + --color-components-chat-input-bg-mask-2: var(--color-components-chat-input-bg-mask-2); + --color-components-chat-input-border: var(--color-components-chat-input-border); + + --color-components-label-gray: var(--color-components-label-gray); + + --color-components-premium-badge-highlight-stop-0: var(--color-components-premium-badge-highlight-stop-0); + --color-components-premium-badge-highlight-stop-100: var(--color-components-premium-badge-highlight-stop-100); + --color-components-premium-badge-orange-bg-stop-0: var(--color-components-premium-badge-orange-bg-stop-0); + --color-components-premium-badge-orange-bg-stop-100: var(--color-components-premium-badge-orange-bg-stop-100); + --color-components-premium-badge-orange-stroke-stop-0: var(--color-components-premium-badge-orange-stroke-stop-0); + --color-components-premium-badge-orange-stroke-stop-100: var(--color-components-premium-badge-orange-stroke-stop-100); + --color-components-premium-badge-orange-text-stop-0: var(--color-components-premium-badge-orange-text-stop-0); + --color-components-premium-badge-orange-text-stop-100: var(--color-components-premium-badge-orange-text-stop-100); + --color-components-premium-badge-orange-glow: var(--color-components-premium-badge-orange-glow); + --color-components-premium-badge-orange-glow-hover: var(--color-components-premium-badge-orange-glow-hover); + --color-components-premium-badge-orange-bg-stop-0-hover: var(--color-components-premium-badge-orange-bg-stop-0-hover); + --color-components-premium-badge-orange-bg-stop-100-hover: var(--color-components-premium-badge-orange-bg-stop-100-hover); + --color-components-premium-badge-orange-stroke-stop-0-hover: var(--color-components-premium-badge-orange-stroke-stop-0-hover); + --color-components-premium-badge-orange-stroke-stop-100-hover: var(--color-components-premium-badge-orange-stroke-stop-100-hover); + + --color-components-premium-badge-blue-bg-stop-0: var(--color-components-premium-badge-blue-bg-stop-0); + --color-components-premium-badge-blue-bg-stop-100: var(--color-components-premium-badge-blue-bg-stop-100); + --color-components-premium-badge-blue-stroke-stop-0: var(--color-components-premium-badge-blue-stroke-stop-0); + --color-components-premium-badge-blue-stroke-stop-100: var(--color-components-premium-badge-blue-stroke-stop-100); + --color-components-premium-badge-blue-text-stop-0: var(--color-components-premium-badge-blue-text-stop-0); + --color-components-premium-badge-blue-text-stop-100: var(--color-components-premium-badge-blue-text-stop-100); + --color-components-premium-badge-blue-glow: var(--color-components-premium-badge-blue-glow); + --color-components-premium-badge-blue-glow-hover: var(--color-components-premium-badge-blue-glow-hover); + --color-components-premium-badge-blue-bg-stop-0-hover: var(--color-components-premium-badge-blue-bg-stop-0-hover); + --color-components-premium-badge-blue-bg-stop-100-hover: var(--color-components-premium-badge-blue-bg-stop-100-hover); + --color-components-premium-badge-blue-stroke-stop-0-hover: var(--color-components-premium-badge-blue-stroke-stop-0-hover); + --color-components-premium-badge-blue-stroke-stop-100-hover: var(--color-components-premium-badge-blue-stroke-stop-100-hover); + + --color-components-premium-badge-indigo-bg-stop-0: var(--color-components-premium-badge-indigo-bg-stop-0); + --color-components-premium-badge-indigo-bg-stop-100: var(--color-components-premium-badge-indigo-bg-stop-100); + --color-components-premium-badge-indigo-stroke-stop-0: var(--color-components-premium-badge-indigo-stroke-stop-0); + --color-components-premium-badge-indigo-stroke-stop-100: var(--color-components-premium-badge-indigo-stroke-stop-100); + --color-components-premium-badge-indigo-text-stop-0: var(--color-components-premium-badge-indigo-text-stop-0); + --color-components-premium-badge-indigo-text-stop-100: var(--color-components-premium-badge-indigo-text-stop-100); + --color-components-premium-badge-indigo-glow: var(--color-components-premium-badge-indigo-glow); + --color-components-premium-badge-indigo-glow-hover: var(--color-components-premium-badge-indigo-glow-hover); + --color-components-premium-badge-indigo-bg-stop-0-hover: var(--color-components-premium-badge-indigo-bg-stop-0-hover); + --color-components-premium-badge-indigo-bg-stop-100-hover: var(--color-components-premium-badge-indigo-bg-stop-100-hover); + --color-components-premium-badge-indigo-stroke-stop-0-hover: var(--color-components-premium-badge-indigo-stroke-stop-0-hover); + --color-components-premium-badge-indigo-stroke-stop-100-hover: var(--color-components-premium-badge-indigo-stroke-stop-100-hover); + + --color-components-premium-badge-grey-bg-stop-0: var(--color-components-premium-badge-grey-bg-stop-0); + --color-components-premium-badge-grey-bg-stop-100: var(--color-components-premium-badge-grey-bg-stop-100); + --color-components-premium-badge-grey-stroke-stop-0: var(--color-components-premium-badge-grey-stroke-stop-0); + --color-components-premium-badge-grey-stroke-stop-100: var(--color-components-premium-badge-grey-stroke-stop-100); + --color-components-premium-badge-grey-text-stop-0: var(--color-components-premium-badge-grey-text-stop-0); + --color-components-premium-badge-grey-text-stop-100: var(--color-components-premium-badge-grey-text-stop-100); + --color-components-premium-badge-grey-glow: var(--color-components-premium-badge-grey-glow); + --color-components-premium-badge-grey-glow-hover: var(--color-components-premium-badge-grey-glow-hover); + --color-components-premium-badge-grey-bg-stop-0-hover: var(--color-components-premium-badge-grey-bg-stop-0-hover); + --color-components-premium-badge-grey-bg-stop-100-hover: var(--color-components-premium-badge-grey-bg-stop-100-hover); + --color-components-premium-badge-grey-stroke-stop-0-hover: var(--color-components-premium-badge-grey-stroke-stop-0-hover); + --color-components-premium-badge-grey-stroke-stop-100-hover: var(--color-components-premium-badge-grey-stroke-stop-100-hover); + + --color-components-dropzone-bg: var(--color-components-dropzone-bg); + --color-components-dropzone-bg-alt: var(--color-components-dropzone-bg-alt); + --color-components-dropzone-bg-accent: var(--color-components-dropzone-bg-accent); + --color-components-dropzone-border: var(--color-components-dropzone-border); + --color-components-dropzone-border-alt: var(--color-components-dropzone-border-alt); + --color-components-dropzone-border-accent: var(--color-components-dropzone-border-accent); + + --color-components-panel-bg: var(--color-components-panel-bg); + --color-components-panel-bg-transparent: var(--color-components-panel-bg-transparent); + --color-components-panel-bg-alt: var(--color-components-panel-bg-alt); + --color-components-panel-bg-blur: var(--color-components-panel-bg-blur); + --color-components-panel-bg-blur-burn: var(--color-components-panel-bg-blur-burn); + --color-components-panel-border: var(--color-components-panel-border); + --color-components-panel-border-subtle: var(--color-components-panel-border-subtle); + --color-components-panel-gradient-1: var(--color-components-panel-gradient-1); + --color-components-panel-gradient-2: var(--color-components-panel-gradient-2); + --color-components-panel-on-panel-item-bg: var(--color-components-panel-on-panel-item-bg); + --color-components-panel-on-panel-item-bg-transparent: var(--color-components-panel-on-panel-item-bg-transparent); + --color-components-panel-on-panel-item-bg-hover: var(--color-components-panel-on-panel-item-bg-hover); + --color-components-panel-on-panel-item-bg-hover-transparent: var(--color-components-panel-on-panel-item-bg-hover-transparent); + --color-components-panel-on-panel-item-bg-destructive-hover-transparent: var(--color-components-panel-on-panel-item-bg-destructive-hover-transparent); + --color-components-panel-on-panel-item-bg-alt: var(--color-components-panel-on-panel-item-bg-alt); + + --color-components-marketplace-header-bg: var(--color-components-marketplace-header-bg); + + --color-components-card-bg: var(--color-components-card-bg); + --color-components-card-bg-alt: var(--color-components-card-bg-alt); + --color-components-card-bg-transparent: var(--color-components-card-bg-transparent); + --color-components-card-bg-alt-transparent: var(--color-components-card-bg-alt-transparent); + --color-components-card-border: var(--color-components-card-border); + + --color-components-actionbar-border: var(--color-components-actionbar-border); + --color-components-actionbar-bg: var(--color-components-actionbar-bg); + --color-components-actionbar-bg-accent: var(--color-components-actionbar-bg-accent); + --color-components-actionbar-border-accent: var(--color-components-actionbar-border-accent); + --color-components-input-bg-normal: var(--color-components-input-bg-normal); - --color-components-input-text-placeholder: var(--color-components-input-text-placeholder); --color-components-input-bg-hover: var(--color-components-input-bg-hover); --color-components-input-bg-active: var(--color-components-input-bg-active); - --color-components-input-border-active: var(--color-components-input-border-active); - --color-components-input-border-destructive: var(--color-components-input-border-destructive); - --color-components-input-text-filled: var(--color-components-input-text-filled); - --color-components-input-bg-destructive: var(--color-components-input-bg-destructive); --color-components-input-bg-disabled: var(--color-components-input-bg-disabled); + --color-components-input-bg-destructive: var(--color-components-input-bg-destructive); + --color-components-input-text-filled: var(--color-components-input-text-filled); + --color-components-input-text-placeholder: var(--color-components-input-text-placeholder); --color-components-input-text-disabled: var(--color-components-input-text-disabled); --color-components-input-text-filled-disabled: var(--color-components-input-text-filled-disabled); + --color-components-input-text-filled-blue: var(--color-components-input-text-filled-blue); --color-components-input-border-hover: var(--color-components-input-border-hover); + --color-components-input-border-active: var(--color-components-input-border-active); --color-components-input-border-active-prompt-1: var(--color-components-input-border-active-prompt-1); --color-components-input-border-active-prompt-2: var(--color-components-input-border-active-prompt-2); + --color-components-input-border-destructive: var(--color-components-input-border-destructive); - --color-components-kbd-bg-gray: var(--color-components-kbd-bg-gray); - --color-components-kbd-bg-white: var(--color-components-kbd-bg-white); + --color-components-progress-brand-progress: var(--color-components-progress-brand-progress); + --color-components-progress-brand-border: var(--color-components-progress-brand-border); + --color-components-progress-brand-bg: var(--color-components-progress-brand-bg); - --color-components-tooltip-bg: var(--color-components-tooltip-bg); + --color-components-progress-white-progress: var(--color-components-progress-white-progress); + --color-components-progress-white-border: var(--color-components-progress-white-border); + --color-components-progress-white-bg: var(--color-components-progress-white-bg); + --color-components-progress-gray-progress: var(--color-components-progress-gray-progress); + --color-components-progress-gray-border: var(--color-components-progress-gray-border); + --color-components-progress-gray-bg: var(--color-components-progress-gray-bg); + + --color-components-progress-warning-progress: var(--color-components-progress-warning-progress); + --color-components-progress-warning-border: var(--color-components-progress-warning-border); + --color-components-progress-warning-bg: var(--color-components-progress-warning-bg); + + --color-components-progress-error-progress: var(--color-components-progress-error-progress); + --color-components-progress-error-border: var(--color-components-progress-error-border); + --color-components-progress-error-bg: var(--color-components-progress-error-bg); + + --color-components-progress-bar-progress: var(--color-components-progress-bar-progress); + --color-components-progress-bar-progress-highlight: var(--color-components-progress-bar-progress-highlight); + --color-components-progress-bar-progress-solid: var(--color-components-progress-bar-progress-solid); + --color-components-progress-bar-border: var(--color-components-progress-bar-border); + --color-components-progress-bar-bg: var(--color-components-progress-bar-bg); + + --color-components-button-button-seam: var(--color-components-button-button-seam); --color-components-button-primary-text: var(--color-components-button-primary-text); - --color-components-button-primary-bg: var(--color-components-button-primary-bg); - --color-components-button-primary-border: var(--color-components-button-primary-border); - --color-components-button-primary-bg-hover: var(--color-components-button-primary-bg-hover); - --color-components-button-primary-border-hover: var(--color-components-button-primary-border-hover); - --color-components-button-primary-bg-disabled: var(--color-components-button-primary-bg-disabled); - --color-components-button-primary-border-disabled: var(--color-components-button-primary-border-disabled); --color-components-button-primary-text-disabled: var(--color-components-button-primary-text-disabled); + --color-components-button-primary-bg: var(--color-components-button-primary-bg); + --color-components-button-primary-bg-hover: var(--color-components-button-primary-bg-hover); + --color-components-button-primary-bg-disabled: var(--color-components-button-primary-bg-disabled); + --color-components-button-primary-border: var(--color-components-button-primary-border); + --color-components-button-primary-border-hover: var(--color-components-button-primary-border-hover); + --color-components-button-primary-border-disabled: var(--color-components-button-primary-border-disabled); --color-components-button-secondary-text: var(--color-components-button-secondary-text); --color-components-button-secondary-text-disabled: var(--color-components-button-secondary-text-disabled); @@ -45,6 +379,15 @@ --color-components-button-secondary-border-hover: var(--color-components-button-secondary-border-hover); --color-components-button-secondary-border-disabled: var(--color-components-button-secondary-border-disabled); + --color-components-button-secondary-accent-text: var(--color-components-button-secondary-accent-text); + --color-components-button-secondary-accent-text-disabled: var(--color-components-button-secondary-accent-text-disabled); + --color-components-button-secondary-accent-bg: var(--color-components-button-secondary-accent-bg); + --color-components-button-secondary-accent-bg-hover: var(--color-components-button-secondary-accent-bg-hover); + --color-components-button-secondary-accent-bg-disabled: var(--color-components-button-secondary-accent-bg-disabled); + --color-components-button-secondary-accent-border: var(--color-components-button-secondary-accent-border); + --color-components-button-secondary-accent-border-hover: var(--color-components-button-secondary-accent-border-hover); + --color-components-button-secondary-accent-border-disabled: var(--color-components-button-secondary-accent-border-disabled); + --color-components-button-tertiary-text: var(--color-components-button-tertiary-text); --color-components-button-tertiary-text-disabled: var(--color-components-button-tertiary-text-disabled); --color-components-button-tertiary-bg: var(--color-components-button-tertiary-bg); @@ -83,42 +426,43 @@ --color-components-button-destructive-ghost-text-disabled: var(--color-components-button-destructive-ghost-text-disabled); --color-components-button-destructive-ghost-bg-hover: var(--color-components-button-destructive-ghost-bg-hover); - --color-components-button-secondary-accent-text: var(--color-components-button-secondary-accent-text); - --color-components-button-secondary-accent-text-disabled: var(--color-components-button-secondary-accent-text-disabled); - --color-components-button-secondary-accent-bg: var(--color-components-button-secondary-accent-bg); - --color-components-button-secondary-accent-bg-hover: var(--color-components-button-secondary-accent-bg-hover); - --color-components-button-secondary-accent-bg-disabled: var(--color-components-button-secondary-accent-bg-disabled); - --color-components-button-secondary-accent-border: var(--color-components-button-secondary-accent-border); - --color-components-button-secondary-accent-border-hover: var(--color-components-button-secondary-accent-border-hover); - --color-components-button-secondary-accent-border-disabled: var(--color-components-button-secondary-accent-border-disabled); - --color-components-button-indigo-bg: var(--color-components-button-indigo-bg); --color-components-button-indigo-bg-hover: var(--color-components-button-indigo-bg-hover); --color-components-button-indigo-bg-disabled: var(--color-components-button-indigo-bg-disabled); + --color-components-button-debug-text: var(--color-components-button-debug-text); + --color-components-button-debug-text-disabled: var(--color-components-button-debug-text-disabled); + --color-components-button-debug-bg: var(--color-components-button-debug-bg); + --color-components-button-debug-bg-hover: var(--color-components-button-debug-bg-hover); + --color-components-button-debug-bg-disabled: var(--color-components-button-debug-bg-disabled); + --color-components-button-debug-border: var(--color-components-button-debug-border); + --color-components-button-debug-border-hover: var(--color-components-button-debug-border-hover); + --color-components-button-debug-border-disabled: var(--color-components-button-debug-border-disabled); + --color-components-checkbox-icon: var(--color-components-checkbox-icon); --color-components-checkbox-icon-disabled: var(--color-components-checkbox-icon-disabled); --color-components-checkbox-bg: var(--color-components-checkbox-bg); --color-components-checkbox-bg-hover: var(--color-components-checkbox-bg-hover); --color-components-checkbox-bg-disabled: var(--color-components-checkbox-bg-disabled); + --color-components-checkbox-bg-disabled-checked: var(--color-components-checkbox-bg-disabled-checked); + --color-components-checkbox-bg-unchecked: var(--color-components-checkbox-bg-unchecked); + --color-components-checkbox-bg-unchecked-hover: var(--color-components-checkbox-bg-unchecked-hover); --color-components-checkbox-border: var(--color-components-checkbox-border); --color-components-checkbox-border-hover: var(--color-components-checkbox-border-hover); --color-components-checkbox-border-disabled: var(--color-components-checkbox-border-disabled); - --color-components-checkbox-bg-unchecked: var(--color-components-checkbox-bg-unchecked); - --color-components-checkbox-bg-unchecked-hover: var(--color-components-checkbox-bg-unchecked-hover); - --color-components-checkbox-bg-disabled-checked: var(--color-components-checkbox-bg-disabled-checked); + --color-components-radio-bg: var(--color-components-radio-bg); + --color-components-radio-bg-hover: var(--color-components-radio-bg-hover); + --color-components-radio-bg-disabled: var(--color-components-radio-bg-disabled); --color-components-radio-border-checked: var(--color-components-radio-border-checked); --color-components-radio-border-checked-hover: var(--color-components-radio-border-checked-hover); --color-components-radio-border-checked-disabled: var(--color-components-radio-border-checked-disabled); - --color-components-radio-bg-disabled: var(--color-components-radio-bg-disabled); --color-components-radio-border: var(--color-components-radio-border); --color-components-radio-border-hover: var(--color-components-radio-border-hover); --color-components-radio-border-disabled: var(--color-components-radio-border-disabled); - --color-components-radio-bg: var(--color-components-radio-bg); - --color-components-radio-bg-hover: var(--color-components-radio-bg-hover); --color-components-toggle-knob: var(--color-components-toggle-knob); + --color-components-toggle-knob-hover: var(--color-components-toggle-knob-hover); --color-components-toggle-knob-disabled: var(--color-components-toggle-knob-disabled); --color-components-toggle-bg: var(--color-components-toggle-bg); --color-components-toggle-bg-hover: var(--color-components-toggle-bg-hover); @@ -126,90 +470,55 @@ --color-components-toggle-bg-unchecked: var(--color-components-toggle-bg-unchecked); --color-components-toggle-bg-unchecked-hover: var(--color-components-toggle-bg-unchecked-hover); --color-components-toggle-bg-unchecked-disabled: var(--color-components-toggle-bg-unchecked-disabled); - --color-components-toggle-knob-hover: var(--color-components-toggle-knob-hover); - --color-components-card-bg: var(--color-components-card-bg); - --color-components-card-border: var(--color-components-card-border); - --color-components-card-bg-alt: var(--color-components-card-bg-alt); - --color-components-card-bg-transparent: var(--color-components-card-bg-transparent); - --color-components-card-bg-alt-transparent: var(--color-components-card-bg-alt-transparent); - - --color-components-menu-item-text: var(--color-components-menu-item-text); - --color-components-menu-item-text-active: var(--color-components-menu-item-text-active); - --color-components-menu-item-text-hover: var(--color-components-menu-item-text-hover); - --color-components-menu-item-text-active-accent: var(--color-components-menu-item-text-active-accent); - --color-components-menu-item-bg-active: var(--color-components-menu-item-bg-active); - --color-components-menu-item-bg-hover: var(--color-components-menu-item-bg-hover); - - --color-components-panel-bg: var(--color-components-panel-bg); - --color-components-panel-bg-blur: var(--color-components-panel-bg-blur); - --color-components-panel-border: var(--color-components-panel-border); - --color-components-panel-border-subtle: var(--color-components-panel-border-subtle); - --color-components-panel-gradient-2: var(--color-components-panel-gradient-2); - --color-components-panel-gradient-1: var(--color-components-panel-gradient-1); - --color-components-panel-bg-alt: var(--color-components-panel-bg-alt); - --color-components-panel-on-panel-item-bg: var(--color-components-panel-on-panel-item-bg); - --color-components-panel-on-panel-item-bg-hover: var(--color-components-panel-on-panel-item-bg-hover); - --color-components-panel-on-panel-item-bg-alt: var(--color-components-panel-on-panel-item-bg-alt); - --color-components-panel-on-panel-item-bg-transparent: var(--color-components-panel-on-panel-item-bg-transparent); - --color-components-panel-on-panel-item-bg-hover-transparent: var(--color-components-panel-on-panel-item-bg-hover-transparent); - --color-components-panel-on-panel-item-bg-destructive-hover-transparent: var(--color-components-panel-on-panel-item-bg-destructive-hover-transparent); - - --color-components-panel-bg-transparent: var(--color-components-panel-bg-transparent); - - --color-components-main-nav-nav-button-text: var(--color-components-main-nav-nav-button-text); - --color-components-main-nav-nav-button-text-active: var(--color-components-main-nav-nav-button-text-active); - --color-components-main-nav-nav-button-bg: var(--color-components-main-nav-nav-button-bg); - --color-components-main-nav-nav-button-bg-active: var(--color-components-main-nav-nav-button-bg-active); - --color-components-main-nav-nav-button-border: var(--color-components-main-nav-nav-button-border); - --color-components-main-nav-nav-button-bg-hover: var(--color-components-main-nav-nav-button-bg-hover); - --color-components-main-nav-glass-text-glow: var(--color-components-main-nav-glass-text-glow); - --color-components-main-nav-glass-surface-first: var(--color-components-main-nav-glass-surface-first); - --color-components-main-nav-glass-surface-middle-1: var(--color-components-main-nav-glass-surface-middle-1); - --color-components-main-nav-glass-surface-middle-2: var(--color-components-main-nav-glass-surface-middle-2); - --color-components-main-nav-glass-surface-end: var(--color-components-main-nav-glass-surface-end); - --color-components-main-nav-glass-edge-highlight-first: var(--color-components-main-nav-glass-edge-highlight-first); - --color-components-main-nav-glass-edge-highlight-middle: var(--color-components-main-nav-glass-edge-highlight-middle); - --color-components-main-nav-glass-edge-highlight-end: var(--color-components-main-nav-glass-edge-highlight-end); - --color-components-main-nav-glass-edge-reflection-first: var(--color-components-main-nav-glass-edge-reflection-first); - --color-components-main-nav-glass-edge-reflection-middle: var(--color-components-main-nav-glass-edge-reflection-middle); - --color-components-main-nav-glass-edge-reflection-end: var(--color-components-main-nav-glass-edge-reflection-end); - --color-components-main-nav-glass-inner-glow: var(--color-components-main-nav-glass-inner-glow); - --color-components-main-nav-glass-shadow-reflection: var(--color-components-main-nav-glass-shadow-reflection); - --color-components-main-nav-glass-shadow-reflection-glow: var(--color-components-main-nav-glass-shadow-reflection-glow); - - --color-components-main-nav-nav-user-border: var(--color-components-main-nav-nav-user-border); + --color-components-option-card-option-bg: var(--color-components-option-card-option-bg); + --color-components-option-card-option-border: var(--color-components-option-card-option-border); + --color-components-option-card-option-bg-hover: var(--color-components-option-card-option-bg-hover); + --color-components-option-card-option-border-hover: var(--color-components-option-card-option-border-hover); + --color-components-option-card-option-selected-bg: var(--color-components-option-card-option-selected-bg); + --color-components-option-card-option-selected-border: var(--color-components-option-card-option-selected-border); --color-components-slider-knob: var(--color-components-slider-knob); + --color-components-slider-knob-border: var(--color-components-slider-knob-border); + --color-components-slider-knob-border-hover: var(--color-components-slider-knob-border-hover); --color-components-slider-knob-hover: var(--color-components-slider-knob-hover); --color-components-slider-knob-disabled: var(--color-components-slider-knob-disabled); --color-components-slider-range: var(--color-components-slider-range); --color-components-slider-track: var(--color-components-slider-track); - --color-components-slider-knob-border-hover: var(--color-components-slider-knob-border-hover); - --color-components-slider-knob-border: var(--color-components-slider-knob-border); + --color-components-tooltip-bg: var(--color-components-tooltip-bg); + + --color-components-kbd-bg-gray: var(--color-components-kbd-bg-gray); + --color-components-kbd-bg-white: var(--color-components-kbd-bg-white); + + --color-components-menu-item-text: var(--color-components-menu-item-text); + --color-components-menu-item-text-hover: var(--color-components-menu-item-text-hover); + --color-components-menu-item-text-active: var(--color-components-menu-item-text-active); + --color-components-menu-item-text-active-accent: var(--color-components-menu-item-text-active-accent); + --color-components-menu-item-bg-hover: var(--color-components-menu-item-bg-hover); + --color-components-menu-item-bg-active: var(--color-components-menu-item-bg-active); + + --color-components-segmented-control-bg-normal: var(--color-components-segmented-control-bg-normal); --color-components-segmented-control-item-active-bg: var(--color-components-segmented-control-item-active-bg); --color-components-segmented-control-item-active-border: var(--color-components-segmented-control-item-active-border); - --color-components-segmented-control-bg-normal: var(--color-components-segmented-control-bg-normal); --color-components-segmented-control-item-active-accent-bg: var(--color-components-segmented-control-item-active-accent-bg); --color-components-segmented-control-item-active-accent-border: var(--color-components-segmented-control-item-active-accent-border); - --color-components-option-card-option-bg: var(--color-components-option-card-option-bg); - --color-components-option-card-option-selected-bg: var(--color-components-option-card-option-selected-bg); - --color-components-option-card-option-selected-border: var(--color-components-option-card-option-selected-border); - --color-components-option-card-option-border: var(--color-components-option-card-option-border); - --color-components-option-card-option-bg-hover: var(--color-components-option-card-option-bg-hover); - --color-components-option-card-option-border-hover: var(--color-components-option-card-option-border-hover); - - --color-components-tab-active: var(--color-components-tab-active); - --color-components-badge-white-to-dark: var(--color-components-badge-white-to-dark); + --color-components-badge-white-to-dark-alpha: var(--color-components-badge-white-to-dark-alpha); + --color-components-badge-bg-green-soft: var(--color-components-badge-bg-green-soft); + --color-components-badge-bg-orange-soft: var(--color-components-badge-bg-orange-soft); + --color-components-badge-bg-red-soft: var(--color-components-badge-bg-red-soft); + --color-components-badge-bg-blue-light-soft: var(--color-components-badge-bg-blue-light-soft); + --color-components-badge-bg-gray-soft: var(--color-components-badge-bg-gray-soft); + --color-components-badge-bg-dimm: var(--color-components-badge-bg-dimm); + + --color-components-badge-status-light-border-outer: var(--color-components-badge-status-light-border-outer); + --color-components-badge-status-light-high-light: var(--color-components-badge-status-light-high-light); --color-components-badge-status-light-success-bg: var(--color-components-badge-status-light-success-bg); --color-components-badge-status-light-success-border-inner: var(--color-components-badge-status-light-success-border-inner); --color-components-badge-status-light-success-halo: var(--color-components-badge-status-light-success-halo); - --color-components-badge-status-light-border-outer: var(--color-components-badge-status-light-border-outer); - --color-components-badge-status-light-high-light: var(--color-components-badge-status-light-high-light); --color-components-badge-status-light-warning-bg: var(--color-components-badge-status-light-warning-bg); --color-components-badge-status-light-warning-border-inner: var(--color-components-badge-status-light-warning-border-inner); --color-components-badge-status-light-warning-halo: var(--color-components-badge-status-light-warning-halo); @@ -226,221 +535,60 @@ --color-components-badge-status-light-disabled-border-inner: var(--color-components-badge-status-light-disabled-border-inner); --color-components-badge-status-light-disabled-halo: var(--color-components-badge-status-light-disabled-halo); - --color-components-badge-bg-green-soft: var(--color-components-badge-bg-green-soft); - --color-components-badge-bg-orange-soft: var(--color-components-badge-bg-orange-soft); - --color-components-badge-bg-red-soft: var(--color-components-badge-bg-red-soft); - --color-components-badge-bg-blue-light-soft: var(--color-components-badge-bg-blue-light-soft); - --color-components-badge-bg-gray-soft: var(--color-components-badge-bg-gray-soft); - --color-components-badge-bg-dimm: var(--color-components-badge-bg-dimm); + --color-components-tab-active: var(--color-components-tab-active); + --color-components-main-nav-text: var(--color-components-main-nav-text); + --color-components-main-nav-text-active: var(--color-components-main-nav-text-active); + --color-components-main-nav-nav-button-border: var(--color-components-main-nav-nav-button-border); + --color-components-main-nav-nav-button-text: var(--color-components-main-nav-nav-button-text); + --color-components-main-nav-nav-button-text-active: var(--color-components-main-nav-nav-button-text-active); + --color-components-main-nav-nav-button-bg: var(--color-components-main-nav-nav-button-bg); + --color-components-main-nav-nav-button-bg-hover: var(--color-components-main-nav-nav-button-bg-hover); + --color-components-main-nav-nav-button-bg-active: var(--color-components-main-nav-nav-button-bg-active); + + --color-components-main-nav-nav-user-border: var(--color-components-main-nav-nav-user-border); + + --color-components-main-nav-glass-inner-glow: var(--color-components-main-nav-glass-inner-glow); + --color-components-main-nav-glass-shadow-reflection: var(--color-components-main-nav-glass-shadow-reflection); + --color-components-main-nav-glass-shadow-reflection-glow: var(--color-components-main-nav-glass-shadow-reflection-glow); + --color-components-main-nav-glass-text-glow: var(--color-components-main-nav-glass-text-glow); + --color-components-main-nav-glass-surface-first: var(--color-components-main-nav-glass-surface-first); + --color-components-main-nav-glass-surface-middle-1: var(--color-components-main-nav-glass-surface-middle-1); + --color-components-main-nav-glass-surface-middle-2: var(--color-components-main-nav-glass-surface-middle-2); + --color-components-main-nav-glass-surface-end: var(--color-components-main-nav-glass-surface-end); + + --color-components-main-nav-glass-edge-reflection-first: var(--color-components-main-nav-glass-edge-reflection-first); + --color-components-main-nav-glass-edge-reflection-middle: var(--color-components-main-nav-glass-edge-reflection-middle); + --color-components-main-nav-glass-edge-reflection-end: var(--color-components-main-nav-glass-edge-reflection-end); + + --color-components-main-nav-glass-edge-highlight-first: var(--color-components-main-nav-glass-edge-highlight-first); + --color-components-main-nav-glass-edge-highlight-middle: var(--color-components-main-nav-glass-edge-highlight-middle); + --color-components-main-nav-glass-edge-highlight-end: var(--color-components-main-nav-glass-edge-highlight-end); + + --color-components-chart-bg: var(--color-components-chart-bg); --color-components-chart-line: var(--color-components-chart-line); - --color-components-chart-area-1: var(--color-components-chart-area-1); - --color-components-chart-area-2: var(--color-components-chart-area-2); --color-components-chart-current-1: var(--color-components-chart-current-1); --color-components-chart-current-2: var(--color-components-chart-current-2); - --color-components-chart-bg: var(--color-components-chart-bg); + --color-components-chart-area-1: var(--color-components-chart-area-1); + --color-components-chart-area-2: var(--color-components-chart-area-2); - --color-components-actionbar-bg: var(--color-components-actionbar-bg); - --color-components-actionbar-border: var(--color-components-actionbar-border); - --color-components-actionbar-bg-accent: var(--color-components-actionbar-bg-accent); - --color-components-actionbar-border-accent: var(--color-components-actionbar-border-accent); + --color-divider-subtle: var(--color-divider-subtle); + --color-divider-regular: var(--color-divider-regular); + --color-divider-deep: var(--color-divider-deep); + --color-divider-intense: var(--color-divider-intense); + --color-divider-burn: var(--color-divider-burn); + --color-divider-solid: var(--color-divider-solid); + --color-divider-solid-alt: var(--color-divider-solid-alt); + --color-divider-accent: var(--color-divider-accent); - --color-components-dropzone-bg-alt: var(--color-components-dropzone-bg-alt); - --color-components-dropzone-bg: var(--color-components-dropzone-bg); - --color-components-dropzone-bg-accent: var(--color-components-dropzone-bg-accent); - --color-components-dropzone-border: var(--color-components-dropzone-border); - --color-components-dropzone-border-alt: var(--color-components-dropzone-border-alt); - --color-components-dropzone-border-accent: var(--color-components-dropzone-border-accent); - - --color-components-progress-brand-progress: var(--color-components-progress-brand-progress); - --color-components-progress-brand-border: var(--color-components-progress-brand-border); - --color-components-progress-brand-bg: var(--color-components-progress-brand-bg); - - --color-components-progress-white-progress: var(--color-components-progress-white-progress); - --color-components-progress-white-border: var(--color-components-progress-white-border); - --color-components-progress-white-bg: var(--color-components-progress-white-bg); - - --color-components-progress-gray-progress: var(--color-components-progress-gray-progress); - --color-components-progress-gray-border: var(--color-components-progress-gray-border); - --color-components-progress-gray-bg: var(--color-components-progress-gray-bg); - - --color-components-progress-warning-progress: var(--color-components-progress-warning-progress); - --color-components-progress-warning-border: var(--color-components-progress-warning-border); - --color-components-progress-warning-bg: var(--color-components-progress-warning-bg); - - --color-components-progress-error-progress: var(--color-components-progress-error-progress); - --color-components-progress-error-border: var(--color-components-progress-error-border); - --color-components-progress-error-bg: var(--color-components-progress-error-bg); - - --color-components-chat-input-audio-bg: var(--color-components-chat-input-audio-bg); - --color-components-chat-input-audio-wave-default: var(--color-components-chat-input-audio-wave-default); - --color-components-chat-input-bg-mask-1: var(--color-components-chat-input-bg-mask-1); - --color-components-chat-input-bg-mask-2: var(--color-components-chat-input-bg-mask-2); - --color-components-chat-input-border: var(--color-components-chat-input-border); - --color-components-chat-input-audio-wave-active: var(--color-components-chat-input-audio-wave-active); - --color-components-chat-input-audio-bg-alt: var(--color-components-chat-input-audio-bg-alt); - - --color-components-avatar-shape-fill-stop-0: var(--color-components-avatar-shape-fill-stop-0); - --color-components-avatar-shape-fill-stop-100: var(--color-components-avatar-shape-fill-stop-100); - - --color-components-avatar-bg-mask-stop-0: var(--color-components-avatar-bg-mask-stop-0); - --color-components-avatar-bg-mask-stop-100: var(--color-components-avatar-bg-mask-stop-100); - - --color-components-avatar-default-avatar-bg: var(--color-components-avatar-default-avatar-bg); - --color-components-avatar-mask-darkmode-dimmed: var(--color-components-avatar-mask-darkmode-dimmed); - - --color-components-label-gray: var(--color-components-label-gray); - - --color-components-premium-badge-blue-bg-stop-0: var(--color-components-premium-badge-blue-bg-stop-0); - --color-components-premium-badge-blue-bg-stop-100: var(--color-components-premium-badge-blue-bg-stop-100); - --color-components-premium-badge-blue-stroke-stop-0: var(--color-components-premium-badge-blue-stroke-stop-0); - --color-components-premium-badge-blue-stroke-stop-100: var(--color-components-premium-badge-blue-stroke-stop-100); - --color-components-premium-badge-blue-text-stop-0: var(--color-components-premium-badge-blue-text-stop-0); - --color-components-premium-badge-blue-text-stop-100: var(--color-components-premium-badge-blue-text-stop-100); - --color-components-premium-badge-blue-glow: var(--color-components-premium-badge-blue-glow); - --color-components-premium-badge-blue-bg-stop-0-hover: var(--color-components-premium-badge-blue-bg-stop-0-hover); - --color-components-premium-badge-blue-bg-stop-100-hover: var(--color-components-premium-badge-blue-bg-stop-100-hover); - --color-components-premium-badge-blue-glow-hover: var(--color-components-premium-badge-blue-glow-hover); - --color-components-premium-badge-blue-stroke-stop-0-hover: var(--color-components-premium-badge-blue-stroke-stop-0-hover); - --color-components-premium-badge-blue-stroke-stop-100-hover: var(--color-components-premium-badge-blue-stroke-stop-100-hover); - - --color-components-premium-badge-highlight-stop-0: var(--color-components-premium-badge-highlight-stop-0); - --color-components-premium-badge-highlight-stop-100: var(--color-components-premium-badge-highlight-stop-100); - --color-components-premium-badge-indigo-bg-stop-0: var(--color-components-premium-badge-indigo-bg-stop-0); - --color-components-premium-badge-indigo-bg-stop-100: var(--color-components-premium-badge-indigo-bg-stop-100); - --color-components-premium-badge-indigo-stroke-stop-0: var(--color-components-premium-badge-indigo-stroke-stop-0); - --color-components-premium-badge-indigo-stroke-stop-100: var(--color-components-premium-badge-indigo-stroke-stop-100); - --color-components-premium-badge-indigo-text-stop-0: var(--color-components-premium-badge-indigo-text-stop-0); - --color-components-premium-badge-indigo-text-stop-100: var(--color-components-premium-badge-indigo-text-stop-100); - --color-components-premium-badge-indigo-glow: var(--color-components-premium-badge-indigo-glow); - --color-components-premium-badge-indigo-glow-hover: var(--color-components-premium-badge-indigo-glow-hover); - --color-components-premium-badge-indigo-bg-stop-0-hover: var(--color-components-premium-badge-indigo-bg-stop-0-hover); - --color-components-premium-badge-indigo-bg-stop-100-hover: var(--color-components-premium-badge-indigo-bg-stop-100-hover); - --color-components-premium-badge-indigo-stroke-stop-0-hover: var(--color-components-premium-badge-indigo-stroke-stop-0-hover); - --color-components-premium-badge-indigo-stroke-stop-100-hover: var(--color-components-premium-badge-indigo-stroke-stop-100-hover); - - --color-components-premium-badge-grey-bg-stop-0: var(--color-components-premium-badge-grey-bg-stop-0); - --color-components-premium-badge-grey-bg-stop-100: var(--color-components-premium-badge-grey-bg-stop-100); - --color-components-premium-badge-grey-stroke-stop-0: var(--color-components-premium-badge-grey-stroke-stop-0); - --color-components-premium-badge-grey-stroke-stop-100: var(--color-components-premium-badge-grey-stroke-stop-100); - --color-components-premium-badge-grey-text-stop-0: var(--color-components-premium-badge-grey-text-stop-0); - --color-components-premium-badge-grey-text-stop-100: var(--color-components-premium-badge-grey-text-stop-100); - --color-components-premium-badge-grey-glow: var(--color-components-premium-badge-grey-glow); - --color-components-premium-badge-grey-glow-hover: var(--color-components-premium-badge-grey-glow-hover); - --color-components-premium-badge-grey-bg-stop-0-hover: var(--color-components-premium-badge-grey-bg-stop-0-hover); - --color-components-premium-badge-grey-bg-stop-100-hover: var(--color-components-premium-badge-grey-bg-stop-100-hover); - --color-components-premium-badge-grey-stroke-stop-0-hover: var(--color-components-premium-badge-grey-stroke-stop-0-hover); - --color-components-premium-badge-grey-stroke-stop-100-hover: var(--color-components-premium-badge-grey-stroke-stop-100-hover); - - --color-components-premium-badge-orange-bg-stop-0: var(--color-components-premium-badge-orange-bg-stop-0); - --color-components-premium-badge-orange-bg-stop-100: var(--color-components-premium-badge-orange-bg-stop-100); - --color-components-premium-badge-orange-stroke-stop-0: var(--color-components-premium-badge-orange-stroke-stop-0); - --color-components-premium-badge-orange-stroke-stop-100: var(--color-components-premium-badge-orange-stroke-stop-100); - --color-components-premium-badge-orange-text-stop-0: var(--color-components-premium-badge-orange-text-stop-0); - --color-components-premium-badge-orange-text-stop-100: var(--color-components-premium-badge-orange-text-stop-100); - --color-components-premium-badge-orange-glow: var(--color-components-premium-badge-orange-glow); - --color-components-premium-badge-orange-glow-hover: var(--color-components-premium-badge-orange-glow-hover); - --color-components-premium-badge-orange-bg-stop-0-hover: var(--color-components-premium-badge-orange-bg-stop-0-hover); - --color-components-premium-badge-orange-bg-stop-100-hover: var(--color-components-premium-badge-orange-bg-stop-100-hover); - --color-components-premium-badge-orange-stroke-stop-0-hover: var(--color-components-premium-badge-orange-stroke-stop-0-hover); - --color-components-premium-badge-orange-stroke-stop-100-hover: var(--color-components-premium-badge-orange-stroke-stop-100-hover); - - --color-components-progress-bar-bg: var(--color-components-progress-bar-bg); - --color-components-progress-bar-progress: var(--color-components-progress-bar-progress); - --color-components-progress-bar-border: var(--color-components-progress-bar-border); - --color-components-progress-bar-progress-solid: var(--color-components-progress-bar-progress-solid); - --color-components-progress-bar-progress-highlight: var(--color-components-progress-bar-progress-highlight); - - --color-components-icon-bg-red-solid: var(--color-components-icon-bg-red-solid); - --color-components-icon-bg-rose-solid: var(--color-components-icon-bg-rose-solid); - --color-components-icon-bg-pink-solid: var(--color-components-icon-bg-pink-solid); - --color-components-icon-bg-orange-dark-solid: var(--color-components-icon-bg-orange-dark-solid); - --color-components-icon-bg-yellow-solid: var(--color-components-icon-bg-yellow-solid); - --color-components-icon-bg-green-solid: var(--color-components-icon-bg-green-solid); - --color-components-icon-bg-teal-solid: var(--color-components-icon-bg-teal-solid); - --color-components-icon-bg-blue-light-solid: var(--color-components-icon-bg-blue-light-solid); - --color-components-icon-bg-blue-solid: var(--color-components-icon-bg-blue-solid); - --color-components-icon-bg-indigo-solid: var(--color-components-icon-bg-indigo-solid); - --color-components-icon-bg-violet-solid: var(--color-components-icon-bg-violet-solid); - --color-components-icon-bg-midnight-solid: var(--color-components-icon-bg-midnight-solid); - --color-components-icon-bg-rose-soft: var(--color-components-icon-bg-rose-soft); - --color-components-icon-bg-pink-soft: var(--color-components-icon-bg-pink-soft); - --color-components-icon-bg-orange-dark-soft: var(--color-components-icon-bg-orange-dark-soft); - --color-components-icon-bg-yellow-soft: var(--color-components-icon-bg-yellow-soft); - --color-components-icon-bg-green-soft: var(--color-components-icon-bg-green-soft); - --color-components-icon-bg-teal-soft: var(--color-components-icon-bg-teal-soft); - --color-components-icon-bg-blue-light-soft: var(--color-components-icon-bg-blue-light-soft); - --color-components-icon-bg-blue-soft: var(--color-components-icon-bg-blue-soft); - --color-components-icon-bg-indigo-soft: var(--color-components-icon-bg-indigo-soft); - --color-components-icon-bg-violet-soft: var(--color-components-icon-bg-violet-soft); - --color-components-icon-bg-midnight-soft: var(--color-components-icon-bg-midnight-soft); - --color-components-icon-bg-red-soft: var(--color-components-icon-bg-red-soft); - --color-components-icon-bg-orange-solid: var(--color-components-icon-bg-orange-solid); - --color-components-icon-bg-orange-soft: var(--color-components-icon-bg-orange-soft); - - --color-text-primary: var(--color-text-primary); - --color-text-secondary: var(--color-text-secondary); - --color-text-tertiary: var(--color-text-tertiary); - --color-text-quaternary: var(--color-text-quaternary); - --color-text-destructive: var(--color-text-destructive); - --color-text-success: var(--color-text-success); - --color-text-warning: var(--color-text-warning); - --color-text-destructive-secondary: var(--color-text-destructive-secondary); - --color-text-success-secondary: var(--color-text-success-secondary); - --color-text-warning-secondary: var(--color-text-warning-secondary); - --color-text-accent: var(--color-text-accent); - --color-text-primary-on-surface: var(--color-text-primary-on-surface); - --color-text-placeholder: var(--color-text-placeholder); - --color-text-disabled: var(--color-text-disabled); - --color-text-accent-secondary: var(--color-text-accent-secondary); - --color-text-accent-light-mode-only: var(--color-text-accent-light-mode-only); - --color-text-text-selected: var(--color-text-text-selected); - --color-text-secondary-on-surface: var(--color-text-secondary-on-surface); - --color-text-logo-text: var(--color-text-logo-text); - --color-text-empty-state-icon: var(--color-text-empty-state-icon); - --color-text-inverted: var(--color-text-inverted); - --color-text-inverted-dimmed: var(--color-text-inverted-dimmed); - - --color-background-body: var(--color-background-body); - --color-background-default-subtle: var(--color-background-default-subtle); - --color-background-neutral-subtle: var(--color-background-neutral-subtle); - --color-background-sidenav-bg: var(--color-background-sidenav-bg); - --color-background-default: var(--color-background-default); - --color-background-soft: var(--color-background-soft); - --color-background-gradient-bg-fill-chat-bg-1: var(--color-background-gradient-bg-fill-chat-bg-1); - --color-background-gradient-bg-fill-chat-bg-2: var(--color-background-gradient-bg-fill-chat-bg-2); - --color-background-gradient-bg-fill-chat-bubble-bg-1: var(--color-background-gradient-bg-fill-chat-bubble-bg-1); - --color-background-gradient-bg-fill-chat-bubble-bg-2: var(--color-background-gradient-bg-fill-chat-bubble-bg-2); - --color-background-gradient-bg-fill-debug-bg-1: var(--color-background-gradient-bg-fill-debug-bg-1); - --color-background-gradient-bg-fill-debug-bg-2: var(--color-background-gradient-bg-fill-debug-bg-2); - - --color-background-gradient-mask-gray: var(--color-background-gradient-mask-gray); - --color-background-gradient-mask-transparent: var(--color-background-gradient-mask-transparent); - --color-background-gradient-mask-input-clear-2: var(--color-background-gradient-mask-input-clear-2); - --color-background-gradient-mask-input-clear-1: var(--color-background-gradient-mask-input-clear-1); - --color-background-gradient-mask-transparent-dark: var(--color-background-gradient-mask-transparent-dark); - --color-background-gradient-mask-side-panel-2: var(--color-background-gradient-mask-side-panel-2); - --color-background-gradient-mask-side-panel-1: var(--color-background-gradient-mask-side-panel-1); - - --color-background-default-burn: var(--color-background-default-burn); - --color-background-overlay-fullscreen: var(--color-background-overlay-fullscreen); - --color-background-default-lighter: var(--color-background-default-lighter); - --color-background-section: var(--color-background-section); - --color-background-interaction-from-bg-1: var(--color-background-interaction-from-bg-1); - --color-background-interaction-from-bg-2: var(--color-background-interaction-from-bg-2); - --color-background-section-burn: var(--color-background-section-burn); - --color-background-default-dodge: var(--color-background-default-dodge); - --color-background-overlay: var(--color-background-overlay); - --color-background-default-dimmed: var(--color-background-default-dimmed); - --color-background-default-hover: var(--color-background-default-hover); - --color-background-overlay-alt: var(--color-background-overlay-alt); - --color-background-surface-white: var(--color-background-surface-white); - --color-background-overlay-destructive: var(--color-background-overlay-destructive); - --color-background-overlay-backdrop: var(--color-background-overlay-backdrop); - --color-background-body-transparent: var(--color-background-body-transparent); - --color-background-section-burn-inverted: var(--color-background-section-burn-inverted); + --color-effects-highlight: var(--color-effects-highlight); + --color-effects-highlight-subtle: var(--color-effects-highlight-subtle); + --color-effects-highlight-lightmode-off: var(--color-effects-highlight-lightmode-off); + --color-effects-image-frame: var(--color-effects-image-frame); + --color-effects-icon-border: var(--color-effects-icon-border); --color-shadow-shadow-1: var(--color-shadow-shadow-1); + --color-shadow-shadow-2: var(--color-shadow-shadow-2); --color-shadow-shadow-3: var(--color-shadow-shadow-3); --color-shadow-shadow-4: var(--color-shadow-shadow-4); --color-shadow-shadow-5: var(--color-shadow-shadow-5); @@ -448,124 +596,18 @@ --color-shadow-shadow-7: var(--color-shadow-shadow-7); --color-shadow-shadow-8: var(--color-shadow-shadow-8); --color-shadow-shadow-9: var(--color-shadow-shadow-9); - --color-shadow-shadow-2: var(--color-shadow-shadow-2); --color-shadow-shadow-10: var(--color-shadow-shadow-10); - --color-workflow-block-border: var(--color-workflow-block-border); - --color-workflow-block-parma-bg: var(--color-workflow-block-parma-bg); - --color-workflow-block-bg: var(--color-workflow-block-bg); - --color-workflow-block-bg-transparent: var(--color-workflow-block-bg-transparent); - --color-workflow-block-border-highlight: var(--color-workflow-block-border-highlight); - --color-workflow-block-wrapper-bg-1: var(--color-workflow-block-wrapper-bg-1); - --color-workflow-block-wrapper-bg-2: var(--color-workflow-block-wrapper-bg-2); - - --color-workflow-canvas-workflow-dot-color: var(--color-workflow-canvas-workflow-dot-color); - --color-workflow-canvas-workflow-bg: var(--color-workflow-canvas-workflow-bg); - --color-workflow-canvas-workflow-top-bar-1: var(--color-workflow-canvas-workflow-top-bar-1); - --color-workflow-canvas-workflow-top-bar-2: var(--color-workflow-canvas-workflow-top-bar-2); - --color-workflow-canvas-canvas-overlay: var(--color-workflow-canvas-canvas-overlay); - - --color-workflow-link-line-active: var(--color-workflow-link-line-active); - --color-workflow-link-line-normal: var(--color-workflow-link-line-normal); - --color-workflow-link-line-handle: var(--color-workflow-link-line-handle); - --color-workflow-link-line-normal-transparent: var(--color-workflow-link-line-normal-transparent); - --color-workflow-link-line-failure-active: var(--color-workflow-link-line-failure-active); - --color-workflow-link-line-failure-handle: var(--color-workflow-link-line-failure-handle); - --color-workflow-link-line-failure-button-bg: var(--color-workflow-link-line-failure-button-bg); - --color-workflow-link-line-failure-button-hover: var(--color-workflow-link-line-failure-button-hover); - - --color-workflow-link-line-success-active: var(--color-workflow-link-line-success-active); - --color-workflow-link-line-success-handle: var(--color-workflow-link-line-success-handle); - - --color-workflow-link-line-error-active: var(--color-workflow-link-line-error-active); - --color-workflow-link-line-error-handle: var(--color-workflow-link-line-error-handle); - - --color-workflow-minimap-bg: var(--color-workflow-minimap-bg); - --color-workflow-minimap-block: var(--color-workflow-minimap-block); - - --color-workflow-display-success-bg: var(--color-workflow-display-success-bg); - --color-workflow-display-success-border-1: var(--color-workflow-display-success-border-1); - --color-workflow-display-success-border-2: var(--color-workflow-display-success-border-2); - --color-workflow-display-success-vignette-color: var(--color-workflow-display-success-vignette-color); - --color-workflow-display-success-bg-line-pattern: var(--color-workflow-display-success-bg-line-pattern); - - --color-workflow-display-glass-1: var(--color-workflow-display-glass-1); - --color-workflow-display-glass-2: var(--color-workflow-display-glass-2); - --color-workflow-display-vignette-dark: var(--color-workflow-display-vignette-dark); - --color-workflow-display-highlight: var(--color-workflow-display-highlight); - --color-workflow-display-outline: var(--color-workflow-display-outline); - --color-workflow-display-error-bg: var(--color-workflow-display-error-bg); - --color-workflow-display-error-bg-line-pattern: var(--color-workflow-display-error-bg-line-pattern); - --color-workflow-display-error-border-1: var(--color-workflow-display-error-border-1); - --color-workflow-display-error-border-2: var(--color-workflow-display-error-border-2); - --color-workflow-display-error-vignette-color: var(--color-workflow-display-error-vignette-color); - - --color-workflow-display-warning-bg: var(--color-workflow-display-warning-bg); - --color-workflow-display-warning-bg-line-pattern: var(--color-workflow-display-warning-bg-line-pattern); - --color-workflow-display-warning-border-1: var(--color-workflow-display-warning-border-1); - --color-workflow-display-warning-border-2: var(--color-workflow-display-warning-border-2); - --color-workflow-display-warning-vignette-color: var(--color-workflow-display-warning-vignette-color); - - --color-workflow-display-normal-bg: var(--color-workflow-display-normal-bg); - --color-workflow-display-normal-bg-line-pattern: var(--color-workflow-display-normal-bg-line-pattern); - --color-workflow-display-normal-border-1: var(--color-workflow-display-normal-border-1); - --color-workflow-display-normal-border-2: var(--color-workflow-display-normal-border-2); - --color-workflow-display-normal-vignette-color: var(--color-workflow-display-normal-vignette-color); - - --color-workflow-display-disabled-bg: var(--color-workflow-display-disabled-bg); - --color-workflow-display-disabled-bg-line-pattern: var(--color-workflow-display-disabled-bg-line-pattern); - --color-workflow-display-disabled-border-1: var(--color-workflow-display-disabled-border-1); - --color-workflow-display-disabled-border-2: var(--color-workflow-display-disabled-border-2); - --color-workflow-display-disabled-vignette-color: var(--color-workflow-display-disabled-vignette-color); - --color-workflow-display-disabled-outline: var(--color-workflow-display-disabled-outline); - - --color-workflow-workflow-progress-bg-1: var(--color-workflow-workflow-progress-bg-1); - --color-workflow-workflow-progress-bg-2: var(--color-workflow-workflow-progress-bg-2); - - --color-divider-subtle: var(--color-divider-subtle); - --color-divider-regular: var(--color-divider-regular); - --color-divider-deep: var(--color-divider-deep); - --color-divider-burn: var(--color-divider-burn); - --color-divider-intense: var(--color-divider-intense); - --color-divider-solid: var(--color-divider-solid); - --color-divider-solid-alt: var(--color-divider-solid-alt); - --color-divider-accent: var(--color-divider-accent); - - --color-state-base-hover: var(--color-state-base-hover); - --color-state-base-active: var(--color-state-base-active); - --color-state-base-hover-alt: var(--color-state-base-hover-alt); - --color-state-base-handle: var(--color-state-base-handle); - --color-state-base-handle-hover: var(--color-state-base-handle-hover); - --color-state-base-hover-subtle: var(--color-state-base-hover-subtle); - - --color-state-accent-hover: var(--color-state-accent-hover); - --color-state-accent-active: var(--color-state-accent-active); - --color-state-accent-hover-alt: var(--color-state-accent-hover-alt); - --color-state-accent-solid: var(--color-state-accent-solid); - --color-state-accent-active-alt: var(--color-state-accent-active-alt); - - --color-state-destructive-hover: var(--color-state-destructive-hover); - --color-state-destructive-hover-alt: var(--color-state-destructive-hover-alt); - --color-state-destructive-active: var(--color-state-destructive-active); - --color-state-destructive-solid: var(--color-state-destructive-solid); - --color-state-destructive-border: var(--color-state-destructive-border); - --color-state-destructive-hover-transparent: var(--color-state-destructive-hover-transparent); - - --color-state-success-hover: var(--color-state-success-hover); - --color-state-success-hover-alt: var(--color-state-success-hover-alt); - --color-state-success-active: var(--color-state-success-active); - --color-state-success-solid: var(--color-state-success-solid); - - --color-state-warning-hover: var(--color-state-warning-hover); - --color-state-warning-hover-alt: var(--color-state-warning-hover-alt); - --color-state-warning-active: var(--color-state-warning-active); - --color-state-warning-solid: var(--color-state-warning-solid); - --color-state-warning-hover-transparent: var(--color-state-warning-hover-transparent); - - --color-effects-highlight: var(--color-effects-highlight); - --color-effects-highlight-lightmode-off: var(--color-effects-highlight-lightmode-off); - --color-effects-image-frame: var(--color-effects-image-frame); - --color-effects-icon-border: var(--color-effects-icon-border); + --color-third-party-LangChain: var(--color-third-party-LangChain); + --color-third-party-Langfuse: var(--color-third-party-Langfuse); + --color-third-party-Github: var(--color-third-party-Github); + --color-third-party-Github-tertiary: var(--color-third-party-Github-tertiary); + --color-third-party-Github-secondary: var(--color-third-party-Github-secondary); + --color-third-party-aws: var(--color-third-party-aws); + --color-third-party-aws-alt: var(--color-third-party-aws-alt); + --color-third-party-model-bg-openai: var(--color-third-party-model-bg-openai); + --color-third-party-model-bg-anthropic: var(--color-third-party-model-bg-anthropic); + --color-third-party-model-bg-default: var(--color-third-party-model-bg-default); --color-util-colors-orange-dark-orange-dark-50: var(--color-util-colors-orange-dark-orange-dark-50); --color-util-colors-orange-dark-orange-dark-100: var(--color-util-colors-orange-dark-orange-dark-100); @@ -640,15 +682,6 @@ --color-util-colors-blue-light-blue-light-600: var(--color-util-colors-blue-light-blue-light-600); --color-util-colors-blue-light-blue-light-700: var(--color-util-colors-blue-light-blue-light-700); - --color-util-colors-gray-blue-gray-blue-50: var(--color-util-colors-gray-blue-gray-blue-50); - --color-util-colors-gray-blue-gray-blue-100: var(--color-util-colors-gray-blue-gray-blue-100); - --color-util-colors-gray-blue-gray-blue-200: var(--color-util-colors-gray-blue-gray-blue-200); - --color-util-colors-gray-blue-gray-blue-300: var(--color-util-colors-gray-blue-gray-blue-300); - --color-util-colors-gray-blue-gray-blue-400: var(--color-util-colors-gray-blue-gray-blue-400); - --color-util-colors-gray-blue-gray-blue-500: var(--color-util-colors-gray-blue-gray-blue-500); - --color-util-colors-gray-blue-gray-blue-600: var(--color-util-colors-gray-blue-gray-blue-600); - --color-util-colors-gray-blue-gray-blue-700: var(--color-util-colors-gray-blue-gray-blue-700); - --color-util-colors-blue-brand-blue-brand-50: var(--color-util-colors-blue-brand-blue-brand-50); --color-util-colors-blue-brand-blue-brand-100: var(--color-util-colors-blue-brand-blue-brand-100); --color-util-colors-blue-brand-blue-brand-200: var(--color-util-colors-blue-brand-blue-brand-200); @@ -658,6 +691,15 @@ --color-util-colors-blue-brand-blue-brand-600: var(--color-util-colors-blue-brand-blue-brand-600); --color-util-colors-blue-brand-blue-brand-700: var(--color-util-colors-blue-brand-blue-brand-700); + --color-util-colors-gray-blue-gray-blue-50: var(--color-util-colors-gray-blue-gray-blue-50); + --color-util-colors-gray-blue-gray-blue-100: var(--color-util-colors-gray-blue-gray-blue-100); + --color-util-colors-gray-blue-gray-blue-200: var(--color-util-colors-gray-blue-gray-blue-200); + --color-util-colors-gray-blue-gray-blue-300: var(--color-util-colors-gray-blue-gray-blue-300); + --color-util-colors-gray-blue-gray-blue-400: var(--color-util-colors-gray-blue-gray-blue-400); + --color-util-colors-gray-blue-gray-blue-500: var(--color-util-colors-gray-blue-gray-blue-500); + --color-util-colors-gray-blue-gray-blue-600: var(--color-util-colors-gray-blue-gray-blue-600); + --color-util-colors-gray-blue-gray-blue-700: var(--color-util-colors-gray-blue-gray-blue-700); + --color-util-colors-red-red-50: var(--color-util-colors-red-red-50); --color-util-colors-red-red-100: var(--color-util-colors-red-red-100); --color-util-colors-red-red-200: var(--color-util-colors-red-red-200); @@ -676,6 +718,15 @@ --color-util-colors-green-green-600: var(--color-util-colors-green-green-600); --color-util-colors-green-green-700: var(--color-util-colors-green-green-700); + --color-util-colors-green-light-green-light-50: var(--color-util-colors-green-light-green-light-50); + --color-util-colors-green-light-green-light-100: var(--color-util-colors-green-light-green-light-100); + --color-util-colors-green-light-green-light-200: var(--color-util-colors-green-light-green-light-200); + --color-util-colors-green-light-green-light-300: var(--color-util-colors-green-light-green-light-300); + --color-util-colors-green-light-green-light-400: var(--color-util-colors-green-light-green-light-400); + --color-util-colors-green-light-green-light-500: var(--color-util-colors-green-light-green-light-500); + --color-util-colors-green-light-green-light-600: var(--color-util-colors-green-light-green-light-600); + --color-util-colors-green-light-green-light-700: var(--color-util-colors-green-light-green-light-700); + --color-util-colors-warning-warning-50: var(--color-util-colors-warning-warning-50); --color-util-colors-warning-warning-100: var(--color-util-colors-warning-warning-100); --color-util-colors-warning-warning-200: var(--color-util-colors-warning-warning-200); @@ -730,15 +781,6 @@ --color-util-colors-gray-gray-600: var(--color-util-colors-gray-gray-600); --color-util-colors-gray-gray-700: var(--color-util-colors-gray-gray-700); - --color-util-colors-green-light-green-light-50: var(--color-util-colors-green-light-green-light-50); - --color-util-colors-green-light-green-light-100: var(--color-util-colors-green-light-green-light-100); - --color-util-colors-green-light-green-light-200: var(--color-util-colors-green-light-green-light-200); - --color-util-colors-green-light-green-light-300: var(--color-util-colors-green-light-green-light-300); - --color-util-colors-green-light-green-light-500: var(--color-util-colors-green-light-green-light-500); - --color-util-colors-green-light-green-light-400: var(--color-util-colors-green-light-green-light-400); - --color-util-colors-green-light-green-light-600: var(--color-util-colors-green-light-green-light-600); - --color-util-colors-green-light-green-light-700: var(--color-util-colors-green-light-green-light-700); - --color-util-colors-rose-rose-50: var(--color-util-colors-rose-rose-50); --color-util-colors-rose-rose-100: var(--color-util-colors-rose-rose-100); --color-util-colors-rose-rose-200: var(--color-util-colors-rose-rose-200); @@ -757,19 +799,9 @@ --color-util-colors-midnight-midnight-600: var(--color-util-colors-midnight-midnight-600); --color-util-colors-midnight-midnight-700: var(--color-util-colors-midnight-midnight-700); - --color-third-party-LangChain: var(--color-third-party-LangChain); - --color-third-party-Langfuse: var(--color-third-party-Langfuse); - --color-third-party-Github: var(--color-third-party-Github); - --color-third-party-Github-tertiary: var(--color-third-party-Github-tertiary); - --color-third-party-Github-secondary: var(--color-third-party-Github-secondary); - --color-third-party-model-bg-openai: var(--color-third-party-model-bg-openai); - --color-third-party-model-bg-anthropic: var(--color-third-party-model-bg-anthropic); - --color-third-party-model-bg-default: var(--color-third-party-model-bg-default); - - --color-third-party-aws: var(--color-third-party-aws); - --color-third-party-aws-alt: var(--color-third-party-aws-alt); - --color-saas-background: var(--color-saas-background); + --color-saas-background-inverted: var(--color-saas-background-inverted); + --color-saas-background-inverted-hover: var(--color-saas-background-inverted-hover); --color-saas-pricing-grid-bg: var(--color-saas-pricing-grid-bg); --color-saas-dify-blue-static: var(--color-saas-dify-blue-static); --color-saas-dify-blue-static-hover: var(--color-saas-dify-blue-static-hover); @@ -777,10 +809,21 @@ --color-saas-dify-blue-inverted: var(--color-saas-dify-blue-inverted); --color-saas-dify-blue-inverted-dimmed: var(--color-saas-dify-blue-inverted-dimmed); - --color-saas-background-inverted: var(--color-saas-background-inverted); - --color-saas-background-inverted-hover: var(--color-saas-background-inverted-hover); - --color-dify-logo-blue: var(--color-dify-logo-blue); --color-dify-logo-black: var(--color-dify-logo-black); + --color-dify-logo-outline-1: var(--color-dify-logo-outline-1); + --color-dify-logo-outline-2: var(--color-dify-logo-outline-2); + + --color-brand-color-opacity-50: var(--color-brand-color-opacity-50); + --color-brand-color-opacity-100: var(--color-brand-color-opacity-100); + --color-brand-color-opacity-200: var(--color-brand-color-opacity-200); + --color-brand-color-opacity-300: var(--color-brand-color-opacity-300); + --color-brand-color-opacity-400: var(--color-brand-color-opacity-400); + --color-brand-color-opacity-500: var(--color-brand-color-opacity-500); + --color-brand-color-opacity-600: var(--color-brand-color-opacity-600); + --color-brand-color-opacity-700: var(--color-brand-color-opacity-700); + --color-brand-color-opacity-800: var(--color-brand-color-opacity-800); + --color-brand-color-opacity-900: var(--color-brand-color-opacity-900); + --color-brand-color-opacity-1000: var(--color-brand-color-opacity-1000); } From 62cb5b586504b730731e28e743c038b72950058e Mon Sep 17 00:00:00 2001 From: WH-2099 Date: Tue, 30 Jun 2026 16:19:58 +0800 Subject: [PATCH 08/54] fix(api): scope nested resource lookups by owner refs (#38177) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/controllers/console/app/annotation.py | 35 ++- api/controllers/console/app/audio.py | 13 +- api/controllers/console/app/generator.py | 5 +- api/controllers/console/app/mcp_server.py | 13 +- api/controllers/console/app/workflow.py | 9 +- .../console/datasets/datasets_document.py | 4 +- .../console/datasets/datasets_segments.py | 149 ++++-------- .../rag_pipeline/rag_pipeline_workflow.py | 9 +- api/controllers/console/explore/audio.py | 13 +- api/controllers/console/explore/trial.py | 11 +- api/controllers/console/tag/tags.py | 3 +- api/controllers/service_api/app/annotation.py | 11 +- api/controllers/service_api/app/audio.py | 11 +- .../service_api/dataset/dataset.py | 8 +- .../service_api/dataset/segment.py | 101 +++----- api/controllers/web/audio.py | 11 +- api/core/llm_generator/llm_generator.py | 8 +- api/services/annotation_service.py | 98 ++++---- api/services/app_ref_service.py | 70 ++++++ api/services/audio_service.py | 19 +- api/services/dataset_ref_service.py | 56 +++++ api/services/dataset_service.py | 52 +++- api/services/rag_pipeline/rag_pipeline.py | 34 ++- api/services/tag_service.py | 34 ++- api/services/workflow_ref_service.py | 26 ++ api/services/workflow_service.py | 32 ++- .../service_api/dataset/test_dataset.py | 2 +- .../services/test_annotation_service.py | 41 +++- .../services/test_audio_service_db.py | 19 +- .../services/test_dataset_service_document.py | 12 +- .../services/test_workflow_service.py | 26 +- .../workflow/test_workflow_deletion.py | 19 +- .../console/app/test_annotation_api.py | 127 ++++++++++ .../controllers/console/app/test_audio.py | 28 +++ .../console/app/test_generator_api.py | 16 +- .../console/app/test_mcp_server_response.py | 47 ++++ .../datasets/test_datasets_document.py | 1 + .../datasets/test_datasets_segments.py | 224 +++++++++++++++++- .../controllers/console/explore/test_audio.py | 18 +- .../controllers/console/explore/test_trial.py | 22 ++ .../controllers/console/tag/test_tags.py | 29 +++ .../service_api/app/test_annotation.py | 6 +- .../controllers/service_api/app/test_audio.py | 30 ++- .../dataset/test_dataset_segment.py | 96 ++++---- .../unit_tests/controllers/web/test_audio.py | 19 ++ .../core/llm_generator/test_llm_generator.py | 21 ++ .../rag_pipeline/test_rag_pipeline_service.py | 55 ++++- .../services/test_annotation_service.py | 157 +++++------- .../unit_tests/services/test_audio_service.py | 26 +- .../services/test_dataset_service_document.py | 35 +++ .../services/test_dataset_service_segment.py | 80 +++++++ .../unit_tests/services/test_tag_service.py | 67 +++++- .../services/test_workflow_service.py | 80 ++++++- scripts/stress-test/test_setup_scripts.py | 4 +- 54 files changed, 1636 insertions(+), 506 deletions(-) create mode 100644 api/services/app_ref_service.py create mode 100644 api/services/dataset_ref_service.py create mode 100644 api/services/workflow_ref_service.py diff --git a/api/controllers/console/app/annotation.py b/api/controllers/console/app/annotation.py index 48fb4aedc63..d03bd15046b 100644 --- a/api/controllers/console/app/annotation.py +++ b/api/controllers/console/app/annotation.py @@ -4,6 +4,8 @@ from uuid import UUID from flask import abort, make_response, request from flask_restx import Resource from pydantic import BaseModel, Field, TypeAdapter, field_validator +from sqlalchemy import select +from werkzeug.exceptions import NotFound from controllers.common.errors import NoFileUploadedError, TooManyFilesError from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models @@ -30,7 +32,8 @@ from fields.annotation_fields import ( ) from fields.base import ResponseModel from libs.helper import uuid_value -from libs.login import login_required +from libs.login import current_account_with_tenant, login_required +from models.model import App from services.annotation_service import ( AppAnnotationService, EnableAnnotationArgs, @@ -38,6 +41,17 @@ from services.annotation_service import ( UpdateAnnotationSettingArgs, UpsertAnnotationArgs, ) +from services.app_ref_service import AppRef, AppRefService + + +def _get_app_ref(app_id: str) -> AppRef: + _, current_tenant_id = current_account_with_tenant() + app = db.session.scalar( + select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) + ) + if app is None: + raise NotFound("App not found") + return AppRefService.create_app_ref(app) class AnnotationReplyPayload(BaseModel): @@ -330,7 +344,8 @@ class AnnotationApi(Resource): "message": "annotation_ids are required if the parameter is provided.", }, 400 - AppAnnotationService.delete_app_annotations_in_batch(str(app_id), annotation_ids) + app_ref = _get_app_ref(str(app_id)) + AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids) return "", 204 # If no annotation_ids are provided, handle clearing all annotations else: @@ -389,9 +404,9 @@ class AnnotationUpdateDeleteApi(Resource): update_args["answer"] = args.answer if args.question is not None: update_args["question"] = args.question - annotation = AppAnnotationService.update_app_annotation_directly( - update_args, str(app_id), str(annotation_id), db.session - ) + app_ref = _get_app_ref(str(app_id)) + annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) + annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session) return Annotation.model_validate(annotation, from_attributes=True).model_dump(mode="json") @setup_required @@ -401,7 +416,9 @@ class AnnotationUpdateDeleteApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @console_ns.response(204, "Annotation deleted successfully") def delete(self, app_id: UUID, annotation_id: UUID): - AppAnnotationService.delete_app_annotation(str(app_id), str(annotation_id), db.session) + app_ref = _get_app_ref(str(app_id)) + annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) + AppAnnotationService.delete_app_annotation(annotation_ref, db.session) return "", 204 @@ -514,8 +531,12 @@ class AnnotationHitHistoryListApi(Resource): def get(self, app_id: UUID, annotation_id: UUID): page = request.args.get("page", default=1, type=int) limit = request.args.get("limit", default=20, type=int) + app_ref = _get_app_ref(str(app_id)) + annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) annotation_hit_history_list, total = AppAnnotationService.get_annotation_hit_histories( - str(app_id), str(annotation_id), page, limit + annotation_ref, + page, + limit, ) history_models = TypeAdapter(list[AnnotationHitHistory]).validate_python( annotation_hit_history_list, from_attributes=True diff --git a/api/controllers/console/app/audio.py b/api/controllers/console/app/audio.py index b66c97c274c..33e0efeef51 100644 --- a/api/controllers/console/app/audio.py +++ b/api/controllers/console/app/audio.py @@ -32,8 +32,9 @@ from controllers.console.wraps import ( from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError -from libs.login import login_required +from libs.login import current_user, login_required from models import App, AppMode +from services.app_ref_service import AppRefService from services.audio_service import AudioService from services.errors.audio import ( AudioTooLargeServiceError, @@ -140,13 +141,21 @@ class ChatMessageTextApi(Resource): def post(self, app_model: App): try: payload = TextToSpeechPayload.model_validate(console_ns.payload) + message_ref = None + if payload.message_id: + app_ref = AppRefService.create_app_ref(app_model) + message_ref = AppRefService.create_message_ref( + app_ref, + payload.message_id, + account_id=current_user.id, + ) response = AudioService.transcript_tts( app_model=app_model, session=db.session, text=payload.text, voice=payload.voice, - message_id=payload.message_id, + message_ref=message_ref, is_draft=True, ) return response diff --git a/api/controllers/console/app/generator.py b/api/controllers/console/app/generator.py index 1cf0ec82eb7..634b1dc4151 100644 --- a/api/controllers/console/app/generator.py +++ b/api/controllers/console/app/generator.py @@ -3,6 +3,7 @@ from typing import Any, Literal from flask_restx import Resource from pydantic import BaseModel, Field, RootModel +from sqlalchemy import select from sqlalchemy.orm import Session from controllers.common.fields import SimpleDataResponse @@ -216,7 +217,9 @@ class InstructionGenerateApi(Resource): try: # Generate from nothing for a workflow node if (args.current in (code_template, "")) and args.node_id != "": - app = session.get(App, args.flow_id) + app = session.scalar( + select(App).where(App.id == args.flow_id, App.tenant_id == current_tenant_id).limit(1) + ) if not app: return {"error": f"app {args.flow_id} not found"}, 400 workflow = WorkflowService().get_draft_workflow(app_model=app, session=session) diff --git a/api/controllers/console/app/mcp_server.py b/api/controllers/console/app/mcp_server.py index 6d6a56b5e1d..dacd386f0b0 100644 --- a/api/controllers/console/app/mcp_server.py +++ b/api/controllers/console/app/mcp_server.py @@ -26,6 +26,7 @@ from libs.helper import to_timestamp from libs.login import login_required from models.enums import AppMCPServerStatus from models.model import App, AppMCPServer +from services.app_ref_service import AppRefService class MCPServerCreatePayload(BaseModel): @@ -146,7 +147,17 @@ class AppMCPServerController(Resource): @get_app_model def put(self, app_model: App): payload = MCPServerUpdatePayload.model_validate(console_ns.payload or {}) - server = db.session.get(AppMCPServer, payload.id) + app_ref = AppRefService.create_app_ref(app_model) + server_ref = AppRefService.create_mcp_server_ref(app_ref, payload.id) + server = db.session.scalar( + select(AppMCPServer) + .where( + AppMCPServer.id == server_ref.server_id, + AppMCPServer.tenant_id == server_ref.tenant_id, + AppMCPServer.app_id == server_ref.app_id, + ) + .limit(1) + ) if not server: raise NotFound() diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index aff32035233..c56bfb63478 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -78,6 +78,7 @@ from repositories.workflow_collaboration_repository import WORKFLOW_ONLINE_USERS from services.app_generate_service import AppGenerateService from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError from services.errors.llm import InvokeRateLimitError +from services.workflow_ref_service import WorkflowRefService from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService logger = logging.getLogger(__name__) @@ -1406,15 +1407,15 @@ class WorkflowByIdApi(Resource): return {"message": "No valid fields to update"}, 400 workflow_service = WorkflowService() + workflow_ref = WorkflowRefService.create_app_workflow_ref(app_model, workflow_id) # Create a session and manage the transaction with sessionmaker(db.engine, expire_on_commit=False).begin() as session: workflow = workflow_service.update_workflow( session=session, - workflow_id=workflow_id, - tenant_id=app_model.tenant_id, account_id=current_user.id, data=update_data, + workflow_ref=workflow_ref, ) if not workflow: @@ -1434,12 +1435,14 @@ class WorkflowByIdApi(Resource): Delete workflow """ workflow_service = WorkflowService() + workflow_ref = WorkflowRefService.create_app_workflow_ref(app_model, workflow_id) # Create a session and manage the transaction with sessionmaker(db.engine).begin() as session: try: workflow_service.delete_workflow( - session=session, workflow_id=workflow_id, tenant_id=app_model.tenant_id + session=session, + workflow_ref=workflow_ref, ) except WorkflowInUseError as e: abort(400, description=str(e)) diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index 37f64f54d24..afa617535e1 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -49,6 +49,7 @@ from libs.login import login_required from models import Account, DatasetProcessRule, Document, DocumentSegment, UploadFile from models.dataset import DocumentPipelineExecutionLog from models.enums import IndexingStatus, SegmentStatus +from services.dataset_ref_service import DatasetRefService from services.dataset_service import DatasetService, DocumentService from services.entities.knowledge_entities.knowledge_entities import KnowledgeConfig, ProcessRule, RetrievalModel from services.file_service import FileService @@ -474,7 +475,8 @@ class DatasetDocumentListApi(Resource): try: document_ids = request.args.getlist("document_id") - DocumentService.delete_documents(dataset, document_ids, db.session) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + DocumentService.delete_documents(dataset_ref, document_ids, dataset.doc_form, db.session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 1a6f4c7a712..40ae9b207fd 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -58,8 +58,9 @@ from graphon.model_runtime.entities.model_entities import ModelType from libs.helper import dump_response, escape_like_pattern from libs.login import login_required from models import Account -from models.dataset import ChildChunk, DocumentSegment +from models.dataset import Dataset, Document, DocumentSegment from models.model import UploadFile +from services.dataset_ref_service import DatasetRefService, SegmentRef from services.dataset_service import DatasetService, DocumentService, SegmentService from services.entities.knowledge_entities.knowledge_entities import ChildChunkUpdateArgs, SegmentUpdateArgs from services.errors.chunk import ChildChunkDeleteIndexError as ChildChunkDeleteIndexServiceError @@ -162,6 +163,21 @@ register_response_schema_models( ) +def _get_segment_for_document( + dataset: Dataset, document: Document, segment_id: str +) -> tuple[SegmentRef, DocumentSegment]: + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + document_ref = DatasetRefService.create_document_ref(dataset_ref, document) + if document_ref is None: + raise NotFound("Document not found.") + + segment_ref = DatasetRefService.create_segment_ref(document_ref, segment_id) + segment = SegmentService.get_segment_by_ref(segment_ref) + if not segment: + raise NotFound("Segment not found.") + return segment_ref, segment + + @console_ns.route("/datasets//documents//segments") class DatasetDocumentSegmentListApi(Resource): @console_ns.doc(params=SegmentDocParams.DATASET_DOCUMENT) @@ -465,6 +481,13 @@ class DatasetDocumentSegmentUpdateApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") + # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor + if not current_user.is_dataset_editor: + raise Forbidden() + try: + DatasetService.check_dataset_permission(dataset, current_user, db.session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: # check embedding model setting try: @@ -481,22 +504,8 @@ class DatasetDocumentSegmentUpdateApi(Resource): ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) - # check segment segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") - # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor - if not current_user.is_dataset_editor: - raise Forbidden() - try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) - except services.errors.account.NoPermissionError as e: - raise Forbidden(str(e)) + _, segment = _get_segment_for_document(dataset, document, segment_id_str) # validate args payload = SegmentUpdatePayload.model_validate(console_ns.payload or {}) payload_dict = payload.model_dump(exclude_none=True) @@ -541,15 +550,6 @@ class DatasetDocumentSegmentUpdateApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") - # check segment - segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() @@ -557,6 +557,8 @@ class DatasetDocumentSegmentUpdateApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) + segment_id_str = str(segment_id) + _, segment = _get_segment_for_document(dataset, document, segment_id_str) SegmentService.delete_segment(segment, document, dataset, db.session) return "", 204 @@ -663,17 +665,12 @@ class ChildChunkAddApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") - # check segment - segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") if not current_user.is_dataset_editor: raise Forbidden() + try: + DatasetService.check_dataset_permission(dataset, current_user, db.session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) # check embedding model setting if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: try: @@ -690,10 +687,8 @@ class ChildChunkAddApi(Resource): ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) - try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) - except services.errors.account.NoPermissionError as e: - raise Forbidden(str(e)) + segment_id_str = str(segment_id) + _, segment = _get_segment_for_document(dataset, document, segment_id_str) # validate args try: payload = ChildChunkCreatePayload.model_validate(console_ns.payload or {}) @@ -723,15 +718,8 @@ class ChildChunkAddApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") - # check segment segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") + _get_segment_for_document(dataset, document, segment_id_str) args = query_params_from_request(ChildChunkListQuery, use_defaults_for_malformed_ints=True) page = args.page @@ -780,15 +768,6 @@ class ChildChunkAddApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") - # check segment - segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() @@ -796,6 +775,8 @@ class ChildChunkAddApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) + segment_id_str = str(segment_id) + _, segment = _get_segment_for_document(dataset, document, segment_id_str) # validate args payload = ChildChunkBatchUpdatePayload.model_validate(console_ns.payload or {}) try: @@ -839,29 +820,6 @@ class ChildChunkUpdateApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") - # check segment - segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") - # check child chunk - child_chunk_id_str = str(child_chunk_id) - child_chunk = db.session.scalar( - select(ChildChunk) - .where( - ChildChunk.id == child_chunk_id_str, - ChildChunk.tenant_id == current_tenant_id, - ChildChunk.segment_id == segment.id, - ChildChunk.document_id == document_id_str, - ) - .limit(1) - ) - if not child_chunk: - raise NotFound("Child chunk not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() @@ -869,6 +827,12 @@ class ChildChunkUpdateApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) + segment_id_str = str(segment_id) + segment_ref, _ = _get_segment_for_document(dataset, document, segment_id_str) + child_chunk_id_str = str(child_chunk_id) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + if not child_chunk: + raise NotFound("Child chunk not found.") try: SegmentService.delete_child_chunk(child_chunk, dataset, db.session) except ChildChunkDeleteIndexServiceError as e: @@ -907,29 +871,6 @@ class ChildChunkUpdateApi(Resource): document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) if not document: raise NotFound("Document not found.") - # check segment - segment_id_str = str(segment_id) - segment = db.session.scalar( - select(DocumentSegment) - .where(DocumentSegment.id == segment_id_str, DocumentSegment.tenant_id == current_tenant_id) - .limit(1) - ) - if not segment: - raise NotFound("Segment not found.") - # check child chunk - child_chunk_id_str = str(child_chunk_id) - child_chunk = db.session.scalar( - select(ChildChunk) - .where( - ChildChunk.id == child_chunk_id_str, - ChildChunk.tenant_id == current_tenant_id, - ChildChunk.segment_id == segment.id, - ChildChunk.document_id == document_id_str, - ) - .limit(1) - ) - if not child_chunk: - raise NotFound("Child chunk not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() @@ -937,6 +878,12 @@ class ChildChunkUpdateApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) + segment_id_str = str(segment_id) + segment_ref, segment = _get_segment_for_document(dataset, document, segment_id_str) + child_chunk_id_str = str(child_chunk_id) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + if not child_chunk: + raise NotFound("Child chunk not found.") # validate args try: payload = ChildChunkUpdatePayload.model_validate(console_ns.payload or {}) diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py index fdc55ea9737..e0a6ee0a83e 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py @@ -64,6 +64,7 @@ from services.rag_pipeline.pipeline_generate_service import PipelineGenerateServ from services.rag_pipeline.rag_pipeline import RagPipelineService from services.rag_pipeline.rag_pipeline_manage_service import RagPipelineManageService from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService +from services.workflow_ref_service import WorkflowRefService from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService logger = logging.getLogger(__name__) @@ -738,15 +739,15 @@ class RagPipelineByIdApi(Resource): return {"message": "No valid fields to update"}, 400 rag_pipeline_service = RagPipelineService() + workflow_ref = WorkflowRefService.create_pipeline_workflow_ref(pipeline, workflow_id) # Create a session and manage the transaction with sessionmaker(db.engine, expire_on_commit=False).begin() as session: workflow = rag_pipeline_service.update_workflow( session=session, - workflow_id=workflow_id, - tenant_id=pipeline.tenant_id, account_id=current_user.id, data=update_data, + workflow_ref=workflow_ref, ) if not workflow: @@ -769,13 +770,13 @@ class RagPipelineByIdApi(Resource): abort(400, description=f"Cannot delete workflow that is currently in use by pipeline '{pipeline.id}'") workflow_service = WorkflowService() + workflow_ref = WorkflowRefService.create_pipeline_workflow_ref(pipeline, workflow_id) with sessionmaker(db.engine).begin() as session: try: workflow_service.delete_workflow( session=session, - workflow_id=workflow_id, - tenant_id=pipeline.tenant_id, + workflow_ref=workflow_ref, ) except WorkflowInUseError as e: abort(400, description=str(e)) diff --git a/api/controllers/console/explore/audio.py b/api/controllers/console/explore/audio.py index c2104ccfc61..c0b86c19e43 100644 --- a/api/controllers/console/explore/audio.py +++ b/api/controllers/console/explore/audio.py @@ -22,7 +22,9 @@ from controllers.console.explore.wraps import InstalledAppResource from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError +from libs.login import current_account_with_tenant from models.model import InstalledApp +from services.app_ref_service import AppRefService from services.audio_service import AudioService from services.errors.audio import ( AudioTooLargeServiceError, @@ -99,13 +101,22 @@ class ChatTextApi(InstalledAppResource): message_id = payload.message_id text = payload.text voice = payload.voice + message_ref = None + if message_id: + current_user, _ = current_account_with_tenant() + app_ref = AppRefService.create_app_ref(app_model) + message_ref = AppRefService.create_message_ref( + app_ref, + message_id, + account_id=current_user.id, + ) response = AudioService.transcript_tts( app_model=app_model, session=db.session, text=text, voice=voice, - message_id=message_id, + message_ref=message_ref, ) return response except services.errors.app_model_config.AppModelConfigBrokenError: diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index 6aef9129780..9cf9aee66d9 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -81,6 +81,7 @@ from models.account import TenantStatus from models.model import AppMode, Site from models.workflow import Workflow from services.app_generate_service import AppGenerateService +from services.app_ref_service import AppRefService from services.app_service import AppService from services.audio_service import AudioService from services.dataset_service import DatasetService @@ -414,6 +415,14 @@ class TrialChatTextApi(TrialAppResource): message_id = request_data.message_id text = request_data.text voice = request_data.voice + message_ref = None + if message_id: + app_ref = AppRefService.create_app_ref(app_model) + message_ref = AppRefService.create_message_ref( + app_ref, + message_id, + account_id=current_user.id, + ) # Get IDs before they might be detached from session app_id = app_model.id @@ -424,7 +433,7 @@ class TrialChatTextApi(TrialAppResource): session=db.session, text=text, voice=voice, - message_id=message_id, + message_ref=message_ref, ) RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) return response diff --git a/api/controllers/console/tag/tags.py b/api/controllers/console/tag/tags.py index 1af5b113f1d..8d1f7491f0f 100644 --- a/api/controllers/console/tag/tags.py +++ b/api/controllers/console/tag/tags.py @@ -117,7 +117,8 @@ def _enforce_snippet_tag_rbac_by_tag_id(tag_id: str) -> None: if not dify_config.RBAC_ENABLED: return - tag_type = db.session.scalar(select(Tag.type).where(Tag.id == tag_id).limit(1)) + _, current_tenant_id = current_account_with_tenant() + tag_type = db.session.scalar(select(Tag.type).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id).limit(1)) _enforce_snippet_tag_rbac_if_needed(tag_type) diff --git a/api/controllers/service_api/app/annotation.py b/api/controllers/service_api/app/annotation.py index 8a57ec9818a..0fbf8125ed9 100644 --- a/api/controllers/service_api/app/annotation.py +++ b/api/controllers/service_api/app/annotation.py @@ -21,6 +21,7 @@ from services.annotation_service import ( InsertAnnotationArgs, UpdateAnnotationArgs, ) +from services.app_ref_service import AppRefService class AnnotationCreatePayload(BaseModel): @@ -282,9 +283,9 @@ class AnnotationUpdateDeleteApi(Resource): """Update an existing annotation.""" payload = AnnotationCreatePayload.model_validate(service_api_ns.payload or {}) update_args: UpdateAnnotationArgs = {"question": payload.question, "answer": payload.answer} - annotation = AppAnnotationService.update_app_annotation_directly( - update_args, app_model.id, str(annotation_id), db.session - ) + app_ref = AppRefService.create_app_ref(app_model) + annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) + annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session) response = Annotation.model_validate(annotation, from_attributes=True) return response.model_dump(mode="json") @@ -313,5 +314,7 @@ class AnnotationUpdateDeleteApi(Resource): @edit_permission_required def delete(self, app_model: App, annotation_id: UUID): """Delete an annotation.""" - AppAnnotationService.delete_app_annotation(app_model.id, str(annotation_id), db.session) + app_ref = AppRefService.create_app_ref(app_model) + annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) + AppAnnotationService.delete_app_annotation(annotation_ref, db.session) return "", 204 diff --git a/api/controllers/service_api/app/audio.py b/api/controllers/service_api/app/audio.py index 59ed4b4a4b1..53b31c8e6c4 100644 --- a/api/controllers/service_api/app/audio.py +++ b/api/controllers/service_api/app/audio.py @@ -26,6 +26,7 @@ from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotIni from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from models.model import App, EndUser +from services.app_ref_service import AppRefService from services.audio_service import AudioService from services.errors.audio import ( AudioTooLargeServiceError, @@ -177,13 +178,21 @@ class TextApi(Resource): message_id = payload.message_id text = payload.text voice = payload.voice + message_ref = None + if message_id: + app_ref = AppRefService.create_app_ref(app_model) + message_ref = AppRefService.create_message_ref( + app_ref, + message_id, + end_user_id=end_user.id, + ) response = AudioService.transcript_tts( app_model=app_model, session=db.session, text=text, voice=voice, end_user=end_user.external_user_id, - message_id=message_id, + message_ref=message_ref, ) return response diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index e903e92e7a6..0d52b7a25e6 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -943,9 +943,11 @@ class DatasetTagsApi(DatasetApiResource): payload = TagUpdatePayload.model_validate(service_api_ns.payload or {}) tag_id = payload.tag_id - tag = TagService.update_tags(UpdateTagServicePayload(name=payload.name), tag_id, db.session) + tag = TagService.update_tags( + UpdateTagServicePayload(name=payload.name), tag_id, db.session, tag_type=TagType.KNOWLEDGE + ) - binding_count = TagService.get_tag_binding_count(tag_id, db.session) + binding_count = TagService.get_tag_binding_count(tag_id, db.session, tag_type=TagType.KNOWLEDGE) response = dump_response( KnowledgeTagResponse, @@ -975,7 +977,7 @@ class DatasetTagsApi(DatasetApiResource): def delete(self, _): """Delete a knowledge type tag.""" payload = TagDeletePayload.model_validate(service_api_ns.payload or {}) - TagService.delete_tag(payload.tag_id, db.session) + TagService.delete_tag(payload.tag_id, db.session, tag_type=TagType.KNOWLEDGE) return "", 204 diff --git a/api/controllers/service_api/dataset/segment.py b/api/controllers/service_api/dataset/segment.py index 7b0e31952c7..41fbc709fdd 100644 --- a/api/controllers/service_api/dataset/segment.py +++ b/api/controllers/service_api/dataset/segment.py @@ -37,7 +37,8 @@ from fields.segment_fields import ( from graphon.model_runtime.entities.model_entities import ModelType from libs.helper import dump_response from libs.login import current_account_with_tenant -from models.dataset import Dataset, DocumentSegment +from models.dataset import Dataset, Document, DocumentSegment +from services.dataset_ref_service import DatasetRefService, SegmentRef from services.dataset_service import DatasetService, DocumentService, SegmentService from services.entities.knowledge_entities.knowledge_entities import SegmentUpdateArgs from services.errors.chunk import ChildChunkDeleteIndexError, ChildChunkIndexingError @@ -127,6 +128,21 @@ register_response_schema_models( ) +def _get_segment_for_document( + dataset: Dataset, document: Document, segment_id: str +) -> tuple[SegmentRef, DocumentSegment]: + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + document_ref = DatasetRefService.create_document_ref(dataset_ref, document) + if document_ref is None: + raise NotFound("Document not found.") + + segment_ref = DatasetRefService.create_segment_ref(document_ref, segment_id) + segment = SegmentService.get_segment_by_ref(segment_ref) + if not segment: + raise NotFound("Segment not found.") + return segment_ref, segment + + @service_api_ns.route("/datasets//documents//segments") class SegmentApi(DatasetApiResource): """Resource for segments.""" @@ -339,7 +355,7 @@ class DatasetSegmentApi(DatasetApiResource): ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def delete(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): - _, current_tenant_id = current_account_with_tenant() + current_account_with_tenant() dataset_id_str = str(dataset_id) # check dataset dataset = db.session.scalar( @@ -355,12 +371,7 @@ class DatasetSegmentApi(DatasetApiResource): if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) - # check segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") + _, segment = _get_segment_for_document(dataset, document, segment_id_str) SegmentService.delete_segment(segment, document, dataset, db.session) return "", 204 @@ -419,12 +430,7 @@ class DatasetSegmentApi(DatasetApiResource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) segment_id_str = str(segment_id) - # check segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") + _, segment = _get_segment_for_document(dataset, document, segment_id_str) payload = SegmentUpdatePayload.model_validate(service_api_ns.payload or {}) @@ -463,7 +469,7 @@ class DatasetSegmentApi(DatasetApiResource): service_api_ns.models[SegmentDetailResponse.__name__], ) def get(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): - _, current_tenant_id = current_account_with_tenant() + current_account_with_tenant() dataset_id_str = str(dataset_id) # check dataset dataset = db.session.scalar( @@ -479,12 +485,7 @@ class DatasetSegmentApi(DatasetApiResource): if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) - # check segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") + _, segment = _get_segment_for_document(dataset, document, segment_id_str) summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) response = { @@ -546,12 +547,7 @@ class ChildChunkApi(DatasetApiResource): raise NotFound("Document not found.") segment_id_str = str(segment_id) - # check segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") + _, segment = _get_segment_for_document(dataset, document, segment_id_str) # check embedding model setting if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -605,7 +601,7 @@ class ChildChunkApi(DatasetApiResource): service_api_ns.models[ChildChunkListResponse.__name__], ) def get(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): - _, current_tenant_id = current_account_with_tenant() + current_account_with_tenant() """Get child chunks.""" dataset_id_str = str(dataset_id) # check dataset @@ -622,12 +618,7 @@ class ChildChunkApi(DatasetApiResource): raise NotFound("Document not found.") segment_id_str = str(segment_id) - # check segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") + _get_segment_for_document(dataset, document, segment_id_str) args = query_params_from_request(ChildChunkListQuery, use_defaults_for_malformed_ints=True) @@ -677,7 +668,7 @@ class DatasetChildChunkApi(DatasetApiResource): @cloud_edition_billing_knowledge_limit_check("add_segment", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def delete(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID, child_chunk_id: UUID): - _, current_tenant_id = current_account_with_tenant() + current_account_with_tenant() """Delete child chunk.""" dataset_id_str = str(dataset_id) # check dataset @@ -694,29 +685,14 @@ class DatasetChildChunkApi(DatasetApiResource): raise NotFound("Document not found.") segment_id_str = str(segment_id) - # check segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") - - # validate segment belongs to the specified document - if segment.document_id != document_id_str: - raise NotFound("Document not found.") + segment_ref, _ = _get_segment_for_document(dataset, document, segment_id_str) child_chunk_id_str = str(child_chunk_id) # check child chunk - child_chunk = SegmentService.get_child_chunk_by_id( - child_chunk_id=child_chunk_id_str, tenant_id=current_tenant_id, session=db.session - ) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) if not child_chunk: raise NotFound("Child chunk not found.") - # validate child chunk belongs to the specified segment - if child_chunk.segment_id != segment.id: - raise NotFound("Child chunk not found.") - try: SegmentService.delete_child_chunk(child_chunk, dataset, db.session) except ChildChunkDeleteIndexServiceError as e: @@ -753,7 +729,7 @@ class DatasetChildChunkApi(DatasetApiResource): @cloud_edition_billing_knowledge_limit_check("add_segment", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def patch(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID, child_chunk_id: UUID): - _, current_tenant_id = current_account_with_tenant() + current_account_with_tenant() """Update child chunk.""" dataset_id_str = str(dataset_id) # check dataset @@ -770,29 +746,14 @@ class DatasetChildChunkApi(DatasetApiResource): raise NotFound("Document not found.") segment_id_str = str(segment_id) - # get segment - segment = SegmentService.get_segment_by_id( - segment_id=segment_id_str, tenant_id=current_tenant_id, session=db.session - ) - if not segment: - raise NotFound("Segment not found.") - - # validate segment belongs to the specified document - if segment.document_id != document_id_str: - raise NotFound("Segment not found.") + segment_ref, segment = _get_segment_for_document(dataset, document, segment_id_str) child_chunk_id_str = str(child_chunk_id) # get child chunk - child_chunk = SegmentService.get_child_chunk_by_id( - child_chunk_id=child_chunk_id_str, tenant_id=current_tenant_id, session=db.session - ) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) if not child_chunk: raise NotFound("Child chunk not found.") - # validate child chunk belongs to the specified segment - if child_chunk.segment_id != segment.id: - raise NotFound("Child chunk not found.") - # validate args payload = ChildChunkUpdatePayload.model_validate(service_api_ns.payload or {}) diff --git a/api/controllers/web/audio.py b/api/controllers/web/audio.py index 801c1f5a629..47e72ff95a5 100644 --- a/api/controllers/web/audio.py +++ b/api/controllers/web/audio.py @@ -26,6 +26,7 @@ from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from libs.helper import uuid_value from models.model import App, EndUser +from services.app_ref_service import AppRefService from services.audio_service import AudioService from services.errors.audio import ( AudioTooLargeServiceError, @@ -130,13 +131,21 @@ class TextApi(WebApiResource): message_id = payload.message_id text = payload.text voice = payload.voice + message_ref = None + if message_id: + app_ref = AppRefService.create_app_ref(app_model) + message_ref = AppRefService.create_message_ref( + app_ref, + message_id, + end_user_id=end_user.id, + ) response = AudioService.transcript_tts( app_model=app_model, session=db.session, text=text, voice=voice, end_user=end_user.external_user_id, - message_id=message_id, + message_ref=message_ref, ) return response diff --git a/api/core/llm_generator/llm_generator.py b/api/core/llm_generator/llm_generator.py index b2073716d13..6b447c6ce42 100644 --- a/api/core/llm_generator/llm_generator.py +++ b/api/core/llm_generator/llm_generator.py @@ -498,7 +498,11 @@ class LLMGenerator: ideal_output: str | None, ): last_run: Message | None = db.session.scalar( - select(Message).where(Message.app_id == flow_id).order_by(Message.created_at.desc()).limit(1) + select(Message) + .join(App, App.id == Message.app_id) + .where(Message.app_id == flow_id, App.tenant_id == tenant_id) + .order_by(Message.created_at.desc()) + .limit(1) ) if not last_run: return LLMGenerator.__instruction_modify_common( @@ -540,7 +544,7 @@ class LLMGenerator: ): session = db.session() - app: App | None = session.scalar(select(App).where(App.id == flow_id).limit(1)) + app: App | None = session.scalar(select(App).where(App.id == flow_id, App.tenant_id == tenant_id).limit(1)) if not app: raise ValueError("App not found.") workflow = workflow_service.get_draft_workflow(app_model=app) diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index 4f69c4b44a9..080527d9769 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -14,6 +14,7 @@ from extensions.ext_redis import redis_client from libs.datetime_utils import naive_utc_now from libs.login import current_account_with_tenant from models.model import App, AppAnnotationHitHistory, AppAnnotationSetting, Message, MessageAnnotation +from services.app_ref_service import AnnotationRef, AppRef from services.feature_service import FeatureService from tasks.annotation.add_annotation_to_index_task import add_annotation_to_index_task from tasks.annotation.batch_import_annotations_task import batch_import_annotations_task @@ -88,6 +89,17 @@ class UpdateAnnotationSettingArgs(TypedDict): class AppAnnotationService: + @staticmethod + def _get_annotation_by_ref(annotation_ref: AnnotationRef, session: scoped_session) -> MessageAnnotation | None: + return session.scalar( + select(MessageAnnotation) + .where( + MessageAnnotation.id == annotation_ref.annotation_id, + MessageAnnotation.app_id == annotation_ref.app_id, + ) + .limit(1) + ) + @classmethod def up_insert_app_annotation_from_message(cls, args: UpsertAnnotationArgs, app_id: str) -> MessageAnnotation: # get app info @@ -302,18 +314,9 @@ class AppAnnotationService: @classmethod def update_app_annotation_directly( - cls, args: UpdateAnnotationArgs, app_id: str, annotation_id: str, session: scoped_session + cls, args: UpdateAnnotationArgs, annotation_ref: AnnotationRef, session: scoped_session ): - # get app info - _, current_tenant_id = current_account_with_tenant() - app = session.scalar( - select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) - ) - - if not app: - raise NotFound("App not found") - - annotation = session.get(MessageAnnotation, annotation_id) + annotation = cls._get_annotation_by_ref(annotation_ref, session) if not annotation: raise NotFound("Annotation not found") @@ -332,32 +335,23 @@ class AppAnnotationService: session.commit() # if annotation reply is enabled , add annotation to index app_annotation_setting = session.scalar( - select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) + select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == annotation_ref.app_id).limit(1) ) if app_annotation_setting: update_annotation_to_index_task.delay( annotation.id, annotation.question_text, - current_tenant_id, - app_id, + annotation_ref.tenant_id, + annotation_ref.app_id, app_annotation_setting.collection_binding_id, ) return annotation @classmethod - def delete_app_annotation(cls, app_id: str, annotation_id: str, session: scoped_session): - # get app info - _, current_tenant_id = current_account_with_tenant() - app = session.scalar( - select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) - ) - - if not app: - raise NotFound("App not found") - - annotation = session.get(MessageAnnotation, annotation_id) + def delete_app_annotation(cls, annotation_ref: AnnotationRef, session: scoped_session): + annotation = cls._get_annotation_by_ref(annotation_ref, session) if not annotation: raise NotFound("Annotation not found") @@ -365,7 +359,10 @@ class AppAnnotationService: session.delete(annotation) annotation_hit_histories = session.scalars( - select(AppAnnotationHitHistory).where(AppAnnotationHitHistory.annotation_id == annotation_id) + select(AppAnnotationHitHistory).where( + AppAnnotationHitHistory.app_id == annotation_ref.app_id, + AppAnnotationHitHistory.annotation_id == annotation_ref.annotation_id, + ) ).all() if annotation_hit_histories: for annotation_hit_history in annotation_hit_histories: @@ -374,30 +371,24 @@ class AppAnnotationService: session.commit() # if annotation reply is enabled , delete annotation index app_annotation_setting = session.scalar( - select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) + select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == annotation_ref.app_id).limit(1) ) if app_annotation_setting: delete_annotation_index_task.delay( - annotation.id, app_id, current_tenant_id, app_annotation_setting.collection_binding_id + annotation.id, + annotation_ref.app_id, + annotation_ref.tenant_id, + app_annotation_setting.collection_binding_id, ) @classmethod - def delete_app_annotations_in_batch(cls, app_id: str, annotation_ids: list[str]): - # get app info - _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( - select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) - ) - - if not app: - raise NotFound("App not found") - + def delete_app_annotations_in_batch(cls, app_ref: AppRef, annotation_ids: list[str]): # Fetch annotations and their settings in a single query annotations_to_delete = db.session.execute( select(MessageAnnotation, AppAnnotationSetting) .outerjoin(AppAnnotationSetting, MessageAnnotation.app_id == AppAnnotationSetting.app_id) - .where(MessageAnnotation.id.in_(annotation_ids)) + .where(MessageAnnotation.id.in_(annotation_ids), MessageAnnotation.app_id == app_ref.app_id) ).all() if not annotations_to_delete: @@ -408,19 +399,25 @@ class AppAnnotationService: # Step 2: Bulk delete hit histories in a single query db.session.execute( - delete(AppAnnotationHitHistory).where(AppAnnotationHitHistory.annotation_id.in_(annotation_ids_to_delete)) + delete(AppAnnotationHitHistory).where( + AppAnnotationHitHistory.app_id == app_ref.app_id, + AppAnnotationHitHistory.annotation_id.in_(annotation_ids_to_delete), + ) ) # Step 3: Trigger async tasks for search index deletion for annotation, annotation_setting in annotations_to_delete: if annotation_setting: delete_annotation_index_task.delay( - annotation.id, app_id, current_tenant_id, annotation_setting.collection_binding_id + annotation.id, app_ref.app_id, app_ref.tenant_id, annotation_setting.collection_binding_id ) # Step 4: Bulk delete annotations in a single query delete_result = db.session.execute( - delete(MessageAnnotation).where(MessageAnnotation.id.in_(annotation_ids_to_delete)) + delete(MessageAnnotation).where( + MessageAnnotation.id.in_(annotation_ids_to_delete), + MessageAnnotation.app_id == app_ref.app_id, + ) ) deleted_count = getattr(delete_result, "rowcount", 0) @@ -562,17 +559,8 @@ class AppAnnotationService: return {"job_id": job_id, "job_status": "waiting", "record_count": len(result)} @classmethod - def get_annotation_hit_histories(cls, app_id: str, annotation_id: str, page, limit): - _, current_tenant_id = current_account_with_tenant() - # get app info - app = db.session.scalar( - select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) - ) - - if not app: - raise NotFound("App not found") - - annotation = db.session.get(MessageAnnotation, annotation_id) + def get_annotation_hit_histories(cls, annotation_ref: AnnotationRef, page, limit): + annotation = cls._get_annotation_by_ref(annotation_ref, db.session) if not annotation: raise NotFound("Annotation not found") @@ -580,8 +568,8 @@ class AppAnnotationService: stmt = ( select(AppAnnotationHitHistory) .where( - AppAnnotationHitHistory.app_id == app_id, - AppAnnotationHitHistory.annotation_id == annotation_id, + AppAnnotationHitHistory.app_id == annotation_ref.app_id, + AppAnnotationHitHistory.annotation_id == annotation_ref.annotation_id, ) .order_by(AppAnnotationHitHistory.created_at.desc()) ) diff --git a/api/services/app_ref_service.py b/api/services/app_ref_service.py new file mode 100644 index 00000000000..5f794491811 --- /dev/null +++ b/api/services/app_ref_service.py @@ -0,0 +1,70 @@ +"""Typed resource references for app ownership chains.""" + +from typing import NamedTuple + +from models.model import App + + +class AppRef(NamedTuple): + """App identifiers used to scope downstream resource lookups.""" + + tenant_id: str + app_id: str + + +class MessageRef(NamedTuple): + """Message identifiers used to scope downstream resource lookups.""" + + tenant_id: str + app_id: str + message_id: str + end_user_id: str | None = None + account_id: str | None = None + + +class AnnotationRef(NamedTuple): + """Annotation identifiers used to scope downstream resource lookups.""" + + tenant_id: str + app_id: str + annotation_id: str + + +class AppMCPServerRef(NamedTuple): + """MCP server identifiers used to scope downstream resource lookups.""" + + tenant_id: str + app_id: str + server_id: str + + +class AppRefService: + """Factory helpers for app and child resource refs.""" + + @staticmethod + def create_app_ref(app: App) -> AppRef: + return AppRef(tenant_id=app.tenant_id, app_id=app.id) + + @staticmethod + def create_message_ref( + app_ref: AppRef, + message_id: str, + *, + end_user_id: str | None = None, + account_id: str | None = None, + ) -> MessageRef: + return MessageRef( + tenant_id=app_ref.tenant_id, + app_id=app_ref.app_id, + message_id=message_id, + end_user_id=end_user_id, + account_id=account_id, + ) + + @staticmethod + def create_annotation_ref(app_ref: AppRef, annotation_id: str) -> AnnotationRef: + return AnnotationRef(tenant_id=app_ref.tenant_id, app_id=app_ref.app_id, annotation_id=annotation_id) + + @staticmethod + def create_mcp_server_ref(app_ref: AppRef, server_id: str) -> AppMCPServerRef: + return AppMCPServerRef(tenant_id=app_ref.tenant_id, app_id=app_ref.app_id, server_id=server_id) diff --git a/api/services/audio_service.py b/api/services/audio_service.py index 14c5c0111e5..86c56e60a13 100644 --- a/api/services/audio_service.py +++ b/api/services/audio_service.py @@ -5,6 +5,7 @@ from collections.abc import Generator from typing import cast from flask import Response, stream_with_context +from sqlalchemy import select from sqlalchemy.orm import Session, scoped_session from werkzeug.datastructures import FileStorage @@ -13,6 +14,7 @@ from core.model_manager import ModelManager from graphon.model_runtime.entities.model_entities import ModelType from models.enums import MessageStatus from models.model import App, AppMode, Message +from services.app_ref_service import MessageRef from services.errors.audio import ( AudioTooLargeServiceError, NoAudioUploadedServiceError, @@ -29,6 +31,15 @@ logger = logging.getLogger(__name__) class AudioService: + @staticmethod + def _get_message_by_ref(session: Session | scoped_session, message_ref: MessageRef) -> Message | None: + stmt = select(Message).where(Message.id == message_ref.message_id, Message.app_id == message_ref.app_id) + if message_ref.end_user_id is not None: + stmt = stmt.where(Message.from_end_user_id == message_ref.end_user_id) + if message_ref.account_id is not None: + stmt = stmt.where(Message.from_account_id == message_ref.account_id) + return session.scalar(stmt.limit(1)) + @classmethod def transcript_asr(cls, app_model: App, file: FileStorage | None, end_user: str | None = None): if app_model.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: @@ -82,7 +93,7 @@ class AudioService: text: str | None = None, voice: str | None = None, end_user: str | None = None, - message_id: str | None = None, + message_ref: MessageRef | None = None, is_draft: bool = False, ): def invoke_tts(text_content: str, app_model: App, voice: str | None = None, is_draft: bool = False): @@ -129,12 +140,12 @@ class AudioService: except Exception as e: raise e - if message_id: + if message_ref: try: - uuid.UUID(message_id) + uuid.UUID(message_ref.message_id) except ValueError: return None - message = session.get(Message, message_id) + message = cls._get_message_by_ref(session, message_ref) if message is None: return None if message.answer == "" and message.status in {MessageStatus.NORMAL, MessageStatus.PAUSED}: diff --git a/api/services/dataset_ref_service.py b/api/services/dataset_ref_service.py new file mode 100644 index 00000000000..049a8d05759 --- /dev/null +++ b/api/services/dataset_ref_service.py @@ -0,0 +1,56 @@ +"""Typed resource references for dataset ownership chains.""" + +from typing import NamedTuple + +from models.dataset import Dataset, Document + + +class DatasetRef(NamedTuple): + """Dataset identifiers used to scope downstream resource lookups.""" + + tenant_id: str + dataset_id: str + + +class DocumentRef(NamedTuple): + """Document identifiers used to scope downstream resource lookups.""" + + tenant_id: str + dataset_id: str + document_id: str + + +class SegmentRef(NamedTuple): + """Segment identifiers used to scope downstream resource lookups.""" + + tenant_id: str + dataset_id: str + document_id: str + segment_id: str + + +class DatasetRefService: + """Factory helpers for dataset, document, and segment refs.""" + + @staticmethod + def create_dataset_ref(dataset: Dataset) -> DatasetRef: + return DatasetRef(tenant_id=dataset.tenant_id, dataset_id=dataset.id) + + @staticmethod + def create_document_ref(dataset_ref: DatasetRef, document: Document) -> DocumentRef | None: + if document.tenant_id != dataset_ref.tenant_id or document.dataset_id != dataset_ref.dataset_id: + return None + return DocumentRef( + tenant_id=dataset_ref.tenant_id, + dataset_id=dataset_ref.dataset_id, + document_id=document.id, + ) + + @staticmethod + def create_segment_ref(document_ref: DocumentRef, segment_id: str) -> SegmentRef: + return SegmentRef( + tenant_id=document_ref.tenant_id, + dataset_id=document_ref.dataset_id, + document_id=document_ref.document_id, + segment_id=segment_id, + ) diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index 4dbcf372bb0..17e1531db5f 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -64,6 +64,7 @@ from models.model import UploadFile from models.provider_ids import ModelProviderID from models.source import DataSourceOauthBinding from models.workflow import Workflow +from services.dataset_ref_service import DatasetRef, SegmentRef from services.document_indexing_proxy.document_indexing_task_proxy import DocumentIndexingTaskProxy from services.document_indexing_proxy.duplicate_document_indexing_task_proxy import DuplicateDocumentIndexingTaskProxy from services.enterprise import rbac_service as enterprise_rbac_service @@ -1970,11 +1971,23 @@ class DocumentService: session.commit() @staticmethod - def delete_documents(dataset: Dataset, document_ids: list[str], session: scoped_session | Session): + def delete_documents( + dataset_ref: DatasetRef, + document_ids: list[str], + doc_form: str | None, + session: scoped_session | Session, + ): # Check if document_ids is not empty to avoid WHERE false condition if not document_ids or len(document_ids) == 0: return - documents = session.scalars(select(Document).where(Document.id.in_(document_ids))).all() + documents = session.scalars( + select(Document).where( + Document.id.in_(document_ids), + Document.tenant_id == dataset_ref.tenant_id, + Document.dataset_id == dataset_ref.dataset_id, + ) + ).all() + deleted_document_ids = [document.id for document in documents] file_ids = [ document.data_source_info_dict.get("upload_file_id", "") for document in documents @@ -1989,8 +2002,8 @@ class DocumentService: # Dispatch cleanup task after commit to avoid lock contention # Task cleans up segments, files, and vector indexes - if dataset.doc_form is not None: - batch_clean_document_task.delay(document_ids, dataset.id, dataset.doc_form, file_ids) + if deleted_document_ids and doc_form is not None: + batch_clean_document_task.delay(deleted_document_ids, dataset_ref.dataset_id, doc_form, file_ids) @staticmethod def rename_document(dataset_id: str, document_id: str, name: str, session: scoped_session | Session) -> Document: @@ -4120,6 +4133,22 @@ class SegmentService: ) return result if isinstance(result, ChildChunk) else None + @classmethod + def get_child_chunk_by_segment_ref(cls, child_chunk_id: str, segment_ref: SegmentRef) -> ChildChunk | None: + """Get a child chunk through the full tenant/dataset/document/segment chain.""" + result = db.session.scalar( + select(ChildChunk) + .where( + ChildChunk.id == child_chunk_id, + ChildChunk.tenant_id == segment_ref.tenant_id, + ChildChunk.dataset_id == segment_ref.dataset_id, + ChildChunk.document_id == segment_ref.document_id, + ChildChunk.segment_id == segment_ref.segment_id, + ) + .limit(1) + ) + return result if isinstance(result, ChildChunk) else None + @classmethod def get_segments( cls, @@ -4160,6 +4189,21 @@ class SegmentService: ) return result if isinstance(result, DocumentSegment) else None + @classmethod + def get_segment_by_ref(cls, segment_ref: SegmentRef) -> DocumentSegment | None: + """Get a segment through the full tenant/dataset/document ownership chain.""" + result = db.session.scalar( + select(DocumentSegment) + .where( + DocumentSegment.id == segment_ref.segment_id, + DocumentSegment.tenant_id == segment_ref.tenant_id, + DocumentSegment.dataset_id == segment_ref.dataset_id, + DocumentSegment.document_id == segment_ref.document_id, + ) + .limit(1) + ) + return result if isinstance(result, DocumentSegment) else None + @classmethod def get_segments_by_document_and_dataset( cls, diff --git a/api/services/rag_pipeline/rag_pipeline.py b/api/services/rag_pipeline/rag_pipeline.py index abab174b3d9..10112ef19a0 100644 --- a/api/services/rag_pipeline/rag_pipeline.py +++ b/api/services/rag_pipeline/rag_pipeline.py @@ -82,6 +82,7 @@ from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, from services.rag_pipeline.pipeline_template.pipeline_template_factory import PipelineTemplateRetrievalFactory from services.tools.builtin_tools_manage_service import BuiltinToolManageService from services.workflow_draft_variable_service import DraftVariableSaver, DraftVarLoader +from services.workflow_ref_service import WorkflowRef from services.workflow_restore import apply_published_workflow_snapshot_to_draft logger = logging.getLogger(__name__) @@ -984,8 +985,21 @@ class RagPipelineService: if invoke_from: if invoke_from.value == InvokeFrom.PUBLISHED_PIPELINE: document_id = get_system_segment(variable_pool, SystemVariableKey.DOCUMENT_ID) - if document_id: - document = db.session.get(Document, document_id.value) + dataset_id = get_system_segment(variable_pool, SystemVariableKey.DATASET_ID) + pipeline_id = get_system_segment(variable_pool, SystemVariableKey.APP_ID) + if document_id and dataset_id and pipeline_id: + document = db.session.scalar( + select(Document) + .join(Dataset, Dataset.id == Document.dataset_id) + .where( + Document.id == document_id.value, + Document.tenant_id == tenant_id, + Document.dataset_id == dataset_id.value, + Dataset.tenant_id == tenant_id, + Dataset.pipeline_id == pipeline_id.value, + ) + .limit(1) + ) if document: document.indexing_status = IndexingStatus.ERROR document.error = error @@ -995,19 +1009,27 @@ class RagPipelineService: return workflow_node_execution def update_workflow( - self, *, session: Session, workflow_id: str, tenant_id: str, account_id: str, data: dict[str, Any] + self, + *, + session: Session, + account_id: str, + data: dict[str, Any], + workflow_ref: WorkflowRef, ) -> Workflow | None: """ Update workflow attributes :param session: SQLAlchemy database session - :param workflow_id: Workflow ID - :param tenant_id: Tenant ID :param account_id: Account ID (for permission check) :param data: Dictionary containing fields to update + :param workflow_ref: Owner-bound workflow reference :return: Updated workflow or None if not found """ - stmt = select(Workflow).where(Workflow.id == workflow_id, Workflow.tenant_id == tenant_id) + stmt = select(Workflow).where( + Workflow.id == workflow_ref.workflow_id, + Workflow.tenant_id == workflow_ref.tenant_id, + Workflow.app_id == workflow_ref.owner_id, + ) workflow = session.scalar(stmt) if not workflow: diff --git a/api/services/tag_service.py b/api/services/tag_service.py index b144c15e85e..45deedb5b93 100644 --- a/api/services/tag_service.py +++ b/api/services/tag_service.py @@ -147,8 +147,14 @@ class TagService: return tag @staticmethod - def update_tags(payload: UpdateTagPayload, tag_id: str, session: scoped_session) -> Tag: - tag = session.scalar(select(Tag).where(Tag.id == tag_id).limit(1)) + def update_tags( + payload: UpdateTagPayload, tag_id: str, session: scoped_session, *, tag_type: TagType | None = None + ) -> Tag: + current_tenant_id = current_user.current_tenant_id + stmt = select(Tag).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id) + if tag_type is not None: + stmt = stmt.where(Tag.type == tag_type) + tag = session.scalar(stmt.limit(1)) if not tag: raise NotFound("Tag not found") if payload.name != tag.name: @@ -169,18 +175,32 @@ class TagService: return tag @staticmethod - def get_tag_binding_count(tag_id: str, session: scoped_session) -> int: - count = session.scalar(select(func.count(TagBinding.id)).where(TagBinding.tag_id == tag_id)) or 0 + def get_tag_binding_count(tag_id: str, session: scoped_session, *, tag_type: TagType | None = None) -> int: + current_tenant_id = current_user.current_tenant_id + stmt = ( + select(func.count(TagBinding.id)) + .join(Tag, Tag.id == TagBinding.tag_id) + .where(TagBinding.tag_id == tag_id, Tag.tenant_id == current_tenant_id) + ) + if tag_type is not None: + stmt = stmt.where(Tag.type == tag_type) + count = session.scalar(stmt) or 0 return count @staticmethod - def delete_tag(tag_id: str, session: scoped_session): - tag = session.scalar(select(Tag).where(Tag.id == tag_id).limit(1)) + def delete_tag(tag_id: str, session: scoped_session, *, tag_type: TagType | None = None): + current_tenant_id = current_user.current_tenant_id + stmt = select(Tag).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id) + if tag_type is not None: + stmt = stmt.where(Tag.type == tag_type) + tag = session.scalar(stmt.limit(1)) if not tag: raise NotFound("Tag not found") session.delete(tag) # delete tag binding - tag_bindings = session.scalars(select(TagBinding).where(TagBinding.tag_id == tag_id)).all() + tag_bindings = session.scalars( + select(TagBinding).where(TagBinding.tag_id == tag_id, TagBinding.tenant_id == current_tenant_id) + ).all() if tag_bindings: for tag_binding in tag_bindings: session.delete(tag_binding) diff --git a/api/services/workflow_ref_service.py b/api/services/workflow_ref_service.py new file mode 100644 index 00000000000..c28428c7156 --- /dev/null +++ b/api/services/workflow_ref_service.py @@ -0,0 +1,26 @@ +"""Typed resource references for workflow ownership chains.""" + +from typing import NamedTuple + +from models.dataset import Pipeline +from models.model import App + + +class WorkflowRef(NamedTuple): + """Workflow identifiers used to scope downstream resource lookups.""" + + tenant_id: str + owner_id: str + workflow_id: str + + +class WorkflowRefService: + """Factory helpers for app and RAG pipeline workflow refs.""" + + @staticmethod + def create_app_workflow_ref(app: App, workflow_id: str) -> WorkflowRef: + return WorkflowRef(tenant_id=app.tenant_id, owner_id=app.id, workflow_id=workflow_id) + + @staticmethod + def create_pipeline_workflow_ref(pipeline: Pipeline, workflow_id: str) -> WorkflowRef: + return WorkflowRef(tenant_id=pipeline.tenant_id, owner_id=pipeline.id, workflow_id=workflow_id) diff --git a/api/services/workflow_service.py b/api/services/workflow_service.py index 262ccc18f83..c3327f5787d 100644 --- a/api/services/workflow_service.py +++ b/api/services/workflow_service.py @@ -84,6 +84,7 @@ from services.errors.app import ( ) from services.human_input_service import HumanInputService from services.workflow.workflow_converter import WorkflowConverter +from services.workflow_ref_service import WorkflowRef from .errors.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError from .human_input_delivery_test_service import ( @@ -1593,19 +1594,27 @@ class WorkflowService: raise ValueError(f"Invalid HumanInput node data: {str(e)}") def update_workflow( - self, *, session: Session, workflow_id: str, tenant_id: str, account_id: str, data: dict[str, Any] + self, + *, + session: Session, + account_id: str, + data: dict[str, Any], + workflow_ref: WorkflowRef, ) -> Workflow | None: """ Update workflow attributes :param session: SQLAlchemy database session - :param workflow_id: Workflow ID - :param tenant_id: Tenant ID :param account_id: Account ID (for permission check) :param data: Dictionary containing fields to update + :param workflow_ref: Owner-bound workflow reference :return: Updated workflow or None if not found """ - stmt = select(Workflow).where(Workflow.id == workflow_id, Workflow.tenant_id == tenant_id) + stmt = select(Workflow).where( + Workflow.id == workflow_ref.workflow_id, + Workflow.tenant_id == workflow_ref.tenant_id, + Workflow.app_id == workflow_ref.owner_id, + ) workflow = session.scalar(stmt) if not workflow: @@ -1622,30 +1631,33 @@ class WorkflowService: return workflow - def delete_workflow(self, *, session: Session, workflow_id: str, tenant_id: str) -> bool: + def delete_workflow(self, *, session: Session, workflow_ref: WorkflowRef) -> bool: """ Delete a workflow :param session: SQLAlchemy database session - :param workflow_id: Workflow ID - :param tenant_id: Tenant ID + :param workflow_ref: Owner-bound workflow reference :return: True if successful :raises: ValueError if workflow not found :raises: WorkflowInUseError if workflow is in use :raises: DraftWorkflowDeletionError if workflow is a draft version """ - stmt = select(Workflow).where(Workflow.id == workflow_id, Workflow.tenant_id == tenant_id) + stmt = select(Workflow).where( + Workflow.id == workflow_ref.workflow_id, + Workflow.tenant_id == workflow_ref.tenant_id, + Workflow.app_id == workflow_ref.owner_id, + ) workflow = session.scalar(stmt) if not workflow: - raise ValueError(f"Workflow with ID {workflow_id} not found") + raise ValueError(f"Workflow with ID {workflow_ref.workflow_id} not found") # Check if workflow is a draft version if workflow.version == Workflow.VERSION_DRAFT: raise DraftWorkflowDeletionError("Cannot delete draft workflow versions") # Check if this workflow is currently referenced by an app - app_stmt = select(App).where(App.workflow_id == workflow_id) + app_stmt = select(App).where(App.workflow_id == workflow_ref.workflow_id) app = session.scalar(app_stmt) if app: # Cannot delete a workflow that's currently in use by an app diff --git a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py index 6e8af5ca43e..35e17a12b3a 100644 --- a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py +++ b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py @@ -1188,7 +1188,7 @@ class TestDatasetTagsApiDelete: result = api.delete(_=None) assert result == ("", 204) - mock_tag_svc.delete_tag.assert_called_once_with("tag-1", ANY) + mock_tag_svc.delete_tag.assert_called_once_with("tag-1", ANY, tag_type=TagType.KNOWLEDGE) @patch("libs.login.current_user") def test_delete_tag_forbidden(self, mock_current_user, app: Flask): diff --git a/api/tests/test_containers_integration_tests/services/test_annotation_service.py b/api/tests/test_containers_integration_tests/services/test_annotation_service.py index 6c6b9338d7d..94d72b19be8 100644 --- a/api/tests/test_containers_integration_tests/services/test_annotation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_annotation_service.py @@ -9,6 +9,7 @@ from models import Account from models.enums import ConversationFromSource, InvokeFrom from models.model import MessageAnnotation from services.annotation_service import AppAnnotationService +from services.app_ref_service import AnnotationRef from services.app_service import AppService, CreateAppParams from tests.test_containers_integration_tests.helpers import generate_valid_password @@ -119,6 +120,10 @@ class TestAnnotationService: tenant_id, ) + @staticmethod + def _annotation_ref(app, annotation_id: str) -> AnnotationRef: + return AnnotationRef(tenant_id=app.tenant_id, app_id=app.id, annotation_id=annotation_id) + def _create_test_conversation(self, db_session_with_containers: Session, app, account, fake): """ Helper method to create a test conversation with all required fields. @@ -282,7 +287,9 @@ class TestAnnotationService: "answer": fake.text(max_nb_chars=200), } updated_annotation = AppAnnotationService.update_app_annotation_directly( - updated_args, app.id, annotation.id, db_session_with_containers + updated_args, + self._annotation_ref(app, annotation.id), + db_session_with_containers, ) # Verify annotation was updated correctly @@ -570,7 +577,7 @@ class TestAnnotationService: annotation_id = annotation.id # Delete the annotation - AppAnnotationService.delete_app_annotation(app.id, annotation_id, db_session_with_containers) + AppAnnotationService.delete_app_annotation(self._annotation_ref(app, annotation_id), db_session_with_containers) # Verify annotation was deleted @@ -583,22 +590,27 @@ class TestAnnotationService: # Note: In this test, no annotation setting exists, so task should not be called mock_external_service_dependencies["delete_task"].delay.assert_not_called() - def test_delete_app_annotation_app_not_found( + def test_delete_app_annotation_annotation_not_found_for_wrong_app( self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test deletion of app annotation when app is not found. + Test deletion of app annotation when the annotation does not belong to the supplied app ref. """ fake = Faker() non_existent_app_id = fake.uuid4() annotation_id = fake.uuid4() + app_ref = AnnotationRef( + tenant_id=fake.uuid4(), + app_id=non_existent_app_id, + annotation_id=annotation_id, + ) # Mock random current user to avoid dependency issues self._mock_current_user(mock_external_service_dependencies, fake.uuid4(), fake.uuid4()) - # Try to delete annotation with non-existent app - with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.delete_app_annotation(non_existent_app_id, annotation_id, db_session_with_containers) + # Try to delete annotation with a ref that cannot match any annotation row + with pytest.raises(NotFound, match="Annotation not found"): + AppAnnotationService.delete_app_annotation(app_ref, db_session_with_containers) def test_delete_app_annotation_annotation_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -612,7 +624,10 @@ class TestAnnotationService: # Try to delete non-existent annotation with pytest.raises(NotFound, match="Annotation not found"): - AppAnnotationService.delete_app_annotation(app.id, non_existent_annotation_id, db_session_with_containers) + AppAnnotationService.delete_app_annotation( + self._annotation_ref(app, non_existent_annotation_id), + db_session_with_containers, + ) def test_enable_app_annotation_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -731,7 +746,9 @@ class TestAnnotationService: # Get hit histories hit_histories, total = AppAnnotationService.get_annotation_hit_histories( - app.id, annotation.id, page=1, limit=10 + self._annotation_ref(app, annotation.id), + page=1, + limit=10, ) # Verify results @@ -1229,7 +1246,9 @@ class TestAnnotationService: "answer": fake.text(max_nb_chars=200), } updated_annotation = AppAnnotationService.update_app_annotation_directly( - updated_args, app.id, annotation.id, db_session_with_containers + updated_args, + self._annotation_ref(app, annotation.id), + db_session_with_containers, ) # Verify annotation was updated correctly @@ -1300,7 +1319,7 @@ class TestAnnotationService: mock_external_service_dependencies["delete_task"].delay.reset_mock() # Delete the annotation - AppAnnotationService.delete_app_annotation(app.id, annotation_id, db_session_with_containers) + AppAnnotationService.delete_app_annotation(self._annotation_ref(app, annotation_id), db_session_with_containers) # Verify annotation was deleted deleted_annotation = ( diff --git a/api/tests/test_containers_integration_tests/services/test_audio_service_db.py b/api/tests/test_containers_integration_tests/services/test_audio_service_db.py index c9cf60bcfb1..d1a025c2067 100644 --- a/api/tests/test_containers_integration_tests/services/test_audio_service_db.py +++ b/api/tests/test_containers_integration_tests/services/test_audio_service_db.py @@ -24,6 +24,7 @@ from core.app.entities.app_invoke_entities import InvokeFrom from models.account import TenantAccountJoin from models.enums import ConversationFromSource, MessageStatus from models.model import App, AppMode, Conversation, Message +from services.app_ref_service import MessageRef from services.audio_service import AudioService from tests.test_containers_integration_tests.controllers.console.helpers import ( create_console_account_and_tenant, @@ -159,7 +160,12 @@ class TestAudioServiceTranscriptTTSMessageLookup: result = AudioService.transcript_tts( app_model=app, session=db_session_with_containers, - message_id=message.id, + message_ref=MessageRef( + tenant_id=app.tenant_id, + app_id=app.id, + message_id=message.id, + account_id=account_id, + ), voice="en-US-Neural", ) @@ -176,7 +182,7 @@ class TestAudioServiceTranscriptTTSMessageLookup: result = AudioService.transcript_tts( app_model=app, session=db_session_with_containers, - message_id="invalid-uuid", + message_ref=MessageRef(tenant_id=app.tenant_id, app_id=app.id, message_id="invalid-uuid"), ) assert result is None @@ -188,7 +194,7 @@ class TestAudioServiceTranscriptTTSMessageLookup: result = AudioService.transcript_tts( app_model=app, session=db_session_with_containers, - message_id=str(uuid4()), + message_ref=MessageRef(tenant_id=app.tenant_id, app_id=app.id, message_id=str(uuid4())), ) assert result is None @@ -209,7 +215,12 @@ class TestAudioServiceTranscriptTTSMessageLookup: result = AudioService.transcript_tts( app_model=app, session=db_session_with_containers, - message_id=message.id, + message_ref=MessageRef( + tenant_id=app.tenant_id, + app_id=app.id, + message_id=message.id, + account_id=account_id, + ), ) assert result is None diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py index ab5ba173792..7f4bd96b28f 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py @@ -15,6 +15,7 @@ from models import Account from models.dataset import Dataset, Document from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, IndexingStatus from models.model import UploadFile +from services.dataset_ref_service import DatasetRef from services.dataset_service import DocumentService from services.errors.account import NoPermissionError @@ -615,9 +616,10 @@ def test_delete_document_emits_signal_and_commits(db_session_with_containers: Se def test_delete_documents_ignores_empty_input(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) + dataset_ref = DatasetRef(tenant_id=dataset.tenant_id, dataset_id=dataset.id) with patch("services.dataset_service.batch_clean_document_task.delay") as delay: - DocumentService.delete_documents(dataset, [], session=db_session_with_containers) + DocumentService.delete_documents(dataset_ref, [], dataset.doc_form, session=db_session_with_containers) delay.assert_not_called() @@ -649,9 +651,15 @@ def test_delete_documents_deletes_rows_and_dispatches_cleanup_task(db_session_wi position=2, data_source_info={"upload_file_id": upload_file_b.id}, ) + dataset_ref = DatasetRef(tenant_id=dataset.tenant_id, dataset_id=dataset.id) with patch("services.dataset_service.batch_clean_document_task.delay") as delay: - DocumentService.delete_documents(dataset, [document_a.id, document_b.id], session=db_session_with_containers) + DocumentService.delete_documents( + dataset_ref, + [document_a.id, document_b.id], + dataset.doc_form, + session=db_session_with_containers, + ) assert db_session_with_containers.get(Document, document_a.id) is None assert db_session_with_containers.get(Document, document_b.id) is None diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_service.py index 9ba1fda08b8..349aac1be36 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_service.py @@ -15,6 +15,7 @@ from sqlalchemy.orm import Session from models import Account, AccountStatus, App, TenantStatus, Workflow from models.model import AppMode from models.workflow import WorkflowType +from services.workflow_ref_service import WorkflowRef from services.workflow_service import WorkflowService @@ -1319,10 +1320,9 @@ class TestWorkflowService: # Act result = workflow_service.update_workflow( session=db_session_with_containers, - workflow_id=workflow.id, - tenant_id=workflow.tenant_id, account_id=account.id, data=update_data, + workflow_ref=WorkflowRef(tenant_id=workflow.tenant_id, owner_id=app.id, workflow_id=workflow.id), ) # Assert @@ -1350,10 +1350,9 @@ class TestWorkflowService: # Act result = workflow_service.update_workflow( session=db_session_with_containers, - workflow_id=non_existent_workflow_id, - tenant_id=app.tenant_id, account_id=account.id, data=update_data, + workflow_ref=WorkflowRef(tenant_id=app.tenant_id, owner_id=app.id, workflow_id=non_existent_workflow_id), ) # Assert @@ -1385,10 +1384,9 @@ class TestWorkflowService: # Act result = workflow_service.update_workflow( session=db_session_with_containers, - workflow_id=workflow.id, - tenant_id=workflow.tenant_id, account_id=account.id, data=update_data, + workflow_ref=WorkflowRef(tenant_id=workflow.tenant_id, owner_id=app.id, workflow_id=workflow.id), ) # Assert @@ -1421,7 +1419,8 @@ class TestWorkflowService: # Act result = workflow_service.delete_workflow( - session=db_session_with_containers, workflow_id=workflow.id, tenant_id=workflow.tenant_id + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=workflow.tenant_id, owner_id=app.id, workflow_id=workflow.id), ) # Assert @@ -1456,7 +1455,8 @@ class TestWorkflowService: with pytest.raises(DraftWorkflowDeletionError, match="Cannot delete draft workflow versions"): workflow_service.delete_workflow( - session=db_session_with_containers, workflow_id=workflow.id, tenant_id=workflow.tenant_id + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=workflow.tenant_id, owner_id=app.id, workflow_id=workflow.id), ) def test_delete_workflow_in_use_error(self, db_session_with_containers: Session): @@ -1487,7 +1487,8 @@ class TestWorkflowService: with pytest.raises(WorkflowInUseError, match="Cannot delete workflow that is currently in use by app"): workflow_service.delete_workflow( - session=db_session_with_containers, workflow_id=workflow.id, tenant_id=workflow.tenant_id + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=workflow.tenant_id, owner_id=app.id, workflow_id=workflow.id), ) def test_delete_workflow_not_found_error(self, db_session_with_containers: Session): @@ -1507,7 +1508,12 @@ class TestWorkflowService: # Act & Assert with pytest.raises(ValueError, match=f"Workflow with ID {non_existent_workflow_id} not found"): workflow_service.delete_workflow( - session=db_session_with_containers, workflow_id=non_existent_workflow_id, tenant_id=app.tenant_id + session=db_session_with_containers, + workflow_ref=WorkflowRef( + tenant_id=app.tenant_id, + owner_id=app.id, + workflow_id=non_existent_workflow_id, + ), ) def test_run_free_workflow_node_success(self, db_session_with_containers: Session): diff --git a/api/tests/test_containers_integration_tests/services/workflow/test_workflow_deletion.py b/api/tests/test_containers_integration_tests/services/workflow/test_workflow_deletion.py index afc4908c159..d2fc873cf51 100644 --- a/api/tests/test_containers_integration_tests/services/workflow/test_workflow_deletion.py +++ b/api/tests/test_containers_integration_tests/services/workflow/test_workflow_deletion.py @@ -11,6 +11,7 @@ from models.account import Account, Tenant, TenantAccountJoin from models.model import App from models.tools import WorkflowToolProvider from models.workflow import Workflow +from services.workflow_ref_service import WorkflowRef from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService @@ -111,7 +112,8 @@ class TestWorkflowDeletion: service = WorkflowService(sessionmaker(bind=db.engine)) result = service.delete_workflow( - session=db_session_with_containers, workflow_id=workflow_id, tenant_id=tenant.id + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=tenant.id, owner_id=app.id, workflow_id=workflow_id), ) assert result is True @@ -128,7 +130,10 @@ class TestWorkflowDeletion: service = WorkflowService(sessionmaker(bind=db.engine)) with pytest.raises(DraftWorkflowDeletionError): - service.delete_workflow(session=db_session_with_containers, workflow_id=workflow.id, tenant_id=tenant.id) + service.delete_workflow( + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=tenant.id, owner_id=app.id, workflow_id=workflow.id), + ) def test_delete_workflow_in_use_by_app_raises_error(self, db_session_with_containers: Session): tenant, account = self._create_tenant_and_account(db_session_with_containers) @@ -142,7 +147,10 @@ class TestWorkflowDeletion: service = WorkflowService(sessionmaker(bind=db.engine)) with pytest.raises(WorkflowInUseError, match="currently in use by app"): - service.delete_workflow(session=db_session_with_containers, workflow_id=workflow.id, tenant_id=tenant.id) + service.delete_workflow( + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=tenant.id, owner_id=app.id, workflow_id=workflow.id), + ) def test_delete_workflow_published_as_tool_raises_error(self, db_session_with_containers: Session): tenant, account = self._create_tenant_and_account(db_session_with_containers) @@ -155,4 +163,7 @@ class TestWorkflowDeletion: service = WorkflowService(sessionmaker(bind=db.engine)) with pytest.raises(WorkflowInUseError, match="published as a tool"): - service.delete_workflow(session=db_session_with_containers, workflow_id=workflow.id, tenant_id=tenant.id) + service.delete_workflow( + session=db_session_with_containers, + workflow_ref=WorkflowRef(tenant_id=tenant.id, owner_id=app.id, workflow_id=workflow.id), + ) diff --git a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py index fecbd7f7b06..8a6094b94b8 100644 --- a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py @@ -1,6 +1,23 @@ from __future__ import annotations +from inspect import unwrap +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import pytest +from flask import Flask +from werkzeug.exceptions import NotFound + from controllers.console.app import annotation as annotation_module +from services.app_ref_service import AnnotationRef, AppRef + + +def _app_model() -> SimpleNamespace: + return SimpleNamespace(id="app-1", tenant_id="tenant-1", status="normal") + + +def _annotation_model(annotation_id: str = "ann-1") -> SimpleNamespace: + return SimpleNamespace(id=annotation_id, question="q", content="a", hit_count=0, created_at=None) def test_annotation_reply_payload_valid(): @@ -90,3 +107,113 @@ def test_annotation_file_payload_valid(): """Test AnnotationFilePayload with valid message ID.""" payload = annotation_module.AnnotationFilePayload(message_id="550e8400-e29b-41d4-a716-446655440000") assert payload.message_id == "550e8400-e29b-41d4-a716-446655440000" + + +def test_get_app_ref_raises_not_found_when_app_is_not_in_current_tenant(): + with ( + patch.object( + annotation_module, + "current_account_with_tenant", + return_value=(SimpleNamespace(id="account-1"), "tenant-1"), + ), + patch.object(annotation_module.db.session, "scalar", return_value=None), + ): + with pytest.raises(NotFound): + annotation_module._get_app_ref("app-1") + + +class TestConsoleAnnotationRefBoundaries: + def test_batch_delete_uses_app_ref(self, app: Flask): + api = annotation_module.AnnotationApi() + handler = unwrap(api.delete) + delete_mock = Mock() + + with ( + app.test_request_context("/?annotation_id=ann-1&annotation_id=ann-2", method="DELETE"), + patch.object( + annotation_module, + "current_account_with_tenant", + return_value=(SimpleNamespace(id="account-1"), "tenant-1"), + ), + patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), + patch.object(annotation_module.AppAnnotationService, "delete_app_annotations_in_batch", delete_mock), + ): + response, status = handler(api, "app-1") + + assert response == "" + assert status == 204 + delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"]) + + def test_update_uses_annotation_ref(self, app: Flask): + api = annotation_module.AnnotationUpdateDeleteApi() + handler = unwrap(api.post) + update_mock = Mock(return_value=_annotation_model()) + payload = {"question": "updated"} + + with ( + app.test_request_context("/annotations/ann-1", method="POST", json=payload), + patch.object(type(annotation_module.console_ns), "payload", payload), + patch.object( + annotation_module, + "current_account_with_tenant", + return_value=(SimpleNamespace(id="account-1"), "tenant-1"), + ), + patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), + patch.object(annotation_module.AppAnnotationService, "update_app_annotation_directly", update_mock), + ): + response = handler(api, "app-1", "ann-1") + + assert response["question"] == "q" + update_mock.assert_called_once() + assert update_mock.call_args.args[1] == AnnotationRef("tenant-1", "app-1", "ann-1") + + def test_delete_uses_annotation_ref(self, app: Flask): + api = annotation_module.AnnotationUpdateDeleteApi() + handler = unwrap(api.delete) + delete_mock = Mock() + + with ( + app.test_request_context("/annotations/ann-1", method="DELETE"), + patch.object( + annotation_module, + "current_account_with_tenant", + return_value=(SimpleNamespace(id="account-1"), "tenant-1"), + ), + patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), + patch.object(annotation_module.AppAnnotationService, "delete_app_annotation", delete_mock), + ): + response, status = handler(api, "app-1", "ann-1") + + assert response == "" + assert status == 204 + delete_mock.assert_called_once() + assert delete_mock.call_args.args[0] == AnnotationRef("tenant-1", "app-1", "ann-1") + + def test_hit_history_uses_annotation_ref(self, app: Flask): + api = annotation_module.AnnotationHitHistoryListApi() + handler = unwrap(api.get) + history = SimpleNamespace( + id="history-1", + source="hit-testing", + score=0.9, + question="q", + annotation_question="q", + annotation_content="a", + created_at=None, + ) + hit_history_mock = Mock(return_value=([history], 1)) + + with ( + app.test_request_context("/hit-histories?page=2&limit=5", method="GET"), + patch.object( + annotation_module, + "current_account_with_tenant", + return_value=(SimpleNamespace(id="account-1"), "tenant-1"), + ), + patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), + patch.object(annotation_module.AppAnnotationService, "get_annotation_hit_histories", hit_history_mock), + ): + response = handler(api, "app-1", "ann-1") + + assert response["total"] == 1 + hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5) diff --git a/api/tests/unit_tests/controllers/console/app/test_audio.py b/api/tests/unit_tests/controllers/console/app/test_audio.py index 82b9b68247f..5dfc5b271eb 100644 --- a/api/tests/unit_tests/controllers/console/app/test_audio.py +++ b/api/tests/unit_tests/controllers/console/app/test_audio.py @@ -3,6 +3,7 @@ from __future__ import annotations import io from inspect import unwrap from types import SimpleNamespace +from unittest.mock import patch import pytest from flask import Flask @@ -23,6 +24,7 @@ from controllers.console.app.error import ( ) from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from graphon.model_runtime.errors.invoke import InvokeError +from services.app_ref_service import MessageRef from services.audio_service import AudioService from services.errors.app_model_config import AppModelConfigBrokenError from services.errors.audio import ( @@ -103,6 +105,32 @@ def test_console_text_api_success(app: Flask, monkeypatch: pytest.MonkeyPatch) - assert response == {"audio": "ok"} +def test_console_text_api_builds_message_ref(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + api = ChatMessageTextApi() + handler = unwrap(api.post) + app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") + calls = {} + + def fake_transcript_tts(**kwargs): + calls.update(kwargs) + return {"audio": "ok"} + + monkeypatch.setattr(AudioService, "transcript_tts", fake_transcript_tts) + + with ( + app.test_request_context( + "/console/api/apps/app-1/text-to-audio", + method="POST", + json={"text": "hello", "message_id": "message-1"}, + ), + patch("controllers.console.app.audio.current_user", SimpleNamespace(id="account-1")), + ): + response = handler(api, app_model=app_model) + + assert response == {"audio": "ok"} + assert calls["message_ref"] == MessageRef("tenant-1", "app-1", "message-1", account_id="account-1") + + def test_console_text_api_error_mapping(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(AudioService, "transcript_tts", lambda **_kwargs: (_ for _ in ()).throw(QuotaExceededError())) diff --git a/api/tests/unit_tests/controllers/console/app/test_generator_api.py b/api/tests/unit_tests/controllers/console/app/test_generator_api.py index 308089b8489..6b06ea99d5b 100644 --- a/api/tests/unit_tests/controllers/console/app/test_generator_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_generator_api.py @@ -70,7 +70,7 @@ def test_instruction_generate_app_not_found(app: Flask, monkeypatch: pytest.Monk method = unwrap(api.post) session = MagicMock() - session.get.return_value = None + session.scalar.return_value = None with app.test_request_context( "/console/api/instruction-generate", @@ -86,7 +86,13 @@ def test_instruction_generate_app_not_found(app: Flask, monkeypatch: pytest.Monk assert status == 400 assert response["error"] == "app app-1 not found" - session.get.assert_called_once_with(generator_module.App, "app-1") + stmt = session.scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "apps.id" in statement + assert "apps.tenant_id" in statement + assert "app-1" in compiled.params.values() + assert "t1" in compiled.params.values() def test_instruction_generate_workflow_not_found(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -94,7 +100,7 @@ def test_instruction_generate_workflow_not_found(app: Flask, monkeypatch: pytest method = unwrap(api.post) app_model = SimpleNamespace(id="app-1") - session = SimpleNamespace(get=lambda *_args, **_kwargs: app_model) + session = SimpleNamespace(scalar=lambda *_args, **_kwargs: app_model) _install_workflow_service(monkeypatch, workflow=None) with app.test_request_context( @@ -118,7 +124,7 @@ def test_instruction_generate_node_missing(app: Flask, monkeypatch: pytest.Monke method = unwrap(api.post) app_model = SimpleNamespace(id="app-1") - session = SimpleNamespace(get=lambda *_args, **_kwargs: app_model) + session = SimpleNamespace(scalar=lambda *_args, **_kwargs: app_model) workflow = SimpleNamespace(graph_dict={"nodes": []}) _install_workflow_service(monkeypatch, workflow=workflow) @@ -144,7 +150,7 @@ def test_instruction_generate_code_node(app: Flask, monkeypatch: pytest.MonkeyPa method = unwrap(api.post) app_model = SimpleNamespace(id="app-1") - session = SimpleNamespace(get=lambda *_args, **_kwargs: app_model) + session = SimpleNamespace(scalar=lambda *_args, **_kwargs: app_model) workflow = SimpleNamespace( graph_dict={ diff --git a/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py b/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py index aa248180bca..1b392d5185e 100644 --- a/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py +++ b/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py @@ -130,3 +130,50 @@ class TestAppMCPServerController: assert response == {"id": "server-1"} assert status_code == 201 + + def test_put_binds_server_lookup_to_app_ref(self): + api = AppMCPServerController() + method = unwrap(api.put) + payload = {"id": "server-1", "description": "Updated", "parameters": {"timeout": 30}, "status": "active"} + app = Flask(__name__) + app.config["TESTING"] = True + server = SimpleNamespace( + id="server-1", + tenant_id="tenant-1", + app_id="app-1", + name="Old", + description="Old", + parameters="{}", + status="active", + ) + + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch("controllers.console.app.mcp_server.db.session.scalar", return_value=server) as scalar, + patch("controllers.console.app.mcp_server.db.session.get") as get_mock, + patch("controllers.console.app.mcp_server.db.session.commit") as commit, + patch( + "controllers.console.app.mcp_server.AppMCPServerResponse.model_validate", + return_value=_ValidatedResponse({"id": "server-1"}), + ), + ): + response = method( + api, + app_model=SimpleNamespace( + id="app-1", tenant_id="tenant-1", name="Demo App", description="App description" + ), + ) + + stmt = scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "app_mcp_servers.id" in statement + assert "app_mcp_servers.tenant_id" in statement + assert "app_mcp_servers.app_id" in statement + assert payload["id"] in compiled.params.values() + assert "tenant-1" in compiled.params.values() + assert "app-1" in compiled.params.values() + get_mock.assert_not_called() + commit.assert_called_once() + assert response == {"id": "server-1"} diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py index e8bb6f88674..99d7cf626f0 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py @@ -83,6 +83,7 @@ def make_dataset(**overrides): "tenant_id": "tenant-1", "name": "Dataset", "indexing_technique": "economy", + "chunk_structure": IndexStructureType.PARAGRAPH_INDEX, "created_by": "u1", "summary_index_setting": {"enable": True}, } diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py index 09d4da94748..2e54c977260 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py @@ -108,6 +108,15 @@ def _segment_response_dict(): } +def _bind_dataset_document(dataset, document, dataset_id: str = "ds-1", document_id: str = "doc-1"): + dataset.id = dataset_id + dataset.tenant_id = "tenant-1" + document.id = document_id + document.dataset_id = dataset_id + document.tenant_id = "tenant-1" + return document + + def test_segment_response_with_summary(): segment = _segment() @@ -380,6 +389,7 @@ class TestDatasetDocumentSegmentAddApi: document = MagicMock() document.doc_form = IndexStructureType.PARAGRAPH_INDEX + _bind_dataset_document(dataset, document) segment = _segment() @@ -504,6 +514,7 @@ class TestDatasetDocumentSegmentUpdateApi: document = MagicMock() document.doc_form = IndexStructureType.PARAGRAPH_INDEX + _bind_dataset_document(dataset, document) segment = _segment() @@ -519,8 +530,8 @@ class TestDatasetDocumentSegmentUpdateApi: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - side_effect=[segment, None], + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", + return_value=segment, ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -538,6 +549,7 @@ class TestDatasetDocumentSegmentUpdateApi: "controllers.console.datasets.datasets_segments.SummaryIndexService.get_segment_summary", return_value=None, ), + patch("models.dataset.db.session.scalar", return_value=None), patch("models.dataset.db.session.execute", return_value=MagicMock(all=MagicMock(return_value=[]))), ): response, status = method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") @@ -545,6 +557,75 @@ class TestDatasetDocumentSegmentUpdateApi: assert status == 200 assert "data" in response + def test_patch_document_outside_dataset_is_not_found(self, app: Flask): + api = DatasetDocumentSegmentUpdateApi() + method = inspect.unwrap(api.patch) + + payload = {"content": "updated"} + user = MagicMock(is_dataset_editor=True) + dataset = MagicMock(id="ds-1", tenant_id="tenant-1", indexing_technique="economy") + document = MagicMock(id="doc-1", dataset_id="other-dataset", tenant_id="tenant-1") + + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", + return_value=None, + ), + patch( + "controllers.console.datasets.datasets_segments.DocumentService.get_document", + return_value=document, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", + return_value=None, + ), + ): + with pytest.raises(NotFound): + method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + + def test_patch_segment_not_found(self, app: Flask): + api = DatasetDocumentSegmentUpdateApi() + method = inspect.unwrap(api.patch) + + payload = {"content": "updated"} + user = MagicMock(is_dataset_editor=True) + dataset = MagicMock(indexing_technique="economy") + document = MagicMock() + _bind_dataset_document(dataset, document) + + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", + return_value=None, + ), + patch( + "controllers.console.datasets.datasets_segments.DocumentService.get_document", + return_value=document, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", + return_value=None, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", + return_value=None, + ), + ): + with pytest.raises(NotFound): + method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + def test_patch_llm_bad_request(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() method = inspect.unwrap(api.patch) @@ -576,6 +657,10 @@ class TestDatasetDocumentSegmentUpdateApi: "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", + return_value=None, + ), patch( "controllers.console.datasets.datasets_segments.ModelManager.get_model_instance", side_effect=LLMBadRequestError(), @@ -781,13 +866,15 @@ class TestChildChunkAddApi: api = ChildChunkAddApi() method = inspect.unwrap(api.get) + dataset = MagicMock() + document = _bind_dataset_document(dataset, MagicMock()) pagination = MagicMock(items=[], total=0, pages=0) with ( app.test_request_context("/?page=bad&limit="), patch( "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), + return_value=dataset, ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", @@ -795,10 +882,10 @@ class TestChildChunkAddApi: ), patch( "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=MagicMock(), + return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=MagicMock(), ), patch( @@ -826,6 +913,7 @@ class TestChildChunkAddApi: dataset.indexing_technique = "economy" document = MagicMock() + _bind_dataset_document(dataset, document) segment = MagicMock() child_chunk = _child_chunk() @@ -841,7 +929,7 @@ class TestChildChunkAddApi: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment, ), patch( @@ -868,6 +956,7 @@ class TestChildChunkAddApi: dataset = MagicMock(indexing_technique="economy") document = MagicMock() + _bind_dataset_document(dataset, document) segment = MagicMock() with ( @@ -882,7 +971,7 @@ class TestChildChunkAddApi: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment, ), patch( @@ -897,6 +986,35 @@ class TestChildChunkAddApi: with pytest.raises(ChildChunkIndexingError): method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + def test_post_permission_denied(self, app: Flask): + api = ChildChunkAddApi() + method = inspect.unwrap(api.post) + + payload = {"content": "child"} + user = MagicMock(is_dataset_editor=True) + dataset = MagicMock(indexing_technique="economy") + document = MagicMock() + _bind_dataset_document(dataset, document) + + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_segments.DocumentService.get_document", + return_value=document, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", + side_effect=services.errors.account.NoPermissionError("no access"), + ), + ): + with pytest.raises(Forbidden): + method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + class TestChildChunkUpdateApi: def test_delete_success(self, app: Flask): @@ -908,6 +1026,7 @@ class TestChildChunkUpdateApi: dataset = MagicMock() document = MagicMock() + _bind_dataset_document(dataset, document) segment = MagicMock() child_chunk = MagicMock() @@ -922,8 +1041,12 @@ class TestChildChunkUpdateApi: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - side_effect=[segment, child_chunk], + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", + return_value=segment, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", + return_value=child_chunk, ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -947,6 +1070,7 @@ class TestChildChunkUpdateApi: dataset = MagicMock() document = MagicMock() + _bind_dataset_document(dataset, document) segment = MagicMock() child_chunk = MagicMock() @@ -961,8 +1085,12 @@ class TestChildChunkUpdateApi: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - side_effect=[segment, child_chunk], + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", + return_value=segment, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", + return_value=child_chunk, ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -976,6 +1104,80 @@ class TestChildChunkUpdateApi: with pytest.raises(ChildChunkDeleteIndexError): method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + def test_delete_child_chunk_not_found(self, app: Flask): + api = ChildChunkUpdateApi() + method = inspect.unwrap(api.delete) + + user = MagicMock(is_dataset_editor=True) + dataset = MagicMock() + document = MagicMock() + _bind_dataset_document(dataset, document) + segment = MagicMock() + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_segments.DocumentService.get_document", + return_value=document, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", + return_value=segment, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", + return_value=None, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", + return_value=None, + ), + ): + with pytest.raises(NotFound): + method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + + def test_patch_child_chunk_not_found(self, app: Flask): + api = ChildChunkUpdateApi() + method = inspect.unwrap(api.patch) + + payload = {"content": "updated child"} + user = MagicMock(is_dataset_editor=True) + dataset = MagicMock() + document = MagicMock() + _bind_dataset_document(dataset, document) + segment = MagicMock() + + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_segments.DocumentService.get_document", + return_value=document, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", + return_value=segment, + ), + patch( + "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", + return_value=None, + ), + patch( + "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", + return_value=None, + ), + ): + with pytest.raises(NotFound): + method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + class TestSegmentListAdvancedCases: def test_segment_list_with_keyword_filter(self, app: Flask): diff --git a/api/tests/unit_tests/controllers/console/explore/test_audio.py b/api/tests/unit_tests/controllers/console/explore/test_audio.py index a6642f85825..c1930cfbcf5 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_audio.py +++ b/api/tests/unit_tests/controllers/console/explore/test_audio.py @@ -21,6 +21,7 @@ from core.errors.error import ( QuotaExceededError, ) from graphon.model_runtime.errors.invoke import InvokeError +from services.app_ref_service import MessageRef from services.errors.audio import ( AudioTooLargeServiceError, NoAudioUploadedServiceError, @@ -40,6 +41,8 @@ def unwrap(func): def installed_app(): app = MagicMock() app.app = MagicMock() + app.app.id = "app-1" + app.app.tenant_id = "tenant-1" return app @@ -237,20 +240,29 @@ class TestChatTextApi: self.method = unwrap(self.api.post) def test_post_success(self, app: Flask, installed_app): + transcript_tts = MagicMock(return_value={"audio": "ok"}) + with ( app.test_request_context( "/", json={"message_id": "m1", "text": "hello", "voice": "v1"}, ), patch.object( - audio_module.AudioService, - "transcript_tts", - return_value={"audio": "ok"}, + audio_module, + "current_account_with_tenant", + return_value=(MagicMock(id="account-1"), "tenant-1"), ), + patch.object(audio_module.AudioService, "transcript_tts", transcript_tts), ): resp = self.method(installed_app) assert resp == {"audio": "ok"} + assert transcript_tts.call_args.kwargs["message_ref"] == MessageRef( + tenant_id="tenant-1", + app_id="app-1", + message_id="m1", + account_id="account-1", + ) def test_provider_not_initialized(self, app: Flask, installed_app): with ( diff --git a/api/tests/unit_tests/controllers/console/explore/test_trial.py b/api/tests/unit_tests/controllers/console/explore/test_trial.py index be68a3beed6..db800a23b84 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -32,6 +32,7 @@ from graphon.model_runtime.errors.invoke import InvokeError from models import Account from models.account import TenantStatus from models.model import AppMode +from services.app_ref_service import MessageRef from services.errors.conversation import ConversationNotExistsError from services.errors.llm import InvokeRateLimitError @@ -774,6 +775,27 @@ class TestTrialChatTextApi: assert result == {"audio": "base64_data"} + def test_success_with_message_ref(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: + api = module.TrialChatTextApi() + method = unwrap(api.post) + transcript_tts = MagicMock(return_value={"audio": "base64_data"}) + trial_app_chat.tenant_id = "tenant-1" + + with ( + app.test_request_context("/", json={"text": "hello", "message_id": "message-1"}), + patch.object(module.AudioService, "transcript_tts", transcript_tts), + patch.object(module.RecommendedAppService, "add_trial_app_record"), + ): + result = method(api, account, trial_app_chat) + + assert result == {"audio": "base64_data"} + assert transcript_tts.call_args.kwargs["message_ref"] == MessageRef( + "tenant-1", + "a-chat", + "message-1", + account_id="u1", + ) + def test_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() method = unwrap(api.post) diff --git a/api/tests/unit_tests/controllers/console/tag/test_tags.py b/api/tests/unit_tests/controllers/console/tag/test_tags.py index 8f5bb176b8e..2da11afa1f7 100644 --- a/api/tests/unit_tests/controllers/console/tag/test_tags.py +++ b/api/tests/unit_tests/controllers/console/tag/test_tags.py @@ -253,6 +253,35 @@ class TestTagUpdateDeleteApi: delete_mock.assert_called_once_with("tag-1", module.db.session) assert status == 204 + def test_delete_snippet_tag_checks_type_in_current_tenant(self, app: Flask, admin_user): + api = TagUpdateDeleteApi() + method = unwrap(api.delete) + + with ( + app.test_request_context("/"), + patch("controllers.console.tag.tags.dify_config.RBAC_ENABLED", True), + patch( + "controllers.console.tag.tags.current_account_with_tenant", + return_value=(SimpleNamespace(id="user-1"), "tenant-1"), + ), + patch.object(module.db.session, "scalar", return_value=TagType.SNIPPET) as scalar_mock, + patch("controllers.console.tag.tags.enforce_rbac_access") as enforce_mock, + patch("controllers.console.tag.tags.TagService.delete_tag") as delete_mock, + ): + result, status = method(api, "tag-1") + + scalar_mock.assert_called_once() + enforce_mock.assert_called_once_with( + tenant_id="tenant-1", + account_id="user-1", + resource_type=module.RBACResourceScope.WORKSPACE, + scene=module.RBACPermission.SNIPPETS_CREATE_AND_MODIFY, + resource_required=False, + ) + delete_mock.assert_called_once_with("tag-1", module.db.session) + assert result == "" + assert status == 204 + class TestTagBindingCollectionApi: def test_create_success(self, app: Flask, admin_user, payload_patch): diff --git a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py index b4dd5e957c1..810101fb0a5 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py @@ -188,7 +188,7 @@ class TestAnnotationReplyActionApi: api = AnnotationReplyActionApi() handler = unwrap(api.post) - app_model = SimpleNamespace(id="app") + app_model = SimpleNamespace(id="app", tenant_id="tenant") with app.test_request_context( "/apps/annotation-reply/enable", @@ -206,7 +206,7 @@ class TestAnnotationReplyActionApi: api = AnnotationReplyActionApi() handler = unwrap(api.post) - app_model = SimpleNamespace(id="app") + app_model = SimpleNamespace(id="app", tenant_id="tenant") with app.test_request_context( "/apps/annotation-reply/disable", @@ -333,7 +333,7 @@ class TestAnnotationUpdateDeleteApi: api = AnnotationUpdateDeleteApi() put_handler = unwrap(api.put) delete_handler = unwrap(api.delete) - app_model = SimpleNamespace(id="app") + app_model = SimpleNamespace(id="app", tenant_id="tenant") with app.test_request_context("/apps/annotations/1", method="PUT", json={"question": "q", "answer": "a"}): response = put_handler(api, app_model=app_model, annotation_id="1") diff --git a/api/tests/unit_tests/controllers/service_api/app/test_audio.py b/api/tests/unit_tests/controllers/service_api/app/test_audio.py index 52d050ff55a..3615eb6d209 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_audio.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_audio.py @@ -32,6 +32,7 @@ from controllers.service_api.app.error import ( ) from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from graphon.model_runtime.errors.invoke import InvokeError +from services.app_ref_service import MessageRef from services.audio_service import AudioService from services.errors.app_model_config import AppModelConfigBrokenError from services.errors.audio import ( @@ -180,7 +181,6 @@ class TestAudioServiceMockedBehavior: text="Hello world", voice="nova", end_user="user_123", - message_id="msg_123", ) assert result["audio"] == "base64_audio_data" @@ -245,7 +245,7 @@ class TestTextApi: api = TextApi() handler = unwrap(api.post) app_model = SimpleNamespace(id="a1") - end_user = SimpleNamespace(external_user_id="ext") + end_user = SimpleNamespace(id="end-user-1", external_user_id="ext") with app.test_request_context( "/text-to-audio", @@ -256,6 +256,30 @@ class TestTextApi: assert response == {"audio": "ok"} + def test_success_with_message_ref(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + calls = {} + + def fake_transcript_tts(**kwargs): + calls.update(kwargs) + return {"audio": "ok"} + + monkeypatch.setattr(AudioService, "transcript_tts", fake_transcript_tts) + + api = TextApi() + handler = unwrap(api.post) + app_model = SimpleNamespace(id="a1", tenant_id="tenant-1") + end_user = SimpleNamespace(id="end-user-1", external_user_id="ext") + + with app.test_request_context( + "/text-to-audio", + method="POST", + json={"text": "hello", "message_id": "message-1"}, + ): + response = handler(api, app_model=app_model, end_user=end_user) + + assert response == {"audio": "ok"} + assert calls["message_ref"] == MessageRef("tenant-1", "a1", "message-1", end_user_id="end-user-1") + def test_error_mapping(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( AudioService, "transcript_tts", lambda **_kwargs: (_ for _ in ()).throw(QuotaExceededError()) @@ -264,7 +288,7 @@ class TestTextApi: api = TextApi() handler = unwrap(api.post) app_model = SimpleNamespace(id="a1") - end_user = SimpleNamespace(external_user_id="ext") + end_user = SimpleNamespace(id="end-user-1", external_user_id="ext") with app.test_request_context("/text-to-audio", method="POST", json={"text": "hello"}): with pytest.raises(ProviderQuotaExceededError): diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py index 9170f38df2a..a95baf1b482 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py @@ -90,6 +90,19 @@ def _child_chunk() -> ChildChunk: return child_chunk +def _document_for_dataset( + dataset: Dataset, document_id: str = "doc-id", doc_form: str = IndexStructureType.PARAGRAPH_INDEX +): + document = Mock() + document.id = document_id + document.dataset_id = dataset.id + document.tenant_id = dataset.tenant_id + document.indexing_status = "completed" + document.enabled = True + document.doc_form = doc_form + return document + + class TestSegmentCreatePayload: """Test suite for SegmentCreatePayload Pydantic model.""" @@ -890,7 +903,9 @@ class TestSegmentApiGet: # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) + mock_doc_svc.get_document.return_value = _document_for_dataset( + mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX + ) mock_seg_svc.get_segments.return_value = ([mock_segment], 1) mock_get_summaries.return_value = {} mock_dump_segments.return_value = [_segment_response_dict()] @@ -1010,7 +1025,7 @@ class TestSegmentApiPost: mock_dataset.indexing_technique = "economy" mock_db.session.scalar.return_value = mock_dataset - mock_doc = Mock() + mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "completed" mock_doc.enabled = True mock_doc.doc_form = IndexStructureType.PARAGRAPH_INDEX @@ -1062,7 +1077,7 @@ class TestSegmentApiPost: mock_dataset.indexing_technique = "economy" mock_db.session.scalar.return_value = mock_dataset - mock_doc = Mock() + mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "completed" mock_doc.enabled = True mock_doc_svc.get_document.return_value = mock_doc @@ -1104,7 +1119,7 @@ class TestSegmentApiPost: mock_db.session.scalar.return_value = mock_dataset - mock_doc = Mock() + mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "indexing" # Not completed mock_doc_svc.get_document.return_value = mock_doc @@ -1156,10 +1171,10 @@ class TestDatasetSegmentApiDelete: mock_db.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc = Mock() + mock_doc = _document_for_dataset(mock_dataset) mock_doc_svc.get_document.return_value = mock_doc - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_seg_svc.get_segment_by_ref.return_value = mock_segment mock_seg_svc.delete_segment.return_value = None # Act @@ -1199,13 +1214,13 @@ class TestDatasetSegmentApiDelete: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc = Mock() + mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "completed" mock_doc.enabled = True mock_doc.doc_form = IndexStructureType.PARAGRAPH_INDEX mock_doc_svc.get_document.return_value = mock_doc - mock_seg_svc.get_segment_by_id.return_value = None # Segment not found + mock_seg_svc.get_segment_by_ref.return_value = None # Segment not found # Act & Assert with app.test_request_context( @@ -1351,8 +1366,10 @@ class TestDatasetSegmentApiUpdate: mock_dataset.indexing_technique = "economy" mock_db.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc_svc.get_document.return_value = Mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_doc_svc.get_document.return_value = _document_for_dataset( + mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX + ) + mock_seg_svc.get_segment_by_ref.return_value = mock_segment updated = Mock() updated.id = "updated-seg" mock_seg_svc.update_segment.return_value = updated @@ -1441,8 +1458,8 @@ class TestDatasetSegmentApiUpdate: mock_dataset.indexing_technique = "economy" mock_db.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc_svc.get_document.return_value = Mock() - mock_seg_svc.get_segment_by_id.return_value = None + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) + mock_seg_svc.get_segment_by_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id", @@ -1492,9 +1509,9 @@ class TestDatasetSegmentApiGetSingle: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc = Mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) + mock_doc = _document_for_dataset(mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX) mock_doc_svc.get_document.return_value = mock_doc - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_seg_svc.get_segment_by_ref.return_value = mock_segment mock_get_summary.return_value = None mock_dump_segment.return_value = _segment_response_dict() @@ -1539,9 +1556,9 @@ class TestDatasetSegmentApiGetSingle: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc = Mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) + mock_doc = _document_for_dataset(mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX) mock_doc_svc.get_document.return_value = mock_doc - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_seg_svc.get_segment_by_ref.return_value = mock_segment mock_summary_record = Mock(summary_content="This is the segment summary") mock_get_summary.return_value = mock_summary_record mock_dump_segment.return_value = _segment_response_dict("This is the segment summary") @@ -1641,8 +1658,8 @@ class TestDatasetSegmentApiGetSingle: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc_svc.get_document.return_value = Mock() - mock_seg_svc.get_segment_by_id.return_value = None + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) + mock_seg_svc.get_segment_by_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id", @@ -1682,8 +1699,8 @@ class TestChildChunkApiGet: """Test successful child chunk list retrieval.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() - mock_seg_svc.get_segment_by_id.return_value = Mock() + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) + mock_seg_svc.get_segment_by_ref.return_value = Mock() mock_pagination = Mock() mock_pagination.items = [_child_chunk(), _child_chunk()] @@ -1781,8 +1798,8 @@ class TestChildChunkApiGet: """Test 404 when segment not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() - mock_seg_svc.get_segment_by_id.return_value = None + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) + mock_seg_svc.get_segment_by_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id/child_chunks", @@ -1844,8 +1861,8 @@ class TestChildChunkApiPost: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_dataset.indexing_technique = "economy" mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() - mock_seg_svc.get_segment_by_id.return_value = Mock() + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) + mock_seg_svc.get_segment_by_ref.return_value = Mock() mock_child = _child_chunk() mock_seg_svc.create_child_chunk.return_value = mock_child @@ -1922,8 +1939,8 @@ class TestChildChunkApiPost: self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() - mock_seg_svc.get_segment_by_id.return_value = None + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) + mock_seg_svc.get_segment_by_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id/child_chunks", @@ -1976,19 +1993,19 @@ class TestDatasetChildChunkApiDelete: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc = Mock() + mock_doc = _document_for_dataset(mock_dataset) mock_doc_svc.get_document.return_value = mock_doc segment_id = str(uuid.uuid4()) mock_segment = Mock() mock_segment.id = segment_id mock_segment.document_id = "doc-id" - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_seg_svc.get_segment_by_ref.return_value = mock_segment child_chunk_id = str(uuid.uuid4()) mock_child = Mock() mock_child.segment_id = segment_id - mock_seg_svc.get_child_chunk_by_id.return_value = mock_child + mock_seg_svc.get_child_chunk_by_segment_ref.return_value = mock_child mock_seg_svc.delete_child_chunk.return_value = None with app.test_request_context( @@ -2025,14 +2042,14 @@ class TestDatasetChildChunkApiDelete: """Test 404 when child chunk not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) segment_id = str(uuid.uuid4()) mock_segment = Mock() mock_segment.id = segment_id mock_segment.document_id = "doc-id" - mock_seg_svc.get_segment_by_id.return_value = mock_segment - mock_seg_svc.get_child_chunk_by_id.return_value = None + mock_seg_svc.get_segment_by_ref.return_value = mock_segment + mock_seg_svc.get_child_chunk_by_segment_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/{segment_id}/child_chunks/cc-id", @@ -2066,13 +2083,10 @@ class TestDatasetChildChunkApiDelete: """Test 404 when segment does not belong to the document.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) segment_id = str(uuid.uuid4()) - mock_segment = Mock() - mock_segment.id = segment_id - mock_segment.document_id = "different-doc-id" - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_seg_svc.get_segment_by_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/{segment_id}/child_chunks/cc-id", @@ -2106,17 +2120,15 @@ class TestDatasetChildChunkApiDelete: """Test 404 when child chunk does not belong to the segment.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_db.session.scalar.return_value = mock_dataset - mock_doc_svc.get_document.return_value = Mock() + mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) segment_id = str(uuid.uuid4()) mock_segment = Mock() mock_segment.id = segment_id mock_segment.document_id = "doc-id" - mock_seg_svc.get_segment_by_id.return_value = mock_segment + mock_seg_svc.get_segment_by_ref.return_value = mock_segment - mock_child = Mock() - mock_child.segment_id = "different-segment-id" - mock_seg_svc.get_child_chunk_by_id.return_value = mock_child + mock_seg_svc.get_child_chunk_by_segment_ref.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/{segment_id}/child_chunks/cc-id", diff --git a/api/tests/unit_tests/controllers/web/test_audio.py b/api/tests/unit_tests/controllers/web/test_audio.py index a6ca441801b..a3f773f10ff 100644 --- a/api/tests/unit_tests/controllers/web/test_audio.py +++ b/api/tests/unit_tests/controllers/web/test_audio.py @@ -22,6 +22,7 @@ from controllers.web.error import ( ) from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from graphon.model_runtime.errors.invoke import InvokeError +from services.app_ref_service import MessageRef from services.errors.audio import ( AudioTooLargeServiceError, NoAudioUploadedServiceError, @@ -122,6 +123,24 @@ class TestTextApi: assert result == "audio-bytes" mock_tts.assert_called_once() + @patch("controllers.web.audio.AudioService.transcript_tts", return_value="audio-bytes") + @patch("controllers.web.audio.web_ns") + def test_happy_path_with_message_ref(self, mock_ns: MagicMock, mock_tts: MagicMock, app: Flask) -> None: + message_id = "550e8400-e29b-41d4-a716-446655440000" + mock_ns.payload = {"text": "hello", "message_id": message_id} + app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="chat") + + with app.test_request_context("/text-to-audio", method="POST"): + result = TextApi().post(app_model, _end_user()) + + assert result == "audio-bytes" + assert mock_tts.call_args.kwargs["message_ref"] == MessageRef( + "tenant-1", + "app-1", + message_id, + end_user_id="eu-1", + ) + @patch( "controllers.web.audio.AudioService.transcript_tts", side_effect=InvokeError(description="invoke failed"), diff --git a/api/tests/unit_tests/core/llm_generator/test_llm_generator.py b/api/tests/unit_tests/core/llm_generator/test_llm_generator.py index c4e610d5b07..290f58365d6 100644 --- a/api/tests/unit_tests/core/llm_generator/test_llm_generator.py +++ b/api/tests/unit_tests/core/llm_generator/test_llm_generator.py @@ -422,6 +422,13 @@ class TestLLMGenerator: "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" ) assert result == {"modified": "prompt"} + stmt = mock_scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "messages.app_id" in statement + assert "apps.tenant_id" in statement + assert "flow_id" in compiled.params.values() + assert "tenant_id" in compiled.params.values() def test_instruction_modify_legacy_with_last_run(self, mock_model_instance, model_config_entity): with patch("extensions.ext_database.db.session.scalar") as mock_scalar: @@ -439,12 +446,26 @@ class TestLLMGenerator: "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" ) assert result == {"modified": "prompt"} + stmt = mock_scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "messages.app_id" in statement + assert "apps.tenant_id" in statement + assert "flow_id" in compiled.params.values() + assert "tenant_id" in compiled.params.values() def test_instruction_modify_workflow_app_not_found(self): with patch("extensions.ext_database.db.session") as mock_session: mock_session.return_value.scalar.return_value = None with pytest.raises(ValueError, match="App not found."): LLMGenerator.instruction_modify_workflow("t", "f", "n", "c", "i", MagicMock(), "o", MagicMock()) + stmt = mock_session.return_value.scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "apps.id" in statement + assert "apps.tenant_id" in statement + assert "f" in compiled.params.values() + assert "t" in compiled.params.values() def test_instruction_modify_workflow_no_workflow(self): with patch("extensions.ext_database.db.session") as mock_session: diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py index 669794ac6d4..b2daaa42b66 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py @@ -12,6 +12,7 @@ from models.dataset import Dataset, Pipeline, PipelineCustomizedTemplate, Pipeli from models.workflow import Workflow from services.entities.knowledge_entities.rag_pipeline_entities import IconInfo, PipelineTemplateInfoEntity from services.rag_pipeline.rag_pipeline import RagPipelineService +from services.workflow_ref_service import WorkflowRef @pytest.fixture @@ -335,15 +336,15 @@ def test_update_workflow_updates_allowed_fields( workflow = SimpleNamespace( id="wf-1", marked_name="", marked_comment="", updated_by=None, updated_at=None, disallowed="original" ) + workflow_ref = WorkflowRef(tenant_id="t1", owner_id="pipeline-1", workflow_id="wf-1") session = mocker.Mock() session.scalar.return_value = workflow result = rag_pipeline_service.update_workflow( session=session, - workflow_id="wf-1", - tenant_id="t1", account_id="u1", data={"marked_name": "v1", "marked_comment": "release", "disallowed": "hacked"}, + workflow_ref=workflow_ref, ) assert result.marked_name == "v1" @@ -360,15 +361,43 @@ def test_update_workflow_returns_none_when_not_found( result = rag_pipeline_service.update_workflow( session=session, - workflow_id="wf-missing", - tenant_id="t1", account_id="u1", data={"marked_name": "v1"}, + workflow_ref=WorkflowRef(tenant_id="t1", owner_id="pipeline-1", workflow_id="wf-missing"), ) assert result is None +def test_update_workflow_with_ref_scopes_lookup_to_pipeline( + mocker: MockerFixture, rag_pipeline_service: RagPipelineService +) -> None: + workflow = SimpleNamespace( + id="wf-1", marked_name="", marked_comment="", updated_by=None, updated_at=None, disallowed="original" + ) + workflow_ref = WorkflowRef(tenant_id="t1", owner_id="pipeline-1", workflow_id="wf-1") + session = mocker.Mock() + session.scalar.return_value = workflow + + result = rag_pipeline_service.update_workflow( + session=session, + account_id="u1", + data={"marked_name": "v1"}, + workflow_ref=workflow_ref, + ) + + stmt = session.scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "workflows.id" in statement + assert "workflows.tenant_id" in statement + assert "workflows.app_id" in statement + assert "wf-1" in compiled.params.values() + assert "t1" in compiled.params.values() + assert "pipeline-1" in compiled.params.values() + assert result is workflow + + # --- get_rag_pipeline_paginate_workflow_runs --- @@ -1627,6 +1656,8 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( def __init__(self): self._values = { ("sys", "invoke_from"): SimpleNamespace(value=InvokeFrom.PUBLISHED_PIPELINE), + ("sys", "app_id"): SimpleNamespace(value="pipeline-1"), + ("sys", "dataset_id"): SimpleNamespace(value="dataset-1"), ("sys", "document_id"): SimpleNamespace(value="doc-1"), } @@ -1660,7 +1691,8 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( ) document = SimpleNamespace(indexing_status="waiting", error=None) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", return_value=document) + scalar_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=document) + get_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get") add_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.add") commit_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.commit") @@ -1672,6 +1704,19 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( ) assert result.status == WorkflowNodeExecutionStatus.FAILED + stmt = scalar_mock.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "documents.id" in statement + assert "documents.tenant_id" in statement + assert "documents.dataset_id" in statement + assert "datasets.tenant_id" in statement + assert "datasets.pipeline_id" in statement + assert "doc-1" in compiled.params.values() + assert "t1" in compiled.params.values() + assert "dataset-1" in compiled.params.values() + assert "pipeline-1" in compiled.params.values() + get_mock.assert_not_called() assert document.indexing_status == "error" assert document.error == "boom" add_mock.assert_called_once_with(document) diff --git a/api/tests/unit_tests/services/test_annotation_service.py b/api/tests/unit_tests/services/test_annotation_service.py index 7574c13342a..c483440dd2c 100644 --- a/api/tests/unit_tests/services/test_annotation_service.py +++ b/api/tests/unit_tests/services/test_annotation_service.py @@ -15,6 +15,7 @@ from werkzeug.exceptions import NotFound from models.model import App, AppAnnotationHitHistory, AppAnnotationSetting, Message, MessageAnnotation from services.annotation_service import AppAnnotationService +from services.app_ref_service import AnnotationRef, AppRef def _make_app(app_id: str = "app-1", tenant_id: str = "tenant-1") -> MagicMock: @@ -25,6 +26,14 @@ def _make_app(app_id: str = "app-1", tenant_id: str = "tenant-1") -> MagicMock: return app +def _make_app_ref(app: MagicMock) -> AppRef: + return AppRef(tenant_id=app.tenant_id, app_id=app.id) + + +def _make_annotation_ref(app: MagicMock, annotation_id: str = "ann-1") -> AnnotationRef: + return AnnotationRef(tenant_id=app.tenant_id, app_id=app.id, annotation_id=annotation_id) + + def _make_user(user_id: str = "user-1") -> MagicMock: user = MagicMock() user.id = user_id @@ -41,9 +50,10 @@ def _make_message(message_id: str = "msg-1", app_id: str = "app-1") -> MagicMock return message -def _make_annotation(annotation_id: str = "ann-1") -> MagicMock: +def _make_annotation(annotation_id: str = "ann-1", app_id: str = "app-1") -> MagicMock: annotation = MagicMock(spec=MessageAnnotation) annotation.id = annotation_id + annotation.app_id = app_id annotation.content = "" annotation.question = "" annotation.question_text = "" @@ -66,6 +76,15 @@ def _make_file(content: bytes) -> FileStorage: return FileStorage(stream=BytesIO(content)) +def _assert_statement_binds_annotation(stmt: Any, annotation_id: str, app_id: str) -> None: + compiled = stmt.compile() + statement = str(compiled) + assert "message_annotations.id" in statement + assert "message_annotations.app_id" in statement + assert annotation_id in compiled.params.values() + assert app_id in compiled.params.values() + + class TestAppAnnotationServiceUpInsert: """Test suite for up_insert_app_annotation_from_message.""" @@ -537,23 +556,6 @@ class TestAppAnnotationServiceDirectManipulation: tenant_id = "tenant-1" app = _make_app() - with ( - patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, - ): - mock_db.session.scalar.return_value = app - mock_db.session.get.return_value = None - - # Act & Assert - with pytest.raises(NotFound): - AppAnnotationService.update_app_annotation_directly(args, app.id, "ann-1", mock_db.session) - - def test_update_app_annotation_directly_should_raise_not_found_when_app_missing(self) -> None: - """Test missing app raises NotFound in update path.""" - # Arrange - args = {"answer": "hello", "question": "q1"} - tenant_id = "tenant-1" - with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, @@ -562,7 +564,11 @@ class TestAppAnnotationServiceDirectManipulation: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.update_app_annotation_directly(args, "app-1", "ann-1", mock_db.session) + AppAnnotationService.update_app_annotation_directly( + args, + _make_annotation_ref(app, "ann-1"), + mock_db.session, + ) def test_update_app_annotation_directly_should_raise_value_error_when_question_missing(self) -> None: """Test missing question raises ValueError.""" @@ -576,12 +582,13 @@ class TestAppAnnotationServiceDirectManipulation: patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, ): - mock_db.session.scalar.return_value = app - mock_db.session.get.return_value = annotation + mock_db.session.scalar.return_value = annotation # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.update_app_annotation_directly(args, app.id, annotation.id, mock_db.session) + AppAnnotationService.update_app_annotation_directly( + args, _make_annotation_ref(app, annotation.id), mock_db.session + ) def test_update_app_annotation_directly_should_update_annotation_and_index(self) -> None: """Test update changes fields and triggers index update.""" @@ -598,16 +605,19 @@ class TestAppAnnotationServiceDirectManipulation: patch("services.annotation_service.db") as mock_db, patch("services.annotation_service.update_annotation_to_index_task") as mock_task, ): - mock_db.session.scalar.side_effect = [app, setting] - mock_db.session.get.return_value = annotation + mock_db.session.scalar.side_effect = [annotation, setting] # Act - result = AppAnnotationService.update_app_annotation_directly(args, app.id, annotation.id, mock_db.session) + result = AppAnnotationService.update_app_annotation_directly( + args, _make_annotation_ref(app, annotation.id), mock_db.session + ) # Assert assert result == annotation assert annotation.content == "hello" assert annotation.question == "q1" + _assert_statement_binds_annotation(mock_db.session.scalar.call_args_list[0].args[0], annotation.id, app.id) + mock_db.session.get.assert_not_called() mock_db.session.commit.assert_called_once() mock_task.delay.assert_called_once_with( annotation.id, @@ -632,17 +642,18 @@ class TestAppAnnotationServiceDirectManipulation: patch("services.annotation_service.db") as mock_db, patch("services.annotation_service.delete_annotation_index_task") as mock_task, ): - mock_db.session.scalar.side_effect = [app, setting] - mock_db.session.get.return_value = annotation + mock_db.session.scalar.side_effect = [annotation, setting] scalars_result = MagicMock() scalars_result.all.return_value = [history1, history2] mock_db.session.scalars.return_value = scalars_result # Act - AppAnnotationService.delete_app_annotation(app.id, annotation.id, mock_db.session) + AppAnnotationService.delete_app_annotation(_make_annotation_ref(app, annotation.id), mock_db.session) # Assert + _assert_statement_binds_annotation(mock_db.session.scalar.call_args_list[0].args[0], annotation.id, app.id) + mock_db.session.get.assert_not_called() mock_db.session.delete.assert_any_call(annotation) mock_db.session.delete.assert_any_call(history1) mock_db.session.delete.assert_any_call(history2) @@ -654,21 +665,6 @@ class TestAppAnnotationServiceDirectManipulation: setting.collection_binding_id, ) - def test_delete_app_annotation_should_raise_not_found_when_app_missing(self) -> None: - """Test delete raises NotFound when app is missing.""" - # Arrange - tenant_id = "tenant-1" - - with ( - patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, - ): - mock_db.session.scalar.return_value = None - - # Act & Assert - with pytest.raises(NotFound): - AppAnnotationService.delete_app_annotation("app-1", "ann-1", mock_db.session) - def test_delete_app_annotation_should_raise_not_found_when_annotation_missing(self) -> None: """Test delete raises NotFound when annotation is missing.""" # Arrange @@ -679,12 +675,11 @@ class TestAppAnnotationServiceDirectManipulation: patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, ): - mock_db.session.scalar.return_value = app - mock_db.session.get.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.delete_app_annotation(app.id, "ann-1", mock_db.session) + AppAnnotationService.delete_app_annotation(_make_annotation_ref(app, "ann-1"), mock_db.session) def test_delete_app_annotations_in_batch_should_return_zero_when_none_found(self) -> None: """Test batch delete returns zero when no annotations found.""" @@ -696,30 +691,14 @@ class TestAppAnnotationServiceDirectManipulation: patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, ): - mock_db.session.scalar.return_value = app mock_db.session.execute.return_value.all.return_value = [] # Act - result = AppAnnotationService.delete_app_annotations_in_batch(app.id, ["ann-1"]) + result = AppAnnotationService.delete_app_annotations_in_batch(_make_app_ref(app), ["ann-1"]) # Assert assert result == {"deleted_count": 0} - def test_delete_app_annotations_in_batch_should_raise_not_found_when_app_missing(self) -> None: - """Test batch delete raises NotFound when app is missing.""" - # Arrange - tenant_id = "tenant-1" - - with ( - patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, - ): - mock_db.session.scalar.return_value = None - - # Act & Assert - with pytest.raises(NotFound): - AppAnnotationService.delete_app_annotations_in_batch("app-1", ["ann-1"]) - def test_delete_app_annotations_in_batch_should_delete_annotations_and_histories(self) -> None: """Test batch delete removes annotations and triggers index deletion.""" # Arrange @@ -734,8 +713,6 @@ class TestAppAnnotationServiceDirectManipulation: patch("services.annotation_service.db") as mock_db, patch("services.annotation_service.delete_annotation_index_task") as mock_task, ): - mock_db.session.scalar.return_value = app - # First execute().all() for multi-column query, subsequent execute() calls for deletes execute_result_multi = MagicMock() execute_result_multi.all.return_value = [(annotation1, setting), (annotation2, None)] @@ -744,10 +721,17 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.execute.side_effect = [execute_result_multi, MagicMock(), execute_result_delete] # Act - result = AppAnnotationService.delete_app_annotations_in_batch(app.id, ["ann-1", "ann-2"]) + result = AppAnnotationService.delete_app_annotations_in_batch(_make_app_ref(app), ["ann-1", "ann-2"]) # Assert assert result == {"deleted_count": 2} + fetch_stmt = mock_db.session.execute.call_args_list[0].args[0] + compiled = fetch_stmt.compile() + statement = str(compiled) + assert "message_annotations.id IN" in statement + assert "message_annotations.app_id" in statement + assert ["ann-1", "ann-2"] in compiled.params.values() + assert app.id in compiled.params.values() mock_task.delay.assert_called_once_with(annotation1.id, app.id, tenant_id, setting.collection_binding_id) mock_db.session.commit.assert_called_once() @@ -1094,20 +1078,17 @@ class TestAppAnnotationServiceBatchImport: class TestAppAnnotationServiceHitHistoryAndSettings: """Test suite for hit history and settings methods.""" - def test_get_annotation_hit_histories_should_raise_not_found_when_app_missing(self) -> None: - """Test missing app raises NotFound.""" + def test_get_annotation_hit_histories_should_raise_not_found_when_annotation_missing(self) -> None: + """Test missing annotation raises NotFound.""" # Arrange - tenant_id = "tenant-1" + app = _make_app() - with ( - patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, - ): + with patch("services.annotation_service.db") as mock_db: mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_annotation_hit_histories("app-1", "ann-1", 1, 10) + AppAnnotationService.get_annotation_hit_histories(_make_annotation_ref(app, "ann-1"), 1, 10) def test_get_annotation_hit_histories_should_return_items_and_total(self) -> None: """Test hit histories pagination returns items and total.""" @@ -1121,33 +1102,21 @@ class TestAppAnnotationServiceHitHistoryAndSettings: patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, ): - mock_db.session.scalar.return_value = app - mock_db.session.get.return_value = annotation + mock_db.session.scalar.return_value = annotation mock_db.paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_hit_histories(app.id, annotation.id, 1, 10) + items, total = AppAnnotationService.get_annotation_hit_histories( + _make_annotation_ref(app, annotation.id), + 1, + 10, + ) # Assert assert items == ["h1"] assert total == 2 - - def test_get_annotation_hit_histories_should_raise_not_found_when_annotation_missing(self) -> None: - """Test missing annotation raises NotFound.""" - # Arrange - tenant_id = "tenant-1" - app = _make_app() - - with ( - patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, - ): - mock_db.session.scalar.return_value = app - mock_db.session.get.return_value = None - - # Act & Assert - with pytest.raises(NotFound): - AppAnnotationService.get_annotation_hit_histories(app.id, "ann-1", 1, 10) + _assert_statement_binds_annotation(mock_db.session.scalar.call_args_list[0].args[0], annotation.id, app.id) + mock_db.session.get.assert_not_called() def test_get_annotation_by_id_should_return_none_when_missing(self) -> None: """Test get_annotation_by_id returns None when not found.""" diff --git a/api/tests/unit_tests/services/test_audio_service.py b/api/tests/unit_tests/services/test_audio_service.py index 788a47c5c31..bd865dd5dd3 100644 --- a/api/tests/unit_tests/services/test_audio_service.py +++ b/api/tests/unit_tests/services/test_audio_service.py @@ -62,6 +62,7 @@ from werkzeug.datastructures import FileStorage from models.enums import MessageStatus from models.model import App, AppMode, AppModelConfig, Message from models.workflow import Workflow +from services.app_ref_service import MessageRef from services.audio_service import AudioService from services.errors.audio import ( AudioTooLargeServiceError, @@ -521,9 +522,16 @@ class TestAudioServiceTTS: # Arrange app = factory.create_app_mock(mode=AppMode.CHAT) message_id = "00000000-0000-0000-0000-000000000001" + message_ref = MessageRef( + tenant_id=app.tenant_id, + app_id=app.id, + message_id=message_id, + end_user_id="end-user-1", + account_id="account-1", + ) message = factory.create_message_mock(message_id=message_id, answer="Message answer") session = MagicMock() - session.get.return_value = message + session.scalar.return_value = message mock_model_manager = mock_model_manager_class.return_value mock_model_instance = MagicMock() @@ -534,13 +542,25 @@ class TestAudioServiceTTS: result = AudioService.transcript_tts( app_model=app, session=session, - message_id=message_id, + message_ref=message_ref, voice="message-voice", ) # Assert assert result == b"message audio" - session.get.assert_called_once_with(Message, message_id) + session.scalar.assert_called_once() + session.get.assert_not_called() + stmt = session.scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "messages.id" in statement + assert "messages.app_id" in statement + assert "messages.from_end_user_id" in statement + assert "messages.from_account_id" in statement + assert message_id in compiled.params.values() + assert app.id in compiled.params.values() + assert "end-user-1" in compiled.params.values() + assert "account-1" in compiled.params.values() mock_model_instance.invoke_tts.assert_called_once_with( content_text="Message answer", voice="message-voice", diff --git a/api/tests/unit_tests/services/test_dataset_service_document.py b/api/tests/unit_tests/services/test_dataset_service_document.py index c108b06ac6b..02661fbe1f3 100644 --- a/api/tests/unit_tests/services/test_dataset_service_document.py +++ b/api/tests/unit_tests/services/test_dataset_service_document.py @@ -1,5 +1,7 @@ """Unit tests for DocumentService behaviors in dataset_service.""" +from services.dataset_ref_service import DatasetRef + from .dataset_service_test_helpers import ( Account, BuiltInField, @@ -103,6 +105,39 @@ class TestDocumentServiceMutations: assert DocumentService.check_archived(document) is expected + def test_delete_documents_limits_query_and_cleanup_to_dataset_ref(self): + dataset = _make_dataset(dataset_id="dataset-1", tenant_id="tenant-1") + dataset.doc_form = "paragraph_index" + document = _make_document(document_id="doc-1", dataset_id=dataset.id, tenant_id=dataset.tenant_id) + document.data_source_info_dict = {} + + with ( + patch("services.dataset_service.db") as mock_db, + patch("services.dataset_service.batch_clean_document_task") as clean_task, + ): + mock_db.session.scalars.return_value.all.return_value = [document] + + dataset_ref = DatasetRef(tenant_id=dataset.tenant_id, dataset_id=dataset.id) + DocumentService.delete_documents( + dataset_ref, + ["doc-1", "other-doc"], + dataset.doc_form, + mock_db.session, + ) + + stmt = mock_db.session.scalars.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "documents.id IN" in statement + assert "documents.tenant_id" in statement + assert "documents.dataset_id" in statement + assert ["doc-1", "other-doc"] in compiled.params.values() + assert dataset.tenant_id in compiled.params.values() + assert dataset.id in compiled.params.values() + mock_db.session.delete.assert_called_once_with(document) + mock_db.session.commit.assert_called_once() + clean_task.delay.assert_called_once_with(["doc-1"], dataset.id, dataset.doc_form, []) + def test_rename_document_raises_when_dataset_is_missing(self, rename_account_context): session = MagicMock() diff --git a/api/tests/unit_tests/services/test_dataset_service_segment.py b/api/tests/unit_tests/services/test_dataset_service_segment.py index a625f17ef37..f2c08324774 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -1,5 +1,7 @@ """Unit tests for SegmentService behaviors in dataset_service.""" +from services.dataset_ref_service import DatasetRef, DatasetRefService + from .dataset_service_test_helpers import ( Account, ChildChunk, @@ -24,6 +26,41 @@ from .dataset_service_test_helpers import ( ) +def _make_segment_ref(segment_id: str = "segment-1"): + dataset = _make_dataset() + document = _make_document(dataset_id=dataset.id, tenant_id=dataset.tenant_id) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + document_ref = DatasetRefService.create_document_ref(dataset_ref, document) + assert document_ref is not None + return DatasetRefService.create_segment_ref(document_ref, segment_id) + + +class TestDatasetRefService: + """Unit tests for typed dataset resource refs.""" + + def test_dataset_ref_is_plain_named_tuple(self): + dataset_ref = DatasetRef("tenant-1", "dataset-1") + + assert dataset_ref.tenant_id == "tenant-1" + assert dataset_ref.dataset_id == "dataset-1" + assert tuple(dataset_ref) == ("tenant-1", "dataset-1") + + def test_create_document_ref_rejects_document_outside_dataset(self): + dataset = _make_dataset(dataset_id="dataset-1", tenant_id="tenant-1") + document = _make_document(document_id="doc-1", dataset_id="other-dataset", tenant_id="tenant-1") + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + + assert DatasetRefService.create_document_ref(dataset_ref, document) is None + + def test_create_segment_ref_carries_full_parent_chain(self): + segment_ref = _make_segment_ref() + + assert segment_ref.tenant_id == "tenant-1" + assert segment_ref.dataset_id == "dataset-1" + assert segment_ref.document_id == "doc-1" + assert segment_ref.segment_id == "segment-1" + + class TestSegmentServiceChildChunks: """Unit tests for child-chunk CRUD helpers.""" @@ -265,6 +302,23 @@ class TestSegmentServiceQueries: assert result is None + def test_get_child_chunk_by_segment_ref_uses_full_ownership_chain(self): + child_chunk = _make_child_chunk() + segment_ref = _make_segment_ref() + + with patch("services.dataset_service.db") as mock_db: + mock_db.session.scalar.return_value = child_chunk + result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref) + + assert result is child_chunk + stmt = mock_db.session.scalar.call_args.args[0] + sql = str(stmt.compile(compile_kwargs={"literal_binds": True})) + assert "child_chunks.id = 'child-a'" in sql + assert "child_chunks.tenant_id = 'tenant-1'" in sql + assert "child_chunks.dataset_id = 'dataset-1'" in sql + assert "child_chunks.document_id = 'doc-1'" in sql + assert "child_chunks.segment_id = 'segment-1'" in sql + def test_get_segments_uses_status_and_keyword_filters(self): paginated = SimpleNamespace(items=["segment"], total=1) @@ -312,6 +366,32 @@ class TestSegmentServiceQueries: assert result is None + def test_get_segment_by_ref_uses_full_ownership_chain(self): + segment = DocumentSegment( + tenant_id="tenant-1", + dataset_id="dataset-1", + document_id="doc-1", + position=1, + content="segment", + word_count=7, + tokens=2, + created_by="user-1", + ) + segment.id = "segment-1" + segment_ref = _make_segment_ref() + + with patch("services.dataset_service.db") as mock_db: + mock_db.session.scalar.return_value = segment + result = SegmentService.get_segment_by_ref(segment_ref) + + assert result is segment + stmt = mock_db.session.scalar.call_args.args[0] + sql = str(stmt.compile(compile_kwargs={"literal_binds": True})) + assert "document_segments.id = 'segment-1'" in sql + assert "document_segments.tenant_id = 'tenant-1'" in sql + assert "document_segments.dataset_id = 'dataset-1'" in sql + assert "document_segments.document_id = 'doc-1'" in sql + def test_get_segments_by_document_and_dataset_returns_scalars_result(self): segment = DocumentSegment( tenant_id="tenant-1", diff --git a/api/tests/unit_tests/services/test_tag_service.py b/api/tests/unit_tests/services/test_tag_service.py index 73df7cc2673..d304ab23cea 100644 --- a/api/tests/unit_tests/services/test_tag_service.py +++ b/api/tests/unit_tests/services/test_tag_service.py @@ -5,7 +5,7 @@ from pytest_mock import MockerFixture from werkzeug.exceptions import NotFound from models.enums import TagType -from services.tag_service import TagBindingCreatePayload, TagBindingDeletePayload, TagService +from services.tag_service import TagBindingCreatePayload, TagBindingDeletePayload, TagService, UpdateTagPayload @pytest.fixture @@ -78,6 +78,71 @@ def test_delete_tag_binding_does_not_commit_when_no_rows_deleted(mocker: MockerF db_session.commit.assert_not_called() +def test_update_tags_scopes_lookup_to_current_tenant_and_type(current_user, db_session): + tag = SimpleNamespace(id="tag-1", name="old", type=TagType.KNOWLEDGE) + db_session.scalar.side_effect = [tag, None] + + result = TagService.update_tags(UpdateTagPayload(name="new"), "tag-1", db_session, tag_type=TagType.KNOWLEDGE) + + stmt = db_session.scalar.call_args_list[0].args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "tags.id" in statement + assert "tags.tenant_id" in statement + assert "tags.type" in statement + assert "tag-1" in compiled.params.values() + assert current_user.current_tenant_id in compiled.params.values() + assert TagType.KNOWLEDGE in compiled.params.values() + assert result is tag + assert tag.name == "new" + db_session.commit.assert_called_once() + + +def test_get_tag_binding_count_scopes_lookup_to_current_tenant_and_type(current_user, db_session): + db_session.scalar.return_value = 3 + + result = TagService.get_tag_binding_count("tag-1", db_session, tag_type=TagType.KNOWLEDGE) + + stmt = db_session.scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "tag_bindings.tag_id" in statement + assert "tags.tenant_id" in statement + assert "tags.type" in statement + assert "tag-1" in compiled.params.values() + assert current_user.current_tenant_id in compiled.params.values() + assert TagType.KNOWLEDGE in compiled.params.values() + assert result == 3 + + +def test_delete_tag_scopes_lookup_and_bindings_to_current_tenant(current_user, db_session): + tag = SimpleNamespace(id="tag-1", name="old", type=TagType.KNOWLEDGE) + binding = SimpleNamespace(id="binding-1") + db_session.scalar.return_value = tag + db_session.scalars.return_value.all.return_value = [binding] + + TagService.delete_tag("tag-1", db_session, tag_type=TagType.KNOWLEDGE) + + tag_stmt = db_session.scalar.call_args.args[0] + tag_compiled = tag_stmt.compile() + assert "tags.id" in str(tag_compiled) + assert "tags.tenant_id" in str(tag_compiled) + assert "tags.type" in str(tag_compiled) + assert "tag-1" in tag_compiled.params.values() + assert current_user.current_tenant_id in tag_compiled.params.values() + assert TagType.KNOWLEDGE in tag_compiled.params.values() + + binding_stmt = db_session.scalars.call_args.args[0] + binding_compiled = binding_stmt.compile() + assert "tag_bindings.tag_id" in str(binding_compiled) + assert "tag_bindings.tenant_id" in str(binding_compiled) + assert "tag-1" in binding_compiled.params.values() + assert current_user.current_tenant_id in binding_compiled.params.values() + db_session.delete.assert_any_call(tag) + db_session.delete.assert_any_call(binding) + db_session.commit.assert_called_once() + + def test_get_target_ids_by_tag_ids_returns_empty_without_query_for_empty_input(db_session): result = TagService.get_target_ids_by_tag_ids(TagType.SNIPPET, "tenant-1", [], db_session) diff --git a/api/tests/unit_tests/services/test_workflow_service.py b/api/tests/unit_tests/services/test_workflow_service.py index 1199fc773f2..b322909a128 100644 --- a/api/tests/unit_tests/services/test_workflow_service.py +++ b/api/tests/unit_tests/services/test_workflow_service.py @@ -36,6 +36,7 @@ from models.model import App, AppMode from models.workflow import Workflow, WorkflowType from services.errors.app import IsDraftWorkflowError, TriggerNodeLimitExceededError, WorkflowHashNotEqualError from services.errors.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError +from services.workflow_ref_service import WorkflowRef from services.workflow_service import ( WorkflowService, _rebuild_file_for_user_inputs_in_start_node, @@ -1008,6 +1009,8 @@ class TestWorkflowService: """ workflow_id = "workflow-123" tenant_id = "tenant-456" + app_id = "app-789" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id=app_id, workflow_id=workflow_id) account_id = "user-123" mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock(workflow_id=workflow_id) @@ -1021,10 +1024,9 @@ class TestWorkflowService: result = workflow_service.update_workflow( session=mock_session, - workflow_id=workflow_id, - tenant_id=tenant_id, account_id=account_id, data={"marked_name": "Updated Name", "marked_comment": "Updated Comment"}, + workflow_ref=workflow_ref, ) assert result == mock_workflow @@ -1044,14 +1046,42 @@ class TestWorkflowService: result = workflow_service.update_workflow( session=mock_session, - workflow_id="nonexistent", - tenant_id="tenant-456", account_id="user-123", data={"marked_name": "Test"}, + workflow_ref=WorkflowRef(tenant_id="tenant-456", owner_id="app-789", workflow_id="nonexistent"), ) assert result is None + def test_update_workflow_with_ref_scopes_lookup_to_app(self, workflow_service: WorkflowService): + """Test update_workflow includes the trusted app owner in the lookup.""" + workflow_id = "workflow-123" + tenant_id = "tenant-456" + app_id = "app-789" + account_id = "user-123" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id=app_id, workflow_id=workflow_id) + mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock(workflow_id=workflow_id) + mock_session = MagicMock() + mock_session.scalar.return_value = mock_workflow + + result = workflow_service.update_workflow( + session=mock_session, + account_id=account_id, + data={"marked_name": "Updated Name"}, + workflow_ref=workflow_ref, + ) + + stmt = mock_session.scalar.call_args.args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "workflows.id" in statement + assert "workflows.tenant_id" in statement + assert "workflows.app_id" in statement + assert workflow_id in compiled.params.values() + assert tenant_id in compiled.params.values() + assert app_id in compiled.params.values() + assert result == mock_workflow + # ==================== Delete Workflow Tests ==================== # These tests verify workflow deletion with safety checks @@ -1064,6 +1094,8 @@ class TestWorkflowService: """ workflow_id = "workflow-123" tenant_id = "tenant-456" + app_id = "app-789" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id=app_id, workflow_id=workflow_id) mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock(workflow_id=workflow_id, version="v1") mock_session = MagicMock() @@ -1078,13 +1110,35 @@ class TestWorkflowService: mock_select.return_value = mock_stmt mock_stmt.where.return_value = mock_stmt - result = workflow_service.delete_workflow( - session=mock_session, workflow_id=workflow_id, tenant_id=tenant_id - ) + result = workflow_service.delete_workflow(session=mock_session, workflow_ref=workflow_ref) assert result is True mock_session.delete.assert_called_once_with(mock_workflow) + def test_delete_workflow_with_ref_scopes_lookup_to_app(self, workflow_service: WorkflowService): + """Test delete_workflow includes the trusted app owner in the lookup.""" + workflow_id = "workflow-123" + tenant_id = "tenant-456" + app_id = "app-789" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id=app_id, workflow_id=workflow_id) + mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock(workflow_id=workflow_id, version="v1") + mock_session = MagicMock() + mock_session.scalar.side_effect = [mock_workflow, None, None] + + result = workflow_service.delete_workflow(session=mock_session, workflow_ref=workflow_ref) + + stmt = mock_session.scalar.call_args_list[0].args[0] + compiled = stmt.compile() + statement = str(compiled) + assert "workflows.id" in statement + assert "workflows.tenant_id" in statement + assert "workflows.app_id" in statement + assert workflow_id in compiled.params.values() + assert tenant_id in compiled.params.values() + assert app_id in compiled.params.values() + assert result is True + mock_session.delete.assert_called_once_with(mock_workflow) + def test_delete_workflow_draft_raises_error(self, workflow_service: WorkflowService): """ Test delete_workflow raises error when trying to delete draft. @@ -1094,6 +1148,7 @@ class TestWorkflowService: """ workflow_id = "workflow-123" tenant_id = "tenant-456" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id="app-789", workflow_id=workflow_id) mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock( workflow_id=workflow_id, version=Workflow.VERSION_DRAFT ) @@ -1107,7 +1162,7 @@ class TestWorkflowService: mock_stmt.where.return_value = mock_stmt with pytest.raises(DraftWorkflowDeletionError, match="Cannot delete draft workflow"): - workflow_service.delete_workflow(session=mock_session, workflow_id=workflow_id, tenant_id=tenant_id) + workflow_service.delete_workflow(session=mock_session, workflow_ref=workflow_ref) def test_delete_workflow_in_use_by_app_raises_error(self, workflow_service: WorkflowService): """ @@ -1118,6 +1173,7 @@ class TestWorkflowService: """ workflow_id = "workflow-123" tenant_id = "tenant-456" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id="app-789", workflow_id=workflow_id) mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock(workflow_id=workflow_id, version="v1") mock_app = TestWorkflowAssociatedDataFactory.create_app_mock(workflow_id=workflow_id) @@ -1130,7 +1186,7 @@ class TestWorkflowService: mock_stmt.where.return_value = mock_stmt with pytest.raises(WorkflowInUseError, match="currently in use by app"): - workflow_service.delete_workflow(session=mock_session, workflow_id=workflow_id, tenant_id=tenant_id) + workflow_service.delete_workflow(session=mock_session, workflow_ref=workflow_ref) def test_delete_workflow_published_as_tool_raises_error(self, workflow_service: WorkflowService): """ @@ -1142,6 +1198,7 @@ class TestWorkflowService: """ workflow_id = "workflow-123" tenant_id = "tenant-456" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id="app-789", workflow_id=workflow_id) mock_workflow = TestWorkflowAssociatedDataFactory.create_workflow_mock(workflow_id=workflow_id, version="v1") mock_tool_provider = MagicMock() @@ -1154,12 +1211,13 @@ class TestWorkflowService: mock_stmt.where.return_value = mock_stmt with pytest.raises(WorkflowInUseError, match="published as a tool"): - workflow_service.delete_workflow(session=mock_session, workflow_id=workflow_id, tenant_id=tenant_id) + workflow_service.delete_workflow(session=mock_session, workflow_ref=workflow_ref) def test_delete_workflow_not_found_raises_error(self, workflow_service: WorkflowService): """Test delete_workflow raises error when workflow not found.""" workflow_id = "nonexistent" tenant_id = "tenant-456" + workflow_ref = WorkflowRef(tenant_id=tenant_id, owner_id="app-789", workflow_id=workflow_id) mock_session = MagicMock() mock_session.scalar.return_value = None @@ -1170,7 +1228,7 @@ class TestWorkflowService: mock_stmt.where.return_value = mock_stmt with pytest.raises(ValueError, match="not found"): - workflow_service.delete_workflow(session=mock_session, workflow_id=workflow_id, tenant_id=tenant_id) + workflow_service.delete_workflow(session=mock_session, workflow_ref=workflow_ref) # ==================== Get Default Block Config Tests ==================== # These tests verify retrieval of default node configurations diff --git a/scripts/stress-test/test_setup_scripts.py b/scripts/stress-test/test_setup_scripts.py index 01ddd2b8fb9..bdca551c682 100644 --- a/scripts/stress-test/test_setup_scripts.py +++ b/scripts/stress-test/test_setup_scripts.py @@ -113,8 +113,6 @@ def test_plugin_install_response_without_task_is_non_blocking(): def test_import_response_with_warnings_and_app_id_is_success(): import_workflow_app = _load_setup_module("import_workflow_app") - assert import_workflow_app.is_successful_import_response( - {"status": "completed-with-warnings", "app_id": "app-id"} - ) + assert import_workflow_app.is_successful_import_response({"status": "completed-with-warnings", "app_id": "app-id"}) assert not import_workflow_app.is_successful_import_response({"status": "failed", "app_id": "app-id"}) assert not import_workflow_app.is_successful_import_response({"status": "completed"}) From 5ce13d1773b676304a7a6516fb4e287763f43541 Mon Sep 17 00:00:00 2001 From: Xiyuan Chen <52963600+GareArc@users.noreply.github.com> Date: Tue, 30 Jun 2026 01:20:24 -0700 Subject: [PATCH 09/54] fix(api): register rbac-migrate-dataset-permissions CLI command (#38204) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/extensions/ext_commands.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/api/extensions/ext_commands.py b/api/extensions/ext_commands.py index a85f6569978..ac96e3031a8 100644 --- a/api/extensions/ext_commands.py +++ b/api/extensions/ext_commands.py @@ -27,6 +27,7 @@ def init_app(app: DifyApp): install_plugins, install_rag_pipeline_plugins, migrate_data_for_plugin, + migrate_dataset_permissions_to_rbac, migrate_member_roles_to_rbac, migrate_oss, migration_data_wizard, @@ -56,6 +57,7 @@ def init_app(app: DifyApp): upgrade_db, fix_app_site_missing, migrate_data_for_plugin, + migrate_dataset_permissions_to_rbac, migrate_member_roles_to_rbac, backfill_plugin_auto_upgrade, extract_plugins, From 44e85c0023b8e7efa485a2057740cc974761db97 Mon Sep 17 00:00:00 2001 From: chariri Date: Tue, 30 Jun 2026 17:51:32 +0900 Subject: [PATCH 10/54] fix(api): Fixing API contract generation infrastructure (#38042) --- api/dev/lint_response_contracts.py | 8 ++- .../commands/test_lint_response_contracts.py | 23 ++++++ packages/contracts/openapi-ts.api.config.ts | 72 +++++++++---------- 3 files changed, 65 insertions(+), 38 deletions(-) diff --git a/api/dev/lint_response_contracts.py b/api/dev/lint_response_contracts.py index 77a2d2dc818..6cfc1c6b446 100644 --- a/api/dev/lint_response_contracts.py +++ b/api/dev/lint_response_contracts.py @@ -595,9 +595,15 @@ def line_has_ignore_marker(line: str) -> bool: return any(ignore_marker in normalized for ignore_marker in IGNORE_COMMENT_MARKERS) +def node_ignore_scan_end_lineno(node: ast.ClassDef | MethodNode) -> int: + if isinstance(node, ast.ClassDef): + return node.lineno + return node.end_lineno or node.lineno + + def node_has_ignore_comment(lines: Sequence[str], node: ast.ClassDef | MethodNode) -> bool: start = node_start_lineno(node) - end = node.end_lineno or node.lineno + end = node_ignore_scan_end_lineno(node) if any(line_has_ignore_marker(line) for line in lines[start - 1 : end]): return True diff --git a/api/tests/unit_tests/commands/test_lint_response_contracts.py b/api/tests/unit_tests/commands/test_lint_response_contracts.py index 68c4cbf966e..17f3156d127 100644 --- a/api/tests/unit_tests/commands/test_lint_response_contracts.py +++ b/api/tests/unit_tests/commands/test_lint_response_contracts.py @@ -168,6 +168,29 @@ class RegularApi(Resource): assert checks[0].classification == "valid" +def test_method_ignore_comment_does_not_skip_sibling_methods(tmp_path: Path): + checks = _checks_for_source( + tmp_path, + """ +@ns.route("/mixed") +class MixedApi(Resource): + # response-contract:ignore binary response + @ns.response(200, "Binary file") + def get(self): + return send_file(path) + + @ns.response(200, "OK", ns.models[RegularResponse.__name__]) + def post(self): + return dump_response(RegularResponse, {}) +""", + ) + + assert len(checks) == 1 + assert checks[0].class_name == "MixedApi" + assert checks[0].method == "post" + assert checks[0].classification == "valid" + + def test_main_is_report_only_by_default_for_mismatches(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): module = _load_lint_response_contracts_module() controller_path = tmp_path / "controllers" / "sample.py" diff --git a/packages/contracts/openapi-ts.api.config.ts b/packages/contracts/openapi-ts.api.config.ts index 1adbf4fda8e..99e1f79ee2b 100644 --- a/packages/contracts/openapi-ts.api.config.ts +++ b/packages/contracts/openapi-ts.api.config.ts @@ -63,15 +63,6 @@ const operationMethods = new Set(['delete', 'get', 'patch', 'post', 'put']) const pydanticDecimalStringPattern = '^(?!^[-+.]*$)[+-]?0*\\d*\\.?\\d*$' const codegenSafeDecimalStringPattern = '^(?![-+.]*$)[+-]?0*\\d*\\.?\\d*$' -const opaqueJsonContent = (): Record => ({ - 'application/json': { - schema: { - additionalProperties: true, - type: 'object', - }, - }, -}) - const apiSpecs: ApiSpec[] = [ { filename: 'console-openapi.json', name: 'console' }, { filename: 'web-openapi.json', name: 'web' }, @@ -201,42 +192,49 @@ const addOperationIds = (document: SwaggerDocument) => { } } -const isOpaqueContractResponse = (response: OpenApiResponse) => { - const content = response.content - if (!isObject(content)) - return false - - return Object.entries(content).some(([mediaType, media]) => { - if (!isObject(media)) - return false - - return (mediaType === 'application/json' || mediaType === 'text/event-stream') && !('schema' in media) - }) -} - -const hasOpaqueContractSuccessResponse = (operation: SwaggerOperation) => { - return Object.entries(operation.responses ?? {}).some(([status, response]) => { - return /^2\d\d$/.test(status) && isObject(response) && isOpaqueContractResponse(response) - }) -} - const normalizeOpaqueContractResponses = (document: SwaggerDocument) => { - // Some backend endpoints has no schema (e.g. external) and will trap heyapi here - // So we forge an opaque schema here + // This runs before contract filtering. Flask-RESTX often emits plain success responses + // without a body schema; give those routes an opaque output so they stay in oRPC. for (const pathItem of Object.values(document.paths ?? {})) { for (const [method, operation] of Object.entries(pathItem)) { if (!operationMethods.has(method) || !isObject(operation)) continue const swaggerOperation = operation as SwaggerOperation - if (!hasOpaqueContractSuccessResponse(swaggerOperation)) - continue + for (const [status, response] of Object.entries(swaggerOperation.responses ?? {})) { + // Ignore non-2xx or 204 or those w/o a response field + if (!/^2\d\d$/.test(status) || status === '204' || !isObject(response)) + continue - Object.values(swaggerOperation.responses ?? {}) - .filter(response => isObject(response) && isOpaqueContractResponse(response)) - .forEach((response) => { - response.content = opaqueJsonContent() - }) + const content = response.content + if (!isObject(content) || Object.keys(content).length === 0) { + // No response specification, fill a dummy opaque resp + response.content = { + 'application/json': { + schema: { + additionalProperties: true, + type: 'object', + }, + }, + } + continue + } + + for (const [mediaType, media] of Object.entries(content)) { + if (mediaType !== 'application/json' && mediaType !== 'text/event-stream') + continue + if (!isObject(media) || isObject(media.schema)) + continue + + // JSON/SSE media without a schema traps heyapi. Patch only that media entry so + // sibling binary media keeps heyapi's Blob | File inference. + // Still a dummy opaque resp + media.schema = { + additionalProperties: true, + type: 'object', + } + } + } } } } From 200f8b800f1849f1354e2931bd9cd4cbdb7866cc Mon Sep 17 00:00:00 2001 From: Pyuyi <136783609@qq.com> Date: Tue, 30 Jun 2026 18:09:05 +0800 Subject: [PATCH 11/54] fix(api): prevent plugin provider cache stampedes (#37388) Co-authored-by: VeraPyuyi <204892921+VeraPyuyi@users.noreply.github.com> --- api/core/plugin/plugin_service.py | 122 ++++++++++--- .../core/plugin/test_model_runtime_adapter.py | 26 ++- .../services/plugin/test_plugin_service.py | 166 ++++++++++++++++++ 3 files changed, 291 insertions(+), 23 deletions(-) diff --git a/api/core/plugin/plugin_service.py b/api/core/plugin/plugin_service.py index 2ab3f87db72..4b749bb4c90 100644 --- a/api/core/plugin/plugin_service.py +++ b/api/core/plugin/plugin_service.py @@ -16,10 +16,11 @@ import logging import time from collections.abc import Mapping, Sequence from mimetypes import guess_type -from typing import ClassVar +from typing import Any, ClassVar from pydantic import BaseModel, TypeAdapter, ValidationError from redis import RedisError +from redis.exceptions import LockError from sqlalchemy import delete, select, update from sqlalchemy.orm import Session from yarl import URL @@ -82,6 +83,10 @@ class PluginService: REDIS_TTL = 60 * 5 # 5 minutes PLUGIN_MODEL_PROVIDERS_REDIS_KEY_PREFIX = "plugin_model_providers:tenant_id:" PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX = "plugin_model_providers_generation:tenant_id:" + PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX = "plugin_model_providers_refresh_lock:tenant_id:" + PLUGIN_MODEL_PROVIDERS_LOCK_TTL = 30 + PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT = 2.0 + PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05 PLUGIN_INSTALL_TASK_TERMINAL_STATUSES = (PluginInstallTaskStatus.Success, PluginInstallTaskStatus.Failed) # Mirror the detail-panel endpoint query size so list reconciliation and # the visible endpoint drawer exercise the same daemon pagination path. @@ -98,6 +103,10 @@ class PluginService: def _get_plugin_model_providers_generation_cache_key(cls, tenant_id: str) -> str: return f"{cls.PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX}{tenant_id}" + @classmethod + def _get_plugin_model_providers_lock_key(cls, tenant_id: str, generation: int) -> str: + return f"{cls.PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX}{tenant_id}:generation:{generation}" + @staticmethod def _get_provider_short_name_alias(provider: PluginModelProviderEntity) -> str: """ @@ -197,32 +206,41 @@ class PluginService: cls, tenant_id: str, *, client: PluginModelClient | None = None ) -> tuple[ProviderEntity, ...] | None: generation = cls._load_plugin_model_providers_generation(tenant_id) + cached_providers, _ = cls._load_cached_plugin_model_providers_for_generation(tenant_id, generation) + return cached_providers + + @classmethod + def _load_cached_plugin_model_providers_for_generation( + cls, tenant_id: str, generation: int | None + ) -> tuple[tuple[ProviderEntity, ...] | None, bool]: if generation is not None: in_memory_cached_providers = cls._load_in_memory_plugin_model_providers(tenant_id, generation) if in_memory_cached_providers is not None: - return in_memory_cached_providers + return in_memory_cached_providers, True + + if generation is None: + return None, False cache_keys = [] - if generation is not None: - cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id, generation)) - if generation == 0: - cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id)) + cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id, generation)) + if generation == 0: + cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id)) if not cache_keys: - return None + return None, True try: cached_provider_entries = redis_client.mget(cache_keys) - except (RedisError, RuntimeError): + except (LockError, RedisError, RuntimeError): logger.warning("Failed to read cached plugin model providers for tenant %s.", tenant_id, exc_info=True) - return None + return None, False if len(cached_provider_entries) != len(cache_keys): logger.warning( "Unexpected cached plugin model providers response size for tenant %s.", tenant_id, ) - return None + return None, False for cache_key, cached_providers in zip(cache_keys, cached_provider_entries): if not cached_providers: @@ -232,7 +250,7 @@ class PluginService: providers = tuple(_provider_entities_adapter.validate_json(cached_providers)) if generation is not None: cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers) - return providers + return providers, True except (TypeError, ValueError, ValidationError): logger.warning( "Invalid cached plugin model providers for tenant %s; deleting cache key %s.", @@ -249,7 +267,7 @@ class PluginService: exc_info=True, ) - return None + return None, True @classmethod def _store_cached_plugin_model_providers( @@ -262,6 +280,49 @@ class PluginService: except (RedisError, RuntimeError): logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True) + @classmethod + def _try_acquire_plugin_model_providers_lock(cls, tenant_id: str, generation: int) -> tuple[Any | None, bool]: + lock_key = cls._get_plugin_model_providers_lock_key(tenant_id, generation) + try: + lock = redis_client.lock(lock_key, timeout=cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, blocking=False) + acquired = lock.acquire(blocking=False) + except (RedisError, RuntimeError): + logger.warning( + "Failed to acquire plugin model providers refresh lock for tenant %s.", + tenant_id, + exc_info=True, + ) + return None, False + + if not acquired: + return None, True + + return lock, True + + @classmethod + def _release_plugin_model_providers_lock(cls, tenant_id: str, lock: Any) -> None: + try: + lock.release() + except (LockError, RedisError, RuntimeError): + logger.warning( + "Failed to release plugin model providers refresh lock for tenant %s.", + tenant_id, + exc_info=True, + ) + + @classmethod + def _wait_for_plugin_model_providers_refresh( + cls, tenant_id: str, *, client: PluginModelClient | None = None + ) -> tuple[ProviderEntity, ...] | None: + deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT + while time.monotonic() < deadline: + time.sleep(cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL) + cached_providers = cls._load_cached_plugin_model_providers(tenant_id, client=client) + if cached_providers is not None: + return cached_providers + + return None + @classmethod def invalidate_plugin_model_providers_cache(cls, tenant_id: str) -> None: """Invalidate tenant-scoped provider metadata across Redis and worker-local mirrors.""" @@ -287,21 +348,38 @@ class PluginService: are intentionally owned by this service so tenant isolation and cache expiry are handled in one place. """ - cached_providers = cls._load_cached_plugin_model_providers(tenant_id, client=client) + generation = cls._load_plugin_model_providers_generation(tenant_id) + cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation( + tenant_id, generation + ) if cached_providers is not None: return cached_providers + refresh_lock: Any | None = None + refresh_generation = generation + if generation is not None and cache_available: + lock_wait_deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL + while time.monotonic() < lock_wait_deadline: + refresh_lock, lock_available = cls._try_acquire_plugin_model_providers_lock(tenant_id, generation) + if refresh_lock is not None or not lock_available: + break + refreshed_providers = cls._wait_for_plugin_model_providers_refresh(tenant_id, client=client) + if refreshed_providers is not None: + return refreshed_providers + model_client = client or PluginModelClient() - providers = tuple( - cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id) - ) - if not providers: + try: + providers = tuple( + cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id) + ) + generation = cls._load_plugin_model_providers_generation(tenant_id) + if generation is not None and generation == refresh_generation: + cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers) + cls._store_cached_plugin_model_providers(tenant_id, generation, providers) return providers - generation = cls._load_plugin_model_providers_generation(tenant_id) - if generation is not None: - cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers) - cls._store_cached_plugin_model_providers(tenant_id, generation, providers) - return providers + finally: + if refresh_lock is not None: + cls._release_plugin_model_providers_lock(tenant_id, refresh_lock) @staticmethod def fetch_latest_plugin_version(plugin_ids: Sequence[str]) -> Mapping[str, LatestPluginCache | None]: diff --git a/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py b/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py index 17973916779..c3ee4227d25 100644 --- a/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py +++ b/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py @@ -44,6 +44,29 @@ class _FakeRedis: def delete(self, key: str) -> None: self._values.pop(key, None) + def lock(self, key: str, *, timeout: int, blocking: bool) -> "_FakeRedisLock": + return _FakeRedisLock(self, key) + + +class _FakeRedisLock: + def __init__(self, redis: _FakeRedis, key: str) -> None: + self._redis = redis + self._key = key + self._acquired = False + + def acquire(self, *, blocking: bool) -> bool: + if self._key in self._redis._values: + return False + + self._redis._values[self._key] = "locked" + self._acquired = True + return True + + def release(self) -> None: + if self._acquired: + self._redis.delete(self._key) + self._acquired = False + @pytest.fixture(autouse=True) def clear_plugin_model_provider_memory_cache() -> None: @@ -416,9 +439,10 @@ class TestPluginModelRuntime: mget=Mock(return_value=[None, None]), delete=Mock(), setex=Mock(), + lock=Mock(return_value=SimpleNamespace(acquire=Mock(return_value=True), release=Mock())), ), ) - monkeypatch.setattr(plugin_service_module.dify_config, "PLUGIN_MODEL_PROVIDERS_CACHE_TTL", 300) + monkeypatch.setattr(plugin_service_module.dify_config, "PLUGIN_MODEL_PROVIDERS_CACHE_TTL", 0) runtime = PluginModelRuntime(tenant_id="tenant", user_id="user", client=client, plugin_service=PluginService) runtime.fetch_model_providers() diff --git a/api/tests/unit_tests/services/plugin/test_plugin_service.py b/api/tests/unit_tests/services/plugin/test_plugin_service.py index fca05f94fc7..c7bd4dff08e 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service.py @@ -235,6 +235,172 @@ class TestPluginModelProviderCache: client.fetch_model_providers.assert_called_once_with("tenant-1") assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + def test_fetch_plugin_model_providers_waits_for_concurrent_refresh_cache_fill(self) -> None: + """A cache miss waits for the active tenant refresh instead of stampeding the daemon.""" + cached_provider = _build_provider_entity() + cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]).decode("utf-8") + cache_key = _provider_cache_key("tenant-1", 0) + legacy_cache_key = _provider_cache_key("tenant-1") + + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.sleep") as sleep, + ): + redis_client.get.return_value = None + redis_client.mget.side_effect = [[None, None], [cached_payload, None]] + redis_client.lock.return_value.acquire.return_value = False + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider(provider="anthropic")] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + redis_client.lock.assert_called_once_with( + PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), + timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + blocking=False, + ) + redis_client.lock.return_value.acquire.assert_called_once_with(blocking=False) + assert redis_client.mget.call_args_list == [ + call([cache_key, legacy_cache_key]), + call([cache_key, legacy_cache_key]), + ] + sleep.assert_called() + client.fetch_model_providers.assert_not_called() + assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + + def test_fetch_plugin_model_providers_retries_lock_after_wait_timeout(self) -> None: + """Only a lock owner should refresh the daemon when the first refresh takes too long.""" + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.sleep"), + patch(f"{MODULE}.PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT", 0), + ): + redis_client.get.return_value = None + redis_client.mget.return_value = [None, None] + redis_client.lock.return_value.acquire.side_effect = [False, True] + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider()] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert redis_client.lock.return_value.acquire.call_args_list == [ + call(blocking=False), + call(blocking=False), + ] + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.lock.return_value.release.assert_called_once_with() + assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + + def test_fetch_plugin_model_providers_releases_owned_refresh_lock_after_store(self) -> None: + """The refresh owner releases only its token after storing provider metadata.""" + cache_key = _provider_cache_key("tenant-1", 0) + legacy_cache_key = _provider_cache_key("tenant-1") + + with patch(f"{MODULE}.redis_client") as redis_client: + redis_client.get.return_value = None + redis_client.mget.return_value = [None, None] + redis_client.lock.return_value.acquire.return_value = True + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider()] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + redis_client.lock.assert_called_once_with( + PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), + timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + blocking=False, + ) + redis_client.lock.return_value.acquire.assert_called_once_with(blocking=False) + redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key]) + redis_client.lock.return_value.release.assert_called_once_with() + redis_client.eval.assert_not_called() + client.fetch_model_providers.assert_called_once_with("tenant-1") + assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + + def test_fetch_plugin_model_providers_skips_wait_when_refresh_lock_fails(self) -> None: + """Lock API failures should fall back directly instead of adding timeout latency.""" + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.sleep") as sleep, + ): + redis_client.get.return_value = None + redis_client.mget.return_value = [None, None] + redis_client.lock.side_effect = RedisError("redis unavailable") + redis_client.set.side_effect = AssertionError("raw redis set must not be used for refresh locks") + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider()] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + sleep.assert_not_called() + redis_client.lock.assert_called_once() + client.fetch_model_providers.assert_called_once_with("tenant-1") + assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + + def test_fetch_plugin_model_providers_caches_empty_provider_list(self) -> None: + """An empty provider list is still a valid refresh result for single-flight waiters.""" + cache_key = _provider_cache_key("tenant-1", 0) + with patch(f"{MODULE}.redis_client") as redis_client: + redis_client.get.return_value = None + redis_client.mget.return_value = [None, None] + redis_client.lock.return_value.acquire.return_value = True + client = Mock() + client.fetch_model_providers.return_value = [] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert result == () + redis_client.setex.assert_called_once() + assert redis_client.setex.call_args.args[0] == cache_key + redis_client.lock.return_value.release.assert_called_once_with() + + def test_fetch_plugin_model_providers_skips_cache_write_when_generation_changes_during_refresh(self) -> None: + """A refresh that started before invalidation must not populate the newer generation cache.""" + with patch(f"{MODULE}.redis_client") as redis_client: + redis_client.get.side_effect = [None, "1"] + redis_client.mget.return_value = [None, None] + redis_client.lock.return_value.acquire.return_value = True + client = Mock() + client.fetch_model_providers.return_value = [] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert result == () + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.setex.assert_not_called() + redis_client.lock.return_value.release.assert_called_once_with() + + def test_fetch_plugin_model_providers_reuses_cached_empty_provider_list(self) -> None: + """A cached empty list should prevent repeated daemon fetches for tenants without plugin models.""" + empty_payload = TypeAdapter(list[ProviderEntity]).dump_json([]).decode("utf-8") + cache_key = _provider_cache_key("tenant-1", 0) + legacy_cache_key = _provider_cache_key("tenant-1") + + with patch(f"{MODULE}.redis_client") as redis_client: + redis_client.get.return_value = None + redis_client.mget.return_value = [empty_payload, None] + client = Mock() + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert result == () + redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key]) + client.fetch_model_providers.assert_not_called() + def test_fetch_plugin_model_providers_creates_default_client_on_cache_miss(self) -> None: """The service owns plugin daemon access when no runtime-provided client is injected.""" with ( From b1bb6ef977fb0e5cbc49c76582904712d2adf6a5 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Tue, 30 Jun 2026 18:36:38 +0800 Subject: [PATCH 12/54] perf(web): improve vinext home startup time (#38219) --- web/app/(commonLayout)/global-mounts.tsx | 21 +++++ web/app/(commonLayout)/layout.tsx | 12 +-- .../main-nav/__tests__/index.spec.tsx | 86 ++++++++++--------- .../components/snippet-detail-section.tsx | 35 ++++++++ web/app/components/main-nav/index.tsx | 41 +++------ web/i18n-config/client.ts | 7 +- web/i18n-config/load-resource.ts | 35 ++++++++ web/i18n-config/locale-resources/ar-TN.ts | 1 + web/i18n-config/locale-resources/de-DE.ts | 1 + web/i18n-config/locale-resources/en-US.ts | 1 + web/i18n-config/locale-resources/es-ES.ts | 1 + web/i18n-config/locale-resources/fa-IR.ts | 1 + web/i18n-config/locale-resources/fr-FR.ts | 1 + web/i18n-config/locale-resources/hi-IN.ts | 1 + web/i18n-config/locale-resources/id-ID.ts | 1 + web/i18n-config/locale-resources/it-IT.ts | 1 + web/i18n-config/locale-resources/ja-JP.ts | 1 + web/i18n-config/locale-resources/ko-KR.ts | 1 + web/i18n-config/locale-resources/nl-NL.ts | 1 + web/i18n-config/locale-resources/pl-PL.ts | 1 + web/i18n-config/locale-resources/pt-BR.ts | 1 + web/i18n-config/locale-resources/ro-RO.ts | 1 + web/i18n-config/locale-resources/ru-RU.ts | 1 + web/i18n-config/locale-resources/sl-SI.ts | 1 + web/i18n-config/locale-resources/th-TH.ts | 1 + web/i18n-config/locale-resources/tr-TR.ts | 1 + web/i18n-config/locale-resources/uk-UA.ts | 1 + web/i18n-config/locale-resources/vi-VN.ts | 1 + web/i18n-config/locale-resources/zh-Hans.ts | 1 + web/i18n-config/locale-resources/zh-Hant.ts | 1 + web/i18n-config/server.ts | 9 +- web/knip.config.ts | 1 + 32 files changed, 178 insertions(+), 92 deletions(-) create mode 100644 web/app/(commonLayout)/global-mounts.tsx create mode 100644 web/app/components/main-nav/components/snippet-detail-section.tsx create mode 100644 web/i18n-config/load-resource.ts create mode 100644 web/i18n-config/locale-resources/ar-TN.ts create mode 100644 web/i18n-config/locale-resources/de-DE.ts create mode 100644 web/i18n-config/locale-resources/en-US.ts create mode 100644 web/i18n-config/locale-resources/es-ES.ts create mode 100644 web/i18n-config/locale-resources/fa-IR.ts create mode 100644 web/i18n-config/locale-resources/fr-FR.ts create mode 100644 web/i18n-config/locale-resources/hi-IN.ts create mode 100644 web/i18n-config/locale-resources/id-ID.ts create mode 100644 web/i18n-config/locale-resources/it-IT.ts create mode 100644 web/i18n-config/locale-resources/ja-JP.ts create mode 100644 web/i18n-config/locale-resources/ko-KR.ts create mode 100644 web/i18n-config/locale-resources/nl-NL.ts create mode 100644 web/i18n-config/locale-resources/pl-PL.ts create mode 100644 web/i18n-config/locale-resources/pt-BR.ts create mode 100644 web/i18n-config/locale-resources/ro-RO.ts create mode 100644 web/i18n-config/locale-resources/ru-RU.ts create mode 100644 web/i18n-config/locale-resources/sl-SI.ts create mode 100644 web/i18n-config/locale-resources/th-TH.ts create mode 100644 web/i18n-config/locale-resources/tr-TR.ts create mode 100644 web/i18n-config/locale-resources/uk-UA.ts create mode 100644 web/i18n-config/locale-resources/vi-VN.ts create mode 100644 web/i18n-config/locale-resources/zh-Hans.ts create mode 100644 web/i18n-config/locale-resources/zh-Hant.ts diff --git a/web/app/(commonLayout)/global-mounts.tsx b/web/app/(commonLayout)/global-mounts.tsx new file mode 100644 index 00000000000..9da151d84f8 --- /dev/null +++ b/web/app/(commonLayout)/global-mounts.tsx @@ -0,0 +1,21 @@ +'use client' + +import dynamic from '@/next/dynamic' + +const InSiteMessageNotification = dynamic(() => import('@/app/components/app/in-site-message/notification'), { ssr: false }) +const PartnerStack = dynamic(() => import('@/app/components/billing/partner-stack'), { ssr: false }) +const ReadmePanel = dynamic(() => import('@/app/components/plugins/readme-panel'), { ssr: false }) +const WorkflowGeneratorMount = dynamic(() => import('@/app/components/workflow/workflow-generator/mount'), { ssr: false }) +const GotoAnything = dynamic(() => import('@/app/components/goto-anything').then(mod => mod.GotoAnything), { ssr: false }) + +export function CommonLayoutGlobalMounts() { + return ( + <> + + + + + + + ) +} diff --git a/web/app/(commonLayout)/layout.tsx b/web/app/(commonLayout)/layout.tsx index 69a72bdb3ee..e2d29f85d30 100644 --- a/web/app/(commonLayout)/layout.tsx +++ b/web/app/(commonLayout)/layout.tsx @@ -1,20 +1,16 @@ import type { ReactNode } from 'react' -import InSiteMessageNotification from '@/app/components/app/in-site-message/notification' import AmplitudeProvider from '@/app/components/base/amplitude' import { GoogleAnalyticsScripts } from '@/app/components/base/ga' import Zendesk from '@/app/components/base/zendesk' import { EducationVerifyActionRecorder } from '@/app/components/education-verify-action-recorder' -import { GotoAnything } from '@/app/components/goto-anything' import MainNavLayout from '@/app/components/main-nav/layout' import { NextRouteStateBridge } from '@/app/components/next-route-state' import { OAuthRegistrationAnalytics } from '@/app/components/oauth-registration-analytics' -import ReadmePanel from '@/app/components/plugins/readme-panel' -import WorkflowGeneratorMount from '@/app/components/workflow/workflow-generator/mount' import { AppContextProvider } from '@/context/app-context-provider' import { EventEmitterContextProvider } from '@/context/event-emitter-provider' import { ModalContextProvider } from '@/context/modal-context-provider' import { ProviderContextProvider } from '@/context/provider-context-provider' -import PartnerStack from '../components/billing/partner-stack' +import { CommonLayoutGlobalMounts } from './global-mounts' import { CommonLayoutHydrationBoundary } from './hydration-boundary' export default async function Layout({ children }: { children: ReactNode }) { @@ -33,11 +29,7 @@ export default async function Layout({ children }: { children: ReactNode }) { {children} - - - - - + diff --git a/web/app/components/main-nav/__tests__/index.spec.tsx b/web/app/components/main-nav/__tests__/index.spec.tsx index 8846220c21c..d1722533b9a 100644 --- a/web/app/components/main-nav/__tests__/index.spec.tsx +++ b/web/app/components/main-nav/__tests__/index.spec.tsx @@ -495,7 +495,7 @@ describe('MainNav', () => { expect(screen.queryByRole('link', { name: /common.menus.deployments/ })).not.toBeInTheDocument() }) - it('aligns the global navigation spacing with the main sidebar design', () => { + it('aligns the global navigation spacing with the main sidebar design', async () => { mockInstalledApps = [createInstalledApp()] renderMainNav() @@ -508,7 +508,7 @@ describe('MainNav', () => { expect(homeLink.closest('nav')).toHaveClass('isolate', 'flex', 'flex-col', 'gap-px', 'p-2') expect(homeLink).toHaveClass('h-8', 'w-full', 'rounded-[10px]', 'px-2', 'py-1.5') - const webAppsButton = screen.getByRole('button', { name: 'explore.sidebar.webApps' }) + const webAppsButton = await screen.findByRole('button', { name: 'explore.sidebar.webApps' }) expect(webAppsButton.parentElement).toHaveClass('py-1', 'pr-2', 'pl-2') const helpButton = screen.getByRole('button', { name: 'common.mainNav.help.openMenu' }) @@ -551,7 +551,7 @@ describe('MainNav', () => { expect(container.querySelector('.relative.z-30')).not.toBeInTheDocument() }) - it('hides the environment tag when app detail navigation is collapsed', () => { + it('hides the environment tag when app detail navigation is collapsed', async () => { mockPathname = '/app/app-1/overview' ;(useAppContext as Mock).mockReturnValue({ ...appContextValue, @@ -562,7 +562,7 @@ describe('MainNav', () => { }) const { container } = renderMainNav() - fireEvent.click(screen.getByTestId('app-detail-toggle')) + fireEvent.click(await screen.findByTestId('app-detail-toggle')) expect(screen.queryByText('common.environment.testing')).not.toBeInTheDocument() expect(container.querySelector('.relative.z-30')).not.toBeInTheDocument() @@ -659,7 +659,7 @@ describe('MainNav', () => { expect(screen.getByRole('link', { name: /common.mainNav.home/ })).not.toHaveAttribute('aria-current') }) - it('replaces global navigation with snippet detail navigation on snippet routes', () => { + it('replaces global navigation with snippet detail navigation on snippet routes', async () => { mockPathname = '/snippets/snippet-1/orchestrate' snippetDraftState.inputFields = snippetFields snippetNavigationState.onFieldsChange = mockSnippetFieldsChange @@ -671,8 +671,8 @@ describe('MainNav', () => { expect(screen.getByRole('complementary')).toHaveClass('w-62') expect(screen.getByRole('complementary')).toHaveClass('p-1') expect(screen.getByRole('complementary')).toHaveClass('bg-background-body') - expect(screen.getByTestId('snippet-detail-top')).toHaveAttribute('data-expand', 'true') - expect(screen.getByTestId('snippet-sidebar-content')).toHaveAttribute('data-readonly', 'false') + expect(await screen.findByTestId('snippet-detail-top')).toHaveAttribute('data-expand', 'true') + expect(await screen.findByTestId('snippet-sidebar-content')).toHaveAttribute('data-readonly', 'false') expect(screen.getByText(snippet.name)).toBeInTheDocument() expect(screen.getByText('query')).toBeInTheDocument() fireEvent.click(screen.getByRole('button', { name: 'change snippet fields' })) @@ -686,14 +686,14 @@ describe('MainNav', () => { expect(screen.getByRole('button', { name: 'common.mainNav.help.openMenu' })).toBeInTheDocument() }) - it('collapses snippet detail navigation from the top-right toggle', () => { + it('collapses snippet detail navigation from the top-right toggle', async () => { mockPathname = '/snippets/snippet-1/orchestrate' snippetDraftState.inputFields = snippetFields snippetNavigationState.onFieldsChange = mockSnippetFieldsChange snippetNavigationState.snippet = snippet renderMainNav() - fireEvent.click(screen.getByTestId('snippet-detail-toggle')) + fireEvent.click(await screen.findByTestId('snippet-detail-toggle')) expect(screen.getByRole('complementary')).toHaveClass('w-16') expect(screen.getByRole('complementary')).toHaveClass('p-1') @@ -704,13 +704,13 @@ describe('MainNav', () => { expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') }) - it('replaces global navigation with app detail navigation on app routes', () => { + it('replaces global navigation with app detail navigation on app routes', async () => { mockPathname = '/app/app-1/overview' renderMainNav() - expect(screen.getByTestId('app-detail-top')).toBeInTheDocument() - expect(screen.getByTestId('app-detail-section')).toBeInTheDocument() + expect(await screen.findByTestId('app-detail-top')).toBeInTheDocument() + expect(await screen.findByTestId('app-detail-section')).toBeInTheDocument() expect(screen.getByTestId('app-detail-top')).toHaveAttribute('data-expand', 'true') expect(screen.getByTestId('app-detail-section')).toHaveAttribute('data-expand', 'true') expect(screen.getByRole('complementary')).toHaveClass('w-62') @@ -742,56 +742,56 @@ describe('MainNav', () => { }) }) - it('collapses app detail navigation from the top-right toggle', () => { + it('collapses app detail navigation from the top-right toggle', async () => { mockPathname = '/app/app-1/overview' renderMainNav() - fireEvent.click(screen.getByTestId('app-detail-toggle')) + fireEvent.click(await screen.findByTestId('app-detail-toggle')) expect(screen.getByRole('complementary')).toHaveClass('w-16') expect(screen.getByRole('complementary')).not.toHaveClass('transition-none') expect(screen.getByRole('complementary')).toHaveClass('p-1') expect(screen.getByTestId('app-detail-top')).toHaveAttribute('data-expand', 'false') - expect(screen.getByTestId('app-detail-section')).toHaveAttribute('data-expand', 'false') + expect(await screen.findByTestId('app-detail-section')).toHaveAttribute('data-expand', 'false') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') }) - it('shows app detail navigation as a floating preview when hovering the collapsed top toggle', () => { + it('shows app detail navigation as a floating preview when hovering the collapsed top toggle', async () => { mockPathname = '/app/app-1/overview' renderMainNav() - fireEvent.click(screen.getByTestId('app-detail-toggle')) + fireEvent.click(await screen.findByTestId('app-detail-toggle')) fireEvent.mouseEnter(screen.getByTestId('app-detail-top').parentElement!) expect(screen.getByRole('complementary')).toHaveClass('w-16', 'overflow-visible') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') expect(screen.getAllByTestId('app-detail-top')).toHaveLength(1) expect(screen.getByTestId('app-detail-top')).toHaveAttribute('data-expand', 'true') - expect(screen.getByTestId('app-detail-section')).toHaveAttribute('data-expand', 'true') + expect(await screen.findByTestId('app-detail-section')).toHaveAttribute('data-expand', 'true') }) - it('persists expanded app detail navigation without width animation when clicking the hovered toggle', () => { + it('persists expanded app detail navigation without width animation when clicking the hovered toggle', async () => { mockPathname = '/app/app-1/overview' renderMainNav() - fireEvent.click(screen.getByTestId('app-detail-toggle')) + fireEvent.click(await screen.findByTestId('app-detail-toggle')) fireEvent.mouseEnter(screen.getByTestId('app-detail-top').parentElement!) fireEvent.click(screen.getByTestId('app-detail-toggle')) expect(screen.getByRole('complementary')).toHaveClass('w-62', 'transition-none') expect(screen.getByRole('complementary')).not.toHaveClass('overflow-visible') expect(screen.getByTestId('app-detail-top')).toHaveAttribute('data-expand', 'true') - expect(screen.getByTestId('app-detail-section')).toHaveAttribute('data-expand', 'true') + expect(await screen.findByTestId('app-detail-section')).toHaveAttribute('data-expand', 'true') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('expand') }) - it('replaces global navigation with dataset detail navigation on dataset routes', () => { + it('replaces global navigation with dataset detail navigation on dataset routes', async () => { mockPathname = '/datasets/dataset-1/documents' renderMainNav() - expect(screen.getByTestId('dataset-detail-top')).toBeInTheDocument() - expect(screen.getByTestId('dataset-detail-section')).toBeInTheDocument() + expect(await screen.findByTestId('dataset-detail-top')).toBeInTheDocument() + expect(await screen.findByTestId('dataset-detail-section')).toBeInTheDocument() expect(screen.getByTestId('dataset-detail-top')).toHaveAttribute('data-expand', 'true') expect(screen.getByTestId('dataset-detail-section')).toHaveAttribute('data-expand', 'true') expect(screen.getByRole('complementary')).toHaveClass('w-62') @@ -802,42 +802,44 @@ describe('MainNav', () => { expect(screen.queryByRole('link', { name: /common.menus.datasets/ })).not.toBeInTheDocument() }) - it('collapses dataset detail navigation from the top-right toggle', () => { + it('collapses dataset detail navigation from the top-right toggle', async () => { mockPathname = '/datasets/dataset-1/documents' renderMainNav() - fireEvent.click(screen.getByTestId('dataset-detail-toggle')) + fireEvent.click(await screen.findByTestId('dataset-detail-toggle')) expect(screen.getByRole('complementary')).toHaveClass('w-16') expect(screen.getByRole('complementary')).toHaveClass('p-1') expect(screen.getByTestId('dataset-detail-top')).toHaveAttribute('data-expand', 'false') - expect(screen.getByTestId('dataset-detail-section')).toHaveAttribute('data-expand', 'false') + expect(await screen.findByTestId('dataset-detail-section')).toHaveAttribute('data-expand', 'false') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') }) - it('shows dataset detail navigation as a floating preview when hovering the collapsed top toggle', () => { + it('shows dataset detail navigation as a floating preview when hovering the collapsed top toggle', async () => { mockPathname = '/datasets/dataset-1/documents' renderMainNav() - fireEvent.click(screen.getByTestId('dataset-detail-toggle')) + fireEvent.click(await screen.findByTestId('dataset-detail-toggle')) fireEvent.mouseEnter(screen.getByTestId('dataset-detail-top').parentElement!) expect(screen.getByRole('complementary')).toHaveClass('w-16', 'overflow-visible') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') expect(screen.getAllByTestId('dataset-detail-top')).toHaveLength(1) expect(screen.getByTestId('dataset-detail-top')).toHaveAttribute('data-expand', 'true') - expect(screen.getByTestId('dataset-detail-section')).toHaveAttribute('data-expand', 'true') + expect(await screen.findByTestId('dataset-detail-section')).toHaveAttribute('data-expand', 'true') }) - it('replaces global navigation with agent detail navigation on roster detail routes', () => { + it('replaces global navigation with agent detail navigation on roster detail routes', async () => { mockPathname = '/roster/agent/agent-1/configure' renderMainNav() - expect(screen.getByTestId('agent-detail-top')).toBeInTheDocument() - expect(screen.getByTestId('agent-detail-section')).toBeInTheDocument() + expect(await screen.findByTestId('agent-detail-top')).toBeInTheDocument() + const agentDetailNavigation = await screen.findByRole('navigation', { name: 'agentV2.agentDetail.navigationLabel' }) + expect(agentDetailNavigation).toBeInTheDocument() + expect(screen.getByRole('link', { name: /agentV2.agentDetail.sections.configure/ })).toHaveAttribute('href', '/roster/agent/agent-1/configure') expect(screen.getByTestId('agent-detail-top')).toHaveAttribute('data-expand', 'true') - expect(screen.getByTestId('agent-detail-section')).toHaveAttribute('data-expand', 'true') + expect(agentDetailNavigation).toHaveClass('px-1') expect(screen.getByRole('complementary')).toHaveClass('w-62') expect(screen.getByRole('complementary')).toHaveClass('p-1') expect(screen.getByRole('complementary')).toHaveClass('bg-background-body') @@ -873,16 +875,16 @@ describe('MainNav', () => { expect(screen.queryByRole('link', { name: /common.menus.deployments/ })).not.toBeInTheDocument() }) - it('collapses agent detail navigation from the top-right toggle', () => { + it('collapses agent detail navigation from the top-right toggle', async () => { mockPathname = '/roster/agent/agent-1/configure' renderMainNav() - fireEvent.click(screen.getByTestId('agent-detail-toggle')) + fireEvent.click(await screen.findByTestId('agent-detail-toggle')) expect(screen.getByRole('complementary')).toHaveClass('w-16') expect(screen.getByRole('complementary')).toHaveClass('p-1') expect(screen.getByTestId('agent-detail-top')).toHaveAttribute('data-expand', 'false') - expect(screen.getByTestId('agent-detail-section')).toHaveAttribute('data-expand', 'false') + expect(await screen.findByRole('navigation', { name: 'agentV2.agentDetail.navigationLabel' })).toHaveClass('px-3') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') }) @@ -923,18 +925,18 @@ describe('MainNav', () => { ) }) - it('shows agent detail navigation as a floating preview when hovering the collapsed top toggle', () => { + it('shows agent detail navigation as a floating preview when hovering the collapsed top toggle', async () => { mockPathname = '/roster/agent/agent-1/configure' renderMainNav() - fireEvent.click(screen.getByTestId('agent-detail-toggle')) + fireEvent.click(await screen.findByTestId('agent-detail-toggle')) fireEvent.mouseEnter(screen.getByTestId('agent-detail-top').parentElement!) expect(screen.getByRole('complementary')).toHaveClass('w-16', 'overflow-visible') expect(localStorage.getItem(DETAIL_SIDEBAR_STORAGE_KEY)).toBe('collapse') expect(screen.getAllByTestId('agent-detail-top')).toHaveLength(1) expect(screen.getByTestId('agent-detail-top')).toHaveAttribute('data-expand', 'true') - expect(screen.getByTestId('agent-detail-section')).toHaveAttribute('data-expand', 'true') + expect(await screen.findByRole('navigation', { name: 'agentV2.agentDetail.navigationLabel' })).toHaveClass('px-1') }) it.each([ @@ -1258,12 +1260,12 @@ describe('MainNav', () => { } }) - it('collapses and expands installed web apps from the section arrow', () => { + it('collapses and expands installed web apps from the section arrow', async () => { mockInstalledApps = [createInstalledApp()] renderMainNav() - const webAppsButton = screen.getByRole('button', { name: 'explore.sidebar.webApps' }) + const webAppsButton = await screen.findByRole('button', { name: 'explore.sidebar.webApps' }) expect(webAppsButton).toHaveAttribute('aria-expanded', 'true') expect(screen.getByText('Alpha App')).toBeInTheDocument() diff --git a/web/app/components/main-nav/components/snippet-detail-section.tsx b/web/app/components/main-nav/components/snippet-detail-section.tsx new file mode 100644 index 00000000000..e9798a903d5 --- /dev/null +++ b/web/app/components/main-nav/components/snippet-detail-section.tsx @@ -0,0 +1,35 @@ +'use client' + +import { useShallow } from 'zustand/react/shallow' +import { SnippetCollapsedPreview } from '@/app/components/snippets/components/snippet-collapsed-preview' +import { SnippetSidebarContent } from '@/app/components/snippets/components/snippet-sidebar' +import { useSnippetDraftStore } from '@/app/components/snippets/draft-store' +import { useSnippetDetailStore } from '@/app/components/snippets/store' + +type SnippetDetailSectionProps = { + expand: boolean +} + +export function SnippetDetailSection({ expand }: SnippetDetailSectionProps) { + const snippetNavigation = useSnippetDetailStore(useShallow(state => ({ + onFieldsChange: state.onFieldsChange, + readonly: state.readonly, + snippet: state.snippet, + }))) + const snippetInputFields = useSnippetDraftStore(state => state.inputFields) + + if (!expand) + return + + if (!snippetNavigation.snippet || !snippetNavigation.onFieldsChange) + return null + + return ( + + ) +} diff --git a/web/app/components/main-nav/index.tsx b/web/app/components/main-nav/index.tsx index 9caa2900bb2..19e0872b7b1 100644 --- a/web/app/components/main-nav/index.tsx +++ b/web/app/components/main-nav/index.tsx @@ -7,31 +7,21 @@ import { useSuspenseQuery } from '@tanstack/react-query' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useShallow } from 'zustand/react/shallow' -import AppDetailSection from '@/app/components/app-sidebar/app-detail-section' -import AppDetailTop from '@/app/components/app-sidebar/app-detail-top' -import DatasetDetailSection from '@/app/components/app-sidebar/dataset-detail-section' -import DatasetDetailTop from '@/app/components/app-sidebar/dataset-detail-top' import { useStore as useAppStore } from '@/app/components/app/store' import Badge from '@/app/components/base/badge' import DifyLogo from '@/app/components/base/logo/dify-logo' import EnvNav from '@/app/components/header/env-nav' -import { SnippetCollapsedPreview } from '@/app/components/snippets/components/snippet-collapsed-preview' -import { SnippetSidebarContent } from '@/app/components/snippets/components/snippet-sidebar' -import { useSnippetDraftStore } from '@/app/components/snippets/draft-store' -import { useSnippetDetailStore } from '@/app/components/snippets/store' import { useAppContext } from '@/context/app-context' -import { AgentDetailSection, AgentDetailTop } from '@/features/agent-v2/agent-detail/navigation' import { isAgentV2Enabled } from '@/features/agent-v2/feature-flag' import { DeploymentDetailSection, DeploymentDetailTop } from '@/features/deployments/detail/deployment-sidebar' import { systemFeaturesQueryOptions } from '@/features/system-features/client' +import dynamic from '@/next/dynamic' import Link from '@/next/link' import { usePathname } from '@/next/navigation' import AccountSection from './components/account-section' import HelpMenu from './components/help-menu' import MainNavLink from './components/nav-link' import { MainNavSearchButton } from './components/search-button' -import SnippetDetailTop from './components/snippet-detail-top' -import WebAppsSection from './components/web-apps-section' import { WorkspaceCard } from './components/workspace-card' import { isMainNavRouteVisible, MAIN_NAV_ROUTES } from './routes' import { useDetailSidebarMode } from './storage' @@ -41,6 +31,16 @@ const DATASET_DOCUMENT_CREATION_ROUTES = new Set(['create', 'create-from-pipelin const DEPLOYMENT_COLLECTION_ROUTES = new Set(['create']) const secondarySidebarHelpTriggerIcon = +const AppDetailSection = dynamic(() => import('@/app/components/app-sidebar/app-detail-section'), { ssr: false }) +const AppDetailTop = dynamic(() => import('@/app/components/app-sidebar/app-detail-top'), { ssr: false }) +const DatasetDetailSection = dynamic(() => import('@/app/components/app-sidebar/dataset-detail-section'), { ssr: false }) +const DatasetDetailTop = dynamic(() => import('@/app/components/app-sidebar/dataset-detail-top'), { ssr: false }) +const AgentDetailSection = dynamic(() => import('@/features/agent-v2/agent-detail/navigation').then(mod => mod.AgentDetailSection), { ssr: false }) +const AgentDetailTop = dynamic(() => import('@/features/agent-v2/agent-detail/navigation').then(mod => mod.AgentDetailTop), { ssr: false }) +const SnippetDetailTop = dynamic(() => import('./components/snippet-detail-top'), { ssr: false }) +const SnippetDetailSection = dynamic(() => import('./components/snippet-detail-section').then(mod => mod.SnippetDetailSection), { ssr: false }) +const WebAppsSection = dynamic(() => import('./components/web-apps-section'), { ssr: false }) + function SecondarySidebarHelpMenu({ triggerClassName, }: { @@ -103,12 +103,6 @@ export function MainNav({ const showDeploymentDetailNavigation = canUseAppDeploy && !isCurrentWorkspaceDatasetOperator && isDeploymentDetailPathname(pathname) const showSnippetDetailNavigation = isSnippetDetailPathname(pathname) const showDetailNavigation = showAppDetailNavigation || showDatasetDetailNavigation || showAgentDetailNavigation || showDeploymentDetailNavigation || showSnippetDetailNavigation - const snippetNavigation = useSnippetDetailStore(useShallow(state => ({ - onFieldsChange: state.onFieldsChange, - readonly: state.readonly, - snippet: state.snippet, - }))) - const snippetInputFields = useSnippetDraftStore(state => state.inputFields) const { hasAppDetail, setAppDetail } = useAppStore(useShallow(state => ({ hasAppDetail: !!state.appDetail, setAppDetail: state.setAppDetail, @@ -307,18 +301,7 @@ export function MainNav({ ? : showDeploymentDetailNavigation ? - : detailNavigationVisibleExpanded - ? snippetNavigation.snippet && snippetNavigation.onFieldsChange - ? ( - - ) - : null - : + : : ( <>