fix: handle unconfigured LLM workflow streams (#39741)

This commit is contained in:
林玮 (Jade Lin) 2026-07-29 14:07:24 +08:00 committed by GitHub
parent d4eab4ae88
commit 14a938c5c5
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 106 additions and 13 deletions

View File

@ -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.")

View File

@ -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,
)

View File

@ -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,
):

View File

@ -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(
[