feat(workflow): support human input in loop and iteration (#39243)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
-LAN-
2026-08-03 02:24:18 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 8dbac96621
commit 351577bdb0
86 changed files with 1454 additions and 762 deletions
@@ -21,6 +21,7 @@ from core.workflow.system_variables import SystemVariableKey, get_system_text
from extensions.ext_database import db
from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelType
from graphon.runtime import VariablePool
from graphon.variables.template_resolution import convert_template
from models.model import Conversation
from .entities import AgentNodeData, AgentOldVersionModelFeatures, ParamsAutoGenerated
@@ -67,7 +68,7 @@ class AgentRuntimeSupport:
except TypeError:
parameter_value = str(agent_input.value)
segment_group = variable_pool.convert_template(parameter_value)
segment_group = convert_template(variable_pool, parameter_value)
parameter_value = segment_group.log if for_log else segment_group.text
try:
if not isinstance(agent_input.value, str):
+14 -17
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
import json
import logging
from collections.abc import Mapping, Sequence
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime, timedelta
from typing import Any
@@ -11,10 +11,10 @@ from core.repositories.human_input_repository import FormCreateParams, HumanInpu
from core.workflow.human_input_adapter import DeliveryChannelConfig
from core.workflow.node_runtime import DifyFileReferenceFactory
from graphon.nodes.human_input.entities import Completed, Expired, HITLContext, HITLDecision, PauseRequested
from graphon.runtime import VariablePool
from graphon.runtime.graph_runtime_state_protocol import ReadOnlyVariablePool
from graphon.variables.factory import build_segment
from graphon.variables.segments import Segment
from graphon.variables.template_resolution import convert_template
from libs.datetime_utils import ensure_naive_utc, naive_utc_now
from .entities import (
@@ -31,24 +31,13 @@ from .session_binding import default_session_binding
logger = logging.getLogger(__name__)
def _require_template_variable_pool(pool: ReadOnlyVariablePool) -> VariablePool:
"""Return the concrete graphon pool required for template expansion."""
if isinstance(pool, VariablePool):
return pool
msg = "human input rendering requires graphon.runtime.VariablePool for template expansion"
raise TypeError(msg)
def render_form_content_before_submission(
node_data: HumanInputNodeData,
*,
variable_pool: ReadOnlyVariablePool,
) -> str:
"""Process form content by substituting runtime variables before pause."""
# NOTE(QuantumGhost): This is not ideal, we should expose
# VariablePool method in Graphon.
rendered_form_content = _require_template_variable_pool(variable_pool).convert_template(node_data.form_content)
rendered_form_content = convert_template(variable_pool, node_data.form_content)
return rendered_form_content.markdown
@@ -91,6 +80,7 @@ class DifyHITLCallback:
delivery_methods: Sequence[DeliveryChannelConfig] = (),
display_in_ui: bool = False,
file_reference_factory: DifyFileReferenceFactory | None = None,
execution_id_getter: Callable[[], str | None] | None = None,
) -> None:
self._form_repository = form_repository
self._session_binding = default_session_binding
@@ -100,11 +90,17 @@ class DifyHITLCallback:
self._delivery_methods = tuple(delivery_methods)
self._display_in_ui = display_in_ui
self._file_reference_factory = file_reference_factory
self._execution_id_getter = execution_id_getter
def __call__(self, ctx: HITLContext) -> HITLDecision:
form = self._form_repository.get_form(ctx.node_id)
form_id = self._execution_id_getter() if self._execution_id_getter is not None else None
form = (
self._form_repository.get_form(ctx.node_id, form_id=form_id)
if form_id is not None
else self._form_repository.get_form(ctx.node_id)
)
if form is None:
created = self._create_form(ctx)
created = self._create_form(ctx, form_id=form_id)
return PauseRequested(session_id=self._session_binding.issue_session_id_for_form(form_id=created.id))
status = self._normalize_status(form.status)
@@ -163,7 +159,7 @@ class DifyHITLCallback:
outputs=outputs,
)
def _create_form(self, ctx: HITLContext) -> HumanInputFormEntity:
def _create_form(self, ctx: HITLContext, *, form_id: str | None = None) -> HumanInputFormEntity:
params = FormCreateParams(
workflow_execution_id=self._workflow_execution_id or ctx.workflow_execution_id,
conversation_id=self._conversation_id,
@@ -181,6 +177,7 @@ class DifyHITLCallback:
variable_pool=ctx.variable_pool,
)
),
form_id=form_id,
)
return self._form_repository.create_form(params)
@@ -25,7 +25,6 @@ from graphon.enums import (
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.model_runtime.utils.encoders import jsonable_encoder
from graphon.node_events import NodeRunResult
from graphon.nodes.base import LLMUsageTrackingMixin
from graphon.nodes.base.node import Node
from graphon.variables import (
ArrayFileSegment,
@@ -33,6 +32,7 @@ from graphon.variables import (
StringSegment,
)
from graphon.variables.segments import ArrayObjectSegment
from graphon.variables.template_resolution import convert_template
from .entities import (
Condition,
@@ -64,7 +64,7 @@ def _normalize_metadata_filter_sequence_item(value: object) -> str:
return value if isinstance(value, str) else str(value)
class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeData]):
class KnowledgeRetrievalNode(Node[KnowledgeRetrievalNodeData]):
node_type = BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL
# Instance attributes specific to LLMNode.
@@ -309,7 +309,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD
resolved_value: str | Sequence[str] | int | float | None
match value:
case str():
segment_group = variable_pool.convert_template(value)
segment_group = convert_template(variable_pool, value)
if len(segment_group.value) == 1:
resolved_value = _normalize_metadata_filter_scalar(segment_group.value[0].to_object())
else:
@@ -317,7 +317,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD
case _ if isinstance(value, Sequence) and all(isinstance(v, str) for v in value):
resolved_values: list[str] = []
for v in value:
segment_group = variable_pool.convert_template(v)
segment_group = convert_template(variable_pool, v)
if len(segment_group.value) == 1:
resolved_values.append(
_normalize_metadata_filter_sequence_item(segment_group.value[0].to_object())