fix(completion): preserve legacy workflow compatibility

This commit is contained in:
-LAN-
2026-07-28 13:04:11 +08:00
parent e2c769c6cc
commit 62518a5e55
19 changed files with 277 additions and 329 deletions
+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,
@@ -1,5 +1,4 @@
from collections.abc import Mapping
from typing import Any, cast
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
@@ -10,34 +9,30 @@ from core.app.entities.queue_entities import (
QueueRetrieverResourcesEvent,
QueueStopEvent,
)
from core.prompt.utils.prompt_message_util import SavedPrompt
from core.rag.entities import RetrievalSourceMetadata
from graphon.enums import BuiltinNodeTypes
from graphon.graph_events import (
GraphEngineEvent,
GraphRunAbortedEvent,
GraphRunFailedEvent,
GraphRunSucceededEvent,
NodeRunExceptionEvent,
NodeRunFailedEvent,
NodeRunRetrieverResourceEvent,
NodeRunStreamChunkEvent,
NodeRunSucceededEvent,
)
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, PromptMessage
_LLM_TEXT_SELECTOR_PREFIX = ("llm", "text")
class CompletionGraphEventAdapter:
"""Translate runtime workflow events into legacy Completion queue events."""
"""Translate one runtime graph run into legacy Completion queue events."""
_application_generate_entity: CompletionAppGenerateEntity
_queue_manager: AppQueueManager
_answer: str
_usage: LLMUsage
_saved_prompt: list[SavedPrompt]
_prompt_messages: list[PromptMessage]
_chunk_index: int
def __init__(
@@ -50,9 +45,13 @@ class CompletionGraphEventAdapter:
self._queue_manager = queue_manager
self._answer = ""
self._usage = LLMUsage.empty_usage()
self._saved_prompt = []
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():
@@ -61,8 +60,6 @@ class CompletionGraphEventAdapter:
self._handle_retriever_resource(event)
case NodeRunSucceededEvent():
self._handle_node_succeeded(event)
case NodeRunFailedEvent() | NodeRunExceptionEvent():
self._publish_error(event.error or event.node_run_result.error or "Node failed")
case GraphRunSucceededEvent():
self._publish_message_end(event.outputs)
case GraphRunFailedEvent():
@@ -86,7 +83,7 @@ class CompletionGraphEventAdapter:
QueueLLMChunkEvent(
chunk=LLMResultChunk(
model=self._application_generate_entity.model_conf.model,
prompt_messages=[],
prompt_messages=self._prompt_messages,
delta=LLMResultChunkDelta(
index=self._chunk_index,
message=AssistantPromptMessage(content=event.chunk),
@@ -98,6 +95,10 @@ class CompletionGraphEventAdapter:
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=[
@@ -110,7 +111,7 @@ class CompletionGraphEventAdapter:
)
def _handle_node_succeeded(self, event: NodeRunSucceededEvent) -> None:
if event.node_type != BuiltinNodeTypes.LLM and event.node_id != "llm":
if event.node_id != "llm":
return
result = event.node_run_result
@@ -119,10 +120,6 @@ class CompletionGraphEventAdapter:
self._answer = text
self._usage = result.llm_usage
prompts = result.process_data.get("prompts")
if isinstance(prompts, list):
self._saved_prompt = cast(list[SavedPrompt], prompts)
def _publish_message_end(self, outputs: Mapping[str, object]) -> None:
result = outputs.get("result")
if isinstance(result, str) and not self._answer:
@@ -132,16 +129,15 @@ class CompletionGraphEventAdapter:
QueueMessageEndEvent(
llm_result=LLMResult(
model=self._application_generate_entity.model_conf.model,
prompt_messages=[],
prompt_messages=self._prompt_messages,
message=AssistantPromptMessage(content=self._answer),
usage=self._usage,
),
saved_prompt=self._saved_prompt,
)
),
PublishFrom.APPLICATION_MANAGER,
)
def _publish_error(self, error: Any) -> None:
def _publish_error(self, error: object) -> None:
self._queue_manager.publish(
QueueErrorEvent(error=ValueError(str(error))),
PublishFrom.APPLICATION_MANAGER,
@@ -18,56 +18,52 @@ class RuntimeCompletionWorkflow:
graph_dict: WorkflowGraph
class RuntimeCompletionWorkflowBuilder:
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 __init__(self, workflow_converter: WorkflowConverter | None = None) -> None:
self._workflow_converter = workflow_converter or WorkflowConverter()
def build(
self,
*,
app_model: App,
app_config: CompletionAppConfig,
session: Session,
) -> RuntimeCompletionWorkflow:
graph, _ = self._workflow_converter.build_graph_from_app_config(
app_model=app_model,
app_config=app_config,
target_app_mode=AppMode.WORKFLOW,
session=session,
)
self._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
@staticmethod
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
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
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
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 = body_data.get("params")
if not isinstance(params, dict) or params.get("query") != "":
continue
params["query"] = "{{#sys.query#}}"
body["data"] = json.dumps(body_data)
params["query"] = "{{#sys.query#}}"
body["data"] = json.dumps(body_data)
+32 -47
View File
@@ -10,23 +10,20 @@ 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 RuntimeCompletionWorkflowBuilder
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.entities import DEFAULT_PLUGIN_ID
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_hosting_provider import hosting_configuration
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
from models.provider import ProviderType
@dataclass(frozen=True, slots=True)
@@ -37,12 +34,7 @@ class ModeratedCompletionInputs:
class CompletionWorkflowRunner(AppRunner):
"""Run Completion through a transient WorkflowEntry graph."""
_runtime_workflow_builder: RuntimeCompletionWorkflowBuilder
def __init__(self, runtime_workflow_builder: RuntimeCompletionWorkflowBuilder | None = None) -> None:
self._runtime_workflow_builder = runtime_workflow_builder or RuntimeCompletionWorkflowBuilder()
"""Run a transient WorkflowEntry graph while the legacy task pipeline owns persistence."""
def run(
self,
@@ -52,7 +44,7 @@ class CompletionWorkflowRunner(AppRunner):
session: Session,
) -> None:
app_config = cast(CompletionAppConfig, application_generate_entity.app_config)
app_record = self._get_app(app_config.app_id, session=session)
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,
@@ -63,7 +55,7 @@ class CompletionWorkflowRunner(AppRunner):
if moderation_result.stopped:
return
runtime_workflow = self._runtime_workflow_builder.build(
runtime_workflow = build_runtime_completion_workflow(
app_model=app_record,
app_config=app_config,
session=session,
@@ -78,12 +70,17 @@ class CompletionWorkflowRunner(AppRunner):
)
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
user_from = self._resolve_user_from(application_generate_entity)
extra_context: dict[str, Any] = {}
if self._should_check_hosting_moderation(application_generate_entity):
extra_context[DIFY_BEFORE_LLM_INVOKE_KEY] = self._build_hosting_moderation_hook(
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,
@@ -116,17 +113,14 @@ class CompletionWorkflowRunner(AppRunner):
graph_runtime_state=graph_runtime_state,
command_channel=command_channel,
)
adapter = CompletionGraphEventAdapter(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
)
# 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, *, session: Session) -> App:
app_record = session.scalar(select(App).where(App.id == app_id))
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
@@ -175,13 +169,18 @@ class CompletionWorkflowRunner(AppRunner):
return ModeratedCompletionInputs(stopped=False, inputs=inputs, query=query)
def _build_hosting_moderation_hook(
def _build_before_llm_invoke_hook(
self,
*,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
) -> Callable[[Sequence[PromptMessage]], None]:
def check(prompt_messages: Sequence[PromptMessage]) -> None:
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,
@@ -189,30 +188,16 @@ class CompletionWorkflowRunner(AppRunner):
):
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 _should_check_hosting_moderation(self, application_generate_entity: CompletionAppGenerateEntity) -> bool:
moderation_config = hosting_configuration.moderation_config
openai_provider_name = f"{DEFAULT_PLUGIN_ID}/openai/openai"
hosting_provider = hosting_configuration.provider_map.get(openai_provider_name)
if not (
moderation_config
and moderation_config.enabled is True
and hosting_provider
and hosting_provider.enabled is True
and hosting_provider.credentials is not None
):
return False
model_config = application_generate_entity.model_conf
provider_model_bundle = getattr(model_config, "provider_model_bundle", None)
configuration = getattr(provider_model_bundle, "configuration", None)
using_provider_type = getattr(configuration, "using_provider_type", None)
return (
using_provider_type == ProviderType.SYSTEM
and getattr(model_config, "provider", None) in moderation_config.providers
)
def _build_variable_pool(
self,
*,
@@ -232,7 +217,7 @@ class CompletionWorkflowRunner(AppRunner):
workflow_execution_id=application_generate_entity.task_id,
timestamp=int(time.time()),
query=query,
conversation_id=getattr(message, "conversation_id", None),
conversation_id=message.conversation_id,
)
add_variables_to_pool(
variable_pool,
-2
View File
@@ -6,7 +6,6 @@ from typing import Any
from pydantic import BaseModel, ConfigDict, Field
from core.app.entities.agent_strategy import AgentStrategyInfo
from core.prompt.utils.prompt_message_util import SavedPrompt
from core.rag.entities import RetrievalSourceMetadata
from core.workflow.nodes.human_input.pause_reason import PauseReason
from graphon.entities import WorkflowStartReason
@@ -274,7 +273,6 @@ class QueueMessageEndEvent(AppQueueEvent):
event: QueueEvent = QueueEvent.MESSAGE_END
llm_result: LLMResult | None = None
saved_prompt: list[SavedPrompt] | None = None
class QueueAdvancedChatMessageEndEvent(AppQueueEvent):
-2
View File
@@ -5,7 +5,6 @@ from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field, JsonValue
from core.app.entities.agent_strategy import AgentStrategyInfo
from core.prompt.utils.prompt_message_util import SavedPrompt
from core.rag.entities import RetrievalSourceMetadata
from core.workflow.nodes.human_input.entities import FormInputConfig, UserActionConfig
from core.workflow.nodes.human_input.pause_reason import DifyHITLEventType
@@ -47,7 +46,6 @@ class EasyUITaskState(TaskState):
"""
llm_result: LLMResult
saved_prompt: list[SavedPrompt] | None = None
class WorkflowTaskState(TaskState):
@@ -278,9 +278,6 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
if isinstance(event, QueueMessageEndEvent):
if event.llm_result:
self._task_state.llm_result = event.llm_result
saved_prompt = getattr(event, "saved_prompt", None)
if saved_prompt is not None:
self._task_state.saved_prompt = saved_prompt
else:
self._handle_stop(event)
@@ -409,11 +406,9 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
if not conversation:
raise ValueError(f"Conversation {self._conversation_id} not found")
saved_prompt = self._task_state.saved_prompt
if saved_prompt is None:
saved_prompt = PromptMessageUtil.prompt_messages_to_prompt_for_saving(
self._model_config.mode, self._task_state.llm_result.prompt_messages
)
saved_prompt = PromptMessageUtil.prompt_messages_to_prompt_for_saving(
self._model_config.mode, self._task_state.llm_result.prompt_messages
)
object.__setattr__(message, "message", saved_prompt)
message.message_tokens = usage.prompt_tokens
message.message_unit_price = usage.prompt_unit_price
+12 -7
View File
@@ -95,7 +95,7 @@ if TYPE_CHECKING:
DIFY_BEFORE_LLM_INVOKE_KEY = "_dify_before_llm_invoke"
BeforeLLMInvoke = Callable[[Sequence[PromptMessage]], None]
BeforeLLMInvoke = Callable[[Sequence[PromptMessage], Mapping[str, Any]], Mapping[str, Any]]
_file_access_controller = DatabaseFileAccessController()
@@ -201,9 +201,14 @@ 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]) -> None:
if self._before_invoke is not None:
self._before_invoke(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(
@@ -237,7 +242,7 @@ class DifyPreparedLLM(LLMProtocol):
stop: Sequence[str] | None,
stream: bool,
) -> LLMResult | Generator[LLMResultChunk, None, None]:
self._run_before_invoke(prompt_messages)
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),
@@ -279,7 +284,7 @@ class DifyPreparedLLM(LLMProtocol):
stop: Sequence[str] | None,
stream: bool,
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]:
self._run_before_invoke(prompt_messages)
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(),
@@ -328,7 +333,7 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
stop: Sequence[str] | None,
json_schema: Mapping[str, Any] | None,
) -> LLMPollingResult:
self._run_before_invoke(prompt_messages)
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,
+1 -3
View File
@@ -19,7 +19,6 @@ from core.helper import encrypter
from core.prompt.simple_prompt_transform import SimplePromptTransform
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
from events.app_event import app_was_created
from extensions.ext_database import db
from graphon.file import FileUploadConfig
from graphon.model_runtime.entities.llm_entities import LLMMode
from graphon.model_runtime.utils.encoders import jsonable_encoder
@@ -596,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,
@@ -342,8 +342,7 @@ def test_console_text_api_accepts_message_id_without_text(app: Flask, monkeypatc
assert response == {"audio": "ok"}
assert calls["text"] == ""
assert calls["message_ref"] == MessageRef(
"tenant-1",
"app-1",
AppRef("tenant-1", "app-1"),
"0f67f8c5-8f7c-4ebd-b549-7ac8e972d37e",
account_id="account-1",
)
@@ -40,7 +40,7 @@ def _build_generate_entity(app_config, file_upload_config=None):
def test_workflow_runner_direct_outputs_on_input_moderation() -> None:
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
runner = CompletionWorkflowRunner()
app_record = MagicMock(id="app1", tenant_id="tenant")
app_generate_entity = _build_generate_entity(_build_app_config())
queue_manager = MagicMock()
@@ -61,7 +61,7 @@ def test_workflow_runner_direct_outputs_on_input_moderation() -> None:
def test_workflow_runner_uses_low_image_detail_default() -> None:
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
runner = CompletionWorkflowRunner()
app_generate_entity = _build_generate_entity(_build_app_config(), file_upload_config=None)
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, "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
@@ -2,7 +2,7 @@ import json
from types import SimpleNamespace
from unittest.mock import MagicMock
from core.app.apps.completion.runtime_workflow_builder import RuntimeCompletionWorkflowBuilder
from core.app.apps.completion.runtime_workflow_builder import build_runtime_completion_workflow
from graphon.nodes import BuiltinNodeTypes
from models.model import AppMode
@@ -15,10 +15,11 @@ def test_builder_returns_runtime_graph_without_workflow_record() -> None:
)
session = MagicMock()
result = RuntimeCompletionWorkflowBuilder(workflow_converter=workflow_converter).build(
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-")
@@ -57,10 +58,11 @@ def test_builder_routes_api_based_variable_query_to_runtime_sys_query() -> None:
)
session = MagicMock()
result = RuntimeCompletionWorkflowBuilder(workflow_converter=workflow_converter).build(
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]
@@ -16,18 +16,23 @@ from graphon.graph_events import (
GraphRunAbortedEvent,
GraphRunFailedEvent,
GraphRunSucceededEvent,
NodeRunFailedEvent,
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() -> tuple[CompletionGraphEventAdapter, MagicMock]:
def _adapter(*, show_retrieve_source: bool = True) -> tuple[CompletionGraphEventAdapter, MagicMock]:
queue_manager = MagicMock()
entity = SimpleNamespace(model_conf=SimpleNamespace(model="model"))
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,
@@ -36,6 +41,8 @@ def _adapter() -> tuple[CompletionGraphEventAdapter, MagicMock]:
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(
@@ -50,6 +57,7 @@ def test_stream_chunk_event_publishes_llm_chunk() -> None:
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
@@ -88,10 +96,27 @@ def test_retriever_resource_event_publishes_legacy_retriever_resources() -> None
assert event.retriever_resources[0].dataset_id == "dataset"
def test_llm_success_then_graph_success_publishes_message_end_with_saved_prompt() -> None:
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()
saved_prompt = [{"role": "user", "text": "saved prompt"}]
prompt_message = UserPromptMessage(content="prompt")
adapter.set_prompt_messages([prompt_message])
adapter.handle_event(
NodeRunSucceededEvent(
id="run",
@@ -101,7 +126,6 @@ def test_llm_success_then_graph_success_publishes_message_end_with_saved_prompt(
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
outputs={"text": "final"},
process_data={"prompts": saved_prompt},
llm_usage=usage,
),
)
@@ -113,8 +137,8 @@ def test_llm_success_then_graph_success_publishes_message_end_with_saved_prompt(
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
assert event.saved_prompt == saved_prompt
def test_graph_success_uses_outputs_result_when_llm_success_was_not_seen() -> None:
@@ -126,25 +150,6 @@ def test_graph_success_uses_outputs_result_when_llm_success_was_not_seen() -> No
assert isinstance(event, QueueMessageEndEvent)
assert event.llm_result is not None
assert event.llm_result.message.content == "final from graph"
assert event.saved_prompt == []
def test_failed_node_publishes_error() -> None:
adapter, queue_manager = _adapter()
adapter.handle_event(
NodeRunFailedEvent(
id="run",
node_id="llm",
node_type=BuiltinNodeTypes.LLM,
error="node boom",
start_at=datetime.now(UTC),
)
)
event = queue_manager.publish.call_args.args[0]
assert isinstance(event, QueueErrorEvent)
assert str(event.error) == "node boom"
def test_failed_graph_publishes_error() -> None:
@@ -10,7 +10,6 @@ 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
from models.provider import ProviderType
def _entity() -> SimpleNamespace:
@@ -43,7 +42,7 @@ def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
root_node_id="start",
graph_dict={"nodes": [{"id": "start", "data": {"type": "start"}}], "edges": []},
)
builder = MagicMock(build=MagicMock(return_value=runtime_workflow))
build_runtime_workflow = MagicMock(return_value=runtime_workflow)
graph = MagicMock()
adapter = MagicMock()
workflow_entry = MagicMock()
@@ -67,6 +66,7 @@ def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
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())
@@ -76,12 +76,11 @@ def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
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(runtime_workflow_builder=builder)
runner = CompletionWorkflowRunner()
monkeypatch.setattr(runner, "_get_app", MagicMock(return_value=app))
hosting_hook = MagicMock()
build_hosting_hook = MagicMock(return_value=hosting_hook)
monkeypatch.setattr(runner, "_should_check_hosting_moderation", MagicMock(return_value=True))
monkeypatch.setattr(runner, "_build_hosting_moderation_hook", build_hosting_hook)
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",
@@ -95,7 +94,11 @@ def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
session=session,
)
builder.build.assert_called_once_with(app_model=app, app_config=entity.app_config, 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"
@@ -107,16 +110,17 @@ def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
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_hosting_hook.assert_called_once_with(
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 hosting_hook
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")
@@ -125,8 +129,12 @@ def test_runner_builds_workflow_entry_and_adapts_events(monkeypatch) -> None:
def test_runner_returns_when_input_moderation_stops(monkeypatch) -> None:
app = SimpleNamespace(id="app", tenant_id="tenant", mode=AppMode.COMPLETION)
entity = _entity()
builder = MagicMock()
runner = CompletionWorkflowRunner(runtime_workflow_builder=builder)
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,
@@ -138,33 +146,32 @@ def test_runner_returns_when_input_moderation_stops(monkeypatch) -> None:
application_generate_entity=entity,
queue_manager=MagicMock(),
message=SimpleNamespace(id="message"),
session=MagicMock(),
)
builder.build.assert_not_called()
build_runtime_workflow.assert_not_called()
def test_runner_get_app_raises_when_record_is_missing(monkeypatch) -> None:
from core.app.apps.completion import workflow_runner as module
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
monkeypatch.setattr(module.db.session, "scalar", MagicMock(return_value=None))
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("missing-app")
runner._get_app(app_id="missing-app", tenant_id="tenant", session=session)
def test_runner_get_app_returns_record(monkeypatch) -> None:
from core.app.apps.completion import workflow_runner as module
def test_runner_get_app_returns_record() -> None:
app = SimpleNamespace(id="app")
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
monkeypatch.setattr(module.db.session, "scalar", MagicMock(return_value=app))
runner = CompletionWorkflowRunner()
session = MagicMock()
session.scalar.return_value = app
assert runner._get_app("app") is 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(runtime_workflow_builder=MagicMock())
runner = CompletionWorkflowRunner()
app_record = SimpleNamespace(id="app", tenant_id="tenant")
entity = _entity()
message = SimpleNamespace(id="message")
@@ -187,7 +194,7 @@ def test_runner_direct_outputs_on_input_moderation() -> None:
def test_runner_returns_moderated_inputs_when_input_moderation_passes() -> None:
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
runner = CompletionWorkflowRunner()
app_record = SimpleNamespace(id="app", tenant_id="tenant")
entity = _entity()
message = SimpleNamespace(id="message")
@@ -204,20 +211,23 @@ def test_runner_returns_moderated_inputs_when_input_moderation_passes() -> None:
assert result == ModeratedCompletionInputs(stopped=False, inputs={"name": "Grace"}, query="moderated query")
def test_runner_hosting_moderation_hook_uses_final_prompt() -> None:
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
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_hosting_moderation_hook(
hook = runner._build_before_llm_invoke_hook(
application_generate_entity=entity,
queue_manager=queue_manager,
adapter=adapter,
)
with pytest.raises(GenerateTaskStoppedError):
hook(["final prompt"])
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,
@@ -225,48 +235,31 @@ def test_runner_hosting_moderation_hook_uses_final_prompt() -> None:
)
def test_runner_should_not_check_hosting_moderation_when_config_is_disabled(monkeypatch) -> None:
from core.app.apps.completion import workflow_runner as module
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
monkeypatch.setattr(
module,
"hosting_configuration",
SimpleNamespace(
moderation_config=SimpleNamespace(enabled=False),
provider_map={},
),
)
assert runner._should_check_hosting_moderation(_entity()) is False
def test_runner_should_check_hosting_moderation_for_system_provider(monkeypatch) -> None:
from core.app.apps.completion import workflow_runner as module
def test_runner_before_llm_invoke_hook_recalculates_graph_model_parameters() -> None:
runner = CompletionWorkflowRunner()
entity = _entity()
entity.model_conf = SimpleNamespace(
provider="openai",
provider_model_bundle=SimpleNamespace(
configuration=SimpleNamespace(using_provider_type=ProviderType.SYSTEM),
),
)
runner = CompletionWorkflowRunner(runtime_workflow_builder=MagicMock())
monkeypatch.setattr(
module,
"hosting_configuration",
SimpleNamespace(
moderation_config=SimpleNamespace(enabled=True, providers=["openai"]),
provider_map={
f"{module.DEFAULT_PLUGIN_ID}/openai/openai": SimpleNamespace(
enabled=True,
credentials={"api_key": "secret"},
)
},
),
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,
)
assert runner._should_check_hosting_moderation(entity) is True
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:
@@ -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()
@@ -405,51 +405,6 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert isinstance(responses[-1], MessageEndStreamResponse)
assert pipeline._task_state.llm_result.message.content == "done"
def test_process_stream_response_carries_saved_prompt_from_message_end(self, monkeypatch: pytest.MonkeyPatch):
pipeline, _ = _make_pipeline()
saved_prompt = [{"role": "user", "text": "serialized by graphon"}]
llm_result = LLMResult(
model="mock",
prompt_messages=[],
message=AssistantPromptMessage(content="done"),
usage=LLMUsage.empty_usage(),
)
_set_queue_events(
pipeline,
[_queue_message(QueueMessageEndEvent(llm_result=llm_result, saved_prompt=saved_prompt))],
)
_set_method(pipeline, "handle_output_moderation_when_task_finished", lambda completion: None)
_set_method(
pipeline,
"_message_end_to_stream_response",
lambda: MessageEndStreamResponse(task_id="task", id="msg"),
)
def _save_message(**kwargs):
assert pipeline._task_state.saved_prompt == saved_prompt
_set_method(pipeline, "_save_message", _save_message)
class _Session:
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.sessionmaker",
lambda **kwargs: type("_SessionFactory", (), {"begin": lambda self: _Session()})(),
)
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db",
_FakeDb(),
)
responses = list(pipeline._process_stream_response(publisher=None))
assert isinstance(responses[-1], MessageEndStreamResponse)
def test_handle_output_moderation_chunk_directs_output(self):
conversation = _make_conversation(AppMode.CHAT)
message = _make_message()
@@ -1279,33 +1234,6 @@ class TestEasyUiBasedGenerateTaskPipeline:
assert trace_task.kwargs["trace_session_id"] == "session-1"
assert len(sent_payloads) == 1
def test_save_message_uses_saved_prompt_override(self, monkeypatch: pytest.MonkeyPatch):
pipeline, _ = _make_pipeline()
_set_method(pipeline, "_model_config", _ModelConfigMode(mode="chat"))
pipeline._task_state.saved_prompt = [{"role": "user", "text": "serialized by graphon"}]
pipeline._task_state.llm_result.message = AssistantPromptMessage(content="answer")
pipeline._task_state.llm_result.usage = LLMUsage.empty_usage()
message_obj = _make_message()
conversation_obj = _make_conversation(AppMode.CHAT)
session = Mock()
session.scalar.side_effect = [message_obj, conversation_obj]
serialize_mock = Mock(side_effect=AssertionError("saved prompt override should be used"))
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.PromptMessageUtil.prompt_messages_to_prompt_for_saving",
serialize_mock,
)
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.message_was_created.send",
lambda *args, **kwargs: None,
)
pipeline._save_message(session=session)
assert message_obj.message == [{"role": "user", "text": "serialized by graphon"}]
serialize_mock.assert_not_called()
def test_save_message_raises_when_message_not_found(self):
conversation = _make_conversation(AppMode.CHAT)
message = _make_message()
@@ -228,7 +228,7 @@ 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
before_invoke = Mock()
before_invoke = Mock(return_value={"temperature": 0.05})
prepared = DifyPreparedLLM(
model_instance,
request_metadata={"app_id": "app-id"},
@@ -253,10 +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([])
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,
@@ -275,7 +275,7 @@ 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())
before_invoke = Mock()
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)
@@ -289,14 +289,14 @@ def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytes
)
assert result is sentinel.structured
before_invoke.assert_called_once_with([])
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,
)
@@ -328,7 +328,7 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None:
model_runtime=plugin_runtime,
)
before_invoke = Mock()
before_invoke = Mock(return_value={"temperature": 0.05})
prepared = DifyPreparedPollingLLM(model_instance, before_invoke=before_invoke)
assert isinstance(prepared, LLMPollingCapableProtocol)
@@ -342,7 +342,7 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None:
)
== polling_result
)
before_invoke.assert_called_once_with([])
before_invoke.assert_called_once_with([], {"temperature": 0.1})
assert (
prepared.check_llm_polling(
plugin_state={"task_id": "poll-1"},
@@ -354,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"},
@@ -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(