fix(workflow): coerce non-uuid variable ids and warn on skipped branches (#42750)

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Crazywoola <100913391+crazywoola@users.noreply.github.com>
This commit is contained in:
Andy yang
2026-09-25 08:16:41 +00:00
committed by GitHub
co-authored by Cursor autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Crazywoola
parent f6a1878ceb
commit d86435ee1e
30 changed files with 1307 additions and 36 deletions
@@ -132,6 +132,7 @@ class ConversationVariablesApi(Resource):
"created_at": row.created_at,
"updated_at": row.updated_at,
**row.to_variable().model_dump(),
"id": row.id,
}
)
for row in rows
+29 -1
View File
@@ -101,6 +101,10 @@ from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError,
from services.errors.llm import InvokeRateLimitError
from services.workflow_ref_service import WorkflowRefService
from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService
from services.workflow_variable_reference_validator import (
format_variable_reference_errors,
validate_variable_references,
)
logger = logging.getLogger(__name__)
@@ -401,6 +405,10 @@ class WorkflowOnlineUsersResponse(ResponseModel):
class WorkflowPublishResponse(ResponseModel):
result: str
created_at: int
warning: str | None = Field(
default=None,
description="Advisory warning for variable references that can read a skipped branch. Publish still succeeds.",
)
class SyncDraftWorkflowResponse(ResponseModel):
@@ -1266,6 +1274,21 @@ class DraftWorkflowNodeRunApi(Resource):
).model_dump(mode="json")
def _advisory_variable_reference_warning(graph_text: str | None) -> str | None:
"""Return a non-blocking publish warning. A checker failure must not fail publish."""
if not graph_text:
return None
try:
graph = json.loads(graph_text)
if not isinstance(graph, dict):
return None
issues = validate_variable_references(graph)
return format_variable_reference_errors(issues) if issues else None
except Exception:
logger.warning("Skipped advisory variable reference check", exc_info=True)
return None
@console_ns.route("/apps/<uuid:app_id>/workflows/publish")
class PublishedWorkflowApi(Resource):
@console_ns.doc("get_published_workflow")
@@ -1330,11 +1353,16 @@ class PublishedWorkflowApi(Resource):
app_model_in_session.updated_at = naive_utc_now()
workflow_created_at = TimestampField().format(workflow.created_at)
graph_text = workflow.graph
return {
warning = _advisory_variable_reference_warning(graph_text)
payload: dict[str, object] = {
"result": "success",
"created_at": workflow_created_at,
}
if warning:
payload["warning"] = warning
return payload
@console_ns.route("/apps/<uuid:app_id>/workflows/default-workflow-block-configs")
@@ -479,9 +479,12 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
:param existing_variables: List of existing conversation variables
:return: Updated list including any newly created variables
"""
# Get IDs of existing and workflow variables
# Compare stored primary keys. A non-UUID author id is rewritten by
# ConversationVariable.storage_id, so the workflow id itself will not match.
existing_ids = {var.id for var in existing_variables}
workflow_variables = {var.id: var for var in self._workflow.conversation_variables}
workflow_variables = {
ConversationVariable.storage_id(variable): variable for variable in self._workflow.conversation_variables
}
# Find missing variable IDs
missing_ids = set(workflow_variables.keys()) - existing_ids
+19 -4
View File
@@ -5,7 +5,7 @@ from collections.abc import Generator, Mapping, Sequence
from datetime import datetime
from enum import StrEnum
from typing import TYPE_CHECKING, Any, Optional, TypedDict, cast
from uuid import uuid4
from uuid import NAMESPACE_URL, UUID, uuid4, uuid5
import sqlalchemy as sa
from sqlalchemy import (
@@ -1486,15 +1486,30 @@ class ConversationVariable(TypeBase):
DateTime, nullable=False, server_default=func.current_timestamp(), onupdate=func.current_timestamp(), init=False
)
@classmethod
def storage_id(cls, variable: VariableBase) -> str:
"""UUID primary key for ``variable``.
Draft and DSL ids such as ``opt-comp-prompt-var`` are not UUIDs and cannot
be inserted. Those become a uuid5 of the variable name, so the same variable
keeps one row. An id that is already a UUID is stored unchanged. Callers
still see the author id on the variable payload.
"""
row_id = variable.id
try:
UUID(str(row_id))
except (ValueError, TypeError, AttributeError):
return str(uuid5(NAMESPACE_URL, f"dify:conversation-variable:{variable.name}"))
return str(row_id)
@classmethod
def from_variable(cls, *, app_id: str, conversation_id: str, variable: VariableBase) -> "ConversationVariable":
obj = cls(
id=variable.id,
return cls(
id=cls.storage_id(variable),
app_id=app_id,
conversation_id=conversation_id,
data=variable.model_dump_json(),
)
return obj
def to_variable(self) -> VariableBase:
mapping = json.loads(self.data)
+1
View File
@@ -25238,6 +25238,7 @@ Enabled routes require at least two exits; drafts may omit conditions.
| ---- | ---- | ----------- | -------- |
| created_at | integer | | Yes |
| result | string | | Yes |
| warning | string | Advisory warning for variable references that can read a skipped branch. Publish still succeeds. | No |
#### WorkflowResponse
+19 -1
View File
@@ -89,6 +89,22 @@ IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes
CURRENT_DSL_VERSION = CURRENT_APP_DSL_VERSION
def missing_app_section_error(top_level_keys: list[str]) -> str:
"""Explain a YAML that has no top-level ``app`` mapping.
The found keys are the caller's actual document, so a sketch of nodes is
not reported as a blank import failure.
"""
found = ", ".join(key for key in top_level_keys if key != "app")
if len(found) > 80:
found = found[:80].rstrip(", ") + "…"
return (
"Missing app data in YAML content. "
"Not a valid Dify app DSL: the top-level 'app' section is required "
f"(found: {found or 'none'})."
)
class PendingData(PendingImportOwner):
import_mode: str
yaml_content: str
@@ -208,6 +224,8 @@ class AppDslService:
error="Invalid YAML format: content must be a mapping",
)
original_top_level_keys = [key for key in data if isinstance(key, str)]
# Validate and fix DSL version
if not data.get("version"):
data["version"] = "0.1.0"
@@ -226,7 +244,7 @@ class AppDslService:
return Import(
id=import_id,
status=ImportStatus.FAILED,
error="Missing app data in YAML content",
error=missing_app_section_error(original_top_level_keys),
)
if package is not None and package.has_resources:
+2
View File
@@ -295,6 +295,7 @@ class ConversationService:
"created_at": row.created_at,
"updated_at": row.updated_at,
**row.to_variable().model_dump(),
"id": row.id,
}
for row in rows
]
@@ -379,4 +380,5 @@ class ConversationService:
"created_at": existing_variable.created_at,
"updated_at": naive_utc_now(), # Update timestamp
**updated_variable.model_dump(),
"id": existing_variable.id,
}
@@ -15,7 +15,8 @@ class ConversationVariableUpdater:
def update(self, conversation_id: str, variable: VariableBase) -> None:
stmt = select(ConversationVariable).where(
ConversationVariable.id == variable.id, ConversationVariable.conversation_id == conversation_id
ConversationVariable.id == ConversationVariable.storage_id(variable),
ConversationVariable.conversation_id == conversation_id,
)
with self._session_maker() as session:
row = session.scalar(stmt)
@@ -0,0 +1,261 @@
from __future__ import annotations
import re
from collections import defaultdict, deque
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from typing import Any
from core.trigger.constants import TRIGGER_NODE_TYPES
from core.workflow.nodes.human_input.constants import TIMEOUT_HANDLE
from graphon.enums import BuiltinNodeTypes, ErrorStrategy
_RESERVED_SELECTOR_HEADS: frozenset[str] = frozenset({"sys", "env", "conversation", "start"})
_REFERENCE_EXEMPT_NODE_TYPES: frozenset[str] = frozenset(
{
BuiltinNodeTypes.VARIABLE_AGGREGATOR,
BuiltinNodeTypes.LEGACY_VARIABLE_AGGREGATOR,
}
)
_BRANCH_NODE_TYPES: frozenset[str] = frozenset(
{
BuiltinNodeTypes.IF_ELSE,
BuiltinNodeTypes.QUESTION_CLASSIFIER,
BuiltinNodeTypes.HUMAN_INPUT,
}
)
_TEMPLATE_REFERENCE_PATTERN = re.compile(
r"\{\{#([a-zA-Z0-9_]{1,50})\.[a-zA-Z_][a-zA-Z0-9_]{0,29}(?:\.[a-zA-Z_][a-zA-Z0-9_]{0,29}){0,9}#\}\}"
)
_MAX_REPORTED_ISSUES = 10
@dataclass(frozen=True)
class VariableReferenceIssue:
"""A single unsafe variable reference: ``node`` reads output from ``referenced_node``."""
node_id: str
node_title: str
referenced_node_id: str
referenced_node_title: str
def validate_variable_references(graph: Mapping[str, Any]) -> list[VariableReferenceIssue]:
"""Flag references where the consumer can run without the producer.
Uses engine execution semantics: nodes fire on any active inbound edge; branch
nodes (if-else / question-classifier / human-input / fail-branch) take one handle
(possibly unwired); other nodes fan out all outbound edges. Returns [] when clean.
"""
nodes = graph.get("nodes") or []
edges = graph.get("edges") or []
if not isinstance(nodes, list) or not isinstance(edges, list) or not nodes:
return []
node_type: dict[str, str] = {}
node_title: dict[str, str] = {}
node_parent: dict[str, str | None] = {}
node_data: dict[str, Mapping[str, Any]] = {}
for node in nodes:
if not isinstance(node, Mapping):
continue
node_id = node.get("id")
if not isinstance(node_id, str):
continue
raw_data = node.get("data")
data: Mapping[str, Any] = raw_data if isinstance(raw_data, Mapping) else {}
node_type[node_id] = str(data.get("type") or "")
title = data.get("title")
node_title[node_id] = str(title) if title else node_id
parent = node.get("parentId")
node_parent[node_id] = parent if isinstance(parent, str) else None
node_data[node_id] = data
node_ids = set(node_type)
out_targets_by_handle: dict[str, dict[str | None, list[str]]] = defaultdict(lambda: defaultdict(list))
predecessors: dict[str, list[str]] = defaultdict(list)
successors: dict[str, list[str]] = defaultdict(list)
for edge in edges:
if not isinstance(edge, Mapping):
continue
source = edge.get("source")
target = edge.get("target")
if not isinstance(source, str) or not isinstance(target, str):
continue
out_targets_by_handle[source][edge.get("sourceHandle")].append(target)
predecessors[target].append(source)
successors[source].append(target)
entries = [
nid
for nid in node_ids
if node_parent[nid] is None and node_type[nid] in (BuiltinNodeTypes.START, *TRIGGER_NODE_TYPES)
]
reachable = _reachable_from(entries, successors)
exclusive = {
nid
for nid in node_ids
if node_type[nid] in _BRANCH_NODE_TYPES or node_data[nid].get("error_strategy") == ErrorStrategy.FAIL_BRANCH
}
# A selectable handle can be unwired. It still lets the branch skip a producer.
for nid in exclusive:
data = node_data[nid]
handles: list[str] = []
if node_type[nid] == BuiltinNodeTypes.IF_ELSE:
cases = data.get("cases")
handles = [case["case_id"] for case in cases] if isinstance(cases, list) else ["true"]
handles.append("false")
elif node_type[nid] == BuiltinNodeTypes.QUESTION_CLASSIFIER:
handles = [item["id"] for item in data.get("classes", [])]
elif node_type[nid] == BuiltinNodeTypes.HUMAN_INPUT:
handles = [action["id"] for action in data.get("user_actions", [])]
handles.append(TIMEOUT_HANDLE)
if data.get("error_strategy") == ErrorStrategy.FAIL_BRANCH:
handles.extend(["source", "fail-branch"])
for handle in handles:
out_targets_by_handle[nid].setdefault(handle, [])
consumers_by_producer: dict[str, set[str]] = defaultdict(set)
for node_id in node_ids:
if node_parent[node_id] is not None or node_id not in reachable:
continue
if node_type[node_id] in _REFERENCE_EXEMPT_NODE_TYPES:
continue
for referenced_id in _referenced_node_ids(node_data[node_id]):
if referenced_id == node_id or referenced_id not in node_ids:
continue
if node_parent[referenced_id] is not None:
continue
consumers_by_producer[referenced_id].add(node_id)
issues: list[VariableReferenceIssue] = []
for producer, consumers in sorted(consumers_by_producer.items()):
runnable_without_producer = _nodes_runnable_without(
producer,
entries=entries,
predecessors=predecessors,
successors=successors,
out_targets_by_handle=out_targets_by_handle,
exclusive=exclusive,
)
for consumer in sorted(consumers):
if consumer in runnable_without_producer:
issues.append(
VariableReferenceIssue(
node_id=consumer,
node_title=node_title[consumer],
referenced_node_id=producer,
referenced_node_title=node_title[producer],
)
)
return issues
def format_variable_reference_errors(issues: Sequence[VariableReferenceIssue]) -> str:
"""Build a single-line warning listing unsafe reader ← producer pairs."""
count = len(issues)
shown = issues[:_MAX_REPORTED_ISSUES]
pairs = "; ".join(f'"{issue.node_title}" ← "{issue.referenced_node_title}"' for issue in shown)
if count > len(shown):
pairs += f"; +{count - len(shown)} more"
return (
f"{count} variable reference{'s' if count != 1 else ''} may read a skipped branch "
f"output. Use a Variable Aggregator or a default value — {pairs}."
)
def _reachable_from(starts: Sequence[str], successors: Mapping[str, list[str]]) -> set[str]:
visited: set[str] = set()
queue: deque[str] = deque(starts)
while queue:
current = queue.popleft()
if current in visited:
continue
visited.add(current)
for nxt in successors.get(current, ()):
if nxt not in visited:
queue.append(nxt)
return visited
def _nodes_runnable_without(
producer: str,
*,
entries: Sequence[str],
predecessors: Mapping[str, list[str]],
successors: Mapping[str, list[str]],
out_targets_by_handle: Mapping[str, Mapping[str | None, list[str]]],
exclusive: set[str],
) -> set[str]:
"""Return nodes that can execute in some run where ``producer`` does not."""
forbidden = {producer}
queue: deque[str] = deque([producer])
while queue:
node = queue.popleft()
for pred in predecessors.get(node, ()):
if pred in forbidden:
continue
if pred not in exclusive or (
out_targets_by_handle[pred]
and all(
any(target in forbidden for target in targets) for targets in out_targets_by_handle[pred].values()
)
):
forbidden.add(pred)
queue.append(pred)
runnable: set[str] = set()
work: deque[str] = deque(entry for entry in entries if entry not in forbidden)
while work:
node = work.popleft()
if node in runnable:
continue
runnable.add(node)
# All edges sharing the selected handle activate together. If one would
# execute the producer, none of that handle's siblings can avoid it.
targets = (
[
target
for branch_targets in out_targets_by_handle[node].values()
if not any(target in forbidden for target in branch_targets)
for target in branch_targets
]
if node in exclusive
else successors.get(node, ())
)
for nxt in targets:
if nxt not in forbidden and nxt not in runnable:
work.append(nxt)
return runnable
def _referenced_node_ids(value: Any) -> Iterator[str]:
"""Yield node ids referenced via selectors and {{#node.field#}} templates."""
if isinstance(value, Mapping):
variable = value.get("value")
if value.get("type") == "variable" and isinstance(variable, list) and variable:
yield from _reference_head(variable[0])
for key, child in value.items():
if _is_selector_key(key) and isinstance(child, list) and child:
yield from _reference_head(child[0])
yield from _referenced_node_ids(child)
elif isinstance(value, list):
for item in value:
yield from _referenced_node_ids(item)
elif isinstance(value, str):
for match in _TEMPLATE_REFERENCE_PATTERN.finditer(value):
yield from _reference_head(match.group(1))
def _is_selector_key(key: object) -> bool:
return isinstance(key, str) and (key == "selector" or key.endswith("_selector"))
def _reference_head(head: object) -> Iterator[str]:
if isinstance(head, str) and head not in _RESERVED_SELECTOR_HEADS:
yield head
@@ -3,6 +3,7 @@ from __future__ import annotations
from datetime import UTC, datetime
from inspect import unwrap
from unittest.mock import PropertyMock, patch
from uuid import UUID
import pytest
from flask import Flask
@@ -67,7 +68,9 @@ def test_get_conversation_variables_returns_paginated_response(
assert response["limit"] == 100
assert response["total"] == 1
assert response["has_more"] is False
assert response["data"][0]["id"] == "var-1"
assert response["data"][0]["id"] == row.id
UUID(response["data"][0]["id"])
assert row.to_variable().id == "var-1"
assert response["data"][0]["created_at"] == expected_created_at
assert response["data"][0]["updated_at"] == expected_updated_at
@@ -134,13 +134,43 @@ def _make_workflow(**overrides) -> Workflow:
return workflow
@pytest.mark.parametrize(
"advisory", ["clean", "warning", "checker-error", "formatter-error", "empty", "non-object", "invalid-json"]
)
def test_publish_workflow_returns_success(
app: Flask,
monkeypatch: pytest.MonkeyPatch,
advisory: str,
) -> None:
current_user = SimpleNamespace(id="account-1")
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1")
workflow = SimpleNamespace(id="published-workflow", created_at=datetime(2026, 8, 17, 12, 0, 0))
graph = {
"nodes": [
{"id": "start", "data": {"type": "start"}},
{"id": "branch", "data": {"type": "if-else"}},
{"id": "producer", "data": {"type": "code", "title": "Producer"}},
{
"id": "consumer",
"data": {"type": "answer", "title": "Consumer", "answer": "{{#producer.text#}}"},
},
],
"edges": [
{"source": "start", "target": "branch"},
{"source": "branch", "target": "producer", "sourceHandle": "true"},
{"source": "branch", "target": "consumer", "sourceHandle": "false"},
],
}
workflow = SimpleNamespace(
id="published-workflow",
created_at=datetime(2026, 8, 17, 12, 0, 0),
graph={"clean": "{}", "empty": None, "non-object": "[]", "invalid-json": "{"}.get(advisory, json.dumps(graph)),
)
if advisory == "checker-error":
monkeypatch.setattr(workflow_module, "validate_variable_references", Mock(side_effect=RuntimeError("checker")))
elif advisory == "formatter-error":
monkeypatch.setattr(
workflow_module, "format_variable_reference_errors", Mock(side_effect=RuntimeError("format"))
)
session = Mock()
session.get.return_value = app_model
monkeypatch.setattr(
@@ -163,6 +193,14 @@ def test_publish_workflow_returns_success(
)
assert response["result"] == "success"
assert app_model.workflow_id == workflow.id
assert isinstance(response["created_at"], int)
if advisory == "warning":
assert "Consumer" in response["warning"]
assert "Producer" in response["warning"]
assert "skipped branch" in response["warning"]
else:
assert "warning" not in response
@pytest.mark.parametrize("transaction_fails", [False, True], ids=["commit-succeeds", "commit-fails"])
@@ -128,3 +128,24 @@ def test_all_variables_exist_no_changes(monkeypatch: pytest.MonkeyPatch, sqlite_
assert [variable.id for variable in variables] == [VAR_1_ID, VAR_2_ID]
persisted = sqlite_session.scalars(select(ConversationVariable)).all()
assert len(persisted) == 2
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
def test_non_uuid_conversation_variable_is_created_once(
monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
) -> None:
author_id = "opt-comp-prompt-var"
variable = _variable(author_id, "optimization_comparison_prompt", "-")
_bind_runner_sessions(monkeypatch, sqlite_session)
runner = _runner([variable])
first = runner._initialize_conversation_variables()
second = runner._initialize_conversation_variables()
expected_row_id = ConversationVariable.storage_id(variable)
assert variable.id == author_id
assert [item.id for item in first] == [author_id]
assert [item.id for item in second] == [author_id]
persisted = sqlite_session.scalars(select(ConversationVariable)).all()
assert [row.id for row in persisted] == [expected_row_id]
assert persisted[0].to_variable().id == author_id
@@ -1,10 +1,29 @@
from uuid import uuid4
from uuid import NAMESPACE_URL, UUID, uuid4, uuid5
from factories import variable_factory
from graphon.variables import SegmentType
from models import ConversationVariable
def test_from_variable_coerces_non_uuid_primary_key():
variable = variable_factory.build_conversation_variable_from_mapping(
{
"id": "opt-comp-prompt-var",
"name": "optimization_comparison_prompt",
"value_type": SegmentType.STRING,
"value": "-",
}
)
row = ConversationVariable.from_variable(app_id="app_id", conversation_id="conversation_id", variable=variable)
expected = str(uuid5(NAMESPACE_URL, "dify:conversation-variable:optimization_comparison_prompt"))
UUID(row.id)
assert row.id == expected
assert variable.id == "opt-comp-prompt-var"
assert row.to_variable().id == "opt-comp-prompt-var"
def test_from_variable_and_to_variable():
variable = variable_factory.build_conversation_variable_from_mapping(
{
@@ -852,3 +852,25 @@ def test_overwrite_rejects_incompatible_nodes_before_mutation(mode: AppMode, nod
assert "incompatible" in result.error
assert target.name == "Original"
session.add.assert_not_called()
@pytest.mark.parametrize(
("content", "found"),
[
({"meta": {}, "nodes": [], "edges": []}, "meta, nodes, edges"),
({"app": None}, "none"),
({"x" * 81: {}}, "x" * 80 + "…"),
],
ids=["original-keys", "empty-app", "bounded-key-list"],
)
def test_missing_app_section_names_the_keys_that_were_present(
unbound_session: Session, content: dict[str, object], found: str
) -> None:
result = AppDslService(session=unbound_session).import_app(
account=_account(), import_mode="yaml-content", yaml_content=yaml.safe_dump(content, sort_keys=False)
)
assert result.status == ImportStatus.FAILED
assert result.error is not None
assert result.error.startswith("Missing app data in YAML content.")
assert result.error.endswith(f"(found: {found}).")
assert not unbound_session.in_transaction()
@@ -9,8 +9,10 @@ in-memory SQLite sessions with persisted ORM rows.
import json
from dataclasses import replace
from datetime import timedelta
from decimal import Decimal
from unittest.mock import MagicMock
from uuid import UUID
import pytest
from sqlalchemy import asc, desc, event
@@ -19,6 +21,7 @@ from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from core.credit_usage import CreditUsageAppType
from core.model_context import get_credit_usage_metadata, use_credit_usage_metadata
from factories import variable_factory
from libs.datetime_utils import naive_utc_now
from models import Account, ConversationVariable
from models.agent import (
@@ -591,6 +594,46 @@ class TestConversationServiceHelpers:
class TestConversationServiceConversationalVariable:
"""Test conversational variable operations."""
@pytest.mark.parametrize("sqlite_session", [(Conversation, ConversationVariable)], indirect=True)
def test_imported_variable_id_can_be_used_for_pagination_and_update(self, sqlite_session: Session):
app_model = ConversationServiceTestDataFactory.create_app()
user = ConversationServiceTestDataFactory.create_account()
conversation = ConversationServiceTestDataFactory.create_conversation()
variable = variable_factory.build_conversation_variable_from_mapping(
{"id": "imported-topic", "name": "topic", "value_type": "string", "value": "original"}
)
imported_row = ConversationVariable.from_variable(
app_id=APP_ID, conversation_id=CONVERSATION_ID, variable=variable
)
imported_row.created_at = naive_utc_now()
next_row = _conversation_variable(variable_id=OTHER_VARIABLE_ID, name="next", value="next value")
next_row.created_at = imported_row.created_at + timedelta(seconds=1)
sqlite_session.add_all([conversation, imported_row, next_row])
sqlite_session.commit()
first_page = ConversationService.get_conversational_variable(
app_model, CONVERSATION_ID, user, limit=1, last_id=None, session=sqlite_session
)
returned_id = first_page.data[0]["id"]
assert returned_id == imported_row.id
UUID(returned_id)
assert first_page.has_more is True
second_page = ConversationService.get_conversational_variable(
app_model, CONVERSATION_ID, user, limit=1, last_id=returned_id, session=sqlite_session
)
assert [item["id"] for item in second_page.data] == [OTHER_VARIABLE_ID]
assert second_page.has_more is False
updated = ConversationService.update_conversation_variable(
app_model, CONVERSATION_ID, returned_id, user, "updated", session=sqlite_session
)
assert updated["id"] == returned_id
assert updated["value"] == "updated"
sqlite_session.refresh(imported_row)
assert imported_row.to_variable().id == "imported-topic"
assert imported_row.to_variable().value == "updated"
@pytest.mark.parametrize("sqlite_session", [(Conversation, ConversationVariable)], indirect=True)
def test_get_conversational_variable_with_name_filter_mysql(
self,
@@ -0,0 +1,34 @@
import pytest
from sqlalchemy.orm import Session, sessionmaker
from graphon.variables import StringVariable
from models import ConversationVariable
from services.conversation_variable_updater import ConversationVariableUpdater
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
@pytest.mark.parametrize("variable_id", ["imported-topic", "33333333-3333-3333-3333-333333333333"])
def test_runtime_update_preserves_authored_id_and_other_conversations(
sqlite_session: Session, variable_id: str
) -> None:
variable = StringVariable(id=variable_id, name="topic", value="original", selector=["conversation", "topic"])
conversation_id = "22222222-2222-2222-2222-222222222222"
other_conversation_id = "22222222-2222-2222-2222-222222222223"
rows = [
ConversationVariable.from_variable(
app_id="11111111-1111-1111-1111-111111111111",
conversation_id=owner,
variable=variable,
)
for owner in (conversation_id, other_conversation_id)
]
sqlite_session.add_all(rows)
sqlite_session.commit()
updater = ConversationVariableUpdater(sessionmaker(bind=sqlite_session.get_bind()))
updater.update(conversation_id, variable.model_copy(update={"value": "updated"}))
sqlite_session.expire_all()
assert rows[0].to_variable().id == variable_id
assert rows[0].to_variable().value == "updated"
assert rows[1].to_variable().value == "original"
@@ -0,0 +1,480 @@
"""Unit tests for the preflight workflow variable-reference validator.
See GitHub issue #34358. A node that can read a producer on a skipped branch
is flagged. A producer that always runs, a parallel branch, and an aggregator
are not.
"""
from services.workflow_variable_reference_validator import (
format_variable_reference_errors,
validate_variable_references,
)
def _node(node_id: str, node_type: str, title: str = "", *, selector: list[str] | None = None) -> dict[str, object]:
data: dict[str, object] = {"type": node_type, "title": title or node_id}
if selector is not None:
data["variables"] = [{"value_selector": selector, "variable": "x"}]
return {"id": node_id, "data": data}
def _edge(source: str, target: str, handle: str = "source") -> dict[str, str]:
return {"source": source, "target": target, "sourceHandle": handle}
def _node_data(node: dict[str, object]) -> dict[str, object]:
raw = node["data"]
if not isinstance(raw, dict):
raise AssertionError("node data must be a mapping")
return raw
class TestValidateVariableReferences:
def test_same_selected_handle_runs_producer_and_consumer_together(self) -> None:
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else"),
_node("prod", "tool"),
_node("cons", "llm", selector=["prod", "text"]),
_node("end", "end"),
],
"edges": [
_edge("start", "router"),
_edge("router", "prod", "true"),
_edge("router", "cons", "true"),
_edge("router", "end", "false"),
],
}
assert validate_variable_references(graph) == []
def test_unwired_else_can_skip_all_wired_cases(self) -> None:
router = _node("router", "if-else")
_node_data(router)["cases"] = [{"case_id": "true"}, {"case_id": "elif"}]
graph = {
"nodes": [
_node("start", "start"),
router,
_node("prod", "tool"),
_node("cons", "llm", selector=["prod", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("start", "cons"),
_edge("router", "prod", "true"),
_edge("router", "prod", "elif"),
],
}
assert [(issue.node_id, issue.referenced_node_id) for issue in validate_variable_references(graph)] == [
("cons", "prod")
]
def test_trigger_can_run_without_the_manual_start_path(self) -> None:
graph = {
"nodes": [
_node("start", "start"),
_node("trigger", "trigger-webhook"),
_node("prod", "tool"),
_node("cons", "llm", selector=["prod", "text"]),
],
"edges": [_edge("start", "prod"), _edge("trigger", "cons")],
}
assert [(issue.node_id, issue.referenced_node_id) for issue in validate_variable_references(graph)] == [
("cons", "prod")
]
def test_disconnected_consumer_is_not_an_entry(self) -> None:
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else"),
_node("prod", "tool"),
_node("cons", "llm", selector=["prod", "text"]),
],
"edges": [_edge("start", "router"), _edge("router", "prod", "true")],
}
assert validate_variable_references(graph) == []
def test_empty_graph_returns_no_issues(self) -> None:
"""An empty or node-less graph is accepted without issues."""
assert validate_variable_references({}) == []
assert validate_variable_references({"nodes": [], "edges": []}) == []
def test_upstream_always_run_reference_is_safe(self) -> None:
"""Reading a node that runs before the branch split is safe."""
graph = {
"nodes": [
_node("start", "start"),
_node("norm", "tool", "Normalizer"),
_node("router", "if-else", "Router"),
_node("b", "llm", "B", selector=["norm", "text"]),
],
"edges": [
_edge("start", "norm"),
_edge("norm", "router"),
_edge("router", "b", "true"),
],
}
assert validate_variable_references(graph) == []
def test_cross_branch_reference_is_flagged(self) -> None:
"""Reading a producer that sits on the opposite branch is flagged."""
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "tool", "A"),
_node("b", "llm", "B", selector=["a", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "a", "true"),
_edge("router", "b", "false"),
],
}
issues = validate_variable_references(graph)
assert len(issues) == 1
assert issues[0].node_id == "b"
assert issues[0].referenced_node_id == "a"
def test_join_after_conditional_skip_is_flagged(self) -> None:
"""A join reading a node skipped on the taken branch is flagged."""
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "llm", "A"),
_node("join", "llm", "Join", selector=["a", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "join", "true"),
_edge("router", "a", "false"),
_edge("a", "join"),
],
}
issues = validate_variable_references(graph)
assert len(issues) == 1
assert issues[0].node_id == "join"
assert issues[0].referenced_node_id == "a"
def test_single_wired_if_else_branch_is_flagged(self) -> None:
"""An if-else with only one branch wired skips it when the other is taken."""
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("prod", "tool", "Producer"),
_node("cons", "llm", "Consumer", selector=["prod", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "prod", "true"),
_edge("start", "cons"),
],
}
issues = validate_variable_references(graph)
assert len(issues) == 1
assert issues[0].referenced_node_id == "prod"
def test_fail_branch_producer_is_flagged(self) -> None:
"""A producer behind a fail-branch node's success handle is skipped on failure."""
risky = _node("risky", "tool", "Risky")
_node_data(risky)["error_strategy"] = "fail-branch"
graph = {
"nodes": [
_node("start", "start"),
risky,
_node("prod", "tool", "Producer"),
_node("cons", "llm", "Consumer", selector=["prod", "text"]),
],
"edges": [
_edge("start", "risky"),
_edge("risky", "prod", "success"),
_edge("start", "cons"),
],
}
issues = validate_variable_references(graph)
assert len(issues) == 1
assert issues[0].referenced_node_id == "prod"
def test_branches_converging_into_producer_are_not_flagged(self) -> None:
"""When every wired branch leads into the producer it runs on every path."""
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("prod", "tool", "Producer"),
_node("cons", "llm", "Consumer", selector=["prod", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "prod", "true"),
_edge("router", "prod", "false"),
_edge("start", "cons"),
],
}
assert validate_variable_references(graph) == []
def test_template_reference_is_detected(self) -> None:
"""A reference expressed as a {{#node.field#}} template is detected."""
b = _node("b", "llm", "B")
_node_data(b)["prompt"] = "Use {{#a.text#}} here"
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "tool", "A"),
b,
],
"edges": [
_edge("start", "router"),
_edge("router", "a", "true"),
_edge("router", "b", "false"),
],
}
issues = validate_variable_references(graph)
assert len(issues) == 1
assert issues[0].referenced_node_id == "a"
def test_variable_typed_parameter_reference_is_detected(self) -> None:
"""A {type: variable} parameter selector is detected."""
consumer = _node("c", "tool", "Consumer")
_node_data(consumer)["parameters"] = [{"type": "variable", "value": ["a", "text"]}]
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "tool", "A"),
consumer,
],
"edges": [
_edge("start", "router"),
_edge("router", "a", "true"),
_edge("router", "c", "false"),
],
}
issues = validate_variable_references(graph)
assert len(issues) == 1
assert issues[0].referenced_node_id == "a"
def test_bare_hash_text_is_not_a_placeholder(self) -> None:
"""Bare #a.text# text without braces is not treated as a reference."""
b = _node("b", "llm", "B")
_node_data(b)["prompt"] = "see #a.text# (not a placeholder)"
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "tool", "A"),
b,
],
"edges": [
_edge("start", "router"),
_edge("router", "a", "true"),
_edge("router", "b", "false"),
],
}
assert validate_variable_references(graph) == []
def test_parallel_branches_are_not_flagged(self) -> None:
"""Parallel branches both run, so a join reading one of them is safe."""
graph = {
"nodes": [
_node("start", "start"),
_node("split", "tool", "Split"),
_node("a", "tool", "A"),
_node("b", "tool", "B"),
_node("join", "llm", "Join", selector=["a", "text"]),
],
"edges": [
_edge("start", "split"),
_edge("split", "a"),
_edge("split", "b"),
_edge("a", "join"),
_edge("b", "join"),
],
}
assert validate_variable_references(graph) == []
def test_producer_on_an_independent_always_on_path_is_not_flagged(self) -> None:
"""A producer that also runs on an unconditional path is never skipped."""
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("p", "tool", "Producer"),
_node("c", "llm", "Consumer", selector=["p", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "p", "true"),
_edge("router", "c", "false"),
_edge("p", "c"),
_edge("start", "p"),
],
}
assert validate_variable_references(graph) == []
def test_fan_out_forced_producer_is_not_flagged(self) -> None:
"""A fan-out that forces the producer in on every run is safe."""
graph = {
"nodes": [
_node("start", "start"),
_node("m", "tool", "Fan-out"),
_node("p", "tool", "Producer"),
_node("c", "llm", "Consumer", selector=["p", "text"]),
],
"edges": [
_edge("start", "m"),
_edge("m", "p"),
_edge("m", "c"),
],
}
assert validate_variable_references(graph) == []
def test_variable_aggregator_reference_is_exempt(self) -> None:
"""A Variable Aggregator merging branch outputs is exempt."""
agg = _node("agg", "variable-aggregator", "Aggregator")
_node_data(agg)["variables"] = [["a", "text"], ["b", "text"]]
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "tool", "A"),
_node("b", "tool", "B"),
agg,
],
"edges": [
_edge("start", "router"),
_edge("router", "a", "true"),
_edge("router", "b", "false"),
_edge("a", "agg"),
_edge("b", "agg"),
],
}
assert validate_variable_references(graph) == []
def test_reserved_selector_heads_are_ignored(self) -> None:
"""sys/conversation/env selector heads are always available and never flagged."""
b = _node("b", "llm", "B", selector=["sys", "query"])
_node_data(b)["prompt"] = "{{#conversation.foo#}} {{#env.bar#}}"
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
b,
],
"edges": [_edge("start", "router"), _edge("router", "b", "true")],
}
assert validate_variable_references(graph) == []
def test_question_classifier_unwired_class_can_skip_the_producer(self) -> None:
"""A declared class with no edge is still an exit. The consumer is reached from start."""
router = _node("router", "question-classifier", "Router")
_node_data(router)["classes"] = [{"id": "billing"}, {"id": "shipping"}]
graph = {
"nodes": [
_node("start", "start"),
router,
_node("prod", "tool", "Producer"),
_node("cons", "llm", "Consumer", selector=["prod", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("start", "cons"),
_edge("router", "prod", "billing"),
],
}
assert [(issue.node_id, issue.referenced_node_id) for issue in validate_variable_references(graph)] == [
("cons", "prod")
]
def test_human_input_unwired_timeout_can_skip_the_wired_action(self) -> None:
"""The implicit timeout handle is an exit even when no edge uses it."""
review = _node("review", "human-input", "Review")
_node_data(review)["user_actions"] = [{"id": "approve"}]
graph = {
"nodes": [
_node("start", "start"),
review,
_node("prod", "tool", "Producer"),
_node("cons", "llm", "Consumer", selector=["prod", "text"]),
],
"edges": [
_edge("start", "review"),
_edge("start", "cons"),
_edge("review", "prod", "approve"),
],
}
assert [(issue.node_id, issue.referenced_node_id) for issue in validate_variable_references(graph)] == [
("cons", "prod")
]
def test_nested_parent_nodes_are_not_root_references(self) -> None:
"""Nodes inside a nested graph are neither consumers nor producers of the root check."""
nested_prod = _node("nested-prod", "tool", "Nested producer")
nested_prod["parentId"] = "loop-1"
nested_cons = _node("nested-cons", "llm", "Nested consumer", selector=["prod", "text"])
nested_cons["parentId"] = "loop-1"
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("prod", "tool", "Producer"),
nested_prod,
nested_cons,
_node("root-cons", "llm", "Root consumer", selector=["nested-prod", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "prod", "true"),
_edge("router", "nested-prod", "true"),
_edge("router", "nested-cons", "false"),
_edge("router", "root-cons", "false"),
],
}
assert validate_variable_references(graph) == []
class TestFormatVariableReferenceErrors:
def test_message_includes_titles_and_count(self) -> None:
"""The rendered message names both nodes and reports the issue count."""
graph = {
"nodes": [
_node("start", "start"),
_node("router", "if-else", "Router"),
_node("a", "tool", "Producer"),
_node("b", "llm", "Consumer", selector=["a", "text"]),
],
"edges": [
_edge("start", "router"),
_edge("router", "a", "true"),
_edge("router", "b", "false"),
],
}
issues = validate_variable_references(graph)
message = format_variable_reference_errors(issues)
assert "Consumer" in message
assert "Producer" in message
assert "1 variable reference " in message
def test_message_truncates_after_ten_issues(self) -> None:
"""Only the first ten pairs are listed; the rest are a count."""
nodes: list[dict[str, object]] = [_node("start", "start"), _node("router", "if-else", "Router")]
edges = [_edge("start", "router")]
for index in range(11):
producer_id = f"p{index:02d}"
consumer_id = f"c{index:02d}"
nodes.append(_node(producer_id, "tool", f"Producer {index:02d}"))
nodes.append(_node(consumer_id, "llm", f"Consumer {index:02d}", selector=[producer_id, "text"]))
edges.append(_edge("router", producer_id, "true"))
edges.append(_edge("router", consumer_id, "false"))
issues = validate_variable_references({"nodes": nodes, "edges": edges})
assert len(issues) == 11
message = format_variable_reference_errors(issues)
assert "11 variable references " in message
assert message.count("←") == 10
assert "+1 more" in message
assert "Consumer 10" not in message
assert "Producer 10" not in message
@@ -1156,6 +1156,7 @@ export type PublishWorkflowPayload = {
export type WorkflowPublishResponse = {
created_at: number
result: string
warning?: string | null
}
export type WebhookTriggerResponse = {
@@ -681,6 +681,7 @@ export const zPublishWorkflowPayload = z.object({
export const zWorkflowPublishResponse = z.object({
created_at: z.int(),
result: z.string(),
warning: z.string().nullish(),
})
/**
@@ -267,6 +267,7 @@ export type PublishWorkflowPayload = {
export type WorkflowPublishResponse = {
created_at: number
result: string
warning?: string | null
}
export type WorkflowUpdatePayload = {
@@ -115,6 +115,7 @@ export const zPublishWorkflowPayload = z.object({
export const zWorkflowPublishResponse = z.object({
created_at: z.int(),
result: z.string(),
warning: z.string().nullish(),
})
/**
@@ -1002,6 +1002,47 @@ describe('CreateFromDSLModal', () => {
)
})
it.each([
[
'JSON',
() => Response.json({ error: 'Import confirmation expired' }, { status: 400 }),
'Import confirmation expired',
],
['empty', () => new Response(null, { status: 500 }), undefined],
] as const)(
'reports a %s HTTP failure when confirming an import',
async (_, response, description) => {
const user = userEvent.setup()
mockImportDSL.mockResolvedValue({
id: 'pending-import',
status: DSLImportStatus.PENDING,
imported_dsl_version: '1.0.0',
current_dsl_version: '2.0.0',
})
mockImportDSLConfirm.mockRejectedValueOnce(response())
const onClose = vi.fn()
render(
<CreateFromDSLModal
show
onClose={onClose}
activeTab={CreateFromDSLModalTab.FROM_URL}
dslUrl="https://example.com/app.yml"
/>,
)
await user.click(getCreateButton())
await user.click(await screen.findByRole('button', { name: /newApp\.Confirm/ }))
await waitFor(() => {
expect(toastMocks.error).toHaveBeenCalledExactlyOnceWith(
expect.stringMatching(/newApp\.appCreateFailed/),
{ description },
)
})
expect(onClose).not.toHaveBeenCalled()
},
)
it('should handle pending import confirmation failures and cancellation', async () => {
mockImportDSL.mockResolvedValue({
id: 'import-4',
@@ -261,7 +261,9 @@ function CreateFromDSLModal({
} catch (error) {
toast.error(
t(($) => $['newApp.appCreateFailed'], { ns: 'app' }),
{ description: await getAppTransferErrorMessage(error) },
{
description: await getAppTransferErrorMessage(error),
},
)
return
}
@@ -268,16 +268,24 @@ it('keeps the pending import confirmation and tracks only after successful confi
})
})
it('reports import failure without tracking or navigating', async () => {
const user = userEvent.setup()
importResponse = { id: 'import-1', status: 'failed' }
setup({ cloud: true })
await openCreate(user, true)
await submit(user)
await waitFor(() => expect(toast.error).toHaveBeenCalledWith('app.newApp.appCreateFailed'))
expect(redirect).not.toHaveBeenCalled()
expect(trackCreateApp).not.toHaveBeenCalled()
})
it.each([undefined, 'Missing app data in YAML content'])(
'reports import failure (%s) once without tracking or navigating',
async (error) => {
const user = userEvent.setup()
importResponse = { id: 'import-1', status: 'failed', error }
setup({ cloud: true })
await openCreate(user, true)
await submit(user)
await waitFor(() =>
expect(toast.error).toHaveBeenCalledWith('app.newApp.appCreateFailed', {
description: error,
}),
)
expect(toast.error).toHaveBeenCalledOnce()
expect(redirect).not.toHaveBeenCalled()
expect(trackCreateApp).not.toHaveBeenCalled()
},
)
it('keeps nullable template metadata editable through the actual modal', async () => {
const user = userEvent.setup()
@@ -585,6 +585,26 @@ describe('FeaturesTrigger', () => {
expect(mockPublishWorkflow).not.toHaveBeenCalled()
})
it('should show an advisory warning when publish reports a skipped branch reference', async () => {
const user = userEvent.setup()
mockPublishWorkflow.mockResolvedValueOnce({
created_at: '2024-01-01T00:00:00Z',
warning: '"Answer" ← "Producer"',
})
mockUseNodes.mockReturnValue([{ id: 'start', data: { type: BlockEnum.Start } }])
mockUseEdges.mockReturnValue([{ source: 'start' }])
renderWithToast(<FeaturesTrigger />)
await user.click(screen.getByRole('button', { name: 'publisher-publish' }))
await waitFor(() => {
expect(toastMocks.call).toHaveBeenCalledWith({
type: 'warning',
message: '"Answer" ← "Producer"',
})
})
})
it('should publish workflow and update related stores when validation passes', async () => {
// Arrange
const user = userEvent.setup()
@@ -202,6 +202,7 @@ const FeaturesTrigger = () => {
if (options?.showSuccessToast !== false) {
toast.success(t(($) => $['api.actionSuccess'], { ns: 'common' }))
}
if (res.warning) toast.warning(res.warning)
updatePublishedWorkflow(appID!)
updateAppDetail()
invalidateAppTriggers(appID!)
+139
View File
@@ -218,4 +218,143 @@ describe('useImportDSL', () => {
expect(mockGetRedirection).toHaveBeenCalledTimes(1)
expect(result.current.isFetching).toBe(false)
})
it('should toast the backend error when import status is failed', async () => {
const importError =
'Missing app data in YAML content. ' +
"Not a valid Dify app DSL: the top-level 'app' section is required (found: meta)."
mockImportDSL.mockResolvedValue({
id: 'import-failed',
status: DSLImportStatus.FAILED,
error: importError,
})
const onFailed = vi.fn()
const { result } = renderHookWithConsoleQuery(() => useImportDSL())
await act(async () => {
await result.current.handleImportDSL(
{
mode: DSLImportMode.YAML_CONTENT,
yaml_content: 'meta: {}\n',
},
{ onFailed },
)
})
expect(toastMocks.error).toHaveBeenCalledExactlyOnceWith('app.newApp.appCreateFailed', {
description: importError,
})
expect(onFailed).toHaveBeenCalled()
})
it.each([
[
'JSON',
() => Response.json({ message: 'Missing app section' }, { status: 400 }),
'Missing app section',
],
['empty', () => new Response(null, { status: 500 }), undefined],
['HTML', () => new Response('<html>Bad gateway</html>', { status: 502 }), undefined],
] as const)(
'shows one failure toast for a %s import response',
async (_, response, description) => {
mockImportDSL.mockRejectedValue(response())
const onFailed = vi.fn()
const { result } = renderHookWithConsoleQuery(() => useImportDSL())
await act(async () => {
await result.current.handleImportDSL(
{
mode: DSLImportMode.YAML_CONTENT,
yaml_content: 'meta: {}\n',
},
{ onFailed },
)
})
expect(toastMocks.error).toHaveBeenCalledExactlyOnceWith('app.newApp.appCreateFailed', {
description,
})
expect(onFailed).toHaveBeenCalled()
},
)
it('should toast a generic failure when import throws before a response', async () => {
mockImportDSL.mockRejectedValue(new Error('network'))
const onFailed = vi.fn()
const { result } = renderHookWithConsoleQuery(() => useImportDSL())
await act(async () => {
await result.current.handleImportDSL(
{
mode: DSLImportMode.YAML_CONTENT,
yaml_content: 'meta: {}\n',
},
{ onFailed },
)
})
expect(toastMocks.error).toHaveBeenCalledExactlyOnceWith('app.newApp.appCreateFailed', {
description: 'network',
})
expect(onFailed).toHaveBeenCalled()
})
it.each([undefined, 'Import confirmation expired'])(
'shows one error when a confirmed import fails (%s)',
async (importError) => {
mockImportDSL.mockResolvedValue({
id: 'import-1',
status: DSLImportStatus.PENDING,
})
mockImportDSLConfirm.mockResolvedValue({
id: 'import-1',
status: DSLImportStatus.FAILED,
error: importError,
})
const onFailed = vi.fn()
const { result } = renderHookWithConsoleQuery(() => useImportDSL())
await act(async () => {
await result.current.handleImportDSL(
{
mode: DSLImportMode.YAML_CONTENT,
yaml_content: 'app: demo',
},
{},
)
})
await act(async () => {
await result.current.handleImportDSLConfirm({ onFailed })
})
expect(toastMocks.error).toHaveBeenCalledExactlyOnceWith('app.newApp.appCreateFailed', {
description: importError,
})
expect(onFailed).toHaveBeenCalled()
},
)
it('shows the backend error when confirmation rejects with an HTTP response', async () => {
mockImportDSL.mockResolvedValue({ id: 'import-1', status: DSLImportStatus.PENDING })
mockImportDSLConfirm.mockRejectedValue(
Response.json({ error: 'Import confirmation expired' }, { status: 400 }),
)
const onFailed = vi.fn()
const { result } = renderHookWithConsoleQuery(() => useImportDSL())
await act(async () => {
await result.current.handleImportDSL(
{ mode: DSLImportMode.YAML_CONTENT, yaml_content: 'app: demo' },
{},
)
await result.current.handleImportDSLConfirm({ onFailed })
})
expect(toastMocks.error).toHaveBeenCalledExactlyOnceWith('app.newApp.appCreateFailed', {
description: 'Import confirmation expired',
})
expect(onFailed).toHaveBeenCalledExactlyOnceWith()
expect(result.current.isFetching).toBe(false)
})
})
+38 -10
View File
@@ -5,6 +5,7 @@ import { useAtomValue } from 'jotai'
import { createElement, useCallback, useRef, useState } from 'react'
import { useTranslation } from 'react-i18next'
import DSLImportWarningDescription from '@/app/components/app/create-from-dsl-modal/dsl-import-warning-description'
import { getAppTransferErrorMessage } from '@/app/components/app/transfer-error'
import { usePluginDependencies } from '@/app/components/workflow/plugin-dependency/hooks'
import { toast } from '@/app/notifications'
import { workspacePermissionKeysAtom } from '@/context/permission-state'
@@ -25,13 +26,18 @@ type ResponseCallback = {
onFailed?: () => void
skipRedirectOnSuccess?: boolean
}
export const useImportDSL = () => {
const { t } = useTranslation(['app'])
const { handleCheckPluginDependencies } = usePluginDependencies()
const { push } = useRouter()
const { mutateAsync: importApp } = useMutation(consoleQuery.apps.imports.post.mutationOptions())
const { mutateAsync: importApp } = useMutation(
consoleQuery.apps.imports.post.mutationOptions({ context: { silent: true } }),
)
const { mutateAsync: confirmImport } = useMutation(
consoleQuery.apps.imports.byImportId.confirm.post.mutationOptions(),
consoleQuery.apps.imports.byImportId.confirm.post.mutationOptions({
context: { silent: true },
}),
)
const actionInFlightRef = useRef(false)
const [isFetching, setIsFetching] = useState(false)
@@ -112,11 +118,21 @@ export const useImportDSL = () => {
importIdRef.current = id
onPending?.(response)
} else {
toast.error(t(($) => $['newApp.appCreateFailed'], { ns: 'app' }))
toast.error(
t(($) => $['newApp.appCreateFailed'], { ns: 'app' }),
{
description: response.error || undefined,
},
)
onFailed?.()
}
} catch {
toast.error(t(($) => $['newApp.appCreateFailed'], { ns: 'app' }))
} catch (error) {
toast.error(
t(($) => $['newApp.appCreateFailed'], { ns: 'app' }),
{
description: await getAppTransferErrorMessage(error),
},
)
onFailed?.()
} finally {
actionInFlightRef.current = false
@@ -151,6 +167,16 @@ export const useImportDSL = () => {
})
const { status, app_id, app_mode, permission_keys } = response
if (status === DSLImportStatus.FAILED) {
toast.error(
t(($) => $['newApp.appCreateFailed'], { ns: 'app' }),
{
description: response.error || undefined,
},
)
onFailed?.()
return
}
if (!app_id) return
if (
@@ -188,12 +214,14 @@ export const useImportDSL = () => {
isRbacEnabled,
})
}
} else if (status === DSLImportStatus.FAILED) {
toast.error(t(($) => $['newApp.appCreateFailed'], { ns: 'app' }))
onFailed?.()
}
} catch {
toast.error(t(($) => $['newApp.appCreateFailed'], { ns: 'app' }))
} catch (error) {
toast.error(
t(($) => $['newApp.appCreateFailed'], { ns: 'app' }),
{
description: await getAppTransferErrorMessage(error),
},
)
onFailed?.()
} finally {
actionInFlightRef.current = false
+39 -2
View File
@@ -22,7 +22,7 @@ import { act, renderHook } from '@testing-library/react'
import { beforeEach, describe, expect, it, vi } from 'vite-plus/test'
import { consoleQuery } from '@/service/console'
import { AppModeEnum } from '@/types/app'
import { useUpdateWorkflow } from '../use-workflow'
import { usePublishWorkflow, useUpdateWorkflow } from '../use-workflow'
import {
appWorkflowQueryOptions,
appWorkflowVersionsInfiniteQueryKey,
@@ -30,12 +30,13 @@ import {
} from '../workflow-queries'
const mockPatch = vi.hoisted(() => vi.fn())
const mockPost = vi.hoisted(() => vi.fn())
vi.mock('../base', () => ({
del: vi.fn(),
get: vi.fn(),
patch: (...args: unknown[]) => mockPatch(...args),
post: vi.fn(),
post: (...args: unknown[]) => mockPost(...args),
put: vi.fn(),
}))
@@ -109,6 +110,42 @@ const createEnvironmentDeployment = ({
},
})
describe('usePublishWorkflow', () => {
beforeEach(() => {
vi.clearAllMocks()
})
it('returns the server warning with the publish result', async () => {
const published = {
result: 'success',
created_at: 1_710_000_100,
warning: '"Answer" ← "Producer"',
}
mockPost.mockResolvedValue(published)
const queryClient = createQueryClient()
const { result } = renderHook(() => usePublishWorkflow(), {
wrapper: createWrapper(queryClient),
})
let response: unknown
await act(async () => {
response = await result.current.mutateAsync({
url: '/apps/app-1/workflows/publish',
title: 'Release 1',
releaseNotes: 'Notes',
})
})
expect(response).toEqual(published)
expect(mockPost).toHaveBeenCalledWith('/apps/app-1/workflows/publish', {
body: {
marked_name: 'Release 1',
marked_comment: 'Notes',
},
})
})
})
describe('useUpdateWorkflow', () => {
beforeEach(() => {
vi.clearAllMocks()
+2 -1
View File
@@ -1,5 +1,6 @@
import type {
WorkflowPaginationResponse,
WorkflowPublishResponse,
WorkflowResponse,
} from '@dify/contracts/api/console/apps/types.gen'
import type {
@@ -317,7 +318,7 @@ export const usePublishWorkflow = () => {
return useMutation({
mutationKey: [NAME_SPACE, 'publish'],
mutationFn: (params: PublishWorkflowParams) =>
post<CommonResponse & { created_at: number }>(params.url, {
post<WorkflowPublishResponse>(params.url, {
body: {
marked_name: params.title,
marked_comment: params.releaseNotes,