mirror of
https://github.com/langgenius/dify.git
synced 2026-09-28 06:13:22 +08:00
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:
co-authored by
Cursor
autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Crazywoola
parent
f6a1878ceb
commit
d86435ee1e
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"])
|
||||
|
||||
+21
@@ -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()
|
||||
|
||||
+20
@@ -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!)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user