diff --git a/api/core/app/llm/model_access.py b/api/core/app/llm/model_access.py index 765268f7a0e..d2b8e3539fa 100644 --- a/api/core/app/llm/model_access.py +++ b/api/core/app/llm/model_access.py @@ -151,6 +151,9 @@ def fetch_model_config( credentials_provider: CredentialsProvider, model_factory: DifyModelFactory, ) -> tuple[ModelInstance, ModelConfigWithCredentialsEntity]: + if not node_data_model.provider or not node_data_model.name: + raise ValueError("LLM provider and model are required.") + if not node_data_model.mode: raise LLMModeRequiredError("LLM mode is required.") diff --git a/api/tasks/app_generate/workflow_execute_task.py b/api/tasks/app_generate/workflow_execute_task.py index 9bc09bac781..8a88ff4dfa6 100644 --- a/api/tasks/app_generate/workflow_execute_task.py +++ b/api/tasks/app_generate/workflow_execute_task.py @@ -311,34 +311,45 @@ def _publish_failed_workflow_terminal_events(exc: Exception, exec_params: AppExe topic.publish(json.dumps(finished_payload.model_dump(mode="json"), ensure_ascii=False).encode()) -def _get_event_name(event: str | Mapping[str, Any] | BaseModel) -> str | None: +def _get_event_data(event: str | Mapping[str, Any] | BaseModel) -> Mapping[str, Any] | None: if isinstance(event, BaseModel): # Temporary compatibility for legacy BaseModel stream events; remove after confirming generators always emit # str / Mapping responses. - event_name = getattr(event, "event", None) - elif isinstance(event, Mapping): - event_name = event.get("event") - else: + return event.model_dump() + if isinstance(event, Mapping): + return event + return None + + +def _get_event_name(event: str | Mapping[str, Any] | BaseModel) -> str | None: + event_data = _get_event_data(event) + if event_data is None: return None + event_name = event_data.get("event") if event_name is None: return None return str(event_name) def _get_task_id(event: str | Mapping[str, Any] | BaseModel) -> str | None: - if isinstance(event, BaseModel): - # Temporary compatibility for legacy BaseModel stream events; remove after confirming generators always emit - # str / Mapping responses. - task_id = getattr(event, "task_id", None) - elif isinstance(event, Mapping): - task_id = event.get("task_id") - else: + event_data = _get_event_data(event) + if event_data is None: return None + task_id = event_data.get("task_id") return task_id if isinstance(task_id, str) and task_id else None +def _get_error_message(event: str | Mapping[str, Any] | BaseModel) -> str | None: + event_data = _get_event_data(event) + if event_data is None: + return None + + message = event_data.get("message") + return message if isinstance(message, str) and message else None + + def _publish_streaming_response( response_stream: Generator[str | Mapping[str, Any] | BaseModel, None, None], workflow_run_id: str | uuid.UUID, @@ -406,6 +417,7 @@ def _publish_streaming_response( started_published = False terminal_published = False last_task_id = normalized_workflow_run_id + stream_error_message: str | None = None try: for event in response_stream: @@ -429,6 +441,8 @@ def _publish_streaming_response( started_published = True elif event_name in terminal_events: terminal_published = True + elif event_name == "error": + stream_error_message = _get_error_message(event) or stream_error_message except Exception as exc: if not terminal_published: logger.exception( @@ -448,7 +462,7 @@ def _publish_streaming_response( normalized_workflow_run_id, ) _publish_failed_terminal_event( - error_message=unexpected_stream_end_message, + error_message=stream_error_message or unexpected_stream_end_message, task_id=last_task_id, publish_started=not started_published, ) diff --git a/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py b/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py index d437c565949..d6f6771a5c2 100644 --- a/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py @@ -351,6 +351,33 @@ def test_fetch_model_config_hydrates_model_instance_runtime_settings(model_confi provider_model.raise_for_status.assert_called_once() +@pytest.mark.parametrize( + ("provider", "model_name"), + [ + ("", "gpt-3.5-turbo"), + ("openai", ""), + ], +) +def test_fetch_model_config_rejects_unconfigured_model(provider: str, model_name: str): + credentials_provider = mock.MagicMock(spec=CredentialsProvider) + model_factory = mock.MagicMock(spec=DifyModelFactory) + + with pytest.raises(ValueError, match="LLM provider and model are required"): + fetch_model_config( + node_data_model=ModelConfig( + provider=provider, + name=model_name, + mode="chat", + completion_params={}, + ), + credentials_provider=credentials_provider, + model_factory=model_factory, + ) + + credentials_provider.fetch.assert_not_called() + model_factory.init_model_instance.assert_not_called() + + def test_fetch_model_config_reuses_validated_provider_model_from_dify_credentials_provider( model_config: ModelConfigWithCredentialsEntity, ): diff --git a/api/tests/unit_tests/tasks/test_workflow_execute_task.py b/api/tests/unit_tests/tasks/test_workflow_execute_task.py index 3b9cad30018..f99d4a1d942 100644 --- a/api/tests/unit_tests/tasks/test_workflow_execute_task.py +++ b/api/tests/unit_tests/tasks/test_workflow_execute_task.py @@ -3,6 +3,7 @@ from __future__ import annotations import json import logging import uuid +from collections.abc import Generator, Mapping from contextlib import nullcontext from datetime import datetime from decimal import Decimal @@ -36,6 +37,7 @@ from tasks.app_generate.workflow_execute_task import ( class _StreamEventModel(BaseModel): event: object | None = None task_id: object | None = None + message: object | None = None def _build_advanced_chat_generate_entity(conversation_id: str | None) -> AdvancedChatAppGenerateEntity: @@ -248,6 +250,21 @@ def test_get_task_id(event: object, expected: str | None): assert workflow_execute_task_module._get_task_id(event) == expected +@pytest.mark.parametrize( + ("event", "expected"), + [ + ({"message": "workflow error"}, "workflow error"), + (_StreamEventModel(message="workflow error"), "workflow error"), + ({"message": ""}, None), + ({"message": 123}, None), + ({}, None), + ("workflow error", None), + ], +) +def test_get_error_message(event: str | Mapping[str, object] | BaseModel, expected: str | None): + assert workflow_execute_task_module._get_error_message(event) == expected + + @pytest.fixture def mock_topic(monkeypatch: pytest.MonkeyPatch) -> MagicMock: topic = MagicMock() @@ -486,6 +503,38 @@ def test_publish_streaming_response_publishes_failed_terminal_on_exhaustion_with assert "ended without a terminal event" in caplog.text +def test_publish_streaming_response_uses_error_message_for_failed_terminal(mock_topic: MagicMock): + def response_stream() -> Generator[str | Mapping[str, object] | BaseModel, None, None]: + yield { + "event": "error", + "workflow_run_id": "workflow-run-id", + "code": "invalid_param", + "message": "LLM provider and model are required.", + "status": 400, + } + + _publish_streaming_response( + response_stream(), + "workflow-run-id", + app_mode=AppMode.WORKFLOW, + workflow_id="workflow-id", + inputs={}, + started_reason=WorkflowStartReason.INITIAL, + ) + + payloads = _published_payloads(mock_topic) + error_payload = payloads[0] + finished_payload = payloads[-1] + assert isinstance(error_payload, dict) + assert isinstance(finished_payload, dict) + assert error_payload["status"] == 400 + assert error_payload["message"] == "LLM provider and model are required." + finished_data = finished_payload["data"] + assert isinstance(finished_data, dict) + assert finished_data["status"] == WorkflowExecutionStatus.FAILED + assert finished_data["error"] == "LLM provider and model are required." + + def test_publish_streaming_response_does_not_publish_synthetic_failure_after_terminal_event(mock_topic: MagicMock): response_stream = iter( [