Compare commits

..
Author SHA1 Message Date
yyh 5d790d634d Merge remote-tracking branch 'origin/main' into deploy/saas 2026-07-27 15:16:43 +08:00
林玮 (Jade Lin) 2b343f25d4 Merge remote-tracking branch 'origin/linw/snp-581-celery-task-should-join-thread' into deploy/saas 2026-07-27 14:59:40 +08:00
林玮 (Jade Lin) 45317a11b9 fix: bound app worker cleanup wait 2026-07-27 14:54:04 +08:00
林玮 (Jade Lin) 0e6af39553 Merge remote-tracking branch 'origin/linw/snp-581-celery-task-should-join-thread' into deploy/saas 2026-07-27 13:57:41 +08:00
yyh b81d521e0b perf(web): prefetch workspace list on intent 2026-07-27 12:47:19 +08:00
yyh c487dc2d7c fix(web): refine workspace switcher loading state 2026-07-27 12:27:12 +08:00
autofix-ci[bot]andyyh 0ed9fb4db6 [autofix.ci] apply automated fixes 2026-07-27 11:56:34 +08:00
yyh 52b825a4a1 fix: align workspace card plan ownership 2026-07-27 11:56:33 +08:00
yyhyyhautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
ffb1e29527 refactor(workflow): remove hook barrel exports (#39588)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-27 11:42:40 +08:00
林玮 (Jade Lin) bedb3c571b fix: correct workflow task return type 2026-07-27 11:40:34 +08:00
林玮 (Jade Lin) 05df469e28 test: join workflow resume thread stub 2026-07-27 11:35:56 +08:00
林玮 (Jade Lin) 07ab81c9f0 test: join workflow generation thread stub 2026-07-27 11:35:23 +08:00
林玮 (Jade Lin) 1319f485a1 Merge remote-tracking branch 'origin/linw/snp-581-celery-task-should-join-thread' into deploy/saas 2026-07-27 11:09:17 +08:00
林玮 (Jade Lin) b12ad644ee Merge remote-tracking branch 'origin/linw/snp-560-webapp-site-response-with-mode' into deploy/saas 2026-07-27 10:47:28 +08:00
林玮 (Jade Lin) 8073ba2fbe refactor(web): remove redundant app mode propagation 2026-07-27 10:42:07 +08:00
autofix-ci[bot]andGitHub 4919cc2395 [autofix.ci] apply automated fixes 2026-07-27 02:03:29 +00:00
林玮 (Jade Lin) e3dcf2efdb Merge remote-tracking branch 'origin/linw/snp-560-webapp-site-response-with-mode' into deploy/saas 2026-07-27 09:43:03 +08:00
hjlarry de7d4d4ab2 Merge branch 'build/plugin-list-improve' into deploy/saas 2026-07-25 15:36:41 +08:00
hjlarry a1efdeeabb fix: simplify plugin list API contracts 2026-07-25 15:34:38 +08:00
autofix-ci[bot]andGitHub 15cea0c08d [autofix.ci] apply automated fixes 2026-07-25 07:10:01 +00:00
hjlarry 543642542f Merge remote-tracking branch 'myori/deploy/saas' into deploy/saas 2026-07-25 15:07:07 +08:00
hjlarry 0d71273e30 Merge branch 'build/plugin-list-improve' into deploy/saas 2026-07-25 14:56:43 +08:00
hjlarry 671742d5b6 chore: improve summary api 2026-07-25 11:04:22 +08:00
林玮 (Jade Lin) e67d45cf7c fix: wait for workflow worker cleanup 2026-07-24 20:24:30 +08:00
林玮 (Jade Lin) 5d5439a169 feat: expose app mode in webapp site response 2026-07-24 20:15:10 +08:00
林玮 (Jade Lin) 6506d06f61 Merge remote-tracking branch 'origin/main' into deploy/saas 2026-07-24 16:00:54 +08:00
林玮 (Jade Lin) 689da4782f Merge remote-tracking branch 'origin/main' into deploy/saas 2026-07-24 16:00:43 +08:00
林玮 (Jade Lin) 3efe5f0da8 refactor(web): use semantic labelling for explore banner slides
(cherry picked from commit 8e89bd38ae)
2026-07-24 15:48:27 +08:00
林玮 (Jade Lin) 3665f6cb73 refactor(web): add accessible label to plan range switcher
(cherry picked from commit 810d028a3c)
2026-07-24 15:48:27 +08:00
林玮 (Jade Lin) 39ffee49ff refactor(web): use semantic markup for workflow save status
(cherry picked from commit a689278288)
2026-07-24 15:48:27 +08:00
林玮 (Jade Lin) 952d3482bb refactor(web): use semantic markup for billing usage info cards
(cherry picked from commit 87b9873bb7)
2026-07-24 15:48:27 +08:00
autofix-ci[bot]andGitHub 6fc6ac3d4c [autofix.ci] apply automated fixes 2026-07-24 05:42:34 +00:00
hjlarry c1647cd2a6 feat: model plugin summary api 2026-07-24 13:38:35 +08:00
林玮 (Jade Lin) 9b257f7df2 Merge remote-tracking branch 'origin/main' into deploy/saas 2026-07-23 18:18:26 +08:00
autofix-ci[bot]andGitHub aff527ce39 [autofix.ci] apply automated fixes 2026-07-23 08:36:56 +00:00
hjlarry 4ca9c36421 Merge remote-tracking branch 'origin/build/plugin-list-improve' into deploy/saas 2026-07-23 16:36:16 +08:00
hjlarry 00e7f4d984 feat: add category-scoped installed plugin ids 2026-07-23 16:13:38 +08:00
hjlarry 4375c7b988 feat: plugin list query improve 2026-07-23 10:32:26 +08:00
464 changed files with 20033 additions and 29470 deletions
-43
View File
@@ -1,7 +1,6 @@
name: Deploy Knowledge
permissions:
actions: read
contents: read
on:
@@ -19,48 +18,6 @@ jobs:
github.event.workflow_run.conclusion == 'success' &&
github.event.workflow_run.head_branch == 'deploy/konwledge'
steps:
- name: Wait for KnowledgeFS CI
uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0
timeout-minutes: 35
with:
github-token: ${{ secrets.GITHUB_TOKEN }}
script: |
const workflowId = "knowledge-fs-ci.yml";
const headBranch = context.payload.workflow_run.head_branch;
const headSha = context.payload.workflow_run.head_sha;
const deadline = Date.now() + 30 * 60 * 1000;
const pollIntervalMs = 15 * 1000;
while (Date.now() < deadline) {
const { data } = await github.rest.actions.listWorkflowRuns({
owner: context.repo.owner,
repo: context.repo.repo,
workflow_id: workflowId,
branch: headBranch,
event: "push",
head_sha: headSha,
per_page: 10,
});
const run = data.workflow_runs[0];
if (!run) {
core.info(`Waiting for ${workflowId} to start for ${headSha}.`);
} else if (run.status !== "completed") {
core.info(`Waiting for ${run.html_url}; current status is ${run.status}.`);
} else if (run.conclusion !== "success") {
throw new Error(
`${workflowId} did not succeed for ${headSha}: ${run.conclusion} (${run.html_url})`,
);
} else {
core.info(`KnowledgeFS CI succeeded for ${headSha}: ${run.html_url}`);
return;
}
await new Promise((resolve) => setTimeout(resolve, pollIntervalMs));
}
throw new Error(`Timed out waiting for ${workflowId} to succeed for ${headSha}.`);
- name: Deploy to server
uses: appleboy/ssh-action@0ff4204d59e8e51228ff73bce53f80d53301dee2 # v1.2.5
with:
-1
View File
@@ -63,7 +63,6 @@ jobs:
E2E_ADMIN_PASSWORD: E2eAdmin12345
E2E_FORCE_WEB_BUILD: "1"
E2E_INIT_PASSWORD: E2eInit12345
E2E_START_AGENT_BACKEND: "1"
run: vp run e2e:full
- name: Preserve Chromium E2E report and logs
-3
View File
@@ -666,7 +666,6 @@ PLUGIN_REMOTE_INSTALL_PORT=5003
PLUGIN_REMOTE_INSTALL_HOST=localhost
PLUGIN_MAX_PACKAGE_SIZE=15728640
PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600
PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED=true
PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400
# Comma-separated marketplace plugin IDs whose latest versions are installed for newly registered users.
# Example: langgenius/openai,langgenius/gemini
@@ -678,8 +677,6 @@ INNER_API_KEY_FOR_PLUGIN=QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y
# Dify Agent backend
AGENT_BACKEND_BASE_URL=http://localhost:5050
# Bearer token sent to the Agent backend /runs API. Must match DIFY_AGENT_API_TOKEN on the server side.
AGENT_BACKEND_API_TOKEN=dify-agent-run-token-for-dev-only
AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30
AGENT_BACKEND_STREAM_MAX_RECONNECTS=3
AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200
+2 -2
View File
@@ -99,9 +99,9 @@ ENV VIRTUAL_ENV=/app/api/.venv
COPY --from=packages --chown=dify:dify ${VIRTUAL_ENV} ${VIRTUAL_ENV}
ENV PATH="${VIRTUAL_ENV}/bin:${PATH}"
# Download nltk data
RUN mkdir -p /usr/local/share/nltk_data \
&& NLTK_DATA=/usr/local/share/nltk_data python -m nltk.downloader punkt_tab averaged_perceptron_tagger_eng stopwords \
&& NLTK_DATA=/usr/local/share/nltk_data python -c "import nltk; nltk.data.find('tokenizers/punkt_tab'); nltk.data.find('taggers/averaged_perceptron_tagger_eng'); nltk.data.find('corpora/stopwords')" \
&& NLTK_DATA=/usr/local/share/nltk_data python -c "import nltk; nltk.download('punkt'); nltk.download('averaged_perceptron_tagger'); nltk.download('stopwords')" \
&& chmod -R 755 /usr/local/share/nltk_data
ENV TIKTOKEN_CACHE_DIR=/app/api/.tiktoken_cache
+30 -59
View File
@@ -1,13 +1,10 @@
import logging
import time
from collections.abc import Callable
from typing import NamedTuple
import socketio
from flask import request
from opentelemetry.trace import get_current_span
from opentelemetry.trace.span import INVALID_SPAN_ID, INVALID_TRACE_ID
from werkzeug.exceptions import Forbidden, HTTPException, ServiceUnavailable
from configs import dify_config
from contexts.wrapper import RecyclableContextVar
@@ -45,53 +42,6 @@ _CONSOLE_EXEMPT_PREFIXES = (
"/console/api/activate/check",
)
_WEBAPP_EXEMPT_PREFIXES = ("/api/system-features",)
_INVALID_LICENSE_STATUSES = (LicenseStatus.INACTIVE, LicenseStatus.EXPIRED, LicenseStatus.LOST)
def _session_surface_error(license_status: LicenseStatus | None) -> HTTPException:
if license_status is None:
return UnauthorizedAndForceLogout("Unable to verify enterprise license. Please contact your administrator.")
return UnauthorizedAndForceLogout(f"Enterprise license is {license_status}. Please contact your administrator.")
def _bearer_surface_error(license_status: LicenseStatus | None) -> HTTPException:
"""Token-authed: forcing a logout is meaningless and license state must not leak."""
return Forbidden(description="license_required")
def _retryable_surface_error(license_status: LicenseStatus | None) -> HTTPException:
"""Webhook senders retry on 5xx but treat 4xx as permanent, disabling the subscription."""
return ServiceUnavailable(description="license_required")
class _LicenseGatedSurface(NamedTuple):
prefix: str
exempt_prefixes: tuple[str, ...]
build_error: Callable[[LicenseStatus | None], HTTPException]
# /files (plugin-daemon data plane), /inner/api (enterprise control plane) and /health
# stay ungated: blocking them breaks workflow execution or license recovery itself.
_LICENSE_GATED_SURFACES = (
_LicenseGatedSurface("/console/api/", _CONSOLE_EXEMPT_PREFIXES, _session_surface_error),
_LicenseGatedSurface("/api/", _WEBAPP_EXEMPT_PREFIXES, _session_surface_error),
_LicenseGatedSurface("/v1", (), _bearer_surface_error),
_LicenseGatedSurface("/mcp", (), _bearer_surface_error),
_LicenseGatedSurface("/triggers", (), _retryable_surface_error),
)
def _match_license_gated_surface(path: str) -> _LicenseGatedSurface | None:
for surface in _LICENSE_GATED_SURFACES:
if not path.startswith(surface.prefix):
continue
if any(path.startswith(exempt) for exempt in surface.exempt_prefixes):
return None
return surface
return None
# ----------------------------
# Application Factory Function
@@ -112,17 +62,38 @@ def create_flask_app_with_configs() -> DifyApp:
init_request_context()
RecyclableContextVar.increment_thread_recycles()
# Enterprise license validation for API endpoints (both console and webapp)
# When license expires, block all API access except bootstrap endpoints needed
# for the frontend to load the license expiration page without infinite reloads.
if dify_config.ENTERPRISE_ENABLED:
surface = _match_license_gated_surface(request.path)
if surface is not None:
try:
license_status = EnterpriseService.get_cached_license_status()
except Exception:
logger.exception("Failed to check enterprise license status")
license_status = None
is_console_api = request.path.startswith("/console/api/")
is_webapp_api = request.path.startswith("/api/")
if license_status is None or license_status in _INVALID_LICENSE_STATUSES:
raise surface.build_error(license_status)
if is_console_api or is_webapp_api:
if is_console_api:
is_exempt = any(request.path.startswith(p) for p in _CONSOLE_EXEMPT_PREFIXES)
else: # webapp API
is_exempt = request.path.startswith("/api/system-features")
if not is_exempt:
try:
# Check license status (cached — see EnterpriseService for TTL details)
license_status = EnterpriseService.get_cached_license_status()
if license_status in (LicenseStatus.INACTIVE, LicenseStatus.EXPIRED, LicenseStatus.LOST):
raise UnauthorizedAndForceLogout(
f"Enterprise license is {license_status}. Please contact your administrator."
)
if license_status is None:
raise UnauthorizedAndForceLogout(
"Unable to verify enterprise license. Please contact your administrator."
)
except UnauthorizedAndForceLogout:
raise
except Exception:
logger.exception("Failed to check enterprise license status")
raise UnauthorizedAndForceLogout(
"Unable to verify enterprise license. Please contact your administrator."
)
# add after request hook for injecting trace headers from OpenTelemetry span context
# Only adds headers when OTEL is enabled and has valid context
+12
View File
@@ -5,6 +5,8 @@ API adapters: request building from Dify product concepts, a thin client wrapper
event adaptation for future workflow integration, and deterministic fakes.
"""
from dify_agent.protocol import RuntimeLayerSpec, extract_runtime_layer_specs
from clients.agent_backend.client import AgentBackendRunClient, DifyAgentBackendRunClient
from clients.agent_backend.errors import (
AgentBackendError,
@@ -45,6 +47,11 @@ from clients.agent_backend.request_builder import (
AgentBackendWorkflowNodeRunInput,
redact_for_agent_backend_log,
)
from clients.agent_backend.session_cleanup import (
AgentBackendSessionCleanupPayload,
AgentBackendSessionCleanupResult,
cleanup_agent_backend_session,
)
__all__ = [
"AGENT_SOUL_PROMPT_LAYER_ID",
@@ -73,6 +80,8 @@ __all__ = [
"AgentBackendRunRequestBuilder",
"AgentBackendRunStartedInternalEvent",
"AgentBackendRunSucceededInternalEvent",
"AgentBackendSessionCleanupPayload",
"AgentBackendSessionCleanupResult",
"AgentBackendStreamError",
"AgentBackendStreamInternalEvent",
"AgentBackendTransportError",
@@ -81,6 +90,9 @@ __all__ = [
"DifyAgentBackendRunClient",
"FakeAgentBackendRunClient",
"FakeAgentBackendScenario",
"RuntimeLayerSpec",
"cleanup_agent_backend_session",
"create_agent_backend_run_client",
"extract_runtime_layer_specs",
"redact_for_agent_backend_log",
]
+1 -5
View File
@@ -11,7 +11,6 @@ from clients.agent_backend.fake_client import FakeAgentBackendRunClient, FakeAge
def create_agent_backend_run_client(
*,
base_url: str | None = None,
api_token: str | None = None,
use_fake: bool = False,
fake_scenario: str | FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS,
stream_read_timeout_seconds: float = 30,
@@ -23,11 +22,8 @@ def create_agent_backend_run_client(
return FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario(fake_scenario))
if base_url is None:
raise ValueError("base_url is required when creating a real Agent backend client")
headers: dict[str, str] = {}
if api_token:
headers["Authorization"] = f"Bearer {api_token}"
return DifyAgentBackendRunClient(
Client(base_url=base_url, stream_timeout=stream_read_timeout_seconds, headers=headers),
Client(base_url=base_url, stream_timeout=stream_read_timeout_seconds),
stream_max_reconnects=stream_max_reconnects,
stream_timeout_seconds=stream_run_timeout_seconds,
)
+76 -30
View File
@@ -16,6 +16,7 @@ from collections.abc import Mapping
from typing import ClassVar, Literal
from agenton.compositor import CompositorSessionSnapshot
from agenton.compositor.schemas import LayerSessionSnapshot
from agenton.layers import ExitIntent
from agenton_collections.layers.plain import PLAIN_PROMPT_LAYER_TYPE_ID, PromptLayerConfig
from agenton_collections.layers.pydantic_ai import PYDANTIC_AI_HISTORY_LAYER_TYPE_ID
@@ -36,7 +37,6 @@ from dify_agent.layers.execution_context import (
)
from dify_agent.layers.knowledge import DIFY_KNOWLEDGE_BASE_LAYER_TYPE_ID, DifyKnowledgeBaseLayerConfig
from dify_agent.layers.output import DIFY_OUTPUT_LAYER_TYPE_ID, DifyOutputLayerConfig
from dify_agent.layers.runtime import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig
from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig
from dify_agent.protocol import (
DIFY_AGENT_HISTORY_LAYER_ID,
@@ -47,6 +47,7 @@ from dify_agent.protocol import (
LayerExitSignals,
RunComposition,
RunLayerSpec,
RuntimeLayerSpec,
)
from pydantic import BaseModel, ConfigDict, Field, JsonValue, field_validator
@@ -55,7 +56,6 @@ WORKFLOW_NODE_JOB_PROMPT_LAYER_ID = "workflow_node_job_prompt"
WORKFLOW_USER_PROMPT_LAYER_ID = "workflow_user_prompt"
AGENT_APP_USER_PROMPT_LAYER_ID = "agent_app_user_prompt"
DIFY_EXECUTION_CONTEXT_LAYER_ID = "execution_context"
DIFY_RUNTIME_LAYER_ID = "runtime"
DIFY_CONFIG_LAYER_ID = "config"
DIFY_DRIVE_LAYER_ID = "drive"
DIFY_PLUGIN_TOOLS_LAYER_ID = "tools"
@@ -66,11 +66,25 @@ DIFY_SHELL_LAYER_ID = "shell"
type AgentConfigVersionKind = Literal["snapshot", "draft", "build_draft"]
def _filter_snapshot_to_specs(
snapshot: CompositorSessionSnapshot,
specs: list[RuntimeLayerSpec],
) -> CompositorSessionSnapshot:
"""Keep only snapshot layers whose names appear in the cleanup spec list.
The agenton compositor rejects a snapshot whose layer-name sequence does
not match the active composition exactly. Cleanup-replay drops plugin
layers, so we must drop the matching snapshot entries here.
"""
kept_names = {spec.name for spec in specs}
filtered_layers: list[LayerSessionSnapshot] = [layer for layer in snapshot.layers if layer.name in kept_names]
if len(filtered_layers) == len(snapshot.layers):
return snapshot
return CompositorSessionSnapshot(schema_version=snapshot.schema_version, layers=filtered_layers)
def _shell_layer_deps() -> dict[str, str]:
return {
"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID,
"runtime": DIFY_RUNTIME_LAYER_ID,
}
return {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID}
def _drive_layer_deps() -> dict[str, str]:
@@ -200,7 +214,6 @@ class AgentBackendWorkflowNodeRunInput(BaseModel):
model: AgentBackendModelConfig
execution_context: DifyExecutionContextLayerConfig
backend_binding_ref: str = Field(min_length=1)
workflow_node_job_prompt: str
user_prompt: str
agent_soul_prompt: str | None = None
@@ -218,8 +231,8 @@ class AgentBackendWorkflowNodeRunInput(BaseModel):
# the Agent Soul configures human involvement; a deferred call ends the run and
# the workflow pauses via the existing HITL form mechanism (ENG-635).
ask_human_config: DifyAskHumanLayerConfig | None = None
# Inject the sandboxed shell graph. Requires a deployment-selected runtime
# backend plus the product-resolved persistent Binding.
# Inject the sandboxed shell layer (dify.shell). Requires the agent backend
# to be wired with a shellctl entrypoint; see configs AGENT_SHELL_ENABLED.
include_shell: bool = False
shell_config: DifyShellLayerConfig | None = None
session_snapshot: CompositorSessionSnapshot | None = None
@@ -227,6 +240,7 @@ class AgentBackendWorkflowNodeRunInput(BaseModel):
# (ENG-638). Keyed by the original deferred tool_call_id.
deferred_tool_results: DeferredToolResultsPayload | None = None
include_history: bool = True
suspend_on_exit: bool = True
metadata: dict[str, JsonValue] = Field(default_factory=dict)
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", arbitrary_types_allowed=True)
@@ -250,7 +264,6 @@ class AgentBackendAgentAppRunInput(BaseModel):
model: AgentBackendModelConfig
execution_context: DifyExecutionContextLayerConfig
backend_binding_ref: str = Field(min_length=1)
user_prompt: str
agent_soul_prompt: str | None = None
agent_config_version_kind: AgentConfigVersionKind = "snapshot"
@@ -266,8 +279,8 @@ class AgentBackendAgentAppRunInput(BaseModel):
# Human-in-the-loop ask_human deferred tool (dify.ask_human). Present only when
# the Agent Soul configures human involvement (ENG-635).
ask_human_config: DifyAskHumanLayerConfig | None = None
# Inject the sandboxed shell graph. Requires a deployment-selected runtime
# backend plus the product-resolved persistent Binding.
# Inject the sandboxed shell layer (dify.shell). Requires the agent backend
# to be wired with a shellctl entrypoint; see configs AGENT_SHELL_ENABLED.
include_shell: bool = False
shell_config: DifyShellLayerConfig | None = None
session_snapshot: CompositorSessionSnapshot | None = None
@@ -275,6 +288,7 @@ class AgentBackendAgentAppRunInput(BaseModel):
# (ENG-638). Keyed by the original deferred tool_call_id.
deferred_tool_results: DeferredToolResultsPayload | None = None
include_history: bool = True
suspend_on_exit: bool = True
metadata: dict[str, JsonValue] = Field(default_factory=dict)
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", arbitrary_types_allowed=True)
@@ -336,14 +350,6 @@ class AgentBackendRunRequestBuilder:
run_input.include_shell or run_input.config_layer_config is not None or run_input.drive_config is not None
)
if include_shell:
layers.append(
RunLayerSpec(
name=DIFY_RUNTIME_LAYER_ID,
type=DIFY_RUNTIME_LAYER_TYPE_ID,
metadata=run_input.metadata,
config=DifyRuntimeLayerConfig(backend_binding_ref=run_input.backend_binding_ref),
)
)
# Sandboxed bash workspace (dify.shell). It enters before config/drive
# so eager pulls materialize content in the same filesystem used by
# model commands.
@@ -475,7 +481,53 @@ class AgentBackendRunRequestBuilder:
metadata=run_input.metadata,
session_snapshot=run_input.session_snapshot,
deferred_tool_results=run_input.deferred_tool_results,
on_exit=LayerExitSignals(default=ExitIntent.SUSPEND),
on_exit=LayerExitSignals(
default=ExitIntent.SUSPEND if run_input.suspend_on_exit else ExitIntent.DELETE,
),
)
def build_cleanup_request(
self,
*,
session_snapshot: CompositorSessionSnapshot,
runtime_layer_specs: list[RuntimeLayerSpec],
idempotency_key: str | None = None,
metadata: dict[str, JsonValue] | None = None,
) -> CreateRunRequest:
"""Build a lifecycle-only cleanup request that replays the prior layers.
The agenton compositor enforces that the session snapshot's layer names
match the active composition in order, so cleanup must replay the same
non-plugin layer graph that produced the snapshot. Plugin layers
(``dify.plugin.llm``, ``dify.plugin.tools``) are excluded from both the
composition and the snapshot before submission because their configs
may carry credentials or runtime-only declarations that are not
persisted between runs.
"""
if not runtime_layer_specs:
raise ValueError(
"build_cleanup_request requires runtime_layer_specs; an empty "
"composition would fail the agent backend's snapshot validation."
)
request_metadata = dict(metadata or {})
request_metadata["agent_backend_lifecycle"] = "session_cleanup"
layers = [
RunLayerSpec(
name=spec.name,
type=spec.type,
deps=dict(spec.deps),
metadata=dict(spec.metadata),
config=spec.config,
)
for spec in runtime_layer_specs
]
filtered_snapshot = _filter_snapshot_to_specs(session_snapshot, runtime_layer_specs)
return CreateRunRequest(
composition=RunComposition(layers=layers),
idempotency_key=idempotency_key,
metadata=request_metadata,
session_snapshot=filtered_snapshot,
on_exit=LayerExitSignals(default=ExitIntent.DELETE),
)
def build_for_workflow_node(self, run_input: AgentBackendWorkflowNodeRunInput) -> CreateRunRequest:
@@ -528,14 +580,6 @@ class AgentBackendRunRequestBuilder:
run_input.include_shell or run_input.config_layer_config is not None or run_input.drive_config is not None
)
if include_shell:
layers.append(
RunLayerSpec(
name=DIFY_RUNTIME_LAYER_ID,
type=DIFY_RUNTIME_LAYER_TYPE_ID,
metadata=run_input.metadata,
config=DifyRuntimeLayerConfig(backend_binding_ref=run_input.backend_binding_ref),
)
)
# Sandboxed bash workspace (dify.shell). It enters before drive so
# drive can materialize mentioned targets with `dify-agent drive pull`
# in the same shell-visible filesystem used by model commands.
@@ -669,7 +713,9 @@ class AgentBackendRunRequestBuilder:
metadata=run_input.metadata,
session_snapshot=run_input.session_snapshot,
deferred_tool_results=run_input.deferred_tool_results,
on_exit=LayerExitSignals(default=ExitIntent.SUSPEND),
on_exit=LayerExitSignals(
default=ExitIntent.SUSPEND if run_input.suspend_on_exit else ExitIntent.DELETE,
),
)
@@ -0,0 +1,100 @@
"""Shared API-side helper for Agent backend lifecycle-only session cleanup.
Product code owns local row retirement and background-task dispatch. This module
only adapts persisted cleanup inputs into the public ``dify-agent`` run
protocol, performs the synchronous ``create_run + wait_run`` loop used by Celery
workers, and reports whether the backend cleanup succeeded, was skipped, or
failed.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import ClassVar, Literal
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.protocol import RuntimeLayerSpec
from pydantic import BaseModel, ConfigDict, Field, JsonValue
from clients.agent_backend.client import AgentBackendRunClient
from clients.agent_backend.errors import AgentBackendError
from clients.agent_backend.request_builder import AgentBackendRunRequestBuilder
class AgentBackendSessionCleanupPayload(BaseModel):
"""Serialized cleanup inputs preserved across API and Celery boundaries."""
session_snapshot: CompositorSessionSnapshot | None = None
runtime_layer_specs: list[RuntimeLayerSpec] = Field(default_factory=list)
idempotency_key: str | None = None
metadata: dict[str, JsonValue] = Field(default_factory=dict)
timeout_seconds: float = 30.0
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid")
@dataclass(frozen=True, slots=True)
class AgentBackendSessionCleanupResult:
"""Terminal outcome of one backend cleanup attempt."""
status: Literal["succeeded", "skipped", "failed"]
reason: str | None = None
cleanup_run_id: str | None = None
@classmethod
def succeeded(cls, cleanup_run_id: str) -> AgentBackendSessionCleanupResult:
return cls(status="succeeded", cleanup_run_id=cleanup_run_id)
@classmethod
def skipped(cls, reason: str) -> AgentBackendSessionCleanupResult:
return cls(status="skipped", reason=reason)
@classmethod
def failed(cls, reason: str, cleanup_run_id: str | None = None) -> AgentBackendSessionCleanupResult:
return cls(status="failed", reason=reason, cleanup_run_id=cleanup_run_id)
def cleanup_agent_backend_session(
*,
payload: AgentBackendSessionCleanupPayload,
client: AgentBackendRunClient | None,
request_builder: AgentBackendRunRequestBuilder | None = None,
) -> AgentBackendSessionCleanupResult:
"""Run lifecycle-only cleanup against the Agent backend and report status."""
if client is None:
return AgentBackendSessionCleanupResult.skipped("no_agent_backend_client")
if payload.session_snapshot is None:
return AgentBackendSessionCleanupResult.skipped("missing_session_snapshot")
if not payload.runtime_layer_specs:
return AgentBackendSessionCleanupResult.skipped("missing_runtime_layer_specs")
builder = request_builder or AgentBackendRunRequestBuilder()
request = builder.build_cleanup_request(
session_snapshot=payload.session_snapshot,
runtime_layer_specs=payload.runtime_layer_specs,
idempotency_key=payload.idempotency_key,
metadata=payload.metadata,
)
try:
response = client.create_run(request)
except AgentBackendError as exc:
return AgentBackendSessionCleanupResult.failed(str(exc))
try:
status_response = client.wait_run(response.run_id, timeout_seconds=payload.timeout_seconds)
except AgentBackendError as exc:
return AgentBackendSessionCleanupResult.failed(str(exc), cleanup_run_id=response.run_id)
if status_response.status != "succeeded":
reason = status_response.error or f"cleanup run ended with status {status_response.status}"
return AgentBackendSessionCleanupResult.failed(reason, cleanup_run_id=response.run_id)
return AgentBackendSessionCleanupResult.succeeded(response.run_id)
__all__ = [
"AgentBackendSessionCleanupPayload",
"AgentBackendSessionCleanupResult",
"cleanup_agent_backend_session",
]
+3 -7
View File
@@ -12,11 +12,6 @@ class AgentBackendConfig(BaseSettings):
default=None,
)
AGENT_BACKEND_API_TOKEN: str | None = Field(
description="Bearer token for authenticating with the Agent backend /runs API.",
default=None,
)
AGENT_BACKEND_USE_FAKE: bool = Field(
description="Use the deterministic in-process fake Agent backend client.",
default=False,
@@ -44,8 +39,9 @@ class AgentBackendConfig(BaseSettings):
AGENT_SHELL_ENABLED: bool = Field(
description=(
"Inject the Home, Workspace, Sandbox, and Shell runtime layers into Agent runs. "
"Requires Dify Agent to have a deployment-selected runtime backend."
"Inject the dify.shell layer (sandboxed bash workspace) into Agent runs. "
"Requires the agent backend to be wired with a shellctl entrypoint before "
"shell-using Agent runs are executed."
),
default=True,
)
-6
View File
@@ -266,12 +266,6 @@ class PluginConfig(BaseSettings):
default=60 * 60,
)
PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED: bool = Field(
description="Whether tenant plugin model providers are cached in Redis. Disable when plugins are installed "
"by a system other than this one, which cannot invalidate the cache when a tenant's plugins change.",
default=True,
)
PLUGIN_MODEL_PROVIDERS_CACHE_TTL: PositiveInt = Field(
description="TTL in seconds for caching tenant plugin model providers in Redis",
default=60 * 60 * 24,
+1 -1
View File
@@ -52,7 +52,7 @@ def with_session[T, **P, R](
session.commit()
return result
except Exception:
session.rollback() # guard-ignore: no-new-controller-sqlalchemy -- decorator owns rollback
session.rollback() # noqa: no-new-controller-sqlalchemy decorator owns transaction rollback
raise
with session_factory.create_session() as session:
+29 -7
View File
@@ -62,7 +62,7 @@ from libs.datetime_utils import parse_time_range
from libs.helper import dump_response
from libs.login import login_required
from models import Account
from models.agent import Agent, AgentStatus
from models.agent import Agent, AgentConfigDraftType, AgentStatus
from models.agent_config_entities import AgentSoulConfig
from models.enums import ApiTokenType
from models.model import ApiToken, App, IconType
@@ -265,6 +265,13 @@ class AgentDebugConversationRefreshResponse(BaseModel):
debug_conversation_message_count: int = 0
class AgentDebugConversationRefreshPayload(BaseModel):
draft_type: AgentConfigDraftType = Field(
default=AgentConfigDraftType.DEBUG_BUILD,
description="Agent draft surface whose conversation should be refreshed",
)
class AgentPublishPayload(BaseModel):
version_note: str | None = Field(default=None, description="Optional note for this published Agent version")
@@ -308,6 +315,7 @@ register_schema_models(
AgentAppCopyPayload,
AgentPublishPayload,
AgentBuildDraftCheckoutPayload,
AgentDebugConversationRefreshPayload,
ComposerSavePayload,
AgentApiStatusPayload,
AgentInviteOptionsQuery,
@@ -387,10 +395,11 @@ def _serialize_agent_app_detail(
payload["backing_app_id"] = roster_service.runtime_backing_app_id(agent)
payload["hidden_app_backed"] = bool(agent.backing_app_id and agent.backing_app_id != agent.app_id)
payload["id"] = agent.id
debug_conversation_id = roster_service.get_or_create_build_conversation(
debug_conversation_id = roster_service.get_or_create_agent_app_debug_conversation_id(
tenant_id=app_model.tenant_id,
agent_id=agent.id,
account_id=current_user.id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
commit=False,
)
message_count = roster_service.count_agent_app_debug_conversation_messages(
@@ -430,10 +439,11 @@ def _serialize_agent_app_pagination(session: Session, app_pagination, *, tenant_
tenant_id=tenant_id,
agent_ids=[agent.id for agent in agents_by_app_id.values()],
)
debug_conversation_ids_by_agent_id = roster_service.load_or_create_build_conversation_ids_by_agent_id(
debug_conversation_ids_by_agent_id = roster_service.load_or_create_agent_app_debug_conversation_ids_by_agent_id(
tenant_id=tenant_id,
agents=list(agents_by_app_id.values()),
account_id=current_user.id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
)
payload = AgentAppPagination.model_validate(
app_pagination,
@@ -668,6 +678,16 @@ class AgentAppApi(Resource):
@console_ns.route("/agent/<uuid:agent_id>/debug-conversation/refresh")
class AgentDebugConversationRefreshApi(Resource):
@console_ns.expect(console_ns.models[AgentDebugConversationRefreshPayload.__name__])
@console_ns.doc(
params={
"payload": {
"in": "body",
"required": False,
"schema": {"$ref": f"#/components/schemas/{AgentDebugConversationRefreshPayload.__name__}"},
}
}
)
@console_ns.response(
200,
"Agent debug conversation refreshed",
@@ -682,10 +702,12 @@ class AgentDebugConversationRefreshApi(Resource):
@with_current_tenant_id
@with_session
def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID):
debug_conversation_id = _agent_roster_service(session).reset_build_conversation(
args = AgentDebugConversationRefreshPayload.model_validate(request.get_json(silent=True) or {})
debug_conversation_id = _agent_roster_service(session).refresh_agent_app_debug_conversation_id(
tenant_id=tenant_id,
agent_id=str(agent_id),
account_id=current_user.id,
draft_type=args.draft_type,
)
return AgentDebugConversationRefreshResponse(
debug_conversation_id=debug_conversation_id,
@@ -729,7 +751,7 @@ class AgentBuildDraftCheckoutApi(Resource):
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@with_session
def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID):
args = AgentBuildDraftCheckoutPayload.model_validate(console_ns.payload or {})
return AgentComposerService.checkout_agent_app_build_draft(
@@ -786,7 +808,7 @@ class AgentBuildDraftApi(Resource):
@edit_permission_required
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@with_session
def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID):
return AgentComposerService.discard_agent_app_build_draft(
session=session,
@@ -806,7 +828,7 @@ class AgentBuildDraftApplyApi(Resource):
@rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False)
@with_current_user
@with_current_tenant_id
@with_session(write=False)
@with_session
def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID):
return AgentComposerService.apply_agent_app_build_draft(
session=session,
@@ -1,7 +1,7 @@
"""Console routes for Agent App and workflow Agent sandbox file access.
The API accepts product-facing Conversation, Build Draft, or Workflow Node
Execution locators and proxies list/read/upload to the agent backend's
The API keeps product-facing locators (conversation or workflow node identity)
on this public boundary and proxies list/read/upload to the agent backend's new
``/sandbox`` contract.
"""
@@ -26,16 +26,10 @@ from controllers.common.session import with_session
from controllers.console import console_ns
from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model
from controllers.console.app.wraps import get_app_model
from controllers.console.wraps import (
account_initialization_required,
setup_required,
with_current_tenant_id,
with_current_user,
)
from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id
from extensions.ext_database import db
from fields.base import ResponseModel
from libs.login import login_required
from models import Account
from models.model import App, AppMode
from services.agent_app_sandbox_service import (
AgentAppSandboxService,
@@ -43,43 +37,52 @@ from services.agent_app_sandbox_service import (
WorkflowAgentSandboxService,
)
_NODE_EXECUTION_ID_DESCRIPTION = (
"Optional workflow node execution ID. When omitted, the latest active session for the node is used."
)
class AgentSandboxListQuery(BaseModel):
caller_type: Literal["conversation", "build_draft"]
caller_id: str = Field(min_length=1, description="Agent App caller ID")
conversation_id: str = Field(min_length=1, description="Agent App conversation ID")
path: str = Field(default=".", description="Directory path relative to the sandbox workspace")
class AgentSandboxInfoQuery(BaseModel):
caller_type: Literal["conversation", "build_draft"]
caller_id: str = Field(min_length=1, description="Agent App caller ID")
conversation_id: str = Field(min_length=1, description="Agent App conversation ID")
class AgentSandboxFileQuery(BaseModel):
caller_type: Literal["conversation", "build_draft"]
caller_id: str = Field(min_length=1, description="Agent App caller ID")
conversation_id: str = Field(min_length=1, description="Agent App conversation ID")
path: str = Field(min_length=1, description="File path relative to the sandbox workspace")
class AgentSandboxUploadPayload(BaseModel):
caller_type: Literal["conversation", "build_draft"]
caller_id: str = Field(min_length=1, description="Agent App caller ID")
conversation_id: str = Field(min_length=1, description="Agent App conversation ID")
path: str = Field(min_length=1, description="File path relative to the sandbox workspace")
class WorkflowAgentSandboxListQuery(BaseModel):
node_execution_id: str = Field(min_length=1, description="Workflow node execution ID")
path: str = Field(default=".", description="Directory path relative to the sandbox workspace")
node_execution_id: str | None = Field(
default=None,
description=_NODE_EXECUTION_ID_DESCRIPTION,
)
class WorkflowAgentSandboxFileQuery(BaseModel):
node_execution_id: str = Field(min_length=1, description="Workflow node execution ID")
path: str = Field(min_length=1, description="File path relative to the sandbox workspace")
node_execution_id: str | None = Field(
default=None,
description=_NODE_EXECUTION_ID_DESCRIPTION,
)
class WorkflowAgentSandboxUploadPayload(BaseModel):
node_execution_id: str = Field(min_length=1, description="Workflow node execution ID")
path: str = Field(min_length=1, description="File path relative to the sandbox workspace")
node_execution_id: str | None = Field(
default=None,
description=_NODE_EXECUTION_ID_DESCRIPTION,
)
class SandboxFileEntryResponse(ResponseModel):
@@ -96,6 +99,7 @@ class SandboxListResponse(ResponseModel):
class SandboxInfoResponse(ResponseModel):
session_id: str
workspace_cwd: str
@@ -151,19 +155,15 @@ class AgentAppSandboxInfoResource(Resource):
@login_required
@account_initialization_required
@with_current_tenant_id
@with_current_user
@with_session(write=False)
def get(self, session: Session, current_user: Account, tenant_id: str, agent_id: UUID):
def get(self, session: Session, tenant_id: str, agent_id: UUID):
app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id)
query = query_params_from_request(AgentSandboxInfoQuery)
try:
result = AgentAppSandboxService().get_info(
tenant_id=tenant_id,
app_id=app_model.id,
agent_id=str(agent_id),
caller_type=query.caller_type,
caller_id=query.caller_id,
account_id=current_user.id,
conversation_id=query.conversation_id,
)
except Exception as exc:
return _handle(exc)
@@ -180,19 +180,15 @@ class AgentAppSandboxListResource(Resource):
@login_required
@account_initialization_required
@with_current_tenant_id
@with_current_user
@with_session(write=False)
def get(self, session: Session, current_user: Account, tenant_id: str, agent_id: UUID):
def get(self, session: Session, tenant_id: str, agent_id: UUID):
app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id)
query = query_params_from_request(AgentSandboxListQuery)
try:
result = AgentAppSandboxService().list_files(
tenant_id=tenant_id,
app_id=app_model.id,
agent_id=str(agent_id),
caller_type=query.caller_type,
caller_id=query.caller_id,
account_id=current_user.id,
conversation_id=query.conversation_id,
path=query.path,
)
except Exception as exc:
@@ -210,19 +206,15 @@ class AgentAppSandboxReadResource(Resource):
@login_required
@account_initialization_required
@with_current_tenant_id
@with_current_user
@with_session(write=False)
def get(self, session: Session, current_user: Account, tenant_id: str, agent_id: UUID):
def get(self, session: Session, tenant_id: str, agent_id: UUID):
app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id)
query = query_params_from_request(AgentSandboxFileQuery)
try:
result = AgentAppSandboxService().read_file(
tenant_id=tenant_id,
app_id=app_model.id,
agent_id=str(agent_id),
caller_type=query.caller_type,
caller_id=query.caller_id,
account_id=current_user.id,
conversation_id=query.conversation_id,
path=query.path,
)
except Exception as exc:
@@ -240,19 +232,15 @@ class AgentAppSandboxUploadResource(Resource):
@login_required
@account_initialization_required
@with_current_tenant_id
@with_current_user
@with_session(write=False)
def post(self, session: Session, current_user: Account, tenant_id: str, agent_id: UUID):
def post(self, session: Session, tenant_id: str, agent_id: UUID):
app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id)
payload = AgentSandboxUploadPayload.model_validate(request.get_json(silent=True) or {})
try:
result = AgentAppSandboxService().upload_file(
tenant_id=tenant_id,
app_id=app_model.id,
agent_id=str(agent_id),
caller_type=payload.caller_type,
caller_id=payload.caller_id,
account_id=current_user.id,
conversation_id=payload.conversation_id,
path=payload.path,
)
except Exception as exc:
-77
View File
@@ -58,7 +58,6 @@ from services.app_service import (
AppResponseView,
AppService,
CreateAppParams,
RecentAppMode,
StarredAppListParams,
)
from services.enterprise import rbac_service as enterprise_rbac_service
@@ -140,10 +139,6 @@ class AppListBaseQuery(BaseModel):
raise ValueError("Invalid UUID format in creator_ids.") from exc
class RecentAppListQuery(BaseModel):
limit: int = Field(default=8, ge=1, le=8, description="Number of recently modified apps to return (1-8)")
class AppListQuery(AppListBaseQuery):
pass
@@ -416,33 +411,6 @@ class AppPartial(AppResponseModel):
return to_timestamp(value)
class RecentAppResponse(ResponseModel):
id: str
name: str
icon_type: IconType | None = None
icon: str | None = None
icon_background: str | None = None
mode: RecentAppMode
author_name: str | None = None
updated_at: int
permission_keys: list[str] = Field(default_factory=list)
maintainer: str | None = None
@computed_field(return_type=str | None) # type: ignore[prop-decorator]
@property
def icon_url(self) -> str | None:
return build_icon_url(self.icon_type, self.icon)
@field_validator("updated_at", mode="before")
@classmethod
def _normalize_timestamp(cls, value: datetime | int) -> int:
return to_timestamp(value)
class RecentAppListResponse(ResponseModel):
data: list[RecentAppResponse]
class AppDetail(AppResponseModel):
id: str
name: str
@@ -607,8 +575,6 @@ register_schema_models(
register_response_schema_models(
console_ns,
AppPartial,
RecentAppResponse,
RecentAppListResponse,
AppDetailWithSite,
AppPagination,
)
@@ -733,49 +699,6 @@ class AppListApi(Resource):
return app_detail.model_dump(mode="json"), 201
@console_ns.route("/apps/recent")
class RecentAppListApi(Resource):
@console_ns.doc("list_recent_apps")
@console_ns.doc(description="Get recently modified apps for the home Continue Work section")
@console_ns.doc(params=query_params_from_model(RecentAppListQuery))
@console_ns.response(200, "Success", console_ns.models[RecentAppListResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@enterprise_license_required
@with_session(write=False)
@with_current_user_id
@with_current_tenant_id
def get(self, current_tenant_id: str, current_user_id: str, session: Session):
"""Return the lightweight app cards needed by the Explore home page."""
args = query_params_from_request(RecentAppListQuery)
params = AppListParams(limit=args.limit)
permissions = enterprise_rbac_service.RBACService.MyPermissions.get(
current_tenant_id,
current_user_id,
session=session,
)
if dify_config.RBAC_ENABLED:
access_filter = resolve_app_access_filter(
current_tenant_id,
current_user_id,
session=session,
permissions=permissions,
)
access_filter.apply_to_params(params)
recent_apps = AppService().get_recent_apps(current_user_id, current_tenant_id, params, session)
permission_keys_map = permissions.app.permission_keys_by_resource_ids([app.id for app in recent_apps])
response_items = [
RecentAppResponse.model_validate(app, from_attributes=True).model_copy(
update={"permission_keys": permission_keys_map.get(app.id, [])}
)
for app in recent_apps
]
return dump_response(RecentAppListResponse, {"data": response_items}), 200
@console_ns.route("/apps/starred")
class StarredAppListApi(Resource):
@console_ns.doc("list_starred_apps")
+15 -18
View File
@@ -36,7 +36,7 @@ from controllers.console.wraps import (
with_current_user_id,
)
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.app_invoke_entities import AGENT_RUNTIME_EXIT_INTENT_ARG, InvokeFrom
from core.app.features.rate_limiting.rate_limit import RateLimitGenerator
from core.errors.error import (
ModelCurrentlyNotSupportError,
@@ -353,7 +353,12 @@ def _resolve_current_user_agent_debug_conversation_id(
draft_type: AgentConfigDraftType,
start_new: bool = False,
) -> str:
"""Resolve the current editor's Build or Preview conversation."""
"""Resolve or rotate the current editor's conversation within one draft surface.
``start_new`` rotates the scoped mapping through ``AgentRosterService`` so
the old runtime session is retired before the new conversation is used.
Continuations and Build chat keep resolving the existing mapping.
"""
roster_service = AgentRosterService(session)
resolved_agent_id = agent_id
@@ -363,26 +368,17 @@ def _resolve_current_user_agent_debug_conversation_id(
raise AgentNotFoundError()
resolved_agent_id = agent.id
if draft_type == AgentConfigDraftType.DEBUG_BUILD:
return roster_service.get_or_create_build_conversation(
tenant_id=current_tenant_id,
agent_id=resolved_agent_id,
account_id=current_user.id,
)
if start_new:
return roster_service.rotate_preview_conversation(
tenant_id=current_tenant_id,
agent_id=resolved_agent_id,
account_id=current_user.id,
)
conversation_id = roster_service.get_current_preview_conversation(
resolve_conversation = (
roster_service.refresh_agent_app_debug_conversation_id
if start_new
else roster_service.get_or_create_agent_app_debug_conversation_id
)
return resolve_conversation(
tenant_id=current_tenant_id,
agent_id=resolved_agent_id,
account_id=current_user.id,
draft_type=draft_type,
)
if conversation_id is None:
raise NotFound("Conversation Not Exists.")
return conversation_id
def _create_chat_message(
@@ -454,6 +450,7 @@ def _create_build_chat_finalization_message(
"draft_type": "debug_build",
"conversation_id": debug_conversation_id,
"auto_generate_name": False,
AGENT_RUNTIME_EXIT_INTENT_ARG: "delete",
}
external_trace_id = get_external_trace_id(request)
if external_trace_id:
+2 -14
View File
@@ -76,13 +76,11 @@ from models import Account, App
from models.model import AppMode
from models.workflow import Workflow
from repositories.workflow_collaboration_repository import WORKFLOW_ONLINE_USERS_PREFIX
from services.agent.retirement_service import WorkflowAgentRetirementService
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
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
logger = logging.getLogger(__name__)
@@ -319,7 +317,7 @@ class _WorkflowResponseSource:
self._session = session
def __getattr__(self, name: str) -> object:
return getattr(self._workflow, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self._workflow, name) # noqa: no-new-getattr response adapter delegates model fields
@property
def created_by_account(self) -> Account | None:
@@ -1247,7 +1245,7 @@ class PublishedWorkflowApi(Resource):
workflow_service = WorkflowService()
with sessionmaker(db.engine).begin() as session:
workflow, retirement_candidates = workflow_service.publish_workflow(
workflow = workflow_service.publish_workflow(
session=session,
app_model=app_model,
account=current_user,
@@ -1264,16 +1262,6 @@ class PublishedWorkflowApi(Resource):
workflow_created_at = TimestampField().format(workflow.created_at)
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=app_model.tenant_id,
agent_ids=retirement_candidates,
account_id=current_user.id,
)
enqueue_agent_resource_collection(
tenant_id=app_model.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
return {
"result": "success",
"created_at": workflow_created_at,
+2 -2
View File
@@ -218,7 +218,7 @@ class _DatasetQueryResponseSource:
return self.query.get_queries(session=self.session)
def __getattr__(self, name: str) -> Any:
return getattr(self.query, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self.query, name) # noqa: no-new-getattr response adapter delegates model fields
class DatasetQueryListResponse(ResponseModel):
@@ -257,7 +257,7 @@ class _RelatedAppResponseSource:
return self.app.mode_compatible_with_agent_with_session(session=self.session)
def __getattr__(self, name: str) -> Any:
return getattr(self.app, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self.app, name) # noqa: no-new-getattr response adapter delegates model fields
class RelatedAppListResponse(ResponseModel):
+1 -1
View File
@@ -106,7 +106,7 @@ class ExternalKnowledgeApiResponseSource:
return self.external_knowledge_api.get_dataset_bindings(session=self.session)
def __getattr__(self, name: str) -> Any:
return getattr(self.external_knowledge_api, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self.external_knowledge_api, name) # noqa: no-new-getattr response adapter delegates model fields
def external_knowledge_api_response(
+1 -1
View File
@@ -404,7 +404,7 @@ class TrialWorkflowResponseSource:
return self.workflow.get_tool_published(session=self.session)
def __getattr__(self, name: str) -> Any:
return getattr(self.workflow, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self.workflow, name) # noqa: no-new-getattr response adapter delegates model fields
register_schema_models(
@@ -56,12 +56,10 @@ from libs.helper import TimestampField
from libs.login import current_account_with_tenant, login_required
from models import Account
from models.snippet import CustomizedSnippet
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.agent.workflow_publish_service import WorkflowAgentPublishService
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
from services.snippet_generate_service import SnippetGenerateService
from services.snippet_service import SnippetService
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
logger = logging.getLogger(__name__)
@@ -297,9 +295,8 @@ class SnippetPublishedWorkflowApi(Resource):
with Session(db.engine) as session:
snippet = session.merge(snippet)
tenant_id = snippet.tenant_id
try:
workflow, retirement_candidates = snippet_service.publish_workflow(
workflow = snippet_service.publish_workflow(
session=session,
snippet=snippet,
account=current_user,
@@ -309,16 +306,6 @@ class SnippetPublishedWorkflowApi(Resource):
except ValueError as e:
return {"message": str(e)}, 400
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=tenant_id,
agent_ids=retirement_candidates,
account_id=current_user.id,
)
enqueue_agent_resource_collection(
tenant_id=tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
return {
"result": "success",
"created_at": workflow_created_at,
@@ -26,7 +26,11 @@ from libs.helper import dump_response, uuid_value
from libs.login import login_required
from models import Account
from services.billing_service import BillingService
from services.entities.model_provider_entities import ProviderResponse
from services.entities.model_provider_entities import (
ModelProviderPluginSummaryResponse,
ModelProviderSummaryResponse,
ProviderResponse,
)
from services.model_provider_service import ModelProviderService
@@ -91,6 +95,11 @@ class ModelProviderListResponse(ResponseModel):
data: list[ProviderResponse]
class ModelProviderSummaryListResponse(ResponseModel):
data: list[ModelProviderSummaryResponse]
plugins: dict[str, ModelProviderPluginSummaryResponse]
class ProviderCredentialsResponse(ResponseModel):
credentials: dict[str, Any] | None = None
@@ -114,6 +123,7 @@ register_response_schema_models(
console_ns,
SimpleResultResponse,
ModelProviderListResponse,
ModelProviderSummaryListResponse,
ProviderCredentialsResponse,
ValidationResultResponse,
ModelProviderPaymentCheckoutUrlResponse,
@@ -140,6 +150,25 @@ class ModelProviderListApi(Resource):
return ModelProviderListResponse(data=provider_list).model_dump(mode="json")
@console_ns.route("/workspaces/current/model-providers/summary")
class ModelProviderSummaryListApi(Resource):
@console_ns.response(
200,
"Model provider summaries retrieved successfully",
console_ns.models[ModelProviderSummaryListResponse.__name__],
)
@setup_required
@login_required
@account_initialization_required
@with_current_tenant_id
def get(self, tenant_id: str):
providers, plugins = ModelProviderService().get_provider_summary_list(tenant_id=tenant_id)
return dump_response(
ModelProviderSummaryListResponse,
{"data": providers, "plugins": plugins},
)
@console_ns.route("/workspaces/current/model-providers/<path:provider>/credentials")
class ModelProviderCredentialApi(Resource):
@console_ns.doc(params=query_params_from_model(ParserCredentialId))
+96 -5
View File
@@ -1,5 +1,5 @@
import io
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, Literal, TypedDict
@@ -43,6 +43,7 @@ from core.plugin.entities.plugin_daemon import PluginDecodeResponse, PluginInsta
from core.plugin.impl.exc import PluginDaemonClientSideError
from core.plugin.plugin_service import PluginService
from core.tools.builtin_tool.providers._positions import BuiltinToolProviderSort
from core.tools.entities.api_entities import ToolProviderApiEntity
from core.tools.entities.common_entities import I18nObject
from core.tools.entities.tool_entities import ToolProviderType
from core.tools.tool_manager import ToolManager
@@ -90,9 +91,21 @@ class ParserList(BaseModel):
page_size: int = Field(default=256, ge=1, le=256, description="Page size (1-256)")
type PluginCategoryListLanguage = Literal["en_US", "zh_Hans", "ja_JP", "pt_BR"]
class PluginCategoryListQuery(BaseModel):
page: int = Field(default=1, ge=1, description="Page number")
page_size: int = Field(default=256, ge=1, le=256, description="Page size (1-256)")
query: str = Field(default="", max_length=256, description="Case-insensitive search query")
tags: list[str] = Field(default_factory=list, max_length=128, description="Match any plugin tag")
language: Literal["en_US", "zh_Hans", "ja_JP", "pt_BR"] = Field(
default="en_US", description="Language used for localized label and description search"
)
class PluginInstalledIdsQuery(BaseModel):
category: PluginCategory = Field(description="Plugin category to include")
class ParserLatest(BaseModel):
@@ -325,6 +338,10 @@ class PluginListResponse(ResponseModel):
total: int
class PluginInstalledIdsResponse(ResponseModel):
plugin_ids: list[str]
class PluginVersionsResponse(ResponseModel):
versions: Mapping[str, PluginService.LatestPluginCache | None]
@@ -384,6 +401,7 @@ register_schema_models(
console_ns,
ParserList,
PluginCategoryListQuery,
PluginInstalledIdsQuery,
PluginAutoUpgradeSettingsPayload,
PluginPermissionSettingsPayload,
ParserLatest,
@@ -420,6 +438,7 @@ register_response_schema_models(
PluginDebuggingKeyResponse,
PluginDynamicOptionsResponse,
PluginInstallationsResponse,
PluginInstalledIdsResponse,
PluginInstallTaskStartResponse,
PluginListResponse,
PluginManifestResponse,
@@ -477,7 +496,39 @@ def _read_upload_content(file: FileStorage, max_size: int) -> bytes:
return content
def _list_hardcoded_builtin_tool_providers(tenant_id: str) -> list[dict[str, Any]]:
def _localized_builtin_tool_text(value: I18nObject, language: PluginCategoryListLanguage) -> str:
return value.to_dict()[language] or value.en_US
def _builtin_tool_provider_matches_filters(
provider: ToolProviderApiEntity,
*,
query: str,
tags: Sequence[str],
language: PluginCategoryListLanguage,
) -> bool:
if tags and not any(tag in provider.labels for tag in tags):
return False
if not query:
return True
lower_query = query.lower()
candidates = (
provider.name,
_localized_builtin_tool_text(provider.label, language),
_localized_builtin_tool_text(provider.description, language),
)
return any(lower_query in candidate.lower() for candidate in candidates)
def _list_hardcoded_builtin_tool_providers(
tenant_id: str,
*,
query: str = "",
tags: Sequence[str] = (),
language: PluginCategoryListLanguage = "en_US",
) -> list[dict[str, Any]]:
"""List builtin providers using the same search and tag semantics as category plugins."""
db_builtin_providers = {
str(ToolProviderID(provider.provider)): provider
for provider in ToolManager.list_default_builtin_providers(tenant_id)
@@ -498,6 +549,13 @@ def _list_hardcoded_builtin_tool_providers(tenant_id: str) -> list[dict[str, Any
db_provider=db_builtin_providers.get(provider.entity.identity.name),
decrypt_credentials=False,
)
if not _builtin_tool_provider_matches_filters(
user_provider,
query=query,
tags=tags,
language=language,
):
continue
ToolTransformService.repack_provider(tenant_id=tenant_id, provider=user_provider)
builtin_providers.append(user_provider)
@@ -552,7 +610,9 @@ class PluginCategoryListApi(Resource):
@account_initialization_required
@with_current_tenant_id
def get(self, tenant_id: str, category: str):
args = PluginCategoryListQuery.model_validate(request.args.to_dict(flat=True))
args = PluginCategoryListQuery.model_validate(
{**request.args.to_dict(flat=True), "tags": request.args.getlist("tags")}
)
try:
plugin_category = PluginCategory(category)
@@ -560,13 +620,26 @@ class PluginCategoryListApi(Resource):
return {"code": "invalid_param", "message": "invalid plugin category"}, 400
try:
plugins = PluginService.list_by_category(tenant_id, plugin_category, args.page, args.page_size)
plugins = PluginService.list_by_category(
tenant_id,
plugin_category,
args.page,
args.page_size,
query=args.query,
tags=args.tags,
language=args.language,
)
except PluginDaemonClientSideError as e:
return {"code": "plugin_error", "message": e.description}, 400
builtin_tools = []
if plugin_category == PluginCategory.Tool:
builtin_tools = _list_hardcoded_builtin_tool_providers(tenant_id)
builtin_tools = _list_hardcoded_builtin_tool_providers(
tenant_id,
query=args.query,
tags=args.tags,
language=args.language,
)
return dump_response(
PluginCategoryListResponse,
@@ -578,6 +651,24 @@ class PluginCategoryListApi(Resource):
)
@console_ns.route("/workspaces/current/plugin/installed-ids")
class PluginInstalledIdsApi(Resource):
@console_ns.doc(params=query_params_from_model(PluginInstalledIdsQuery))
@console_ns.response(200, "Success", console_ns.models[PluginInstalledIdsResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@with_current_tenant_id
def get(self, tenant_id: str):
args = PluginInstalledIdsQuery.model_validate(request.args.to_dict(flat=True))
try:
plugin_ids = PluginService.list_installed_plugin_ids(tenant_id, args.category)
except PluginDaemonClientSideError as e:
return {"code": "plugin_error", "message": e.description}, 400
return dump_response(PluginInstalledIdsResponse, {"plugin_ids": plugin_ids})
@console_ns.route("/workspaces/current/plugin/list/latest-versions")
class PluginListLatestVersionsApi(Resource):
@console_ns.expect(console_ns.models[ParserLatest.__name__])
+2
View File
@@ -23,6 +23,7 @@ from .knowledge import retrieval as _knowledge_retrieval
from .plugin import agent_config as _agent_config
from .plugin import agent_drive as _agent_drive
from .plugin import plugin as _plugin
from .workspace import plugin_model_providers as _plugin_model_providers
from .workspace import workspace as _workspace
api.add_namespace(inner_api_ns)
@@ -35,6 +36,7 @@ __all__ = [
"_knowledge_retrieval",
"_mail",
"_plugin",
"_plugin_model_providers",
"_runtime_credentials",
"_workspace",
"api",
@@ -0,0 +1,39 @@
from flask_restx import Resource
from pydantic import BaseModel, ConfigDict, Field
from controllers.common.schema import register_schema_model
from controllers.console.wraps import setup_required
from controllers.inner_api import inner_api_ns
from controllers.inner_api.wraps import enterprise_inner_api_only
from core.plugin.plugin_service import PluginService
class InvalidatePluginModelProvidersCachePayload(BaseModel):
model_config = ConfigDict(extra="forbid")
tenant_ids: list[str] = Field(default_factory=list, description="Workspace ids whose cache should be invalidated")
register_schema_model(inner_api_ns, InvalidatePluginModelProvidersCachePayload)
@inner_api_ns.route("/enterprise/workspace/plugin-model-providers/invalidate")
class EnterprisePluginModelProvidersCacheInvalidate(Resource):
@setup_required
@enterprise_inner_api_only
@inner_api_ns.doc(
"enterprise_invalidate_plugin_model_providers_cache",
responses={
200: "Cache invalidated",
400: "Invalid request",
401: "Unauthorized - invalid API key",
},
)
@inner_api_ns.expect(inner_api_ns.models[InvalidatePluginModelProvidersCachePayload.__name__])
def post(self):
args = InvalidatePluginModelProvidersCachePayload.model_validate(inner_api_ns.payload or {})
for tenant_id in args.tenant_ids:
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
return {"result": "success"}, 200
+2 -14
View File
@@ -16,7 +16,6 @@ from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import (
AdvancedChatAppGenerateEntity,
AppGenerateEntity,
DifyRunContext,
InvokeFrom,
)
from core.app.entities.queue_entities import (
@@ -32,7 +31,7 @@ from core.moderation.base import ModerationError
from core.moderation.input_moderation import InputModeration
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.node_factory import get_default_root_node_id
from core.workflow.nodes.agent_v2.workspace_retirement_layer import build_workflow_agent_workspace_retirement_layer
from core.workflow.nodes.agent_v2.session_cleanup_layer import build_workflow_agent_session_cleanup_layer
from core.workflow.system_variables import (
build_bootstrap_variables,
build_system_variables,
@@ -268,18 +267,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
)
workflow_entry.graph_engine.layer(persistence_layer)
workflow_entry.graph_engine.layer(
build_workflow_agent_workspace_retirement_layer(
dify_run_context=DifyRunContext(
tenant_id=self._workflow.tenant_id,
app_id=self._workflow.app_id,
user_id=self.application_generate_entity.user_id,
user_from=user_from,
invoke_from=invoke_from,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
)
)
workflow_entry.graph_engine.layer(build_workflow_agent_session_cleanup_layer())
conversation_variable_layer = ConversationVariablePersistenceLayer(
ConversationVariableUpdater(session_factory.get_session_maker())
)
+65 -135
View File
@@ -33,13 +33,15 @@ from core.app.apps.agent_app.app_runner import AgentAppRunner
from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError
from core.app.apps.agent_app.generate_response_converter import AgentAppGenerateResponseConverter
from core.app.apps.agent_app.runtime_request_builder import AgentAppRuntimeRequestBuilder
from core.app.apps.agent_app.session_store import AgentAppWorkspaceStore
from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import (
AGENT_RUNTIME_EXIT_INTENT_ARG,
AgentAppGenerateEntity,
AgentRuntimeExitIntent,
DifyRunContext,
InvokeFrom,
UserFrom,
@@ -49,24 +51,19 @@ from core.db.session_factory import session_factory
from core.ops.ops_trace_manager import TraceQueueManager
from core.workflow.file_reference import build_file_reference, is_canonical_file_reference
from extensions.ext_database import db
from models import Account, App, AppModelConfig, Conversation, EndUser, Message, MessageAnnotation
from models import Account, App, AppModelConfig, EndUser, Message, MessageAnnotation
from models.agent import (
APP_BACKED_AGENT_SOURCES,
Agent,
AgentConfigDraft,
AgentConfigDraftType,
AgentConfigSnapshot,
AgentConfigVersionKind,
AgentScope,
AgentSource,
AgentStatus,
AgentWorkingResourceStatus,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
)
from models.agent_config_entities import AgentSoulConfig
from models.model import load_annotation_reply_config
from services.agent.workspace_service import AgentWorkspaceService, WorkspaceOwnerScope
from services.conversation_service import ConversationService
logger = logging.getLogger(__name__)
@@ -153,28 +150,26 @@ class AgentAppGenerator(MessageBasedAppGenerator):
inputs = args["inputs"]
prompt_file_mappings = args.get("files") or []
conversation = None
conversation_id = args.get("conversation_id")
if conversation_id:
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
# New conversations use the current Agent generation. Existing
# conversations use the immutable generation named by their Binding.
# Resolve the bound roster Agent + its current Agent Soul snapshot.
agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent(
app_model,
invoke_from=invoke_from,
draft_type=args.get("draft_type"),
user=user,
session=session,
conversation=conversation,
)
session_scope_config_version_id = self._session_scope_config_version_id(
runtime_session_snapshot_id = self._runtime_session_snapshot_id(
invoke_from=invoke_from,
config_version_id=agent_config_id,
snapshot_id=agent_config_id,
)
conversation = None
conversation_id = args.get("conversation_id")
if conversation_id:
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
# Build the EasyUI-shaped config from the Agent Soul so the chat pipeline
# can persist usage; the answer itself comes from the agent backend.
app_model_config = (
@@ -191,6 +186,8 @@ class AgentAppGenerator(MessageBasedAppGenerator):
model_conf = ModelConfigConverter.convert(app_config)
trace_manager = TraceQueueManager(app_model.id, user.id if isinstance(user, Account) else user.session_id)
agent_runtime_exit_intent = self._resolve_agent_runtime_exit_intent(args)
application_generate_entity = AgentAppGenerateEntity(
task_id=str(uuid.uuid4()),
app_config=app_config,
@@ -218,7 +215,8 @@ class AgentAppGenerator(MessageBasedAppGenerator):
agent_id=agent.id,
agent_config_snapshot_id=agent_config_id,
agent_config_version_kind=agent_config_version_kind,
agent_session_scope_config_version_id=session_scope_config_version_id,
agent_runtime_session_snapshot_id=runtime_session_snapshot_id,
agent_runtime_exit_intent=agent_runtime_exit_intent,
)
conversation, message = self._init_generate_records(
@@ -267,7 +265,6 @@ class AgentAppGenerator(MessageBasedAppGenerator):
app_model: App,
user: Account | EndUser,
conversation_id: str,
form_id: str,
invoke_from: InvokeFrom,
session: Session,
) -> None:
@@ -282,21 +279,14 @@ class AgentAppGenerator(MessageBasedAppGenerator):
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
draft_type, draft_id = self._resolve_resume_draft(
app_model=app_model,
conversation=conversation,
user=user,
form_id=form_id,
session=session,
)
agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent(
app_model,
invoke_from=invoke_from,
draft_type=draft_type,
draft_id=draft_id,
draft_type=self._resume_draft_type(
app_model=app_model, conversation=conversation, user=user, session=session
),
user=user,
session=session,
conversation=conversation,
)
app_model_config = (
@@ -394,37 +384,30 @@ class AgentAppGenerator(MessageBasedAppGenerator):
)
@staticmethod
def _resolve_resume_draft(
*,
app_model: App,
conversation: Any,
user: Account | EndUser,
form_id: str,
session: Session,
) -> tuple[str | None, str | None]:
def _resume_draft_type(
*, app_model: App, conversation: Any, user: Account | EndUser, session: Session
) -> str | None:
if conversation.invoke_from != InvokeFrom.DEBUGGER:
return None, None
if not isinstance(user, Account):
return AgentConfigDraftType.DRAFT.value, None
build_draft = session.scalar(
select(AgentConfigDraft)
.join(
AgentWorkspaceBinding,
AgentWorkspaceBinding.id == AgentConfigDraft.agent_workspace_binding_id,
)
.where(
AgentConfigDraft.tenant_id == app_model.tenant_id,
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
AgentConfigDraft.account_id == user.id,
AgentWorkspaceBinding.tenant_id == app_model.tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
AgentWorkspaceBinding.pending_form_id == form_id,
)
return None
active_session = AgentAppRuntimeSessionStore().load_active_session_for_conversation(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
conversation_id=conversation.id,
)
if build_draft is not None:
return AgentConfigDraftType.DEBUG_BUILD.value, build_draft.id
return AgentConfigDraftType.DRAFT.value, None
snapshot_id = active_session.scope.agent_config_snapshot_id if active_session is not None else None
if snapshot_id and isinstance(user, Account):
draft = session.scalar(
select(AgentConfigDraft).where(
AgentConfigDraft.tenant_id == app_model.tenant_id,
AgentConfigDraft.id == snapshot_id,
)
)
if draft is not None:
if draft.draft_type == AgentConfigDraftType.DEBUG_BUILD and draft.account_id == user.id:
return AgentConfigDraftType.DEBUG_BUILD.value
if draft.draft_type == AgentConfigDraftType.DRAFT and draft.account_id is None:
return AgentConfigDraftType.DRAFT.value
return AgentConfigDraftType.DRAFT.value
def _generate_worker(
self,
@@ -499,7 +482,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
invoke_from=application_generate_entity.invoke_from,
)
with session_factory.create_session() as session:
agent, config_version, agent_soul = self._resolve_agent_by_id(
_, _, agent_soul = self._resolve_agent_by_id(
tenant_id=app_config.tenant_id,
agent_id=application_generate_entity.agent_id,
snapshot_id=application_generate_entity.agent_config_snapshot_id,
@@ -513,18 +496,13 @@ class AgentAppGenerator(MessageBasedAppGenerator):
agent_config_snapshot_id=application_generate_entity.agent_config_snapshot_id,
agent_config_version_kind=application_generate_entity.agent_config_version_kind,
agent_soul=agent_soul,
home_snapshot_id=config_version.home_snapshot_id,
conversation_id=conversation.id,
query=query,
message_id=message.id,
model_name=application_generate_entity.model_conf.model,
queue_manager=queue_manager,
session_scope_snapshot_id=application_generate_entity.agent_session_scope_config_version_id,
build_draft_id=(
application_generate_entity.agent_config_snapshot_id
if application_generate_entity.agent_config_version_kind == AgentConfigVersionKind.BUILD_DRAFT
else None
),
session_scope_snapshot_id=application_generate_entity.agent_runtime_session_snapshot_id,
agent_runtime_exit_intent=application_generate_entity.agent_runtime_exit_intent,
)
except GenerateTaskStoppedError:
pass
@@ -541,6 +519,18 @@ class AgentAppGenerator(MessageBasedAppGenerator):
raise AgentAppGeneratorError("query is required")
return query.replace("\x00", "")
@staticmethod
def _resolve_agent_runtime_exit_intent(args: Mapping[str, Any]) -> AgentRuntimeExitIntent:
"""Resolve API-internal runtime exit policy from controller-owned args.
Only the private controller-injected "delete" value changes behavior.
Normal chat and resume flows default/fallback to "suspend" so public
payloads and invalid internal values preserve existing semantics.
"""
if args.get(AGENT_RUNTIME_EXIT_INTENT_ARG) == "delete":
return "delete"
return "suspend"
@staticmethod
def _build_runner(dify_context: DifyRunContext) -> AgentAppRunner:
credentials_provider, _ = build_dify_model_access(dify_context)
@@ -548,7 +538,6 @@ class AgentAppGenerator(MessageBasedAppGenerator):
request_builder=AgentAppRuntimeRequestBuilder(credentials_provider=credentials_provider),
agent_backend_client=create_agent_backend_run_client(
base_url=dify_config.AGENT_BACKEND_BASE_URL,
api_token=dify_config.AGENT_BACKEND_API_TOKEN,
use_fake=dify_config.AGENT_BACKEND_USE_FAKE,
fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO,
stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS,
@@ -556,7 +545,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS,
),
event_adapter=AgentBackendRunEventAdapter(),
session_store=AgentAppWorkspaceStore(),
session_store=AgentAppRuntimeSessionStore(),
text_delta_debounce_seconds=dify_config.AGENT_APP_TEXT_DELTA_DEBOUNCE_SECONDS,
)
@@ -620,10 +609,8 @@ class AgentAppGenerator(MessageBasedAppGenerator):
*,
invoke_from: InvokeFrom,
draft_type: Any,
draft_id: str | None = None,
user: Account | EndUser,
session: Session,
conversation: Conversation | None = None,
) -> tuple[Agent, str, Literal["snapshot", "draft", "build_draft"], AgentSoulConfig]:
agent = session.scalar(
select(Agent)
@@ -655,7 +642,6 @@ class AgentAppGenerator(MessageBasedAppGenerator):
tenant_id=app_model.tenant_id,
agent=agent,
draft_type=draft_type,
draft_id=draft_id,
account_id=user.id if isinstance(user, Account) else None,
session=session,
)
@@ -668,82 +654,28 @@ class AgentAppGenerator(MessageBasedAppGenerator):
# Public runtime must keep serving the active snapshot even when unpublished draft edits exist.
if not agent.active_config_snapshot_id:
raise AgentAppNotPublishedError("Agent has not been published")
conversation_binding = self._resolve_conversation_binding(
session=session,
tenant_id=app_model.tenant_id,
app_id=app_model.id,
agent_id=agent.id,
conversation=conversation,
)
snapshot_id = (
conversation_binding.agent_config_version_id
if conversation_binding is not None
else agent.active_config_snapshot_id
)
_, snapshot, agent_soul = self._resolve_agent_by_id(
tenant_id=app_model.tenant_id,
agent_id=agent.id,
snapshot_id=snapshot_id,
snapshot_id=agent.active_config_snapshot_id,
session=session,
)
if conversation_binding is not None:
AgentWorkspaceService.validate_binding_generation(
conversation_binding,
base_home_snapshot_id=snapshot.home_snapshot_id,
agent_config_version_id=snapshot.id,
agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT,
)
return agent, snapshot.id, "snapshot", agent_soul
@staticmethod
def _resolve_conversation_binding(
*,
session: Session,
tenant_id: str,
app_id: str,
agent_id: str,
conversation: Conversation | None,
) -> AgentWorkspaceBinding | None:
"""Resolve the exact participant generation owned by an existing conversation."""
if conversation is None or conversation.agent_workspace_binding_id is None:
return None
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=tenant_id,
binding_id=conversation.agent_workspace_binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=app_id,
owner_type=AgentWorkspaceOwnerType.CONVERSATION,
owner_id=conversation.id,
),
)
if binding is None or binding.agent_id != agent_id:
raise AgentAppGeneratorError("Conversation participant Binding is unavailable")
return binding
@staticmethod
def _session_scope_config_version_id(*, invoke_from: InvokeFrom, config_version_id: str) -> str | None:
"""Return the config version id that scopes Agent App session reuse.
def _runtime_session_snapshot_id(*, invoke_from: InvokeFrom, snapshot_id: str) -> str | None:
"""Return the session scope snapshot id for Agent App runtime state.
Console preview/debug chat uses a stable Agent draft row id; build mode
uses the current user's build-draft row id. Published/web/API runs use
immutable published snapshot ids. This keeps Workspace Binding continuity
immutable published snapshot ids. This keeps runtime session continuity
inside one editable surface without mixing draft/build/published state.
"""
del invoke_from
return config_version_id
return snapshot_id
@staticmethod
def _resolve_debug_draft(
*,
tenant_id: str,
agent: Agent,
draft_type: Any,
account_id: str | None,
session: Session,
draft_id: str | None = None,
*, tenant_id: str, agent: Agent, draft_type: Any, account_id: str | None, session: Session
) -> AgentConfigDraft:
effective_draft_type = (
AgentConfigDraftType.DEBUG_BUILD
@@ -767,8 +699,6 @@ class AgentAppGenerator(MessageBasedAppGenerator):
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
AgentConfigDraft.account_id == account_id,
)
if draft_id is not None:
stmt = stmt.where(AgentConfigDraft.id == draft_id)
draft = session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1))
if draft is not None:
return draft
+143 -45
View File
@@ -3,7 +3,8 @@
Unlike the legacy ``AgentChatAppRunner`` (which runs an in-process ReAct loop),
this runner delegates to the Agent backend, consumes the streamed event flow,
republishes the assistant answer through the existing EasyUI chat task
pipeline, and saves the latest Agenton snapshot on the persistent Binding.
pipeline, and then either saves or retires the conversation-owned runtime
session depending on the turn's exit policy.
"""
from __future__ import annotations
@@ -30,20 +31,22 @@ from clients.agent_backend import (
AgentBackendRunFailedInternalEvent,
AgentBackendRunSucceededInternalEvent,
AgentBackendStreamInternalEvent,
extract_runtime_layer_specs,
)
from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload
from core.app.apps.agent_app.runtime_request_builder import (
AgentAppRuntimeBuildContext,
AgentAppRuntimeRequest,
AgentAppRuntimeRequestBuilder,
)
from core.app.apps.agent_app.session_store import (
AgentAppRuntimeSessionStore,
AgentAppSessionScope,
AgentAppWorkspaceStore,
StoredAgentAppSession,
)
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.entities.app_invoke_entities import DifyRunContext
from core.app.entities.app_invoke_entities import AgentRuntimeExitIntent, DifyRunContext
from core.app.entities.queue_entities import (
QueueAgentMessageEvent,
QueueAgentThoughtEvent,
@@ -64,10 +67,10 @@ from graphon.model_runtime.errors.invoke import (
InvokeRateLimitError,
InvokeServerUnavailableError,
)
from models.agent import AgentConfigVersionKind
from models.agent_config_entities import AgentSoulConfig
from models.enums import CreatorUserRole
from models.model import MessageAgentThought
from tasks.agent_backend_session_cleanup_task import cleanup_conversation_agent_runtime_session
logger = logging.getLogger(__name__)
@@ -617,7 +620,7 @@ class AgentAppRunner:
request_builder: AgentAppRuntimeRequestBuilder,
agent_backend_client: AgentBackendRunClient,
event_adapter: AgentBackendRunEventAdapter,
session_store: AgentAppWorkspaceStore,
session_store: AgentAppRuntimeSessionStore,
text_delta_debounce_seconds: float,
) -> None:
self._request_builder = request_builder
@@ -634,41 +637,37 @@ class AgentAppRunner:
agent_config_snapshot_id: str,
agent_config_version_kind: Literal["snapshot", "draft", "build_draft"] = "snapshot",
agent_soul: AgentSoulConfig,
home_snapshot_id: str,
conversation_id: str,
query: str,
message_id: str,
model_name: str,
queue_manager: AppQueueManager,
session_scope_snapshot_id: str | None | _DefaultSessionScopeSnapshotId = _DEFAULT_SESSION_SCOPE_SNAPSHOT_ID,
build_draft_id: str | None = None,
agent_runtime_exit_intent: AgentRuntimeExitIntent = "suspend",
) -> None:
preserve_session = agent_runtime_exit_intent == "suspend"
scope = self._build_session_scope(
dify_context=dify_context,
agent_id=agent_id,
agent_config_snapshot_id=agent_config_snapshot_id,
home_snapshot_id=home_snapshot_id,
conversation_id=conversation_id,
session_scope_snapshot_id=session_scope_snapshot_id,
agent_config_version_kind=AgentConfigVersionKind(agent_config_version_kind),
build_draft_id=build_draft_id,
)
# ENG-638: if a prior turn paused on ask_human and the form is now answered,
# resume by threading the human's reply into this run as deferred_tool_results.
stored = self._session_store.load_or_create(scope)
stored = self._session_store.load_active_session(scope)
runtime = self._build_runtime(
dify_context=dify_context,
agent_id=agent_id,
agent_config_snapshot_id=agent_config_snapshot_id,
agent_config_version_kind=agent_config_version_kind,
agent_soul=agent_soul,
binding_id=stored.binding_id,
backend_binding_ref=stored.backend_binding_ref,
conversation_id=conversation_id,
query=query,
idempotency_key=message_id,
stored=stored,
message_id=message_id,
suspend_on_exit=preserve_session,
)
create_response = self._agent_backend_client.create_run(runtime.request)
@@ -682,6 +681,9 @@ class AgentAppRunner:
)
if isinstance(terminal, AgentBackendDeferredToolCallInternalEvent):
if not preserve_session:
self._mark_session_cleaned(scope=scope, backend_run_id=terminal.run_id)
raise AgentBackendError("Agent App finalization cannot pause for human input.")
# ENG-635: the agent asked a human. End this turn with the question and
# a conversation-owned HITL form; a form submission resumes the run.
self._pause_for_ask_human(
@@ -701,8 +703,8 @@ class AgentAppRunner:
if not isinstance(terminal, AgentBackendRunSucceededInternalEvent):
if isinstance(terminal, AgentBackendRunFailedInternalEvent):
reason = terminal.reason
if reason == "binding_lost":
raise AgentBackendError("The retained agent working environment is no longer available.")
if reason == "sandbox_expired":
raise AgentBackendError("The agent session sandbox has expired. Please start a new conversation.")
raise _agent_backend_failure_to_exception(terminal)
raise AgentBackendError("Agent backend run did not complete successfully.")
@@ -717,18 +719,38 @@ class AgentAppRunner:
message_id,
exc_info=True,
)
self._publish_terminal_answer(
queue_manager=queue_manager,
model_name=model_name,
answer=answer,
query=query,
usage=_llm_usage_from_agent_backend(terminal.usage),
)
self._save_session(
scope=scope,
binding_id=runtime.binding_id,
snapshot=terminal.session_snapshot,
)
if preserve_session:
superseded_sessions = self._load_superseded_sessions(scope=scope)
self._publish_terminal_answer(
queue_manager=queue_manager,
model_name=model_name,
answer=answer,
query=query,
usage=_llm_usage_from_agent_backend(terminal.usage),
)
session_saved = self._save_session(
scope=scope,
backend_run_id=terminal.run_id,
snapshot=terminal.session_snapshot,
runtime_layer_specs=extract_runtime_layer_specs(runtime.request.composition),
)
if session_saved:
self._cleanup_superseded_sessions(superseded_sessions)
else:
# The backend has already accepted a terminal success with
# delete-on-exit semantics. Local publish/persistence errors must
# not keep the API-side session row active, and cleanup failures
# must not replace the original publish/error outcome.
try:
self._publish_terminal_answer(
queue_manager=queue_manager,
model_name=model_name,
answer=answer,
query=query,
usage=_llm_usage_from_agent_backend(terminal.usage),
)
finally:
self._mark_session_cleaned(scope=scope, backend_run_id=terminal.run_id)
def _build_session_scope(
self,
@@ -736,11 +758,8 @@ class AgentAppRunner:
dify_context: DifyRunContext,
agent_id: str,
agent_config_snapshot_id: str,
home_snapshot_id: str,
conversation_id: str,
session_scope_snapshot_id: str | None | _DefaultSessionScopeSnapshotId,
agent_config_version_kind: AgentConfigVersionKind,
build_draft_id: str | None = None,
) -> AgentAppSessionScope:
if isinstance(session_scope_snapshot_id, _DefaultSessionScopeSnapshotId):
effective_session_scope_snapshot_id: str | None = agent_config_snapshot_id
@@ -751,10 +770,7 @@ class AgentAppRunner:
app_id=dify_context.app_id,
conversation_id=conversation_id,
agent_id=agent_id,
agent_config_snapshot_id=effective_session_scope_snapshot_id or agent_config_snapshot_id,
home_snapshot_id=home_snapshot_id,
agent_config_version_kind=agent_config_version_kind,
build_draft_id=build_draft_id,
agent_config_snapshot_id=effective_session_scope_snapshot_id,
)
def _build_runtime(
@@ -765,15 +781,14 @@ class AgentAppRunner:
agent_config_snapshot_id: str,
agent_config_version_kind: Literal["snapshot", "draft", "build_draft"],
agent_soul: AgentSoulConfig,
binding_id: str,
backend_binding_ref: str,
conversation_id: str,
query: str,
idempotency_key: str,
stored: StoredAgentAppSession,
stored: StoredAgentAppSession | None,
message_id: str | None,
suspend_on_exit: bool,
) -> AgentAppRuntimeRequest:
session_snapshot = stored.session_snapshot
session_snapshot = stored.session_snapshot if stored is not None else None
deferred_tool_results = (
self._resolve_pending_ask_human(stored=stored, dify_context=dify_context, message_id=message_id)
if message_id is not None
@@ -789,10 +804,9 @@ class AgentAppRunner:
conversation_id=conversation_id,
user_query=query,
idempotency_key=idempotency_key,
binding_id=binding_id,
backend_binding_ref=backend_binding_ref,
session_snapshot=session_snapshot,
deferred_tool_results=deferred_tool_results,
suspend_on_exit=suspend_on_exit,
)
)
@@ -829,8 +843,9 @@ class AgentAppRunner:
# second run with the human's answer (ENG-637/638 columns, conversation owner).
self._save_session(
scope=scope,
binding_id=runtime.binding_id,
backend_run_id=terminal.run_id,
snapshot=terminal.session_snapshot,
runtime_layer_specs=extract_runtime_layer_specs(runtime.request.composition),
pending_form_id=created.form_id,
pending_tool_call_id=terminal.deferred_tool_call.tool_call_id,
)
@@ -847,12 +862,12 @@ class AgentAppRunner:
def _resolve_pending_ask_human(
self,
*,
stored: StoredAgentAppSession,
stored: StoredAgentAppSession | None,
dify_context: DifyRunContext,
message_id: str,
) -> DeferredToolResultsPayload | None:
"""Build deferred_tool_results when a pending ask_human form is answered."""
if stored.pending_form_id is None or stored.pending_tool_call_id is None:
if stored is None or stored.pending_form_id is None or stored.pending_tool_call_id is None:
return None
outcome = resolve_ask_human_form(
form_id=stored.pending_form_id,
@@ -1021,16 +1036,18 @@ class AgentAppRunner:
self,
*,
scope: AgentAppSessionScope,
binding_id: str,
backend_run_id: str,
snapshot: Any,
runtime_layer_specs: Any,
pending_form_id: str | None = None,
pending_tool_call_id: str | None = None,
) -> bool:
try:
self._session_store.save_active_snapshot(
scope=scope,
binding_id=binding_id,
backend_run_id=backend_run_id,
snapshot=snapshot,
runtime_layer_specs=runtime_layer_specs,
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
@@ -1047,6 +1064,87 @@ class AgentAppRunner:
)
return False
def _load_superseded_sessions(self, *, scope: AgentAppSessionScope) -> list[StoredAgentAppSession]:
try:
stored_sessions = self._session_store.list_active_sessions_for_conversation(
tenant_id=scope.tenant_id,
app_id=scope.app_id,
conversation_id=scope.conversation_id,
)
except Exception:
logger.warning(
"Failed to load existing Agent App conversation sessions before snapshot save: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s",
scope.tenant_id,
scope.app_id,
scope.conversation_id,
scope.agent_id,
exc_info=True,
)
return []
return [stored for stored in stored_sessions if stored.scope != scope]
def _cleanup_superseded_sessions(self, stored_sessions: list[StoredAgentAppSession]) -> None:
for stored_session in stored_sessions:
try:
if stored_session.runtime_layer_specs:
payload = AgentBackendSessionCleanupPayload(
session_snapshot=stored_session.session_snapshot,
runtime_layer_specs=stored_session.runtime_layer_specs,
idempotency_key=(
f"{stored_session.scope.tenant_id}:{stored_session.scope.app_id}:"
f"{stored_session.scope.conversation_id}:{stored_session.scope.agent_id}:"
f"{stored_session.scope.agent_config_snapshot_id or 'no-config'}:"
f"superseded-session-cleanup:{stored_session.backend_run_id or 'no-run'}"
),
metadata={
"tenant_id": stored_session.scope.tenant_id,
"app_id": stored_session.scope.app_id,
"conversation_id": stored_session.scope.conversation_id,
"agent_id": stored_session.scope.agent_id,
"agent_config_snapshot_id": stored_session.scope.agent_config_snapshot_id,
"previous_agent_backend_run_id": stored_session.backend_run_id,
},
)
cleanup_conversation_agent_runtime_session.delay(payload.model_dump(mode="json"))
except Exception:
logger.warning(
"Failed to enqueue Agent backend cleanup for superseded Agent App session: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s",
stored_session.scope.tenant_id,
stored_session.scope.app_id,
stored_session.scope.conversation_id,
stored_session.scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
def _mark_session_cleaned(
self,
*,
scope: AgentAppSessionScope,
backend_run_id: str,
) -> None:
"""Best-effort delete-on-exit cleanup for the API-side session row.
Once the Agent backend reaches a terminal event, cleanup persistence
must not replace the original publish/error outcome for that turn.
"""
try:
self._session_store.mark_cleaned(scope=scope, backend_run_id=backend_run_id)
except Exception:
logger.warning(
"Failed to retire Agent App conversation session after delete-on-exit: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s",
scope.tenant_id,
scope.app_id,
scope.conversation_id,
scope.agent_id,
backend_run_id,
exc_info=True,
)
@staticmethod
def _terminal_output_to_answer(output: JsonValue) -> str:
"""Normalize the backend's terminal output to assistant text.
@@ -70,12 +70,11 @@ class AgentAppRuntimeBuildContext:
conversation_id: str
user_query: str
idempotency_key: str
binding_id: str
backend_binding_ref: str
agent_config_version_kind: Literal["snapshot", "draft", "build_draft"] = "snapshot"
session_snapshot: CompositorSessionSnapshot | None = None
# ENG-638: set when resuming a chat turn after a submitted ask_human form.
deferred_tool_results: DeferredToolResultsPayload | None = None
suspend_on_exit: bool = True
@dataclass(frozen=True, slots=True)
@@ -83,7 +82,6 @@ class AgentAppRuntimeRequest:
request: CreateRunRequest
redacted_request: dict[str, Any]
metadata: dict[str, Any]
binding_id: str
class AgentAppRuntimeRequestBuilder:
@@ -162,7 +160,6 @@ class AgentAppRuntimeRequestBuilder:
invoke_from=cast(DifyExecutionContextInvokeFrom, context.dify_context.invoke_from.value),
agent_mode="agent_app",
),
backend_binding_ref=context.backend_binding_ref,
# ENG-616: expand slash-menu mention tokens to canonical names so
# no frontend-internal {{#…#}} marker ever reaches the model.
agent_soul_prompt=expand_prompt_mentions(agent_soul.prompt.system_prompt, soul_prompt_resolver).strip()
@@ -178,17 +175,13 @@ class AgentAppRuntimeRequestBuilder:
shell_config=build_shell_layer_config(agent_soul),
session_snapshot=context.session_snapshot,
deferred_tool_results=context.deferred_tool_results,
suspend_on_exit=context.suspend_on_exit,
idempotency_key=context.idempotency_key,
metadata=metadata,
)
)
redacted = cast(dict[str, Any], redact_for_agent_backend_log(request))
return AgentAppRuntimeRequest(
request=request,
redacted_request=redacted,
metadata=metadata,
binding_id=context.binding_id,
)
return AgentAppRuntimeRequest(request=request, redacted_request=redacted, metadata=metadata)
def _build_tool_layers(
self,
+204 -125
View File
@@ -1,176 +1,255 @@
"""Persist and resolve the exact participant owned by an Agent App caller."""
"""Conversation-keyed Agent backend session store for the Agent App type.
Shares the unified ``agent_runtime_sessions`` table with the workflow Agent
Node store, but owns rows with ``owner_type = conversation``: one Agent App
conversation maps to one Agent session, so multi-turn chat re-enters the same
``session_snapshot``. Cross-conversation memory (PRD Global / Per app) is a
phase-2 concern and not modeled here.
"""
from __future__ import annotations
from dataclasses import dataclass
from dataclasses import dataclass, field
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.protocol import RuntimeLayerSpec
from pydantic import TypeAdapter
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
from libs.datetime_utils import naive_utc_now
from models.agent import (
AgentConfigDraft,
AgentConfigDraftType,
AgentConfigVersionKind,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
)
from models.model import App, Conversation
from services.agent.workspace_service import (
AgentWorkspaceNotFoundError,
AgentWorkspaceService,
WorkspaceOwnerScope,
AgentRuntimeSession,
AgentRuntimeSessionOwnerType,
AgentRuntimeSessionStatus,
)
_RUNTIME_LAYER_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec])
def _serialize_runtime_layer_specs(specs: list[RuntimeLayerSpec]) -> str:
return _RUNTIME_LAYER_SPECS_ADAPTER.dump_json(specs).decode()
def _deserialize_runtime_layer_specs(value: str | None) -> list[RuntimeLayerSpec]:
if not value:
return []
return _RUNTIME_LAYER_SPECS_ADAPTER.validate_json(value)
@dataclass(frozen=True, slots=True)
class AgentAppSessionScope:
"""Identity of one Agent App conversation session."""
tenant_id: str
app_id: str
conversation_id: str
agent_id: str
agent_config_snapshot_id: str
home_snapshot_id: str
agent_config_version_kind: AgentConfigVersionKind = AgentConfigVersionKind.SNAPSHOT
build_draft_id: str | None = None
@property
def workspace_owner(self) -> WorkspaceOwnerScope:
owner_type = (
AgentWorkspaceOwnerType.BUILD_DRAFT if self.build_draft_id else AgentWorkspaceOwnerType.CONVERSATION
)
return WorkspaceOwnerScope(
tenant_id=self.tenant_id,
app_id=self.app_id,
owner_type=owner_type,
owner_id=self.build_draft_id or self.conversation_id,
)
agent_config_snapshot_id: str | None
@dataclass(frozen=True, slots=True)
class StoredAgentAppSession:
"""Persisted Agent App conversation session with reusable runtime specs."""
scope: AgentAppSessionScope
binding_id: str
workspace_id: str
backend_binding_ref: str
session_snapshot: CompositorSessionSnapshot | None
session_snapshot: CompositorSessionSnapshot
backend_run_id: str | None
runtime_layer_specs: list[RuntimeLayerSpec] = field(default_factory=list)
# ENG-635: set while the conversation turn is paused on a dify.ask_human
# deferred call, awaiting a HITL form submission.
pending_form_id: str | None = None
pending_tool_call_id: str | None = None
class AgentAppWorkspaceStore:
"""Resolve Agent App sessions through a caller-owned Binding pointer."""
class AgentAppRuntimeSessionStore:
"""Persists Agent backend session snapshots for Agent App conversations."""
def load_or_create(self, scope: AgentAppSessionScope) -> StoredAgentAppSession:
def load_active_snapshot(self, scope: AgentAppSessionScope) -> CompositorSessionSnapshot | None:
stored = self.load_active_session(scope)
return stored.session_snapshot if stored is not None else None
def load_active_session(self, scope: AgentAppSessionScope) -> StoredAgentAppSession | None:
with session_factory.create_session() as session:
caller = self._load_caller(session=session, scope=scope)
binding_id = caller.agent_workspace_binding_id
if binding_id is None:
binding = AgentWorkspaceService.create_binding(
session=session,
scope=scope.workspace_owner,
agent_id=scope.agent_id,
base_home_snapshot_id=scope.home_snapshot_id,
agent_config_version_id=scope.agent_config_snapshot_id,
agent_config_version_kind=scope.agent_config_version_kind,
)
caller.agent_workspace_binding_id = binding.id
session.commit()
else:
binding = self._get_binding(session=session, scope=scope, binding_id=binding_id)
return self._stored(scope, binding)
@staticmethod
def _load_caller(*, session: Session, scope: AgentAppSessionScope) -> Conversation | AgentConfigDraft:
if scope.build_draft_id is not None:
if scope.agent_config_version_kind != AgentConfigVersionKind.BUILD_DRAFT:
raise AgentWorkspaceNotFoundError("Build Draft caller requires build_draft generation")
draft = session.scalar(
select(AgentConfigDraft).where(
AgentConfigDraft.id == scope.build_draft_id,
AgentConfigDraft.tenant_id == scope.tenant_id,
AgentConfigDraft.agent_id == scope.agent_id,
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
)
row = session.scalar(self._active_stmt(scope))
if row is None:
return None
return StoredAgentAppSession(
scope=scope,
session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
backend_run_id=row.backend_run_id,
runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs),
pending_form_id=row.pending_form_id,
pending_tool_call_id=row.pending_tool_call_id,
)
if draft is None:
raise AgentWorkspaceNotFoundError("Build Draft caller is unavailable")
return draft
if scope.agent_config_version_kind == AgentConfigVersionKind.BUILD_DRAFT:
raise AgentWorkspaceNotFoundError("Build Draft caller ID is required")
conversation = session.scalar(
select(Conversation)
.join(App, App.id == Conversation.app_id)
def load_active_session_for_conversation(
self, *, tenant_id: str, app_id: str, conversation_id: str
) -> StoredAgentAppSession | None:
"""Load the latest ACTIVE session for one conversation-level sandbox lookup.
Sandbox inspection only knows the product locator
``tenant_id + app_id + conversation_id``; it does not know which
``agent_id`` or Agent Soul snapshot produced the active shell session.
This method therefore resolves the newest ACTIVE conversation-owned row
for that conversation and returns both the resumable snapshot and the
persisted non-sensitive runtime layer specs needed to build a
``SandboxLocator``.
"""
stmt = (
select(AgentRuntimeSession)
.where(
App.tenant_id == scope.tenant_id,
Conversation.id == scope.conversation_id,
Conversation.app_id == scope.app_id,
Conversation.is_deleted.is_(False),
AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION,
AgentRuntimeSession.tenant_id == tenant_id,
AgentRuntimeSession.app_id == app_id,
AgentRuntimeSession.conversation_id == conversation_id,
AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE,
)
.order_by(AgentRuntimeSession.updated_at.desc())
)
if conversation is None:
raise AgentWorkspaceNotFoundError("Conversation caller is unavailable")
return conversation
with session_factory.create_session() as session:
row = session.scalar(stmt)
if row is None:
return None
return StoredAgentAppSession(
scope=AgentAppSessionScope(
tenant_id=row.tenant_id,
app_id=row.app_id,
conversation_id=row.conversation_id or "",
agent_id=row.agent_id,
agent_config_snapshot_id=row.agent_config_snapshot_id or "",
),
session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
backend_run_id=row.backend_run_id,
runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs),
)
@staticmethod
def _get_binding(
*,
session: Session,
scope: AgentAppSessionScope,
binding_id: str,
) -> AgentWorkspaceBinding:
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=scope.tenant_id,
binding_id=binding_id,
expected_owner_scope=scope.workspace_owner,
def list_active_sessions_for_conversation(
self, *, tenant_id: str, app_id: str, conversation_id: str
) -> list[StoredAgentAppSession]:
"""List all ACTIVE conversation-owned sessions for lifecycle cleanup."""
stmt = (
select(AgentRuntimeSession)
.where(
AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION,
AgentRuntimeSession.tenant_id == tenant_id,
AgentRuntimeSession.app_id == app_id,
AgentRuntimeSession.conversation_id == conversation_id,
AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE,
)
.order_by(AgentRuntimeSession.updated_at.desc())
)
if binding is None or binding.agent_id != scope.agent_id:
raise AgentWorkspaceNotFoundError("Caller participant Binding is unavailable")
AgentWorkspaceService.validate_binding_generation(
binding,
base_home_snapshot_id=scope.home_snapshot_id,
agent_config_version_id=scope.agent_config_snapshot_id,
agent_config_version_kind=scope.agent_config_version_kind,
)
return binding
with session_factory.create_session() as session:
rows = session.scalars(stmt).all()
return [
StoredAgentAppSession(
scope=AgentAppSessionScope(
tenant_id=row.tenant_id,
app_id=row.app_id,
conversation_id=row.conversation_id or "",
agent_id=row.agent_id,
agent_config_snapshot_id=row.agent_config_snapshot_id,
),
session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
backend_run_id=row.backend_run_id,
runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs),
pending_form_id=row.pending_form_id,
pending_tool_call_id=row.pending_tool_call_id,
)
for row in rows
]
def save_active_snapshot(
self,
*,
scope: AgentAppSessionScope,
binding_id: str,
backend_run_id: str,
snapshot: CompositorSessionSnapshot | None,
runtime_layer_specs: list[RuntimeLayerSpec],
pending_form_id: str | None = None,
pending_tool_call_id: str | None = None,
) -> None:
"""Persist the current conversation snapshot and enforce one ACTIVE row.
Agent App chat treats one conversation as one resumable runtime shell.
Saving the latest snapshot therefore upserts the scoped row back to
ACTIVE and retires any other ACTIVE conversation-owned rows for the
same ``tenant_id + app_id + conversation_id`` so later lookups see a
single active session.
"""
if snapshot is None:
return
AgentWorkspaceService.save_binding_session_snapshot(
tenant_id=scope.tenant_id,
binding_id=binding_id,
session_snapshot=snapshot.model_dump_json(),
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
snapshot_json = snapshot.model_dump_json()
runtime_layer_specs_json = _serialize_runtime_layer_specs(runtime_layer_specs)
with session_factory.create_session() as session:
row = session.scalar(self._scope_stmt(scope))
if row is None:
row = AgentRuntimeSession(
tenant_id=scope.tenant_id,
app_id=scope.app_id,
owner_type=AgentRuntimeSessionOwnerType.CONVERSATION,
agent_id=scope.agent_id,
agent_config_snapshot_id=scope.agent_config_snapshot_id,
conversation_id=scope.conversation_id,
backend_run_id=backend_run_id,
session_snapshot=snapshot_json,
composition_layer_specs=runtime_layer_specs_json,
status=AgentRuntimeSessionStatus.ACTIVE,
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
session.add(row)
else:
row.backend_run_id = backend_run_id
row.session_snapshot = snapshot_json
row.composition_layer_specs = runtime_layer_specs_json
row.status = AgentRuntimeSessionStatus.ACTIVE
row.cleaned_at = None
# Set (or clear, when omitted) the ask_human pause correlation.
row.pending_form_id = pending_form_id
row.pending_tool_call_id = pending_tool_call_id
session.flush()
other_rows = session.scalars(
select(AgentRuntimeSession).where(
AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION,
AgentRuntimeSession.tenant_id == scope.tenant_id,
AgentRuntimeSession.app_id == scope.app_id,
AgentRuntimeSession.conversation_id == scope.conversation_id,
AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE,
AgentRuntimeSession.id != row.id,
)
).all()
for other_row in other_rows:
other_row.status = AgentRuntimeSessionStatus.CLEANED
other_row.cleaned_at = naive_utc_now()
session.commit()
def mark_cleaned(self, *, scope: AgentAppSessionScope, backend_run_id: str | None = None) -> None:
with session_factory.create_session() as session:
row = session.scalar(self._active_stmt(scope))
if row is None:
return
if backend_run_id is not None:
row.backend_run_id = backend_run_id
row.status = AgentRuntimeSessionStatus.CLEANED
row.cleaned_at = naive_utc_now()
session.commit()
@staticmethod
def _stored(scope: AgentAppSessionScope, binding: AgentWorkspaceBinding) -> StoredAgentAppSession:
snapshot = (
CompositorSessionSnapshot.model_validate_json(binding.session_snapshot)
if binding.session_snapshot
else None
)
return StoredAgentAppSession(
scope=scope,
binding_id=binding.id,
workspace_id=binding.workspace_id,
backend_binding_ref=binding.backend_binding_ref,
session_snapshot=snapshot,
pending_form_id=binding.pending_form_id,
pending_tool_call_id=binding.pending_tool_call_id,
def _scope_stmt(scope: AgentAppSessionScope):
stmt = select(AgentRuntimeSession).where(
AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION,
AgentRuntimeSession.tenant_id == scope.tenant_id,
AgentRuntimeSession.conversation_id == scope.conversation_id,
AgentRuntimeSession.agent_id == scope.agent_id,
)
if scope.agent_config_snapshot_id is None:
return stmt.where(AgentRuntimeSession.agent_config_snapshot_id.is_(None))
return stmt.where(AgentRuntimeSession.agent_config_snapshot_id == scope.agent_config_snapshot_id)
@classmethod
def _active_stmt(cls, scope: AgentAppSessionScope):
return cls._scope_stmt(scope).where(AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE)
__all__ = ["AgentAppSessionScope", "AgentAppWorkspaceStore", "StoredAgentAppSession"]
__all__ = ["AgentAppRuntimeSessionStore", "AgentAppSessionScope", "StoredAgentAppSession"]
+3 -14
View File
@@ -10,11 +10,11 @@ from core.app.apps.workflow.command_channels import (
CombinedCommandChannel,
)
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.node_factory import get_default_root_node_id
from core.workflow.nodes.agent_v2.workspace_retirement_layer import build_workflow_agent_workspace_retirement_layer
from core.workflow.nodes.agent_v2.session_cleanup_layer import build_workflow_agent_session_cleanup_layer
from core.workflow.snippet_start import get_compatible_start_aliases
from core.workflow.system_variables import build_bootstrap_variables, build_system_variables
from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool
@@ -197,18 +197,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
)
workflow_entry.graph_engine.layer(persistence_layer)
workflow_entry.graph_engine.layer(
build_workflow_agent_workspace_retirement_layer(
dify_run_context=DifyRunContext(
tenant_id=self._workflow.tenant_id,
app_id=self._workflow.app_id,
user_id=self.application_generate_entity.user_id,
user_from=user_from,
invoke_from=invoke_from,
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
)
)
workflow_entry.graph_engine.layer(build_workflow_agent_session_cleanup_layer())
for layer in self._graph_engine_layers:
workflow_entry.graph_engine.layer(layer)
+10 -3
View File
@@ -15,6 +15,8 @@ if TYPE_CHECKING:
DIFY_RUN_CONTEXT_KEY = "_dify"
AGENT_RUNTIME_EXIT_INTENT_ARG = "_agent_runtime_exit_intent"
type AgentRuntimeExitIntent = Literal["suspend", "delete"]
class UserFrom(StrEnum):
@@ -225,8 +227,12 @@ class AgentAppGenerateEntity(ChatAppGenerateEntity):
backend should read from: immutable snapshot, shared draft, or per-user
build draft.
``agent_session_scope_config_version_id`` identifies the draft or immutable
config version whose Workspace Binding should be reused for this session.
``agent_runtime_session_snapshot_id`` carries the runtime session scope
used to resume or suspend within the same editable config surface.
``agent_runtime_exit_intent`` is API-internal lifecycle policy for the
Agent backend session after this turn finishes. Normal chat/resume turns
suspend on exit; build-chat finalization deletes the backend runtime.
``prompt_file_mappings`` preserves the raw request ``files`` array for the
Agent backend prompt. These references are appended to the backend prompt
@@ -236,7 +242,8 @@ class AgentAppGenerateEntity(ChatAppGenerateEntity):
agent_id: str
agent_config_snapshot_id: str
agent_config_version_kind: Literal["snapshot", "draft", "build_draft"] = "snapshot"
agent_session_scope_config_version_id: str | None = None
agent_runtime_session_snapshot_id: str | None = None
agent_runtime_exit_intent: AgentRuntimeExitIntent = "suspend"
prompt_file_mappings: Sequence[JsonValue] = Field(default_factory=list)
+4 -19
View File
@@ -20,13 +20,11 @@ from core.helper.trace_id_helper import ParentTraceContext
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.node_execution_process_data import preserve_workflow_agent_binding_id
from core.workflow.system_variables import SystemVariableKey
from core.workflow.variable_prefixes import SYSTEM_VARIABLE_NODE_ID
from core.workflow.workflow_run_outputs import project_node_outputs_for_workflow_run
from graphon.entities import WorkflowExecution, WorkflowNodeExecution
from graphon.enums import (
BuiltinNodeTypes,
WorkflowExecutionStatus,
WorkflowNodeExecutionMetadataKey,
WorkflowNodeExecutionStatus,
@@ -243,10 +241,7 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
)
self._node_execution_cache[event.id] = domain_execution
if event.node_type == BuiltinNodeTypes.AGENT and event.node_version == "2":
self._workflow_node_execution_repository.save_synchronously(domain_execution)
else:
self._workflow_node_execution_repository.save(domain_execution)
self._workflow_node_execution_repository.save(domain_execution)
snapshot = _NodeRuntimeSnapshot(
node_id=event.node_id,
@@ -366,11 +361,7 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
def _append_retry_history(self, execution: WorkflowNodeExecution, event: NodeRunRetryEvent) -> None:
"""Append a validated full attempt before repository truncation or offload."""
finished_at = naive_utc_now()
process_data = preserve_workflow_agent_binding_id(
event.node_run_result.process_data,
execution.process_data,
)
process_data = dict(process_data or {})
process_data = dict(execution.process_data or {})
raw_history = process_data.get(RETRY_HISTORY_PROCESS_DATA_KEY)
history = list(raw_history) if isinstance(raw_history, list) else []
projected_outputs = project_node_outputs_for_workflow_run(
@@ -399,12 +390,11 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
next_process_data: Mapping[str, Any] | None,
) -> Mapping[str, Any] | None:
"""Keep internal retry history while replacing node-specific Process Data."""
merged_process_data = preserve_workflow_agent_binding_id(existing_process_data, next_process_data)
raw_history = (existing_process_data or {}).get(RETRY_HISTORY_PROCESS_DATA_KEY)
if not isinstance(raw_history, list) or not raw_history:
return merged_process_data
return next_process_data
merged_process_data = dict(merged_process_data or {})
merged_process_data = dict(next_process_data or {})
merged_process_data[RETRY_HISTORY_PROCESS_DATA_KEY] = raw_history
return merged_process_data
@@ -450,11 +440,6 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
outputs=projected_outputs,
metadata=node_result.metadata,
)
else:
domain_execution.process_data = preserve_workflow_agent_binding_id(
node_result.process_data,
domain_execution.process_data,
)
self._workflow_node_execution_repository.save(domain_execution)
self._workflow_node_execution_repository.save_execution_data(domain_execution)
+17 -1
View File
@@ -12,7 +12,7 @@ from core.agent.plugin_entities import AgentProviderEntityWithPlugin
from core.datasource.entities.datasource_entities import DatasourceProviderEntityWithPlugin
from core.plugin.entities.base import BasePluginEntity
from core.plugin.entities.parameters import PluginParameterOption
from core.plugin.entities.plugin import PluginDeclaration, PluginEntity
from core.plugin.entities.plugin import PluginDeclaration, PluginEntity, PluginInstallationSource
from core.tools.entities.common_entities import I18nObject
from core.tools.entities.tool_entities import ToolProviderEntityWithPlugin
from core.trigger.entities.entities import TriggerProviderEntity
@@ -94,6 +94,18 @@ class PluginModelProviderEntity(BaseModel):
declaration: ProviderEntity = Field(description="The declaration of the model provider.")
class PluginModelProviderBinding(BaseModel):
"""Lightweight installation metadata for one model provider."""
provider: str
installation_id: str
plugin_id: str
plugin_unique_identifier: str
runtime_type: str
source: PluginInstallationSource
version: str
class PluginTextEmbeddingNumTokensResponse(BaseModel):
"""
Response for number of tokens.
@@ -207,6 +219,10 @@ class PluginListResponse(BaseModel):
total: int
class PluginInstalledIdsDaemonResponse(BaseModel):
plugin_ids: list[str]
class PluginListWithoutTotalResponse(BaseModel):
list: list[PluginEntity]
has_more: bool
+9
View File
@@ -6,6 +6,7 @@ from core.plugin.entities.plugin_daemon import (
PluginBasicBooleanResponse,
PluginDaemonInnerError,
PluginLLMNumTokensResponse,
PluginModelProviderBinding,
PluginModelProviderEntity,
PluginModelSchemaEntity,
PluginStringResultResponse,
@@ -47,6 +48,14 @@ class PluginModelClient(BasePluginClient):
)
return response
def fetch_model_provider_bindings(self, tenant_id: str) -> Sequence[PluginModelProviderBinding]:
"""Fetch only model-provider installation identities from the daemon."""
return self._request_with_plugin_daemon_response(
"GET",
f"plugin/{tenant_id}/management/models/bindings",
list[PluginModelProviderBinding],
)
def get_model_schema(
self,
tenant_id: str,
+28 -2
View File
@@ -14,6 +14,7 @@ from core.plugin.entities.plugin import (
)
from core.plugin.entities.plugin_daemon import (
PluginDecodeResponse,
PluginInstalledIdsDaemonResponse,
PluginInstallTask,
PluginInstallTaskStartResponse,
PluginListResponse,
@@ -68,6 +69,16 @@ class PluginInstaller(BasePluginClient):
)
return result.list
def list_installed_plugin_ids(self, tenant_id: str, category: PluginCategory) -> list[str]:
"""List all currently installed plugin IDs in one category."""
result = self._request_with_plugin_daemon_response(
"GET",
f"plugin/{tenant_id}/management/installation/ids",
PluginInstalledIdsDaemonResponse,
params={"category": category.value},
)
return result.plugin_ids
def list_plugins_with_total(self, tenant_id: str, page: int, page_size: int) -> PluginListResponse:
return self._request_with_plugin_daemon_response(
"GET",
@@ -77,13 +88,28 @@ class PluginInstaller(BasePluginClient):
)
def list_plugins_by_category(
self, tenant_id: str, category: PluginCategory, page: int, page_size: int
self,
tenant_id: str,
category: PluginCategory,
page: int,
page_size: int,
*,
query: str = "",
tags: Sequence[str] = (),
language: str = "en_US",
) -> PluginListWithoutTotalResponse:
return self._request_with_plugin_daemon_response(
"GET",
f"plugin/{tenant_id}/management/{category.value}/list",
PluginListWithoutTotalResponse,
params={"page": page, "page_size": page_size, "response_type": "paged"},
params={
"page": page,
"page_size": page_size,
"response_type": "paged",
"query": query,
"tags": list(tags),
"language": language,
},
)
def upload_pkg(
+73 -45
View File
@@ -48,6 +48,7 @@ from core.plugin.entities.plugin_daemon import (
PluginInstallTaskStatus,
PluginListResponse,
PluginListWithoutTotalResponse,
PluginModelProviderBinding,
PluginModelProviderEntity,
PluginVerification,
)
@@ -66,7 +67,7 @@ from services.enterprise.plugin_manager_service import (
PreUninstallPluginRequest,
)
from services.errors.plugin import PluginInstallationForbiddenError
from services.feature_service import FeatureService, PluginInstallationPermissionModel, PluginInstallationScope
from services.feature_service import FeatureService, PluginInstallationScope
logger = logging.getLogger(__name__)
_provider_entities_adapter: TypeAdapter[list[ProviderEntity]] = TypeAdapter(list[ProviderEntity])
@@ -78,6 +79,12 @@ class _RedisLock(Protocol):
def release(self) -> None: ...
class _ModelPluginIdentity(Protocol):
plugin_id: str
plugin_unique_identifier: str
source: PluginInstallationSource
class PluginService:
class LatestPluginCache(BaseModel):
plugin_id: str
@@ -265,7 +272,7 @@ class PluginService:
logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True)
@classmethod
def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[PluginEntity]) -> str | None:
def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[_ModelPluginIdentity]) -> str | None:
remote_model_plugins = sorted(
f"{plugin.plugin_id}:{plugin.plugin_unique_identifier}"
for plugin in plugins
@@ -350,7 +357,7 @@ class PluginService:
def _should_invalidate_model_provider_cache_for_remote_model_plugins(
cls,
tenant_id: str,
plugins: Sequence[PluginEntity],
plugins: Sequence[_ModelPluginIdentity],
) -> bool:
remote_model_plugin_marker = cls._get_remote_model_plugin_cache_marker(plugins)
cached_remote_model_plugin_marker = cls._load_cached_remote_model_plugin_marker(tenant_id)
@@ -434,18 +441,14 @@ class PluginService:
exc_info=True,
)
@classmethod
def _fetch_plugin_model_providers_uncached(
cls, tenant_id: str, client: PluginModelClient | None
) -> tuple[ProviderEntity, ...]:
model_client = client or PluginModelClient()
return tuple(cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id))
@classmethod
def _fetch_and_cache_plugin_model_providers(
cls, tenant_id: str, client: PluginModelClient | None, *, refresh_generation: int | None
) -> tuple[ProviderEntity, ...]:
providers = cls._fetch_plugin_model_providers_uncached(tenant_id, client)
model_client = client or PluginModelClient()
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_cached_plugin_model_providers(tenant_id, generation, providers)
@@ -475,9 +478,6 @@ class PluginService:
are intentionally owned by this service so tenant isolation and cache
expiry are handled in one place.
"""
if not dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED:
return cls._fetch_plugin_model_providers_uncached(tenant_id, client)
deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT
while True:
@@ -604,30 +604,22 @@ class PluginService:
return result
@staticmethod
def _check_marketplace_only_permission() -> None:
def _check_marketplace_only_permission():
"""
Check if the marketplace only permission is enabled
"""
permission = PluginService._get_plugin_installation_permission()
if permission.restrict_to_marketplace_only:
features = FeatureService.get_system_features()
if features.plugin_installation_permission.restrict_to_marketplace_only:
raise PluginInstallationForbiddenError("Plugin installation is restricted to marketplace only")
@staticmethod
def _get_plugin_installation_permission() -> PluginInstallationPermissionModel:
"""Resolve the validated policy and reject deny-all before any installation side effect."""
permission = FeatureService.get_plugin_installation_permission()
if permission.plugin_installation_scope == PluginInstallationScope.NONE:
raise PluginInstallationForbiddenError("Installing plugins is not allowed")
return permission
@staticmethod
def _check_plugin_installation_scope(plugin_verification: PluginVerification | None) -> None:
def _check_plugin_installation_scope(plugin_verification: PluginVerification | None):
"""
Check the plugin installation scope
"""
permission = PluginService._get_plugin_installation_permission()
features = FeatureService.get_system_features()
match permission.plugin_installation_scope:
match features.plugin_installation_permission.plugin_installation_scope:
case PluginInstallationScope.OFFICIAL_ONLY:
if (
plugin_verification is None
@@ -642,10 +634,10 @@ class PluginService:
raise PluginInstallationForbiddenError(
"Plugin installation is restricted to official and specific partners"
)
case PluginInstallationScope.NONE:
raise PluginInstallationForbiddenError("Installing plugins is not allowed")
case PluginInstallationScope.ALL:
pass
case _:
raise PluginInstallationForbiddenError("Plugin installation policy is invalid")
@staticmethod
def get_debugging_key(tenant_id: str) -> str:
@@ -671,6 +663,26 @@ class PluginService:
plugins = manager.list_plugins(tenant_id)
return plugins
@staticmethod
def list_installed_plugin_ids(tenant_id: str, category: PluginCategory) -> Sequence[str]:
"""List all currently installed plugin IDs in one category through the daemon's lightweight query."""
manager = PluginInstaller()
return manager.list_installed_plugin_ids(tenant_id, category)
@staticmethod
def list_model_provider_bindings(
tenant_id: str, *, client: PluginModelClient | None = None
) -> Sequence[PluginModelProviderBinding]:
"""Return fresh model bindings and reconcile remote-debug provider metadata before it is read."""
model_client = client or PluginModelClient()
bindings = model_client.fetch_model_provider_bindings(tenant_id)
if PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins(tenant_id, bindings):
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
marker = PluginService._get_remote_model_plugin_cache_marker(bindings)
PluginService._store_cached_remote_model_plugin_marker(tenant_id, marker)
return bindings
@staticmethod
def list_with_total(tenant_id: str, user_id: str, page: int, page_size: int) -> PluginListResponse:
"""List tenant plugins with endpoint counts reconciled from live records.
@@ -687,17 +699,33 @@ class PluginService:
@staticmethod
def list_by_category(
tenant_id: str, category: PluginCategory, page: int, page_size: int
tenant_id: str,
category: PluginCategory,
page: int,
page_size: int,
*,
query: str = "",
tags: Sequence[str] = (),
language: str = "en_US",
) -> PluginListWithoutTotalResponse:
"""
List plugins in one category with a has-more cursor signal and without calculating total.
The daemon scans tenant installations in the existing list order and stops once it finds one extra match.
This keeps pagination usable before category is persisted on installation rows.
The daemon applies category, search, and tag filters before pagination, then stops once it finds one extra
match. Only a complete, unfiltered first page may reconcile the model-provider cache; the unpaginated model
binding read is the authoritative marker source for larger result sets.
"""
manager = PluginInstaller()
plugins = manager.list_plugins_by_category(tenant_id, category, page, page_size)
if category == PluginCategory.Model:
plugins = manager.list_plugins_by_category(
tenant_id,
category,
page,
page_size,
query=query,
tags=tags,
language=language,
)
if category == PluginCategory.Model and page == 1 and not plugins.has_more and not query and not tags:
should_invalidate_model_provider_cache = (
PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins(
tenant_id,
@@ -915,7 +943,7 @@ class PluginService:
# check if plugin pkg is already downloaded
manager = PluginInstaller()
permission = PluginService._get_plugin_installation_permission()
features = FeatureService.get_system_features()
try:
manager.fetch_plugin_manifest(tenant_id, new_plugin_unique_identifier)
@@ -927,7 +955,7 @@ class PluginService:
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=permission.restrict_to_marketplace_only,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
)
# check if the plugin is available to install
@@ -982,11 +1010,11 @@ class PluginService:
"""
PluginService._check_marketplace_only_permission()
manager = PluginInstaller()
permission = PluginService._get_plugin_installation_permission()
features = FeatureService.get_system_features()
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=permission.restrict_to_marketplace_only,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
)
PluginService._check_plugin_installation_scope(response.verification)
@@ -1004,13 +1032,13 @@ class PluginService:
pkg = download_with_size_limit(
f"https://github.com/{repo}/releases/download/{version}/{package}", dify_config.PLUGIN_MAX_PACKAGE_SIZE
)
permission = PluginService._get_plugin_installation_permission()
features = FeatureService.get_system_features()
manager = PluginInstaller()
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=permission.restrict_to_marketplace_only,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
)
PluginService._check_plugin_installation_scope(response.verification)
@@ -1084,7 +1112,7 @@ class PluginService:
if not dify_config.MARKETPLACE_ENABLED:
raise ValueError("marketplace is not enabled")
permission = PluginService._get_plugin_installation_permission()
features = FeatureService.get_system_features()
manager = PluginInstaller()
try:
@@ -1094,7 +1122,7 @@ class PluginService:
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=permission.restrict_to_marketplace_only,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
)
# check if the plugin is available to install
PluginService._check_plugin_installation_scope(response.verification)
@@ -1116,7 +1144,7 @@ class PluginService:
# collect actual plugin_unique_identifiers
actual_plugin_unique_identifiers = []
metas = []
permission = PluginService._get_plugin_installation_permission()
features = FeatureService.get_system_features()
# check if already downloaded
for plugin_unique_identifier in plugin_unique_identifiers:
@@ -1134,7 +1162,7 @@ class PluginService:
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=permission.restrict_to_marketplace_only,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
)
# check if the plugin is available to install
PluginService._check_plugin_installation_scope(response.verification)
@@ -2023,8 +2023,6 @@ class DatasetRetrieval:
redis_client.zremrangebyscore(key, 0, current_time - 60000)
request_count = redis_client.zcard(key)
if request_count > knowledge_rate_limit.limit:
# The rate-limit exception is raised after this block, so commit the audit row
# explicitly instead of relying on the Session context, which only closes it.
with session_factory.create_session() as session:
rate_limit_log = RateLimitLog(
tenant_id=tenant_id,
@@ -2032,7 +2030,6 @@ class DatasetRetrieval:
operation="knowledge",
)
session.add(rate_limit_log)
session.commit()
raise exc.RateLimitExceededError(
"you have reached the knowledge base request rate limit of your subscription."
)
@@ -16,9 +16,6 @@ from core.repositories.factory import (
OrderConfig,
WorkflowNodeExecutionRepository,
)
from core.repositories.sqlalchemy_workflow_node_execution_repository import (
SQLAlchemyWorkflowNodeExecutionRepository,
)
from graphon.entities import WorkflowNodeExecution
from models import Account, CreatorUserRole, EndUser
from models.workflow import WorkflowNodeExecutionTriggeredFrom
@@ -52,7 +49,6 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
_creator_user_role: CreatorUserRole
_execution_cache: dict[str, WorkflowNodeExecution]
_workflow_execution_mapping: dict[str, list[str]]
_sql_repository: SQLAlchemyWorkflowNodeExecutionRepository
def __init__(
self,
@@ -102,13 +98,6 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
# Cache for mapping workflow_execution_ids to execution IDs for efficient retrieval
self._workflow_execution_mapping = {}
self._sql_repository = SQLAlchemyWorkflowNodeExecutionRepository(
session_factory=session_factory,
tenant_id=tenant_id,
user=user,
app_id=app_id,
triggered_from=triggered_from,
)
logger.info(
"Initialized CeleryWorkflowNodeExecutionRepository for tenant %s, app %s, triggered_from %s",
@@ -160,17 +149,6 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
# For now, we'll re-raise the exception
raise
@override
def save_synchronously(self, execution: WorkflowNodeExecution) -> None:
"""Create the Agent v2 caller row before runtime participant allocation."""
self._sql_repository.save_synchronously(execution)
self._execution_cache[execution.id] = execution
if execution.workflow_execution_id:
execution_ids = self._workflow_execution_mapping.setdefault(execution.workflow_execution_id, [])
if execution.id not in execution_ids:
execution_ids.append(execution.id)
@override
def get_by_workflow_execution(
self,
-2
View File
@@ -35,8 +35,6 @@ class WorkflowExecutionRepository(Protocol):
class WorkflowNodeExecutionRepository(Protocol):
def save(self, execution: WorkflowNodeExecution): ...
def save_synchronously(self, execution: WorkflowNodeExecution) -> None: ...
def save_execution_data(self, execution: WorkflowNodeExecution): ...
def get_by_workflow_execution(
@@ -18,7 +18,6 @@ from tenacity import before_sleep_log, retry, retry_if_exception, stop_after_att
from configs import dify_config
from core.repositories.factory import OrderConfig, WorkflowNodeExecutionRepository
from core.workflow.node_execution_process_data import preserve_workflow_agent_binding_id
from extensions.ext_storage import storage
from graphon.entities import WorkflowNodeExecution
from graphon.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
@@ -373,12 +372,6 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
logger.exception("Failed to save workflow node execution after all retries")
raise
@override
def save_synchronously(self, execution: WorkflowNodeExecution) -> None:
"""Persist a caller row before an Agent v2 participant is materialized."""
self.save(execution)
def _persist_to_database(self, db_model: WorkflowNodeExecutionModel):
"""
Persist the database model to the database.
@@ -393,13 +386,6 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
existing = session.get(WorkflowNodeExecutionModel, db_model.id)
if existing:
merged_process_data = preserve_workflow_agent_binding_id(
existing.process_data_dict,
db_model.process_data_dict,
)
db_model.process_data = (
_deterministic_json_dump(merged_process_data) if merged_process_data is not None else None
)
# Update existing record by copying all non-private attributes
for key, value in db_model.__dict__.items():
if not key.startswith("_"):
@@ -456,25 +442,18 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
else:
db_model.outputs = self._json_encode(domain_model.outputs)
process_data = preserve_workflow_agent_binding_id(db_model.process_data_dict, domain_model.process_data)
if process_data is not None:
if domain_model.process_data is not None:
result = self._truncate_and_upload(
process_data,
domain_model.process_data,
domain_model.id,
ExecutionOffLoadType.PROCESS_DATA,
)
if result is not None:
truncated_process_data = preserve_workflow_agent_binding_id(
process_data,
result.truncated_value,
)
if truncated_process_data is None:
raise ValueError("truncated process data is unavailable")
db_model.process_data = self._json_encode(truncated_process_data)
domain_model.set_truncated_process_data(truncated_process_data)
db_model.process_data = self._json_encode(result.truncated_value)
domain_model.set_truncated_process_data(result.truncated_value)
offload_data = _replace_or_append_offload(offload_data, result.offload)
else:
db_model.process_data = self._json_encode(process_data)
db_model.process_data = self._json_encode(domain_model.process_data)
db_model.offload_data = offload_data
with self._session_factory() as session, session.begin():
-2
View File
@@ -107,8 +107,6 @@ class MCPTool(Tool):
if self.entity.output_schema and result.structuredContent:
for k, v in result.structuredContent.items():
yield self.create_variable_message(k, v)
elif result.structuredContent:
yield self.create_json_message(result.structuredContent)
def _process_text_content(self, content: TextContent) -> Generator[ToolInvokeMessage, None, None]:
"""Process text content and yield appropriate messages."""
@@ -1,27 +0,0 @@
from collections.abc import Mapping
from typing import Any
WORKFLOW_AGENT_BINDING_ID_KEY = "workflow_agent_binding_id"
def preserve_workflow_agent_binding_id(
identity_source: Mapping[str, Any] | None,
process_data: Mapping[str, Any] | None,
) -> dict[str, Any] | None:
source_id = (identity_source or {}).get(WORKFLOW_AGENT_BINDING_ID_KEY)
target_id = (process_data or {}).get(WORKFLOW_AGENT_BINDING_ID_KEY)
for value in (source_id, target_id):
if value is not None and not isinstance(value, str):
raise ValueError("workflow_agent_binding_id must be a string")
if source_id is not None and target_id is not None and source_id != target_id:
raise ValueError("workflow_agent_binding_id does not match")
if process_data is None and source_id is None:
return None
merged = dict(process_data or {})
if source_id is not None:
merged[WORKFLOW_AGENT_BINDING_ID_KEY] = source_id
return merged
__all__ = ["WORKFLOW_AGENT_BINDING_ID_KEY", "preserve_workflow_agent_binding_id"]
+2 -3
View File
@@ -487,7 +487,7 @@ class DifyNodeFactory(NodeFactory):
from core.workflow.nodes.agent_v2.output_failure_orchestrator import OutputFailureOrchestrator
from core.workflow.nodes.agent_v2.output_file_rebacker import reback_tool_file_output
from core.workflow.nodes.agent_v2.output_type_checker import PerOutputTypeChecker
from core.workflow.nodes.agent_v2.session_store import WorkflowAgentWorkspaceStore
from core.workflow.nodes.agent_v2.session_store import WorkflowAgentRuntimeSessionStore
return {
"binding_resolver": WorkflowAgentBindingResolver(),
@@ -497,7 +497,6 @@ class DifyNodeFactory(NodeFactory):
),
"agent_backend_client": create_agent_backend_run_client(
base_url=dify_config.AGENT_BACKEND_BASE_URL,
api_token=dify_config.AGENT_BACKEND_API_TOKEN,
use_fake=dify_config.AGENT_BACKEND_USE_FAKE,
fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO,
stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS,
@@ -512,7 +511,7 @@ class DifyNodeFactory(NodeFactory):
# tenant validator resolves ToolFile (canonical) + UploadFile refs.
"type_checker": PerOutputTypeChecker(file_validator=AgentOutputFileTenantValidator()),
"failure_orchestrator": OutputFailureOrchestrator(),
"session_store": WorkflowAgentWorkspaceStore(),
"session_store": WorkflowAgentRuntimeSessionStore(),
}
return {
"strategy_resolver": self._agent_strategy_resolver,
@@ -198,23 +198,8 @@ class AgentRuntimeSupport:
if model_schema:
model_schema = self._remove_unsupported_model_features_for_old_version(model_schema)
value["entity"] = model_schema.model_dump(mode="json")
# The model selector value from the workflow frontend only
# carries provider/model/mode — it does NOT include
# completion_params. AgentStrategy plugins (cot_agent,
# function_calling) read completion_params to build the
# LLMModelConfig that is backwards-invoked, and some model
# providers raise KeyError('required') when
# completion_params is empty because their parameter_rules
# declare required fields with no default. Populate
# completion_params with the defaults declared in the model
# schema so the plugin daemon always receives a valid set
# of model parameters.
if "completion_params" not in value:
value["completion_params"] = self._extract_default_completion_params(model_schema)
else:
value["entity"] = None
if "completion_params" not in value:
value["completion_params"] = {}
result[parameter_name] = value
return result
@@ -290,24 +275,6 @@ class AgentRuntimeSupport:
model_schema.features.remove(feature)
return model_schema
@staticmethod
def _extract_default_completion_params(model_schema: AIModelEntity) -> dict[str, Any]:
"""Build a completion_params dict from the model schema's parameter_rules.
The workflow Agent node's model-selector parameter only stores
provider/model/mode it never carries completion_params. When the
value is forwarded to the plugin daemon, AgentModelConfig defaults
completion_params to ``{}``, which causes some model providers to fail
because their parameter_rules declare required fields. This helper
collects the ``default`` value of every parameter_rule that has one so
the plugin daemon receives a valid, non-empty set of model parameters.
"""
completion_params: dict[str, Any] = {}
for rule in model_schema.parameter_rules:
if rule.default is not None:
completion_params[rule.name] = rule.default
return completion_params
@staticmethod
def _filter_mcp_type_tool(
strategy: ResolvedAgentStrategy,
+154 -148
View File
@@ -18,10 +18,13 @@ from clients.agent_backend import (
AgentBackendRunEventAdapter,
AgentBackendRunFailedInternalEvent,
AgentBackendRunSucceededInternalEvent,
AgentBackendSessionCleanupPayload,
AgentBackendStreamError,
AgentBackendStreamInternalEvent,
AgentBackendTransportError,
AgentBackendValidationError,
RuntimeLayerSpec,
extract_runtime_layer_specs,
)
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
from core.repositories.human_input_repository import HumanInputFormRepository, HumanInputFormRepositoryImpl
@@ -30,12 +33,11 @@ from core.workflow.nodes.human_input.session_binding import default_session_bind
from core.workflow.system_variables import SystemVariableKey, get_system_text
from graphon.entities.pause_reason import HitlRequired, SchedulingPause
from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
from graphon.graph_events import NodeRunPauseRequestedEvent
from graphon.node_events import NodeEventBase, NodeRunResult, StreamCompletedEvent
from graphon.node_events import NodeEventBase, NodeRunResult, PauseRequestedEvent, StreamCompletedEvent
from graphon.nodes.base.node import Node
from models.agent_config_entities import AgentSoulConfig, WorkflowNodeJobConfig
from services.agent.prompt_mentions import extract_workflow_node_output_selectors
from services.agent.workspace_service import AgentWorkspaceNotFoundError
from tasks.agent_backend_session_cleanup_task import cleanup_workflow_agent_runtime_session
from .ask_human_hitl import AskHumanFormBuildError, build_ask_human_pause_reason
from .ask_human_resume import build_deferred_tool_results, resolve_ask_human_form
@@ -54,7 +56,7 @@ from .runtime_request_builder import (
WorkflowAgentRuntimeRequestBuilder,
WorkflowAgentRuntimeRequestBuildError,
)
from .session_store import WorkflowAgentSessionScope, WorkflowAgentWorkspaceStore
from .session_store import WorkflowAgentRuntimeSessionStore, WorkflowAgentSessionScope
if TYPE_CHECKING:
from graphon.entities import GraphInitParams
@@ -66,7 +68,7 @@ logger = logging.getLogger(__name__)
# Stage 4 §5+§7: the terminal events that `_consume_event_stream` may return.
# Stream + started events are filtered out before we yield; transport errors
# are surfaced as a separate StreamCompletedEvent in the second tuple slot.
type _TerminalAgentBackendEvent = (
_TerminalAgentBackendEvent = (
AgentBackendRunSucceededInternalEvent
| AgentBackendRunFailedInternalEvent
| AgentBackendRunCancelledInternalEvent
@@ -91,7 +93,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
output_adapter: WorkflowAgentOutputAdapter,
type_checker: PerOutputTypeChecker,
failure_orchestrator: OutputFailureOrchestrator,
session_store: WorkflowAgentWorkspaceStore,
session_store: WorkflowAgentRuntimeSessionStore | None = None,
) -> None:
super().__init__(
node_id=node_id,
@@ -128,34 +130,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
return reason
@override
def _run(self) -> Generator[NodeEventBase | NodeRunPauseRequestedEvent, None, None]:
inputs: dict[str, Any] = {}
process_data: dict[str, Any] = {}
metadata: dict[str, Any] = {
"agent_backend": {
"status": "not_started",
}
}
try:
yield from self._run_inner(inputs=inputs, process_data=process_data, metadata=metadata)
except Exception as error:
if not process_data:
raise
yield self._failure_event(
inputs=inputs,
process_data=process_data,
metadata=metadata,
error=str(error),
error_type="agent_workflow_node_runtime_error",
)
def _run_inner(
self,
*,
inputs: dict[str, Any],
process_data: dict[str, Any],
metadata: dict[str, Any],
) -> Generator[NodeEventBase | NodeRunPauseRequestedEvent, None, None]:
def _run(self) -> Generator[NodeEventBase, None, None]:
dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY))
workflow_id = self.graph_init_params.workflow_id
workflow_run_id = get_system_text(
@@ -168,24 +143,21 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
self.graph_runtime_state.variable_pool,
SystemVariableKey.CONVERSATION_ID,
)
inputs: dict[str, Any] = {}
process_data: dict[str, Any] = {}
metadata: dict[str, Any] = {
"agent_backend": {
"status": "not_started",
}
}
# ──── Setup: resolve binding once + extract declared outputs for stage 4 checks ────
try:
existing_scope = self._session_store.load_existing_node_execution_scope(
tenant_id=dify_ctx.tenant_id,
app_id=dify_ctx.app_id,
workflow_id=workflow_id,
workflow_run_id=workflow_run_id,
node_id=self._node_id,
node_execution_id=self.execution_id,
)
bundle = self._binding_resolver.resolve(
tenant_id=dify_ctx.tenant_id,
app_id=dify_ctx.app_id,
workflow_id=workflow_id,
node_id=self._node_id,
binding_id=existing_scope.workflow_agent_binding_id if existing_scope is not None else None,
snapshot_id=existing_scope.agent_config_snapshot_id if existing_scope is not None else None,
)
except WorkflowAgentBindingError as error:
yield self._failure_event(
@@ -196,31 +168,20 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
error_type=error.error_code,
)
return
except AgentWorkspaceNotFoundError as error:
yield self._failure_event(
inputs=inputs,
process_data=process_data,
metadata=metadata,
error=str(error),
error_type="agent_workflow_node_runtime_error",
)
return
process_data.update(
{
"agent_id": bundle.agent.id,
"agent_config_snapshot_id": bundle.snapshot.id,
"workflow_agent_binding_id": bundle.binding.id,
}
)
session_scope = existing_scope or WorkflowAgentSessionScope(
process_data = {
"agent_id": bundle.agent.id,
"agent_config_snapshot_id": bundle.snapshot.id,
"binding_id": bundle.binding.id,
}
session_scope = WorkflowAgentSessionScope(
tenant_id=dify_ctx.tenant_id,
app_id=dify_ctx.app_id,
workflow_id=workflow_id,
workflow_run_id=workflow_run_id,
node_id=self._node_id,
node_execution_id=self.execution_id,
workflow_agent_binding_id=bundle.binding.id,
node_execution_id=self.id,
binding_id=bundle.binding.id,
agent_id=bundle.agent.id,
agent_config_snapshot_id=bundle.snapshot.id,
)
@@ -239,53 +200,47 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
# the second Agent run as deferred_tool_results; if it is somehow still
# waiting, re-emit the same pause defensively.
deferred_tool_results = None
stored_session = self._session_store.load_or_create_node_execution_session(
session_scope,
home_snapshot_id=bundle.snapshot.home_snapshot_id,
)
if stored_session.pending_form_id is not None:
resume_outcome = resolve_ask_human_form(
form_id=stored_session.pending_form_id,
tenant_id=dify_ctx.tenant_id,
node_id=self._node_id,
)
if resume_outcome is not None and resume_outcome.repause is not None:
yield self._pause_event(
reason=resume_outcome.repause,
inputs=inputs,
process_data=process_data,
metadata=metadata,
)
return
if (
resume_outcome is not None
and resume_outcome.deferred_result is not None
and stored_session.pending_tool_call_id is not None
):
deferred_tool_results = build_deferred_tool_results(
tool_call_id=stored_session.pending_tool_call_id,
result=resume_outcome.deferred_result,
if self._session_store is not None:
stored_session = self._session_store.load_active_session(session_scope)
if stored_session is not None and stored_session.pending_form_id is not None:
resume_outcome = resolve_ask_human_form(
form_id=stored_session.pending_form_id,
tenant_id=dify_ctx.tenant_id,
node_id=self._node_id,
)
if resume_outcome is not None and resume_outcome.repause is not None:
yield PauseRequestedEvent(reason=self._to_graph_pause_reason(resume_outcome.repause))
return
if (
resume_outcome is not None
and resume_outcome.deferred_result is not None
and stored_session.pending_tool_call_id is not None
):
deferred_tool_results = build_deferred_tool_results(
tool_call_id=stored_session.pending_tool_call_id,
result=resume_outcome.deferred_result,
)
# ──── Retry loop (Stage 4 §7) ────
attempt = 0
while True:
try:
session_snapshot = None
if self._session_store is not None:
session_snapshot = self._session_store.load_active_snapshot(session_scope)
runtime_request = self._runtime_request_builder.build(
WorkflowAgentRuntimeBuildContext(
dify_context=dify_ctx,
workflow_id=workflow_id,
workflow_run_id=workflow_run_id,
node_id=self._node_id,
node_execution_id=self.execution_id,
node_execution_id=self.id,
variable_pool=self.graph_runtime_state.variable_pool,
binding=bundle.binding,
agent=bundle.agent,
snapshot=bundle.snapshot,
binding_id=stored_session.binding_id,
backend_binding_ref=stored_session.backend_binding_ref,
attempt=attempt,
session_snapshot=stored_session.session_snapshot,
session_snapshot=session_snapshot,
deferred_tool_results=deferred_tool_results,
)
)
@@ -311,9 +266,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
# Capture inputs only from the first attempt so retry doesn't churn the
# node's "inputs" payload that ends up in the workflow detail view.
if attempt == 0:
inputs["agent_backend_request"] = runtime_request.redacted_request
metadata.clear()
metadata.update(runtime_request.metadata)
inputs = {"agent_backend_request": runtime_request.redacted_request}
metadata = dict(runtime_request.metadata)
metadata["attempt"] = attempt
try:
@@ -334,12 +288,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
"status": create_response.status,
}
terminal_event, exhausted = self._consume_event_stream(
create_response.run_id,
inputs=inputs,
process_data=process_data,
metadata=metadata,
)
terminal_event, exhausted = self._consume_event_stream(create_response.run_id, metadata)
if exhausted is not None:
# Streaming error / unexpected end — surface immediately without
# retrying because the failure is transport-level.
@@ -400,23 +349,29 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
)
self._save_session_snapshot(
session_scope=session_scope,
binding_id=stored_session.binding_id,
backend_run_id=terminal_event.run_id,
snapshot=terminal_event.session_snapshot,
runtime_layer_specs=extract_runtime_layer_specs(runtime_request.request.composition),
metadata=metadata,
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
yield self._pause_event(
reason=pause_reason,
inputs=inputs,
process_data=process_data,
metadata=metadata,
)
yield PauseRequestedEvent(reason=self._to_graph_pause_reason(pause_reason))
return
# A failed attempt does not retire the product-owned Binding. The
# Workflow Run terminal lifecycle event owns that transition.
# Non-success terminal (failed / cancelled) skips per-output
# post-processing — the backend itself already failed. We also retire
# the local ACTIVE session row so a workflow loop back into the same
# Agent node cannot resume from a stale snapshot. The failed agent
# backend layers (suspended per ``on_exit``) are left for agent
# backend's own GC; this row will no longer be picked up by the
# workflow-terminal cleanup layer.
if not isinstance(terminal_event, AgentBackendRunSucceededInternalEvent):
self._mark_session_cleaned_on_failure(
session_scope=session_scope,
backend_run_id=terminal_event.run_id,
metadata=metadata,
)
yield StreamCompletedEvent(
node_run_result=self._output_adapter.build_failure_result(
event=terminal_event,
@@ -429,8 +384,9 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
self._save_session_snapshot(
session_scope=session_scope,
binding_id=stored_session.binding_id,
backend_run_id=terminal_event.run_id,
snapshot=terminal_event.session_snapshot,
runtime_layer_specs=extract_runtime_layer_specs(runtime_request.request.composition),
metadata=metadata,
)
@@ -502,9 +458,6 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
def _consume_event_stream(
self,
run_id: str,
*,
inputs: dict[str, Any],
process_data: dict[str, Any],
metadata: dict[str, Any],
) -> tuple[
_TerminalAgentBackendEvent | None,
@@ -554,8 +507,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
return internal_event, None
self._cancel_backend_run(run_id, reason="unexpected_event")
return None, self._failure_event(
inputs=inputs,
process_data=process_data,
inputs={},
process_data={},
metadata=metadata,
error=f"Unexpected internal event type {internal_event.type!r}",
error_type="agent_backend_stream_error",
@@ -563,8 +516,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
except AgentBackendError as error:
self._cancel_backend_run(run_id, reason=self._stream_stop_reason())
return None, self._failure_event(
inputs=inputs,
process_data=process_data,
inputs={},
process_data={},
metadata=metadata,
error=str(error),
error_type=self._agent_backend_error_type(error),
@@ -572,8 +525,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
except Exception as error:
self._cancel_backend_run(run_id, reason=self._stream_stop_reason())
return None, self._failure_event(
inputs=inputs,
process_data=process_data,
inputs={},
process_data={},
metadata=metadata,
error=str(error),
error_type="agent_backend_stream_error",
@@ -643,17 +596,21 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
self,
*,
session_scope: WorkflowAgentSessionScope,
binding_id: str,
backend_run_id: str,
snapshot: CompositorSessionSnapshot | None,
runtime_layer_specs: list[RuntimeLayerSpec],
metadata: dict[str, Any],
pending_form_id: str | None = None,
pending_tool_call_id: str | None = None,
) -> None:
if self._session_store is None:
return
try:
self._session_store.save_active_snapshot(
scope=session_scope,
binding_id=binding_id,
backend_run_id=backend_run_id,
snapshot=snapshot,
runtime_layer_specs=runtime_layer_specs,
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
@@ -662,18 +619,88 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
metadata["agent_backend"] = agent_backend
except Exception:
logger.warning(
"Failed to persist workflow Agent Binding session snapshot: "
"tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s",
"Failed to persist workflow Agent runtime session snapshot: "
"tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s backend_run_id=%s",
session_scope.tenant_id,
session_scope.workflow_run_id,
session_scope.node_id,
session_scope.workflow_agent_binding_id,
session_scope.binding_id,
session_scope.agent_id,
backend_run_id,
exc_info=True,
)
agent_backend = dict(metadata.get("agent_backend") or {})
agent_backend["session_snapshot_persisted"] = False
agent_backend["session_snapshot_persist_error"] = "workflow_agent_workspace_store_error"
agent_backend["session_snapshot_persist_error"] = "workflow_agent_runtime_session_store_error"
metadata["agent_backend"] = agent_backend
def _mark_session_cleaned_on_failure(
self,
*,
session_scope: WorkflowAgentSessionScope,
backend_run_id: str,
metadata: dict[str, Any],
) -> None:
if self._session_store is None:
return
stored_session = self._session_store.load_active_session(session_scope)
try:
if stored_session is not None and stored_session.runtime_layer_specs:
payload = AgentBackendSessionCleanupPayload(
session_snapshot=stored_session.session_snapshot,
runtime_layer_specs=stored_session.runtime_layer_specs,
idempotency_key=(
f"{session_scope.tenant_id}:{session_scope.workflow_run_id}:{session_scope.node_id}:"
f"{session_scope.binding_id}:workflow-agent-failure-cleanup:"
f"{stored_session.backend_run_id or 'no-stored-run'}:{backend_run_id}"
),
metadata={
"tenant_id": session_scope.tenant_id,
"app_id": session_scope.app_id,
"workflow_id": session_scope.workflow_id,
"workflow_run_id": session_scope.workflow_run_id,
"node_id": session_scope.node_id,
"node_execution_id": session_scope.node_execution_id,
"binding_id": session_scope.binding_id,
"agent_id": session_scope.agent_id,
"agent_config_snapshot_id": session_scope.agent_config_snapshot_id,
"previous_agent_backend_run_id": stored_session.backend_run_id,
"failed_agent_backend_run_id": backend_run_id,
},
)
cleanup_workflow_agent_runtime_session.delay(payload.model_dump(mode="json"))
except Exception:
logger.warning(
"Failed to enqueue workflow Agent backend cleanup on agent run failure: "
"tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s backend_run_id=%s",
session_scope.tenant_id,
session_scope.workflow_run_id,
session_scope.node_id,
session_scope.binding_id,
session_scope.agent_id,
backend_run_id,
exc_info=True,
)
try:
self._session_store.mark_cleaned(scope=session_scope, backend_run_id=backend_run_id)
agent_backend = dict(metadata.get("agent_backend") or {})
agent_backend["session_snapshot_cleaned_on_failure"] = True
metadata["agent_backend"] = agent_backend
except Exception:
logger.warning(
"Failed to mark workflow Agent runtime session cleaned on agent run failure: "
"tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s backend_run_id=%s",
session_scope.tenant_id,
session_scope.workflow_run_id,
session_scope.node_id,
session_scope.binding_id,
session_scope.agent_id,
backend_run_id,
exc_info=True,
)
agent_backend = dict(metadata.get("agent_backend") or {})
agent_backend["session_snapshot_cleaned_on_failure"] = False
agent_backend["session_snapshot_cleanup_error"] = "workflow_agent_runtime_session_store_error"
metadata["agent_backend"] = agent_backend
@staticmethod
@@ -715,27 +742,6 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
)
)
def _pause_event(
self,
*,
reason: HumanInputRequired | SchedulingPause,
inputs: dict[str, Any],
process_data: dict[str, Any],
metadata: dict[str, Any],
) -> NodeRunPauseRequestedEvent:
return NodeRunPauseRequestedEvent(
id=self.execution_id,
node_id=self._node_id,
node_type=self.node_type,
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.PAUSED,
inputs=inputs,
process_data=process_data,
metadata={WorkflowNodeExecutionMetadataKey.AGENT_LOG: metadata},
),
reason=self._to_graph_pause_reason(reason),
)
@staticmethod
def _agent_backend_error_type(error: AgentBackendError) -> str:
if isinstance(error, AgentBackendValidationError):
@@ -41,27 +41,18 @@ class WorkflowAgentBindingResolver:
app_id: str,
workflow_id: str,
node_id: str,
binding_id: str | None = None,
snapshot_id: str | None = None,
) -> WorkflowAgentBindingBundle:
"""Resolve the current binding, optionally at a generation pinned by an existing execution."""
if (binding_id is None) != (snapshot_id is None):
raise WorkflowAgentBindingError(
"agent_binding_generation_invalid",
"Workflow Agent binding and config snapshot must be pinned together.",
)
with session_factory.create_session() as session:
binding_stmt = select(WorkflowAgentNodeBinding).where(
WorkflowAgentNodeBinding.tenant_id == tenant_id,
WorkflowAgentNodeBinding.app_id == app_id,
WorkflowAgentNodeBinding.workflow_id == workflow_id,
WorkflowAgentNodeBinding.node_id == node_id,
binding = session.scalar(
select(WorkflowAgentNodeBinding)
.where(
WorkflowAgentNodeBinding.tenant_id == tenant_id,
WorkflowAgentNodeBinding.app_id == app_id,
WorkflowAgentNodeBinding.workflow_id == workflow_id,
WorkflowAgentNodeBinding.node_id == node_id,
)
.limit(1)
)
if binding_id is not None:
binding_stmt = binding_stmt.where(WorkflowAgentNodeBinding.id == binding_id)
binding = session.scalar(binding_stmt.limit(1))
if binding is None:
raise WorkflowAgentBindingError(
"agent_binding_not_found",
@@ -86,16 +77,12 @@ class WorkflowAgentBindingResolver:
f"Agent {binding.agent_id} is not available or has not been published.",
)
effective_snapshot_id = (
(
agent.active_config_snapshot_id
if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT
else binding.current_snapshot_id
)
if snapshot_id is None
else snapshot_id
snapshot_id = (
agent.active_config_snapshot_id
if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT
else binding.current_snapshot_id
)
if effective_snapshot_id is None:
if snapshot_id is None:
raise WorkflowAgentBindingError(
"agent_config_snapshot_not_found",
"Workflow Agent binding has no current config snapshot.",
@@ -106,14 +93,14 @@ class WorkflowAgentBindingResolver:
.where(
AgentConfigSnapshot.tenant_id == tenant_id,
AgentConfigSnapshot.agent_id == agent.id,
AgentConfigSnapshot.id == effective_snapshot_id,
AgentConfigSnapshot.id == snapshot_id,
)
.limit(1)
)
if snapshot is None:
raise WorkflowAgentBindingError(
"agent_config_snapshot_not_found",
f"Agent config snapshot {effective_snapshot_id} not found.",
f"Agent config snapshot {snapshot_id} not found.",
)
session.expunge(binding)
@@ -33,6 +33,7 @@ from dify_agent.layers.shell import (
DifyShellCliToolConfig,
DifyShellEnvVarConfig,
DifyShellLayerConfig,
DifyShellSandboxConfig,
DifyShellSecretRefConfig,
)
from dify_agent.protocol import CreateRunRequest, DeferredToolResultsPayload
@@ -135,8 +136,6 @@ class WorkflowAgentRuntimeBuildContext:
binding: WorkflowAgentNodeBinding
agent: Agent
snapshot: AgentConfigSnapshot
binding_id: str
backend_binding_ref: str
# Stage 4 §7 / D-4: 0 for the first run, then incremented per retry. Drives the
# idempotency key so the backend treats each retry as a fresh request.
attempt: int = 0
@@ -252,7 +251,6 @@ class WorkflowAgentRuntimeRequestBuilder:
agent_mode=self._agent_backend_agent_mode(context.dify_context.invoke_from),
invoke_from=cast(DifyExecutionContextInvokeFrom, context.dify_context.invoke_from.value),
),
backend_binding_ref=context.backend_binding_ref,
agent_soul_prompt=soul_prompt or None,
workflow_node_job_prompt=workflow_job_prompt,
user_prompt=user_prompt,
@@ -741,6 +739,7 @@ class WorkflowAgentRuntimeRequestBuilder:
def build_shell_layer_config(agent_soul: AgentSoulConfig) -> DifyShellLayerConfig:
"""Map Agent Soul shell-adjacent fields into the Agent backend shell config."""
sandbox_config = _plain_mapping(agent_soul.sandbox.config)
return DifyShellLayerConfig(
cli_tools=[
tool
@@ -751,6 +750,12 @@ def build_shell_layer_config(agent_soul: AgentSoulConfig) -> DifyShellLayerConfi
secret_refs=[
secret for secret in (_shell_secret_ref(item) for item in agent_soul.env.secret_refs) if secret is not None
],
sandbox=DifyShellSandboxConfig(
provider=agent_soul.sandbox.provider,
config=sandbox_config,
)
if agent_soul.sandbox.provider or sandbox_config
else None,
)
@@ -0,0 +1,126 @@
"""Workflow terminal layer that retires Agent backend sessions asynchronously."""
from __future__ import annotations
import logging
from typing import override
from clients.agent_backend import AgentBackendSessionCleanupPayload
from core.workflow.system_variables import SystemVariableKey, get_system_text
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.graph_events import (
GraphEngineEvent,
GraphRunAbortedEvent,
GraphRunFailedEvent,
GraphRunPartialSucceededEvent,
GraphRunSucceededEvent,
)
from tasks.agent_backend_session_cleanup_task import cleanup_workflow_agent_runtime_session
from .session_store import StoredWorkflowAgentSession, WorkflowAgentRuntimeSessionStore
logger = logging.getLogger(__name__)
class WorkflowAgentSessionCleanupLayer(GraphEngineLayer):
"""Retire workflow-owned Agent runtime sessions when the workflow ends.
Workflow termination is a product-lifecycle boundary: once the run reaches a
terminal graph event, the local session row must no longer be resumable. The
actual Agent backend cleanup is therefore dispatched asynchronously with the
persisted snapshot/specs payload, while the local row is marked CLEANED
immediately afterwards regardless of enqueue outcome.
"""
_TERMINAL_EVENTS = (
GraphRunSucceededEvent,
GraphRunPartialSucceededEvent,
GraphRunFailedEvent,
GraphRunAbortedEvent,
)
def __init__(self, *, session_store: WorkflowAgentRuntimeSessionStore) -> None:
super().__init__()
self._session_store = session_store
@override
def on_graph_start(self) -> None:
return
@override
def on_event(self, event: GraphEngineEvent) -> None:
if not isinstance(event, self._TERMINAL_EVENTS):
return
workflow_run_id = get_system_text(
self.graph_runtime_state.variable_pool,
SystemVariableKey.WORKFLOW_EXECUTION_ID,
)
if not workflow_run_id:
logger.warning("Skipping workflow Agent session cleanup: workflow_run_id is missing.")
return
for stored_session in self._session_store.list_active_sessions(workflow_run_id=workflow_run_id):
self._cleanup_session(stored_session)
@override
def on_graph_end(self, error: Exception | None) -> None:
return
def _cleanup_session(self, stored_session: StoredWorkflowAgentSession) -> None:
scope = stored_session.scope
try:
if stored_session.runtime_layer_specs:
payload = AgentBackendSessionCleanupPayload(
session_snapshot=stored_session.session_snapshot,
runtime_layer_specs=stored_session.runtime_layer_specs,
idempotency_key=f"{scope.workflow_run_id}:{scope.node_id}:{scope.binding_id}:agent-session-cleanup",
metadata={
"tenant_id": scope.tenant_id,
"app_id": scope.app_id,
"workflow_id": scope.workflow_id,
"workflow_run_id": scope.workflow_run_id,
"node_id": scope.node_id,
"node_execution_id": scope.node_execution_id,
"binding_id": scope.binding_id,
"agent_id": scope.agent_id,
"agent_config_snapshot_id": scope.agent_config_snapshot_id,
"previous_agent_backend_run_id": stored_session.backend_run_id,
},
)
cleanup_workflow_agent_runtime_session.delay(payload.model_dump(mode="json"))
else:
logger.warning(
"Skipping workflow Agent backend cleanup enqueue: no runtime_layer_specs persisted. "
"workflow_run_id=%s node_id=%s agent_id=%s",
scope.workflow_run_id,
scope.node_id,
scope.agent_id,
)
except Exception:
logger.warning(
"Failed to enqueue workflow Agent backend cleanup: "
"workflow_run_id=%s node_id=%s agent_id=%s previous_run_id=%s",
scope.workflow_run_id,
scope.node_id,
scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
finally:
try:
self._session_store.mark_cleaned(scope=scope, backend_run_id=stored_session.backend_run_id)
except Exception:
logger.warning(
"Failed to retire workflow Agent runtime session after cleanup enqueue: "
"workflow_run_id=%s node_id=%s agent_id=%s previous_run_id=%s",
scope.workflow_run_id,
scope.node_id,
scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
def build_workflow_agent_session_cleanup_layer() -> WorkflowAgentSessionCleanupLayer:
"""Wire the cleanup layer with the standard workflow-owned session store."""
return WorkflowAgentSessionCleanupLayer(session_store=WorkflowAgentRuntimeSessionStore())
+146 -227
View File
@@ -1,32 +1,31 @@
"""Workflow Agent participant persistence keyed by node execution."""
from __future__ import annotations
import json
import time
from dataclasses import dataclass
from dataclasses import dataclass, field
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.protocol import RuntimeLayerSpec
from pydantic import TypeAdapter
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
from libs.datetime_utils import naive_utc_now
from models.agent import (
AgentConfigVersionKind,
AgentWorkingResourceStatus,
AgentWorkspace,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
)
from models.workflow import WorkflowNodeExecutionModel
from services.agent.workspace_service import (
AgentWorkspaceNotFoundError,
AgentWorkspaceService,
WorkspaceOwnerScope,
AgentRuntimeSessionOwnerType,
WorkflowAgentRuntimeSession,
WorkflowAgentRuntimeSessionStatus,
)
_CALLER_VISIBILITY_ATTEMPTS = 60
_CALLER_VISIBILITY_INTERVAL_SECONDS = 0.05
_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec])
def _serialize_specs(specs: list[RuntimeLayerSpec]) -> str:
return _SPECS_ADAPTER.dump_json(specs).decode()
def _deserialize_specs(value: str | None) -> list[RuntimeLayerSpec]:
if not value:
return []
return _SPECS_ADAPTER.validate_json(value)
@dataclass(frozen=True, slots=True)
@@ -37,251 +36,171 @@ class WorkflowAgentSessionScope:
workflow_run_id: str | None
node_id: str
node_execution_id: str
workflow_agent_binding_id: str
binding_id: str
agent_id: str
agent_config_snapshot_id: str
@property
def workspace_owner(self) -> WorkspaceOwnerScope:
return WorkspaceOwnerScope(
tenant_id=self.tenant_id,
app_id=self.app_id,
owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN,
owner_id=self.workflow_run_id or self.node_execution_id,
owner_scope_key=f"{self.node_id}:{self.workflow_agent_binding_id}",
)
@dataclass(frozen=True, slots=True)
class StoredWorkflowAgentSession:
scope: WorkflowAgentSessionScope
binding_id: str
workspace_id: str
backend_binding_ref: str
session_snapshot: CompositorSessionSnapshot | None
session_snapshot: CompositorSessionSnapshot
backend_run_id: str | None
runtime_layer_specs: list[RuntimeLayerSpec] = field(default_factory=list)
# ENG-637: set while the session is paused on a dify.ask_human deferred call.
pending_form_id: str | None = None
pending_tool_call_id: str | None = None
class WorkflowAgentWorkspaceStore:
"""Load or create the participant named by a node execution caller row."""
class WorkflowAgentRuntimeSessionStore:
"""Stores Agent backend session snapshots for workflow Agent node re-entry."""
def load_existing_node_execution_scope(
self,
*,
tenant_id: str,
app_id: str,
workflow_id: str,
workflow_run_id: str | None,
node_id: str,
node_execution_id: str,
) -> WorkflowAgentSessionScope | None:
"""Return the generation pinned by an existing node execution participant."""
def load_active_snapshot(self, scope: WorkflowAgentSessionScope) -> CompositorSessionSnapshot | None:
stored = self.load_active_session(scope)
return stored.session_snapshot if stored is not None else None
def load_active_session(self, scope: WorkflowAgentSessionScope) -> StoredWorkflowAgentSession | None:
"""Load the active session row including any pending ask_human correlation."""
if scope.workflow_run_id is None:
return None
with session_factory.create_session() as session:
execution = self._load_execution_by_identity(
session=session,
tenant_id=tenant_id,
app_id=app_id,
workflow_id=workflow_id,
workflow_run_id=workflow_run_id,
node_id=node_id,
node_execution_id=node_execution_id,
row = session.scalar(
select(WorkflowAgentRuntimeSession).where(
WorkflowAgentRuntimeSession.tenant_id == scope.tenant_id,
WorkflowAgentRuntimeSession.workflow_run_id == scope.workflow_run_id,
WorkflowAgentRuntimeSession.node_id == scope.node_id,
WorkflowAgentRuntimeSession.binding_id == scope.binding_id,
WorkflowAgentRuntimeSession.agent_id == scope.agent_id,
WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE,
)
)
binding_id = execution.agent_workspace_binding_id
if binding_id is None:
if row is None:
return None
process_data = execution.process_data_dict
if not isinstance(process_data, dict):
raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is invalid")
workflow_agent_binding_id = process_data.get("workflow_agent_binding_id")
if not isinstance(workflow_agent_binding_id, str):
raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is missing")
owner_scope = WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=app_id,
owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN,
owner_id=workflow_run_id or node_execution_id,
owner_scope_key=f"{node_id}:{workflow_agent_binding_id}",
)
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=tenant_id,
binding_id=binding_id,
expected_owner_scope=owner_scope,
)
if binding is None or binding.agent_config_version_kind != AgentConfigVersionKind.SNAPSHOT:
raise AgentWorkspaceNotFoundError("Workflow node participant Binding is unavailable")
return WorkflowAgentSessionScope(
tenant_id=tenant_id,
app_id=app_id,
workflow_id=workflow_id,
workflow_run_id=workflow_run_id,
node_id=node_id,
node_execution_id=node_execution_id,
workflow_agent_binding_id=workflow_agent_binding_id,
agent_id=binding.agent_id,
agent_config_snapshot_id=binding.agent_config_version_id,
return StoredWorkflowAgentSession(
scope=scope,
session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
backend_run_id=row.backend_run_id,
runtime_layer_specs=_deserialize_specs(row.composition_layer_specs),
pending_form_id=row.pending_form_id,
pending_tool_call_id=row.pending_tool_call_id,
)
def load_or_create_node_execution_session(
self, scope: WorkflowAgentSessionScope, *, home_snapshot_id: str
) -> StoredWorkflowAgentSession:
def list_active_sessions(self, *, workflow_run_id: str) -> list[StoredWorkflowAgentSession]:
with session_factory.create_session() as session:
execution = self._load_execution(session=session, scope=scope)
process_data = execution.process_data_dict
if process_data is None:
process_data = {}
if not isinstance(process_data, dict):
raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is invalid")
stored_workflow_binding_id = process_data.get("workflow_agent_binding_id")
if stored_workflow_binding_id is not None and stored_workflow_binding_id != scope.workflow_agent_binding_id:
raise AgentWorkspaceNotFoundError("Workflow node execution caller identity does not match")
binding_id = execution.agent_workspace_binding_id
if binding_id is None:
binding = AgentWorkspaceService.create_binding(
session=session,
scope=scope.workspace_owner,
agent_id=scope.agent_id,
base_home_snapshot_id=home_snapshot_id,
agent_config_version_id=scope.agent_config_snapshot_id,
agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT,
rows = session.scalars(
select(WorkflowAgentRuntimeSession).where(
WorkflowAgentRuntimeSession.workflow_run_id == workflow_run_id,
WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE,
)
execution.agent_workspace_binding_id = binding.id
execution.process_data = json.dumps(
{
**process_data,
"workflow_agent_binding_id": scope.workflow_agent_binding_id,
},
ensure_ascii=False,
).all()
return [
StoredWorkflowAgentSession(
scope=WorkflowAgentSessionScope(
tenant_id=row.tenant_id,
app_id=row.app_id,
# These columns are nullable on the unified runtime-session
# table (workflow_run ⊕ conversation owner), but are always
# populated for a workflow-owned row; coerce for the typed scope.
workflow_id=row.workflow_id or "",
workflow_run_id=row.workflow_run_id,
node_id=row.node_id or "",
node_execution_id=row.node_execution_id or "",
binding_id=row.binding_id or "",
agent_id=row.agent_id,
agent_config_snapshot_id=row.agent_config_snapshot_id or "",
),
session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
backend_run_id=row.backend_run_id,
runtime_layer_specs=_deserialize_specs(row.composition_layer_specs),
)
session.commit()
else:
if stored_workflow_binding_id is None:
raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is missing")
resolved_binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=scope.tenant_id,
binding_id=binding_id,
expected_owner_scope=scope.workspace_owner,
)
if resolved_binding is None or resolved_binding.agent_id != scope.agent_id:
raise AgentWorkspaceNotFoundError("Workflow node participant Binding is unavailable")
binding = resolved_binding
AgentWorkspaceService.validate_binding_generation(
binding,
base_home_snapshot_id=home_snapshot_id,
agent_config_version_id=scope.agent_config_snapshot_id,
agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT,
)
return self._stored(scope, binding)
for row in rows
]
def save_active_snapshot(
self,
*,
scope: WorkflowAgentSessionScope,
binding_id: str,
backend_run_id: str,
snapshot: CompositorSessionSnapshot | None,
runtime_layer_specs: list[RuntimeLayerSpec],
pending_form_id: str | None = None,
pending_tool_call_id: str | None = None,
) -> None:
if snapshot is None:
if scope.workflow_run_id is None or snapshot is None:
return
AgentWorkspaceService.save_binding_session_snapshot(
tenant_id=scope.tenant_id,
binding_id=binding_id,
session_snapshot=snapshot.model_dump_json(),
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
def retire_workflow_run(self, *, tenant_id: str, app_id: str, workflow_run_id: str) -> list[str]:
"""Retire active Workspaces, commit, and return active or already-retired IDs for collection."""
retired: list[str] = []
snapshot_json = snapshot.model_dump_json()
specs_json = _serialize_specs(runtime_layer_specs)
with session_factory.create_session() as session:
workspaces = session.scalars(
select(AgentWorkspace).where(
AgentWorkspace.tenant_id == tenant_id,
AgentWorkspace.app_id == app_id,
AgentWorkspace.owner_type == AgentWorkspaceOwnerType.WORKFLOW_RUN,
AgentWorkspace.owner_id == workflow_run_id,
AgentWorkspace.status.in_((AgentWorkingResourceStatus.ACTIVE, AgentWorkingResourceStatus.RETIRED)),
row = session.scalar(
select(WorkflowAgentRuntimeSession).where(
WorkflowAgentRuntimeSession.tenant_id == scope.tenant_id,
WorkflowAgentRuntimeSession.workflow_run_id == scope.workflow_run_id,
WorkflowAgentRuntimeSession.node_id == scope.node_id,
WorkflowAgentRuntimeSession.binding_id == scope.binding_id,
WorkflowAgentRuntimeSession.agent_id == scope.agent_id,
)
).all()
for workspace in workspaces:
if workspace.status == AgentWorkingResourceStatus.RETIRED:
retired.append(workspace.id)
continue
workspace_id = AgentWorkspaceService.retire_workspace(
session=session,
tenant_id=tenant_id,
workspace_id=workspace.id,
)
if row is None:
row = WorkflowAgentRuntimeSession(
tenant_id=scope.tenant_id,
app_id=scope.app_id,
owner_type=AgentRuntimeSessionOwnerType.WORKFLOW_RUN,
workflow_id=scope.workflow_id,
workflow_run_id=scope.workflow_run_id,
node_id=scope.node_id,
node_execution_id=scope.node_execution_id,
binding_id=scope.binding_id,
agent_id=scope.agent_id,
agent_config_snapshot_id=scope.agent_config_snapshot_id,
backend_run_id=backend_run_id,
session_snapshot=snapshot_json,
composition_layer_specs=specs_json,
status=WorkflowAgentRuntimeSessionStatus.ACTIVE,
pending_form_id=pending_form_id,
pending_tool_call_id=pending_tool_call_id,
)
if workspace_id is not None:
retired.append(workspace_id)
session.add(row)
else:
row.node_execution_id = scope.node_execution_id
row.agent_config_snapshot_id = scope.agent_config_snapshot_id
row.backend_run_id = backend_run_id
row.session_snapshot = snapshot_json
row.composition_layer_specs = specs_json
row.status = WorkflowAgentRuntimeSessionStatus.ACTIVE
row.cleaned_at = None
# Set (or clear, when omitted) the ask_human pause correlation.
row.pending_form_id = pending_form_id
row.pending_tool_call_id = pending_tool_call_id
session.commit()
return retired
@staticmethod
def _load_execution(*, session: Session, scope: WorkflowAgentSessionScope) -> WorkflowNodeExecutionModel:
return WorkflowAgentWorkspaceStore._load_execution_by_identity(
session=session,
tenant_id=scope.tenant_id,
app_id=scope.app_id,
workflow_id=scope.workflow_id,
workflow_run_id=scope.workflow_run_id,
node_id=scope.node_id,
node_execution_id=scope.node_execution_id,
)
def mark_cleaned(self, *, scope: WorkflowAgentSessionScope, backend_run_id: str | None = None) -> None:
if scope.workflow_run_id is None:
return
@staticmethod
def _load_execution_by_identity(
*,
session: Session,
tenant_id: str,
app_id: str,
workflow_id: str,
workflow_run_id: str | None,
node_id: str,
node_execution_id: str,
) -> WorkflowNodeExecutionModel:
"""Wait briefly for the already-emitted node-start event to persist its caller row."""
stmt = select(WorkflowNodeExecutionModel).where(
WorkflowNodeExecutionModel.id == node_execution_id,
WorkflowNodeExecutionModel.tenant_id == tenant_id,
WorkflowNodeExecutionModel.app_id == app_id,
WorkflowNodeExecutionModel.workflow_id == workflow_id,
WorkflowNodeExecutionModel.node_id == node_id,
WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id,
)
for attempt in range(_CALLER_VISIBILITY_ATTEMPTS):
execution = session.scalar(stmt)
if execution is not None:
return execution
if attempt < _CALLER_VISIBILITY_ATTEMPTS - 1:
time.sleep(_CALLER_VISIBILITY_INTERVAL_SECONDS)
raise AgentWorkspaceNotFoundError("Workflow node execution caller is unavailable")
@staticmethod
def _stored(scope: WorkflowAgentSessionScope, binding: AgentWorkspaceBinding) -> StoredWorkflowAgentSession:
snapshot = (
CompositorSessionSnapshot.model_validate_json(binding.session_snapshot)
if binding.session_snapshot
else None
)
return StoredWorkflowAgentSession(
scope=scope,
binding_id=binding.id,
workspace_id=binding.workspace_id,
backend_binding_ref=binding.backend_binding_ref,
session_snapshot=snapshot,
pending_form_id=binding.pending_form_id,
pending_tool_call_id=binding.pending_tool_call_id,
)
with session_factory.create_session() as session:
row = session.scalar(
select(WorkflowAgentRuntimeSession).where(
WorkflowAgentRuntimeSession.tenant_id == scope.tenant_id,
WorkflowAgentRuntimeSession.workflow_run_id == scope.workflow_run_id,
WorkflowAgentRuntimeSession.node_id == scope.node_id,
WorkflowAgentRuntimeSession.binding_id == scope.binding_id,
WorkflowAgentRuntimeSession.agent_id == scope.agent_id,
WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE,
)
)
if row is None:
return
if backend_run_id is not None:
row.backend_run_id = backend_run_id
row.status = WorkflowAgentRuntimeSessionStatus.CLEANED
row.cleaned_at = naive_utc_now()
session.commit()
__all__ = ["StoredWorkflowAgentSession", "WorkflowAgentSessionScope", "WorkflowAgentWorkspaceStore"]
__all__ = [
"StoredWorkflowAgentSession",
"WorkflowAgentRuntimeSessionStore",
"WorkflowAgentSessionScope",
]
@@ -1,89 +0,0 @@
"""Retire Workflow Agent Workspaces when the Workflow Run terminates."""
from __future__ import annotations
import logging
from typing import override
from core.app.entities.app_invoke_entities import DifyRunContext
from core.workflow.nodes.agent_v2.session_store import WorkflowAgentWorkspaceStore
from core.workflow.system_variables import SystemVariableKey, get_system_text
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.graph_events import (
GraphEngineEvent,
GraphRunAbortedEvent,
GraphRunFailedEvent,
GraphRunPartialSucceededEvent,
GraphRunSucceededEvent,
)
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
logger = logging.getLogger(__name__)
class WorkflowAgentWorkspaceRetirementLayer(GraphEngineLayer):
"""Synchronously retire run Workspaces, then enqueue physical collection."""
_TERMINAL_EVENTS = (
GraphRunSucceededEvent,
GraphRunPartialSucceededEvent,
GraphRunFailedEvent,
GraphRunAbortedEvent,
)
def __init__(
self,
*,
dify_run_context: DifyRunContext,
) -> None:
super().__init__()
self._dify_run_context = dify_run_context
@override
def on_graph_start(self) -> None:
return
@override
def on_event(self, event: GraphEngineEvent) -> None:
if not isinstance(event, self._TERMINAL_EVENTS):
return
workflow_run_id = get_system_text(
self.graph_runtime_state.variable_pool,
SystemVariableKey.WORKFLOW_EXECUTION_ID,
)
if not workflow_run_id:
logger.warning("Skipping Workflow Agent Workspace retirement: workflow_run_id is missing")
return
try:
workspace_ids = WorkflowAgentWorkspaceStore().retire_workflow_run(
tenant_id=self._dify_run_context.tenant_id,
app_id=self._dify_run_context.app_id,
workflow_run_id=workflow_run_id,
)
except Exception:
logger.exception(
"Failed to retire Workflow Agent Workspaces",
extra={
"tenant_id": self._dify_run_context.tenant_id,
"app_id": self._dify_run_context.app_id,
"workflow_run_id": workflow_run_id,
},
)
return
enqueue_agent_resource_collection(
tenant_id=self._dify_run_context.tenant_id,
workspace_ids=workspace_ids,
)
@override
def on_graph_end(self, error: Exception | None) -> None:
return
def build_workflow_agent_workspace_retirement_layer(
*, dify_run_context: DifyRunContext
) -> WorkflowAgentWorkspaceRetirementLayer:
return WorkflowAgentWorkspaceRetirementLayer(dify_run_context=dify_run_context)
__all__ = ["WorkflowAgentWorkspaceRetirementLayer", "build_workflow_agent_workspace_retirement_layer"]
-1
View File
@@ -166,7 +166,6 @@ def init_app(app: DifyApp) -> Celery:
imports = [
"tasks.async_workflow_tasks", # trigger workers
"tasks.collect_agent_resources_task", # retired Agent resource collection
"tasks.trigger_processing_tasks", # async trigger processing
"tasks.generate_summary_index_task", # summary index generation
"tasks.regenerate_summary_index_task", # summary index regeneration
@@ -277,12 +277,6 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
logger.exception("Failed to dual-write node execution to SQL database: id=%s", execution.id)
# Don't raise - LogStore write succeeded, SQL is just a backup
@override
def save_synchronously(self, execution: WorkflowNodeExecution) -> None:
"""Create the SQL caller row required by Agent v2 participant ownership."""
self.sql_repository.save_synchronously(execution)
@override
def save_execution_data(self, execution: WorkflowNodeExecution) -> None:
"""
+1 -1
View File
@@ -34,7 +34,7 @@ class _SessionResponseSource[SourceT]:
self._session = session
def __getattr__(self, name: str) -> object:
return getattr(self._source, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self._source, name) # noqa: no-new-getattr response adapter delegates model fields
class _FeedbackResponseSource(_SessionResponseSource[MessageFeedback]):
+1 -1
View File
@@ -227,7 +227,7 @@ class DatasetDetailResponseSource:
return self.dataset.get_total_available_documents(session=self.session)
def __getattr__(self, name: str) -> Any:
return getattr(self.dataset, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self.dataset, name) # noqa: no-new-getattr response adapter delegates model fields
def dataset_detail_response_source(dataset: Any, *, session: Session) -> DatasetDetailResponseSource:
+1 -1
View File
@@ -90,7 +90,7 @@ class DocumentWithSession:
return self.document.get_doc_metadata_details(session=self.session)
def __getattr__(self, name: str) -> Any:
return getattr(self.document, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self.document, name) # noqa: no-new-getattr response adapter delegates model fields
def document_response(document: Document, *, session: Session) -> DocumentResponse:
+1 -5
View File
@@ -289,11 +289,7 @@ UUIDStr = Annotated[str, AfterValidator(_strict_uuid)]
def alphanumeric(value: str):
# check if the value is alphanumeric and underlined
# Use re.fullmatch instead of re.match to reject trailing newlines.
# In Python, '$' matches at end-of-string OR just before a trailing newline,
# so re.match accepts "tool_name\n". re.fullmatch requires the entire
# string to match. Regression for #39666 (sibling of #39234 / #39548).
if re.fullmatch(r"^[a-zA-Z0-9_]+$", value):
if re.match(r"^[a-zA-Z0-9_]+$", value):
return value
raise ValueError(f"{value} is not a valid alphanumeric value")
@@ -1,64 +0,0 @@
"""add agent home snapshot ledger
Revision ID: 2f39536b3feb
Revises: 6f5a9c2d8e1b
Create Date: 2026-07-21 22:51:07.268658
"""
from alembic import op
import models as models
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision = '2f39536b3feb'
down_revision = '6f5a9c2d8e1b'
branch_labels = None
depends_on = None
def upgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.create_table('agent_home_snapshots',
sa.Column('id', models.types.StringUUID(), nullable=False),
sa.Column('tenant_id', models.types.StringUUID(), nullable=False),
sa.Column('agent_id', models.types.StringUUID(), nullable=False),
sa.Column('snapshot_ref', sa.String(length=255), nullable=False),
sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.PrimaryKeyConstraint('id', name='agent_home_snapshot_pkey')
)
with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op:
batch_op.create_index('agent_home_snapshot_tenant_agent_idx', ['tenant_id', 'agent_id'], unique=False)
with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op:
batch_op.add_column(sa.Column('home_snapshot_id', models.types.StringUUID(), nullable=False))
with op.batch_alter_table('agent_config_snapshots', schema=None) as batch_op:
batch_op.add_column(sa.Column('home_snapshot_id', models.types.StringUUID(), nullable=False))
with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op:
batch_op.add_column(sa.Column('home_snapshot_id', models.types.StringUUID(), nullable=False))
batch_op.drop_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), postgresql_where='(conversation_id IS NOT NULL)')
batch_op.create_index('agent_runtime_session_conversation_scope_unique', ['tenant_id', 'conversation_id', 'agent_id', 'agent_config_snapshot_id', 'home_snapshot_id'], unique=True, postgresql_where=sa.text('conversation_id IS NOT NULL'))
# ### end Alembic commands ###
def downgrade():
# ### commands auto generated by Alembic - please adjust! ###
with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op:
batch_op.drop_index('agent_runtime_session_conversation_scope_unique', postgresql_where=sa.text('conversation_id IS NOT NULL'))
batch_op.create_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), ['tenant_id', 'conversation_id', 'agent_id', 'agent_config_snapshot_id'], unique=True, postgresql_where='(conversation_id IS NOT NULL)')
batch_op.drop_column('home_snapshot_id')
with op.batch_alter_table('agent_config_snapshots', schema=None) as batch_op:
batch_op.drop_column('home_snapshot_id')
with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op:
batch_op.drop_column('home_snapshot_id')
with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op:
batch_op.drop_index('agent_home_snapshot_tenant_agent_idx')
op.drop_table('agent_home_snapshots')
# ### end Alembic commands ###
@@ -1,153 +0,0 @@
"""replace agent runtime sessions with workspaces and bindings
Revision ID: f6e4c5686857
Revises: 2f39536b3feb
Create Date: 2026-07-23 02:03:05.641638
"""
from alembic import op
import models as models
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision = 'f6e4c5686857'
down_revision = '2f39536b3feb'
branch_labels = None
depends_on = None
def upgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.create_table('agent_workspace_bindings',
sa.Column('tenant_id', models.types.StringUUID(), nullable=False),
sa.Column('app_id', models.types.StringUUID(), nullable=False),
sa.Column('workspace_id', models.types.StringUUID(), nullable=False),
sa.Column('agent_id', models.types.StringUUID(), nullable=False),
sa.Column('base_home_snapshot_id', models.types.StringUUID(), nullable=False),
sa.Column('agent_config_version_id', models.types.StringUUID(), nullable=False),
sa.Column('agent_config_version_kind', sa.String(length=32), nullable=False),
sa.Column('backend_binding_ref', sa.String(length=255), nullable=False),
sa.Column('session_snapshot', models.types.LongText(), nullable=True),
sa.Column('status', sa.String(length=32), server_default='active', nullable=False),
sa.Column('retired_at', sa.DateTime(), nullable=True),
sa.Column('pending_form_id', models.types.StringUUID(), nullable=True),
sa.Column('pending_tool_call_id', sa.String(length=255), nullable=True),
sa.Column('id', models.types.StringUUID(), nullable=False),
sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.Column('updated_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.PrimaryKeyConstraint('id', name='agent_workspace_binding_pkey')
)
with op.batch_alter_table('agent_workspace_bindings', schema=None) as batch_op:
batch_op.create_index('agent_workspace_binding_agent_status_idx', ['tenant_id', 'agent_id', 'status'], unique=False)
batch_op.create_index('agent_workspace_binding_status_retired_idx', ['status', 'retired_at'], unique=False)
batch_op.create_index('agent_workspace_binding_workspace_status_idx', ['tenant_id', 'workspace_id', 'status'], unique=False)
op.create_table('agent_workspaces',
sa.Column('tenant_id', models.types.StringUUID(), nullable=False),
sa.Column('app_id', models.types.StringUUID(), nullable=False),
sa.Column('owner_type', sa.String(length=32), nullable=False),
sa.Column('owner_id', models.types.StringUUID(), nullable=False),
sa.Column('owner_scope_key', sa.String(length=255), nullable=False),
sa.Column('backend_workspace_ref', sa.String(length=255), nullable=False),
sa.Column('status', sa.String(length=32), server_default='active', nullable=False),
sa.Column('active_guard', sa.SmallInteger(), server_default='1', nullable=True),
sa.Column('retired_at', sa.DateTime(), nullable=True),
sa.Column('id', models.types.StringUUID(), nullable=False),
sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.Column('updated_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.PrimaryKeyConstraint('id', name='agent_workspace_pkey')
)
with op.batch_alter_table('agent_workspaces', schema=None) as batch_op:
batch_op.create_index('agent_workspace_owner_active_unique', ['tenant_id', 'owner_type', 'owner_id', 'owner_scope_key', 'active_guard'], unique=True)
batch_op.create_index('agent_workspace_status_retired_idx', ['status', 'retired_at'], unique=False)
batch_op.create_index('agent_workspace_tenant_app_status_idx', ['tenant_id', 'app_id', 'status'], unique=False)
batch_op.create_index('agent_workspace_tenant_status_idx', ['tenant_id', 'status'], unique=False)
with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op:
batch_op.drop_index(batch_op.f('agent_runtime_session_backend_run_idx'))
batch_op.drop_index(batch_op.f('agent_runtime_session_conversation_lookup_idx'))
batch_op.drop_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), postgresql_where='(conversation_id IS NOT NULL)')
batch_op.drop_index(batch_op.f('agent_runtime_session_workflow_lookup_idx'))
batch_op.drop_index(batch_op.f('agent_runtime_session_workflow_scope_unique'), postgresql_where='(workflow_run_id IS NOT NULL)')
op.drop_table('agent_runtime_sessions')
with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op:
batch_op.add_column(sa.Column('status', sa.String(length=32), server_default='active', nullable=False))
batch_op.add_column(sa.Column('retired_at', sa.DateTime(), nullable=True))
batch_op.create_index('agent_home_snapshot_status_retired_idx', ['status', 'retired_at'], unique=False)
with op.batch_alter_table('conversations', schema=None) as batch_op:
batch_op.add_column(sa.Column('agent_workspace_binding_id', models.types.StringUUID(), nullable=True))
with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op:
batch_op.add_column(sa.Column('agent_workspace_binding_id', models.types.StringUUID(), nullable=True))
with op.batch_alter_table('workflow_node_executions', schema=None) as batch_op:
batch_op.add_column(sa.Column('agent_workspace_binding_id', models.types.StringUUID(), nullable=True))
# ### end Alembic commands ###
def downgrade():
# ### commands auto generated by Alembic - please adjust! ###
with op.batch_alter_table('workflow_node_executions', schema=None) as batch_op:
batch_op.drop_column('agent_workspace_binding_id')
with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op:
batch_op.drop_column('agent_workspace_binding_id')
with op.batch_alter_table('conversations', schema=None) as batch_op:
batch_op.drop_column('agent_workspace_binding_id')
with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op:
batch_op.drop_index('agent_home_snapshot_status_retired_idx')
batch_op.drop_column('retired_at')
batch_op.drop_column('status')
op.create_table('agent_runtime_sessions',
sa.Column('id', sa.UUID(), server_default=sa.text('uuidv7()'), autoincrement=False, nullable=False),
sa.Column('tenant_id', sa.UUID(), autoincrement=False, nullable=False),
sa.Column('app_id', sa.UUID(), autoincrement=False, nullable=False),
sa.Column('owner_type', sa.VARCHAR(length=32), autoincrement=False, nullable=False),
sa.Column('agent_id', sa.UUID(), autoincrement=False, nullable=False),
sa.Column('backend_run_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True),
sa.Column('session_snapshot', sa.TEXT(), autoincrement=False, nullable=False),
sa.Column('workflow_id', sa.UUID(), autoincrement=False, nullable=True),
sa.Column('workflow_run_id', sa.UUID(), autoincrement=False, nullable=True),
sa.Column('node_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True),
sa.Column('node_execution_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True),
sa.Column('binding_id', sa.UUID(), autoincrement=False, nullable=True),
sa.Column('agent_config_snapshot_id', sa.UUID(), autoincrement=False, nullable=True),
sa.Column('composition_layer_specs', sa.TEXT(), autoincrement=False, nullable=False),
sa.Column('conversation_id', sa.UUID(), autoincrement=False, nullable=True),
sa.Column('status', sa.VARCHAR(length=32), server_default=sa.text("'active'::character varying"), autoincrement=False, nullable=False),
sa.Column('cleaned_at', postgresql.TIMESTAMP(), autoincrement=False, nullable=True),
sa.Column('created_at', postgresql.TIMESTAMP(), server_default=sa.text('CURRENT_TIMESTAMP'), autoincrement=False, nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(), server_default=sa.text('CURRENT_TIMESTAMP'), autoincrement=False, nullable=False),
sa.Column('pending_form_id', sa.UUID(), autoincrement=False, nullable=True),
sa.Column('pending_tool_call_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True),
sa.Column('home_snapshot_id', sa.UUID(), autoincrement=False, nullable=False),
sa.PrimaryKeyConstraint('id', name=op.f('agent_runtime_session_pkey'))
)
with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op:
batch_op.create_index(batch_op.f('agent_runtime_session_workflow_scope_unique'), ['tenant_id', 'workflow_run_id', 'node_id', 'binding_id', 'agent_id'], unique=True, postgresql_where='(workflow_run_id IS NOT NULL)')
batch_op.create_index(batch_op.f('agent_runtime_session_workflow_lookup_idx'), ['tenant_id', 'workflow_run_id', 'node_id', 'status'], unique=False)
batch_op.create_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), ['tenant_id', 'conversation_id', 'agent_id', 'agent_config_snapshot_id', 'home_snapshot_id'], unique=True, postgresql_where='(conversation_id IS NOT NULL)')
batch_op.create_index(batch_op.f('agent_runtime_session_conversation_lookup_idx'), ['tenant_id', 'conversation_id', 'status'], unique=False)
batch_op.create_index(batch_op.f('agent_runtime_session_backend_run_idx'), ['backend_run_id'], unique=False)
with op.batch_alter_table('agent_workspaces', schema=None) as batch_op:
batch_op.drop_index('agent_workspace_tenant_status_idx')
batch_op.drop_index('agent_workspace_tenant_app_status_idx')
batch_op.drop_index('agent_workspace_status_retired_idx')
batch_op.drop_index('agent_workspace_owner_active_unique')
op.drop_table('agent_workspaces')
with op.batch_alter_table('agent_workspace_bindings', schema=None) as batch_op:
batch_op.drop_index('agent_workspace_binding_workspace_status_idx')
batch_op.drop_index('agent_workspace_binding_status_retired_idx')
batch_op.drop_index('agent_workspace_binding_agent_status_idx')
op.drop_table('agent_workspace_bindings')
# ### end Alembic commands ###
+10 -12
View File
@@ -15,22 +15,21 @@ from .agent import (
AgentConfigRevision,
AgentConfigRevisionOperation,
AgentConfigSnapshot,
AgentConfigVersionKind,
AgentDebugConversation,
AgentDriveFile,
AgentDriveFileKind,
AgentHomeSnapshot,
AgentIconType,
AgentKind,
AgentRuntimeSession,
AgentRuntimeSessionOwnerType,
AgentRuntimeSessionStatus,
AgentScope,
AgentSource,
AgentStatus,
AgentWorkingResourceStatus,
AgentWorkspace,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
WorkflowAgentBindingType,
WorkflowAgentNodeBinding,
WorkflowAgentRuntimeSession,
WorkflowAgentRuntimeSessionStatus,
)
from .api_based_extension import APIBasedExtension, APIBasedExtensionPoint
from .comment import (
@@ -165,20 +164,17 @@ __all__ = [
"AgentConfigRevision",
"AgentConfigRevisionOperation",
"AgentConfigSnapshot",
"AgentConfigVersionKind",
"AgentDebugConversation",
"AgentDriveFile",
"AgentDriveFileKind",
"AgentHomeSnapshot",
"AgentIconType",
"AgentKind",
"AgentRuntimeSession",
"AgentRuntimeSessionOwnerType",
"AgentRuntimeSessionStatus",
"AgentScope",
"AgentSource",
"AgentStatus",
"AgentWorkingResourceStatus",
"AgentWorkspace",
"AgentWorkspaceBinding",
"AgentWorkspaceOwnerType",
"ApiRequest",
"ApiToken",
"ApiToolProvider",
@@ -275,6 +271,8 @@ __all__ = [
"Workflow",
"WorkflowAgentBindingType",
"WorkflowAgentNodeBinding",
"WorkflowAgentRuntimeSession",
"WorkflowAgentRuntimeSessionStatus",
"WorkflowAppLog",
"WorkflowAppLogCreatedFrom",
"WorkflowArchiveLog",
+104 -114
View File
@@ -116,29 +116,35 @@ class WorkflowAgentBindingType(StrEnum):
INLINE_AGENT = "inline_agent"
class AgentWorkingResourceStatus(StrEnum):
"""Product lifecycle state for a persistent working-environment resource."""
class AgentRuntimeSessionStatus(StrEnum):
"""Lifecycle state of an Agent backend session snapshot.
Owner-agnostic: applies both to workflow Agent Node runs (owner =
workflow_run) and to Agent App conversations (owner = conversation).
"""
# Snapshot can be reused by a later Agent run in the same session.
ACTIVE = "active"
RETIRED = "retired"
# Snapshot has been retired and must not be submitted to Agent backend again.
CLEANED = "cleaned"
class AgentWorkspaceOwnerType(StrEnum):
"""Product scope that owns a Workspace."""
class AgentRuntimeSessionOwnerType(StrEnum):
"""Which product surface owns an Agent runtime session row."""
# Owned by one workflow Agent Node execution scope.
WORKFLOW_RUN = "workflow_run"
# Owned by one Agent App conversation (multi-turn chat).
CONVERSATION = "conversation"
BUILD_DRAFT = "build_draft"
class AgentConfigVersionKind(StrEnum):
SNAPSHOT = "snapshot"
DRAFT = "draft"
BUILD_DRAFT = "build_draft"
# Back-compat alias: the workflow lifecycle code (shipped in PR #36724) imports
# the old name. Kept so unifying the table does not churn that path.
WorkflowAgentRuntimeSessionStatus = AgentRuntimeSessionStatus
class Agent(DefaultFieldsMixin, Base):
"""Agent Soul and source lineage; ``AgentWorkspaceBinding.id`` identifies each materialized participant."""
"""Workspace-scoped Agent identity used by Agent Roster and workflow-only agents."""
__tablename__ = "agents"
__table_args__ = (
@@ -215,42 +221,14 @@ class Agent(DefaultFieldsMixin, Base):
archived_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
class AgentHomeSnapshot(Base):
"""Append-only mapping from one Agent-owned Home identity to its backend ref.
Product tables reference ``id``. ``snapshot_ref`` remains an opaque
deployment-specific handle and is only consumed at Dify Agent boundaries.
Snapshot bytes and ``snapshot_ref`` are immutable. Lifecycle metadata can
transition ACTIVE -> RETIRED; successful physical collection deletes row.
"""
__tablename__ = "agent_home_snapshots"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="agent_home_snapshot_pkey"),
Index("agent_home_snapshot_tenant_agent_idx", "tenant_id", "agent_id"),
Index("agent_home_snapshot_status_retired_idx", "status", "retired_at"),
)
id: Mapped[str] = mapped_column(StringUUID, default=lambda: str(uuidv7()))
tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
snapshot_ref: Mapped[str] = mapped_column(String(255), nullable=False)
status: Mapped[AgentWorkingResourceStatus] = mapped_column(
EnumText(AgentWorkingResourceStatus, length=32),
nullable=False,
default=AgentWorkingResourceStatus.ACTIVE,
server_default=AgentWorkingResourceStatus.ACTIVE.value,
)
retired_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, server_default=func.current_timestamp())
class AgentDebugConversation(DefaultFieldsMixin, Base):
"""Current console Conversation pointer for one account and draft surface.
"""Per-account, per-draft console debug conversation for an Agent App.
This row owns no Binding or runtime. A Preview Conversation holds its
CONVERSATION Binding pointer, while a DEBUG_BUILD AgentConfigDraft holds its
BUILD_DRAFT Binding pointer.
Agent App preview state must be isolated by editor account. The Agent row is
shared by everyone in the workspace, so this table owns the user-specific
conversation pointers used by console debug chat. ``draft`` is the Preview
conversation and ``debug_build`` is the Build conversation; they must never
share persisted messages or runtime sessions.
"""
__tablename__ = "agent_debug_conversations"
@@ -281,11 +259,7 @@ class AgentDebugConversation(DefaultFieldsMixin, Base):
class AgentConfigDraft(DefaultFieldsMixin, Base):
"""Editable Agent Soul draft separated from immutable published snapshots.
A DEBUG_BUILD draft owns its materialized participant through
``agent_workspace_binding_id``. Normal drafts leave that pointer unset.
"""
"""Editable Agent Soul draft separated from immutable published snapshots."""
__tablename__ = "agent_config_drafts"
__table_args__ = (
@@ -307,8 +281,6 @@ class AgentConfigDraft(DefaultFieldsMixin, Base):
account_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
draft_owner_key: Mapped[str] = mapped_column(String(255), nullable=False, default="")
base_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
home_snapshot_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
agent_workspace_binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
config_snapshot: Mapped[Any] = mapped_column(JSONModelColumn(AgentSoulConfig), nullable=False)
created_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
updated_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
@@ -344,7 +316,6 @@ class AgentConfigSnapshot(DefaultFieldsMixin, Base):
agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
version: Mapped[int] = mapped_column(sa.Integer, nullable=False)
config_snapshot: Mapped[Any] = mapped_column(JSONModelColumn(AgentSoulConfig), nullable=False)
home_snapshot_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
summary: Mapped[str | None] = mapped_column(LongText, nullable=True)
version_note: Mapped[str | None] = mapped_column(LongText, nullable=True)
created_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
@@ -461,83 +432,102 @@ class WorkflowAgentNodeBinding(DefaultFieldsMixin, Base):
return dict(self.node_job_config)
class AgentWorkspace(DefaultFieldsMixin, Base):
"""Mutable Workspace owned by one product scope, independent of Agents."""
class AgentRuntimeSession(DefaultFieldsMixin, Base):
"""Persisted Agent backend session snapshot, owner-agnostic.
__tablename__ = "agent_workspaces"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="agent_workspace_pkey"),
Index(
"agent_workspace_owner_active_unique",
"tenant_id",
"owner_type",
"owner_id",
"owner_scope_key",
"active_guard",
unique=True,
),
Index("agent_workspace_tenant_status_idx", "tenant_id", "status"),
Index("agent_workspace_tenant_app_status_idx", "tenant_id", "app_id", "status"),
Index("agent_workspace_status_retired_idx", "status", "retired_at"),
)
One unified table serves both owners (decision Q2):
- workflow Agent Node runs: ``owner_type = workflow_run``; the
``workflow_id / workflow_run_id / node_id / binding_id /
agent_config_snapshot_id / composition_layer_specs`` columns are set.
- Agent App conversations: ``owner_type = conversation``; the
``conversation_id`` column is set and the workflow columns stay NULL.
Runtime state is scoped by ``agent_config_snapshot_id``. For published
web/API runs this points to an immutable AgentConfigSnapshot; for console
debugger/build runs it points to the editable AgentConfigDraft row.
tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
app_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
owner_type: Mapped[AgentWorkspaceOwnerType] = mapped_column(
EnumText(AgentWorkspaceOwnerType, length=32), nullable=False
)
owner_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
owner_scope_key: Mapped[str] = mapped_column(String(255), nullable=False)
backend_workspace_ref: Mapped[str] = mapped_column(String(255), nullable=False)
status: Mapped[AgentWorkingResourceStatus] = mapped_column(
EnumText(AgentWorkingResourceStatus, length=32),
nullable=False,
default=AgentWorkingResourceStatus.ACTIVE,
server_default=AgentWorkingResourceStatus.ACTIVE.value,
)
active_guard: Mapped[int | None] = mapped_column(sa.SmallInteger, nullable=True, default=1, server_default="1")
retired_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
class AgentWorkspaceBinding(DefaultFieldsMixin, Base):
"""One materialized Agent participant and session attached to a Workspace.
All resource IDs are logical associations rather than database foreign
keys, so RETIRED rows can outlive their Workspace or base Home Snapshot.
``agent_id`` identifies the source Agent Soul; this row's ``id`` identifies
the participant and its private Materialized Home.
The snapshot is runtime state returned by Agent backend, kept separate from
Agent Soul snapshots and workflow node-job config.
"""
__tablename__ = "agent_workspace_bindings"
__tablename__ = "agent_runtime_sessions"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="agent_workspace_binding_pkey"),
Index("agent_workspace_binding_workspace_status_idx", "tenant_id", "workspace_id", "status"),
Index("agent_workspace_binding_agent_status_idx", "tenant_id", "agent_id", "status"),
Index("agent_workspace_binding_status_retired_idx", "status", "retired_at"),
sa.PrimaryKeyConstraint("id", name="agent_runtime_session_pkey"),
# Workflow owner uniqueness (partial: only rows with a workflow_run_id).
Index(
"agent_runtime_session_workflow_scope_unique",
"tenant_id",
"workflow_run_id",
"node_id",
"binding_id",
"agent_id",
unique=True,
postgresql_where=sa.text("workflow_run_id IS NOT NULL"),
),
# Conversation owner uniqueness (partial: only rows with a conversation_id).
Index(
"agent_runtime_session_conversation_scope_unique",
"tenant_id",
"conversation_id",
"agent_id",
"agent_config_snapshot_id",
unique=True,
postgresql_where=sa.text("conversation_id IS NOT NULL"),
),
Index(
"agent_runtime_session_workflow_lookup_idx",
"tenant_id",
"workflow_run_id",
"node_id",
"status",
),
Index(
"agent_runtime_session_conversation_lookup_idx",
"tenant_id",
"conversation_id",
"status",
),
Index("agent_runtime_session_backend_run_idx", "backend_run_id"),
)
tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
app_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
workspace_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
owner_type: Mapped[AgentRuntimeSessionOwnerType] = mapped_column(
EnumText(AgentRuntimeSessionOwnerType, length=32), nullable=False
)
agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
base_home_snapshot_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
agent_config_version_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
agent_config_version_kind: Mapped[AgentConfigVersionKind] = mapped_column(
EnumText(AgentConfigVersionKind, length=32), nullable=False
)
backend_binding_ref: Mapped[str] = mapped_column(String(255), nullable=False)
session_snapshot: Mapped[str | None] = mapped_column(LongText, nullable=True)
status: Mapped[AgentWorkingResourceStatus] = mapped_column(
EnumText(AgentWorkingResourceStatus, length=32),
backend_run_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
session_snapshot: Mapped[str] = mapped_column(LongText, nullable=False)
# Workflow-owner columns (NULL for conversation owner).
workflow_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
workflow_run_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
node_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
node_execution_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
agent_config_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
# JSON-encoded list of non-sensitive runtime layer specs ({name, type, deps,
# config}). The persisted schema keeps its original name because the sandbox
# refactor intentionally avoids a storage migration.
composition_layer_specs: Mapped[str] = mapped_column(LongText, nullable=False, server_default="[]")
# Conversation-owner column (NULL for workflow owner).
conversation_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
status: Mapped[AgentRuntimeSessionStatus] = mapped_column(
EnumText(AgentRuntimeSessionStatus, length=32),
nullable=False,
default=AgentWorkingResourceStatus.ACTIVE,
server_default=AgentWorkingResourceStatus.ACTIVE.value,
default=AgentRuntimeSessionStatus.ACTIVE,
)
retired_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
cleaned_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
# ENG-637: when a run pauses for a dify.ask_human deferred call, these link
# the session to the awaiting HITL form and the deferred tool_call_id, so a
# resumed node can map the submitted form back into deferred_tool_results.
# Both NULL whenever the session is not paused on human input.
pending_form_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
pending_tool_call_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
# Back-compat alias for the shipped workflow lifecycle code (PR #36724).
WorkflowAgentRuntimeSession = AgentRuntimeSession
class AgentDriveFileKind(StrEnum):
"""Kind of existing file record an agent-drive KV entry points at."""
+2 -10
View File
@@ -1117,14 +1117,14 @@ class ExporleBanner(TypeBase):
status: Mapped[BannerStatus] = mapped_column(
EnumText(BannerStatus, length=255),
nullable=False,
server_default=sa.text("'enabled'"),
server_default=sa.text("'enabled'::character varying"),
default=BannerStatus.ENABLED,
)
created_at: Mapped[datetime] = mapped_column(
sa.DateTime, nullable=False, server_default=func.current_timestamp(), init=False
)
language: Mapped[str] = mapped_column(
String(255), nullable=False, server_default=sa.text("'en-US'"), default="en-US"
String(255), nullable=False, server_default=sa.text("'en-US'::character varying"), default="en-US"
)
@@ -1160,13 +1160,6 @@ class OAuthProviderApp(TypeBase):
class Conversation(Base):
"""Conversation state, including the exact Agent participant when applicable.
``agent_workspace_binding_id`` is a logical pointer rather than a foreign
key because retired Binding ledger rows may be collected before the
conversation history is deleted.
"""
__tablename__ = "conversations"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="conversation_pkey"),
@@ -1188,7 +1181,6 @@ class Conversation(Base):
id: Mapped[str] = mapped_column(StringUUID, default=lambda: str(uuid4()))
app_id = mapped_column(StringUUID, nullable=False)
app_model_config_id = mapped_column(StringUUID, nullable=True)
agent_workspace_binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
model_provider = mapped_column(String(255), nullable=True)
override_model_configs = mapped_column(LongText)
model_id = mapped_column(String(255), nullable=True)
-1
View File
@@ -1031,7 +1031,6 @@ class WorkflowNodeExecutionModel(Base): # This model is expected to have `offlo
node_id: Mapped[str] = mapped_column(String(255))
node_type: Mapped[str] = mapped_column(String(255))
title: Mapped[str] = mapped_column(String(255))
agent_workspace_binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True)
inputs: Mapped[str | None] = mapped_column(LongText)
process_data: Mapped[str | None] = mapped_column(LongText)
outputs: Mapped[str | None] = mapped_column(LongText)
+111 -50
View File
@@ -964,6 +964,12 @@ Stop a running Agent App chat message generation
| ---- | ---------- | ----------- | -------- | ------ |
| agent_id | path | | Yes | string (uuid) |
#### Request Body
| Required | Schema |
| -------- | ------ |
| No | **application/json**: [AgentDebugConversationRefreshPayload](#agentdebugconversationrefreshpayload)<br> |
#### Responses
| Code | Description | Schema |
@@ -1255,8 +1261,7 @@ Get basic information for an Agent App conversation sandbox
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| agent_id | path | Agent ID | Yes | string (uuid) |
| caller_id | query | Agent App caller ID | Yes | string |
| caller_type | query | | Yes | string, <br>**Available values:** "build_draft", "conversation" |
| conversation_id | query | Agent App conversation ID | Yes | string |
#### Responses
@@ -1272,8 +1277,7 @@ List a directory in an Agent App conversation sandbox
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| agent_id | path | Agent ID | Yes | string (uuid) |
| caller_id | query | Agent App caller ID | Yes | string |
| caller_type | query | | Yes | string, <br>**Available values:** "build_draft", "conversation" |
| conversation_id | query | Agent App conversation ID | Yes | string |
| path | query | Directory path relative to the sandbox workspace | No | string, <br>**Default:** . |
#### Responses
@@ -1290,8 +1294,7 @@ Read a text/binary preview file in an Agent App conversation sandbox
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| agent_id | path | Agent ID | Yes | string (uuid) |
| caller_id | query | Agent App caller ID | Yes | string |
| caller_type | query | | Yes | string, <br>**Available values:** "build_draft", "conversation" |
| conversation_id | query | Agent App conversation ID | Yes | string |
| path | query | File path relative to the sandbox workspace | Yes | string |
#### Responses
@@ -1669,23 +1672,6 @@ Create a new application
| 200 | Import confirmed | **application/json**: [Import](#import)<br> |
| 400 | Import failed | **application/json**: [Import](#import)<br> |
### [GET] /apps/recent
**Return the lightweight app cards needed by the Explore home page**
Get recently modified apps for the home Continue Work section
#### Parameters
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| limit | query | Number of recently modified apps to return (1-8) | No | integer, <br>**Default:** 8 |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [RecentAppListResponse](#recentapplistresponse)<br> |
### [GET] /apps/starred
Get applications starred by the current account
@@ -3804,7 +3790,7 @@ List a directory in a workflow Agent node sandbox
| app_id | path | Application ID | Yes | string (uuid) |
| node_id | path | Workflow Agent node ID | Yes | string |
| workflow_run_id | path | Workflow run ID | Yes | string (uuid) |
| node_execution_id | query | Workflow node execution ID | Yes | string |
| node_execution_id | query | Optional workflow node execution ID. When omitted, the latest active session for the node is used. | No | string |
| path | query | Directory path relative to the sandbox workspace | No | string, <br>**Default:** . |
#### Responses
@@ -3823,7 +3809,7 @@ Read a text/binary preview file in a workflow Agent node sandbox
| app_id | path | Application ID | Yes | string (uuid) |
| node_id | path | Workflow Agent node ID | Yes | string |
| workflow_run_id | path | Workflow run ID | Yes | string (uuid) |
| node_execution_id | query | Workflow node execution ID | Yes | string |
| node_execution_id | query | Optional workflow node execution ID. When omitted, the latest active session for the node is used. | No | string |
| path | query | File path relative to the sandbox workspace | Yes | string |
#### Responses
@@ -10581,6 +10567,13 @@ Update a plugin endpoint
| ---- | ----------- | ------ |
| 200 | Model providers retrieved successfully | **application/json**: [ModelProviderListResponse](#modelproviderlistresponse)<br> |
### [GET] /workspaces/current/model-providers/summary
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Model provider summaries retrieved successfully | **application/json**: [ModelProviderSummaryListResponse](#modelprovidersummarylistresponse)<br> |
### [GET] /workspaces/current/model-providers/{provider}/checkout-url
#### Parameters
@@ -11126,6 +11119,19 @@ Returns permission flags that control workspace features like member invitations
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [PluginInstallTaskStartResponse](#plugininstalltaskstartresponse)<br> |
### [GET] /workspaces/current/plugin/installed-ids
#### Parameters
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| category | query | Plugin category to include | Yes | string, <br>**Available values:** "agent-strategy", "datasource", "extension", "model", "tool", "trigger" |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [PluginInstalledIdsResponse](#plugininstalledidsresponse)<br> |
### [GET] /workspaces/current/plugin/list
#### Parameters
@@ -11384,8 +11390,11 @@ Returns permission flags that control workspace features like member invitations
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| language | query | Language used for localized label and description search | No | string, <br>**Available values:** "en_US", "ja_JP", "pt_BR", "zh_Hans", <br>**Default:** en_US |
| page | query | Page number | No | integer, <br>**Default:** 1 |
| page_size | query | Page size (1-256) | No | integer, <br>**Default:** 256 |
| query | query | Case-insensitive search query | No | string |
| tags | query | Match any plugin tag | No | [ string ] |
| category | path | | Yes | string |
#### Responses
@@ -13940,6 +13949,12 @@ Stable Agent Soul reference to one normalized skill archive.
| date | string | | Yes |
| message_count | integer | | Yes |
#### AgentDebugConversationRefreshPayload
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| draft_type | [AgentConfigDraftType](#agentconfigdrafttype) | Agent draft surface whose conversation should be refreshed | No |
#### AgentDebugConversationRefreshResponse
| Name | Type | Description | Required |
@@ -14657,8 +14672,7 @@ section may be empty, which is how callers express "no knowledge layer".
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| caller_id | string | Agent App caller ID | Yes |
| caller_type | string, <br>**Available values:** "build_draft", "conversation" | *Enum:* `"build_draft"`, `"conversation"` | Yes |
| conversation_id | string | Agent App conversation ID | Yes |
| path | string | File path relative to the sandbox workspace | Yes |
#### AgentScope
@@ -19316,6 +19330,16 @@ Enum class for model property key.
| ---- | ---- | ----------- | -------- |
| ModelPropertyKey | string | Enum class for model property key. | |
#### ModelProviderCustomConfigurationSummaryResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| available_credentials | [ [CredentialConfiguration](#credentialconfiguration) ] | | Yes |
| current_credential_id | string | | No |
| current_credential_name | string | | No |
| current_credential_usable | boolean | | Yes |
| status | [CustomConfigurationStatus](#customconfigurationstatus) | | Yes |
#### ModelProviderListResponse
| Name | Type | Description | Required |
@@ -19328,6 +19352,49 @@ Enum class for model property key.
| ---- | ---- | ----------- | -------- |
| payment_link | string | | Yes |
#### ModelProviderPluginSummaryResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| installation_id | string | | Yes |
| plugin_id | string | | Yes |
| plugin_unique_identifier | string | | Yes |
| runtime_type | string | | Yes |
| source | [PluginInstallationSource](#plugininstallationsource) | | Yes |
| version | string | | Yes |
#### ModelProviderSummaryListResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| data | [ [ModelProviderSummaryResponse](#modelprovidersummaryresponse) ] | | Yes |
| plugins | object | | Yes |
#### ModelProviderSummaryResponse
Fields required to render the collapsed model-provider list.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| configurate_methods | [ [ConfigurateMethod](#configuratemethod) ] | | Yes |
| custom_configuration | [ModelProviderCustomConfigurationSummaryResponse](#modelprovidercustomconfigurationsummaryresponse) | | Yes |
| description | [I18nObject](#i18nobject) | | No |
| icon_small | [I18nObject](#i18nobject) | | No |
| icon_small_dark | [I18nObject](#i18nobject) | | No |
| is_configured | boolean | | Yes |
| label | [I18nObject](#i18nobject) | | Yes |
| plugin_id | string | | Yes |
| preferred_provider_type | [ProviderType](#providertype) | | Yes |
| provider | string | | Yes |
| supported_model_types | [ [ModelType](#modeltype) ] | | Yes |
| system_configuration | [ModelProviderSystemConfigurationSummaryResponse](#modelprovidersystemconfigurationsummaryresponse) | | Yes |
#### ModelProviderSystemConfigurationSummaryResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| enabled | boolean | | Yes |
#### ModelSelectorScope
| Name | Type | Description | Required |
@@ -20316,8 +20383,11 @@ Shared permission levels for resources (datasets, credentials, etc.)
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| language | string, <br>**Available values:** "en_US", "ja_JP", "pt_BR", "zh_Hans", <br>**Default:** en_US | Language used for localized label and description search<br>*Enum:* `"en_US"`, `"ja_JP"`, `"pt_BR"`, `"zh_Hans"` | No |
| page | integer, <br>**Default:** 1 | Page number | No |
| page_size | integer, <br>**Default:** 256 | Page size (1-256) | No |
| query | string | Case-insensitive search query | No |
| tags | [ string ] | Match any plugin tag | No |
#### PluginCategoryListResponse
@@ -20518,6 +20588,18 @@ Shared permission levels for resources (datasets, credentials, etc.)
| ---- | ---- | ----------- | -------- |
| plugins | [ [PluginInstallationItemResponse](#plugininstallationitemresponse) ] | | Yes |
#### PluginInstalledIdsQuery
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| category | [PluginCategory](#plugincategory) | Plugin category to include | Yes |
#### PluginInstalledIdsResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| plugin_ids | [ string ] | | Yes |
#### PluginListResponse
| Name | Type | Description | Required |
@@ -21027,28 +21109,6 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs.
| result | string | | Yes |
| updated_at | integer | | Yes |
#### RecentAppListResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| data | [ [RecentAppResponse](#recentappresponse) ] | | Yes |
#### RecentAppResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| author_name | string | | No |
| icon | string | | No |
| icon_background | string | | No |
| icon_type | [IconType](#icontype) | | No |
| icon_url | string | | Yes |
| id | string | | Yes |
| maintainer | string | | No |
| mode | string, <br>**Available values:** "advanced-chat", "agent-chat", "chat", "completion", "workflow" | *Enum:* `"advanced-chat"`, `"agent-chat"`, `"chat"`, `"completion"`, `"workflow"` | Yes |
| name | string | | Yes |
| permission_keys | [ string ] | | No |
| updated_at | integer | | Yes |
#### RecommendedAppDetailNullableResponse
| Name | Type | Description | Required |
@@ -21339,6 +21399,7 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| session_id | string | | Yes |
| workspace_cwd | string | | Yes |
#### SandboxListResponse
@@ -23158,7 +23219,7 @@ How a workflow node is bound to an Agent.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| node_execution_id | string | Workflow node execution ID | Yes |
| node_execution_id | string | Optional workflow node execution ID. When omitted, the latest active session for the node is used. | No |
| path | string | File path relative to the sandbox workspace | Yes |
#### WorkflowAppLogPaginationResponse
@@ -1,5 +1,3 @@
"""Unit tests for Aliyun trace utility transformations and database lookups."""
import json
from collections.abc import Mapping
from typing import Any, cast
@@ -27,13 +25,11 @@ from dify_trace_aliyun.utils import (
serialize_json_data,
)
from opentelemetry.trace import Link, StatusCode
from sqlalchemy.orm import Session
from core.rag.models.document import Document
from graphon.entities import WorkflowNodeExecution
from graphon.enums import WorkflowNodeExecutionStatus
from models import EndUser
from models.enums import EndUserType
def test_get_user_id_from_message_data_no_end_user(monkeypatch: pytest.MonkeyPatch):
@@ -44,40 +40,35 @@ def test_get_user_id_from_message_data_no_end_user(monkeypatch: pytest.MonkeyPat
assert get_user_id_from_message_data(message_data) == "account_id"
@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True)
def test_get_user_id_from_message_data_with_end_user(monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None:
def test_get_user_id_from_message_data_with_end_user(monkeypatch: pytest.MonkeyPatch):
message_data = MagicMock()
message_data.from_account_id = "account_id"
message_data.from_end_user_id = "end_user_id"
end_user_data = EndUser(
id="end_user_id",
tenant_id="tenant_id",
app_id="app_id",
type=EndUserType.BROWSER,
session_id="session_id",
)
sqlite3_session.add(end_user_data)
sqlite3_session.commit()
end_user_data = MagicMock(spec=EndUser)
end_user_data.session_id = "session_id"
mock_session = MagicMock()
mock_session.get.return_value = end_user_data
from dify_trace_aliyun.utils import db
monkeypatch.setattr(db, "session", sqlite3_session)
monkeypatch.setattr(db, "session", mock_session)
assert get_user_id_from_message_data(message_data) == "session_id"
@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True)
def test_get_user_id_from_message_data_end_user_not_found(
monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session
) -> None:
def test_get_user_id_from_message_data_end_user_not_found(monkeypatch: pytest.MonkeyPatch):
message_data = MagicMock()
message_data.from_account_id = "account_id"
message_data.from_end_user_id = "end_user_id"
mock_session = MagicMock()
mock_session.get.return_value = None
from dify_trace_aliyun.utils import db
monkeypatch.setattr(db, "session", sqlite3_session)
monkeypatch.setattr(db, "session", mock_session)
assert get_user_id_from_message_data(message_data) == "account_id"
@@ -1,8 +1,5 @@
"""Unit tests for LangSmith trace translation with SQLite-backed lookups."""
import collections
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import override
from unittest.mock import MagicMock
@@ -14,7 +11,6 @@ from dify_trace_langsmith.entities.langsmith_trace_entity import (
LangSmithRunUpdateModel,
)
from dify_trace_langsmith.langsmith_trace import LangSmithDataTrace
from sqlalchemy.orm import Session
from core.ops.entities.trace_entity import (
DatasetRetrievalTraceInfo,
@@ -28,7 +24,6 @@ from core.ops.entities.trace_entity import (
)
from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey
from models import EndUser
from models.enums import EndUserType
def _dt() -> datetime:
@@ -113,8 +108,7 @@ def test_trace_dispatch(trace_instance, monkeypatch: pytest.MonkeyPatch):
mocks["generate_name_trace"].assert_called_once_with(info)
@pytest.mark.parametrize("sqlite3_session", [()], indirect=True)
def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None:
def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch):
# Setup trace info
workflow_data = MagicMock()
workflow_data.created_at = _dt()
@@ -143,10 +137,10 @@ def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3
workflow_data=workflow_data,
)
monkeypatch.setattr(
"dify_trace_langsmith.langsmith_trace.db",
SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session),
)
# Mock dependencies
mock_session = MagicMock()
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: mock_session)
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine"))
# Mock node executions
node_llm = MagicMock()
@@ -234,10 +228,7 @@ def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3
assert call_args[4].run_type == LangSmithRunType.retriever
@pytest.mark.parametrize("sqlite3_session", [()], indirect=True)
def test_workflow_trace_no_start_time(
trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session
) -> None:
def test_workflow_trace_no_start_time(trace_instance, monkeypatch: pytest.MonkeyPatch):
workflow_data = MagicMock()
workflow_data.created_at = _dt()
workflow_data.finished_at = _dt() + timedelta(seconds=1)
@@ -265,10 +256,9 @@ def test_workflow_trace_no_start_time(
workflow_data=workflow_data,
)
monkeypatch.setattr(
"dify_trace_langsmith.langsmith_trace.db",
SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session),
)
mock_session = MagicMock()
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: mock_session)
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine"))
repo = MagicMock()
repo.get_by_workflow_execution.return_value = []
mock_factory = MagicMock()
@@ -281,10 +271,7 @@ def test_workflow_trace_no_start_time(
assert trace_instance.add_run.called
@pytest.mark.parametrize("sqlite3_session", [()], indirect=True)
def test_workflow_trace_missing_app_id(
trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session
) -> None:
def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.MonkeyPatch):
trace_info = MagicMock(spec=WorkflowTraceInfo)
trace_info.trace_id = "trace-1"
trace_info.message_id = None
@@ -300,17 +287,15 @@ def test_workflow_trace_missing_app_id(
trace_info.workflow_run_outputs = {}
trace_info.error = ""
monkeypatch.setattr(
"dify_trace_langsmith.langsmith_trace.db",
SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session),
)
mock_session = MagicMock()
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: mock_session)
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine"))
with pytest.raises(ValueError, match="No app_id found in trace_info metadata"):
trace_instance.workflow_trace(trace_info)
@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True)
def test_message_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None:
def test_message_trace(trace_instance, monkeypatch: pytest.MonkeyPatch):
message_data = MagicMock()
message_data.id = "msg-1"
message_data.from_account_id = "acc-1"
@@ -336,19 +321,10 @@ def test_message_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_
message_file_data=MagicMock(url="file-url"),
)
end_user = EndUser(
id="end-user-1",
tenant_id="tenant-1",
app_id="app-1",
type=EndUserType.BROWSER,
session_id="session-id-123",
)
sqlite3_session.add(end_user)
sqlite3_session.commit()
monkeypatch.setattr(
"dify_trace_langsmith.langsmith_trace.db",
SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session),
)
# Mock EndUser lookup
mock_end_user = MagicMock(spec=EndUser)
mock_end_user.session_id = "session-id-123"
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db.session.get", lambda model, pk: mock_end_user)
trace_instance.add_run = MagicMock()
@@ -545,13 +521,9 @@ def test_update_run_error(trace_instance):
trace_instance.update_run(update_data)
@pytest.mark.parametrize("sqlite3_session", [()], indirect=True)
def test_workflow_trace_usage_extraction_error(
trace_instance,
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
sqlite3_session: Session,
) -> None:
trace_instance, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
):
workflow_data = MagicMock()
workflow_data.created_at = _dt()
workflow_data.finished_at = _dt() + timedelta(seconds=1)
@@ -604,10 +576,8 @@ def test_workflow_trace_usage_extraction_error(
mock_factory = MagicMock()
mock_factory.create_workflow_node_execution_repository.return_value = repo
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.DifyCoreRepositoryFactory", mock_factory)
monkeypatch.setattr(
"dify_trace_langsmith.langsmith_trace.db",
SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session),
)
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: MagicMock())
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine"))
monkeypatch.setattr(trace_instance, "get_service_account_with_tenant", lambda app_id: MagicMock())
trace_instance.add_run = MagicMock()
@@ -674,11 +644,9 @@ def _make_workflow_trace_info(
)
def _patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session: Session) -> None:
monkeypatch.setattr(
"dify_trace_langsmith.langsmith_trace.db",
SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session),
)
def _patch_workflow_trace_deps(monkeypatch, trace_instance):
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: MagicMock())
monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine"))
repo = MagicMock()
repo.get_by_workflow_execution.return_value = []
factory = MagicMock()
@@ -688,17 +656,14 @@ def _patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session: Ses
trace_instance.add_run = MagicMock()
@pytest.mark.parametrize("sqlite3_session", [()], indirect=True)
def test_workflow_trace_id_uses_message_id_not_external(
trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session
) -> None:
def test_workflow_trace_id_uses_message_id_not_external(trace_instance, monkeypatch: pytest.MonkeyPatch):
"""Chatflow with external trace_id: LangSmith trace_id must be message_id, not external."""
trace_info = _make_workflow_trace_info(
message_id="msg-abc",
workflow_run_id="run-xyz",
trace_id="external-999",
)
_patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session)
_patch_workflow_trace_deps(monkeypatch, trace_instance)
trace_instance.workflow_trace(trace_info)
@@ -712,17 +677,14 @@ def test_workflow_trace_id_uses_message_id_not_external(
assert trace_info.metadata.get("external_trace_id") == "external-999"
@pytest.mark.parametrize("sqlite3_session", [()], indirect=True)
def test_workflow_trace_id_pure_workflow_uses_run_id(
trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session
) -> None:
def test_workflow_trace_id_pure_workflow_uses_run_id(trace_instance, monkeypatch: pytest.MonkeyPatch):
"""Pure workflow (no message_id) with external trace_id: trace_id must be workflow_run_id."""
trace_info = _make_workflow_trace_info(
message_id=None,
workflow_run_id="run-xyz",
trace_id="external-999",
)
_patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session)
_patch_workflow_trace_deps(monkeypatch, trace_instance)
trace_instance.workflow_trace(trace_info)
@@ -316,10 +316,10 @@ class OracleVector(BaseVector):
entities.append(current_entity)
else:
try:
nltk.data.find("tokenizers/punkt_tab")
nltk.data.find("tokenizers/punkt")
nltk.data.find("corpora/stopwords")
except LookupError:
raise LookupError("Unable to find the required NLTK data package: punkt_tab and stopwords")
raise LookupError("Unable to find the required NLTK data package: punkt and stopwords")
e_str = re.sub(r"[^\w ]", "", query)
all_tokens = nltk.word_tokenize(e_str)
stop_words = stopwords.words("english")
+5 -5
View File
@@ -1,12 +1,12 @@
[project]
name = "dify-api"
version = "1.16.1"
version = "1.16.0"
requires-python = "~=3.12.0"
dependencies = [
# Legacy: mature and widely deployed
"bleach>=6.4.0,<7.0.0",
"boto3>=1.43.56,<2.0.0",
"boto3>=1.43.46,<2.0.0",
"celery>=5.6.3,<6.0.0",
"croniter>=6.2.2,<7.0.0",
"dify-agent",
@@ -193,10 +193,10 @@ dev = [
############################################################
storage = [
"azure-storage-blob>=12.30.0,<13.0.0",
"bce-python-sdk==0.9.76",
"bce-python-sdk==0.9.72",
"cos-python-sdk-v5>=1.9.44,<2.0.0",
"esdk-obs-python>=3.26.6,<4.0.0",
"google-cloud-storage>=3.13.0,<4.0.0",
"google-cloud-storage>=3.12.1,<4.0.0",
"opendal==0.46.0",
"oss2>=2.19.1,<3.0.0",
"supabase>=2.31.0,<3.0.0",
@@ -206,7 +206,7 @@ storage = [
############################################################
# [ Tools ] dependency group
############################################################
tools = ["cloudscraper>=1.2.71,<2.0.0", "nltk>=3.10.0,<4.0.0"]
tools = ["cloudscraper>=1.2.71,<2.0.0", "nltk>=3.9.1,<4.0.0"]
############################################################
# [ VDB ] workspace plugins — hollow packages under providers/vdb/*
+32 -365
View File
@@ -9,7 +9,7 @@ from sqlalchemy.sql.elements import ColumnElement
from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot
from libs.helper import to_timestamp
from models import Account, Conversation
from models import Account
from models.agent import (
APP_BACKED_AGENT_SOURCES,
Agent,
@@ -18,15 +18,12 @@ from models.agent import (
AgentConfigRevision,
AgentConfigRevisionOperation,
AgentConfigSnapshot,
AgentConfigVersionKind,
AgentDebugConversation,
AgentDriveFile,
AgentIconType,
AgentKind,
AgentScope,
AgentSource,
AgentStatus,
AgentWorkspaceOwnerType,
WorkflowAgentBindingType,
WorkflowAgentNodeBinding,
)
@@ -38,7 +35,6 @@ from models.workflow import Workflow
from services.agent.agent_soul_state import agent_soul_has_model
from services.agent.composer_validator import ComposerConfigValidator
from services.agent.errors import (
AgentBuildSandboxNotFoundError,
AgentModelNotConfiguredError,
AgentNameConflictError,
AgentNotFoundError,
@@ -46,17 +42,11 @@ from services.agent.errors import (
AgentVersionNotFoundError,
InvalidComposerConfigError,
)
from services.agent.home_snapshot_service import (
AgentHomeSnapshotService,
validate_home_snapshot_binding,
)
from services.agent.knowledge_datasets import (
get_tenant_knowledge_dataset_rows,
list_missing_tenant_knowledge_dataset_ids,
)
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.agent.roster_service import AgentRosterService
from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope
from services.app_service import AppService, CreateAppParams
from services.entities.agent_entities import (
AgentSoulConfig,
@@ -66,7 +56,6 @@ from services.entities.agent_entities import (
ComposerVariant,
WorkflowNodeJobConfig,
)
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
# WorkflowAgentNodeBinding.workflow_version tag for the draft workflow row.
# Mirrors Workflow.version when it is "draft" (see models/workflow.py).
@@ -209,18 +198,6 @@ class AgentComposerService:
binding = cls._get_workflow_binding(
session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id
)
retirement_candidates = (
{binding.agent_id}
if binding is not None
and binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT
and binding.agent_id
and payload.save_strategy
in {
ComposerSaveStrategy.SAVE_AS_NEW_AGENT,
ComposerSaveStrategy.SAVE_TO_ROSTER,
}
else set()
)
match payload.save_strategy:
case ComposerSaveStrategy.NODE_JOB_ONLY:
@@ -255,11 +232,7 @@ class AgentComposerService:
)
case ComposerSaveStrategy.SAVE_TO_ROSTER:
binding = cls._save_to_roster(
session=session,
tenant_id=tenant_id,
account_id=account_id,
binding=binding,
payload=payload,
session=session, tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload
)
session.flush()
@@ -284,17 +257,6 @@ class AgentComposerService:
payload=payload,
agent_id=binding.agent_id,
)
session.commit()
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=tenant_id,
agent_ids=retirement_candidates,
account_id=account_id,
)
enqueue_agent_resource_collection(
tenant_id=tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
return state
@classmethod
@@ -428,10 +390,12 @@ class AgentComposerService:
@classmethod
def _load_agent_composer_for_agent(cls, *, session: Session, tenant_id: str, agent: Agent) -> dict[str, Any]:
draft = cls.get_or_create_normal_agent_draft(
draft = cls._get_or_create_agent_draft(
session=session,
tenant_id=tenant_id,
agent=agent,
draft_type=AgentConfigDraftType.DRAFT,
account_id=None,
created_by=agent.updated_by or agent.created_by,
)
version = cls._get_version_if_present(
@@ -453,34 +417,7 @@ class AgentComposerService:
@classmethod
def save_agent_app_composer(
cls,
*,
session: Session,
tenant_id: str,
app_id: str,
account_id: str,
payload: ComposerSavePayload,
) -> dict[str, Any]:
try:
return cls._save_agent_app_composer_impl(
session=session,
tenant_id=tenant_id,
app_id=app_id,
account_id=account_id,
payload=payload,
)
except IntegrityError as exc:
raise AgentNameConflictError() from exc
@classmethod
def _save_agent_app_composer_impl(
cls,
*,
session: Session,
tenant_id: str,
app_id: str,
account_id: str,
payload: ComposerSavePayload,
cls, *, session: Session, tenant_id: str, app_id: str, account_id: str, payload: ComposerSavePayload
) -> dict[str, Any]:
if payload.variant != ComposerVariant.AGENT_APP:
raise ValueError("Agent App composer endpoint only accepts agent_app variant")
@@ -509,25 +446,11 @@ class AgentComposerService:
updated_by=account_id,
)
session.add(agent)
session.flush()
home_snapshot = AgentHomeSnapshotService.create_initial(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
)
initial_version = cls._create_config_version(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
account_id=account_id,
agent_soul=AgentSoulConfig(),
operation=AgentConfigRevisionOperation.CREATE_VERSION,
version_note=None,
home_snapshot_id=home_snapshot.id,
)
agent.active_config_snapshot_id = initial_version.id
agent.active_config_has_model = False
agent.active_config_is_published = False
try:
session.flush()
except IntegrityError as exc:
session.rollback()
raise AgentNameConflictError() from exc
return cls._save_agent_composer_for_agent(
session=session,
tenant_id=tenant_id,
@@ -565,7 +488,7 @@ class AgentComposerService:
) -> dict[str, Any]:
if payload.agent_soul is None:
raise ValueError("agent_soul is required")
draft = cls._save_agent_draft(
cls._save_agent_draft(
session=session,
tenant_id=tenant_id,
agent=agent,
@@ -580,7 +503,6 @@ class AgentComposerService:
tenant_id=tenant_id,
agent=agent,
agent_soul=payload.agent_soul,
home_snapshot_id=draft.home_snapshot_id,
)
session.flush()
@@ -601,7 +523,6 @@ class AgentComposerService:
tenant_id: str,
agent: Agent,
agent_soul: AgentSoulConfig,
home_snapshot_id: str,
) -> bool:
if not agent.active_config_snapshot_id:
return False
@@ -617,9 +538,7 @@ class AgentComposerService:
if not agent_has_workflow_callable_active_snapshot(session=session, agent=agent):
return False
return home_snapshot_id == active_version.home_snapshot_id and _agent_soul_config_json(
agent_soul
) == _agent_soul_config_json(active_version.config_snapshot_dict)
return _agent_soul_config_json(agent_soul) == _agent_soul_config_json(active_version.config_snapshot_dict)
@classmethod
def publish_agent_app_draft(
@@ -648,11 +567,6 @@ class AgentComposerService:
if not agent_soul_has_model(agent_soul):
raise AgentModelNotConfiguredError()
cls.validate_knowledge_datasets(session=session, tenant_id=tenant_id, agent_soul=agent_soul)
validate_home_snapshot_binding(
session=session,
agent=agent,
home_snapshot_id=draft.home_snapshot_id,
)
version = cls._create_config_version(
session=session,
tenant_id=tenant_id,
@@ -662,7 +576,6 @@ class AgentComposerService:
operation=AgentConfigRevisionOperation.PUBLISH_DRAFT,
version_note=version_note,
previous_snapshot_id=agent.active_config_snapshot_id,
home_snapshot_id=draft.home_snapshot_id,
)
agent.active_config_snapshot_id = version.id
agent.active_config_has_model = agent_soul_has_model(agent_soul)
@@ -682,29 +595,6 @@ class AgentComposerService:
def checkout_agent_app_build_draft(
cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, force: bool = False
) -> dict[str, Any]:
try:
result, retired_binding_id = cls._checkout_agent_app_build_draft_in_transaction(
session=session,
tenant_id=tenant_id,
agent_id=agent_id,
account_id=account_id,
force=force,
)
session.commit()
except Exception:
session.rollback()
raise
if retired_binding_id is not None:
enqueue_agent_resource_collection(
tenant_id=tenant_id,
binding_ids=(retired_binding_id,),
)
return result
@classmethod
def _checkout_agent_app_build_draft_in_transaction(
cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, force: bool
) -> tuple[dict[str, Any], str | None]:
agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id)
normal_draft = cls._get_or_create_agent_draft(
session=session,
@@ -722,23 +612,7 @@ class AgentComposerService:
account_id=account_id,
)
if build_draft is not None and not force:
return cls._serialize_build_draft_state(build_draft), None
retired_binding_id: str | None = None
if build_draft is not None and build_draft.agent_workspace_binding_id is not None:
cls._validate_active_build_draft_binding(
session=session,
tenant_id=tenant_id,
agent=agent,
build_draft=build_draft,
)
retired_binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=tenant_id,
binding_id=build_draft.agent_workspace_binding_id,
)
if retired_binding_id is None:
raise AgentBuildSandboxNotFoundError()
build_draft.agent_workspace_binding_id = None
return cls._serialize_build_draft_state(build_draft)
if build_draft is None:
build_draft = AgentConfigDraft(
tenant_id=tenant_id,
@@ -750,44 +624,10 @@ class AgentComposerService:
)
session.add(build_draft)
build_draft.base_snapshot_id = normal_draft.base_snapshot_id
build_draft.home_snapshot_id = normal_draft.home_snapshot_id
build_draft.config_snapshot = AgentSoulConfig.model_validate(normal_draft.config_snapshot_dict)
build_draft.updated_by = account_id
session.flush()
return cls._serialize_build_draft_state(build_draft), retired_binding_id
@classmethod
def _validate_active_build_draft_binding(
cls,
*,
session: Session,
tenant_id: str,
agent: Agent,
build_draft: AgentConfigDraft,
) -> None:
binding_id = build_draft.agent_workspace_binding_id
runtime_app_id = AgentRosterService.runtime_backing_app_id(agent)
if binding_id is None or runtime_app_id is None:
raise AgentBuildSandboxNotFoundError()
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=tenant_id,
binding_id=binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=runtime_app_id,
owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT,
owner_id=build_draft.id,
),
)
if binding is None or binding.agent_id != agent.id:
raise AgentBuildSandboxNotFoundError()
AgentWorkspaceService.validate_binding_generation(
binding,
base_home_snapshot_id=build_draft.home_snapshot_id,
agent_config_version_id=build_draft.id,
agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT,
)
return cls._serialize_build_draft_state(build_draft)
@classmethod
def load_agent_app_build_draft(
@@ -829,29 +669,6 @@ class AgentComposerService:
def apply_agent_app_build_draft(
cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str
) -> dict[str, Any]:
try:
result, retired_binding_ids = cls._apply_agent_app_build_draft_in_transaction(
session=session,
tenant_id=tenant_id,
agent_id=agent_id,
account_id=account_id,
)
session.commit()
except Exception:
session.rollback()
raise
enqueue_agent_resource_collection(tenant_id=tenant_id, binding_ids=retired_binding_ids)
return result
@classmethod
def _apply_agent_app_build_draft_in_transaction(
cls,
*,
session: Session,
tenant_id: str,
agent_id: str,
account_id: str,
) -> tuple[dict[str, Any], list[str]]:
agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id)
build_draft = cls._get_agent_draft(
session=session,
@@ -863,21 +680,6 @@ class AgentComposerService:
if build_draft is None:
raise AgentVersionNotFoundError()
applied_agent_soul = AgentSoulConfig.model_validate(build_draft.config_snapshot_dict)
ComposerConfigValidator.validate_publish_payload(
ComposerSavePayload(
variant=ComposerVariant.AGENT_APP,
agent_soul=applied_agent_soul,
save_strategy=ComposerSaveStrategy.SAVE_AS_NEW_VERSION,
)
)
cls.validate_knowledge_datasets(session=session, tenant_id=tenant_id, agent_soul=applied_agent_soul)
source_binding_id = build_draft.agent_workspace_binding_id
if source_binding_id is None:
raise AgentBuildSandboxNotFoundError()
home_snapshot = AgentHomeSnapshotService.create_for_build_apply(
session=session,
build_draft=build_draft,
)
normal_draft = cls._save_agent_draft(
session=session,
tenant_id=tenant_id,
@@ -888,145 +690,32 @@ class AgentComposerService:
account_id_for_audit=account_id,
base_snapshot_id=build_draft.base_snapshot_id,
)
retired_binding_ids = cls._retire_normal_preview_bindings(
session=session,
tenant_id=tenant_id,
agent=agent,
normal_draft=normal_draft,
)
normal_draft.home_snapshot_id = home_snapshot.id
agent.active_config_is_published = cls._agent_soul_matches_active_config(
session=session,
tenant_id=tenant_id,
agent=agent,
agent_soul=applied_agent_soul,
home_snapshot_id=home_snapshot.id,
)
agent.updated_by = account_id
retired_binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=tenant_id,
binding_id=source_binding_id,
)
if retired_binding_id is None:
raise AgentBuildSandboxNotFoundError()
retired_binding_ids.append(source_binding_id)
session.delete(build_draft)
return {"result": "success", "draft": cls._serialize_draft(normal_draft)}, retired_binding_ids
@classmethod
def _retire_normal_preview_bindings(
cls,
*,
session: Session,
tenant_id: str,
agent: Agent,
normal_draft: AgentConfigDraft,
) -> list[str]:
"""Retire Preview participants before Build Apply replaces the shared Draft Home."""
mappings = session.scalars(
select(AgentDebugConversation).where(
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent.id,
AgentDebugConversation.draft_type == AgentConfigDraftType.DRAFT,
)
).all()
retired_binding_ids: list[str] = []
for mapping in mappings:
conversation = session.scalar(
select(Conversation).where(
Conversation.id == mapping.conversation_id,
Conversation.app_id == mapping.app_id,
Conversation.is_deleted.is_(False),
)
)
if conversation is None or conversation.agent_workspace_binding_id is None:
continue
binding_id = conversation.agent_workspace_binding_id
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=tenant_id,
binding_id=binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=mapping.app_id,
owner_type=AgentWorkspaceOwnerType.CONVERSATION,
owner_id=conversation.id,
),
)
if binding is None or binding.agent_id != agent.id:
raise AgentWorkspaceNotFoundError("Agent Preview participant Binding is unavailable")
AgentWorkspaceService.validate_binding_generation(
binding,
base_home_snapshot_id=normal_draft.home_snapshot_id,
agent_config_version_id=normal_draft.id,
agent_config_version_kind=AgentConfigVersionKind.DRAFT,
)
retired_binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=tenant_id,
binding_id=binding_id,
)
if retired_binding_id is None:
raise AgentWorkspaceNotFoundError("Agent Preview participant Binding is unavailable")
conversation.agent_workspace_binding_id = None
retired_binding_ids.append(binding_id)
return retired_binding_ids
session.flush()
return {"result": "success", "draft": cls._serialize_draft(normal_draft)}
@classmethod
def discard_agent_app_build_draft(
cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str
) -> dict[str, Any]:
try:
result, retired_binding_id = cls._discard_agent_app_build_draft_in_transaction(
session=session,
tenant_id=tenant_id,
agent_id=agent_id,
account_id=account_id,
)
session.commit()
except Exception:
session.rollback()
raise
if retired_binding_id is not None:
enqueue_agent_resource_collection(
tenant_id=tenant_id,
binding_ids=(retired_binding_id,),
)
return result
@classmethod
def _discard_agent_app_build_draft_in_transaction(
cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str
) -> tuple[dict[str, Any], str | None]:
agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id)
build_draft = cls._get_agent_draft(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
agent_id=agent_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
account_id=account_id,
)
if build_draft is None:
return {"result": "success"}, None
retired_binding_id: str | None = None
if build_draft.agent_workspace_binding_id is not None:
cls._validate_active_build_draft_binding(
session=session,
tenant_id=tenant_id,
agent=agent,
build_draft=build_draft,
)
retired_binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=tenant_id,
binding_id=build_draft.agent_workspace_binding_id,
)
if retired_binding_id is None:
raise AgentBuildSandboxNotFoundError()
session.delete(build_draft)
return {"result": "success"}, retired_binding_id
if build_draft is not None:
session.delete(build_draft)
session.flush()
return {"result": "success"}
@classmethod
def collect_validation_findings(
@@ -1590,12 +1279,6 @@ class AgentComposerService:
binding = cls._require_binding(binding)
if not binding.agent_id or payload.agent_soul is None:
raise ValueError("agent_id and agent_soul are required")
current_snapshot = cls._require_version(
session=session,
tenant_id=tenant_id,
agent_id=binding.agent_id,
version_id=binding.current_snapshot_id,
)
version = cls._create_config_version(
session=session,
tenant_id=tenant_id,
@@ -1604,7 +1287,6 @@ class AgentComposerService:
agent_soul=payload.agent_soul,
operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION,
version_note=payload.version_note,
home_snapshot_id=current_snapshot.home_snapshot_id,
)
agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=binding.agent_id)
agent.active_config_snapshot_id = version.id
@@ -1767,11 +1449,6 @@ class AgentComposerService:
)
session.add(agent)
session.flush()
home_snapshot = AgentHomeSnapshotService.create_initial(
session=session,
tenant_id=tenant_id,
agent_id=agent.id,
)
version = cls._create_config_version(
session=session,
tenant_id=tenant_id,
@@ -1780,7 +1457,6 @@ class AgentComposerService:
agent_soul=agent_soul,
operation=AgentConfigRevisionOperation.CREATE_VERSION,
version_note=None,
home_snapshot_id=home_snapshot.id,
)
agent.active_config_snapshot_id = version.id
agent.active_config_has_model = agent_soul_has_model(agent_soul)
@@ -1914,6 +1590,7 @@ class AgentComposerService:
session=session,
)
except IntegrityError as exc:
session.rollback()
raise AgentNameConflictError() from exc
agent = AgentRosterService(session).get_app_backing_agent(tenant_id=tenant_id, app_id=app.id)
@@ -1951,7 +1628,6 @@ class AgentComposerService:
agent_soul: AgentSoulConfig,
operation: AgentConfigRevisionOperation,
version_note: str | None,
home_snapshot_id: str,
previous_snapshot_id: str | None = None,
) -> AgentConfigSnapshot:
next_version = (
@@ -1968,7 +1644,6 @@ class AgentComposerService:
agent_id=agent_id,
version=next_version,
config_snapshot=agent_soul,
home_snapshot_id=home_snapshot_id,
version_note=version_note,
created_by=account_id,
)
@@ -2008,7 +1683,6 @@ class AgentComposerService:
operation=operation,
version_note=version_note,
previous_snapshot_id=current_snapshot.id,
home_snapshot_id=current_snapshot.home_snapshot_id,
)
@classmethod
@@ -2075,10 +1749,7 @@ class AgentComposerService:
agent: Agent,
created_by: str | None,
) -> AgentConfigDraft:
"""Resolve the normal Draft, rebasing only stale WORKFLOW_ONLY DRAFT rows whose account_id is None.
Roster and DEBUG_BUILD Drafts are never rebased.
"""
"""Resolve the shared Preview draft, rebasing inline agents when needed."""
return cls._get_or_create_agent_draft(
session=session,
tenant_id=tenant_id,
@@ -2096,8 +1767,6 @@ class AgentComposerService:
snapshot: AgentConfigSnapshot,
updated_by: str | None,
) -> bool:
"""Sync a stale normal Draft's base_snapshot_id, home_snapshot_id, config_snapshot, and updated_by."""
if (
agent.scope != AgentScope.WORKFLOW_ONLY
or draft.draft_type != AgentConfigDraftType.DRAFT
@@ -2108,7 +1777,6 @@ class AgentComposerService:
):
return False
draft.base_snapshot_id = snapshot.id
draft.home_snapshot_id = snapshot.home_snapshot_id
draft.config_snapshot = AgentSoulConfig.model_validate(snapshot.config_snapshot_dict)
draft.updated_by = updated_by
return True
@@ -2145,9 +1813,7 @@ class AgentComposerService:
agent_id=agent.id,
version_id=agent.active_config_snapshot_id,
)
if active_snapshot is None:
raise AgentVersionNotFoundError()
if cls._rebase_workflow_only_normal_draft(
if active_snapshot is not None and cls._rebase_workflow_only_normal_draft(
agent=agent,
draft=draft,
snapshot=active_snapshot,
@@ -2161,17 +1827,18 @@ class AgentComposerService:
agent_id=agent.id,
version_id=agent.active_config_snapshot_id,
)
if base_snapshot is None:
raise AgentVersionNotFoundError()
agent_soul = AgentSoulConfig.model_validate(base_snapshot.config_snapshot_dict)
agent_soul = (
AgentSoulConfig.model_validate(base_snapshot.config_snapshot_dict)
if base_snapshot is not None
else AgentSoulConfig()
)
draft = AgentConfigDraft(
tenant_id=tenant_id,
agent_id=agent.id,
draft_type=draft_type,
account_id=account_id if draft_type == AgentConfigDraftType.DEBUG_BUILD else None,
draft_owner_key=account_id if draft_type == AgentConfigDraftType.DEBUG_BUILD and account_id else "",
base_snapshot_id=base_snapshot.id,
home_snapshot_id=base_snapshot.home_snapshot_id,
base_snapshot_id=base_snapshot.id if base_snapshot else None,
config_snapshot=agent_soul,
created_by=created_by,
updated_by=created_by,
@@ -2473,7 +2140,7 @@ class AgentComposerService:
from services.agent.roster_service import AgentRosterService
return AgentRosterService(session).get_or_create_build_conversation(
return AgentRosterService(session).get_or_create_agent_app_debug_conversation_id(
tenant_id=tenant_id,
agent_id=agent.id,
account_id=account_id,
+2 -15
View File
@@ -49,7 +49,6 @@ from services.agent.dsl_entities import (
make_portable_agent_package,
portable_ref,
)
from services.agent.home_snapshot_service import AgentHomeSnapshotService
from services.agent.knowledge_datasets import get_tenant_knowledge_dataset_rows
from services.agent.roster_service import AgentRosterService
from services.entities.dsl_entities import DslImportWarning
@@ -225,7 +224,6 @@ class AgentDslService:
account_id=None,
draft_owner_key="",
base_snapshot_id=snapshot.id,
home_snapshot_id=snapshot.home_snapshot_id,
config_snapshot=soul,
created_by=account.id,
updated_by=account.id,
@@ -245,7 +243,7 @@ class AgentDslService:
portable_graph: Mapping[str, Any],
raw_packages: Mapping[str, Any],
account: Account,
) -> tuple[dict[str, Any], list[DslImportWarning], set[str]]:
) -> tuple[dict[str, Any], list[DslImportWarning]]:
"""Materialize every packaged Agent as a node-owned inline Agent."""
graph = copy.deepcopy(dict(portable_graph))
@@ -258,11 +256,6 @@ class AgentDslService:
WorkflowAgentNodeBinding.workflow_version == Workflow.VERSION_DRAFT,
)
).all()
retirement_candidates = {
binding.agent_id
for binding in previous_bindings
if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and binding.agent_id
}
for binding in previous_bindings:
self.session.delete(binding)
self.session.flush()
@@ -319,7 +312,7 @@ class AgentDslService:
workflow.graph = json.dumps(graph)
self.session.flush()
return graph, warnings, retirement_candidates
return graph, warnings
def clone_inline_binding_for_node(
self,
@@ -569,17 +562,11 @@ class AgentDslService:
)
or 0
) + 1
home_snapshot = AgentHomeSnapshotService.create_initial(
session=self.session,
tenant_id=tenant_id,
agent_id=agent.id,
)
snapshot = AgentConfigSnapshot(
tenant_id=tenant_id,
agent_id=agent.id,
version=next_version,
config_snapshot=soul,
home_snapshot_id=home_snapshot.id,
created_by=account_id,
)
self.session.add(snapshot)
-6
View File
@@ -29,12 +29,6 @@ class AgentModelNotConfiguredError(BaseHTTPException):
code = 400
class AgentBuildSandboxNotFoundError(BaseHTTPException):
error_code = "agent_build_sandbox_not_found"
description = "The retained Build Sandbox is no longer available."
code = 404
class AgentSoulLockedError(BadRequest):
description = "Agent Soul is locked for this workflow node."
-238
View File
@@ -1,238 +0,0 @@
"""Own immutable Agent Home Snapshot ledger rows and physical collection."""
from __future__ import annotations
import logging
from dify_agent.client import Client, DifyAgentNotFoundError
from dify_agent.protocol import CreateHomeSnapshotFromBindingRequest, InitializeHomeSnapshotRequest
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.db.session_factory import session_factory
from libs.datetime_utils import naive_utc_now
from libs.uuid_utils import uuidv7
from models.agent import (
Agent,
AgentConfigDraft,
AgentConfigSnapshot,
AgentConfigVersionKind,
AgentHomeSnapshot,
AgentStatus,
AgentWorkingResourceStatus,
AgentWorkspaceOwnerType,
)
from services.agent.errors import AgentBuildSandboxNotFoundError
from services.agent.workspace_service import AgentWorkspaceService, WorkspaceOwnerScope
logger = logging.getLogger(__name__)
class AgentHomeSnapshotUnavailableError(RuntimeError):
"""The requested owner-scoped Home Snapshot cannot be used."""
class AgentHomeSnapshotService:
"""Create, retire, and collect Agent-owned immutable Home Snapshots."""
@classmethod
def create_initial(
cls,
*,
session: Session,
tenant_id: str,
agent_id: str,
) -> AgentHomeSnapshot:
home_snapshot_id = str(uuidv7())
with cls._client() as client:
response = client.initialize_home_snapshot_sync(
InitializeHomeSnapshotRequest(
tenant_id=tenant_id,
agent_id=agent_id,
home_snapshot_id=home_snapshot_id,
)
)
home_snapshot = AgentHomeSnapshot(
id=home_snapshot_id,
tenant_id=tenant_id,
agent_id=agent_id,
snapshot_ref=response.snapshot_ref,
status=AgentWorkingResourceStatus.ACTIVE,
)
session.add(home_snapshot)
session.flush()
return home_snapshot
@classmethod
def create_for_build_apply(
cls,
*,
session: Session,
build_draft: AgentConfigDraft,
) -> AgentHomeSnapshot:
"""Checkpoint the exact participant owned by ``build_draft``."""
source_binding_id = build_draft.agent_workspace_binding_id
if source_binding_id is None:
raise AgentBuildSandboxNotFoundError()
agent = session.scalar(
select(Agent).where(
Agent.id == build_draft.agent_id,
Agent.tenant_id == build_draft.tenant_id,
)
)
if agent is None:
raise AgentBuildSandboxNotFoundError()
from services.agent.roster_service import AgentRosterService
runtime_app_id = AgentRosterService.runtime_backing_app_id(agent)
if runtime_app_id is None:
raise AgentBuildSandboxNotFoundError()
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=build_draft.tenant_id,
binding_id=source_binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=build_draft.tenant_id,
app_id=runtime_app_id,
owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT,
owner_id=build_draft.id,
),
)
if binding is None or binding.agent_id != build_draft.agent_id:
raise AgentBuildSandboxNotFoundError()
AgentWorkspaceService.validate_binding_generation(
binding,
base_home_snapshot_id=build_draft.home_snapshot_id,
agent_config_version_id=build_draft.id,
agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT,
)
home_snapshot_id = str(uuidv7())
try:
with cls._client() as client:
response = client.create_home_snapshot_from_binding_sync(
CreateHomeSnapshotFromBindingRequest(
tenant_id=build_draft.tenant_id,
agent_id=build_draft.agent_id,
home_snapshot_id=home_snapshot_id,
backend_binding_ref=binding.backend_binding_ref,
)
)
except DifyAgentNotFoundError as exc:
raise AgentBuildSandboxNotFoundError() from exc
home_snapshot = AgentHomeSnapshot(
id=home_snapshot_id,
tenant_id=build_draft.tenant_id,
agent_id=build_draft.agent_id,
snapshot_ref=response.snapshot_ref,
status=AgentWorkingResourceStatus.ACTIVE,
)
session.add(home_snapshot)
return home_snapshot
@classmethod
def retire_all_for_agent(cls, *, session: Session, tenant_id: str, agent_id: str) -> list[str]:
rows = session.scalars(
select(AgentHomeSnapshot).where(
AgentHomeSnapshot.tenant_id == tenant_id,
AgentHomeSnapshot.agent_id == agent_id,
AgentHomeSnapshot.status == AgentWorkingResourceStatus.ACTIVE,
)
).all()
now = naive_utc_now()
for row in rows:
row.status = AgentWorkingResourceStatus.RETIRED
row.retired_at = now
return [row.id for row in rows]
@classmethod
def collect_retired_home_snapshot(cls, *, tenant_id: str, home_snapshot_id: str) -> None:
try:
cls._collect_retired_home_snapshot(tenant_id=tenant_id, home_snapshot_id=home_snapshot_id)
except Exception:
logger.exception(
"Failed to collect retired Agent Home Snapshot",
extra={"tenant_id": tenant_id, "home_snapshot_id": home_snapshot_id},
)
@classmethod
def _collect_retired_home_snapshot(cls, *, tenant_id: str, home_snapshot_id: str) -> None:
with session_factory.create_session() as session:
snapshot = session.scalar(
select(AgentHomeSnapshot).where(
AgentHomeSnapshot.id == home_snapshot_id,
AgentHomeSnapshot.tenant_id == tenant_id,
AgentHomeSnapshot.status == AgentWorkingResourceStatus.RETIRED,
)
)
if snapshot is None:
return
referenced = session.scalar(
select(AgentConfigDraft.id).where(AgentConfigDraft.home_snapshot_id == home_snapshot_id).limit(1)
) or session.scalar(
select(AgentConfigSnapshot.id).where(AgentConfigSnapshot.home_snapshot_id == home_snapshot_id).limit(1)
)
if referenced is not None:
return
snapshot_ref = snapshot.snapshot_ref
try:
cls.delete(snapshot_ref=snapshot_ref)
except Exception:
logger.exception(
"Failed to collect retired Agent Home Snapshot",
extra={"tenant_id": tenant_id, "home_snapshot_id": home_snapshot_id},
)
return
with session_factory.create_session() as session:
snapshot = session.scalar(
select(AgentHomeSnapshot).where(
AgentHomeSnapshot.id == home_snapshot_id,
AgentHomeSnapshot.tenant_id == tenant_id,
AgentHomeSnapshot.status == AgentWorkingResourceStatus.RETIRED,
)
)
if snapshot is not None:
session.delete(snapshot)
session.commit()
@classmethod
def delete(cls, *, snapshot_ref: str) -> None:
with cls._client() as client:
client.delete_home_snapshot_sync(snapshot_ref)
@staticmethod
def _client() -> Client:
base_url = dify_config.AGENT_BACKEND_BASE_URL
if not base_url:
raise AgentHomeSnapshotUnavailableError("Dify Agent backend is required for Home Snapshot operations")
return Client(base_url=base_url)
def validate_home_snapshot_binding(*, session: Session, agent: Agent, home_snapshot_id: str) -> None:
_require_owned_home_snapshot(session=session, agent=agent, home_snapshot_id=home_snapshot_id)
def _require_owned_home_snapshot(*, session: Session, agent: Agent, home_snapshot_id: str) -> AgentHomeSnapshot:
if agent.status != AgentStatus.ACTIVE:
raise AgentHomeSnapshotUnavailableError(f"Agent {agent.id} is not active")
home_snapshot = session.scalar(
select(AgentHomeSnapshot).where(
AgentHomeSnapshot.id == home_snapshot_id,
AgentHomeSnapshot.tenant_id == agent.tenant_id,
AgentHomeSnapshot.agent_id == agent.id,
AgentHomeSnapshot.status == AgentWorkingResourceStatus.ACTIVE,
)
)
if home_snapshot is None:
raise AgentHomeSnapshotUnavailableError(f"Home Snapshot {home_snapshot_id} is unavailable for Agent {agent.id}")
return home_snapshot
__all__ = [
"AgentHomeSnapshotService",
"AgentHomeSnapshotUnavailableError",
"validate_home_snapshot_binding",
]
-166
View File
@@ -1,166 +0,0 @@
"""Workflow-only Agent ownership retirement after product transactions commit."""
from __future__ import annotations
import logging
from collections.abc import Iterable
from sqlalchemy import or_, select
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
from libs.datetime_utils import naive_utc_now
from models.agent import (
Agent,
AgentScope,
AgentStatus,
AgentWorkingResourceStatus,
AgentWorkspaceBinding,
WorkflowAgentNodeBinding,
)
from models.enums import AppStatus
from models.model import App
from models.workflow import Workflow
from services.agent.home_snapshot_service import AgentHomeSnapshotService
from services.agent.workspace_service import AgentWorkspaceService
logger = logging.getLogger(__name__)
class WorkflowAgentRetirementService:
"""Archive workflow-only Agents once no effective binding owns them."""
@classmethod
def retire_unowned(
cls,
*,
tenant_id: str,
agent_ids: Iterable[str],
account_id: str | None,
) -> tuple[list[str], list[str]]:
"""Re-check ownership, archive orphans, and commit their resource retirement."""
candidates = tuple(sorted({agent_id for agent_id in agent_ids if agent_id}))
if not candidates:
return [], []
retired_bindings: list[str] = []
retired_snapshots: list[str] = []
try:
with session_factory.create_session() as session:
retired_agent_ids = cls.archive_unowned(
session=session,
tenant_id=tenant_id,
agent_ids=candidates,
account_id=account_id,
)
for agent_id in retired_agent_ids:
bindings = session.scalars(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.agent_id == agent_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
)
).all()
for binding in bindings:
binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=tenant_id,
binding_id=binding.id,
)
if binding_id is not None:
retired_bindings.append(binding_id)
retired_snapshots.extend(
AgentHomeSnapshotService.retire_all_for_agent(
session=session,
tenant_id=tenant_id,
agent_id=agent_id,
)
)
session.commit()
except Exception:
logger.exception(
"Failed to retire unowned Workflow Agents",
extra={
"tenant_id": tenant_id,
"agent_ids": candidates,
},
)
return [], []
return retired_bindings, retired_snapshots
@classmethod
def archive_unowned(
cls,
*,
session: Session,
tenant_id: str,
agent_ids: Iterable[str],
account_id: str | None,
) -> list[str]:
"""Archive active orphans and return every orphan eligible for Home cleanup."""
candidates = tuple(sorted({agent_id for agent_id in agent_ids if agent_id}))
if not candidates:
return []
agents = session.scalars(
select(Agent).where(
Agent.tenant_id == tenant_id,
Agent.id.in_(candidates),
Agent.scope == AgentScope.WORKFLOW_ONLY,
Agent.status.in_((AgentStatus.ACTIVE, AgentStatus.ARCHIVED)),
)
).all()
effective_agent_ids = cls._effective_agent_ids(
session=session,
tenant_id=tenant_id,
agent_ids=[agent.id for agent in agents],
)
now = naive_utc_now()
cleanup_candidates: list[str] = []
for agent in agents:
if agent.id in effective_agent_ids:
continue
if agent.status == AgentStatus.ACTIVE:
agent.status = AgentStatus.ARCHIVED
agent.archived_by = account_id
agent.archived_at = now
agent.updated_by = account_id or agent.updated_by
agent.updated_at = now
cleanup_candidates.append(agent.id)
session.flush()
return cleanup_candidates
@staticmethod
def _effective_agent_ids(
*,
session: Session,
tenant_id: str,
agent_ids: list[str],
) -> set[str]:
if not agent_ids:
return set()
values = session.scalars(
select(WorkflowAgentNodeBinding.agent_id)
.join(
Workflow,
Workflow.id == WorkflowAgentNodeBinding.workflow_id,
)
.join(App, App.id == WorkflowAgentNodeBinding.app_id)
.where(
WorkflowAgentNodeBinding.tenant_id == tenant_id,
WorkflowAgentNodeBinding.agent_id.in_(agent_ids),
Workflow.tenant_id == tenant_id,
Workflow.app_id == WorkflowAgentNodeBinding.app_id,
Workflow.version == WorkflowAgentNodeBinding.workflow_version,
App.tenant_id == tenant_id,
App.status == AppStatus.NORMAL,
or_(
Workflow.version == Workflow.VERSION_DRAFT,
App.workflow_id == Workflow.id,
),
)
.distinct()
).all()
return {agent_id for agent_id in values if agent_id}
__all__ = ["WorkflowAgentRetirementService"]
+171 -258
View File
@@ -4,8 +4,10 @@ from typing import Any, TypedDict
from sqlalchemy import and_, func, or_, select
from sqlalchemy.exc import IntegrityError
from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload
from constants.model_template import default_app_templates
from core.agent.publish_visibility import workflow_callable_active_snapshot_filter
from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore
from core.app.entities.app_invoke_entities import InvokeFrom
from libs.datetime_utils import naive_utc_now
from libs.helper import to_timestamp
@@ -23,9 +25,6 @@ from models.agent import (
AgentScope,
AgentSource,
AgentStatus,
AgentWorkingResourceStatus,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
WorkflowAgentBindingType,
WorkflowAgentNodeBinding,
)
@@ -37,18 +36,15 @@ from services.agent.agent_soul_state import agent_soul_has_model
from services.agent.composer_validator import ComposerConfigValidator
from services.agent.errors import (
AgentArchivedError,
AgentBuildSandboxNotFoundError,
AgentNameConflictError,
AgentNotFoundError,
AgentVersionNotFoundError,
)
from services.agent.home_snapshot_service import AgentHomeSnapshotService
from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope
from services.app_service import AppService, CreateAppParams
from services.enterprise.enterprise_service import EnterpriseService
from services.entities.agent_entities import RosterAgentCreatePayload, RosterAgentUpdatePayload
from services.feature_service import FeatureService
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
from tasks.agent_backend_session_cleanup_task import cleanup_conversation_agent_runtime_session
logger = logging.getLogger(__name__)
@@ -308,27 +304,6 @@ class AgentRosterService:
account_id: str,
payload: RosterAgentCreatePayload,
source: AgentSource = AgentSource.ROSTER,
) -> Agent:
try:
agent = self._create_roster_agent_in_transaction(
tenant_id=tenant_id,
account_id=account_id,
payload=payload,
source=source,
)
self._session.commit()
return agent
except IntegrityError as exc:
self._session.rollback()
raise AgentNameConflictError() from exc
def _create_roster_agent_in_transaction(
self,
*,
tenant_id: str,
account_id: str,
payload: RosterAgentCreatePayload,
source: AgentSource,
) -> Agent:
ComposerConfigValidator.validate_agent_soul(payload.agent_soul)
@@ -348,19 +323,17 @@ class AgentRosterService:
updated_by=account_id,
)
self._session.add(agent)
self._session.flush()
try:
self._session.flush()
except IntegrityError as exc:
self._session.rollback()
raise AgentNameConflictError() from exc
home_snapshot = AgentHomeSnapshotService.create_initial(
session=self._session,
tenant_id=tenant_id,
agent_id=agent.id,
)
version = AgentConfigSnapshot(
tenant_id=tenant_id,
agent_id=agent.id,
version=1,
config_snapshot=payload.agent_soul,
home_snapshot_id=home_snapshot.id,
version_note=payload.version_note,
created_by=account_id,
)
@@ -381,6 +354,11 @@ class AgentRosterService:
agent.active_config_has_model = agent_soul_has_model(payload.agent_soul)
agent.active_config_is_published = True
try:
self._session.commit()
except IntegrityError as exc:
self._session.rollback()
raise AgentNameConflictError() from exc
return agent
def create_backing_agent_for_app(
@@ -427,19 +405,17 @@ class AgentRosterService:
updated_by=account_id,
)
self._session.add(agent)
self._session.flush()
try:
self._session.flush()
except IntegrityError as exc:
self._session.rollback()
raise AgentNameConflictError() from exc
home_snapshot = AgentHomeSnapshotService.create_initial(
session=self._session,
tenant_id=tenant_id,
agent_id=agent.id,
)
version = AgentConfigSnapshot(
tenant_id=tenant_id,
agent_id=agent.id,
version=1,
config_snapshot=soul,
home_snapshot_id=home_snapshot.id,
created_by=account_id,
)
self._session.add(version)
@@ -614,15 +590,16 @@ class AgentRosterService:
self._session.flush()
return conversation_id
def get_or_create_build_conversation(
def get_or_create_agent_app_debug_conversation_id(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
commit: bool = True,
) -> str:
"""Return the current editor's stable Build conversation."""
"""Return the current editor's Build or Preview conversation for an Agent App."""
agent = self._session.scalar(
select(Agent).where(
@@ -637,20 +614,21 @@ class AgentRosterService:
conversation_id = self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
draft_type=draft_type,
)
if commit:
self._session.commit()
return conversation_id
def get_current_preview_conversation(
def load_agent_app_debug_conversation_id(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
) -> str | None:
"""Return the editor's current Preview conversation without creating one."""
"""Return the editor's existing scoped conversation without creating or repairing rows."""
return self._session.scalar(
select(Conversation.id)
@@ -659,7 +637,7 @@ class AgentRosterService:
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent_id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == AgentConfigDraftType.DRAFT,
AgentDebugConversation.draft_type == draft_type,
AgentDebugConversation.app_id == Conversation.app_id,
Conversation.from_source == ConversationFromSource.CONSOLE,
Conversation.from_account_id == account_id,
@@ -679,12 +657,25 @@ class AgentRosterService:
or 0
)
def rotate_preview_conversation(self, *, tenant_id: str, agent_id: str, account_id: str) -> str:
"""Rotate Preview and retire its exact Conversation-owned Binding.
def refresh_agent_app_debug_conversation_id(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
) -> str:
"""Start a new scoped console conversation for the current Agent App editor.
The mapping update and exact CONVERSATION Binding retirement commit in
one transaction. Validation failures fail fast; collection is enqueued
only after commit.
If this account already has a mapping for the requested draft surface, the previous
conversation is abandoned after the replacement mapping is committed: any ACTIVE
conversation-owned Agent runtime sessions for that old conversation are sent through
best-effort backend cleanup and then retired locally even when enqueueing fails. This
order prevents a failed database commit from retiring the still-current runtime session.
The other draft surface is left untouched.
A user and draft surface own one current mapping. If new-conversation requests overlap,
the last committed rotation becomes current and earlier response IDs cannot be continued.
"""
agent = self._session.scalar(
@@ -703,194 +694,142 @@ class AgentRosterService:
if not backing_app_id:
raise AgentNotFoundError()
retired_binding_id: str | None = None
try:
conversation_id = self._create_agent_app_debug_conversation(
app_id=backing_app_id,
account_id=account_id,
conversation_id = self._create_agent_app_debug_conversation(
app_id=backing_app_id,
account_id=account_id,
)
previous_conversation: tuple[str, str] | None = None
mapping = self._session.scalar(
select(AgentDebugConversation).where(
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent_id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == draft_type,
)
mapping = self._session.scalar(
select(AgentDebugConversation).where(
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent_id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == AgentConfigDraftType.DRAFT,
)
if mapping is None:
self._session.add(
AgentDebugConversation(
tenant_id=tenant_id,
agent_id=agent_id,
app_id=backing_app_id,
account_id=account_id,
draft_type=draft_type,
conversation_id=conversation_id,
)
)
if mapping is None:
self._session.add(
AgentDebugConversation(
tenant_id=tenant_id,
agent_id=agent_id,
app_id=backing_app_id,
account_id=account_id,
draft_type=AgentConfigDraftType.DRAFT,
conversation_id=conversation_id,
)
)
else:
previous_app_id = mapping.app_id or backing_app_id
previous_conversation_id = mapping.conversation_id
if previous_conversation_id:
previous_conversation = self._session.scalar(
select(Conversation).where(
Conversation.id == previous_conversation_id,
Conversation.app_id == previous_app_id,
Conversation.from_source == ConversationFromSource.CONSOLE,
Conversation.from_account_id == account_id,
Conversation.is_deleted.is_(False),
)
)
if (
previous_conversation is not None
and previous_conversation.agent_workspace_binding_id is not None
):
binding_id = previous_conversation.agent_workspace_binding_id
binding = AgentWorkspaceService.get_active_binding(
session=self._session,
tenant_id=tenant_id,
binding_id=binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=previous_app_id,
owner_type=AgentWorkspaceOwnerType.CONVERSATION,
owner_id=previous_conversation.id,
),
)
if binding is None or binding.agent_id != agent_id:
raise AgentWorkspaceNotFoundError(
"Agent debug Conversation participant Binding is unavailable"
)
retired_binding_id = AgentWorkspaceService.retire_binding(
session=self._session,
tenant_id=tenant_id,
binding_id=binding_id,
)
if retired_binding_id is None:
raise AgentWorkspaceNotFoundError(
"Agent debug Conversation participant Binding is unavailable"
)
mapping.app_id = backing_app_id
mapping.conversation_id = conversation_id
self._session.flush()
self._session.commit()
except Exception:
self._session.rollback()
raise
if retired_binding_id is not None:
enqueue_agent_resource_collection(
else:
previous_app_id = mapping.app_id
previous_conversation_id = mapping.conversation_id
if previous_conversation_id:
previous_conversation = (previous_app_id or backing_app_id, previous_conversation_id)
mapping.app_id = backing_app_id
mapping.conversation_id = conversation_id
self._session.flush()
self._session.commit()
if previous_conversation:
previous_app_id, previous_conversation_id = previous_conversation
self._cleanup_debug_conversation_runtime_sessions(
tenant_id=tenant_id,
binding_ids=(retired_binding_id,),
agent_id=agent_id,
account_id=account_id,
draft_type=draft_type,
app_id=previous_app_id,
conversation_id=previous_conversation_id,
)
return conversation_id
def reset_build_conversation(self, *, tenant_id: str, agent_id: str, account_id: str) -> str:
"""Reset Build and retire its exact DEBUG_BUILD Draft-owned Binding.
The mapping update, exact BUILD_DRAFT Binding retirement, and Draft
pointer clear commit in one transaction. Validation failures fail fast;
collection is enqueued only after commit.
"""
agent = self._session.scalar(
select(Agent).where(
Agent.tenant_id == tenant_id,
Agent.id == agent_id,
Agent.status == AgentStatus.ACTIVE,
)
)
if agent is None:
raise AgentNotFoundError()
backing_app_id = self._ensure_workflow_agent_backing_app(
agent=agent,
account_id=agent.updated_by or agent.created_by,
)
if not backing_app_id:
raise AgentNotFoundError()
retired_binding_id: str | None = None
def _cleanup_debug_conversation_runtime_sessions(
self,
*,
tenant_id: str,
agent_id: str,
account_id: str,
draft_type: AgentConfigDraftType,
app_id: str,
conversation_id: str,
) -> None:
try:
conversation_id = self._create_agent_app_debug_conversation(
app_id=backing_app_id,
account_id=account_id,
)
mapping = self._session.scalar(
select(AgentDebugConversation).where(
AgentDebugConversation.tenant_id == tenant_id,
AgentDebugConversation.agent_id == agent_id,
AgentDebugConversation.account_id == account_id,
AgentDebugConversation.draft_type == AgentConfigDraftType.DEBUG_BUILD,
)
)
build_draft = self._session.scalar(
select(AgentConfigDraft)
.where(
AgentConfigDraft.tenant_id == tenant_id,
AgentConfigDraft.agent_id == agent_id,
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
AgentConfigDraft.account_id == account_id,
)
.order_by(AgentConfigDraft.updated_at.desc())
.limit(1)
)
if build_draft is not None and build_draft.agent_workspace_binding_id is not None:
binding_id = build_draft.agent_workspace_binding_id
binding = AgentWorkspaceService.get_active_binding(
session=self._session,
tenant_id=tenant_id,
binding_id=binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=backing_app_id,
owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT,
owner_id=build_draft.id,
),
)
if binding is None or binding.agent_id != agent_id:
raise AgentBuildSandboxNotFoundError()
retired_binding_id = AgentWorkspaceService.retire_binding(
session=self._session,
tenant_id=tenant_id,
binding_id=binding_id,
)
if retired_binding_id is None:
raise AgentBuildSandboxNotFoundError()
build_draft.agent_workspace_binding_id = None
if mapping is None:
self._session.add(
AgentDebugConversation(
tenant_id=tenant_id,
agent_id=agent_id,
app_id=backing_app_id,
account_id=account_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
conversation_id=conversation_id,
)
)
else:
mapping.app_id = backing_app_id
mapping.conversation_id = conversation_id
self._session.flush()
self._session.commit()
except Exception:
self._session.rollback()
raise
if retired_binding_id is not None:
enqueue_agent_resource_collection(
session_store = AgentAppRuntimeSessionStore()
stored_sessions = session_store.list_active_sessions_for_conversation(
tenant_id=tenant_id,
binding_ids=(retired_binding_id,),
app_id=app_id,
conversation_id=conversation_id,
)
return conversation_id
except Exception:
logger.warning(
"Failed to load Agent App runtime sessions for debug conversation refresh: "
"tenant_id=%s app_id=%s conversation_id=%s",
tenant_id,
app_id,
conversation_id,
exc_info=True,
)
return
def load_or_create_build_conversation_ids_by_agent_id(
for stored_session in stored_sessions:
try:
if stored_session.runtime_layer_specs:
payload = AgentBackendSessionCleanupPayload(
session_snapshot=stored_session.session_snapshot,
runtime_layer_specs=stored_session.runtime_layer_specs,
idempotency_key=(
f"{tenant_id}:{agent_id}:{account_id}:{draft_type.value}:{conversation_id}:"
"debug-session-cleanup:"
f"{stored_session.scope.agent_id}:"
f"{stored_session.scope.agent_config_snapshot_id or 'no-config'}:"
f"{stored_session.backend_run_id or 'no-run'}"
),
metadata={
"tenant_id": stored_session.scope.tenant_id,
"app_id": stored_session.scope.app_id,
"conversation_id": stored_session.scope.conversation_id,
"agent_id": stored_session.scope.agent_id,
"agent_config_snapshot_id": stored_session.scope.agent_config_snapshot_id,
"draft_type": draft_type.value,
"previous_agent_backend_run_id": stored_session.backend_run_id,
},
)
cleanup_conversation_agent_runtime_session.delay(payload.model_dump(mode="json"))
except Exception:
logger.warning(
"Failed to enqueue Agent backend cleanup for debug conversation refresh: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s",
stored_session.scope.tenant_id,
stored_session.scope.app_id,
stored_session.scope.conversation_id,
stored_session.scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
finally:
try:
session_store.mark_cleaned(
scope=stored_session.scope,
backend_run_id=stored_session.backend_run_id,
)
except Exception:
logger.warning(
"Failed to retire Agent App runtime session for debug conversation refresh: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s",
stored_session.scope.tenant_id,
stored_session.scope.app_id,
stored_session.scope.conversation_id,
stored_session.scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
def load_or_create_agent_app_debug_conversation_ids_by_agent_id(
self,
*,
tenant_id: str,
agents: list[Agent],
account_id: str,
draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD,
) -> dict[str, str]:
"""Return per-account Build conversations for a page of Agent Apps."""
"""Return per-account scoped conversations for a page of Agent Apps."""
conversation_ids_by_agent_id: dict[str, str] = {}
changed = False
@@ -900,7 +839,7 @@ class AgentRosterService:
conversation_ids_by_agent_id[agent.id] = self._get_or_create_agent_app_debug_conversation(
agent=agent,
account_id=account_id,
draft_type=AgentConfigDraftType.DEBUG_BUILD,
draft_type=draft_type,
)
changed = True
if changed:
@@ -1118,6 +1057,7 @@ class AgentRosterService:
account_id=account.id,
)
self._session.commit()
if FeatureService.get_system_features().webapp_auth.enabled:
try:
original_settings = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(source_app.id)
@@ -1271,38 +1211,13 @@ class AgentRosterService:
def archive_roster_agent(self, *, tenant_id: str, agent_id: str, account_id: str) -> None:
agent = self._get_agent(tenant_id=tenant_id, agent_id=agent_id, roster_only=True)
retired_binding_ids: list[str] = []
if agent.status != AgentStatus.ARCHIVED:
agent.status = AgentStatus.ARCHIVED
agent.archived_by = account_id
agent.archived_at = naive_utc_now()
agent.updated_by = account_id
bindings = self._session.scalars(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.agent_id == agent_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
)
).all()
for binding in bindings:
retired_id = AgentWorkspaceService.retire_binding(
session=self._session,
tenant_id=tenant_id,
binding_id=binding.id,
)
if retired_id is not None:
retired_binding_ids.append(retired_id)
retired_snapshot_ids = AgentHomeSnapshotService.retire_all_for_agent(
session=self._session,
tenant_id=tenant_id,
agent_id=agent_id,
)
if agent.status == AgentStatus.ARCHIVED:
return
agent.status = AgentStatus.ARCHIVED
agent.archived_by = account_id
agent.archived_at = naive_utc_now()
agent.updated_by = account_id
self._session.commit()
enqueue_agent_resource_collection(
tenant_id=tenant_id,
binding_ids=retired_binding_ids,
home_snapshot_ids=retired_snapshot_ids,
)
@staticmethod
def _visible_version_operations(agent: Agent) -> set[AgentConfigRevisionOperation]:
@@ -1467,11 +1382,9 @@ class AgentRosterService:
account_id=None,
draft_owner_key="",
created_by=account_id,
home_snapshot_id=version.home_snapshot_id,
)
self._session.add(draft)
draft.base_snapshot_id = version.id
draft.home_snapshot_id = version.home_snapshot_id
draft.config_snapshot = AgentSoulConfig.model_validate(version.config_snapshot_dict)
draft.updated_by = account_id
agent.active_config_is_published = version.id == agent.active_config_snapshot_id
+8 -57
View File
@@ -24,7 +24,6 @@ from models.agent_config_entities import (
WorkflowNodeJobConfig,
WorkflowPreviousNodeOutputRef,
)
from models.model import App
from models.workflow import Workflow
from services.agent.composer_validator import ComposerConfigValidator
from services.agent.prompt_mentions import (
@@ -225,7 +224,7 @@ class WorkflowAgentPublishService:
session: Session,
draft_workflow: Workflow,
account_id: str,
) -> set[str]:
) -> None:
agent_nodes = dict(WorkflowAgentNodeValidator.iter_agent_v2_nodes(draft_workflow.graph_dict))
existing_bindings = list(
session.scalars(
@@ -238,12 +237,9 @@ class WorkflowAgentPublishService:
).all()
)
existing_by_node_id = {binding.node_id: binding for binding in existing_bindings}
retirement_candidates: set[str] = set()
for binding in existing_bindings:
if binding.node_id not in agent_nodes:
if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and binding.agent_id:
retirement_candidates.add(binding.agent_id)
session.delete(binding)
for node_id, node_data in agent_nodes.items():
@@ -256,34 +252,16 @@ class WorkflowAgentPublishService:
not binding_payload.get("agent_id") or not binding_payload.get("current_snapshot_id")
):
continue
existing_binding = existing_by_node_id.get(node_id)
replaced_inline_agent_id = (
existing_binding.agent_id
if existing_binding is not None
and existing_binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT
and existing_binding.agent_id
else None
)
cls._sync_agent_binding_for_node(
session=session,
draft_workflow=draft_workflow,
node_id=node_id,
node_data=node_data,
node_binding=binding_payload,
existing_binding=existing_binding,
existing_binding=existing_by_node_id.get(node_id),
account_id=account_id,
)
if (
replaced_inline_agent_id
and existing_binding is not None
and (
existing_binding.binding_type != WorkflowAgentBindingType.INLINE_AGENT
or existing_binding.agent_id != replaced_inline_agent_id
)
):
retirement_candidates.add(replaced_inline_agent_id)
session.flush()
return retirement_candidates
@classmethod
def sync_roster_agent_bindings_for_draft(
@@ -292,8 +270,8 @@ class WorkflowAgentPublishService:
session: Session,
draft_workflow: Workflow,
account_id: str,
) -> set[str]:
return cls.sync_agent_bindings_for_draft(
) -> None:
cls.sync_agent_bindings_for_draft(
session=session,
draft_workflow=draft_workflow,
account_id=account_id,
@@ -583,32 +561,12 @@ class WorkflowAgentPublishService:
session: Session,
draft_workflow: Workflow,
published_workflow: Workflow,
) -> set[str]:
current_workflow_id = session.scalar(
select(App.workflow_id).where(
App.tenant_id == draft_workflow.tenant_id,
App.id == draft_workflow.app_id,
)
)
retirement_candidates: set[str] = set()
if current_workflow_id:
retirement_candidates = {
agent_id
for agent_id in session.scalars(
select(WorkflowAgentNodeBinding.agent_id).where(
WorkflowAgentNodeBinding.tenant_id == draft_workflow.tenant_id,
WorkflowAgentNodeBinding.app_id == draft_workflow.app_id,
WorkflowAgentNodeBinding.workflow_id == current_workflow_id,
WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.INLINE_AGENT,
)
).all()
if agent_id
}
) -> None:
node_ids = {
node_id for node_id, _node_data in WorkflowAgentNodeValidator.iter_agent_v2_nodes(draft_workflow.graph_dict)
}
if not node_ids:
return retirement_candidates
return
bindings = session.scalars(
select(WorkflowAgentNodeBinding).where(
@@ -620,7 +578,7 @@ class WorkflowAgentPublishService:
)
).all()
if not bindings:
return retirement_candidates
return
agents_by_id = {
agent.id: agent
@@ -653,7 +611,6 @@ class WorkflowAgentPublishService:
updated_by=binding.updated_by,
)
session.add(copied)
return retirement_candidates
@classmethod
def restore_agent_node_bindings_to_draft(
@@ -663,7 +620,7 @@ class WorkflowAgentPublishService:
source_workflow: Workflow,
draft_workflow: Workflow,
account_id: str,
) -> set[str]:
) -> None:
"""Replace draft bindings with the frozen bindings of a published workflow."""
existing = session.scalars(
@@ -674,11 +631,6 @@ class WorkflowAgentPublishService:
WorkflowAgentNodeBinding.workflow_version == cls._DRAFT_WORKFLOW_VERSION,
)
).all()
retirement_candidates = {
binding.agent_id
for binding in existing
if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and binding.agent_id
}
for binding in existing:
session.delete(binding)
@@ -729,4 +681,3 @@ class WorkflowAgentPublishService:
)
)
session.flush()
return retirement_candidates
-469
View File
@@ -1,469 +0,0 @@
"""Own Workspace and AgentWorkspaceBinding product lifecycle.
Dify API is the lifecycle ledger. Dify Agent only executes physical create,
acquire, and destroy operations selected by this service. Retire methods only
mutate the caller's transaction; collection performs network I/O after commit
and deletes ledger rows only after idempotent physical cleanup succeeds.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from dify_agent.client import Client
from dify_agent.protocol import CreateExecutionBindingRequest, DestroyExecutionBindingRequest
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.db.session_factory import session_factory
from libs.datetime_utils import naive_utc_now
from libs.uuid_utils import uuidv7
from models.agent import (
AgentConfigVersionKind,
AgentHomeSnapshot,
AgentWorkingResourceStatus,
AgentWorkspace,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
)
logger = logging.getLogger(__name__)
class AgentWorkspaceError(RuntimeError):
pass
class AgentWorkspaceNotFoundError(AgentWorkspaceError):
pass
class AgentWorkspaceBindingGenerationMismatchError(AgentWorkspaceError):
pass
@dataclass(frozen=True, slots=True)
class WorkspaceOwnerScope:
tenant_id: str
app_id: str
owner_type: AgentWorkspaceOwnerType
owner_id: str
owner_scope_key: str = "root"
class AgentWorkspaceService:
"""Allocate and manage working-environment resources.
A Binding ID is the participant identity. Product callers persist that ID
and use :meth:`get_active_binding`; Agent and Workspace attributes are not
participant lookup keys.
"""
@classmethod
def resolve_active_workspace(cls, *, session: Session, scope: WorkspaceOwnerScope) -> AgentWorkspace | None:
return session.scalar(
select(AgentWorkspace).where(
AgentWorkspace.tenant_id == scope.tenant_id,
AgentWorkspace.app_id == scope.app_id,
AgentWorkspace.owner_type == scope.owner_type,
AgentWorkspace.owner_id == scope.owner_id,
AgentWorkspace.owner_scope_key == scope.owner_scope_key,
AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE,
)
)
@classmethod
def get_active_binding(
cls,
*,
session: Session,
tenant_id: str,
binding_id: str,
expected_owner_scope: WorkspaceOwnerScope,
) -> AgentWorkspaceBinding | None:
return session.scalar(
select(AgentWorkspaceBinding)
.join(
AgentWorkspace,
(AgentWorkspace.tenant_id == AgentWorkspaceBinding.tenant_id)
& (AgentWorkspace.id == AgentWorkspaceBinding.workspace_id),
)
.where(
AgentWorkspaceBinding.id == binding_id,
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
AgentWorkspace.tenant_id == expected_owner_scope.tenant_id,
AgentWorkspace.app_id == expected_owner_scope.app_id,
AgentWorkspace.owner_type == expected_owner_scope.owner_type,
AgentWorkspace.owner_id == expected_owner_scope.owner_id,
AgentWorkspace.owner_scope_key == expected_owner_scope.owner_scope_key,
AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE,
)
)
@classmethod
def create_binding(
cls,
*,
session: Session,
scope: WorkspaceOwnerScope,
agent_id: str,
base_home_snapshot_id: str,
agent_config_version_id: str,
agent_config_version_kind: AgentConfigVersionKind,
) -> AgentWorkspaceBinding:
"""Allocate one new participant in the caller-owned transaction.
After backend creation returns successfully, any later Python, flush,
or commit failure may leave an orphan. Dify API does not perform
cross-system compensation; a future global reconciler is responsible
for those orphans. Backend-local cleanup applies only when creation
fails before the backend returns success.
"""
home_snapshot = session.scalar(
select(AgentHomeSnapshot).where(
AgentHomeSnapshot.id == base_home_snapshot_id,
AgentHomeSnapshot.tenant_id == scope.tenant_id,
AgentHomeSnapshot.agent_id == agent_id,
AgentHomeSnapshot.status == AgentWorkingResourceStatus.ACTIVE,
)
)
if home_snapshot is None:
raise AgentWorkspaceNotFoundError("base Home Snapshot is unavailable")
workspace = cls.resolve_active_workspace(session=session, scope=scope)
workspace_id = workspace.id if workspace is not None else str(uuidv7())
binding_id = str(uuidv7())
with cls._client() as client:
allocation = client.create_execution_binding_sync(
CreateExecutionBindingRequest(
tenant_id=scope.tenant_id,
agent_id=agent_id,
binding_id=binding_id,
workspace_id=workspace_id,
existing_workspace_ref=workspace.backend_workspace_ref if workspace is not None else None,
home_snapshot_ref=home_snapshot.snapshot_ref,
)
)
if workspace is not None and allocation.workspace_ref != workspace.backend_workspace_ref:
raise AgentWorkspaceError("backend changed the existing Workspace ref")
if workspace is None:
workspace = AgentWorkspace(
id=workspace_id,
tenant_id=scope.tenant_id,
app_id=scope.app_id,
owner_type=scope.owner_type,
owner_id=scope.owner_id,
owner_scope_key=scope.owner_scope_key,
backend_workspace_ref=allocation.workspace_ref,
status=AgentWorkingResourceStatus.ACTIVE,
active_guard=1,
)
session.add(workspace)
binding = AgentWorkspaceBinding(
id=binding_id,
tenant_id=scope.tenant_id,
app_id=scope.app_id,
workspace_id=workspace_id,
agent_id=agent_id,
base_home_snapshot_id=base_home_snapshot_id,
agent_config_version_id=agent_config_version_id,
agent_config_version_kind=agent_config_version_kind,
backend_binding_ref=allocation.binding_ref,
status=AgentWorkingResourceStatus.ACTIVE,
)
session.add(binding)
return binding
@classmethod
def save_binding_session_snapshot(
cls,
*,
tenant_id: str,
binding_id: str,
session_snapshot: str,
pending_form_id: str | None = None,
pending_tool_call_id: str | None = None,
) -> None:
with session_factory.create_session() as session:
binding = session.scalar(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.id == binding_id,
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
)
)
if binding is None:
raise AgentWorkspaceNotFoundError("ACTIVE Binding is unavailable")
binding.session_snapshot = session_snapshot
binding.pending_form_id = pending_form_id
binding.pending_tool_call_id = pending_tool_call_id
session.commit()
@classmethod
def retire_binding(cls, *, session: Session, tenant_id: str, binding_id: str) -> str | None:
binding = session.scalar(
select(AgentWorkspaceBinding)
.where(
AgentWorkspaceBinding.id == binding_id,
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
)
.with_for_update()
)
if binding is None:
return None
workspace = session.scalar(
select(AgentWorkspace)
.where(
AgentWorkspace.id == binding.workspace_id,
AgentWorkspace.tenant_id == tenant_id,
AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE,
)
.with_for_update()
)
now = naive_utc_now()
binding.status = AgentWorkingResourceStatus.RETIRED
binding.retired_at = now
if workspace is not None:
other_binding = session.scalar(
select(AgentWorkspaceBinding.id).where(
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.workspace_id == workspace.id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
AgentWorkspaceBinding.id != binding.id,
)
)
if other_binding is None:
workspace.status = AgentWorkingResourceStatus.RETIRED
workspace.active_guard = None
workspace.retired_at = now
return binding.id
@classmethod
def retire_workspace(cls, *, session: Session, tenant_id: str, workspace_id: str) -> str | None:
workspace = session.scalar(
select(AgentWorkspace)
.where(
AgentWorkspace.id == workspace_id,
AgentWorkspace.tenant_id == tenant_id,
AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE,
)
.with_for_update()
)
if workspace is None:
return None
now = naive_utc_now()
workspace.status = AgentWorkingResourceStatus.RETIRED
workspace.active_guard = None
workspace.retired_at = now
bindings = session.scalars(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.workspace_id == workspace.id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
)
).all()
for binding in bindings:
binding.status = AgentWorkingResourceStatus.RETIRED
binding.retired_at = now
return workspace.id
@classmethod
def retire_all_for_app(cls, *, session: Session, tenant_id: str, app_id: str) -> list[str]:
"""Retire all ACTIVE Workspaces owned by an App in the caller's transaction."""
workspaces = session.scalars(
select(AgentWorkspace).where(
AgentWorkspace.tenant_id == tenant_id,
AgentWorkspace.app_id == app_id,
AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE,
)
).all()
retired: list[str] = []
for workspace in workspaces:
workspace_id = cls.retire_workspace(
session=session,
tenant_id=tenant_id,
workspace_id=workspace.id,
)
if workspace_id is not None:
retired.append(workspace_id)
return retired
@classmethod
def collect_retired_binding(cls, *, tenant_id: str, binding_id: str) -> None:
try:
cls._collect_retired_binding(tenant_id=tenant_id, binding_id=binding_id)
except Exception:
logger.exception(
"Failed to collect retired Agent Workspace Binding",
extra={"tenant_id": tenant_id, "binding_id": binding_id},
)
@classmethod
def _collect_retired_binding(cls, *, tenant_id: str, binding_id: str) -> None:
with session_factory.create_session() as session:
binding = session.scalar(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.id == binding_id,
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED,
)
)
if binding is None:
return
backend_binding_ref = binding.backend_binding_ref
workspace = session.scalar(
select(AgentWorkspace).where(
AgentWorkspace.id == binding.workspace_id,
AgentWorkspace.tenant_id == tenant_id,
)
)
if workspace is not None and workspace.status == AgentWorkingResourceStatus.RETIRED:
workspace_id = workspace.id
else:
workspace_id = None
if workspace_id is not None:
cls.collect_retired_workspace(tenant_id=tenant_id, workspace_id=workspace_id)
return
try:
with cls._client() as client:
client.destroy_execution_binding_sync(
DestroyExecutionBindingRequest(
binding_ref=backend_binding_ref,
destroy_workspace=False,
)
)
except Exception:
logger.exception(
"Failed to collect retired Agent Workspace Binding",
extra={"tenant_id": tenant_id, "binding_id": binding_id},
)
return
with session_factory.create_session() as session:
binding = session.scalar(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.id == binding_id,
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED,
)
)
if binding is not None:
session.delete(binding)
session.commit()
@classmethod
def collect_retired_workspace(cls, *, tenant_id: str, workspace_id: str) -> None:
try:
cls._collect_retired_workspace(tenant_id=tenant_id, workspace_id=workspace_id)
except Exception:
logger.exception(
"Failed to collect retired Agent Workspace",
extra={"tenant_id": tenant_id, "workspace_id": workspace_id},
)
@classmethod
def _collect_retired_workspace(cls, *, tenant_id: str, workspace_id: str) -> None:
with session_factory.create_session() as session:
workspace = session.scalar(
select(AgentWorkspace).where(
AgentWorkspace.id == workspace_id,
AgentWorkspace.tenant_id == tenant_id,
AgentWorkspace.status == AgentWorkingResourceStatus.RETIRED,
)
)
if workspace is None:
return
bindings = session.scalars(
select(AgentWorkspaceBinding)
.where(
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.workspace_id == workspace_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED,
)
.order_by(AgentWorkspaceBinding.created_at)
).all()
if not bindings:
logger.error(
"RETIRED Workspace has no Binding available for physical collection",
extra={"tenant_id": tenant_id, "workspace_id": workspace_id},
)
return
anchor = bindings[0]
remaining_ids = [binding.id for binding in bindings[1:]]
workspace_ref = workspace.backend_workspace_ref
binding_ref = anchor.backend_binding_ref
anchor_id = anchor.id
try:
with cls._client() as client:
client.destroy_execution_binding_sync(
DestroyExecutionBindingRequest(
binding_ref=binding_ref,
workspace_ref=workspace_ref,
destroy_workspace=True,
)
)
except Exception:
logger.exception(
"Failed to collect retired Agent Workspace",
extra={"tenant_id": tenant_id, "workspace_id": workspace_id, "binding_id": anchor_id},
)
return
with session_factory.create_session() as session:
stored_workspace = session.scalar(
select(AgentWorkspace).where(
AgentWorkspace.id == workspace_id,
AgentWorkspace.tenant_id == tenant_id,
AgentWorkspace.status == AgentWorkingResourceStatus.RETIRED,
)
)
stored_anchor = session.scalar(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.id == anchor_id,
AgentWorkspaceBinding.tenant_id == tenant_id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED,
)
)
if stored_workspace is not None:
session.delete(stored_workspace)
if stored_anchor is not None:
session.delete(stored_anchor)
session.commit()
for remaining_id in remaining_ids:
cls.collect_retired_binding(tenant_id=tenant_id, binding_id=remaining_id)
@staticmethod
def validate_binding_generation(
binding: AgentWorkspaceBinding,
*,
base_home_snapshot_id: str,
agent_config_version_id: str,
agent_config_version_kind: AgentConfigVersionKind,
) -> None:
if (
binding.base_home_snapshot_id != base_home_snapshot_id
or binding.agent_config_version_id != agent_config_version_id
or binding.agent_config_version_kind != agent_config_version_kind
):
raise AgentWorkspaceBindingGenerationMismatchError(
"ACTIVE Binding belongs to a different Agent config/Home generation"
)
@staticmethod
def _client() -> Client:
base_url = dify_config.AGENT_BACKEND_BASE_URL
if not base_url:
raise AgentWorkspaceError("Dify Agent backend is required for Workspace operations")
return Client(base_url=base_url)
__all__ = [
"AgentWorkspaceBindingGenerationMismatchError",
"AgentWorkspaceError",
"AgentWorkspaceNotFoundError",
"AgentWorkspaceService",
"WorkspaceOwnerScope",
]
+169 -269
View File
@@ -1,40 +1,39 @@
"""Resolve product locators to ACTIVE Workspace Bindings and proxy file access."""
"""Resolve and proxy sandbox file access for Agent App and workflow Agent sessions.
These services keep product-facing locators (conversation, workflow run, node)
on the API boundary and translate them into the agent backend's
``SandboxLocator`` using persisted non-sensitive runtime layer specs plus the
saved Agenton session snapshot. Upload responses stay console-facing here: the
agent backend still returns a canonical ToolFile mapping, while this API layer
re-resolves that mapping into a signed browser download URL.
"""
from __future__ import annotations
import urllib.parse
from collections.abc import Callable
from typing import Any, Literal, cast
from typing import Any
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.client import Client
from dify_agent.layers.execution_context import (
DifyExecutionContextAgentConfigVersionKind,
DifyExecutionContextLayerConfig,
)
from dify_agent.protocol import WorkspaceListResponse, WorkspaceReadResponse, WorkspaceUploadRequest
from pydantic import BaseModel
from dify_agent.protocol import RuntimeLayerSpec, SandboxLocator, build_sandbox_locator_from_layer_specs
from pydantic import BaseModel, TypeAdapter
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore
from core.app.file_access import DatabaseFileAccessController
from core.app.workflow.file_runtime import DifyWorkflowFileRuntime
from core.db.session_factory import session_factory
from factories import file_factory
from models.agent import (
Agent,
AgentConfigDraft,
AgentConfigDraftType,
AgentWorkspaceBinding,
AgentWorkspaceOwnerType,
)
from models.model import App, Conversation
from models.workflow import WorkflowNodeExecutionModel
from services.agent.roster_service import AgentRosterService
from services.agent.workspace_service import AgentWorkspaceService, WorkspaceOwnerScope
from models.agent import AgentRuntimeSessionOwnerType, WorkflowAgentRuntimeSession, WorkflowAgentRuntimeSessionStatus
_RUNTIME_LAYER_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec])
class AgentSandboxInspectorError(Exception):
"""A sandbox inspection failure mapped to an HTTP status by the controller."""
code: str
message: str
status_code: int
@@ -47,200 +46,82 @@ class AgentSandboxInspectorError(Exception):
class AgentSandboxInfo(BaseModel):
"""Basic Agent App sandbox metadata returned after a successful availability probe."""
session_id: str
workspace_cwd: str
class AgentSandboxUploadDownload(BaseModel):
"""Signed browser download URL for one sandbox upload result."""
url: str
class AgentAppSandboxService:
def __init__(self, *, client_factory: Callable[[], Client] | None = None) -> None:
"""Inspect and proxy file access for an Agent App conversation sandbox."""
def __init__(
self,
*,
session_store: AgentAppRuntimeSessionStore | None = None,
client_factory: Callable[[], Client] | None = None,
) -> None:
self._session_store = session_store or AgentAppRuntimeSessionStore()
self._client_factory = client_factory or _default_client_factory
def get_info(
self,
*,
tenant_id: str,
app_id: str,
agent_id: str,
caller_type: Literal["conversation", "build_draft"],
caller_id: str,
account_id: str,
) -> AgentSandboxInfo:
self._resolve_binding(
tenant_id=tenant_id,
app_id=app_id,
agent_id=agent_id,
caller_type=caller_type,
caller_id=caller_id,
account_id=account_id,
def get_info(self, *, tenant_id: str, app_id: str, conversation_id: str) -> AgentSandboxInfo:
locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id)
session_id, workspace_cwd = _extract_shell_workspace_or_raise(
snapshot=locator.session_snapshot,
not_found_message="this conversation's agent has no sandbox workspace",
)
return AgentSandboxInfo(workspace_cwd=".")
def list_files(
self,
*,
tenant_id: str,
app_id: str,
agent_id: str,
caller_type: Literal["conversation", "build_draft"],
caller_id: str,
account_id: str,
path: str,
) -> WorkspaceListResponse:
binding = self._resolve_binding(
tenant_id=tenant_id,
app_id=app_id,
agent_id=agent_id,
caller_type=caller_type,
caller_id=caller_id,
account_id=account_id,
return AgentSandboxInfo(
session_id=session_id,
workspace_cwd=workspace_cwd,
)
with self._client_factory() as client:
return client.list_workspace_files_sync(binding.backend_binding_ref, path)
def read_file(
self,
*,
tenant_id: str,
app_id: str,
agent_id: str,
caller_type: Literal["conversation", "build_draft"],
caller_id: str,
account_id: str,
path: str,
) -> WorkspaceReadResponse:
binding = self._resolve_binding(
tenant_id=tenant_id,
app_id=app_id,
agent_id=agent_id,
caller_type=caller_type,
caller_id=caller_id,
account_id=account_id,
)
with self._client_factory() as client:
return client.read_workspace_file_sync(binding.backend_binding_ref, path)
def list_files(self, *, tenant_id: str, app_id: str, conversation_id: str, path: str):
locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id)
return self._client_factory().list_sandbox_files_sync(locator, path)
def read_file(self, *, tenant_id: str, app_id: str, conversation_id: str, path: str):
locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id)
return self._client_factory().read_sandbox_file_sync(locator, path)
def upload_file(
self,
*,
tenant_id: str,
app_id: str,
agent_id: str,
caller_type: Literal["conversation", "build_draft"],
caller_id: str,
account_id: str,
path: str,
self, *, tenant_id: str, app_id: str, conversation_id: str, path: str
) -> AgentSandboxUploadDownload:
binding = self._resolve_binding(
locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id)
uploaded = self._client_factory().upload_sandbox_file_sync(locator, path)
return _upload_download_response(
tenant_id=tenant_id,
file_mapping=uploaded.file.model_dump(mode="python"),
)
def _resolve_locator(self, *, tenant_id: str, app_id: str, conversation_id: str) -> SandboxLocator:
stored = self._session_store.load_active_session_for_conversation(
tenant_id=tenant_id,
app_id=app_id,
agent_id=agent_id,
caller_type=caller_type,
caller_id=caller_id,
account_id=account_id,
conversation_id=conversation_id,
)
with self._client_factory() as client:
uploaded = client.upload_workspace_file_sync(
WorkspaceUploadRequest(
backend_binding_ref=binding.backend_binding_ref,
path=path,
execution_context=DifyExecutionContextLayerConfig(
tenant_id=tenant_id,
app_id=app_id,
conversation_id=caller_id if caller_type == "conversation" else None,
agent_id=agent_id,
agent_config_version_id=binding.agent_config_version_id,
agent_config_version_kind=cast(
DifyExecutionContextAgentConfigVersionKind,
binding.agent_config_version_kind.value,
),
agent_mode="agent_app",
invoke_from="debugger",
),
)
if stored is None:
raise AgentSandboxInspectorError(
"no_active_session",
"this conversation has no active sandbox session yet",
status_code=404,
)
return _upload_download_response(tenant_id=tenant_id, file_mapping=uploaded.file.model_dump(mode="python"))
@staticmethod
def _resolve_binding(
*,
tenant_id: str,
app_id: str,
agent_id: str,
caller_type: Literal["conversation", "build_draft"],
caller_id: str,
account_id: str,
) -> AgentWorkspaceBinding:
with session_factory.create_session() as session:
caller: AgentConfigDraft | Conversation | None
if caller_type == "build_draft":
agent = session.scalar(
select(Agent).where(
Agent.id == agent_id,
Agent.tenant_id == tenant_id,
)
)
if agent is None or AgentRosterService.runtime_backing_app_id(agent) != app_id:
caller = None
else:
caller = session.scalar(
select(AgentConfigDraft).where(
AgentConfigDraft.id == caller_id,
AgentConfigDraft.tenant_id == tenant_id,
AgentConfigDraft.agent_id == agent_id,
AgentConfigDraft.account_id == account_id,
AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD,
)
)
owner_scope = WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=app_id,
owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT,
owner_id=caller_id,
)
else:
caller = session.scalar(
select(Conversation)
.join(App, App.id == Conversation.app_id)
.where(
App.tenant_id == tenant_id,
Conversation.app_id == app_id,
Conversation.id == caller_id,
Conversation.from_account_id == account_id,
Conversation.is_deleted.is_(False),
)
)
owner_scope = WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=app_id,
owner_type=AgentWorkspaceOwnerType.CONVERSATION,
owner_id=caller_id,
)
if caller is None or caller.agent_workspace_binding_id is None:
raise AgentSandboxInspectorError(
"no_active_binding",
"this caller has no active Agent Workspace Binding",
status_code=404,
)
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=tenant_id,
binding_id=caller.agent_workspace_binding_id,
expected_owner_scope=owner_scope,
)
if binding is None or binding.agent_id != agent_id:
raise AgentSandboxInspectorError(
"no_active_binding",
"this caller has no active Agent Workspace Binding",
status_code=404,
)
session.expunge(binding)
return binding
return _build_locator_or_raise(
snapshot=stored.session_snapshot,
runtime_layer_specs=stored.runtime_layer_specs,
not_found_message="this conversation's agent has no sandbox workspace",
)
class WorkflowAgentSandboxService:
"""List/read/upload files in a workflow Agent node sandbox."""
def __init__(self, *, client_factory: Callable[[], Client] | None = None) -> None:
self._client_factory = client_factory or _default_client_factory
@@ -251,11 +132,11 @@ class WorkflowAgentSandboxService:
app_id: str,
workflow_run_id: str,
node_id: str,
node_execution_id: str,
node_execution_id: str | None,
path: str,
session: Session,
) -> WorkspaceListResponse:
binding = self._resolve_binding(
):
locator = self._resolve_locator(
tenant_id=tenant_id,
app_id=app_id,
workflow_run_id=workflow_run_id,
@@ -263,8 +144,7 @@ class WorkflowAgentSandboxService:
node_execution_id=node_execution_id,
session=session,
)
with self._client_factory() as client:
return client.list_workspace_files_sync(binding.backend_binding_ref, path)
return self._client_factory().list_sandbox_files_sync(locator, path)
def read_file(
self,
@@ -273,11 +153,11 @@ class WorkflowAgentSandboxService:
app_id: str,
workflow_run_id: str,
node_id: str,
node_execution_id: str,
node_execution_id: str | None,
path: str,
session: Session,
) -> WorkspaceReadResponse:
binding = self._resolve_binding(
):
locator = self._resolve_locator(
tenant_id=tenant_id,
app_id=app_id,
workflow_run_id=workflow_run_id,
@@ -285,8 +165,7 @@ class WorkflowAgentSandboxService:
node_execution_id=node_execution_id,
session=session,
)
with self._client_factory() as client:
return client.read_workspace_file_sync(binding.backend_binding_ref, path)
return self._client_factory().read_sandbox_file_sync(locator, path)
def upload_file(
self,
@@ -295,11 +174,11 @@ class WorkflowAgentSandboxService:
app_id: str,
workflow_run_id: str,
node_id: str,
node_execution_id: str,
node_execution_id: str | None,
path: str,
session: Session,
) -> AgentSandboxUploadDownload:
binding = self._resolve_binding(
locator = self._resolve_locator(
tenant_id=tenant_id,
app_id=app_id,
workflow_run_id=workflow_run_id,
@@ -307,97 +186,118 @@ class WorkflowAgentSandboxService:
node_execution_id=node_execution_id,
session=session,
)
with self._client_factory() as client:
uploaded = client.upload_workspace_file_sync(
WorkspaceUploadRequest(
backend_binding_ref=binding.backend_binding_ref,
path=path,
execution_context=DifyExecutionContextLayerConfig(
tenant_id=tenant_id,
app_id=app_id,
workflow_run_id=workflow_run_id,
node_id=node_id,
agent_id=binding.agent_id,
agent_config_version_id=binding.agent_config_version_id,
agent_config_version_kind=cast(
DifyExecutionContextAgentConfigVersionKind,
binding.agent_config_version_kind.value,
),
agent_mode="workflow_run",
invoke_from="debugger",
),
)
)
return _upload_download_response(tenant_id=tenant_id, file_mapping=uploaded.file.model_dump(mode="python"))
uploaded = self._client_factory().upload_sandbox_file_sync(locator, path)
return _upload_download_response(
tenant_id=tenant_id,
file_mapping=uploaded.file.model_dump(mode="python"),
)
@staticmethod
def _resolve_binding(
def _resolve_locator(
self,
*,
tenant_id: str,
app_id: str,
workflow_run_id: str,
node_id: str,
node_execution_id: str,
node_execution_id: str | None,
session: Session,
) -> AgentWorkspaceBinding:
execution = session.scalar(
select(WorkflowNodeExecutionModel).where(
WorkflowNodeExecutionModel.id == node_execution_id,
WorkflowNodeExecutionModel.tenant_id == tenant_id,
WorkflowNodeExecutionModel.app_id == app_id,
WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id,
WorkflowNodeExecutionModel.node_id == node_id,
)
) -> SandboxLocator:
"""Resolve one workflow Agent sandbox from product-facing identifiers.
Callers may target either a specific node execution or the current node
as a whole. When ``node_execution_id`` is provided, lookup narrows to
that execution's ACTIVE runtime-session row. When it is omitted, the
service falls back to the most recently updated ACTIVE session for the
same ``workflow_run_id + node_id`` pair so console sandbox inspection can
still work from the broader workflow/node locator.
"""
stmt = select(WorkflowAgentRuntimeSession).where(
WorkflowAgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.WORKFLOW_RUN,
WorkflowAgentRuntimeSession.tenant_id == tenant_id,
WorkflowAgentRuntimeSession.app_id == app_id,
WorkflowAgentRuntimeSession.workflow_run_id == workflow_run_id,
WorkflowAgentRuntimeSession.node_id == node_id,
WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE,
)
process_data = execution.process_data_dict if execution is not None else None
workflow_agent_binding_id = process_data.get("workflow_agent_binding_id") if process_data is not None else None
if (
execution is None
or execution.agent_workspace_binding_id is None
or not isinstance(workflow_agent_binding_id, str)
):
if node_execution_id:
stmt = stmt.where(WorkflowAgentRuntimeSession.node_execution_id == node_execution_id)
stmt = stmt.order_by(WorkflowAgentRuntimeSession.updated_at.desc()).limit(1)
row = session.scalar(stmt)
if row is None:
raise AgentSandboxInspectorError(
"no_active_binding",
"this Workflow Agent node execution has no active Workspace Binding",
"no_active_session",
"this workflow Agent node has no active sandbox session yet",
status_code=404,
)
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=tenant_id,
binding_id=execution.agent_workspace_binding_id,
expected_owner_scope=WorkspaceOwnerScope(
tenant_id=tenant_id,
app_id=app_id,
owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN,
owner_id=workflow_run_id,
owner_scope_key=f"{node_id}:{workflow_agent_binding_id}",
),
return _build_locator_or_raise(
snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs),
not_found_message="this workflow Agent node has no sandbox workspace",
)
if binding is None:
raise AgentSandboxInspectorError(
"no_active_binding",
"this Workflow Agent node execution has no active Workspace Binding",
status_code=404,
)
return binding
def _build_locator_or_raise(
*,
snapshot: CompositorSessionSnapshot,
runtime_layer_specs: list[RuntimeLayerSpec],
not_found_message: str,
) -> SandboxLocator:
try:
return build_sandbox_locator_from_layer_specs(
layer_specs=runtime_layer_specs,
session_snapshot=snapshot,
)
except ValueError as exc:
raise AgentSandboxInspectorError("no_sandbox", not_found_message, status_code=404) from exc
def _extract_shell_workspace_or_raise(
*,
snapshot: CompositorSessionSnapshot,
not_found_message: str,
) -> tuple[str, str]:
shell_layer = next((layer for layer in snapshot.layers if layer.name == "shell"), None)
if shell_layer is None:
raise AgentSandboxInspectorError("no_sandbox", not_found_message, status_code=404)
session_id = shell_layer.runtime_state.get("session_id")
workspace_cwd = shell_layer.runtime_state.get("workspace_cwd")
if not isinstance(session_id, str) or not isinstance(workspace_cwd, str):
raise AgentSandboxInspectorError("no_sandbox", not_found_message, status_code=404)
return session_id, workspace_cwd
def _deserialize_runtime_layer_specs(value: str | None) -> list[RuntimeLayerSpec]:
if not value:
return []
return _RUNTIME_LAYER_SPECS_ADAPTER.validate_json(value)
def _upload_download_response(*, tenant_id: str, file_mapping: dict[str, Any]) -> AgentSandboxUploadDownload:
"""Resolve one uploaded ToolFile mapping into a signed external download URL."""
controller = DatabaseFileAccessController()
runtime = DifyWorkflowFileRuntime(file_access_controller=controller)
try:
file = file_factory.build_from_mapping(mapping=file_mapping, tenant_id=tenant_id, access_controller=controller)
file = file_factory.build_from_mapping(
mapping=file_mapping,
tenant_id=tenant_id,
access_controller=controller,
)
url = runtime.resolve_file_url(file=file, for_external=True)
except ValueError as exc:
raise AgentSandboxInspectorError(
"workspace_upload_download_unavailable",
"uploaded Workspace file could not be converted to a download URL",
"sandbox_upload_download_unavailable",
"uploaded sandbox file could not be converted to a download URL",
status_code=502,
) from exc
if not url:
raise AgentSandboxInspectorError(
"workspace_upload_download_unavailable",
"uploaded Workspace file does not support download URL generation",
"sandbox_upload_download_unavailable",
"uploaded sandbox file does not support download URL generation",
status_code=502,
)
return AgentSandboxUploadDownload(url=_with_as_attachment(url))
@@ -415,7 +315,7 @@ def _default_client_factory() -> Client:
if not base_url:
raise AgentSandboxInspectorError(
"inspector_unavailable",
"the Workspace file inspector is not available (Agent backend not configured)",
"the sandbox file inspector is not available (agent backend not configured)",
status_code=503,
)
return Client(base_url=base_url)
+1 -14
View File
@@ -41,7 +41,6 @@ from models import Account, App, AppMode
from models.model import AppModelConfig, AppModelConfigDict, IconType, load_annotation_reply_config
from models.workflow import Workflow
from services.agent.dsl_service import AgentDslService, AgentPackage
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.agent.workflow_publish_service import WorkflowAgentPublishService
from services.dsl_content import DSL_MAX_SIZE, dsl_content_size
from services.dsl_version import check_version_compatibility
@@ -52,7 +51,6 @@ from services.errors.app import WorkflowNotFoundError
from services.plugin.dependencies_analysis import DependenciesAnalysisService
from services.workflow_draft_variable_service import WorkflowDraftVariableService
from services.workflow_service import WorkflowService
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
logger = logging.getLogger(__name__)
@@ -548,7 +546,7 @@ class AppDslService:
sync_agent_bindings=not raw_agent_packages,
)
if raw_agent_packages:
_, warnings, retirement_candidates = AgentDslService(self._session).import_workflow_packages(
_, warnings = AgentDslService(self._session).import_workflow_packages(
workflow=draft_workflow,
portable_graph=graph,
raw_packages=raw_agent_packages,
@@ -559,17 +557,6 @@ class AppDslService:
session=self._session,
draft_workflow=draft_workflow,
)
self._session.commit()
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=app.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=app.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
case AppMode.CHAT | AppMode.AGENT_CHAT | AppMode.COMPLETION:
# Initialize model config
model_config = data.get("model_config")
+5 -157
View File
@@ -1,7 +1,6 @@
import json
import logging
from collections.abc import Sequence
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Literal, NotRequired, TypedDict, cast, override
@@ -26,48 +25,22 @@ from libs.datetime_utils import naive_utc_now
from libs.login import current_user
from libs.pagination import PaginatedResult, paginate_query
from models import Account, AppStar
from models.agent import (
APP_BACKED_AGENT_SOURCES,
Agent,
AgentIconType,
AgentScope,
AgentStatus,
AgentWorkingResourceStatus,
AgentWorkspaceBinding,
)
from models.agent import APP_BACKED_AGENT_SOURCES, Agent, AgentIconType, AgentScope, AgentStatus
from models.model import App, AppMode, AppModelConfig, IconType, Site, load_annotation_reply_config
from models.tools import ApiToolProvider
from models.workflow import Workflow
from services.agent.errors import AgentNameConflictError
from services.agent.home_snapshot_service import AgentHomeSnapshotService
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.agent.workspace_service import AgentWorkspaceService
from services.billing_service import BillingService
from services.enterprise import rbac_service as enterprise_rbac_service
from services.enterprise.enterprise_service import EnterpriseService
from services.feature_service import FeatureService
from services.openapi.visibility import apply_openapi_gate, is_openapi_visible
from services.tag_service import TagService
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
from tasks.remove_app_and_related_data_task import remove_app_and_related_data_task
logger = logging.getLogger(__name__)
AppListSortBy = Literal["last_modified", "recently_created", "earliest_created"]
RecentAppMode = Literal[
AppMode.COMPLETION,
AppMode.WORKFLOW,
AppMode.CHAT,
AppMode.ADVANCED_CHAT,
AppMode.AGENT_CHAT,
]
RECENT_APP_MODES: tuple[RecentAppMode, ...] = (
AppMode.COMPLETION,
AppMode.WORKFLOW,
AppMode.CHAT,
AppMode.ADVANCED_CHAT,
AppMode.AGENT_CHAT,
)
class AppListBaseParams(BaseModel):
@@ -92,19 +65,6 @@ class StarredAppListParams(AppListBaseParams):
pass
@dataclass(frozen=True)
class RecentAppListItem:
id: str
name: str
icon_type: IconType | None
icon: str | None
icon_background: str | None
mode: RecentAppMode
author_name: str | None
updated_at: datetime
maintainer: str | None
class CreateAppParams(BaseModel):
name: str = Field(min_length=1)
description: str | None = None
@@ -126,7 +86,7 @@ class AppModelConfigResponseView:
self._session = session
def __getattr__(self, name: str) -> Any:
return getattr(self._app_model_config, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self._app_model_config, name) # noqa: no-new-getattr response adapter delegates model fields
@property
def annotation_reply_dict(self) -> Any:
@@ -141,7 +101,7 @@ class AppResponseView:
self._session = session
def __getattr__(self, name: str) -> Any:
return getattr(self._app, name) # guard-ignore: no-new-getattr -- delegates model fields
return getattr(self._app, name) # noqa: no-new-getattr response adapter delegates model fields
@property
def desc_or_prompt(self) -> str:
@@ -363,62 +323,6 @@ class AppService:
return app_models
def get_recent_apps(
self,
user_id: str,
tenant_id: str,
params: AppListParams,
session: Session,
) -> list[RecentAppListItem]:
"""Return recently modified apps as one lightweight, non-paginated projection."""
filters = self._build_app_list_filters(user_id, tenant_id, params, session)
if not filters:
return []
stmt = (
sa.select(
App.id,
App.name,
App.icon_type,
App.icon,
App.icon_background,
App.mode,
Account.name.label("author_name"),
App.updated_at,
App.maintainer,
)
.outerjoin(Account, Account.id == App.created_by)
.where(*filters, App.mode.in_(RECENT_APP_MODES))
.order_by(App.updated_at.desc())
.limit(params.limit)
)
rows = session.execute(stmt).all()
return [
RecentAppListItem(
id=str(app_id),
name=name,
icon_type=icon_type,
icon=icon,
icon_background=icon_background,
mode=cast(RecentAppMode, mode),
author_name=author_name,
updated_at=updated_at,
maintainer=maintainer,
)
for (
app_id,
name,
icon_type,
icon,
icon_background,
mode,
author_name,
updated_at,
maintainer,
) in rows
]
def get_paginate_starred_apps(
self,
user_id: str,
@@ -490,14 +394,7 @@ class AppService:
session.delete(existing_star)
def create_app(
self,
tenant_id: str,
params: CreateAppParams,
account: Account,
*,
session: Session,
) -> App:
def create_app(self, tenant_id: str, params: CreateAppParams, account: Account, *, session: Session) -> App:
"""
Create app
:param tenant_id: tenant id
@@ -962,67 +859,18 @@ class AppService:
app_was_deleted.send(app)
backing_agent = self._get_backing_agent_for_update(app, session=session)
workflow_agent_ids = session.scalars(
select(Agent.id).where(
Agent.tenant_id == app.tenant_id,
Agent.app_id == app.id,
Agent.scope == AgentScope.WORKFLOW_ONLY,
Agent.status == AgentStatus.ACTIVE,
)
).all()
account_id = current_user.id if current_user else None
if backing_agent is not None:
now = naive_utc_now()
account_id = getattr(current_user, "id", None)
backing_agent.status = AgentStatus.ARCHIVED
backing_agent.archived_by = account_id
backing_agent.archived_at = now
backing_agent.updated_by = account_id
backing_agent.updated_at = now
retired_binding_ids: list[str] = []
retired_snapshot_ids: list[str] = []
if backing_agent is not None:
bindings = session.scalars(
select(AgentWorkspaceBinding).where(
AgentWorkspaceBinding.tenant_id == app.tenant_id,
AgentWorkspaceBinding.agent_id == backing_agent.id,
AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE,
)
).all()
for binding in bindings:
binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=app.tenant_id,
binding_id=binding.id,
)
if binding_id is not None:
retired_binding_ids.append(binding_id)
retired_snapshot_ids = AgentHomeSnapshotService.retire_all_for_agent(
session=session,
tenant_id=app.tenant_id,
agent_id=backing_agent.id,
)
retired_workspace_ids = AgentWorkspaceService.retire_all_for_app(
session=session,
tenant_id=app.tenant_id,
app_id=app.id,
)
session.delete(app)
session.commit()
workflow_binding_ids, workflow_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=app.tenant_id,
agent_ids=workflow_agent_ids,
account_id=account_id,
)
enqueue_agent_resource_collection(
tenant_id=app.tenant_id,
workspace_ids=retired_workspace_ids,
binding_ids=[*retired_binding_ids, *workflow_binding_ids],
home_snapshot_ids=[*retired_snapshot_ids, *workflow_snapshot_ids],
)
# clean up web app settings
if FeatureService.get_system_features().webapp_auth.enabled:
EnterpriseService.WebAppAuth.cleanup_webapp(app.id)
+70 -36
View File
@@ -6,7 +6,9 @@ from typing import Any
from sqlalchemy import asc, desc, func, or_, select
from sqlalchemy.orm import Session
from clients.agent_backend import AgentBackendSessionCleanupPayload
from configs import dify_config
from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore
from core.app.entities.app_invoke_entities import InvokeFrom
from core.llm_generator.llm_generator import LLMGenerator
from factories import variable_factory
@@ -14,9 +16,7 @@ from graphon.variables.types import SegmentType
from libs.datetime_utils import naive_utc_now
from libs.infinite_scroll_pagination import InfiniteScrollPagination
from models import Account, ConversationVariable
from models.agent import AgentWorkspaceOwnerType
from models.model import App, Conversation, EndUser, Message
from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope
from services.errors.conversation import (
ConversationNotExistsError,
ConversationVariableNotExistsError,
@@ -24,7 +24,7 @@ from services.errors.conversation import (
LastConversationNotExistsError,
)
from services.errors.message import MessageNotExistsError
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
from tasks.agent_backend_session_cleanup_task import cleanup_conversation_agent_runtime_session
from tasks.delete_conversation_task import delete_conversation_related_data
logger = logging.getLogger(__name__)
@@ -189,30 +189,23 @@ class ConversationService:
"""
Delete a conversation only if it belongs to the given user and app context.
Conversation deletion is the product lifecycle boundary for its
Workspace. Physical collection happens only after the retire commit.
Before removing the conversation row, this best-effort lifecycle path
enumerates any ACTIVE conversation-owned Agent backend runtime sessions,
enqueues asynchronous backend cleanup for rows with persisted runtime
layer specs, and then retires the local session rows even if enqueueing
fails. Conversation deletion and related-data cleanup scheduling still
proceed when that lifecycle bookkeeping only partially succeeds.
Raises:
ConversationNotExistsError: When the conversation is not visible to the current user.
"""
conversation = cls.get_conversation(app_model, conversation_id, user, session=session)
binding_id = conversation.agent_workspace_binding_id
retired_binding_id: str | None = None
if binding_id is not None:
owner_scope = WorkspaceOwnerScope(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
owner_type=AgentWorkspaceOwnerType.CONVERSATION,
owner_id=conversation.id,
)
binding = AgentWorkspaceService.get_active_binding(
session=session,
tenant_id=app_model.tenant_id,
binding_id=binding_id,
expected_owner_scope=owner_scope,
)
if binding is None:
raise AgentWorkspaceNotFoundError("Conversation participant Binding is unavailable")
session_store = AgentAppRuntimeSessionStore()
stored_sessions = session_store.list_active_sessions_for_conversation(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
conversation_id=conversation.id,
)
try:
logger.info(
@@ -220,25 +213,66 @@ class ConversationService:
app_model.name,
conversation_id,
)
if binding_id is not None:
retired_binding_id = AgentWorkspaceService.retire_binding(
session=session,
tenant_id=app_model.tenant_id,
binding_id=binding_id,
)
if retired_binding_id is None:
raise AgentWorkspaceNotFoundError("Conversation participant Binding is unavailable")
for stored_session in stored_sessions:
try:
if stored_session.runtime_layer_specs:
payload = AgentBackendSessionCleanupPayload(
session_snapshot=stored_session.session_snapshot,
runtime_layer_specs=stored_session.runtime_layer_specs,
idempotency_key=(
f"{stored_session.scope.tenant_id}:{stored_session.scope.app_id}:"
f"{stored_session.scope.conversation_id}:agent-runtime-session-cleanup:"
f"{stored_session.scope.agent_id}:"
f"{stored_session.scope.agent_config_snapshot_id or 'no-config'}:"
f"{stored_session.backend_run_id or 'no-run'}"
),
metadata={
"tenant_id": stored_session.scope.tenant_id,
"app_id": stored_session.scope.app_id,
"conversation_id": stored_session.scope.conversation_id,
"agent_id": stored_session.scope.agent_id,
"agent_config_snapshot_id": stored_session.scope.agent_config_snapshot_id,
"previous_agent_backend_run_id": stored_session.backend_run_id,
},
)
cleanup_conversation_agent_runtime_session.delay(payload.model_dump(mode="json"))
except Exception:
logger.warning(
"Failed to enqueue Agent backend cleanup for conversation deletion: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s",
stored_session.scope.tenant_id,
stored_session.scope.app_id,
stored_session.scope.conversation_id,
stored_session.scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
finally:
try:
session_store.mark_cleaned(
scope=stored_session.scope,
backend_run_id=stored_session.backend_run_id,
)
except Exception:
logger.warning(
"Failed to retire Agent App runtime session for conversation deletion: "
"tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s",
stored_session.scope.tenant_id,
stored_session.scope.app_id,
stored_session.scope.conversation_id,
stored_session.scope.agent_id,
stored_session.backend_run_id,
exc_info=True,
)
session.delete(conversation)
session.commit()
delete_conversation_related_data.delay(conversation.id)
except Exception:
session.rollback()
raise
if retired_binding_id is not None:
enqueue_agent_resource_collection(
tenant_id=app_model.tenant_id,
binding_ids=(retired_binding_id,),
)
delete_conversation_related_data.delay(conversation.id)
@classmethod
def get_conversational_variable(
+1 -13
View File
@@ -26,7 +26,6 @@ from models import Account, ApiToken, Tenant, TenantAccountJoin, TenantAccountRo
from models.enums import ApiTokenType
from models.model import App
from models.tools import ApiToolProvider, MCPToolProvider, WorkflowToolProvider
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.app_dsl_service import AppDslService
from services.data_migration.dependency_discovery_service import DependencyDiscoveryService
from services.data_migration.entities import (
@@ -48,7 +47,6 @@ from services.tools.api_tools_manage_service import ApiToolManageService
from services.tools.mcp_tools_manage_service import MCPToolManageService
from services.tools.workflow_tools_manage_service import WorkflowToolManageService
from services.workflow_service import WorkflowService
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
@dataclass(frozen=True)
@@ -713,7 +711,7 @@ class MigrationImportService:
raise MigrationDataError(f"Referenced workflow app was not found in target tenant: {app_id}")
if account_in_session is None:
raise MigrationDataError(f"Operator account not found: {account.id}")
workflow, retirement_candidates = workflow_service.publish_workflow(
workflow = workflow_service.publish_workflow(
session=session,
app_model=app_in_session,
account=account_in_session,
@@ -723,16 +721,6 @@ class MigrationImportService:
app_in_session.workflow_id = workflow.id
app_in_session.updated_by = account.id
app_in_session.updated_at = naive_utc_now()
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=target.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=target.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
def _import_mcp_tools(
self,
@@ -17,6 +17,7 @@ from core.entities.provider_entities import (
QuotaConfiguration,
UnaddedModelConfiguration,
)
from core.plugin.entities.plugin import PluginInstallationSource
from graphon.model_runtime.entities.common_entities import I18nObject
from graphon.model_runtime.entities.model_entities import (
FetchFrom,
@@ -69,6 +70,64 @@ class SystemConfigurationResponse(BaseModel):
quota_configurations: list[QuotaConfiguration] = []
class ModelProviderCustomConfigurationSummaryResponse(BaseModel):
status: CustomConfigurationStatus
available_credentials: list[CredentialConfiguration]
current_credential_id: str | None = None
current_credential_name: str | None = None
current_credential_usable: bool
class ModelProviderSystemConfigurationSummaryResponse(BaseModel):
enabled: bool
class ModelProviderPluginSummaryResponse(BaseModel):
installation_id: str
plugin_id: str
plugin_unique_identifier: str
runtime_type: str
source: PluginInstallationSource
version: str
class ModelProviderSummaryResponse(BaseModel):
"""Fields required to render the collapsed model-provider list."""
tenant_id: str = Field(exclude=True)
provider: str
plugin_id: str
label: I18nObject
description: I18nObject | None = None
icon_small: I18nObject | None = None
icon_small_dark: I18nObject | None = None
supported_model_types: Sequence[ModelType]
configurate_methods: list[ConfigurateMethod]
preferred_provider_type: ProviderType
is_configured: bool
custom_configuration: ModelProviderCustomConfigurationSummaryResponse
system_configuration: ModelProviderSystemConfigurationSummaryResponse
model_config = ConfigDict(protected_namespaces=())
@model_validator(mode="after")
def build_icon_urls(self):
url_prefix = (
dify_config.CONSOLE_API_URL + f"/console/api/workspaces/{self.tenant_id}/model-providers/{self.provider}"
)
if self.icon_small is not None:
self.icon_small = I18nObject(
en_US=f"{url_prefix}/icon_small/en_US",
zh_Hans=f"{url_prefix}/icon_small/zh_Hans",
)
if self.icon_small_dark is not None:
self.icon_small_dark = I18nObject(
en_US=f"{url_prefix}/icon_small_dark/en_US",
zh_Hans=f"{url_prefix}/icon_small_dark/zh_Hans",
)
return self
class ProviderResponse(BaseModel):
"""
Model class for provider response.
+9 -48
View File
@@ -1,8 +1,6 @@
import logging
from collections.abc import Mapping
from enum import StrEnum
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from pydantic import BaseModel, ConfigDict, Field
from configs import dify_config
from constants.dsl_version import CURRENT_APP_DSL_VERSION
@@ -12,8 +10,6 @@ from enums.hosted_provider import HostedTrialProvider
from services.billing_service import BillingInfo, BillingService
from services.enterprise.enterprise_service import EnterpriseService
logger = logging.getLogger(__name__)
class FeatureResponseModel(BaseModel):
model_config = ConfigDict(json_schema_serialization_defaults_required=True, protected_namespaces=())
@@ -135,13 +131,6 @@ class PluginInstallationPermissionModel(FeatureResponseModel):
restrict_to_marketplace_only: bool = False
class _EnterprisePluginInstallationPermission(BaseModel):
model_config = ConfigDict(extra="ignore")
plugin_installation_scope: PluginInstallationScope = Field(alias="pluginInstallationScope")
restrict_to_marketplace_only: bool = Field(alias="restrictToMarketplaceOnly", strict=True)
class FeatureModel(FeatureResponseModel):
billing: BillingModel = BillingModel()
education: EducationModel = EducationModel()
@@ -296,14 +285,6 @@ class FeatureService:
"""Return whether Enterprise plugin credential policies must be enforced."""
return dify_config.ENTERPRISE_ENABLED
@classmethod
def get_plugin_installation_permission(cls) -> PluginInstallationPermissionModel:
"""Resolve the validated deployment-wide plugin installation policy."""
if not dify_config.ENTERPRISE_ENABLED:
return PluginInstallationPermissionModel()
return cls._resolve_plugin_installation_permission(EnterpriseService.get_info())
@classmethod
def get_license(cls) -> LicenseModel:
"""Return full license detail. Enterprise-only; requires an authenticated caller.
@@ -471,33 +452,6 @@ class FeatureService:
)
return license_model
@classmethod
def _resolve_plugin_installation_permission(
cls, enterprise_info: Mapping[str, object]
) -> PluginInstallationPermissionModel:
if "PluginInstallationPermission" not in enterprise_info:
return PluginInstallationPermissionModel()
try:
permission = _EnterprisePluginInstallationPermission.model_validate(
enterprise_info["PluginInstallationPermission"]
)
except ValidationError as exc:
# Do not attach the exception because it may contain raw Enterprise configuration values.
logger.error( # noqa: TRY400
"Invalid Enterprise plugin installation permission; denying all plugin installations: %s",
exc.errors(include_input=False),
)
return PluginInstallationPermissionModel(
plugin_installation_scope=PluginInstallationScope.NONE,
restrict_to_marketplace_only=True,
)
return PluginInstallationPermissionModel(
plugin_installation_scope=permission.plugin_installation_scope,
restrict_to_marketplace_only=permission.restrict_to_marketplace_only,
)
@classmethod
def _fulfill_params_from_enterprise(cls, features: SystemFeatureModel):
enterprise_info = EnterpriseService.get_info()
@@ -545,4 +499,11 @@ class FeatureService:
status=LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE))
)
features.plugin_installation_permission = cls._resolve_plugin_installation_permission(enterprise_info)
if "PluginInstallationPermission" in enterprise_info:
plugin_installation_info = enterprise_info["PluginInstallationPermission"]
features.plugin_installation_permission.plugin_installation_scope = plugin_installation_info[
"pluginInstallationScope"
]
features.plugin_installation_permission.restrict_to_marketplace_only = plugin_installation_info[
"restrictToMarketplaceOnly"
]
+272 -1
View File
@@ -1,18 +1,42 @@
import logging
from collections import defaultdict
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from sqlalchemy import and_, select
if TYPE_CHECKING:
from models.account import Account
from configs import dify_config
from core.db.session_factory import session_factory
from core.entities.model_entities import ModelWithProviderEntity, ProviderModelWithStatusEntity
from core.entities.provider_entities import CredentialConfiguration
from core.helper.position_helper import is_filtered
from core.plugin.entities.plugin import PluginInstallationSource
from core.plugin.entities.plugin_daemon import PluginModelProviderBinding
from core.plugin.impl.model_runtime_factory import create_plugin_model_provider_factory, create_plugin_provider_manager
from core.plugin.plugin_service import PluginService
from core.provider_manager import ProviderManager
from extensions import ext_hosting_provider
from graphon.model_runtime.entities.model_entities import ModelType, ParameterRule
from models.provider import ProviderType
from models.provider import (
Provider,
ProviderCredential,
ProviderModel,
ProviderModelCredential,
ProviderType,
TenantPreferredModelProvider,
)
from models.provider_ids import ModelProviderID
from services.entities.model_provider_entities import (
CustomConfigurationResponse,
CustomConfigurationStatus,
DefaultModelResponse,
ModelProviderCustomConfigurationSummaryResponse,
ModelProviderPluginSummaryResponse,
ModelProviderSummaryResponse,
ModelProviderSystemConfigurationSummaryResponse,
ModelWithProviderEntityResponse,
ProviderResponse,
ProviderWithModelsResponse,
@@ -24,6 +48,17 @@ from services.errors.app_model_config import ProviderNotFoundError
logger = logging.getLogger(__name__)
@dataclass(slots=True)
class _ProviderSummaryState:
has_custom_provider: bool = False
available_credentials: list[CredentialConfiguration] = field(default_factory=list)
has_custom_models: bool = False
current_credential_id: str | None = None
current_credential_name: str | None = None
current_credential_usable: bool = False
preferred_provider_type: ProviderType | None = None
class ModelProviderService:
"""
Model Provider Service
@@ -132,6 +167,242 @@ class ModelProviderService:
return provider_responses
@staticmethod
def _load_provider_summary_states(tenant_id: str) -> dict[str, _ProviderSummaryState]:
"""Load only the workspace columns required by the collapsed provider list."""
with session_factory.create_session() as session:
custom_provider_rows = session.execute(
select(
Provider.provider_name,
Provider.credential_id,
ProviderCredential.provider_name.label("credential_provider_name"),
ProviderCredential.credential_name,
)
.outerjoin(
ProviderCredential,
and_(
ProviderCredential.id == Provider.credential_id,
ProviderCredential.tenant_id == tenant_id,
),
)
.where(
Provider.tenant_id == tenant_id,
Provider.provider_type == ProviderType.CUSTOM,
Provider.is_valid.is_(True),
)
).all()
credential_rows = session.execute(
select(
ProviderCredential.id,
ProviderCredential.provider_name,
ProviderCredential.credential_name,
)
.where(ProviderCredential.tenant_id == tenant_id)
.order_by(
ProviderCredential.created_at.desc(),
ProviderCredential.id.desc(),
)
).all()
custom_model_rows = session.execute(
select(ProviderModel.provider_name.label("provider_name"))
.where(
ProviderModel.tenant_id == tenant_id,
ProviderModel.is_valid.is_(True),
)
.union(
select(ProviderModelCredential.provider_name.label("provider_name")).where(
ProviderModelCredential.tenant_id == tenant_id
)
)
).all()
preferred_provider_rows = session.execute(
select(
TenantPreferredModelProvider.provider_name,
TenantPreferredModelProvider.preferred_provider_type,
).where(TenantPreferredModelProvider.tenant_id == tenant_id)
).all()
states: defaultdict[str, _ProviderSummaryState] = defaultdict(_ProviderSummaryState)
for credential in credential_rows:
provider_name = str(ModelProviderID(credential.provider_name))
states[provider_name].available_credentials.append(
CredentialConfiguration(
credential_id=credential.id,
credential_name=credential.credential_name,
)
)
selected_provider_priorities: dict[str, bool] = {}
for provider in custom_provider_rows:
provider_name = str(ModelProviderID(provider.provider_name))
state = states[provider_name]
state.has_custom_provider = True
is_canonical_row = provider.provider_name == provider_name
if provider_name in selected_provider_priorities and not is_canonical_row:
continue
selected_provider_priorities[provider_name] = is_canonical_row
state.current_credential_id = provider.credential_id
if (
provider.credential_provider_name is not None
and str(ModelProviderID(provider.credential_provider_name)) == provider_name
):
state.current_credential_name = provider.credential_name
state.current_credential_usable = True
else:
state.current_credential_name = None
state.current_credential_usable = False
for model in custom_model_rows:
states[str(ModelProviderID(model.provider_name))].has_custom_models = True
preferred_provider_priorities: dict[str, bool] = {}
for preferred_provider in preferred_provider_rows:
provider_name = str(ModelProviderID(preferred_provider.provider_name))
is_canonical_row = preferred_provider.provider_name == provider_name
if provider_name in preferred_provider_priorities and not is_canonical_row:
continue
preferred_provider_priorities[provider_name] = is_canonical_row
states[provider_name].preferred_provider_type = preferred_provider.preferred_provider_type
return dict(states)
@staticmethod
def _is_system_provider_enabled(provider: str) -> bool:
configuration = ext_hosting_provider.hosting_configuration.provider_map.get(provider)
return bool(configuration and configuration.enabled and configuration.quotas)
@staticmethod
def _select_binding(
current_binding: PluginModelProviderBinding | None,
candidate_binding: PluginModelProviderBinding,
) -> PluginModelProviderBinding:
"""Prefer a remote-debug runtime when one shadows an installed plugin."""
if current_binding is None:
return candidate_binding
if (
candidate_binding.source == PluginInstallationSource.Remote
and current_binding.source != PluginInstallationSource.Remote
):
return candidate_binding
return current_binding
@staticmethod
def _get_preferred_provider_type(
state: _ProviderSummaryState,
*,
custom_present: bool,
system_enabled: bool,
) -> ProviderType:
if state.preferred_provider_type is not None:
return state.preferred_provider_type
if dify_config.EDITION == "CLOUD" and system_enabled:
return ProviderType.SYSTEM
if custom_present:
return ProviderType.CUSTOM
if system_enabled:
return ProviderType.SYSTEM
return ProviderType.CUSTOM
def get_provider_summary_list(
self, tenant_id: str
) -> tuple[list[ModelProviderSummaryResponse], dict[str, ModelProviderPluginSummaryResponse]]:
"""Build the complete first-screen provider projection without assembling provider configurations."""
# Read bindings first: remote-debug identity changes invalidate provider metadata
# before the provider cache is consulted.
bindings = PluginService.list_model_provider_bindings(tenant_id)
provider_entities = PluginService.fetch_plugin_model_providers(tenant_id=tenant_id)
states = self._load_provider_summary_states(tenant_id)
bindings_by_provider: dict[str, PluginModelProviderBinding] = {}
for binding in bindings:
provider_name = (
str(ModelProviderID(binding.provider))
if binding.provider.count("/") == 2
else str(ModelProviderID(f"{binding.plugin_id}/{binding.provider}"))
)
bindings_by_provider[provider_name] = self._select_binding(
bindings_by_provider.get(provider_name),
binding,
)
provider_summaries: list[ModelProviderSummaryResponse] = []
emitted_provider_names: set[str] = set()
for provider_entity in provider_entities:
if is_filtered(
include_set=dify_config.POSITION_PROVIDER_INCLUDES_SET,
exclude_set=dify_config.POSITION_PROVIDER_EXCLUDES_SET,
data=provider_entity,
name_func=lambda provider: provider.provider,
):
continue
provider_id = ModelProviderID(provider_entity.provider)
provider_name = str(provider_id)
if provider_name in emitted_provider_names:
continue
emitted_provider_names.add(provider_name)
state = states.get(provider_name, _ProviderSummaryState())
custom_configured = (
state.has_custom_provider and bool(state.available_credentials)
) or state.has_custom_models
custom_present = state.has_custom_provider or state.has_custom_models
system_enabled = self._is_system_provider_enabled(provider_name)
preferred_provider_type = self._get_preferred_provider_type(
state,
custom_present=custom_present,
system_enabled=system_enabled,
)
provider_summaries.append(
ModelProviderSummaryResponse(
tenant_id=tenant_id,
provider=provider_name,
plugin_id=provider_id.plugin_id,
label=provider_entity.label,
description=provider_entity.description,
icon_small=provider_entity.icon_small,
icon_small_dark=provider_entity.icon_small_dark,
supported_model_types=provider_entity.supported_model_types,
configurate_methods=provider_entity.configurate_methods,
preferred_provider_type=preferred_provider_type,
is_configured=custom_configured or system_enabled,
custom_configuration=ModelProviderCustomConfigurationSummaryResponse(
status=CustomConfigurationStatus.ACTIVE
if custom_configured
else CustomConfigurationStatus.NO_CONFIGURE,
available_credentials=state.available_credentials,
current_credential_id=state.current_credential_id,
current_credential_name=state.current_credential_name,
current_credential_usable=state.current_credential_usable,
),
system_configuration=ModelProviderSystemConfigurationSummaryResponse(
enabled=system_enabled,
),
)
)
plugin_bindings: dict[str, PluginModelProviderBinding] = {}
for binding in bindings_by_provider.values():
plugin_bindings[binding.plugin_id] = self._select_binding(
plugin_bindings.get(binding.plugin_id),
binding,
)
plugin_summaries = {
plugin_id: ModelProviderPluginSummaryResponse(
installation_id=binding.installation_id,
plugin_id=binding.plugin_id,
plugin_unique_identifier=binding.plugin_unique_identifier,
runtime_type=binding.runtime_type,
source=binding.source,
version=binding.version,
)
for plugin_id, binding in plugin_bindings.items()
}
return provider_summaries, plugin_summaries
def get_models_by_provider(self, tenant_id: str, provider: str) -> list[ModelWithProviderEntityResponse]:
"""
get provider models.
+2 -26
View File
@@ -19,14 +19,12 @@ from models import Account
from models.snippet import CustomizedSnippet, SnippetType
from models.workflow import Workflow
from services.agent.dsl_service import AgentDslService
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.agent.workflow_publish_service import WorkflowAgentPublishService
from services.dsl_content import DSL_MAX_SIZE, dsl_content_size
from services.dsl_version import check_version_compatibility
from services.entities.dsl_entities import CheckDependenciesResult, DslImportWarning, ImportMode, ImportStatus
from services.plugin.dependencies_analysis import DependenciesAnalysisService
from services.snippet_service import SNIPPET_FORBIDDEN_NODE_TYPES, SnippetService
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
logger = logging.getLogger(__name__)
@@ -428,7 +426,6 @@ class SnippetDslService:
self._session.flush()
# Create or update draft workflow
retirement_candidates: set[str] = set()
if workflow_data:
graph = workflow_data.get("graph", {})
raw_agent_packages = data.get("agent_packages") or {}
@@ -447,10 +444,10 @@ class SnippetDslService:
unique_hash=unique_hash,
account=account,
input_fields=input_fields,
sync_agent_bindings=False,
sync_agent_bindings=not raw_agent_packages,
)
if raw_agent_packages:
_, warnings, retirement_candidates = AgentDslService(self._session).import_workflow_packages(
_, warnings = AgentDslService(self._session).import_workflow_packages(
workflow=draft_workflow,
portable_graph=graph,
raw_packages=raw_agent_packages,
@@ -461,29 +458,8 @@ class SnippetDslService:
session=self._session,
draft_workflow=draft_workflow,
)
else:
retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft(
session=self._session,
draft_workflow=draft_workflow,
account_id=account.id,
)
WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync(
session=self._session,
draft_workflow=draft_workflow,
)
self._session.commit()
if workflow_data:
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=snippet.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=snippet.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
return snippet
def export_snippet_dsl(self, snippet: CustomizedSnippet, include_secret: bool = False) -> str:
+7 -32
View File
@@ -35,7 +35,6 @@ from models.workflow import (
WorkflowType,
)
from repositories.factory import DifyAPIRepositoryFactory
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
from services.tag_service import TagService
from services.workflow_node_execution_trace_service import (
@@ -43,7 +42,6 @@ from services.workflow_node_execution_trace_service import (
assemble_workflow_node_execution_traces,
)
from services.workflow_restore import apply_published_workflow_snapshot_to_draft
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
logger = logging.getLogger(__name__)
@@ -82,7 +80,7 @@ class SnippetService:
@contextmanager
def _session_scope(self) -> Generator[Session, None, None]:
current_session = self._session
current_session = getattr(self, "_session", None)
if current_session is not None:
yield current_session
return
@@ -91,7 +89,7 @@ class SnippetService:
yield session
def _commit_if_owned(self, session: Session) -> None:
if self._session is None:
if getattr(self, "_session", None) is None:
session.commit()
@staticmethod
@@ -602,13 +600,12 @@ class SnippetService:
from services.agent.workflow_publish_service import WorkflowAgentPublishService
retirement_candidates: set[str] = set()
with self._session_scope() as session:
session.add(workflow)
session.add(snippet)
if sync_agent_bindings:
session.flush()
retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft(
WorkflowAgentPublishService.sync_agent_bindings_for_draft(
session=session,
draft_workflow=workflow,
account_id=account.id,
@@ -618,17 +615,6 @@ class SnippetService:
draft_workflow=workflow,
)
self._commit_if_owned(session)
if self._session is None:
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=snippet.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=snippet.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
return workflow
def restore_published_workflow_to_draft(
@@ -670,24 +656,13 @@ class SnippetService:
session.flush()
from services.agent.workflow_publish_service import WorkflowAgentPublishService
retirement_candidates = WorkflowAgentPublishService.restore_agent_node_bindings_to_draft(
WorkflowAgentPublishService.restore_agent_node_bindings_to_draft(
session=session,
source_workflow=source_workflow,
draft_workflow=draft_workflow,
account_id=account.id,
)
self._commit_if_owned(session)
if self._session is None:
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=snippet.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=snippet.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
return draft_workflow
def publish_workflow(
@@ -696,7 +671,7 @@ class SnippetService:
session: Session,
snippet: CustomizedSnippet,
account: Account,
) -> tuple[Workflow, set[str]]:
) -> Workflow:
"""
Publish the draft workflow as a new version.
@@ -740,7 +715,7 @@ class SnippetService:
kind=WorkflowKind.SNIPPET.value,
)
session.add(workflow)
retirement_candidates = WorkflowAgentPublishService.copy_agent_node_bindings_to_published(
WorkflowAgentPublishService.copy_agent_node_bindings_to_published(
session=session,
draft_workflow=draft_workflow,
published_workflow=workflow,
@@ -753,7 +728,7 @@ class SnippetService:
snippet.updated_by = account.id
session.add(snippet)
return workflow, retirement_candidates
return workflow
def get_all_published_workflows(
self,
+5 -28
View File
@@ -76,7 +76,6 @@ from models.model import App, AppMode
from models.tools import WorkflowToolProvider
from models.workflow import Workflow, WorkflowNodeExecutionModel, WorkflowNodeExecutionTriggeredFrom, WorkflowType
from repositories.factory import DifyAPIRepositoryFactory
from services.agent.retirement_service import WorkflowAgentRetirementService
from services.billing_service import BillingService
from services.errors.app import (
IsDraftWorkflowError,
@@ -84,7 +83,6 @@ from services.errors.app import (
WorkflowHashNotEqualError,
WorkflowNotFoundError,
)
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
@dataclass(frozen=True)
@@ -376,9 +374,8 @@ class WorkflowService:
from services.agent.workflow_publish_service import WorkflowAgentPublishService
session.flush()
retirement_candidates: set[str] = set()
if sync_agent_bindings:
retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft(
WorkflowAgentPublishService.sync_agent_bindings_for_draft(
session=session,
draft_workflow=workflow,
account_id=account.id,
@@ -391,16 +388,6 @@ class WorkflowService:
# commit db session changes
if commit:
session.commit()
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=app_model.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=app_model.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
# trigger app workflow events
if commit:
@@ -522,7 +509,7 @@ class WorkflowService:
from services.agent.workflow_publish_service import WorkflowAgentPublishService
session.flush()
retirement_candidates = WorkflowAgentPublishService.restore_agent_node_bindings_to_draft(
WorkflowAgentPublishService.restore_agent_node_bindings_to_draft(
session=session,
source_workflow=source_workflow,
draft_workflow=draft_workflow,
@@ -530,16 +517,6 @@ class WorkflowService:
)
session.commit()
binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned(
tenant_id=app_model.tenant_id,
agent_ids=retirement_candidates,
account_id=account.id,
)
enqueue_agent_resource_collection(
tenant_id=app_model.tenant_id,
binding_ids=binding_ids,
home_snapshot_ids=home_snapshot_ids,
)
app_draft_workflow_was_synced.send(app_model, synced_draft_workflow=draft_workflow)
return draft_workflow
@@ -552,7 +529,7 @@ class WorkflowService:
account: Account,
marked_name: str = "",
marked_comment: str = "",
) -> tuple[Workflow, set[str]]:
) -> Workflow:
draft_workflow_stmt = select(Workflow).where(
Workflow.tenant_id == app_model.tenant_id,
Workflow.app_id == app_model.id,
@@ -611,7 +588,7 @@ class WorkflowService:
# commit db session changes
session.add(workflow)
retirement_candidates = WorkflowAgentPublishService.copy_agent_node_bindings_to_published(
WorkflowAgentPublishService.copy_agent_node_bindings_to_published(
session=session,
draft_workflow=draft_workflow,
published_workflow=workflow,
@@ -621,7 +598,7 @@ class WorkflowService:
app_published_workflow_was_updated.send(app_model, published_workflow=workflow)
# return new workflow
return workflow, retirement_candidates
return workflow
def _validate_workflow_credentials(self, workflow: Workflow, *, session: Session) -> None:
"""
@@ -0,0 +1,71 @@
"""Celery tasks that execute Agent backend lifecycle-only session cleanup."""
from __future__ import annotations
import logging
from celery import shared_task
from clients.agent_backend.factory import create_agent_backend_run_client
from clients.agent_backend.request_builder import AgentBackendRunRequestBuilder
from clients.agent_backend.session_cleanup import (
AgentBackendSessionCleanupPayload,
cleanup_agent_backend_session,
)
from configs import dify_config
logger = logging.getLogger(__name__)
def _create_agent_backend_client():
if not (dify_config.AGENT_BACKEND_USE_FAKE or dify_config.AGENT_BACKEND_BASE_URL):
return None
return create_agent_backend_run_client(
base_url=dify_config.AGENT_BACKEND_BASE_URL,
use_fake=dify_config.AGENT_BACKEND_USE_FAKE,
fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO,
stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS,
stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS,
stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS,
)
def _run_cleanup_task(payload_dict: dict[str, object]) -> None:
payload = AgentBackendSessionCleanupPayload.model_validate(payload_dict)
result = cleanup_agent_backend_session(
payload=payload,
client=_create_agent_backend_client(),
request_builder=AgentBackendRunRequestBuilder(),
)
if result.status == "succeeded":
return
log_fields = {
"tenant_id": payload.metadata.get("tenant_id"),
"app_id": payload.metadata.get("app_id"),
"workflow_run_id": payload.metadata.get("workflow_run_id"),
"node_id": payload.metadata.get("node_id"),
"conversation_id": payload.metadata.get("conversation_id"),
"agent_id": payload.metadata.get("agent_id"),
"previous_agent_backend_run_id": payload.metadata.get("previous_agent_backend_run_id"),
"failed_agent_backend_run_id": payload.metadata.get("failed_agent_backend_run_id"),
"cleanup_run_id": result.cleanup_run_id,
"reason": result.reason,
}
if result.status == "skipped":
logger.info("Agent backend session cleanup skipped: %s", log_fields)
return
logger.warning("Agent backend session cleanup failed: %s", log_fields)
@shared_task(queue="workflow_storage")
def cleanup_workflow_agent_runtime_session(payload_dict: dict[str, object]) -> None:
"""Run one workflow-owned Agent backend cleanup payload."""
_run_cleanup_task(payload_dict)
@shared_task(queue="conversation")
def cleanup_conversation_agent_runtime_session(payload_dict: dict[str, object]) -> None:
"""Run one conversation-owned Agent backend cleanup payload."""
_run_cleanup_task(payload_dict)
@@ -51,7 +51,6 @@ def resume_agent_app_execution(*, conversation_id: str, form_id: str) -> None:
app_model=app_model,
user=user,
conversation_id=conversation_id,
form_id=form_id,
invoke_from=_resolve_invoke_from(conversation),
session=db.session(),
)
-75
View File
@@ -1,75 +0,0 @@
"""Asynchronously collect retired Agent working resources."""
from __future__ import annotations
import logging
from collections.abc import Iterable
from celery import shared_task
from services.agent.home_snapshot_service import AgentHomeSnapshotService
from services.agent.workspace_service import AgentWorkspaceService
logger = logging.getLogger(__name__)
@shared_task(queue="retention")
def collect_agent_resources(
*,
tenant_id: str,
binding_ids: list[str],
workspace_ids: list[str],
home_snapshot_ids: list[str],
) -> None:
"""Collect only the explicitly identified RETIRED resources."""
collectors = (
(workspace_ids, "workspace_id", AgentWorkspaceService.collect_retired_workspace),
(binding_ids, "binding_id", AgentWorkspaceService.collect_retired_binding),
(
home_snapshot_ids,
"home_snapshot_id",
AgentHomeSnapshotService.collect_retired_home_snapshot,
),
)
for resource_ids, argument_name, collector in collectors:
for resource_id in resource_ids:
try:
collector(tenant_id=tenant_id, **{argument_name: resource_id})
except Exception:
logger.exception(
"Failed to collect retired Agent resource",
extra={
"tenant_id": tenant_id,
"resource_type": argument_name.removesuffix("_id"),
"resource_id": resource_id,
},
)
def enqueue_agent_resource_collection(
*,
tenant_id: str,
binding_ids: Iterable[str] = (),
workspace_ids: Iterable[str] = (),
home_snapshot_ids: Iterable[str] = (),
) -> None:
"""Best-effort enqueue of physical collection after retire has committed."""
payload = {
"binding_ids": sorted({resource_id for resource_id in binding_ids if resource_id}),
"workspace_ids": sorted({resource_id for resource_id in workspace_ids if resource_id}),
"home_snapshot_ids": sorted({resource_id for resource_id in home_snapshot_ids if resource_id}),
}
if not any(payload.values()):
return
try:
collect_agent_resources.delay(tenant_id=tenant_id, **payload)
except Exception:
logger.exception(
"Failed to enqueue retired Agent resource collection",
extra={"tenant_id": tenant_id, **payload},
)
__all__ = ["collect_agent_resources", "enqueue_agent_resource_collection"]
@@ -5,17 +5,25 @@ from typing import Any, cast
import click
import sqlalchemy as sa
from agenton.compositor import CompositorSessionSnapshot
from celery import shared_task
from dify_agent.protocol import RuntimeLayerSpec
from pydantic import JsonValue, TypeAdapter
from sqlalchemy import delete, select
from sqlalchemy.engine import CursorResult
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.orm import sessionmaker
from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload
from configs import dify_config
from core.db.session_factory import session_factory
from extensions.ext_database import db
from libs.archive_storage import ArchiveStorageNotConfiguredError, get_archive_storage
from libs.datetime_utils import naive_utc_now
from models import (
AgentRuntimeSession,
AgentRuntimeSessionOwnerType,
AgentRuntimeSessionStatus,
ApiToken,
AppAnnotationHitHistory,
AppAnnotationSetting,
@@ -50,8 +58,13 @@ from models.workflow import (
)
from repositories.factory import DifyAPIRepositoryFactory
from services.api_token_service import ApiTokenCache
from tasks.agent_backend_session_cleanup_task import (
cleanup_conversation_agent_runtime_session,
cleanup_workflow_agent_runtime_session,
)
logger = logging.getLogger(__name__)
_RUNTIME_LAYER_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec])
@shared_task(queue="app_deletion", bind=True, max_retries=3)
@@ -59,6 +72,7 @@ def remove_app_and_related_data_task(self, tenant_id: str, app_id: str):
logger.info(click.style(f"Start deleting app and related data: {tenant_id}:{app_id}", fg="green"))
start_at = time.perf_counter()
try:
_cleanup_active_agent_runtime_sessions_for_app(tenant_id, app_id)
# Delete related data
_delete_app_model_configs(tenant_id, app_id)
_delete_app_site(tenant_id, app_id)
@@ -99,6 +113,143 @@ def remove_app_and_related_data_task(self, tenant_id: str, app_id: str):
raise self.retry(exc=e, countdown=60) # Retry after 60 seconds
def _cleanup_active_agent_runtime_sessions_for_app(tenant_id: str, app_id: str, *, batch_size: int = 100) -> None:
"""Best-effort fan-out for ACTIVE Agent runtime sessions during app deletion.
App deletion must not block on synchronous Agent backend lifecycle work, so
this helper scans ACTIVE ``agent_runtime_sessions`` rows in batches,
dispatches owner-specific cleanup tasks only when enough persisted data
exists to replay a lifecycle-only run, and then marks each visited row
``CLEANED`` locally regardless of enqueue outcome. The local retirement is
the contract that lets the rest of app deletion continue even when backend
cleanup dispatch is skipped or fails.
"""
if batch_size <= 0:
raise ValueError("batch_size must be positive")
while True:
with session_factory.create_session() as session:
row_ids = session.scalars(
select(AgentRuntimeSession.id)
.where(
AgentRuntimeSession.tenant_id == tenant_id,
AgentRuntimeSession.app_id == app_id,
AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE,
)
.order_by(AgentRuntimeSession.updated_at.asc())
.limit(batch_size)
).all()
if not row_ids:
return
retired_count = 0
for row_id in row_ids:
with session_factory.create_session() as session:
row = session.get(AgentRuntimeSession, row_id)
if row is None or row.status != AgentRuntimeSessionStatus.ACTIVE:
retired_count += 1
continue
try:
payload = _build_agent_runtime_session_cleanup_payload(row)
if payload is not None:
_enqueue_agent_runtime_session_cleanup(row=row, payload=payload)
except Exception:
logger.warning(
"Failed to enqueue Agent backend cleanup during app deletion: "
"tenant_id=%s app_id=%s owner_type=%s conversation_id=%s workflow_run_id=%s "
"node_id=%s agent_id=%s backend_run_id=%s",
row.tenant_id,
row.app_id,
row.owner_type,
row.conversation_id,
row.workflow_run_id,
row.node_id,
row.agent_id,
row.backend_run_id,
exc_info=True,
)
finally:
try:
row.status = AgentRuntimeSessionStatus.CLEANED
row.cleaned_at = naive_utc_now()
session.commit()
retired_count += 1
except Exception:
session.rollback()
logger.warning(
"Failed to retire Agent runtime session during app deletion: "
"tenant_id=%s app_id=%s owner_type=%s conversation_id=%s workflow_run_id=%s "
"node_id=%s agent_id=%s backend_run_id=%s",
row.tenant_id,
row.app_id,
row.owner_type,
row.conversation_id,
row.workflow_run_id,
row.node_id,
row.agent_id,
row.backend_run_id,
exc_info=True,
)
if retired_count == 0:
logger.warning(
"Failed to retire any active Agent runtime sessions during app deletion: tenant_id=%s app_id=%s",
tenant_id,
app_id,
)
return
def _build_agent_runtime_session_cleanup_payload(
row: AgentRuntimeSession,
) -> AgentBackendSessionCleanupPayload | None:
runtime_layer_specs = _RUNTIME_LAYER_SPECS_ADAPTER.validate_json(row.composition_layer_specs or "[]")
if not runtime_layer_specs:
return None
metadata: dict[str, JsonValue] = {
"tenant_id": row.tenant_id,
"app_id": row.app_id,
"agent_id": row.agent_id,
"agent_config_snapshot_id": row.agent_config_snapshot_id,
"previous_agent_backend_run_id": row.backend_run_id,
}
if row.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION:
metadata["conversation_id"] = row.conversation_id
idempotency_key = (
f"{row.tenant_id}:{row.app_id}:{row.conversation_id}:"
f"{row.agent_id}:app-delete-cleanup:{row.id or row.backend_run_id or 'no-session-id'}"
)
else:
metadata["workflow_run_id"] = row.workflow_run_id
metadata["node_id"] = row.node_id
idempotency_key = (
f"{row.tenant_id}:{row.app_id}:{row.workflow_run_id}:{row.node_id}:"
f"{row.agent_id}:app-delete-cleanup:{row.id or row.backend_run_id or 'no-session-id'}"
)
return AgentBackendSessionCleanupPayload(
session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot),
runtime_layer_specs=runtime_layer_specs,
idempotency_key=idempotency_key,
metadata=metadata,
)
def _enqueue_agent_runtime_session_cleanup(
*,
row: AgentRuntimeSession,
payload: AgentBackendSessionCleanupPayload,
) -> None:
payload_dict = payload.model_dump(mode="json")
if row.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION:
cleanup_conversation_agent_runtime_session.delay(payload_dict)
return
cleanup_workflow_agent_runtime_session.delay(payload_dict)
def _delete_app_model_configs(tenant_id: str, app_id: str):
def del_model_config(session, model_config_id: str):
session.execute(
+1 -3
View File
@@ -13,7 +13,6 @@ from celery import shared_task
from sqlalchemy import select
from core.db.session_factory import session_factory
from core.workflow.node_execution_process_data import preserve_workflow_agent_binding_id
from graphon.entities.workflow_node_execution import (
WorkflowNodeExecution,
)
@@ -145,9 +144,8 @@ def _update_node_execution_from_domain(node_execution: WorkflowNodeExecutionMode
# Update serialized data
json_converter = WorkflowRuntimeTypeConverter()
node_execution.inputs = json.dumps(json_converter.to_json_encodable(execution.inputs)) if execution.inputs else "{}"
process_data = preserve_workflow_agent_binding_id(node_execution.process_data_dict, execution.process_data)
node_execution.process_data = (
json.dumps(json_converter.to_json_encodable(process_data)) if process_data is not None else "{}"
json.dumps(json_converter.to_json_encodable(execution.process_data)) if execution.process_data else "{}"
)
node_execution.outputs = (
json.dumps(json_converter.to_json_encodable(execution.outputs)) if execution.outputs else "{}"
@@ -2,94 +2,114 @@ extend = "../../.ruff.toml"
src = ["../.."]
[lint]
extend-select = ["ANN401", "ARG"]
extend-select = ["ANN401", "ARG", "TID251"]
# Existing strict-mode debt. Remove a file entry when bringing it under strict checking.
[lint.per-file-ignores]
"controllers/console/test_apikey.py" = ["ARG002"]
"controllers/openapi/test_app_dsl.py" = ["ARG002"]
"controllers/service_api/dataset/test_dataset.py" = ["ARG002"]
"controllers/web/test_conversation.py" = ["ARG002"]
"controllers/web/test_human_input_form.py" = ["ARG001"]
"controllers/web/test_wraps.py" = ["ARG002"]
"core/app/layers/test_pause_state_persist_layer.py" = ["ARG002"]
"core/rag/pipeline/test_queue_integration.py" = ["ARG002", "TID251"]
"core/rag/retrieval/test_dataset_retrieval_integration.py" = ["ARG002"]
"models/test_conversation_message_inputs.py" = ["ARG001"]
"core/rag/pipeline/test_queue_integration.py" = ["ANN401", "TID251", "ARG"]
"models/test_types_enum_text.py" = ["ANN401", "TID251"]
"repositories/test_sqlalchemy_api_workflow_run_repository.py" = ["ARG002", "ARG005"]
"repositories/test_workflow_run_repository.py" = ["ARG002"]
"services/auth/test_auth_integration.py" = ["ARG002"]
"services/dataset_collection_binding.py" = ["ARG002"]
"services/document_service_status.py" = ["ARG002"]
"services/rag_pipeline/test_rag_pipeline_service_db.py" = ["ARG002"]
"services/recommend_app/test_database_retrieval.py" = ["ARG002"]
"services/test_account_service.py" = ["ARG002"]
"services/test_advanced_prompt_template_service.py" = ["ARG002"]
"services/test_app_dsl_service.py" = ["ANN401", "ARG001", "ARG002", "ARG005", "TID251"]
"services/test_app_service.py" = ["ARG002"]
"services/test_attachment_service.py" = ["ARG002"]
"services/test_conversation_variable_updater.py" = ["ARG002"]
"services/test_dataset_service_batch_update_document_status.py" = ["ARG002"]
"services/test_delete_archived_workflow_run.py" = ["ARG002"]
"services/test_document_service_rename_document.py" = ["ARG001"]
"services/test_end_user_service.py" = ["ARG002"]
"services/test_feature_service.py" = ["ARG002"]
"services/test_file_service.py" = ["ARG002"]
"services/test_file_service_zip_and_lookup.py" = ["TID251"]
"services/test_messages_clean_service.py" = ["ARG002", "S110"]
"services/test_metadata_partial_update.py" = ["ARG002"]
"services/test_metadata_service.py" = ["ARG002"]
"services/test_model_load_balancing_service.py" = ["ARG002"]
"services/test_model_provider_service.py" = ["ARG002"]
"services/test_ops_service.py" = ["ARG002"]
"services/test_webapp_auth_service.py" = ["ARG002"]
"services/test_webhook_service.py" = ["ARG002"]
"services/test_workflow_draft_variable_service.py" = ["ARG002"]
"services/test_workflow_run_service.py" = ["ARG002"]
"services/test_workflow_service.py" = ["ARG002"]
"services/test_workspace_service.py" = ["ARG002"]
"services/tools/test_api_tools_manage_service.py" = ["ARG002"]
"services/tools/test_mcp_tools_manage_service.py" = ["ARG002", "ARG005"]
"services/tools/test_tools_transform_service.py" = ["ARG002"]
"services/workflow/test_workflow_converter.py" = ["ARG002"]
"tasks/test_add_document_to_index_task.py" = ["ARG002"]
"tasks/test_batch_clean_document_task.py" = ["ARG002"]
"tasks/test_batch_create_segment_to_index_task.py" = ["ARG001", "ARG002"]
"tasks/test_clean_dataset_task.py" = ["T201"]
"tasks/test_clean_notion_document_task.py" = ["ARG002"]
"tasks/test_create_segment_to_index_task.py" = ["ARG002"]
"tasks/test_dataset_indexing_task.py" = ["ARG002"]
"tasks/test_deal_dataset_vector_index_task.py" = ["ARG002"]
"tasks/test_delete_segment_from_index_task.py" = ["ARG002"]
"tasks/test_disable_segment_from_index_task.py" = ["ARG002"]
"tasks/test_disable_segments_from_index_task.py" = ["ARG002"]
"tasks/test_document_indexing_sync_task.py" = ["ARG002"]
"tasks/test_document_indexing_task.py" = ["ARG002"]
"tasks/test_document_indexing_update_task.py" = ["ARG002"]
"tasks/test_duplicate_document_indexing_task.py" = ["ARG002"]
"tasks/test_enable_segments_to_index_task.py" = ["ARG002"]
"tasks/test_mail_change_mail_task.py" = ["ARG002"]
"tasks/test_mail_email_code_login_task.py" = ["ARG002"]
"tasks/test_mail_human_input_delivery_task.py" = ["ARG001"]
"tasks/test_mail_inner_task.py" = ["ARG002"]
"tasks/test_mail_invite_member_task.py" = ["ARG002"]
"tasks/test_mail_owner_transfer_task.py" = ["ARG002"]
"tasks/test_mail_register_task.py" = ["ARG002"]
"tasks/test_rag_pipeline_run_tasks.py" = ["ARG002"]
"test_workflow_pause_integration.py" = ["T201"]
"services/test_app_dsl_service.py" = ["ANN401", "TID251", "ARG"]
"services/test_file_service_zip_and_lookup.py" = ["ANN401", "TID251", "ARG"]
"trigger/conftest.py" = ["ANN401", "TID251"]
"trigger/test_trigger_e2e.py" = ["ANN401", "ARG001", "TID251"]
"workflow/nodes/code_executor/test_code_javascript.py" = ["ARG002"]
"workflow/nodes/code_executor/test_code_jinja2.py" = ["ARG002"]
"workflow/nodes/code_executor/test_code_python3.py" = ["ARG002"]
"trigger/test_trigger_e2e.py" = ["ANN401", "TID251", "ARG"]
"controllers/console/app/test_app_apis.py" = ["ARG"]
"controllers/console/app/test_app_import_api.py" = ["ARG"]
"controllers/console/auth/test_oauth.py" = ["ARG"]
"controllers/console/auth/test_password_reset.py" = ["ARG"]
"controllers/console/datasets/test_data_source.py" = ["ARG"]
"controllers/console/test_apikey.py" = ["ARG"]
"controllers/console/workspace/test_tool_provider.py" = ["ARG"]
"controllers/mcp/test_mcp.py" = ["ARG"]
"controllers/openapi/test_app_dsl.py" = ["ARG"]
"controllers/openapi/test_workspaces.py" = ["ARG"]
"controllers/service_api/dataset/test_dataset.py" = ["ARG"]
"controllers/web/test_conversation.py" = ["ARG"]
"controllers/web/test_human_input_form.py" = ["ARG"]
"controllers/web/test_wraps.py" = ["ARG"]
"core/app/layers/test_pause_state_persist_layer.py" = ["ARG"]
"core/rag/retrieval/test_dataset_retrieval_integration.py" = ["ARG"]
"models/test_conversation_message_inputs.py" = ["ARG"]
"models/test_conversation_status_count.py" = ["ARG"]
"repositories/test_sqlalchemy_api_workflow_run_repository.py" = ["ARG"]
"repositories/test_workflow_run_repository.py" = ["ARG"]
"services/auth/test_api_key_auth_service.py" = ["ARG"]
"services/auth/test_auth_integration.py" = ["ARG"]
"services/dataset_collection_binding.py" = ["ARG"]
"services/dataset_service_update_delete.py" = ["ARG"]
"services/document_service_status.py" = ["ARG"]
"services/enterprise/test_account_deletion_sync.py" = ["ARG"]
"services/plugin/test_plugin_parameter_service.py" = ["ARG"]
"services/plugin/test_plugin_service.py" = ["ARG"]
"services/rag_pipeline/test_rag_pipeline_service_db.py" = ["ARG"]
"services/recommend_app/test_database_retrieval.py" = ["ARG"]
"services/test_account_service.py" = ["ARG"]
"services/test_advanced_prompt_template_service.py" = ["ARG"]
"services/test_annotation_service.py" = ["ARG"]
"services/test_api_based_extension_service.py" = ["ARG"]
"services/test_api_token_service.py" = ["ARG"]
"services/test_app_generate_service.py" = ["ARG"]
"services/test_app_service.py" = ["ARG"]
"services/test_attachment_service.py" = ["ARG"]
"services/test_conversation_variable_updater.py" = ["ARG"]
"services/test_dataset_permission_service.py" = ["ARG"]
"services/test_dataset_service_batch_update_document_status.py" = ["ARG"]
"services/test_dataset_service_retrieval.py" = ["ARG"]
"services/test_delete_archived_workflow_run.py" = ["ARG"]
"services/test_document_service_rename_document.py" = ["ARG"]
"services/test_end_user_service.py" = ["ARG"]
"services/test_feature_service.py" = ["ARG"]
"services/test_feedback_service.py" = ["ARG"]
"services/test_file_service.py" = ["ARG"]
"services/test_human_input_delivery_test_service.py" = ["ARG"]
"services/test_message_service.py" = ["ARG"]
"services/test_messages_clean_service.py" = ["ARG", "S110"]
"services/test_metadata_partial_update.py" = ["ARG"]
"services/test_metadata_service.py" = ["ARG"]
"services/test_model_load_balancing_service.py" = ["ARG"]
"services/test_model_provider_service.py" = ["ARG"]
"services/test_oauth_server_service.py" = ["ARG"]
"services/test_ops_service.py" = ["ARG"]
"services/test_saved_message_service.py" = ["ARG"]
"services/test_web_conversation_service.py" = ["ARG"]
"services/test_webapp_auth_service.py" = ["ARG"]
"services/test_webhook_service.py" = ["ARG"]
"services/test_workflow_app_service.py" = ["ARG"]
"services/test_workflow_draft_variable_service.py" = ["ARG"]
"services/test_workflow_run_service.py" = ["ARG"]
"services/test_workflow_service.py" = ["ARG"]
"services/test_workspace_service.py" = ["ARG"]
"services/tools/test_api_tools_manage_service.py" = ["ARG"]
"services/tools/test_mcp_tools_manage_service.py" = ["ARG"]
"services/tools/test_tools_transform_service.py" = ["ARG"]
"services/workflow/test_workflow_converter.py" = ["ARG"]
"tasks/test_add_document_to_index_task.py" = ["ARG"]
"tasks/test_batch_clean_document_task.py" = ["ARG"]
"tasks/test_batch_create_segment_to_index_task.py" = ["ARG"]
"tasks/test_clean_dataset_task.py" = ["T201"]
"tasks/test_clean_notion_document_task.py" = ["ARG"]
"tasks/test_create_segment_to_index_task.py" = ["ARG"]
"tasks/test_dataset_indexing_task.py" = ["ARG"]
"tasks/test_deal_dataset_vector_index_task.py" = ["ARG"]
"tasks/test_delete_segment_from_index_task.py" = ["ARG"]
"tasks/test_disable_segment_from_index_task.py" = ["ARG"]
"tasks/test_disable_segments_from_index_task.py" = ["ARG"]
"tasks/test_document_indexing_sync_task.py" = ["ARG"]
"tasks/test_document_indexing_task.py" = ["ARG"]
"tasks/test_document_indexing_update_task.py" = ["ARG"]
"tasks/test_duplicate_document_indexing_task.py" = ["ARG"]
"tasks/test_enable_segments_to_index_task.py" = ["ARG"]
"tasks/test_mail_change_mail_task.py" = ["ARG"]
"tasks/test_mail_email_code_login_task.py" = ["ARG"]
"tasks/test_mail_human_input_delivery_task.py" = ["ARG"]
"tasks/test_mail_inner_task.py" = ["ARG"]
"tasks/test_mail_invite_member_task.py" = ["ARG"]
"tasks/test_mail_owner_transfer_task.py" = ["ARG"]
"tasks/test_mail_register_task.py" = ["ARG"]
"tasks/test_rag_pipeline_run_tasks.py" = ["ARG"]
"test_workflow_pause_integration.py" = ["T201"]
"workflow/nodes/code_executor/test_code_javascript.py" = ["ARG"]
"workflow/nodes/code_executor/test_code_jinja2.py" = ["ARG"]
"workflow/nodes/code_executor/test_code_python3.py" = ["ARG"]
"workflow/nodes/code_executor/test_utils.py" = ["T201"]
[lint.flake8-tidy-imports.banned-api."flask_restx.reqparse"]
msg = "Use Pydantic payload/query models instead of reqparse."
[lint.flake8-tidy-imports.banned-api."flask_restx.reqparse.RequestParser"]
msg = "Use Pydantic payload/query models instead of reqparse."
[lint.flake8-tidy-imports.banned-api."typing.Any"]
msg = "Use object, Protocol, TypedDict, TypeVar, ParamSpec, or a localized cast instead."
@@ -1,85 +1,494 @@
"""Integration coverage for Notion page bindings backed by persisted documents."""
"""Testcontainers integration tests for controllers.console.datasets.data_source endpoints."""
from inspect import unwrap
from unittest.mock import MagicMock, patch
from __future__ import annotations
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.data_source import DataSourceNotionListApi
from models import Account
from controllers.console.datasets import data_source
from controllers.console.datasets.data_source import (
DataSourceApi,
DataSourceNotionDatasetSyncApi,
DataSourceNotionDocumentSyncApi,
DataSourceNotionIndexingEstimateApi,
DataSourceNotionListApi,
DataSourceNotionPreviewApi,
)
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
def test_notion_page_is_marked_bound_from_persisted_document(
flask_app_with_containers: Flask,
db_session_with_containers: Session,
) -> None:
tenant_id = str(uuid4())
dataset_id = str(uuid4())
account = Account(name="Test User", email="user@example.com")
account.id = str(uuid4())
document = Document(
tenant_id=tenant_id,
dataset_id=dataset_id,
position=1,
data_source_type=DataSourceType.NOTION_IMPORT,
data_source_info='{"notion_page_id": "page-1"}',
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()
runtime = MagicMock(
get_online_document_pages=lambda **_kwargs: iter(
[
MagicMock(
result=[
MagicMock(
workspace_id="workspace-1",
workspace_name="Workspace",
workspace_icon=None,
pages=[
MagicMock(
page_id="page-1",
page_name="Page",
type="page",
parent_id="parent",
page_icon=None,
)
],
)
]
)
]
),
datasource_provider_type=lambda: None,
)
@pytest.fixture
def current_user() -> Account:
account = Account(name="Test User", email="u1@example.com")
account.id = "u1"
return account
with (
flask_app_with_containers.test_request_context(f"/?credential_id=c1&dataset_id={dataset_id}"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "token"},
),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(data_source_type="notion_import"),
),
patch(
"core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime",
return_value=runtime,
),
@pytest.fixture
def mock_engine() -> Iterator[None]:
with patch.object(
type(data_source.db),
"engine",
new_callable=PropertyMock,
return_value=MagicMock(),
):
response, status = unwrap(DataSourceNotionListApi().get)(
DataSourceNotionListApi(), db_session_with_containers, tenant_id, account
yield
class TestDataSourceApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_success(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.get)
binding = DataSourceOauthBinding(
tenant_id="tenant-1",
access_token="token",
provider="notion",
source_info={
"workspace_name": "Workspace",
"workspace_id": "workspace-1",
"workspace_icon": None,
"total": 1,
"pages": [
{
"page_id": "page-1",
"page_name": "Page",
"page_icon": {"type": "emoji", "emoji": "P", "url": None},
"parent_id": "parent-1",
"type": "page",
}
],
},
)
binding.id = "b1"
binding.created_at = datetime(2026, 5, 25, 1, 2, 3, tzinfo=UTC)
binding.disabled = False
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.db.session.scalars",
return_value=MagicMock(all=lambda: [binding]),
),
):
response, status = method(api, "tenant-1")
assert status == 200
assert response["data"][0] == {
"id": "b1",
"provider": "notion",
"created_at": 1779670923,
"is_bound": True,
"disabled": False,
"source_info": {
"workspace_name": "Workspace",
"workspace_id": "workspace-1",
"workspace_icon": None,
"pages": [
{
"page_name": "Page",
"page_id": "page-1",
"page_icon": {"type": "emoji", "url": None, "emoji": "P"},
"parent_id": "parent-1",
"type": "page",
}
],
"total": 1,
},
"link": "http://localhost/console/api/oauth/data-source/notion",
}
def test_get_no_bindings(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.db.session.scalars",
return_value=MagicMock(all=lambda: []),
),
):
response, status = method(api, "tenant-1")
assert status == 200
assert response["data"] == []
def test_patch_enable_binding(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=True)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
response, status = method(api, session, "tenant-1", "b1", "enable")
assert status == 200
assert binding.disabled is False
def test_patch_disable_binding(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=False)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
response, status = method(api, session, "tenant-1", "b1", "disable")
assert status == 200
assert binding.disabled is True
def test_patch_binding_not_found(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
session = MagicMock()
session.scalar.return_value = None
with app.test_request_context("/"):
with pytest.raises(NotFound):
method(api, session, "tenant-1", "b1", "enable")
def test_patch_enable_already_enabled(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=False)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
with pytest.raises(ValueError):
method(api, session, "tenant-1", "b1", "enable")
def test_patch_disable_already_disabled(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=True)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
with pytest.raises(ValueError):
method(api, session, "tenant-1", "b1", "disable")
class TestDataSourceNotionListApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_credential_not_found(self, app: Flask, current_user: Account) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "tenant-1", current_user)
def test_get_success_no_dataset_id(self, app: Flask, current_user: Account, mock_engine: None) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
page = MagicMock(
page_id="p1",
page_name="Page 1",
type="page",
parent_id="parent",
page_icon=None,
)
assert status == 200
assert response["notion_info"][0]["pages"][0]["is_bound"] is True
online_document_message = MagicMock(
result=[
MagicMock(
workspace_id="w1",
workspace_name="My Workspace",
workspace_icon="icon",
pages=[page],
)
]
)
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "t"},
),
patch(
"core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime",
return_value=MagicMock(
get_online_document_pages=lambda **kw: iter([online_document_message]),
datasource_provider_type=lambda: None,
),
),
):
response, status = method(api, MagicMock(), "tenant-1", current_user)
assert status == 200
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",
page_name="Page 1",
type="page",
parent_id="parent",
page_icon=None,
)
online_document_message = MagicMock(
result=[
MagicMock(
workspace_id="w1",
workspace_name="My Workspace",
workspace_icon="icon",
pages=[page],
)
]
)
dataset = MagicMock(data_source_type="notion_import")
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(f"/?credential_id=c1&dataset_id={dataset_id}"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "t"},
),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=dataset,
),
patch(
"core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime",
return_value=MagicMock(
get_online_document_pages=lambda **kw: iter([online_document_message]),
datasource_provider_type=lambda: None,
),
),
):
response, status = method(api, db_session_with_containers, tenant_id, current_user)
assert status == 200
def test_get_invalid_dataset_type(self, app: Flask, current_user: Account, mock_engine: None) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
dataset = MagicMock(data_source_type="other_type")
with (
app.test_request_context("/?credential_id=c1&dataset_id=ds1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "t"},
),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=dataset,
),
):
with pytest.raises(ValueError):
method(api, MagicMock(), "tenant-1", current_user)
class TestDataSourceNotionPreviewApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_preview_success(self, app: Flask) -> None:
api = DataSourceNotionPreviewApi()
method = inspect.unwrap(api.get)
extractor = MagicMock(extract=lambda: [MagicMock(page_content="hello")])
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"integration_secret": "t"},
),
patch(
"controllers.console.datasets.data_source.NotionExtractor",
return_value=extractor,
),
):
response, status = method(api, "tenant-1", "p1", "page")
assert status == 200
class TestDataSourceNotionIndexingEstimateApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_post_indexing_estimate_success(self, app: Flask) -> None:
api = DataSourceNotionIndexingEstimateApi()
method = inspect.unwrap(api.post)
empty_rules: dict[str, object] = {}
payload: dict[str, object] = {
"notion_info_list": [
{
"workspace_id": "w1",
"credential_id": "c1",
"pages": [{"page_id": "p1", "type": "page"}],
}
],
"process_rule": {"rules": empty_rules},
"doc_form": IndexStructureType.PARAGRAPH_INDEX,
"doc_language": "English",
}
with (
app.test_request_context("/", method="POST", json=payload, headers={"Content-Type": "application/json"}),
patch(
"controllers.console.datasets.data_source.DocumentService.estimate_args_validate",
),
patch(
"controllers.console.datasets.data_source.IndexingRunner.indexing_estimate",
return_value=MagicMock(model_dump=lambda: {"total_pages": 1}),
),
):
response, status = method(api, MagicMock(), "tenant-1")
assert status == 200
class TestDataSourceNotionDatasetSyncApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_success(self, app: Flask) -> None:
api = DataSourceNotionDatasetSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document_by_dataset_id",
return_value=[MagicMock(id="d1")],
),
patch(
"controllers.console.datasets.data_source.document_indexing_sync_task.delay",
return_value=None,
),
):
response, status = method(api, MagicMock(), "ds-1")
assert status == 200
def test_get_dataset_not_found(self, app: Flask) -> None:
api = DataSourceNotionDatasetSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "ds-1")
class TestDataSourceNotionDocumentSyncApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_success(self, app: Flask) -> None:
api = DataSourceNotionDocumentSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.document_indexing_sync_task.delay",
return_value=None,
),
):
response, status = method(api, MagicMock(), "ds-1", "doc-1")
assert status == 200
def test_get_document_not_found(self, app: Flask) -> None:
api = DataSourceNotionDocumentSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "ds-1", "doc-1")
@@ -1,24 +1,42 @@
preset = "strict"
strict-callable-subtyping = true
project-includes = ["."]
search-path = ["../.."]
python-platform = "linux"
python-version = "3.12.0"
infer-with-first-use = true
min-severity = "warn"
# Existing strict-mode debt. Remove a file when bringing it under strict checking.
# Verify project-excludes from the repo root:
# tmp_config=$(mktemp --tmpdir=api/tests/test_containers_integration_tests pyrefly-no-excludes.XXXXXX.toml)
# awk 'BEGIN {skip=0} /^project-excludes = \[/ {skip=1; next} skip && /^\]/ {skip=0; next} !skip {print}' api/tests/test_containers_integration_tests/pyrefly.toml > "$tmp_config"
# tmp_name=$(basename "$tmp_config")
# comm -3 <(sed -n 's/^ "\(.*\)",$/\1/p' api/tests/test_containers_integration_tests/pyrefly.toml | sort) <(uv --directory api run pyrefly check --config "tests/test_containers_integration_tests/$tmp_name" --summary=none --output-format=min-text 2>/dev/null | rg '^ERROR ' | sed -E 's#^ERROR (tests/test_containers_integration_tests/[^:]+):.*#\1#' | sed 's#^tests/test_containers_integration_tests/##' | sort -u)
# rm --force "$tmp_config"
project-excludes = [
"commands/test_legacy_model_type_migration.py",
"controllers/console/app/test_app_apis.py",
"controllers/console/app/test_app_import_api.py",
"controllers/console/app/test_chat_conversation_status_count_api.py",
"controllers/console/app/test_conversation_read_timestamp.py",
"controllers/console/app/test_workflow_draft_variable.py",
"controllers/console/auth/test_email_register.py",
"controllers/console/auth/test_forgot_password.py",
"controllers/console/auth/test_oauth.py",
"controllers/console/auth/test_password_reset.py",
"controllers/console/datasets/rag_pipeline/test_rag_pipeline.py",
"controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py",
"controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py",
"controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py",
"controllers/console/datasets/test_data_source.py",
"controllers/console/explore/test_conversation.py",
"controllers/console/test_api_based_extension.py",
"controllers/console/test_apikey.py",
"controllers/console/workspace/test_tool_provider.py",
"controllers/console/workspace/test_trigger_providers.py",
"controllers/console/workspace/test_workspace_wraps.py",
"controllers/mcp/test_mcp.py",
"controllers/service_api/dataset/test_dataset.py",
"controllers/service_api/test_site.py",
"controllers/web/test_conversation.py",
"controllers/web/test_site.py",
"controllers/web/test_web_forgot_password.py",
"controllers/web/test_wraps.py",
"core/app/layers/test_pause_state_persist_layer.py",
"core/rag/pipeline/test_queue_integration.py",
@@ -39,16 +57,22 @@ project-excludes = [
"repositories/test_sqlalchemy_execution_extra_content_repository.py",
"repositories/test_sqlalchemy_workflow_node_execution_repository.py",
"repositories/test_workflow_run_repository.py",
"services/auth/test_api_key_auth_service.py",
"services/auth/test_auth_integration.py",
"services/dataset_collection_binding.py",
"services/dataset_service_update_delete.py",
"services/document_service_status.py",
"services/enterprise/test_account_deletion_sync.py",
"services/plugin/test_plugin_parameter_service.py",
"services/plugin/test_plugin_service.py",
"services/rag_pipeline/test_rag_pipeline_service_db.py",
"services/recommend_app/test_database_retrieval.py",
"services/test_account_service.py",
"services/test_advanced_prompt_template_service.py",
"services/test_agent_service.py",
"services/test_annotation_service.py",
"services/test_api_based_extension_service.py",
"services/test_api_token_service.py",
"services/test_app_dsl_service.py",
"services/test_app_generate_service.py",
"services/test_app_service.py",
@@ -73,15 +97,20 @@ project-excludes = [
"services/test_document_service_rename_document.py",
"services/test_end_user_service.py",
"services/test_feature_service.py",
"services/test_feedback_service.py",
"services/test_file_service.py",
"services/test_human_input_delivery_test.py",
"services/test_human_input_delivery_test_service.py",
"services/test_message_export_service.py",
"services/test_message_service.py",
"services/test_message_service_execution_extra_content.py",
"services/test_message_service_extra_contents.py",
"services/test_messages_clean_service.py",
"services/test_metadata_partial_update.py",
"services/test_metadata_service.py",
"services/test_model_load_balancing_service.py",
"services/test_model_provider_service.py",
"services/test_oauth_server_service.py",
"services/test_ops_service.py",
"services/test_restore_archived_workflow_run.py",
"services/test_saved_message_service.py",
@@ -95,6 +124,7 @@ project-excludes = [
"services/test_workflow_run_service.py",
"services/test_workflow_service.py",
"services/test_workspace_service.py",
"services/tools/test_api_tools_manage_service.py",
"services/tools/test_mcp_tools_manage_service.py",
"services/tools/test_tools_transform_service.py",
"services/tools/test_workflow_tools_manage_service.py",
@@ -131,6 +161,7 @@ project-excludes = [
"test_workflow_pause_integration.py",
"trigger/conftest.py",
"trigger/test_trigger_e2e.py",
"workflow/nodes/code_executor/test_code_executor.py",
"workflow/nodes/code_executor/test_code_javascript.py",
"workflow/nodes/code_executor/test_code_jinja2.py",
"workflow/nodes/code_executor/test_code_python3.py",
@@ -138,7 +169,6 @@ project-excludes = [
]
[errors]
missing-override-decorator = "error"
redundant-cast = true
unannotated-return = true
unnecessary-type-conversion = true
@@ -124,7 +124,6 @@ class TestAppDslService:
patch("services.app_service.ModelManager.for_tenant") as mock_model_manager,
patch("services.app_service.FeatureService") as mock_feature_service,
patch("services.app_service.EnterpriseService") as mock_enterprise_service,
patch("services.agent.home_snapshot_service.AgentHomeSnapshotService._client") as mock_home_snapshot_client,
):
mock_workflow_service.return_value.get_draft_workflow.return_value = None
mock_workflow_service.return_value.sync_draft_workflow.return_value = MagicMock()
@@ -143,9 +142,6 @@ class TestAppDslService:
mock_feature_service.get_system_features.return_value.webapp_auth.enabled = False
mock_enterprise_service.WebAppAuth.update_app_access_mode.return_value = None
mock_enterprise_service.WebAppAuth.cleanup_webapp.return_value = None
mock_home_snapshot_client.return_value.__enter__.return_value.initialize_home_snapshot_sync.side_effect = (
lambda request: SimpleNamespace(snapshot_ref=f"test:{request.home_snapshot_id}")
)
yield {
"workflow_service": mock_workflow_service,
@@ -1038,9 +1034,7 @@ class TestAppDslService:
)
)
imported_graph, warnings, retirement_candidates = AgentDslService(
db_session_with_containers
).import_workflow_packages(
imported_graph, warnings = AgentDslService(db_session_with_containers).import_workflow_packages(
workflow=workflow,
portable_graph=graph,
raw_packages={"agent_1": package.model_dump(mode="json")},
@@ -1049,7 +1043,6 @@ class TestAppDslService:
db_session_with_containers.commit()
assert warnings == []
assert retirement_candidates == set()
graph_bindings = [node["data"]["agent_binding"] for node in imported_graph["nodes"]]
assert all(binding["binding_type"] == WorkflowAgentBindingType.INLINE_AGENT.value for binding in graph_bindings)
assert len({binding["agent_id"] for binding in graph_bindings}) == 2
@@ -1,5 +1,4 @@
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import create_autospec, patch
import pytest
@@ -30,7 +29,6 @@ class TestAppService:
patch("services.app_service.EnterpriseService") as mock_enterprise_service,
patch("services.app_service.ModelManager.for_tenant") as mock_model_manager,
patch("services.account_service.FeatureService") as mock_account_feature_service,
patch("services.agent.home_snapshot_service.AgentHomeSnapshotService._client") as mock_home_snapshot_client,
):
# Setup default mock returns for app service
mock_feature_service.get_system_features.return_value.webapp_auth.enabled = False
@@ -44,9 +42,6 @@ class TestAppService:
mock_model_instance = mock_model_manager.return_value
mock_model_instance.get_default_model_instance.return_value = None
mock_model_instance.get_default_provider_model_name.return_value = ("openai", "gpt-3.5-turbo")
mock_home_snapshot_client.return_value.__enter__.return_value.initialize_home_snapshot_sync.side_effect = (
lambda request: SimpleNamespace(snapshot_ref=f"test:{request.home_snapshot_id}")
)
yield {
"feature_service": mock_feature_service,
@@ -6,11 +6,13 @@ from unittest.mock import patch
from uuid import uuid4
import pytest
from agenton.compositor import CompositorSessionSnapshot
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore
from core.app.entities.app_invoke_entities import InvokeFrom
from models import TenantAccountRole
from models import AgentRuntimeSession, AgentRuntimeSessionOwnerType, AgentRuntimeSessionStatus, TenantAccountRole
from models.account import Account, Tenant, TenantAccountJoin
from models.enums import ConversationFromSource, EndUserType
from models.model import App, Conversation, EndUser, Message, MessageAnnotation
@@ -1075,8 +1077,9 @@ class TestConversationServiceExport:
# Assert
assert result == conversation
@patch("services.conversation_service.cleanup_conversation_agent_runtime_session")
@patch("services.conversation_service.delete_conversation_related_data")
def test_delete_conversation(self, mock_delete_task, db_session_with_containers: Session):
def test_delete_conversation(self, mock_delete_task, mock_cleanup_task, db_session_with_containers: Session):
"""
Test conversation deletion with async cleanup.
@@ -1095,6 +1098,20 @@ class TestConversationServiceExport:
user,
)
conversation_id = conversation.id
runtime_session = AgentRuntimeSession(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
owner_type=AgentRuntimeSessionOwnerType.CONVERSATION,
agent_id=str(uuid4()),
agent_config_snapshot_id=str(uuid4()),
backend_run_id="backend-run-1",
session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(),
composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]',
conversation_id=conversation.id,
status=AgentRuntimeSessionStatus.ACTIVE,
)
db_session_with_containers.add(runtime_session)
db_session_with_containers.commit()
# Act - Delete the conversation
ConversationService.delete(
@@ -1109,11 +1126,27 @@ class TestConversationServiceExport:
# Step 2: Async cleanup task triggered
# The Celery task will handle cleanup of messages, annotations, etc.
mock_delete_task.delay.assert_called_once_with(conversation_id)
mock_cleanup_task.delay.assert_called_once()
cleanup_payload = mock_cleanup_task.delay.call_args.args[0]
assert cleanup_payload["metadata"]["conversation_id"] == conversation_id
assert (
cleanup_payload["idempotency_key"]
== f"{app_model.tenant_id}:{app_model.id}:{conversation_id}:agent-runtime-session-cleanup:"
f"{runtime_session.agent_id}:{runtime_session.agent_config_snapshot_id}:{runtime_session.backend_run_id}"
)
runtime_session_row = db_session_with_containers.scalar(
select(AgentRuntimeSession).where(AgentRuntimeSession.id == runtime_session.id)
)
assert runtime_session_row is not None
assert runtime_session_row.status == AgentRuntimeSessionStatus.CLEANED
@patch("services.conversation_service.cleanup_conversation_agent_runtime_session")
@patch("services.conversation_service.delete_conversation_related_data")
def test_delete_conversation_not_owned_by_account(
self,
mock_delete_task,
mock_cleanup_task,
db_session_with_containers: Session,
):
"""
@@ -1145,17 +1178,22 @@ class TestConversationServiceExport:
not_deleted = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation.id))
assert not_deleted is not None
mock_delete_task.delay.assert_not_called()
mock_cleanup_task.delay.assert_not_called()
@patch("services.conversation_service.cleanup_conversation_agent_runtime_session")
@patch("services.conversation_service.delete_conversation_related_data")
def test_delete_handles_exception_and_rollback(
self,
mock_delete_task,
mock_cleanup_task,
db_session_with_containers: Session,
):
"""
Test that delete propagates exceptions and does not trigger the cleanup task.
When a DB error occurs during deletion, the conversation row stays in place.
When a DB error occurs during deletion, the conversation row stays in
place, but any already-enqueued Agent backend cleanup remains a
best-effort terminal lifecycle action.
"""
# Arrange
app_model, user = ConversationServiceIntegrationTestDataFactory.create_app_and_account(
@@ -1165,6 +1203,20 @@ class TestConversationServiceExport:
db_session_with_containers, app_model, user
)
conversation_id = conversation.id
runtime_session = AgentRuntimeSession(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
owner_type=AgentRuntimeSessionOwnerType.CONVERSATION,
agent_id=str(uuid4()),
agent_config_snapshot_id=str(uuid4()),
backend_run_id="backend-run-rollback",
session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(),
composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]',
conversation_id=conversation.id,
status=AgentRuntimeSessionStatus.ACTIVE,
)
db_session_with_containers.add(runtime_session)
db_session_with_containers.commit()
# Act — force an error during the delete to exercise the rollback path
with patch.object(db_session_with_containers, "delete", side_effect=Exception("DB error")):
@@ -1176,9 +1228,111 @@ class TestConversationServiceExport:
session=db_session_with_containers,
)
# Assert — related-data deletion is not scheduled.
# Assert — related-data deletion is not scheduled, but the backend
# cleanup task was already enqueued before the row delete failed.
mock_delete_task.delay.assert_not_called()
mock_cleanup_task.delay.assert_called_once()
cleanup_payload = mock_cleanup_task.delay.call_args.args[0]
assert (
cleanup_payload["idempotency_key"]
== f"{app_model.tenant_id}:{app_model.id}:{conversation_id}:agent-runtime-session-cleanup:"
f"{runtime_session.agent_id}:{runtime_session.agent_config_snapshot_id}:{runtime_session.backend_run_id}"
)
# Conversation is still present because the deletion was never committed
still_there = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation_id))
assert still_there is not None
@patch("services.conversation_service.cleanup_conversation_agent_runtime_session")
@patch("services.conversation_service.delete_conversation_related_data")
def test_delete_ignores_mark_cleaned_failure(
self,
mock_delete_task,
mock_cleanup_task,
db_session_with_containers: Session,
):
app_model, user = ConversationServiceIntegrationTestDataFactory.create_app_and_account(
db_session_with_containers
)
conversation = ConversationServiceIntegrationTestDataFactory.create_conversation(
db_session_with_containers,
app_model,
user,
)
runtime_session = AgentRuntimeSession(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
owner_type=AgentRuntimeSessionOwnerType.CONVERSATION,
agent_id=str(uuid4()),
agent_config_snapshot_id=str(uuid4()),
backend_run_id="backend-run-cleanup-failure",
session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(),
composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]',
conversation_id=conversation.id,
status=AgentRuntimeSessionStatus.ACTIVE,
)
db_session_with_containers.add(runtime_session)
db_session_with_containers.commit()
with patch.object(AgentAppRuntimeSessionStore, "mark_cleaned", side_effect=RuntimeError("cleanup failed")):
ConversationService.delete(
app_model=app_model,
conversation_id=conversation.id,
user=user,
session=db_session_with_containers,
)
deleted = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation.id))
assert deleted is None
mock_delete_task.delay.assert_called_once_with(conversation.id)
mock_cleanup_task.delay.assert_called_once()
@patch("services.conversation_service.cleanup_conversation_agent_runtime_session")
@patch("services.conversation_service.delete_conversation_related_data")
def test_delete_ignores_cleanup_enqueue_failure_and_still_retires_runtime_session(
self,
mock_delete_task,
mock_cleanup_task,
db_session_with_containers: Session,
):
app_model, user = ConversationServiceIntegrationTestDataFactory.create_app_and_account(
db_session_with_containers
)
conversation = ConversationServiceIntegrationTestDataFactory.create_conversation(
db_session_with_containers,
app_model,
user,
)
conversation_id = conversation.id
runtime_session = AgentRuntimeSession(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
owner_type=AgentRuntimeSessionOwnerType.CONVERSATION,
agent_id=str(uuid4()),
agent_config_snapshot_id=str(uuid4()),
backend_run_id="backend-run-enqueue-failure",
session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(),
composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]',
conversation_id=conversation.id,
status=AgentRuntimeSessionStatus.ACTIVE,
)
db_session_with_containers.add(runtime_session)
db_session_with_containers.commit()
mock_cleanup_task.delay.side_effect = RuntimeError("queue down")
ConversationService.delete(
app_model=app_model,
conversation_id=conversation_id,
user=user,
session=db_session_with_containers,
)
deleted = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation_id))
assert deleted is None
mock_delete_task.delay.assert_called_once_with(conversation_id)
mock_cleanup_task.delay.assert_called_once()
runtime_session_row = db_session_with_containers.scalar(
select(AgentRuntimeSession).where(AgentRuntimeSession.id == runtime_session.id)
)
assert runtime_session_row is not None
assert runtime_session_row.status == AgentRuntimeSessionStatus.CLEANED
@@ -4,7 +4,7 @@ import datetime
import json
import uuid
from decimal import Decimal
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import pytest
from faker import Faker
@@ -1172,8 +1172,65 @@ class TestMessagesCleanServiceIntegration:
# Verify all messages were deleted
assert db_session_with_containers.query(Message).where(Message.id.in_(msg_ids)).count() == 0
def test_from_time_range_validation(self):
"""Test that from_time_range raises ValueError for invalid inputs."""
policy = MagicMock(spec=BillingDisabledPolicy)
now = datetime.datetime.now()
with pytest.raises(ValueError, match="start_from .* must be less than end_before"):
MessagesCleanService.from_time_range(policy, now, now)
with pytest.raises(ValueError, match="batch_size .* must be greater than 0"):
MessagesCleanService.from_time_range(policy, now - datetime.timedelta(days=1), now, batch_size=0)
def test_from_time_range_success(self):
"""Test that from_time_range creates a service with correct parameters."""
policy = MagicMock(spec=BillingDisabledPolicy)
start = datetime.datetime(2024, 1, 1)
end = datetime.datetime(2024, 2, 1)
service = MessagesCleanService.from_time_range(policy, start, end)
assert service._start_from == start
assert service._end_before == end
def test_from_days_validation(self):
"""Test that from_days raises ValueError for invalid inputs."""
policy = MagicMock(spec=BillingDisabledPolicy)
with pytest.raises(ValueError, match="days .* must be greater than or equal to 0"):
MessagesCleanService.from_days(policy, days=-1)
with pytest.raises(ValueError, match="batch_size .* must be greater than 0"):
MessagesCleanService.from_days(policy, days=30, batch_size=0)
def test_from_days_success(self):
"""Test that from_days creates a service with correct parameters."""
policy = MagicMock(spec=BillingDisabledPolicy)
with patch("services.retention.conversation.messages_clean_service.naive_utc_now") as mock_now:
fixed_now = datetime.datetime(2024, 6, 1)
mock_now.return_value = fixed_now
service = MessagesCleanService.from_days(policy, days=10)
assert service._start_from is None
assert service._end_before == fixed_now - datetime.timedelta(days=10)
def test_batch_delete_message_relations_empty(self, db_session_with_containers: Session):
"""Test that batch_delete_message_relations with empty list does nothing."""
# Get execute call count before
MessagesCleanService._batch_delete_message_relations(db_session_with_containers, [])
# No exception means success — empty list is a no-op
def test_run_calls_clean_messages(self):
"""Test that run() delegates to _clean_messages_by_time_range."""
policy = MagicMock(spec=BillingDisabledPolicy)
service = MessagesCleanService(
policy=policy,
end_before=datetime.datetime.now(),
batch_size=10,
)
with patch.object(service, "_clean_messages_by_time_range") as mock_clean:
mock_clean.return_value = {"total_deleted": 5}
result = service.run()
assert result == {"total_deleted": 5}
mock_clean.assert_called_once()
@@ -1,20 +1,16 @@
"""Unit tests for OAuthServerService with SQLite-backed database access."""
"""Testcontainers integration tests for OAuthServerService."""
from __future__ import annotations
import uuid
from collections.abc import Iterator
from typing import cast
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest
from models.engine import db
from models.model import OAuthProviderApp
from services.oauth_server import (
OAUTH_ACCESS_TOKEN_EXPIRES_IN,
@@ -27,24 +23,10 @@ from services.oauth_server import (
)
@pytest.fixture
def oauth_db() -> Iterator[Session]:
"""Provide the production database extension with an isolated SQLite provider table."""
app = Flask(__name__)
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
with app.app_context():
OAuthProviderApp.__table__.create(db.engine)
with Session(db.engine, expire_on_commit=False) as session:
yield session
class TestOAuthServerServiceGetProviderApp:
"""Verify provider lookup against a real SQLAlchemy database."""
"""DB-backed tests for get_oauth_provider_app."""
def test_get_oauth_provider_app_returns_app_when_exists(self, oauth_db: Session) -> None:
client_id = f"client-{uuid4()}"
def _create_oauth_provider_app(self, db_session_with_containers: Session, *, client_id: str) -> OAuthProviderApp:
app = OAuthProviderApp(
app_icon="icon.png",
client_id=client_id,
@@ -53,30 +35,35 @@ class TestOAuthServerServiceGetProviderApp:
redirect_uris=["https://example.com/callback"],
scope="read",
)
oauth_db.add(app)
oauth_db.commit()
db_session_with_containers.add(app)
db_session_with_containers.commit()
return app
def test_get_oauth_provider_app_returns_app_when_exists(self, db_session_with_containers: Session):
client_id = f"client-{uuid4()}"
created = self._create_oauth_provider_app(db_session_with_containers, client_id=client_id)
result = OAuthServerService.get_oauth_provider_app(client_id)
assert result is not None
assert result.client_id == client_id
assert result.id == app.id
assert result.id == created.id
def test_get_oauth_provider_app_returns_none_when_not_exists(self, oauth_db: Session) -> None:
def test_get_oauth_provider_app_returns_none_when_not_exists(self, db_session_with_containers: Session):
result = OAuthServerService.get_oauth_provider_app(f"nonexistent-{uuid4()}")
assert result is None
class TestOAuthServerServiceTokenOperations:
"""Verify Redis-backed token signing and validation branches."""
"""Redis-backed tests for token sign/validate operations."""
@pytest.fixture
def mock_redis(self):
with patch("services.oauth_server.redis_client") as mock:
yield mock
def test_sign_authorization_code_stores_and_returns_code(self, mock_redis) -> None:
def test_sign_authorization_code_stores_and_returns_code(self, mock_redis):
deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000111")
with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid):
code = OAuthServerService.sign_oauth_authorization_code("client-1", "user-1")
@@ -88,7 +75,7 @@ class TestOAuthServerServiceTokenOperations:
ex=600,
)
def test_sign_access_token_raises_bad_request_for_invalid_code(self, mock_redis) -> None:
def test_sign_access_token_raises_bad_request_for_invalid_code(self, mock_redis):
mock_redis.get.return_value = None
with pytest.raises(BadRequest, match="invalid code"):
@@ -98,13 +85,14 @@ class TestOAuthServerServiceTokenOperations:
client_id="client-1",
)
def test_sign_access_token_issues_tokens_for_valid_code(self, mock_redis) -> None:
def test_sign_access_token_issues_tokens_for_valid_code(self, mock_redis):
token_uuids = [
uuid.UUID("00000000-0000-0000-0000-000000000201"),
uuid.UUID("00000000-0000-0000-0000-000000000202"),
]
with patch("services.oauth_server.uuid.uuid4", side_effect=token_uuids):
mock_redis.get.return_value = b"user-1"
access_token, refresh_token = OAuthServerService.sign_oauth_access_token(
grant_type=OAuthGrantType.AUTHORIZATION_CODE,
code="code-1",
@@ -126,7 +114,7 @@ class TestOAuthServerServiceTokenOperations:
ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN,
)
def test_sign_access_token_raises_bad_request_for_invalid_refresh_token(self, mock_redis) -> None:
def test_sign_access_token_raises_bad_request_for_invalid_refresh_token(self, mock_redis):
mock_redis.get.return_value = None
with pytest.raises(BadRequest, match="invalid refresh token"):
@@ -136,10 +124,11 @@ class TestOAuthServerServiceTokenOperations:
client_id="client-1",
)
def test_sign_access_token_issues_new_token_for_valid_refresh(self, mock_redis) -> None:
def test_sign_access_token_issues_new_token_for_valid_refresh(self, mock_redis):
deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000301")
with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid):
mock_redis.get.return_value = b"user-1"
access_token, returned_refresh = OAuthServerService.sign_oauth_access_token(
grant_type=OAuthGrantType.REFRESH_TOKEN,
refresh_token="refresh-1",
@@ -149,14 +138,14 @@ class TestOAuthServerServiceTokenOperations:
assert access_token == str(deterministic_uuid)
assert returned_refresh == "refresh-1"
def test_sign_access_token_returns_none_for_unknown_grant_type(self, mock_redis) -> None:
def test_sign_access_token_returns_none_for_unknown_grant_type(self, mock_redis):
grant_type = cast(OAuthGrantType, "invalid-grant-type")
result = OAuthServerService.sign_oauth_access_token(grant_type=grant_type, client_id="client-1")
assert result is None
def test_sign_refresh_token_stores_with_expected_expiry(self, mock_redis) -> None:
def test_sign_refresh_token_stores_with_expected_expiry(self, mock_redis):
deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000401")
with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid):
refresh_token = OAuthServerService._sign_oauth_refresh_token("client-2", "user-2")
@@ -168,21 +157,22 @@ class TestOAuthServerServiceTokenOperations:
ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN,
)
def test_validate_access_token_returns_none_when_not_found(self, mock_redis, sqlite_engine: Engine) -> None:
def test_validate_access_token_returns_none_when_not_found(self, mock_redis, db_session_with_containers: Session):
mock_redis.get.return_value = None
session = MagicMock()
with Session(sqlite_engine) as session:
result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", session)
result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", db_session_with_containers)
assert result is None
def test_validate_access_token_loads_user_when_exists(self, mock_redis, sqlite_engine: Engine) -> None:
def test_validate_access_token_loads_user_when_exists(self, mock_redis, db_session_with_containers: Session):
mock_redis.get.return_value = b"user-88"
expected_user = MagicMock()
with Session(sqlite_engine) as session:
with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load:
result = OAuthServerService.validate_oauth_access_token("client-1", "access-token", session)
mock_load.assert_called_once_with("user-88", session)
with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load:
result = OAuthServerService.validate_oauth_access_token(
"client-1", "access-token", db_session_with_containers
)
assert result is expected_user
mock_load.assert_called_once_with("user-88", db_session_with_containers)

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