Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
62518a5e55 | ||
|
|
e2c769c6cc | ||
|
|
6ad1a106c0 | ||
|
|
659c486e12 | ||
|
|
eaa55a4292 | ||
|
|
4122233ba0 | ||
|
|
dfc2726e44 | ||
|
|
de8ae2a6ab | ||
|
|
eb82fa8e04 | ||
|
|
8ca84a7b45 |
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
+11
-3
@@ -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
|
||||
|
||||
+70
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user