Compare commits

...
Author SHA1 Message Date
-LAN- 62518a5e55 fix(completion): preserve legacy workflow compatibility 2026-07-28 13:04:11 +08:00
-LAN- e2c769c6cc fix(workflow): keep converter imports sorted 2026-07-28 13:03:12 +08:00
-LAN- 6ad1a106c0 fix(workflow): preserve polling llm invoke hook 2026-07-28 13:03:12 +08:00
-LAN- 659c486e12 test(completion): cover workflow runner compatibility branches 2026-07-28 13:03:12 +08:00
-LAN- eaa55a4292 fix(completion): keep tts text required in schema 2026-07-28 13:03:11 +08:00
-LAN- 4122233ba0 fix(task-pipeline): tolerate message end events without saved prompt 2026-07-28 13:03:11 +08:00
autofix-ci[bot]and-LAN- dfc2726e44 [autofix.ci] apply automated fixes 2026-07-28 13:03:11 +08:00
-LAN- de8ae2a6ab refactor(completion): run through workflow entry
Replace the dedicated Completion runner with a Workflow Entry backed execution path.

Adapt GraphOn events back into legacy Completion queue events and keep existing message persistence in the EasyUI task pipeline.

Refs #37572
2026-07-28 13:03:11 +08:00
-LAN- eb82fa8e04 docs(completion): add workflow entry implementation plan 2026-07-28 13:03:10 +08:00
-LAN- 8ca84a7b45 docs(completion): add workflow entry reuse design 2026-07-28 13:03:10 +08:00
24 changed files with 1428 additions and 458 deletions
+3 -1
View File
@@ -278,7 +278,9 @@ class ChatMessageTextApi(Resource):
@get_app_model
def post(self, app_model: App):
try:
payload = TextToSpeechPayload.model_validate(console_ns.payload)
payload_data = dict(console_ns.payload or {})
payload_data.setdefault("text", "")
payload = TextToSpeechPayload.model_validate(payload_data)
message_ref = None
if payload.message_id:
app_ref = AppRefService.create_app_ref(app_model)
+13 -6
View File
@@ -1,7 +1,7 @@
import base64
import logging
import time
from collections.abc import Generator, Mapping, Sequence
from collections.abc import Generator, Mapping, MutableMapping, Sequence
from mimetypes import guess_extension
from typing import TYPE_CHECKING, Any, Union
@@ -54,9 +54,15 @@ _logger = logging.getLogger(__name__)
class AppRunner:
def recalc_llm_max_tokens(
self, model_config: ModelConfigWithCredentialsEntity, prompt_messages: list[PromptMessage]
):
self,
model_config: ModelConfigWithCredentialsEntity,
prompt_messages: list[PromptMessage],
*,
model_parameters: MutableMapping[str, Any] | None = None,
) -> int | None:
"""Clamp max tokens against the final prompt on the selected parameter mapping."""
# recalc max_tokens if sum(prompt_token + max_tokens) over model token limit
parameters = model_parameters if model_parameters is not None else model_config.parameters
model_instance = ModelInstance(
provider_model_bundle=model_config.provider_model_bundle, model=model_config.model
)
@@ -69,8 +75,7 @@ class AppRunner:
parameter_rule.use_template and parameter_rule.use_template == "max_tokens"
):
max_tokens = (
model_config.parameters.get(parameter_rule.name)
or model_config.parameters.get(parameter_rule.use_template or "")
parameters.get(parameter_rule.name) or parameters.get(parameter_rule.use_template or "")
) or 0
if model_context_tokens is None:
@@ -85,7 +90,9 @@ class AppRunner:
if parameter_rule.name == "max_tokens" or (
parameter_rule.use_template and parameter_rule.use_template == "max_tokens"
):
model_config.parameters[parameter_rule.name] = max_tokens
parameters[parameter_rule.name] = max_tokens
return None
def organize_prompt_messages(
self,
@@ -15,8 +15,8 @@ from core.app.app_config.easy_ui_based_app.model_config.converter import ModelCo
from core.app.app_config.features.file_upload.manager import FileUploadConfigManager
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.completion.app_config_manager import CompletionAppConfigManager
from core.app.apps.completion.app_runner import CompletionAppRunner
from core.app.apps.completion.generate_response_converter import CompletionAppGenerateResponseConverter
from core.app.apps.completion.workflow_runner import CompletionWorkflowRunner
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
@@ -243,7 +243,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
message = self._get_message(message_id)
# chatbot app
runner = CompletionAppRunner()
runner = CompletionWorkflowRunner()
with session_factory.create_session() as session:
runner.run(
application_generate_entity=application_generate_entity,
-208
View File
@@ -1,208 +0,0 @@
import logging
from typing import cast
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.base_app_runner import AppRunner
from core.app.apps.completion.app_config_manager import CompletionAppConfig
from core.app.entities.app_invoke_entities import (
CompletionAppGenerateEntity,
)
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.db.session_factory import create_session
from core.model_manager import ModelInstance
from core.moderation.base import ModerationError
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
from graphon.file import File
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
from models.model import App, Message
logger = logging.getLogger(__name__)
class CompletionAppRunner(AppRunner):
"""
Completion Application Runner
"""
def run(
self,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
message: Message,
session: Session,
):
"""Run the application without retaining ``session`` during model I/O.
Database preparation is committed and the connection is released before
the provider response is requested or consumed.
:param application_generate_entity: application generate entity
:param queue_manager: application queue manager
:param message: message
:return:
"""
app_config = application_generate_entity.app_config
app_config = cast(CompletionAppConfig, app_config)
stmt = select(App).where(App.id == app_config.app_id)
with create_session() as read_session:
app_record = read_session.scalar(stmt)
if app_record:
read_session.expunge(app_record)
if not app_record:
raise ValueError("App not found")
inputs = application_generate_entity.inputs
query = application_generate_entity.query
files = application_generate_entity.files
image_detail_config = (
application_generate_entity.file_upload_config.image_config.detail
if (
application_generate_entity.file_upload_config
and application_generate_entity.file_upload_config.image_config
)
else None
)
image_detail_config = image_detail_config or ImagePromptMessageContent.DETAIL.LOW
# organize all inputs and template to prompt messages
# Include: prompt template, inputs, query(optional), files(optional)
prompt_messages, stop = self.organize_prompt_messages(
app_record=app_record,
model_config=application_generate_entity.model_conf,
prompt_template_entity=app_config.prompt_template,
inputs=inputs,
files=files,
query=query,
image_detail_config=image_detail_config,
)
# moderation
try:
# process sensitive_word_avoidance
_, inputs, query = self.moderation_for_inputs(
app_id=app_record.id,
tenant_id=app_config.tenant_id,
app_generate_entity=application_generate_entity,
inputs=inputs,
query=query or "",
message_id=message.id,
)
except ModerationError as e:
self.direct_output(
queue_manager=queue_manager,
app_generate_entity=application_generate_entity,
prompt_messages=prompt_messages,
text=str(e),
stream=application_generate_entity.stream,
)
return
# fill in variable inputs from external data tools if exists
external_data_tools = app_config.external_data_variables
if external_data_tools:
inputs = self.fill_in_inputs_from_external_data_tools(
tenant_id=app_record.tenant_id,
app_id=app_record.id,
external_data_tools=external_data_tools,
inputs=inputs,
query=query,
)
# get context from datasets
context = None
context_files: list[File] = []
if app_config.dataset and app_config.dataset.dataset_ids:
hit_callback = DatasetIndexToolCallbackHandler(
queue_manager,
app_record.id,
message.id,
application_generate_entity.user_id,
application_generate_entity.invoke_from,
)
dataset_config = app_config.dataset
if dataset_config and dataset_config.retrieve_config.query_variable:
query = inputs.get(dataset_config.retrieve_config.query_variable, "")
dataset_retrieval = DatasetRetrieval(application_generate_entity)
context, retrieved_files = dataset_retrieval.retrieve(
session=session,
app_id=app_record.id,
user_id=application_generate_entity.user_id,
tenant_id=app_record.tenant_id,
model_config=application_generate_entity.model_conf,
config=dataset_config,
query=query or "",
invoke_from=application_generate_entity.invoke_from,
show_retrieve_source=app_config.additional_features.show_retrieve_source
if app_config.additional_features
else False,
hit_callback=hit_callback,
message_id=message.id,
inputs=inputs,
vision_enabled=bool(
application_generate_entity.app_config.app_model_config_dict.get("file_upload", {})
.get("image", {})
.get("enabled", False)
),
)
context_files = retrieved_files or []
session.commit()
session.close()
# reorganize all inputs and template to prompt messages
# Include: prompt template, inputs, query(optional), files(optional)
# memory(optional), external data, dataset context(optional)
prompt_messages, stop = self.organize_prompt_messages(
app_record=app_record,
model_config=application_generate_entity.model_conf,
prompt_template_entity=app_config.prompt_template,
inputs=inputs,
files=files,
query=query,
context=context,
image_detail_config=image_detail_config,
context_files=context_files,
)
# check hosting moderation
hosting_moderation_result = self.check_hosting_moderation(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
prompt_messages=prompt_messages,
)
if hosting_moderation_result:
return
# Re-calculate the max tokens if sum(prompt_token + max_tokens) over model token limit
self.recalc_llm_max_tokens(model_config=application_generate_entity.model_conf, prompt_messages=prompt_messages)
# Invoke model
model_instance = ModelInstance(
provider_model_bundle=application_generate_entity.model_conf.provider_model_bundle,
model=application_generate_entity.model_conf.model,
)
invoke_result = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=application_generate_entity.model_conf.parameters,
stop=stop,
stream=application_generate_entity.stream,
request_metadata={"app_id": app_config.app_id},
)
# handle invoke result
self._handle_invoke_result(
invoke_result=invoke_result,
queue_manager=queue_manager,
stream=application_generate_entity.stream,
message_id=message.id,
user_id=application_generate_entity.user_id,
tenant_id=app_config.tenant_id,
)
@@ -0,0 +1,144 @@
from collections.abc import Mapping, Sequence
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.entities.app_invoke_entities import CompletionAppGenerateEntity
from core.app.entities.queue_entities import (
QueueErrorEvent,
QueueLLMChunkEvent,
QueueMessageEndEvent,
QueueRetrieverResourcesEvent,
QueueStopEvent,
)
from core.rag.entities import RetrievalSourceMetadata
from graphon.graph_events import (
GraphEngineEvent,
GraphRunAbortedEvent,
GraphRunFailedEvent,
GraphRunSucceededEvent,
NodeRunRetrieverResourceEvent,
NodeRunStreamChunkEvent,
NodeRunSucceededEvent,
)
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, PromptMessage
_LLM_TEXT_SELECTOR_PREFIX = ("llm", "text")
class CompletionGraphEventAdapter:
"""Translate one runtime graph run into legacy Completion queue events."""
_application_generate_entity: CompletionAppGenerateEntity
_queue_manager: AppQueueManager
_answer: str
_usage: LLMUsage
_prompt_messages: list[PromptMessage]
_chunk_index: int
def __init__(
self,
*,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
) -> None:
self._application_generate_entity = application_generate_entity
self._queue_manager = queue_manager
self._answer = ""
self._usage = LLMUsage.empty_usage()
self._prompt_messages = []
self._chunk_index = 0
def set_prompt_messages(self, prompt_messages: Sequence[PromptMessage]) -> None:
"""Capture the final GraphOn prompt for legacy chunks and message persistence."""
self._prompt_messages = list(prompt_messages)
def handle_event(self, event: GraphEngineEvent) -> None:
match event:
case NodeRunStreamChunkEvent():
self._handle_stream_chunk(event)
case NodeRunRetrieverResourceEvent():
self._handle_retriever_resource(event)
case NodeRunSucceededEvent():
self._handle_node_succeeded(event)
case GraphRunSucceededEvent():
self._publish_message_end(event.outputs)
case GraphRunFailedEvent():
self._publish_error(event.error)
case GraphRunAbortedEvent():
self._queue_manager.publish(
QueueStopEvent(stopped_by=QueueStopEvent.StopBy.USER_MANUAL),
PublishFrom.APPLICATION_MANAGER,
)
case _:
return
def _handle_stream_chunk(self, event: NodeRunStreamChunkEvent) -> None:
if tuple(event.selector)[:2] != _LLM_TEXT_SELECTOR_PREFIX:
return
if event.is_final and not event.chunk:
return
self._answer += event.chunk
self._queue_manager.publish(
QueueLLMChunkEvent(
chunk=LLMResultChunk(
model=self._application_generate_entity.model_conf.model,
prompt_messages=self._prompt_messages,
delta=LLMResultChunkDelta(
index=self._chunk_index,
message=AssistantPromptMessage(content=event.chunk),
),
)
),
PublishFrom.APPLICATION_MANAGER,
)
self._chunk_index += 1
def _handle_retriever_resource(self, event: NodeRunRetrieverResourceEvent) -> None:
additional_features = self._application_generate_entity.app_config.additional_features
if not additional_features or not additional_features.show_retrieve_source:
return
self._queue_manager.publish(
QueueRetrieverResourcesEvent(
retriever_resources=[
RetrievalSourceMetadata.model_validate(resource) for resource in event.retriever_resources
],
in_iteration_id=event.in_iteration_id,
in_loop_id=event.in_loop_id,
),
PublishFrom.APPLICATION_MANAGER,
)
def _handle_node_succeeded(self, event: NodeRunSucceededEvent) -> None:
if event.node_id != "llm":
return
result = event.node_run_result
text = result.outputs.get("text")
if isinstance(text, str):
self._answer = text
self._usage = result.llm_usage
def _publish_message_end(self, outputs: Mapping[str, object]) -> None:
result = outputs.get("result")
if isinstance(result, str) and not self._answer:
self._answer = result
self._queue_manager.publish(
QueueMessageEndEvent(
llm_result=LLMResult(
model=self._application_generate_entity.model_conf.model,
prompt_messages=self._prompt_messages,
message=AssistantPromptMessage(content=self._answer),
usage=self._usage,
)
),
PublishFrom.APPLICATION_MANAGER,
)
def _publish_error(self, error: object) -> None:
self._queue_manager.publish(
QueueErrorEvent(error=ValueError(str(error))),
PublishFrom.APPLICATION_MANAGER,
)
@@ -0,0 +1,69 @@
import json
from dataclasses import dataclass
from typing import Any
from uuid import uuid4
from sqlalchemy.orm import Session
from core.app.apps.completion.app_config_manager import CompletionAppConfig
from graphon.nodes import BuiltinNodeTypes
from models.model import App, AppMode
from services.workflow.workflow_converter import WorkflowConverter, WorkflowGraph
@dataclass(frozen=True, slots=True)
class RuntimeCompletionWorkflow:
workflow_id: str
root_node_id: str
graph_dict: WorkflowGraph
def build_runtime_completion_workflow(
*,
app_model: App,
app_config: CompletionAppConfig,
session: Session,
workflow_converter: WorkflowConverter | None = None,
) -> RuntimeCompletionWorkflow:
"""Build the transient WorkflowEntry graph used by Completion execution."""
converter = workflow_converter or WorkflowConverter()
graph, _ = converter.build_graph_from_app_config(
app_model=app_model,
app_config=app_config,
target_app_mode=AppMode.WORKFLOW,
session=session,
)
_route_external_data_query_to_sys_query(graph)
return RuntimeCompletionWorkflow(
workflow_id=f"completion-runtime-{uuid4()}",
root_node_id="start",
graph_dict=graph,
)
def _route_external_data_query_to_sys_query(graph: WorkflowGraph) -> None:
"""Preserve Completion API-based variable behavior in the runtime graph."""
for node in graph["nodes"]:
data = node.get("data", {})
if data.get("type") != BuiltinNodeTypes.HTTP_REQUEST:
continue
body = data.get("body")
if not isinstance(body, dict) or body.get("type") != "json":
continue
raw_body_data = body.get("data")
if not isinstance(raw_body_data, str):
continue
try:
body_data: dict[str, Any] = json.loads(raw_body_data)
except json.JSONDecodeError:
continue
params = body_data.get("params")
if not isinstance(params, dict) or params.get("query") != "":
continue
params["query"] = "{{#sys.query#}}"
body["data"] = json.dumps(body_data)
@@ -0,0 +1,242 @@
import time
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Any, cast
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.base_app_runner import AppRunner
from core.app.apps.completion.app_config_manager import CompletionAppConfig
from core.app.apps.completion.graph_event_adapter import CompletionGraphEventAdapter
from core.app.apps.completion.runtime_workflow_builder import build_runtime_completion_workflow
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.workflow_app_runner import init_graph
from core.app.entities.app_invoke_entities import CompletionAppGenerateEntity, UserFrom
from core.moderation.base import ModerationError
from core.workflow.node_runtime import DIFY_BEFORE_LLM_INVOKE_KEY
from core.workflow.system_variables import build_bootstrap_variables, build_system_variables
from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool
from core.workflow.workflow_entry import WorkflowEntry
from extensions.ext_redis import redis_client
from graphon.graph_engine.command_channels import RedisChannel
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessage
from graphon.runtime import GraphRuntimeState, VariablePool
from models.model import App, Message
@dataclass(frozen=True, slots=True)
class ModeratedCompletionInputs:
stopped: bool
inputs: Mapping[str, Any]
query: str
class CompletionWorkflowRunner(AppRunner):
"""Run a transient WorkflowEntry graph while the legacy task pipeline owns persistence."""
def run(
self,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
message: Message,
session: Session,
) -> None:
app_config = cast(CompletionAppConfig, application_generate_entity.app_config)
app_record = self._get_app(app_id=app_config.app_id, tenant_id=app_config.tenant_id, session=session)
moderation_result = self._run_input_moderation(
app_record=app_record,
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
message=message,
)
if moderation_result.stopped:
return
runtime_workflow = build_runtime_completion_workflow(
app_model=app_record,
app_config=app_config,
session=session,
)
variable_pool = self._build_variable_pool(
application_generate_entity=application_generate_entity,
message=message,
workflow_id=runtime_workflow.workflow_id,
root_node_id=runtime_workflow.root_node_id,
inputs=moderation_result.inputs,
query=moderation_result.query,
)
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
user_from = self._resolve_user_from(application_generate_entity)
adapter = CompletionGraphEventAdapter(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
)
extra_context = {
DIFY_BEFORE_LLM_INVOKE_KEY: self._build_before_llm_invoke_hook(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
adapter=adapter,
)
}
graph = init_graph(
app_id=app_config.app_id,
graph_config=runtime_workflow.graph_dict,
graph_runtime_state=graph_runtime_state,
user_from=user_from,
invoke_from=application_generate_entity.invoke_from,
workflow_id=runtime_workflow.workflow_id,
tenant_id=app_config.tenant_id,
user_id=application_generate_entity.user_id,
root_node_id=runtime_workflow.root_node_id,
trace_session_id=application_generate_entity.extras.get("trace_session_id"),
call_depth=application_generate_entity.call_depth,
extra_context=extra_context,
)
queue_manager.graph_runtime_state = graph_runtime_state
command_channel = RedisChannel(redis_client, f"workflow:{application_generate_entity.task_id}:commands")
workflow_entry = WorkflowEntry(
tenant_id=app_config.tenant_id,
app_id=app_config.app_id,
workflow_id=runtime_workflow.workflow_id,
graph_config=runtime_workflow.graph_dict,
graph=graph,
user_id=application_generate_entity.user_id,
user_from=user_from,
invoke_from=application_generate_entity.invoke_from,
call_depth=application_generate_entity.call_depth,
variable_pool=variable_pool,
graph_runtime_state=graph_runtime_state,
command_channel=command_channel,
)
# Do not hold a database connection during graph execution or provider streaming.
session.commit()
session.close()
for event in workflow_entry.run():
adapter.handle_event(event)
def _get_app(self, *, app_id: str, tenant_id: str, session: Session) -> App:
app_record = session.scalar(select(App).where(App.id == app_id, App.tenant_id == tenant_id))
if not app_record:
raise ValueError("App not found")
return app_record
def _run_input_moderation(
self,
*,
app_record: App,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
message: Message,
) -> ModeratedCompletionInputs:
app_config = cast(CompletionAppConfig, application_generate_entity.app_config)
prompt_messages, _ = self.organize_prompt_messages(
app_record=app_record,
model_config=application_generate_entity.model_conf,
prompt_template_entity=app_config.prompt_template,
inputs=application_generate_entity.inputs,
files=application_generate_entity.files,
query=application_generate_entity.query,
image_detail_config=self._resolve_image_detail_config(application_generate_entity),
)
try:
_, inputs, query = self.moderation_for_inputs(
app_id=app_record.id,
tenant_id=app_config.tenant_id,
app_generate_entity=application_generate_entity,
inputs=application_generate_entity.inputs,
query=application_generate_entity.query or "",
message_id=message.id,
)
except ModerationError as exc:
self.direct_output(
queue_manager=queue_manager,
app_generate_entity=application_generate_entity,
prompt_messages=prompt_messages,
text=str(exc),
stream=application_generate_entity.stream,
)
return ModeratedCompletionInputs(
stopped=True,
inputs=application_generate_entity.inputs,
query=application_generate_entity.query or "",
)
return ModeratedCompletionInputs(stopped=False, inputs=inputs, query=query)
def _build_before_llm_invoke_hook(
self,
*,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
adapter: CompletionGraphEventAdapter,
) -> Callable[[Sequence[PromptMessage], Mapping[str, Any]], Mapping[str, Any]]:
def check(
prompt_messages: Sequence[PromptMessage],
model_parameters: Mapping[str, Any],
) -> Mapping[str, Any]:
adapter.set_prompt_messages(prompt_messages)
if self.check_hosting_moderation(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
prompt_messages=list(prompt_messages),
):
raise GenerateTaskStoppedError()
adjusted_parameters = dict(model_parameters)
self.recalc_llm_max_tokens(
model_config=application_generate_entity.model_conf,
prompt_messages=list(prompt_messages),
model_parameters=adjusted_parameters,
)
return adjusted_parameters
return check
def _build_variable_pool(
self,
*,
application_generate_entity: CompletionAppGenerateEntity,
message: Message,
workflow_id: str,
root_node_id: str,
inputs: Mapping[str, Any],
query: str,
) -> VariablePool:
variable_pool = VariablePool()
system_inputs = build_system_variables(
files=application_generate_entity.files,
user_id=application_generate_entity.user_id,
app_id=application_generate_entity.app_config.app_id,
workflow_id=workflow_id,
workflow_execution_id=application_generate_entity.task_id,
timestamp=int(time.time()),
query=query,
conversation_id=message.conversation_id,
)
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(system_variables=system_inputs, environment_variables=[]),
)
add_node_inputs_to_pool(variable_pool, node_id=root_node_id, inputs=inputs)
return variable_pool
@staticmethod
def _resolve_user_from(application_generate_entity: CompletionAppGenerateEntity) -> UserFrom:
if application_generate_entity.invoke_from.runs_as_account():
return UserFrom.ACCOUNT
return UserFrom.END_USER
@staticmethod
def _resolve_image_detail_config(
application_generate_entity: CompletionAppGenerateEntity,
) -> ImagePromptMessageContent.DETAIL:
file_upload_config = application_generate_entity.file_upload_config
if file_upload_config and file_upload_config.image_config:
return file_upload_config.image_config.detail or ImagePromptMessageContent.DETAIL.LOW
return ImagePromptMessageContent.DETAIL.LOW
+61 -37
View File
@@ -93,6 +93,60 @@ from tasks.mail_human_input_delivery_task import dispatch_human_input_email_task
logger = logging.getLogger(__name__)
def init_graph(
*,
app_id: str,
graph_config: Mapping[str, Any],
graph_runtime_state: GraphRuntimeState,
user_from: UserFrom,
invoke_from: InvokeFrom,
workflow_id: str = "",
tenant_id: str = "",
user_id: str = "",
root_node_id: str | None = None,
trace_session_id: str | None = None,
call_depth: int = 0,
extra_context: Mapping[str, Any] | None = None,
) -> Graph:
if "nodes" not in graph_config or "edges" not in graph_config:
raise ValueError("nodes or edges not found in workflow graph")
if not isinstance(graph_config.get("nodes"), list):
raise ValueError("nodes in workflow graph must be a list")
if not isinstance(graph_config.get("edges"), list):
raise ValueError("edges in workflow graph must be a list")
run_context = build_dify_run_context(
tenant_id=tenant_id or "",
app_id=app_id,
user_id=user_id,
user_from=user_from,
invoke_from=invoke_from,
trace_session_id=trace_session_id,
extra_context=extra_context,
)
graph_init_context = DifyGraphInitContext(
workflow_id=workflow_id,
graph_config=graph_config,
run_context=run_context,
call_depth=call_depth,
)
node_factory = DifyNodeFactory.from_graph_init_context(
graph_init_context=graph_init_context,
graph_runtime_state=graph_runtime_state,
)
if root_node_id is None:
root_node_id = get_default_root_node_id(graph_config)
graph = Graph.init(graph_config=graph_config, node_factory=node_factory, root_node_id=root_node_id)
if not graph:
raise ValueError("graph not found in workflow")
return graph
class WorkflowBasedAppRunner:
def __init__(
self,
@@ -128,48 +182,18 @@ class WorkflowBasedAppRunner:
"""
Init graph
"""
if "nodes" not in graph_config or "edges" not in graph_config:
raise ValueError("nodes or edges not found in workflow graph")
if not isinstance(graph_config.get("nodes"), list):
raise ValueError("nodes in workflow graph must be a list")
if not isinstance(graph_config.get("edges"), list):
raise ValueError("edges in workflow graph must be a list")
# Create explicit graph init context for Graph.init.
run_context = build_dify_run_context(
tenant_id=tenant_id or "",
return init_graph(
app_id=self._app_id,
user_id=user_id,
graph_config=graph_config,
graph_runtime_state=graph_runtime_state,
user_from=user_from,
invoke_from=invoke_from,
workflow_id=workflow_id,
tenant_id=tenant_id,
user_id=user_id,
root_node_id=root_node_id,
trace_session_id=trace_session_id,
)
graph_init_context = DifyGraphInitContext(
workflow_id=workflow_id,
graph_config=graph_config,
run_context=run_context,
call_depth=0,
)
# Use the provided graph_runtime_state for consistent state management
node_factory = DifyNodeFactory.from_graph_init_context(
graph_init_context=graph_init_context,
graph_runtime_state=graph_runtime_state,
)
if root_node_id is None:
root_node_id = get_default_root_node_id(graph_config)
# init graph
graph = Graph.init(graph_config=graph_config, node_factory=node_factory, root_node_id=root_node_id)
if not graph:
raise ValueError("graph not found in workflow")
return graph
def _prepare_single_node_execution(
self,
+18 -2
View File
@@ -23,6 +23,8 @@ from core.prompt.entities.advanced_prompt_entities import MemoryConfig
from core.trigger.constants import TRIGGER_NODE_TYPES
from core.workflow.human_input_adapter import adapt_node_config_for_graph
from core.workflow.node_runtime import (
DIFY_BEFORE_LLM_INVOKE_KEY,
BeforeLLMInvoke,
DifyFileReferenceFactory,
DifyHumanInputNodeRuntime,
DifyPreparedLLM,
@@ -549,6 +551,10 @@ class DifyNodeFactory(NodeFactory):
) -> dict[str, object]:
validated_node_data = cast(LLMCompatibleNodeData, node_data)
model_instance = self._build_model_instance_for_llm_node(validated_node_data)
before_llm_invoke = cast(
BeforeLLMInvoke | None,
self.graph_init_params.run_context.get(DIFY_BEFORE_LLM_INVOKE_KEY),
)
node_init_kwargs: dict[str, object] = {
"credentials_provider": self._llm_credentials_provider,
"model_factory": self._llm_model_factory,
@@ -557,6 +563,7 @@ class DifyNodeFactory(NodeFactory):
node_data=validated_node_data,
model_instance=model_instance,
request_metadata={"app_id": self._dify_context.app_id},
before_invoke=before_llm_invoke,
)
if wrap_model_instance
else model_instance
@@ -590,13 +597,22 @@ class DifyNodeFactory(NodeFactory):
node_data: LLMCompatibleNodeData,
model_instance: ModelInstance,
request_metadata: Mapping[str, object] | None = None,
before_invoke: BeforeLLMInvoke | None = None,
) -> DifyPreparedLLM:
# Only graphon's LLM node consumes the polling protocol. Keep classifier
# and extractor nodes on the existing wrapper even if the same model
# advertises polling support.
if node_data.type == BuiltinNodeTypes.LLM and DifyNodeFactory._supports_plugin_llm_polling(model_instance):
return DifyPreparedPollingLLM(model_instance, request_metadata=request_metadata)
return DifyPreparedLLM(model_instance, request_metadata=request_metadata)
return DifyPreparedPollingLLM(
model_instance,
request_metadata=request_metadata,
before_invoke=before_invoke,
)
return DifyPreparedLLM(
model_instance,
request_metadata=request_metadata,
before_invoke=before_invoke,
)
@staticmethod
def _supports_plugin_llm_polling(model_instance: ModelInstance) -> bool:
+28 -3
View File
@@ -94,6 +94,8 @@ if TYPE_CHECKING:
from graphon.nodes.tool.entities import ToolNodeData
DIFY_BEFORE_LLM_INVOKE_KEY = "_dify_before_llm_invoke"
BeforeLLMInvoke = Callable[[Sequence[PromptMessage], Mapping[str, Any]], Mapping[str, Any]]
_file_access_controller = DatabaseFileAccessController()
@@ -150,9 +152,15 @@ class DifyFileReferenceFactory(FileReferenceFactoryProtocol):
class DifyPreparedLLM(LLMProtocol):
"""Workflow-layer adapter that hides the full `ModelInstance` API from `graphon` nodes."""
def __init__(self, model_instance: ModelInstance, request_metadata: Mapping[str, object] | None = None) -> None:
def __init__(
self,
model_instance: ModelInstance,
request_metadata: Mapping[str, object] | None = None,
before_invoke: BeforeLLMInvoke | None = None,
) -> None:
self._model_instance = model_instance
self._request_metadata = request_metadata
self._before_invoke = before_invoke
@property
@override
@@ -193,6 +201,15 @@ class DifyPreparedLLM(LLMProtocol):
def get_llm_num_tokens(self, prompt_messages: Sequence[PromptMessage]) -> int:
return self._model_instance.get_llm_num_tokens(prompt_messages)
def _run_before_invoke(
self,
prompt_messages: Sequence[PromptMessage],
model_parameters: Mapping[str, Any],
) -> Mapping[str, Any]:
if self._before_invoke is None:
return model_parameters
return self._before_invoke(prompt_messages, model_parameters)
@overload
def invoke_llm(
self,
@@ -225,6 +242,7 @@ class DifyPreparedLLM(LLMProtocol):
stop: Sequence[str] | None,
stream: bool,
) -> LLMResult | Generator[LLMResultChunk, None, None]:
model_parameters = self._run_before_invoke(prompt_messages, model_parameters)
return self._model_instance.invoke_llm(
prompt_messages=list(prompt_messages),
model_parameters=dict(model_parameters),
@@ -266,6 +284,7 @@ class DifyPreparedLLM(LLMProtocol):
stop: Sequence[str] | None,
stream: bool,
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]:
model_parameters = self._run_before_invoke(prompt_messages, model_parameters)
return invoke_llm_with_structured_output(
provider=self.provider,
model_schema=self.get_model_schema(),
@@ -285,10 +304,15 @@ class DifyPreparedLLM(LLMProtocol):
class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
"""Prepared workflow LLM adapter that exposes Graphon's polling protocol."""
def __init__(self, model_instance: ModelInstance, request_metadata: Mapping[str, object] | None = None) -> None:
def __init__(
self,
model_instance: ModelInstance,
request_metadata: Mapping[str, object] | None = None,
before_invoke: BeforeLLMInvoke | None = None,
) -> None:
from core.plugin.impl.model_runtime import PluginModelRuntime
super().__init__(model_instance, request_metadata=request_metadata)
super().__init__(model_instance, request_metadata=request_metadata, before_invoke=before_invoke)
model_type_instance = model_instance.model_type_instance
if not isinstance(model_type_instance, LargeLanguageModel):
raise TypeError("Polling wrapper requires a large-language-model instance.")
@@ -309,6 +333,7 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
stop: Sequence[str] | None,
json_schema: Mapping[str, Any] | None,
) -> LLMPollingResult:
model_parameters = self._run_before_invoke(prompt_messages, model_parameters)
return self._plugin_model_runtime.start_llm_polling(
provider=self.provider,
model=self.model_name,
+3 -3
View File
@@ -40,7 +40,7 @@ class AppTaskService:
# Legacy mechanism: Set stop flag in Redis
AppQueueManager.set_stop_flag(task_id, invoke_from, user_id)
# New mechanism: Send stop command via GraphEngine for workflow-based apps
# This ensures proper workflow status recording in the persistence layer
if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW):
# New mechanism: send stop command via GraphEngine for graph-backed apps.
# Completion uses WorkflowEntry at runtime but keeps legacy message persistence.
if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW, AppMode.COMPLETION):
GraphEngineManager(redis_client).send_stop_command(task_id)
+44 -23
View File
@@ -135,7 +135,45 @@ class WorkflowConverter:
app_model=app_model, app_model_config=app_model_config, session=session
)
# init workflow graph
graph, features = self.build_graph_from_app_config(
app_model=app_model,
app_config=app_config,
target_app_mode=new_app_mode,
session=session,
)
# create workflow record
workflow = Workflow(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
type=WorkflowType.from_app_mode(new_app_mode).value,
version=Workflow.VERSION_DRAFT,
graph=json.dumps(graph),
features=json.dumps(features),
created_by=account_id,
environment_variables=[],
conversation_variables=[],
)
session.add(workflow)
session.commit()
return workflow
def build_graph_from_app_config(
self,
*,
app_model: App,
app_config: EasyUIBasedAppConfig,
target_app_mode: AppMode,
session: Session,
) -> tuple[WorkflowGraph, dict[str, Any]]:
"""
Build a workflow graph from an EasyUI app config without persisting it.
This is shared by the persisted app-conversion flow and runtime-only
execution paths that need a graph but must not create a Workflow row.
"""
graph: WorkflowGraph = {"nodes": [], "edges": []}
# Convert list:
@@ -168,7 +206,7 @@ class WorkflowConverter:
# convert to knowledge retrieval node
if app_config.dataset:
knowledge_retrieval_node = self._convert_to_knowledge_retrieval_node(
new_app_mode=new_app_mode, dataset_config=app_config.dataset, model_config=app_config.model
new_app_mode=target_app_mode, dataset_config=app_config.dataset, model_config=app_config.model
)
if knowledge_retrieval_node:
@@ -177,7 +215,7 @@ class WorkflowConverter:
# convert to llm node
llm_node = self._convert_to_llm_node(
original_app_mode=AppMode.value_of(app_model.mode),
new_app_mode=new_app_mode,
new_app_mode=target_app_mode,
graph=graph,
model_config=app_config.model,
prompt_template=app_config.prompt_template,
@@ -189,7 +227,7 @@ class WorkflowConverter:
app_model_config_dict = app_config.app_model_config_dict
match new_app_mode:
match target_app_mode:
case AppMode.WORKFLOW:
end_node = self._convert_to_end_node()
graph = self._append_node(graph, end_node)
@@ -220,23 +258,7 @@ class WorkflowConverter:
"sensitive_word_avoidance": app_model_config_dict.get("sensitive_word_avoidance"),
}
# create workflow record
workflow = Workflow(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
type=WorkflowType.from_app_mode(new_app_mode).value,
version=Workflow.VERSION_DRAFT,
graph=json.dumps(graph),
features=json.dumps(features),
created_by=account_id,
environment_variables=[],
conversation_variables=[],
)
session.add(workflow)
session.commit()
return workflow
return graph, features
def _convert_to_app_config(
self, app_model: App, app_model_config: AppModelConfig, *, session: Session
@@ -573,8 +595,7 @@ class WorkflowConverter:
if new_app_mode == AppMode.ADVANCED_CHAT:
memory = {"role_prefix": role_prefix, "window": {"enabled": False}}
completion_params = model_config.parameters
completion_params.update({"stop": model_config.stop})
completion_params = {**model_config.parameters, "stop": model_config.stop}
return {
"id": "llm",
"position": None,
@@ -321,6 +321,33 @@ def test_console_text_api_builds_message_ref(app: Flask, monkeypatch: pytest.Mon
assert calls["message_ref"] == MessageRef(AppRef("tenant-1", "app-1"), "message-1", account_id="account-1")
def test_console_text_api_accepts_message_id_without_text(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
calls = {}
monkeypatch.setattr(AudioService, "transcript_tts", lambda **kwargs: calls.update(kwargs) or {"audio": "ok"})
api = ChatMessageTextApi()
handler = unwrap(api.post)
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1")
with (
app.test_request_context(
"/console/api/apps/app/text-to-audio",
method="POST",
json={"message_id": "0f67f8c5-8f7c-4ebd-b549-7ac8e972d37e", "streaming": True},
),
patch("controllers.console.app.audio.current_user", SimpleNamespace(id="account-1")),
):
response = handler(api, app_model=app_model)
assert response == {"audio": "ok"}
assert calls["text"] == ""
assert calls["message_ref"] == MessageRef(
AppRef("tenant-1", "app-1"),
"0f67f8c5-8f7c-4ebd-b549-7ac8e972d37e",
account_id="account-1",
)
def test_console_text_api_error_mapping(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(AudioService, "transcript_tts", lambda **_kwargs: (_ for _ in ()).throw(QuotaExceededError()))
@@ -1,21 +1,11 @@
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import ANY, MagicMock, patch
from unittest.mock import MagicMock
import pytest
from pytest_mock import MockerFixture
import core.app.apps.completion.app_runner as module
from core.app.apps.completion.app_runner import CompletionAppRunner
from core.app.apps.completion.workflow_runner import CompletionWorkflowRunner
from core.moderation.base import ModerationError
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
@pytest.fixture
def runner():
return CompletionAppRunner()
def _build_app_config(dataset=None, external_tools=None, additional_features=None):
app_config = MagicMock()
app_config.app_id = "app1"
@@ -24,7 +14,7 @@ def _build_app_config(dataset=None, external_tools=None, additional_features=Non
app_config.dataset = dataset
app_config.external_data_variables = external_tools or []
app_config.additional_features = additional_features
app_config.app_model_config_dict = {"file_upload": {"enabled": True}}
app_config.app_model_config_dict = {"file_upload": {"image": {"enabled": True}}}
return app_config
@@ -38,163 +28,40 @@ def _build_generate_entity(app_config, file_upload_config=None):
return SimpleNamespace(
app_config=app_config,
model_conf=model_conf,
inputs={"qvar": "query_from_input"},
inputs={"qvar": "original_query_from_input"},
query="original_query",
files=[],
file_upload_config=file_upload_config,
stream=True,
user_id="user",
invoke_from=MagicMock(),
trace_manager=None,
)
@contextmanager
def patched_create_session(*, return_value=None):
session = MagicMock()
session.scalar.return_value = return_value
session_context = MagicMock()
session_context.__enter__.return_value = session
with patch.object(module, "create_session", return_value=session_context):
yield session
def test_workflow_runner_direct_outputs_on_input_moderation() -> None:
runner = CompletionWorkflowRunner()
app_record = MagicMock(id="app1", tenant_id="tenant")
app_generate_entity = _build_generate_entity(_build_app_config())
queue_manager = MagicMock()
message = MagicMock(id="msg")
runner.organize_prompt_messages = MagicMock(return_value=([], None))
runner.moderation_for_inputs = MagicMock(side_effect=ModerationError("blocked"))
runner.direct_output = MagicMock()
result = runner._run_input_moderation(
app_record=app_record,
application_generate_entity=app_generate_entity,
queue_manager=queue_manager,
message=message,
)
assert result.stopped is True
runner.direct_output.assert_called_once()
class TestCompletionAppRunner:
def test_run_app_not_found(self, runner, mocker: MockerFixture):
app_config = _build_app_config()
app_generate_entity = _build_generate_entity(app_config)
def test_workflow_runner_uses_low_image_detail_default() -> None:
runner = CompletionWorkflowRunner()
app_generate_entity = _build_generate_entity(_build_app_config(), file_upload_config=None)
with patched_create_session(return_value=None):
with pytest.raises(ValueError):
runner.run(app_generate_entity, MagicMock(), MagicMock(), MagicMock())
def test_run_moderation_error_outputs_direct(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
app_config = _build_app_config()
app_generate_entity = _build_generate_entity(app_config)
runner.organize_prompt_messages = MagicMock(return_value=([], None))
runner.moderation_for_inputs = MagicMock(side_effect=ModerationError("blocked"))
runner.direct_output = MagicMock()
runner._handle_invoke_result = MagicMock()
with patched_create_session(return_value=app_record):
runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock())
runner.direct_output.assert_called_once()
runner._handle_invoke_result.assert_not_called()
def test_run_hosting_moderation_stops(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
app_config = _build_app_config()
app_generate_entity = _build_generate_entity(app_config)
runner.organize_prompt_messages = MagicMock(return_value=([], None))
runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query"))
runner.check_hosting_moderation = MagicMock(return_value=True)
runner._handle_invoke_result = MagicMock()
with patched_create_session(return_value=app_record):
runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock())
runner._handle_invoke_result.assert_not_called()
def test_run_dataset_and_external_tools_flow(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
retrieve_config = MagicMock(query_variable="qvar")
dataset_config = MagicMock(dataset_ids=["ds"], retrieve_config=retrieve_config)
additional_features = MagicMock(show_retrieve_source=True)
app_config = _build_app_config(
dataset=dataset_config,
external_tools=["tool"],
additional_features=additional_features,
)
file_upload_config = MagicMock()
file_upload_config.image_config.detail = ImagePromptMessageContent.DETAIL.HIGH
app_generate_entity = _build_generate_entity(app_config, file_upload_config=file_upload_config)
runner.organize_prompt_messages = MagicMock(side_effect=[(["pm1"], ["stop"]), (["pm2"], ["stop"])])
runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query"))
runner.fill_in_inputs_from_external_data_tools = MagicMock(return_value=app_generate_entity.inputs)
runner.check_hosting_moderation = MagicMock(return_value=False)
runner.recalc_llm_max_tokens = MagicMock()
runner._handle_invoke_result = MagicMock()
dataset_retrieval = MagicMock()
dataset_retrieval.retrieve.return_value = ("ctx", ["file1"])
mocker.patch.object(module, "DatasetRetrieval", return_value=dataset_retrieval)
model_instance = MagicMock()
model_instance.invoke_llm.return_value = "invoke_result"
mocker.patch.object(module, "ModelInstance", return_value=model_instance)
with patched_create_session(return_value=app_record):
runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id="tenant"), MagicMock())
dataset_retrieval.retrieve.assert_called_once()
assert dataset_retrieval.retrieve.call_args.kwargs["query"] == "query_from_input"
runner._handle_invoke_result.assert_called_once()
def test_run_closes_explicit_session_before_stream_consumption(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
app_config = _build_app_config()
app_generate_entity = _build_generate_entity(app_config)
queue_manager = MagicMock()
events = []
session = MagicMock()
session.commit.side_effect = lambda: events.append("commit")
session.close.side_effect = lambda: events.append("close")
runner.organize_prompt_messages = MagicMock(return_value=([], None))
runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query"))
runner.check_hosting_moderation = MagicMock(return_value=False)
runner.recalc_llm_max_tokens = MagicMock()
runner._handle_invoke_result = MagicMock(side_effect=lambda invoke_result, **kwargs: list(invoke_result))
model_instance = MagicMock()
def invoke_stream():
events.append("first-chunk")
yield "chunk"
def invoke_llm(**kwargs):
events.append("invoke")
return invoke_stream()
model_instance.invoke_llm.side_effect = invoke_llm
mocker.patch.object(module, "ModelInstance", return_value=model_instance)
with patched_create_session(return_value=app_record):
runner.run(app_generate_entity, queue_manager, MagicMock(id="msg"), session)
assert events == ["commit", "close", "invoke", "first-chunk"]
runner._handle_invoke_result.assert_called_once_with(
invoke_result=ANY,
queue_manager=queue_manager,
stream=True,
message_id="msg",
user_id="user",
tenant_id="tenant",
)
def test_run_uses_low_image_detail_default(self, runner, mocker: MockerFixture):
app_record = MagicMock(id="app1", tenant_id="tenant")
app_config = _build_app_config()
app_generate_entity = _build_generate_entity(app_config, file_upload_config=None)
runner.organize_prompt_messages = MagicMock(return_value=([], None))
runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query"))
runner.check_hosting_moderation = MagicMock(return_value=True)
with patched_create_session(return_value=app_record):
runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock())
assert (
runner.organize_prompt_messages.call_args.kwargs["image_detail_config"]
== ImagePromptMessageContent.DETAIL.LOW
)
assert runner._resolve_image_detail_config(app_generate_entity) == ImagePromptMessageContent.DETAIL.LOW
@@ -376,18 +376,26 @@ class TestCompletionAppGenerator:
create_session.return_value = session_context
mocker.patch.object(module.db, "session")
mocker.patch.object(generator, "_get_message", return_value=MagicMock())
message = MagicMock()
mocker.patch.object(generator, "_get_message", return_value=message)
runner_instance = MagicMock()
runner_instance.run.side_effect = error
mocker.patch.object(module, "CompletionAppRunner", return_value=runner_instance)
mocker.patch.object(module, "CompletionWorkflowRunner", return_value=runner_instance)
queue_manager = MagicMock()
application_generate_entity = MagicMock()
generator._generate_worker(
flask_app=flask_app,
application_generate_entity=MagicMock(),
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
message_id="msg",
)
runner_instance.run.assert_called_once_with(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
message=message,
session=session,
)
assert queue_manager.publish_error.called is should_publish
@@ -0,0 +1,70 @@
import json
from types import SimpleNamespace
from unittest.mock import MagicMock
from core.app.apps.completion.runtime_workflow_builder import build_runtime_completion_workflow
from graphon.nodes import BuiltinNodeTypes
from models.model import AppMode
def test_builder_returns_runtime_graph_without_workflow_record() -> None:
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
app_config = MagicMock()
workflow_converter = MagicMock(
build_graph_from_app_config=MagicMock(return_value=({"nodes": [{"id": "start"}], "edges": []}, {}))
)
session = MagicMock()
result = build_runtime_completion_workflow(
app_model=app_model,
app_config=app_config,
session=session,
workflow_converter=workflow_converter,
)
assert result.workflow_id.startswith("completion-runtime-")
assert result.root_node_id == "start"
assert result.graph_dict == {"nodes": [{"id": "start"}], "edges": []}
workflow_converter.build_graph_from_app_config.assert_called_once_with(
app_model=app_model,
app_config=app_config,
target_app_mode=AppMode.WORKFLOW,
session=session,
)
def test_builder_routes_api_based_variable_query_to_runtime_sys_query() -> None:
app_model = SimpleNamespace(mode=AppMode.COMPLETION)
app_config = MagicMock()
request_body = {"params": {"query": ""}}
workflow_converter = MagicMock(
build_graph_from_app_config=MagicMock(
return_value=(
{
"nodes": [
{
"id": "http_request_1",
"data": {
"type": BuiltinNodeTypes.HTTP_REQUEST,
"body": {"type": "json", "data": json.dumps(request_body)},
},
}
],
"edges": [],
},
{},
)
)
)
session = MagicMock()
result = build_runtime_completion_workflow(
app_model=app_model,
app_config=app_config,
session=session,
workflow_converter=workflow_converter,
)
http_node = result.graph_dict["nodes"][0]
body = json.loads(http_node["data"]["body"]["data"])
assert body["params"]["query"] == "{{#sys.query#}}"
@@ -0,0 +1,172 @@
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import MagicMock
from core.app.apps.base_app_queue_manager import PublishFrom
from core.app.apps.completion.graph_event_adapter import CompletionGraphEventAdapter
from core.app.entities.queue_entities import (
QueueErrorEvent,
QueueLLMChunkEvent,
QueueMessageEndEvent,
QueueRetrieverResourcesEvent,
QueueStopEvent,
)
from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus
from graphon.graph_events import (
GraphRunAbortedEvent,
GraphRunFailedEvent,
GraphRunSucceededEvent,
NodeRunRetrieverResourceEvent,
NodeRunStreamChunkEvent,
NodeRunSucceededEvent,
)
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.model_runtime.entities.message_entities import UserPromptMessage
from graphon.node_events import NodeRunResult
def _adapter(*, show_retrieve_source: bool = True) -> tuple[CompletionGraphEventAdapter, MagicMock]:
queue_manager = MagicMock()
entity = SimpleNamespace(
model_conf=SimpleNamespace(model="model"),
app_config=SimpleNamespace(
additional_features=SimpleNamespace(show_retrieve_source=show_retrieve_source),
),
)
return (
CompletionGraphEventAdapter(application_generate_entity=entity, queue_manager=queue_manager),
queue_manager,
)
def test_stream_chunk_event_publishes_llm_chunk() -> None:
adapter, queue_manager = _adapter()
prompt_message = UserPromptMessage(content="prompt")
adapter.set_prompt_messages([prompt_message])
adapter.handle_event(
NodeRunStreamChunkEvent(
id="run",
node_id="llm",
node_type=BuiltinNodeTypes.LLM,
selector=["llm", "text"],
chunk="hello",
)
)
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueLLMChunkEvent)
assert event.chunk.delta.message.content == "hello"
assert event.chunk.prompt_messages == [prompt_message]
assert queue_manager.publish.call_args.args[1] == PublishFrom.APPLICATION_MANAGER
def test_stream_chunk_event_skips_final_empty_chunk() -> None:
adapter, queue_manager = _adapter()
adapter.handle_event(
NodeRunStreamChunkEvent(
id="run",
node_id="llm",
node_type=BuiltinNodeTypes.LLM,
selector=["llm", "text"],
chunk="",
is_final=True,
)
)
queue_manager.publish.assert_not_called()
def test_retriever_resource_event_publishes_legacy_retriever_resources() -> None:
adapter, queue_manager = _adapter()
adapter.handle_event(
NodeRunRetrieverResourceEvent(
id="run",
node_id="knowledge_retrieval",
node_type=BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL,
retriever_resources=[{"dataset_id": "dataset", "content": "hit"}],
context="hit",
)
)
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueRetrieverResourcesEvent)
assert event.retriever_resources[0].dataset_id == "dataset"
def test_retriever_resource_event_is_hidden_when_feature_is_disabled() -> None:
adapter, queue_manager = _adapter(show_retrieve_source=False)
adapter.handle_event(
NodeRunRetrieverResourceEvent(
id="run",
node_id="knowledge_retrieval",
node_type=BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL,
retriever_resources=[{"dataset_id": "dataset", "content": "hit"}],
context="hit",
)
)
queue_manager.publish.assert_not_called()
def test_llm_success_then_graph_success_publishes_message_end() -> None:
adapter, queue_manager = _adapter()
usage = LLMUsage.empty_usage()
prompt_message = UserPromptMessage(content="prompt")
adapter.set_prompt_messages([prompt_message])
adapter.handle_event(
NodeRunSucceededEvent(
id="run",
node_id="llm",
node_type=BuiltinNodeTypes.LLM,
start_at=datetime.now(UTC),
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
outputs={"text": "final"},
llm_usage=usage,
),
)
)
adapter.handle_event(GraphRunSucceededEvent(outputs={"result": "final"}))
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueMessageEndEvent)
assert event.llm_result is not None
assert event.llm_result.message.content == "final"
assert event.llm_result.prompt_messages == [prompt_message]
assert event.llm_result.usage is usage
def test_graph_success_uses_outputs_result_when_llm_success_was_not_seen() -> None:
adapter, queue_manager = _adapter()
adapter.handle_event(GraphRunSucceededEvent(outputs={"result": "final from graph"}))
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueMessageEndEvent)
assert event.llm_result is not None
assert event.llm_result.message.content == "final from graph"
def test_failed_graph_publishes_error() -> None:
adapter, queue_manager = _adapter()
adapter.handle_event(GraphRunFailedEvent(error="boom"))
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueErrorEvent)
assert str(event.error) == "boom"
def test_user_abort_publishes_legacy_stop() -> None:
adapter, queue_manager = _adapter()
adapter.handle_event(GraphRunAbortedEvent(reason="Stopped by user."))
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueStopEvent)
assert event.stopped_by == QueueStopEvent.StopBy.USER_MANUAL
@@ -0,0 +1,278 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from core.app.apps.completion.workflow_runner import CompletionWorkflowRunner, ModeratedCompletionInputs
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.moderation.base import ModerationError
from core.workflow.node_runtime import DIFY_BEFORE_LLM_INVOKE_KEY
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
from models.model import AppMode
def _entity() -> SimpleNamespace:
return SimpleNamespace(
app_config=SimpleNamespace(app_id="app", tenant_id="tenant", prompt_template=MagicMock()),
model_conf=SimpleNamespace(model="model"),
user_id="user",
invoke_from=InvokeFrom.SERVICE_API,
task_id="task",
call_depth=2,
inputs={"name": "Ada"},
query="question",
files=[],
file_upload_config=None,
extras={"trace_session_id": "trace"},
stream=True,
trace_manager=None,
)
def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
from core.app.apps.completion import workflow_runner as module
app = SimpleNamespace(id="app", tenant_id="tenant", mode=AppMode.COMPLETION)
entity = _entity()
message = SimpleNamespace(id="message", conversation_id="conv")
queue_manager = SimpleNamespace(graph_runtime_state=None)
runtime_workflow = SimpleNamespace(
workflow_id="completion-runtime-1",
root_node_id="start",
graph_dict={"nodes": [{"id": "start", "data": {"type": "start"}}], "edges": []},
)
build_runtime_workflow = MagicMock(return_value=runtime_workflow)
graph = MagicMock()
adapter = MagicMock()
workflow_entry = MagicMock()
session = MagicMock()
lifecycle_events: list[str] = []
session.commit.side_effect = lambda: lifecycle_events.append("commit")
session.close.side_effect = lambda: lifecycle_events.append("close")
def run_workflow():
lifecycle_events.append("run")
yield "event"
workflow_entry.run.return_value = run_workflow()
init_graph = MagicMock(return_value=graph)
workflow_entry_class = MagicMock(return_value=workflow_entry)
adapter_class = MagicMock(return_value=adapter)
build_system_variables = MagicMock(return_value=["sys"])
build_bootstrap_variables = MagicMock(return_value=["boot"])
add_variables_to_pool = MagicMock()
add_node_inputs_to_pool = MagicMock()
monkeypatch.setattr(module, "init_graph", init_graph)
monkeypatch.setattr(module, "build_runtime_completion_workflow", build_runtime_workflow)
monkeypatch.setattr(module, "WorkflowEntry", workflow_entry_class)
monkeypatch.setattr(module, "CompletionGraphEventAdapter", adapter_class)
monkeypatch.setattr(module, "RedisChannel", MagicMock())
monkeypatch.setattr(module, "redis_client", MagicMock())
monkeypatch.setattr(module, "build_system_variables", build_system_variables)
monkeypatch.setattr(module, "build_bootstrap_variables", build_bootstrap_variables)
monkeypatch.setattr(module, "add_variables_to_pool", add_variables_to_pool)
monkeypatch.setattr(module, "add_node_inputs_to_pool", add_node_inputs_to_pool)
runner = CompletionWorkflowRunner()
monkeypatch.setattr(runner, "_get_app", MagicMock(return_value=app))
before_llm_invoke_hook = MagicMock()
build_before_llm_invoke_hook = MagicMock(return_value=before_llm_invoke_hook)
monkeypatch.setattr(runner, "_build_before_llm_invoke_hook", build_before_llm_invoke_hook)
monkeypatch.setattr(
runner,
"_run_input_moderation",
MagicMock(return_value=ModeratedCompletionInputs(stopped=False, inputs={"name": "Grace"}, query="moderated")),
)
runner.run(
application_generate_entity=entity,
queue_manager=queue_manager,
message=message,
session=session,
)
build_runtime_workflow.assert_called_once_with(
app_model=app,
app_config=entity.app_config,
session=session,
)
assert lifecycle_events == ["commit", "close", "run"]
add_node_inputs_to_pool.assert_called_once()
assert add_node_inputs_to_pool.call_args.kwargs["node_id"] == "start"
assert add_node_inputs_to_pool.call_args.kwargs["inputs"] == {"name": "Grace"}
build_system_variables.assert_called_once()
assert build_system_variables.call_args.kwargs["query"] == "moderated"
assert build_system_variables.call_args.kwargs["conversation_id"] == "conv"
workflow_entry_class.assert_called_once()
assert workflow_entry_class.call_args.kwargs["workflow_id"] == "completion-runtime-1"
assert workflow_entry_class.call_args.kwargs["user_from"] == UserFrom.END_USER
assert workflow_entry_class.call_args.kwargs["call_depth"] == 2
build_before_llm_invoke_hook.assert_called_once_with(
application_generate_entity=entity,
queue_manager=queue_manager,
adapter=adapter,
)
init_graph.assert_called_once()
assert init_graph.call_args.kwargs["app_id"] == "app"
assert init_graph.call_args.kwargs["graph_config"] == runtime_workflow.graph_dict
assert init_graph.call_args.kwargs["root_node_id"] == "start"
assert init_graph.call_args.kwargs["call_depth"] == 2
assert init_graph.call_args.kwargs["extra_context"][DIFY_BEFORE_LLM_INVOKE_KEY] is before_llm_invoke_hook
workflow_entry.graph_engine.layer.assert_not_called()
adapter_class.assert_called_once_with(application_generate_entity=entity, queue_manager=queue_manager)
adapter.handle_event.assert_called_once_with("event")
def test_runner_returns_when_input_moderation_stops(monkeypatch) -> None:
app = SimpleNamespace(id="app", tenant_id="tenant", mode=AppMode.COMPLETION)
entity = _entity()
build_runtime_workflow = MagicMock()
runner = CompletionWorkflowRunner()
monkeypatch.setattr(
"core.app.apps.completion.workflow_runner.build_runtime_completion_workflow",
build_runtime_workflow,
)
monkeypatch.setattr(runner, "_get_app", MagicMock(return_value=app))
monkeypatch.setattr(
runner,
"_run_input_moderation",
MagicMock(return_value=ModeratedCompletionInputs(stopped=True, inputs={}, query="")),
)
runner.run(
application_generate_entity=entity,
queue_manager=MagicMock(),
message=SimpleNamespace(id="message"),
session=MagicMock(),
)
build_runtime_workflow.assert_not_called()
def test_runner_get_app_raises_when_record_is_missing() -> None:
runner = CompletionWorkflowRunner()
session = MagicMock()
session.scalar.return_value = None
with pytest.raises(ValueError, match="App not found"):
runner._get_app(app_id="missing-app", tenant_id="tenant", session=session)
def test_runner_get_app_returns_record() -> None:
app = SimpleNamespace(id="app")
runner = CompletionWorkflowRunner()
session = MagicMock()
session.scalar.return_value = app
assert runner._get_app(app_id="app", tenant_id="tenant", session=session) is app
def test_runner_direct_outputs_on_input_moderation() -> None:
runner = CompletionWorkflowRunner()
app_record = SimpleNamespace(id="app", tenant_id="tenant")
entity = _entity()
message = SimpleNamespace(id="message")
queue_manager = MagicMock()
runner.organize_prompt_messages = MagicMock(return_value=(["prompt"], None))
runner.moderation_for_inputs = MagicMock(side_effect=ModerationError("blocked"))
runner.direct_output = MagicMock()
result = runner._run_input_moderation(
app_record=app_record,
application_generate_entity=entity,
queue_manager=queue_manager,
message=message,
)
assert result.stopped is True
assert result.inputs == {"name": "Ada"}
assert result.query == "question"
runner.direct_output.assert_called_once()
def test_runner_returns_moderated_inputs_when_input_moderation_passes() -> None:
runner = CompletionWorkflowRunner()
app_record = SimpleNamespace(id="app", tenant_id="tenant")
entity = _entity()
message = SimpleNamespace(id="message")
runner.organize_prompt_messages = MagicMock(return_value=(["prompt"], None))
runner.moderation_for_inputs = MagicMock(return_value=(None, {"name": "Grace"}, "moderated query"))
result = runner._run_input_moderation(
app_record=app_record,
application_generate_entity=entity,
queue_manager=MagicMock(),
message=message,
)
assert result == ModeratedCompletionInputs(stopped=False, inputs={"name": "Grace"}, query="moderated query")
def test_runner_before_llm_invoke_hook_captures_and_moderates_final_prompt() -> None:
runner = CompletionWorkflowRunner()
entity = _entity()
queue_manager = MagicMock()
adapter = MagicMock()
runner.check_hosting_moderation = MagicMock(return_value=True)
hook = runner._build_before_llm_invoke_hook(
application_generate_entity=entity,
queue_manager=queue_manager,
adapter=adapter,
)
with pytest.raises(GenerateTaskStoppedError):
hook(["final prompt"], {"max_tokens": 128})
adapter.set_prompt_messages.assert_called_once_with(["final prompt"])
runner.check_hosting_moderation.assert_called_once_with(
application_generate_entity=entity,
queue_manager=queue_manager,
prompt_messages=["final prompt"],
)
def test_runner_before_llm_invoke_hook_recalculates_graph_model_parameters() -> None:
runner = CompletionWorkflowRunner()
entity = _entity()
queue_manager = MagicMock()
adapter = MagicMock()
runner.check_hosting_moderation = MagicMock(return_value=False)
def recalc(*, model_parameters, **kwargs) -> None:
model_parameters["max_tokens"] = 64
runner.recalc_llm_max_tokens = MagicMock(side_effect=recalc)
hook = runner._build_before_llm_invoke_hook(
application_generate_entity=entity,
queue_manager=queue_manager,
adapter=adapter,
)
result = hook(["final prompt"], {"max_tokens": 128, "temperature": 0.2})
assert result == {"max_tokens": 64, "temperature": 0.2}
runner.recalc_llm_max_tokens.assert_called_once_with(
model_config=entity.model_conf,
prompt_messages=["final prompt"],
model_parameters=result,
)
def test_runner_resolves_account_user_from() -> None:
entity = _entity()
entity.invoke_from = InvokeFrom.EXPLORE
assert CompletionWorkflowRunner._resolve_user_from(entity) == UserFrom.ACCOUNT
def test_runner_resolves_configured_image_detail() -> None:
entity = _entity()
entity.file_upload_config = SimpleNamespace(
image_config=SimpleNamespace(detail=ImagePromptMessageContent.DETAIL.HIGH),
)
assert CompletionWorkflowRunner._resolve_image_detail_config(entity) == ImagePromptMessageContent.DETAIL.HIGH
@@ -133,6 +133,33 @@ class TestAppRunner:
assert runner.recalc_llm_max_tokens(model_config, prompt_messages=[]) == -1
def test_recalc_llm_max_tokens_can_update_runtime_parameters(self, monkeypatch: pytest.MonkeyPatch):
runner = AppRunner()
model_schema = AIModelEntity.model_construct(
model_properties={ModelPropertyKey.CONTEXT_SIZE: 100},
parameter_rules=[_DummyParameterRule("max_tokens")],
)
model_config = ModelConfigWithCredentialsEntity.model_construct(
provider_model_bundle=object(),
model="mock",
model_schema=model_schema,
parameters={"max_tokens": 30},
)
runtime_parameters = {"max_tokens": 40}
monkeypatch.setattr(
"core.app.apps.base_app_runner.ModelInstance",
lambda provider_model_bundle, model: _TokenCountingModel(80),
)
runner.recalc_llm_max_tokens(
model_config,
prompt_messages=[AssistantPromptMessage(content="hi")],
model_parameters=runtime_parameters,
)
assert runtime_parameters["max_tokens"] == 20
assert model_config.parameters["max_tokens"] == 30
def test_direct_output_streaming_publishes_chunks_and_end(self):
runner = AppRunner()
queue = _queue_manager()
@@ -5,7 +5,7 @@ from types import SimpleNamespace
import pytest
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner, init_graph
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, InvokeFrom, UserFrom
from core.app.entities.queue_entities import (
QueueAgentLogEvent,
@@ -119,6 +119,40 @@ class TestWorkflowBasedAppRunner:
assert captured["run_context"][DIFY_RUN_CONTEXT_KEY].trace_session_id == "session-1"
def test_init_graph_accepts_call_depth_and_extra_context(self, monkeypatch: pytest.MonkeyPatch):
runtime_state = GraphRuntimeState(
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
hook = object()
captured = {}
def fake_from_graph_init_context(**kwargs):
graph_init_context = kwargs["graph_init_context"]
captured["run_context"] = graph_init_context.run_context
captured["call_depth"] = graph_init_context.call_depth
return SimpleNamespace()
monkeypatch.setattr(
"core.app.apps.workflow_app_runner.DifyNodeFactory.from_graph_init_context",
fake_from_graph_init_context,
)
monkeypatch.setattr("core.app.apps.workflow_app_runner.Graph.init", lambda **_kwargs: SimpleNamespace())
init_graph(
app_id="app",
graph_config={"nodes": [], "edges": []},
graph_runtime_state=runtime_state,
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
root_node_id="root",
call_depth=2,
extra_context={"hook": hook},
)
assert captured["call_depth"] == 2
assert captured["run_context"]["hook"] is hook
def test_prepare_single_node_execution_requires_run(self):
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
@@ -461,6 +461,7 @@ class TestDifyNodeFactoryCreateNode:
def factory(self):
factory = object.__new__(node_factory.DifyNodeFactory)
factory.graph_init_params = sentinel.graph_init_params
factory.graph_init_params.run_context = {}
factory.graph_runtime_state = SimpleNamespace(variable_pool=MagicMock())
factory._dify_context = SimpleNamespace(
tenant_id="tenant-id",
@@ -702,6 +703,44 @@ class TestDifyNodeFactoryCreateNode:
node_data=node_data,
model_instance=sentinel.model_instance,
request_metadata={"app_id": "app-id"},
before_invoke=None,
)
assert kwargs["model_instance"] is wrapped_model_instance
def test_build_llm_compatible_node_init_kwargs_passes_before_llm_hook(self, factory):
before_llm_invoke = MagicMock()
factory.graph_init_params.run_context[node_factory.DIFY_BEFORE_LLM_INVOKE_KEY] = before_llm_invoke
node_data = LLMNodeData.model_validate(
{
"type": BuiltinNodeTypes.LLM,
"title": "LLM",
"model": {"provider": "provider", "name": "model", "mode": "chat", "completion_params": {}},
"prompt_template": [{"role": "system", "text": "x"}],
"context": {"enabled": False, "variable_selector": []},
"vision": {"enabled": False},
}
)
wrapped_model_instance = sentinel.wrapped_model_instance
factory._build_model_instance_for_llm_node = MagicMock(return_value=sentinel.model_instance)
factory._build_memory_for_llm_node = MagicMock(return_value=sentinel.memory)
with patch.object(factory, "_wrap_model_instance_for_node", return_value=wrapped_model_instance) as wrap_model:
kwargs = factory._build_llm_compatible_node_init_kwargs(
node_class=sentinel.node_class,
node_data=node_data,
wrap_model_instance=True,
include_http_client=False,
include_llm_file_saver=False,
include_prompt_message_serializer=False,
include_retriever_attachment_loader=False,
include_jinja2_template_renderer=False,
)
wrap_model.assert_called_once_with(
node_data=node_data,
model_instance=sentinel.model_instance,
request_metadata={"app_id": "app-id"},
before_invoke=before_llm_invoke,
)
assert kwargs["model_instance"] is wrapped_model_instance
@@ -228,7 +228,12 @@ def test_dify_prepared_llm_wraps_model_instance_calls() -> None:
model_schema = _build_model_schema()
model_instance = _ModelInstanceStub(model_schema=model_schema)
model_type_instance = model_instance.model_type_instance
prepared = DifyPreparedLLM(model_instance, request_metadata={"app_id": "app-id"})
before_invoke = Mock(return_value={"temperature": 0.05})
prepared = DifyPreparedLLM(
model_instance,
request_metadata={"app_id": "app-id"},
before_invoke=before_invoke,
)
assert prepared.provider == "langgenius/openai/openai"
assert prepared.model_name == "gpt-4o-mini"
@@ -248,9 +253,10 @@ def test_dify_prepared_llm_wraps_model_instance_calls() -> None:
)
model_type_instance.get_model_schema.assert_called_once_with("gpt-4o-mini", {"api_key": "secret"})
before_invoke.assert_called_once_with([], {"temperature": 0.1})
model_instance.invoke_llm.assert_called_once_with(
prompt_messages=[],
model_parameters={"temperature": 0.1},
model_parameters={"temperature": 0.05},
tools=[],
stop=[],
stream=False,
@@ -269,7 +275,8 @@ def test_dify_prepared_llm_requires_model_schema() -> None:
def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytest.MonkeyPatch) -> None:
model_instance = _ModelInstanceStub(model_schema=_build_model_schema())
prepared = DifyPreparedLLM(model_instance)
before_invoke = Mock(return_value={"temperature": 0.15})
prepared = DifyPreparedLLM(model_instance, before_invoke=before_invoke)
invoke_structured = MagicMock(return_value=sentinel.structured)
monkeypatch.setattr(node_runtime, "invoke_llm_with_structured_output", invoke_structured)
@@ -282,13 +289,14 @@ def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytes
)
assert result is sentinel.structured
before_invoke.assert_called_once_with([], {"temperature": 0.2})
invoke_structured.assert_called_once_with(
provider="langgenius/openai/openai",
model_schema=prepared.get_model_schema(),
model_instance=model_instance,
prompt_messages=[],
json_schema={"type": "object"},
model_parameters={"temperature": 0.2},
model_parameters={"temperature": 0.15},
stop=["done"],
stream=True,
)
@@ -320,7 +328,8 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None:
model_runtime=plugin_runtime,
)
prepared = DifyPreparedPollingLLM(model_instance)
before_invoke = Mock(return_value={"temperature": 0.05})
prepared = DifyPreparedPollingLLM(model_instance, before_invoke=before_invoke)
assert isinstance(prepared, LLMPollingCapableProtocol)
assert (
@@ -333,6 +342,7 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None:
)
== polling_result
)
before_invoke.assert_called_once_with([], {"temperature": 0.1})
assert (
prepared.check_llm_polling(
plugin_state={"task_id": "poll-1"},
@@ -344,7 +354,7 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None:
model="gpt-4o-mini",
credentials={"api_key": "secret"},
prompt_messages=[],
model_parameters={"temperature": 0.1},
model_parameters={"temperature": 0.05},
tools=[],
stop=("END",),
json_schema={"type": "object"},
@@ -14,7 +14,7 @@ class TestAppTaskService:
("app_mode", "should_call_graph_engine"),
[
(AppMode.CHAT, False),
(AppMode.COMPLETION, False),
(AppMode.COMPLETION, True),
(AppMode.AGENT_CHAT, False),
(AppMode.AGENT, False),
(AppMode.CHANNEL, False),
@@ -219,7 +219,13 @@ def test__convert_to_knowledge_retrieval_node_for_workflow_app() -> None:
def test__convert_to_llm_node_for_chatbot_simple_chat_model(default_variables: list[VariableEntity]) -> None:
workflow_converter = WorkflowConverter()
graph = {"nodes": [workflow_converter._convert_to_start_node(default_variables)], "edges": []}
model_config = ModelConfigEntity(provider="openai", model="gpt-4", mode=LLMMode.CHAT.value, parameters={}, stop=[])
model_config = ModelConfigEntity(
provider="openai",
model="gpt-4",
mode=LLMMode.CHAT.value,
parameters={"temperature": 0.2},
stop=["END"],
)
prompt_template = PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.SIMPLE,
simple_prompt_template="You are a helper for {{text_input}} and {{paragraph}}",
@@ -237,6 +243,8 @@ def test__convert_to_llm_node_for_chatbot_simple_chat_model(default_variables: l
assert node["data"]["memory"] is not None
assert node["data"]["prompt_template"][0]["role"] == "user"
assert "{{#start.text_input#}}" in node["data"]["prompt_template"][0]["text"]
assert node["data"]["model"]["completion_params"] == {"temperature": 0.2, "stop": ["END"]}
assert model_config.parameters == {"temperature": 0.2}
def test__convert_to_llm_node_for_chatbot_simple_chat_model_with_empty_template(
@@ -599,6 +607,94 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en
assert set(features.keys()) == {"text_to_speech", "file_upload", "sensitive_word_avoidance"}
def test_build_graph_from_app_config_for_completion_does_not_create_workflow(
converter: WorkflowConverter,
) -> None:
app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.COMPLETION)
app_config = SimpleNamespace(
variables=[],
external_data_variables=[],
dataset=None,
model=_build_model_config(mode=LLMMode.CHAT),
prompt_template=PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.SIMPLE,
simple_prompt_template="Hello",
),
additional_features=None,
app_model_config_dict={
"text_to_speech": {"enabled": False},
"file_upload": {"enabled": False},
"sensitive_word_avoidance": {"enabled": False},
},
)
db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock())
graph, features = converter.build_graph_from_app_config(
app_model=app_model,
app_config=app_config,
target_app_mode=AppMode.WORKFLOW,
session=db_session,
)
assert [node["id"] for node in graph["nodes"]] == ["start", "llm", "end"]
assert features == {
"text_to_speech": {"enabled": False},
"file_upload": {"enabled": False},
"sensitive_word_avoidance": {"enabled": False},
}
db_session.add.assert_not_called()
db_session.commit.assert_not_called()
def test_build_graph_from_app_config_preserves_api_based_variable_nodes(
converter: WorkflowConverter,
monkeypatch: pytest.MonkeyPatch,
) -> None:
app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.COMPLETION)
app_config = SimpleNamespace(
variables=[VariableEntity(variable="city", label="City", type=VariableEntityType.TEXT_INPUT)],
external_data_variables=[
ExternalDataVariableEntity(
variable="weather",
type="api",
config={"api_based_extension_id": "api_based_extension_id"},
)
],
dataset=None,
model=_build_model_config(mode=LLMMode.CHAT),
prompt_template=PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.SIMPLE,
simple_prompt_template="Weather: {{weather}}",
),
additional_features=None,
app_model_config_dict={},
)
extension = SimpleNamespace(
name="Weather API",
api_endpoint="https://example.com/weather",
api_key="encrypted-token",
)
monkeypatch.setattr(converter, "_get_api_based_extension", MagicMock(return_value=extension))
monkeypatch.setattr(converter_module.encrypter, "decrypt_token", MagicMock(return_value="plain-token"))
graph, _ = converter.build_graph_from_app_config(
app_model=app_model,
app_config=app_config,
target_app_mode=AppMode.WORKFLOW,
session=MagicMock(),
)
assert [node["data"]["type"] for node in graph["nodes"]] == [
BuiltinNodeTypes.START,
BuiltinNodeTypes.HTTP_REQUEST,
BuiltinNodeTypes.CODE,
BuiltinNodeTypes.LLM,
BuiltinNodeTypes.END,
]
llm_node = next(node for node in graph["nodes"] if node["id"] == "llm")
assert "{{#code_1.result#}}" in llm_node["data"]["prompt_template"][0]["text"]
def test_convert_to_app_config_should_route_to_correct_manager(
converter: WorkflowConverter,
monkeypatch: pytest.MonkeyPatch,