diff --git a/api/controllers/console/app/conversation_variables.py b/api/controllers/console/app/conversation_variables.py index 00c9620d96e..7bc19874b44 100644 --- a/api/controllers/console/app/conversation_variables.py +++ b/api/controllers/console/app/conversation_variables.py @@ -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 diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index 9a2f85ff90a..96d6564c2c0 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -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//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//workflows/default-workflow-block-configs") diff --git a/api/core/app/apps/advanced_chat/app_runner.py b/api/core/app/apps/advanced_chat/app_runner.py index 5c6d28f524c..30470c0383b 100644 --- a/api/core/app/apps/advanced_chat/app_runner.py +++ b/api/core/app/apps/advanced_chat/app_runner.py @@ -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 diff --git a/api/models/workflow.py b/api/models/workflow.py index 63e324c5833..bc595df4cc2 100644 --- a/api/models/workflow.py +++ b/api/models/workflow.py @@ -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) diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index c8580bcceab..1ff609023b9 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -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 diff --git a/api/services/app_dsl_service.py b/api/services/app_dsl_service.py index 5e828261d5b..5515f3ed170 100644 --- a/api/services/app_dsl_service.py +++ b/api/services/app_dsl_service.py @@ -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: diff --git a/api/services/conversation_service.py b/api/services/conversation_service.py index fda23414051..3842841136b 100644 --- a/api/services/conversation_service.py +++ b/api/services/conversation_service.py @@ -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, } diff --git a/api/services/conversation_variable_updater.py b/api/services/conversation_variable_updater.py index 287d513f480..a821b612b41 100644 --- a/api/services/conversation_variable_updater.py +++ b/api/services/conversation_variable_updater.py @@ -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) diff --git a/api/services/workflow_variable_reference_validator.py b/api/services/workflow_variable_reference_validator.py new file mode 100644 index 00000000000..8ac841724ea --- /dev/null +++ b/api/services/workflow_variable_reference_validator.py @@ -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 diff --git a/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py b/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py index bfd030b2e7a..f64ff3978f3 100644 --- a/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py @@ -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 diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow.py b/api/tests/unit_tests/controllers/console/app/test_workflow.py index c55a923b22f..7c454dd1cf1 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow.py @@ -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"]) diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_conversation_variables.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_conversation_variables.py index 4b34242e5dd..5ed29a74063 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_conversation_variables.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_conversation_variables.py @@ -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 diff --git a/api/tests/unit_tests/models/test_conversation_variable.py b/api/tests/unit_tests/models/test_conversation_variable.py index bb3a6db1a1c..795cbc3f1b1 100644 --- a/api/tests/unit_tests/models/test_conversation_variable.py +++ b/api/tests/unit_tests/models/test_conversation_variable.py @@ -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( { diff --git a/api/tests/unit_tests/services/test_app_dsl_service.py b/api/tests/unit_tests/services/test_app_dsl_service.py index 544142e9d4b..cc2ba0a6c20 100644 --- a/api/tests/unit_tests/services/test_app_dsl_service.py +++ b/api/tests/unit_tests/services/test_app_dsl_service.py @@ -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() diff --git a/api/tests/unit_tests/services/test_conversation_service.py b/api/tests/unit_tests/services/test_conversation_service.py index 0331b02498e..bae347f1c1a 100644 --- a/api/tests/unit_tests/services/test_conversation_service.py +++ b/api/tests/unit_tests/services/test_conversation_service.py @@ -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, diff --git a/api/tests/unit_tests/services/test_conversation_variable_updater.py b/api/tests/unit_tests/services/test_conversation_variable_updater.py new file mode 100644 index 00000000000..6710da8b266 --- /dev/null +++ b/api/tests/unit_tests/services/test_conversation_variable_updater.py @@ -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" diff --git a/api/tests/unit_tests/services/test_workflow_variable_reference_validator.py b/api/tests/unit_tests/services/test_workflow_variable_reference_validator.py new file mode 100644 index 00000000000..6696e304c3c --- /dev/null +++ b/api/tests/unit_tests/services/test_workflow_variable_reference_validator.py @@ -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 diff --git a/packages/contracts/generated/api/console/apps/types.gen.ts b/packages/contracts/generated/api/console/apps/types.gen.ts index 03341a107bd..3ca72d79dcd 100644 --- a/packages/contracts/generated/api/console/apps/types.gen.ts +++ b/packages/contracts/generated/api/console/apps/types.gen.ts @@ -1156,6 +1156,7 @@ export type PublishWorkflowPayload = { export type WorkflowPublishResponse = { created_at: number result: string + warning?: string | null } export type WebhookTriggerResponse = { diff --git a/packages/contracts/generated/api/console/apps/zod.gen.ts b/packages/contracts/generated/api/console/apps/zod.gen.ts index 96cc3ac76d0..e7f2b99cc46 100644 --- a/packages/contracts/generated/api/console/apps/zod.gen.ts +++ b/packages/contracts/generated/api/console/apps/zod.gen.ts @@ -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(), }) /** diff --git a/packages/contracts/generated/api/console/snippets/types.gen.ts b/packages/contracts/generated/api/console/snippets/types.gen.ts index 57c28d99cbe..dd5f456c6ca 100644 --- a/packages/contracts/generated/api/console/snippets/types.gen.ts +++ b/packages/contracts/generated/api/console/snippets/types.gen.ts @@ -267,6 +267,7 @@ export type PublishWorkflowPayload = { export type WorkflowPublishResponse = { created_at: number result: string + warning?: string | null } export type WorkflowUpdatePayload = { diff --git a/packages/contracts/generated/api/console/snippets/zod.gen.ts b/packages/contracts/generated/api/console/snippets/zod.gen.ts index 520173ccdcc..f2108b13f7a 100644 --- a/packages/contracts/generated/api/console/snippets/zod.gen.ts +++ b/packages/contracts/generated/api/console/snippets/zod.gen.ts @@ -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(), }) /** diff --git a/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx b/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx index 55668f02cf1..3fff68ffa61 100644 --- a/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx +++ b/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx @@ -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( + , + ) + + 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', diff --git a/web/app/components/app/create-from-dsl-modal/index.tsx b/web/app/components/app/create-from-dsl-modal/index.tsx index 3ac9910d838..3023ed06d0b 100644 --- a/web/app/components/app/create-from-dsl-modal/index.tsx +++ b/web/app/components/app/create-from-dsl-modal/index.tsx @@ -261,7 +261,9 @@ function CreateFromDSLModal({ } catch (error) { toast.error( t(($) => $['newApp.appCreateFailed'], { ns: 'app' }), - { description: await getAppTransferErrorMessage(error) }, + { + description: await getAppTransferErrorMessage(error), + }, ) return } diff --git a/web/app/components/apps/__tests__/index.spec.tsx b/web/app/components/apps/__tests__/index.spec.tsx index b34d9de6789..59cbac97c5f 100644 --- a/web/app/components/apps/__tests__/index.spec.tsx +++ b/web/app/components/apps/__tests__/index.spec.tsx @@ -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() diff --git a/web/app/components/workflow-app/components/workflow-header/__tests__/features-trigger.spec.tsx b/web/app/components/workflow-app/components/workflow-header/__tests__/features-trigger.spec.tsx index 0bae47e95a5..ae4ea32bc68 100644 --- a/web/app/components/workflow-app/components/workflow-header/__tests__/features-trigger.spec.tsx +++ b/web/app/components/workflow-app/components/workflow-header/__tests__/features-trigger.spec.tsx @@ -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() + + 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() diff --git a/web/app/components/workflow-app/components/workflow-header/features-trigger.tsx b/web/app/components/workflow-app/components/workflow-header/features-trigger.tsx index fa1bd68e3a7..7be97eaccac 100644 --- a/web/app/components/workflow-app/components/workflow-header/features-trigger.tsx +++ b/web/app/components/workflow-app/components/workflow-header/features-trigger.tsx @@ -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!) diff --git a/web/hooks/use-import-dsl.spec.tsx b/web/hooks/use-import-dsl.spec.tsx index d82a98aa12d..543fbe44c3b 100644 --- a/web/hooks/use-import-dsl.spec.tsx +++ b/web/hooks/use-import-dsl.spec.tsx @@ -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('Bad gateway', { 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) + }) }) diff --git a/web/hooks/use-import-dsl.ts b/web/hooks/use-import-dsl.ts index 8ec3fda87f5..7f634f5f475 100644 --- a/web/hooks/use-import-dsl.ts +++ b/web/hooks/use-import-dsl.ts @@ -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 diff --git a/web/service/__tests__/use-workflow.spec.tsx b/web/service/__tests__/use-workflow.spec.tsx index 3972dad07ce..447326d53b9 100644 --- a/web/service/__tests__/use-workflow.spec.tsx +++ b/web/service/__tests__/use-workflow.spec.tsx @@ -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() diff --git a/web/service/use-workflow.ts b/web/service/use-workflow.ts index 139075aff23..cb8211823ee 100644 --- a/web/service/use-workflow.ts +++ b/web/service/use-workflow.ts @@ -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(params.url, { + post(params.url, { body: { marked_name: params.title, marked_comment: params.releaseNotes,