Compare commits

...
Author SHA1 Message Date
autofix-ci[bot]andGitHub ed2a102953 [autofix.ci] apply automated fixes 2026-04-30 11:10:24 +00:00
yunlu.wen 6b87dda00e fix type check 2026-04-30 19:08:12 +08:00
yunlu.wen 67e85b34e3 Merge remote-tracking branch 'upstream/main' into feat/upgrade-graphon 2026-04-30 18:56:49 +08:00
yunlu.wen 28289212fb fix type check 2026-04-30 18:55:57 +08:00
yunlu.wen f8912c920e fix integration tests 2026-04-30 18:17:00 +08:00
yunlu.wen 08d08f0ae3 lint 2026-04-30 18:03:08 +08:00
yunlu.wen 187e12956a fix tests 2026-04-30 18:02:09 +08:00
yunlu.wen 131facbc65 update ModelProviderFactory 2026-04-30 17:04:44 +08:00
yunlu.wen 9acd149469 update node params && VariablePool instantiation 2026-04-30 16:59:08 +08:00
88 changed files with 479 additions and 283 deletions
@@ -175,7 +175,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
# Create a variable pool.
# init variable pool
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
@@ -144,7 +144,7 @@ class PipelineRunner(WorkflowBasedAppRunner):
)
)
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
+1 -1
View File
@@ -106,7 +106,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
workflow_id=app_config.workflow_id,
workflow_execution_id=self.application_generate_entity.workflow_execution_id,
)
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
+1 -1
View File
@@ -188,7 +188,7 @@ class WorkflowBasedAppRunner:
ValueError: If neither single_iteration_run nor single_loop_run is specified
"""
# Create initial runtime state with variable pool containing environment variables
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
+3 -1
View File
@@ -16,6 +16,7 @@ from graphon.graph_engine.entities.commands import AbortCommand, CommandType
from graphon.graph_engine.layers import GraphEngineLayer
from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, NodeRunSucceededEvent
from graphon.nodes.base.node import Node
from graphon.nodes.llm.runtime_protocols import PreparedLLMProtocol
if TYPE_CHECKING:
from graphon.nodes.llm.node import LLMNode
@@ -116,7 +117,8 @@ class LLMQuotaLayer(GraphEngineLayer):
case BuiltinNodeTypes.PARAMETER_EXTRACTOR:
model_instance = cast("ParameterExtractorNode", node).model_instance
case BuiltinNodeTypes.QUESTION_CLASSIFIER:
model_instance = cast("QuestionClassifierNode", node).model_instance
typed_node: QuestionClassifierNode = cast("QuestionClassifierNode", node)
model_instance = cast(PreparedLLMProtocol, typed_node._model_instance)
case _:
return None
except AttributeError:
+7 -5
View File
@@ -24,6 +24,7 @@ from core.entities.provider_entities import (
from core.helper import encrypter
from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType
from core.plugin.impl.model_runtime_factory import create_plugin_model_provider_factory
from graphon.model_runtime import ModelRuntime
from graphon.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType
from graphon.model_runtime.entities.provider_entities import (
ConfigurateMethod,
@@ -33,7 +34,6 @@ from graphon.model_runtime.entities.provider_entities import (
)
from graphon.model_runtime.model_providers.base.ai_model import AIModel
from graphon.model_runtime.model_providers.model_provider_factory import ModelProviderFactory
from graphon.model_runtime.runtime import ModelRuntime
from libs.datetime_utils import naive_utc_now
from models.engine import db
from models.enums import CredentialSourceType
@@ -109,7 +109,7 @@ class ProviderConfiguration(BaseModel):
def get_model_provider_factory(self) -> ModelProviderFactory:
"""Return a provider factory that preserves any request-bound runtime."""
if self._bound_model_runtime is not None:
return ModelProviderFactory(model_runtime=self._bound_model_runtime)
return ModelProviderFactory(runtime=self._bound_model_runtime)
return create_plugin_model_provider_factory(tenant_id=self.tenant_id)
def get_current_credentials(self, model_type: ModelType, model: str) -> dict[str, Any] | None:
@@ -1392,10 +1392,12 @@ class ProviderConfiguration(BaseModel):
:param model_type: model type
:return:
"""
model_provider_factory = self.get_model_provider_factory()
from core.plugin.impl.model_runtime_factory import create_model_type_instance
# Get model instance of LLM
return model_provider_factory.get_model_type_instance(provider=self.provider.provider, model_type=model_type)
model_provider_factory = self.get_model_provider_factory()
return create_model_type_instance(
factory=model_provider_factory, provider=self.provider.provider, model_type=model_type
)
def get_model_schema(
self, model_type: ModelType, model: str, credentials: dict[str, Any] | None
+3 -3
View File
@@ -4,7 +4,7 @@ from typing import cast
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.entities import DEFAULT_PLUGIN_ID
from core.plugin.impl.model_runtime_factory import create_plugin_model_provider_factory
from core.plugin.impl.model_runtime_factory import create_model_type_instance, create_plugin_model_provider_factory
from extensions.ext_hosting_provider import hosting_configuration
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.errors.invoke import InvokeBadRequestError
@@ -44,8 +44,8 @@ def check_moderation(tenant_id: str, model_config: ModelConfigWithCredentialsEnt
model_provider_factory = create_plugin_model_provider_factory(tenant_id=tenant_id)
# Get model instance of LLM
model_type_instance = model_provider_factory.get_model_type_instance(
provider=openai_provider_name, model_type=ModelType.MODERATION
model_type_instance = create_model_type_instance(
factory=model_provider_factory, provider=openai_provider_name, model_type=ModelType.MODERATION
)
model_type_instance = cast(ModerationModel, model_type_instance)
moderation_result = model_type_instance.invoke(
+80 -3
View File
@@ -4,7 +4,7 @@ import hashlib
import logging
from collections.abc import Generator, Iterable, Sequence
from threading import Lock
from typing import IO, Any, Union
from typing import IO, Any, Literal, Union, overload
from pydantic import ValidationError
from redis import RedisError
@@ -14,13 +14,18 @@ from core.plugin.entities.plugin_daemon import PluginModelProviderEntity
from core.plugin.impl.asset import PluginAssetManager
from core.plugin.impl.model import PluginModelClient
from extensions.ext_redis import redis_client
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk
from graphon.model_runtime import ModelRuntime
from graphon.model_runtime.entities.llm_entities import (
LLMResult,
LLMResultChunk,
LLMResultChunkWithStructuredOutput,
LLMResultWithStructuredOutput,
)
from graphon.model_runtime.entities.message_entities import PromptMessage, PromptMessageTool
from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelType
from graphon.model_runtime.entities.provider_entities import ProviderEntity
from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult
from graphon.model_runtime.entities.text_embedding_entities import EmbeddingInputType, EmbeddingResult
from graphon.model_runtime.runtime import ModelRuntime
from models.provider_ids import ModelProviderID
logger = logging.getLogger(__name__)
@@ -195,6 +200,34 @@ class PluginModelRuntime(ModelRuntime):
return schema
@overload
def invoke_llm(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
model_parameters: dict[str, Any],
prompt_messages: Sequence[PromptMessage],
tools: list[PromptMessageTool] | None,
stop: Sequence[str] | None,
stream: Literal[False],
) -> LLMResult: ...
@overload
def invoke_llm(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
model_parameters: dict[str, Any],
prompt_messages: Sequence[PromptMessage],
tools: list[PromptMessageTool] | None,
stop: Sequence[str] | None,
stream: Literal[True],
) -> Generator[LLMResultChunk, None, None]: ...
def invoke_llm(
self,
*,
@@ -222,6 +255,50 @@ class PluginModelRuntime(ModelRuntime):
stream=stream,
)
@overload
def invoke_llm_with_structured_output(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
json_schema: dict[str, Any],
model_parameters: dict[str, Any],
prompt_messages: Sequence[PromptMessage],
stop: Sequence[str] | None,
stream: Literal[False],
) -> LLMResultWithStructuredOutput: ...
@overload
def invoke_llm_with_structured_output(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
json_schema: dict[str, Any],
model_parameters: dict[str, Any],
prompt_messages: Sequence[PromptMessage],
stop: Sequence[str] | None,
stream: Literal[True],
) -> Generator[LLMResultChunkWithStructuredOutput, None, None]: ...
def invoke_llm_with_structured_output(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
json_schema: dict[str, Any],
model_parameters: dict[str, Any],
prompt_messages: Sequence[PromptMessage],
stop: Sequence[str] | None,
stream: bool,
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]:
# TODO: added to pass type check.
# it is a new method from upstream that is not invoked at all.
raise NotImplementedError
def get_llm_num_tokens(
self,
*,
+45 -1
View File
@@ -3,6 +3,14 @@ from __future__ import annotations
from typing import TYPE_CHECKING
from core.plugin.impl.model import PluginModelClient
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.model_providers.base.ai_model import AIModel
from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel
from graphon.model_runtime.model_providers.base.moderation_model import ModerationModel
from graphon.model_runtime.model_providers.base.rerank_model import RerankModel
from graphon.model_runtime.model_providers.base.speech2text_model import Speech2TextModel
from graphon.model_runtime.model_providers.base.text_embedding_model import TextEmbeddingModel
from graphon.model_runtime.model_providers.base.tts_model import TTSModel
from graphon.model_runtime.model_providers.model_provider_factory import ModelProviderFactory
if TYPE_CHECKING:
@@ -10,6 +18,15 @@ if TYPE_CHECKING:
from core.plugin.impl.model_runtime import PluginModelRuntime
from core.provider_manager import ProviderManager
_MODEL_TYPE_CLASS_MAP: dict[ModelType, type[AIModel]] = {
ModelType.LLM: LargeLanguageModel,
ModelType.TEXT_EMBEDDING: TextEmbeddingModel,
ModelType.RERANK: RerankModel,
ModelType.SPEECH2TEXT: Speech2TextModel,
ModelType.MODERATION: ModerationModel,
ModelType.TTS: TTSModel,
}
class PluginModelAssembly:
"""Compose request-scoped model views on top of a single plugin runtime."""
@@ -38,7 +55,7 @@ class PluginModelAssembly:
@property
def model_provider_factory(self) -> ModelProviderFactory:
if self._model_provider_factory is None:
self._model_provider_factory = ModelProviderFactory(model_runtime=self.model_runtime)
self._model_provider_factory = ModelProviderFactory(runtime=self.model_runtime)
return self._model_provider_factory
@property
@@ -87,3 +104,30 @@ def create_plugin_provider_manager(*, tenant_id: str, user_id: str | None = None
def create_plugin_model_manager(*, tenant_id: str, user_id: str | None = None) -> ModelManager:
"""Create a tenant-bound model manager for service flows."""
return create_plugin_model_assembly(tenant_id=tenant_id, user_id=user_id).model_manager
def create_model_type_instance(
factory: ModelProviderFactory,
provider: str,
model_type: ModelType,
) -> AIModel:
"""Instantiate the AIModel subclass for *model_type* backed by *factory*'s runtime.
This replaces ``ModelProviderFactory.get_model_type_instance`` which was
removed in graphon 0.3.0. The mapping from ModelType to concrete AIModel
subclass is maintained here so that callers do not need to know the
subclass constructors.
:param factory: factory whose ``runtime`` and provider resolution are used.
:param provider: provider identifier (canonical or short name).
:param model_type: the model type to instantiate.
:returns: an AIModel subclass instance wired to the factory's runtime.
:raises ValueError: if *model_type* is not supported.
"""
model_class = _MODEL_TYPE_CLASS_MAP.get(model_type)
if model_class is None:
msg = f"Unsupported model type: {model_type}"
raise ValueError(msg)
provider_entity = factory.get_model_provider(provider)
return model_class(provider_schema=provider_entity, model_runtime=factory.runtime)
+3 -3
View File
@@ -56,7 +56,7 @@ from models.provider_ids import ModelProviderID
from services.feature_service import FeatureService
if TYPE_CHECKING:
from graphon.model_runtime.runtime import ModelRuntime
from graphon.model_runtime import ModelRuntime
_credentials_adapter: TypeAdapter[dict[str, Any]] = TypeAdapter(dict[str, Any])
@@ -165,7 +165,7 @@ class ProviderManager:
)
# Get all provider entities
model_provider_factory = ModelProviderFactory(model_runtime=self._model_runtime)
model_provider_factory = ModelProviderFactory(runtime=self._model_runtime)
provider_entities = model_provider_factory.get_providers()
# Get All preferred provider types of the workspace
@@ -362,7 +362,7 @@ class ProviderManager:
if not default_model:
return None
model_provider_factory = ModelProviderFactory(model_runtime=self._model_runtime)
model_provider_factory = ModelProviderFactory(runtime=self._model_runtime)
provider_schema = model_provider_factory.get_provider_schema(provider=default_model.provider_name)
return DefaultModelEntity(
+1 -1
View File
@@ -448,7 +448,7 @@ class DifyNodeFactory(NodeFactory):
node_init_kwargs = node_init_kwargs_factories.get(node_type, lambda: {})()
return node_class(
node_id=node_id,
config=config_for_node_init,
data=config_for_node_init,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
**node_init_kwargs,
+2 -2
View File
@@ -35,7 +35,7 @@ class AgentNode(Node[AgentNodeData]):
def __init__(
self,
node_id: str,
config: AgentNodeData,
data: AgentNodeData,
*,
graph_init_params: GraphInitParams,
graph_runtime_state: GraphRuntimeState,
@@ -46,7 +46,7 @@ class AgentNode(Node[AgentNodeData]):
) -> None:
super().__init__(
node_id=node_id,
config=config,
data=data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
@@ -36,14 +36,14 @@ class DatasourceNode(Node[DatasourceNodeData]):
def __init__(
self,
node_id: str,
config: DatasourceNodeData,
data: DatasourceNodeData,
*,
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
) -> None:
super().__init__(
node_id=node_id,
config=config,
data=data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
@@ -32,14 +32,14 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]):
def __init__(
self,
node_id: str,
config: KnowledgeIndexNodeData,
data: KnowledgeIndexNodeData,
*,
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
) -> None:
super().__init__(
node_id=node_id,
config=config,
data=data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
@@ -71,14 +71,14 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD
def __init__(
self,
node_id: str,
config: KnowledgeRetrievalNodeData,
data: KnowledgeRetrievalNodeData,
*,
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
) -> None:
super().__init__(
node_id=node_id,
config=config,
data=data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
+9 -15
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
from collections import defaultdict
from collections.abc import Mapping, Sequence
from enum import StrEnum
from typing import Any, Protocol, cast
from typing import Any, Protocol
from uuid import uuid4
from graphon.enums import BuiltinNodeTypes
@@ -82,13 +82,10 @@ def build_system_variables(values: Mapping[str, Any] | None = None, /, **kwargs:
normalized = _normalize_system_variable_values(values, **kwargs)
return [
cast(
Variable,
segment_to_variable(
segment=build_segment(value),
selector=system_variable_selector(key),
name=key,
),
segment_to_variable(
segment=build_segment(value),
selector=system_variable_selector(key),
name=key,
)
for key, value in normalized.items()
]
@@ -130,13 +127,10 @@ def build_bootstrap_variables(
for node_id, value in rag_pipeline_variables_map.items():
variables.append(
cast(
Variable,
segment_to_variable(
segment=build_segment(value),
selector=(RAG_PIPELINE_VARIABLE_NODE_ID, node_id),
name=node_id,
),
segment_to_variable(
segment=build_segment(value),
selector=(RAG_PIPELINE_VARIABLE_NODE_ID, node_id),
name=node_id,
)
)
+1 -1
View File
@@ -411,7 +411,7 @@ class WorkflowEntry:
raise ValueError(f"Node class not found for node type {node_type}")
# init variable pool
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(variable_pool, default_system_variables())
# init graph context and runtime state
+1 -1
View File
@@ -45,7 +45,7 @@ dependencies = [
# Emerging: newer and fast-moving, use compatible pins
"fastopenapi[flask]~=0.7.0",
"graphon~=0.2.2",
"graphon~=0.3.0",
"httpx-sse~=0.4.0",
"json-repair~=0.59.4",
]
+1 -1
View File
@@ -88,7 +88,7 @@ logger = logging.getLogger(__name__)
def _build_seeded_variable_pool(variables: Sequence[Variable]) -> VariablePool:
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(variable_pool, variables)
return variable_pool
@@ -32,7 +32,7 @@ from graphon.file import File
from graphon.nodes import BuiltinNodeTypes
from graphon.nodes.variable_assigner.common.helpers import get_updated_variables
from graphon.variable_loader import VariableLoader
from graphon.variables import Segment, StringSegment, VariableBase
from graphon.variables import Segment, StringSegment, Variable, VariableBase
from graphon.variables.consts import SELECTORS_LENGTH
from graphon.variables.segments import (
ArrayFileSegment,
@@ -162,7 +162,7 @@ class DraftVarLoader(VariableLoader):
return list(variable_by_selector.values())
def _load_offloaded_variable(self, draft_var: WorkflowDraftVariable) -> tuple[tuple[str, str], VariableBase]:
def _load_offloaded_variable(self, draft_var: WorkflowDraftVariable) -> tuple[tuple[str, str], Variable]:
# This logic is closely tied to `WorkflowDraftVaribleService._try_offload_large_variable`
# and must remain synchronized with it.
# Ideally, these should be co-located for better maintainability.
+4 -4
View File
@@ -877,7 +877,7 @@ class WorkflowService:
)
else:
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
@@ -1251,7 +1251,7 @@ class WorkflowService:
node_data = HumanInputNode.validate_node_data(adapt_human_input_node_data_for_graph(node_config["data"]))
node = HumanInputNode(
node_id=node_config["id"],
config=node_data,
data=node_data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
runtime=DifyHumanInputNodeRuntime(run_context),
@@ -1271,7 +1271,7 @@ class WorkflowService:
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user_id)
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
@@ -1659,7 +1659,7 @@ def _setup_variable_pool(
system_variable = default_system_variables()
# init variable pool
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
@@ -71,7 +71,7 @@ def test_node_integration_minimal_stream(mocker):
node = DatasourceNode(
node_id="n",
config=DatasourceNodeData(
data=DatasourceNodeData(
type="datasource",
version="1",
title="Datasource",
@@ -4,7 +4,7 @@ from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEnti
from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle
from core.entities.provider_entities import CustomConfiguration, CustomProviderConfiguration, SystemConfiguration
from core.model_manager import ModelInstance
from core.plugin.impl.model_runtime_factory import create_plugin_model_provider_factory
from core.plugin.impl.model_runtime_factory import create_model_type_instance, create_plugin_model_provider_factory
from graphon.model_runtime.entities.model_entities import ModelType
from models.provider import ProviderType
@@ -16,7 +16,11 @@ def get_mocked_fetch_model_config(
credentials: dict,
):
model_provider_factory = create_plugin_model_provider_factory(tenant_id="9d2074fc-6f86-45a9-b09d-6ecc63b9056b")
model_type_instance = model_provider_factory.get_model_type_instance(provider, ModelType.LLM)
model_type_instance = create_model_type_instance(
factory=model_provider_factory,
provider=provider,
model_type=ModelType.LLM,
)
provider_model_bundle = ProviderModelBundle(
configuration=ProviderConfiguration(
tenant_id="1",
@@ -45,7 +45,7 @@ def init_code_node(code_config: dict):
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -66,7 +66,7 @@ def init_code_node(code_config: dict):
node = CodeNode(
node_id=str(uuid.uuid4()),
config=CodeNodeData.model_validate(code_config["data"]),
data=CodeNodeData.model_validate(code_config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
code_executor=node_factory._code_executor,
@@ -55,7 +55,7 @@ def init_http_node(config: dict):
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -76,7 +76,7 @@ def init_http_node(config: dict):
node = HttpRequestNode(
node_id=str(uuid.uuid4()),
config=HttpRequestNodeData.model_validate(config["data"]),
data=HttpRequestNodeData.model_validate(config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
http_request_config=HTTP_REQUEST_CONFIG,
@@ -204,7 +204,7 @@ def test_custom_auth_with_empty_api_key_raises_error(setup_http_mock):
from graphon.runtime import VariablePool
# Create variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="test", files=[]),
user_inputs={},
environment_variables=[],
@@ -702,7 +702,7 @@ def test_nested_object_variable_selector(setup_http_mock):
)
# Create independent variable pool for this test only
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -724,7 +724,7 @@ def test_nested_object_variable_selector(setup_http_mock):
node = HttpRequestNode(
node_id=str(uuid.uuid4()),
config=HttpRequestNodeData.model_validate(graph_config["nodes"][1]["data"]),
data=HttpRequestNodeData.model_validate(graph_config["nodes"][1]["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
http_request_config=HTTP_REQUEST_CONFIG,
@@ -53,7 +53,7 @@ def init_llm_node(config: dict) -> LLMNode:
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="aaa",
app_id=app_id,
@@ -77,7 +77,7 @@ def init_llm_node(config: dict) -> LLMNode:
node = LLMNode(
node_id=str(uuid.uuid4()),
config=LLMNodeData.model_validate(config["data"]),
data=LLMNodeData.model_validate(config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=MagicMock(spec=CredentialsProvider),
@@ -56,7 +56,7 @@ def init_parameter_extractor_node(config: dict, memory=None):
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="aaa", files=[], query="what's the weather in SF", conversation_id="abababa"
),
@@ -71,7 +71,7 @@ def init_parameter_extractor_node(config: dict, memory=None):
node = ParameterExtractorNode(
node_id=str(uuid.uuid4()),
config=ParameterExtractorNodeData.model_validate(config["data"]),
data=ParameterExtractorNodeData.model_validate(config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=MagicMock(spec=CredentialsProvider),
@@ -66,7 +66,7 @@ def test_execute_template_transform():
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -88,7 +88,7 @@ def test_execute_template_transform():
node = TemplateTransformNode(
node_id=str(uuid.uuid4()),
config=TemplateTransformNodeData.model_validate(config["data"]),
data=TemplateTransformNodeData.model_validate(config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
jinja2_template_renderer=_SimpleJinja2Renderer(),
@@ -41,7 +41,7 @@ def init_tool_node(config: dict):
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -62,7 +62,7 @@ def init_tool_node(config: dict):
node = ToolNode(
node_id=str(uuid.uuid4()),
config=ToolNodeData.model_validate(config["data"]),
data=ToolNodeData.model_validate(config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
tool_file_manager_factory=tool_file_manager_factory,
@@ -210,7 +210,9 @@ class TestPauseStatePersistenceLayerTestContainers:
execution_id = workflow_run_id or getattr(self, "test_workflow_run_id", None) or str(uuid.uuid4())
# Create variable pool
variable_pool = VariablePool(system_variables=build_system_variables(workflow_execution_id=execution_id))
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id=execution_id),
)
if variables:
for (node_id, var_key), value in variables.items():
variable_pool.add([node_id, var_key], value)
@@ -66,7 +66,7 @@ def _mock_form_repository_with_submission(action_id: str) -> HumanInputFormRepos
def _build_runtime_state(workflow_execution_id: str, app_id: str, workflow_id: str, user_id: str) -> GraphRuntimeState:
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
workflow_execution_id=workflow_execution_id,
app_id=app_id,
@@ -102,7 +102,7 @@ def _build_graph(
start_data = StartNodeData(title="start", variables=[])
start_node = StartNode(
node_id="start",
config=start_data,
data=start_data,
graph_init_params=params,
graph_runtime_state=runtime_state,
)
@@ -117,7 +117,7 @@ def _build_graph(
)
human_node = HumanInputNode(
node_id="human",
config=human_data,
data=human_data,
graph_init_params=params,
graph_runtime_state=runtime_state,
form_repository=form_repository,
@@ -131,7 +131,7 @@ def _build_graph(
)
end_node = EndNode(
node_id="end",
config=end_data,
data=end_data,
graph_init_params=params,
graph_runtime_state=runtime_state,
)
@@ -186,7 +186,7 @@ class TestEmailDeliveryTestHandler:
handler = EmailDeliveryTestHandler(session_factory=MagicMock())
handler._resolve_recipients = MagicMock(return_value=["test@example.com"])
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
context = DeliveryTestContext(
tenant_id="t1",
app_id="a1",
@@ -177,7 +177,7 @@ def test_dispatch_human_input_email_task_integration(monkeypatch: pytest.MonkeyP
workflow_run_id = str(uuid.uuid4())
workflow_id = str(uuid.uuid4())
app_id = str(uuid.uuid4())
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
variable_pool.add(["node1", "value"], "OK")
_create_workflow_pause_state(
db_session_with_containers,
@@ -236,7 +236,7 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext:
call_depth=0,
workflow_execution_id="run-1",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
runtime_state = GraphRuntimeState(variable_pool=VariablePool.from_bootstrap(), start_at=0.0)
runtime_state.register_paused_node("node-1")
runtime_state.outputs = {"result": "value"}
wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity)
@@ -132,7 +132,9 @@ class TestAdvancedChatGenerateTaskPipeline:
pipeline._task_state.answer = "partial answer"
pipeline._workflow_run_id = "run-id"
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
total_tokens=7,
node_run_steps=3,
@@ -372,7 +374,9 @@ class TestAdvancedChatGenerateTaskPipeline:
pipeline = _make_pipeline()
pipeline._workflow_run_id = "run-id"
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
pipeline._workflow_response_converter.workflow_finish_to_stream_response = lambda **kwargs: "finish"
@@ -583,7 +587,9 @@ class TestAdvancedChatGenerateTaskPipeline:
self.items = items
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
@@ -617,7 +623,9 @@ class TestAdvancedChatGenerateTaskPipeline:
def test_handle_message_end_event_applies_output_moderation(self, monkeypatch):
pipeline = _make_pipeline()
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
pipeline._base_task_pipeline.handle_output_moderation_when_task_finished = lambda answer: "safe"
@@ -9,7 +9,7 @@ from graphon.runtime import GraphRuntimeState, VariablePool
def _make_state(workflow_run_id: str | None) -> GraphRuntimeState:
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(variable_pool, build_system_variables(workflow_execution_id=workflow_run_id))
return GraphRuntimeState(variable_pool=variable_pool, start_at=0.0)
@@ -17,7 +17,7 @@ def _build_converter():
workflow_id="wf-1",
workflow_execution_id="run-1",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
runtime_state = GraphRuntimeState(variable_pool=VariablePool.from_bootstrap(), start_at=0.0)
app_entity = SimpleNamespace(
task_id="task-1",
app_config=SimpleNamespace(app_id="app-1", tenant_id="tenant-1"),
@@ -16,7 +16,7 @@ def _build_converter() -> WorkflowResponseConverter:
workflow_id="wf-1",
workflow_execution_id="run-1",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
runtime_state = GraphRuntimeState(variable_pool=VariablePool.from_bootstrap(), start_at=0.0)
app_entity = SimpleNamespace(
task_id="task-1",
app_config=SimpleNamespace(app_id="app-1", tenant_id="tenant-1"),
@@ -271,7 +271,8 @@ def test_run_normal_path_builds_graph(mocker):
def add(self, selector, value):
return None
mocker.patch.object(module, "VariablePool", return_value=FakeVariablePool())
fake_pool = FakeVariablePool()
mocker.patch("graphon.runtime.VariablePool.from_bootstrap", return_value=fake_pool)
workflow_entry = MagicMock()
workflow_entry.graph_engine = MagicMock()
@@ -58,7 +58,7 @@ class _StubToolNode(Node[_StubToolNodeData]):
def __init__(
self,
node_id: str,
config: _StubToolNodeData,
data: _StubToolNodeData,
*,
graph_init_params,
graph_runtime_state,
@@ -66,7 +66,7 @@ class _StubToolNode(Node[_StubToolNodeData]):
) -> None:
super().__init__(
node_id=node_id,
config=config,
data=data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
@@ -167,7 +167,7 @@ def _build_graph(runtime_state: GraphRuntimeState, *, pause_on: str | None) -> G
def _build_runtime_state(run_id: str) -> GraphRuntimeState:
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="user", app_id="app", workflow_id="workflow"),
user_inputs={},
conversation_variables=[],
@@ -54,7 +54,7 @@ class TestWorkflowBasedAppRunner:
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
@@ -93,7 +93,7 @@ class TestWorkflowBasedAppRunner:
def test_get_graph_and_variable_pool_for_single_node_run(self, monkeypatch):
runner = WorkflowBasedAppRunner(queue_manager=SimpleNamespace(), app_id="app")
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
@@ -162,7 +162,7 @@ class TestWorkflowBasedAppRunner:
app_id="app",
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
@@ -241,7 +241,7 @@ class TestWorkflowBasedAppRunner:
runner = WorkflowBasedAppRunner(queue_manager=_QueueManager(), app_id="app")
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
graph_runtime_state.register_paused_node("node-1")
@@ -284,7 +284,7 @@ class TestWorkflowBasedAppRunner:
runner = WorkflowBasedAppRunner(queue_manager=_QueueManager(), app_id="app")
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
workflow_entry = SimpleNamespace(graph_engine=SimpleNamespace(graph_runtime_state=graph_runtime_state))
@@ -423,7 +423,7 @@ class TestWorkflowBasedAppRunner:
runner = WorkflowBasedAppRunner(queue_manager=_QueueManager(), app_id="app")
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
start_at=0.0,
)
workflow_entry = SimpleNamespace(graph_engine=SimpleNamespace(graph_runtime_state=graph_runtime_state))
@@ -16,7 +16,7 @@ from models.workflow import Workflow
def _make_graph_state():
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
environment_variables=[],
@@ -95,7 +95,9 @@ class TestWorkflowGenerateTaskPipeline:
def test_to_blocking_response_falls_back_to_human_input_required_when_pause_event_missing(self):
pipeline = _make_pipeline()
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
total_tokens=5,
node_run_steps=2,
@@ -283,7 +285,9 @@ class TestWorkflowGenerateTaskPipeline:
pipeline = _make_pipeline()
pipeline._workflow_execution_id = "run-id"
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
pipeline._workflow_response_converter.workflow_finish_to_stream_response = lambda **kwargs: "finish"
@@ -725,7 +729,9 @@ class TestWorkflowGenerateTaskPipeline:
pipeline = _make_pipeline()
pipeline._workflow_execution_id = "run-id"
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
@@ -753,7 +759,9 @@ class TestWorkflowGenerateTaskPipeline:
pipeline = _make_pipeline()
pipeline._workflow_execution_id = "run-id"
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
pipeline._handle_ping_event = lambda event, **kwargs: iter(["ping"])
@@ -769,7 +777,9 @@ class TestWorkflowGenerateTaskPipeline:
def test_process_stream_response_main_match_paths_and_cleanup(self):
pipeline = _make_pipeline()
pipeline._graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-id")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-id")
),
start_at=0.0,
)
pipeline._base_task_pipeline.queue_manager.listen = lambda: iter(
@@ -21,7 +21,9 @@ class TestTriggerPostLayer:
)
runtime_state = SimpleNamespace(
outputs={"answer": "ok"},
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-1")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-1")
),
total_tokens=12,
)
@@ -60,7 +62,9 @@ class TestTriggerPostLayer:
def test_on_event_handles_missing_trigger_log(self):
runtime_state = SimpleNamespace(
outputs={},
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-1")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-1")
),
total_tokens=0,
)
@@ -91,7 +95,9 @@ class TestTriggerPostLayer:
def test_on_event_ignores_non_status_events(self):
runtime_state = SimpleNamespace(
outputs={},
variable_pool=VariablePool(system_variables=build_system_variables(workflow_execution_id="run-1")),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(workflow_execution_id="run-1")
),
total_tokens=0,
)
@@ -8,9 +8,9 @@ from graphon.enums import BuiltinNodeTypes
class DummyNode:
def __init__(self, *, node_id, config, graph_init_params, graph_runtime_state, **kwargs):
def __init__(self, *, node_id, data, graph_init_params, graph_runtime_state, **kwargs):
self.id = node_id
self.config = config
self.data = data
self.graph_init_params = graph_init_params
self.graph_runtime_state = graph_runtime_state
self.kwargs = kwargs
@@ -60,7 +60,9 @@ def _make_layer(
workflow_execution_id="run-id",
conversation_id="conv-id",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(system_variables=system_variables), start_at=0.0)
runtime_state = GraphRuntimeState(
variable_pool=VariablePool.from_bootstrap(system_variables=system_variables), start_at=0.0
)
read_only_state = ReadOnlyGraphRuntimeStateWrapper(runtime_state)
application_generate_entity = WorkflowAppGenerateEntity.model_construct(
@@ -429,20 +429,25 @@ def test_get_model_type_instance_and_schema_delegate_to_factory() -> None:
mock_factory = Mock()
mock_model_type_instance = Mock()
mock_schema = _build_ai_model("gpt-4o")
mock_factory.get_model_type_instance.return_value = mock_model_type_instance
mock_factory.get_model_schema.return_value = mock_schema
with patch(
"core.entities.provider_configuration.create_plugin_model_provider_factory",
return_value=mock_factory,
) as mock_factory_builder:
with (
patch(
"core.entities.provider_configuration.create_plugin_model_provider_factory",
return_value=mock_factory,
) as mock_factory_builder,
patch(
"core.plugin.impl.model_runtime_factory.create_model_type_instance",
return_value=mock_model_type_instance,
) as mock_create_instance,
):
model_type_instance = configuration.get_model_type_instance(ModelType.LLM)
model_schema = configuration.get_model_schema(ModelType.LLM, "gpt-4o", {"api_key": "x"})
assert model_type_instance is mock_model_type_instance
assert model_schema is mock_schema
assert mock_factory_builder.call_count == 2
mock_factory.get_model_type_instance.assert_called_once_with(provider="openai", model_type=ModelType.LLM)
mock_create_instance.assert_called_once_with(factory=mock_factory, provider="openai", model_type=ModelType.LLM)
mock_factory.get_model_schema.assert_called_once_with(
provider="openai",
model_type=ModelType.LLM,
@@ -459,7 +464,6 @@ def test_get_model_type_instance_and_schema_reuse_bound_runtime_factory() -> Non
mock_factory = Mock()
mock_model_type_instance = Mock()
mock_schema = _build_ai_model("gpt-4o")
mock_factory.get_model_type_instance.return_value = mock_model_type_instance
mock_factory.get_model_schema.return_value = mock_schema
with (
@@ -467,6 +471,10 @@ def test_get_model_type_instance_and_schema_reuse_bound_runtime_factory() -> Non
"core.entities.provider_configuration.ModelProviderFactory", return_value=mock_factory
) as mock_factory_cls,
patch("core.entities.provider_configuration.create_plugin_model_provider_factory") as mock_factory_builder,
patch(
"core.plugin.impl.model_runtime_factory.create_model_type_instance",
return_value=mock_model_type_instance,
) as mock_create_instance,
):
model_type_instance = configuration.get_model_type_instance(ModelType.LLM)
model_schema = configuration.get_model_schema(ModelType.LLM, "gpt-4o", {"api_key": "x"})
@@ -474,8 +482,9 @@ def test_get_model_type_instance_and_schema_reuse_bound_runtime_factory() -> Non
assert model_type_instance is mock_model_type_instance
assert model_schema is mock_schema
assert mock_factory_cls.call_count == 2
mock_factory_cls.assert_called_with(model_runtime=bound_runtime)
mock_factory_cls.assert_called_with(runtime=bound_runtime)
mock_factory_builder.assert_not_called()
mock_create_instance.assert_called_once_with(factory=mock_factory, provider="openai", model_type=ModelType.LLM)
def test_get_provider_model_returns_none_when_model_not_found() -> None:
@@ -1,5 +1,6 @@
from types import SimpleNamespace
from typing import cast
from unittest.mock import Mock
import pytest
from pytest_mock import MockerFixture
@@ -68,8 +69,9 @@ def test_check_moderation_returns_true_when_model_accepts_text(mocker: MockerFix
mocker.patch("core.helper.moderation.secrets.choice", return_value="chunk")
moderation_model = SimpleNamespace(invoke=lambda **invoke_kwargs: invoke_kwargs["text"] == "chunk")
factory = SimpleNamespace(get_model_type_instance=lambda **_factory_kwargs: moderation_model)
factory = Mock()
mocker.patch("core.helper.moderation.create_plugin_model_provider_factory", return_value=factory)
mocker.patch("core.helper.moderation.create_model_type_instance", return_value=moderation_model)
assert (
check_moderation(
@@ -119,8 +121,9 @@ def test_check_moderation_returns_false_when_model_rejects_text(mocker: MockerFi
mocker.patch("core.helper.moderation.secrets.choice", return_value="chunk")
moderation_model = SimpleNamespace(invoke=lambda **_invoke_kwargs: False)
factory = SimpleNamespace(get_model_type_instance=lambda **_factory_kwargs: moderation_model)
factory = Mock()
mocker.patch("core.helper.moderation.create_plugin_model_provider_factory", return_value=factory)
mocker.patch("core.helper.moderation.create_model_type_instance", return_value=moderation_model)
assert (
check_moderation(
@@ -147,8 +150,9 @@ def test_check_moderation_raises_bad_request_when_provider_call_fails(mocker: Mo
failing_model = SimpleNamespace(
invoke=lambda **_invoke_kwargs: (_ for _ in ()).throw(RuntimeError("boom")),
)
factory = SimpleNamespace(get_model_type_instance=lambda **_factory_kwargs: failing_model)
factory = Mock()
mocker.patch("core.helper.moderation.create_plugin_model_provider_factory", return_value=factory)
mocker.patch("core.helper.moderation.create_model_type_instance", return_value=failing_model)
with pytest.raises(InvokeBadRequestError, match="Rate limit exceeded, please try again later."):
check_moderation(
@@ -2,6 +2,7 @@ from unittest.mock import Mock
import pytest
from core.plugin.impl.model_runtime_factory import create_model_type_instance
from graphon.model_runtime.entities.common_entities import I18nObject
from graphon.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType
from graphon.model_runtime.entities.provider_entities import (
@@ -73,7 +74,7 @@ def test_model_provider_factory_resolves_runtime_provider_name() -> None:
supported_model_types=[ModelType.LLM],
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL],
)
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime([provider]))
factory = ModelProviderFactory(runtime=_FakeModelRuntime([provider]))
provider_schema = factory.get_model_provider("openai")
@@ -98,7 +99,7 @@ def test_model_provider_factory_resolves_canonical_short_name_independent_of_pro
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL],
),
]
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime(providers))
factory = ModelProviderFactory(runtime=_FakeModelRuntime(providers))
provider_schema = factory.get_model_provider("openai")
@@ -107,8 +108,8 @@ def test_model_provider_factory_resolves_canonical_short_name_independent_of_pro
def test_model_provider_factory_requires_runtime() -> None:
with pytest.raises(ValueError, match="model_runtime is required"):
ModelProviderFactory(model_runtime=None) # type: ignore[arg-type]
with pytest.raises(ValueError, match="runtime is required"):
ModelProviderFactory(runtime=None) # type: ignore[arg-type]
def test_model_provider_factory_get_providers_returns_runtime_providers() -> None:
@@ -119,7 +120,7 @@ def test_model_provider_factory_get_providers_returns_runtime_providers() -> Non
supported_model_types=[ModelType.LLM],
)
]
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime(providers))
factory = ModelProviderFactory(runtime=_FakeModelRuntime(providers))
result = factory.get_providers()
@@ -133,7 +134,7 @@ def test_model_provider_factory_get_provider_schema_delegates_to_provider_lookup
provider_name="openai",
supported_model_types=[ModelType.LLM],
)
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime([provider]))
factory = ModelProviderFactory(runtime=_FakeModelRuntime([provider]))
result = factory.get_provider_schema("openai")
@@ -142,7 +143,7 @@ def test_model_provider_factory_get_provider_schema_delegates_to_provider_lookup
def test_model_provider_factory_raises_for_unknown_provider() -> None:
factory = ModelProviderFactory(
model_runtime=_FakeModelRuntime(
runtime=_FakeModelRuntime(
[
_build_provider(
provider="langgenius/openai/openai",
@@ -172,7 +173,7 @@ def test_model_provider_factory_get_models_filters_provider_and_model_type() ->
models=[_build_model("rerank-v3", ModelType.RERANK)],
),
]
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime(providers))
factory = ModelProviderFactory(runtime=_FakeModelRuntime(providers))
results = factory.get_models(provider="openai", model_type=ModelType.LLM)
@@ -196,7 +197,7 @@ def test_model_provider_factory_get_models_skips_providers_without_requested_mod
models=[_build_model("eleven_multilingual_v2", ModelType.TTS)],
),
]
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime(providers))
factory = ModelProviderFactory(runtime=_FakeModelRuntime(providers))
results = factory.get_models(model_type=ModelType.TTS)
@@ -214,7 +215,7 @@ def test_model_provider_factory_get_models_without_model_type_keeps_all_provider
models=[_build_model("gpt-4o-mini", ModelType.LLM), _build_model("tts-1", ModelType.TTS)],
)
]
factory = ModelProviderFactory(model_runtime=_FakeModelRuntime(providers))
factory = ModelProviderFactory(runtime=_FakeModelRuntime(providers))
results = factory.get_models(provider="openai")
@@ -242,7 +243,7 @@ def test_model_provider_factory_validates_provider_credentials() -> None:
)
]
)
factory = ModelProviderFactory(model_runtime=runtime)
factory = ModelProviderFactory(runtime=runtime)
filtered = factory.provider_credentials_validate(
provider="openai",
@@ -258,7 +259,7 @@ def test_model_provider_factory_validates_provider_credentials() -> None:
def test_model_provider_factory_provider_credentials_validate_requires_schema() -> None:
factory = ModelProviderFactory(
model_runtime=_FakeModelRuntime(
runtime=_FakeModelRuntime(
[
_build_provider(
provider="langgenius/openai/openai",
@@ -294,7 +295,7 @@ def test_model_provider_factory_validates_model_credentials() -> None:
)
]
)
factory = ModelProviderFactory(model_runtime=runtime)
factory = ModelProviderFactory(runtime=runtime)
filtered = factory.model_credentials_validate(
provider="openai",
@@ -314,7 +315,7 @@ def test_model_provider_factory_validates_model_credentials() -> None:
def test_model_provider_factory_model_credentials_validate_requires_schema() -> None:
factory = ModelProviderFactory(
model_runtime=_FakeModelRuntime(
runtime=_FakeModelRuntime(
[
_build_provider(
provider="langgenius/openai/openai",
@@ -346,7 +347,7 @@ def test_model_provider_factory_get_model_schema_and_icon_use_canonical_provider
)
runtime.get_model_schema.return_value = "schema"
runtime.get_provider_icon.return_value = (b"icon", "image/png")
factory = ModelProviderFactory(model_runtime=runtime)
factory = ModelProviderFactory(runtime=runtime)
assert (
factory.get_model_schema(
@@ -387,7 +388,7 @@ def test_model_provider_factory_builds_model_type_instances(
expected_type: type[object],
) -> None:
factory = ModelProviderFactory(
model_runtime=_FakeModelRuntime(
runtime=_FakeModelRuntime(
[
_build_provider(
provider="langgenius/openai/openai",
@@ -398,14 +399,14 @@ def test_model_provider_factory_builds_model_type_instances(
)
)
instance = factory.get_model_type_instance("openai", model_type)
instance = create_model_type_instance(factory=factory, provider="openai", model_type=model_type)
assert isinstance(instance, expected_type)
def test_model_provider_factory_rejects_unsupported_model_type() -> None:
factory = ModelProviderFactory(
model_runtime=_FakeModelRuntime(
runtime=_FakeModelRuntime(
[
_build_provider(
provider="langgenius/openai/openai",
@@ -417,4 +418,4 @@ def test_model_provider_factory_rejects_unsupported_model_type() -> None:
)
with pytest.raises(ValueError, match="Unsupported model type: unsupported"):
factory.get_model_type_instance("openai", "unsupported") # type: ignore[arg-type]
create_model_type_instance(factory=factory, provider="openai", model_type="unsupported") # type: ignore[arg-type]
@@ -31,6 +31,6 @@ def test_plugin_model_assembly_reuses_single_runtime_across_views():
assert assembly.model_manager is model_manager
mock_runtime_factory.assert_called_once_with(tenant_id="tenant-1", user_id="user-1")
mock_provider_factory_cls.assert_called_once_with(model_runtime=runtime)
mock_provider_factory_cls.assert_called_once_with(runtime=runtime)
mock_provider_manager_cls.assert_called_once_with(model_runtime=runtime)
mock_model_manager_cls.assert_called_once_with(provider_manager=provider_manager)
@@ -289,7 +289,7 @@ def test_get_default_model_uses_injected_runtime_for_existing_default_record(moc
result = manager.get_default_model("tenant-id", ModelType.LLM)
mock_factory_cls.assert_called_once_with(model_runtime=manager._model_runtime)
mock_factory_cls.assert_called_once_with(runtime=manager._model_runtime)
assert result is not None
assert result.model == "gpt-4"
assert result.provider.provider == "openai"
@@ -316,7 +316,7 @@ def test_get_configurations_uses_injected_runtime_and_adds_provider_aliases(mock
result = manager.get_configurations("tenant-id")
expected_alias = str(ModelProviderID("openai"))
mock_factory_cls.assert_called_once_with(model_runtime=manager._model_runtime)
mock_factory_cls.assert_called_once_with(runtime=manager._model_runtime)
assert result.tenant_id == "tenant-id"
assert expected_alias in provider_records
assert expected_alias in provider_model_records
@@ -402,7 +402,7 @@ def test_get_configurations_reuses_cached_result_for_same_tenant(mocker: MockerF
assert first is second
mock_get_all_providers.assert_called_once_with("tenant-id")
mock_factory_cls.assert_called_once_with(model_runtime=manager._model_runtime)
mock_factory_cls.assert_called_once_with(runtime=manager._model_runtime)
mock_provider_configuration.assert_called_once()
provider_configuration.bind_model_runtime.assert_called_once_with(manager._model_runtime)
@@ -76,7 +76,7 @@ def _build_variable_pool(
system_variables: list[Variable] | None = None,
environment_variables: list[Variable] | None = None,
) -> VariablePool:
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
@@ -96,7 +96,7 @@ class MockNodeFactory(DifyNodeFactory):
if node_type == BuiltinNodeTypes.CODE:
mock_instance = mock_class(
node_id=node_id,
config=resolved_node_data,
data=resolved_node_data,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
mock_config=self.mock_config,
@@ -106,7 +106,7 @@ class MockNodeFactory(DifyNodeFactory):
elif node_type == BuiltinNodeTypes.HTTP_REQUEST:
mock_instance = mock_class(
node_id=node_id,
config=resolved_node_data,
data=resolved_node_data,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
mock_config=self.mock_config,
@@ -122,7 +122,7 @@ class MockNodeFactory(DifyNodeFactory):
}:
mock_instance = mock_class(
node_id=node_id,
config=resolved_node_data,
data=resolved_node_data,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
mock_config=self.mock_config,
@@ -132,7 +132,7 @@ class MockNodeFactory(DifyNodeFactory):
else:
mock_instance = mock_class(
node_id=node_id,
config=resolved_node_data,
data=resolved_node_data,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
mock_config=self.mock_config,
@@ -56,7 +56,7 @@ class MockNodeMixin:
def __init__(
self,
node_id: str,
config: Any,
data: Any,
*,
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
@@ -98,7 +98,7 @@ class MockNodeMixin:
super().__init__(
node_id=node_id,
config=config,
data=data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
**kwargs,
@@ -111,7 +111,7 @@ class StaticRepo(HumanInputFormRepository):
def _build_runtime_state() -> GraphRuntimeState:
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="user",
app_id="app",
@@ -140,7 +140,7 @@ def _build_graph(runtime_state: GraphRuntimeState, repo: HumanInputFormRepositor
start_config = {"id": "start", "data": StartNodeData(title="Start", variables=[]).model_dump()}
start_node = StartNode(
node_id=start_config["id"],
config=StartNodeData(title="Start", variables=[]),
data=StartNodeData(title="Start", variables=[]),
graph_init_params=graph_init_params,
graph_runtime_state=runtime_state,
)
@@ -155,7 +155,7 @@ def _build_graph(runtime_state: GraphRuntimeState, repo: HumanInputFormRepositor
human_a_config = {"id": "human_a", "data": human_data.model_dump()}
human_a = HumanInputNode(
node_id=human_a_config["id"],
config=human_data,
data=human_data,
graph_init_params=graph_init_params,
graph_runtime_state=runtime_state,
form_repository=repo,
@@ -165,7 +165,7 @@ def _build_graph(runtime_state: GraphRuntimeState, repo: HumanInputFormRepositor
human_b_config = {"id": "human_b", "data": human_data.model_dump()}
human_b = HumanInputNode(
node_id=human_b_config["id"],
config=human_data,
data=human_data,
graph_init_params=graph_init_params,
graph_runtime_state=runtime_state,
form_repository=repo,
@@ -183,7 +183,7 @@ def _build_graph(runtime_state: GraphRuntimeState, repo: HumanInputFormRepositor
end_config = {"id": "end", "data": end_data.model_dump()}
end_node = EndNode(
node_id=end_config["id"],
config=end_data,
data=end_data,
graph_init_params=graph_init_params,
graph_runtime_state=runtime_state,
)
@@ -250,7 +250,7 @@ class WorkflowRunner:
conversation_variables.append(var)
root_node_id = get_default_root_node_id(graph_config)
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_bootstrap_variables(
@@ -48,7 +48,7 @@ def test_execute_answer():
)
# construct variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -71,7 +71,7 @@ def test_execute_answer():
node_id=str(uuid.uuid4()),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
config=AnswerNodeData(
data=AnswerNodeData(
title="123",
type="answer",
answer="Today's weather is {{#start.weather#}}\n{{#llm.text#}}\n{{img}}\nFin.",
@@ -79,7 +79,7 @@ def test_datasource_node_delegates_to_manager_stream(mocker):
node = DatasourceNode(
node_id="n",
config=DatasourceNodeData(
data=DatasourceNodeData(
type="datasource",
version="1",
title="Datasource",
@@ -29,7 +29,7 @@ HTTP_REQUEST_CONFIG = HttpRequestNodeConfig(
def test_executor_with_json_body_and_number_variable():
# Prepare the variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -85,7 +85,7 @@ def test_executor_with_json_body_and_number_variable():
def test_executor_with_json_body_and_object_variable():
# Prepare the variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -143,7 +143,7 @@ def test_executor_with_json_body_and_object_variable():
def test_executor_with_json_body_and_nested_object_variable():
# Prepare the variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -201,7 +201,7 @@ def test_executor_with_json_body_and_nested_object_variable():
def test_extract_selectors_from_template_with_newline():
variable_pool = VariablePool(system_variables=default_system_variables())
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables())
variable_pool.add(("node_id", "custom_query"), "line1\nline2")
node_data = HttpRequestNodeData(
title="Test JSON Body with Nested Object Variable",
@@ -230,7 +230,7 @@ def test_extract_selectors_from_template_with_newline():
def test_executor_with_form_data():
# Prepare the variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -320,7 +320,7 @@ def test_init_headers():
node_data=node_data,
timeout=timeout,
http_request_config=HTTP_REQUEST_CONFIG,
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
http_client=ssrf_proxy,
file_manager=file_manager,
)
@@ -357,7 +357,7 @@ def test_init_params():
node_data=node_data,
timeout=timeout,
http_request_config=HTTP_REQUEST_CONFIG,
variable_pool=VariablePool(system_variables=default_system_variables()),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()),
http_client=ssrf_proxy,
file_manager=file_manager,
)
@@ -390,7 +390,7 @@ def test_init_params():
def test_empty_api_key_raises_error_bearer():
"""Test that empty API key raises AuthorizationConfigError for bearer auth."""
variable_pool = VariablePool(system_variables=default_system_variables())
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables())
node_data = HttpRequestNodeData(
title="test",
method="get",
@@ -417,7 +417,7 @@ def test_empty_api_key_raises_error_bearer():
def test_empty_api_key_raises_error_basic():
"""Test that empty API key raises AuthorizationConfigError for basic auth."""
variable_pool = VariablePool(system_variables=default_system_variables())
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables())
node_data = HttpRequestNodeData(
title="test",
method="get",
@@ -444,7 +444,7 @@ def test_empty_api_key_raises_error_basic():
def test_empty_api_key_raises_error_custom():
"""Test that empty API key raises AuthorizationConfigError for custom auth."""
variable_pool = VariablePool(system_variables=default_system_variables())
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables())
node_data = HttpRequestNodeData(
title="test",
method="get",
@@ -471,7 +471,7 @@ def test_empty_api_key_raises_error_custom():
def test_whitespace_only_api_key_raises_error():
"""Test that whitespace-only API key raises AuthorizationConfigError."""
variable_pool = VariablePool(system_variables=default_system_variables())
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables())
node_data = HttpRequestNodeData(
title="test",
method="get",
@@ -498,7 +498,7 @@ def test_whitespace_only_api_key_raises_error():
def test_valid_api_key_works():
"""Test that valid API key works correctly for bearer auth."""
variable_pool = VariablePool(system_variables=default_system_variables())
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables())
node_data = HttpRequestNodeData(
title="test",
method="get",
@@ -536,7 +536,7 @@ def test_executor_with_json_body_and_unquoted_uuid_variable():
# UUID that triggers the json_repair truncation bug
test_uuid = "57eeeeb1-450b-482c-81b9-4be77e95dee2"
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -583,7 +583,7 @@ def test_executor_with_json_body_and_unquoted_uuid_with_newlines():
"""
test_uuid = "57eeeeb1-450b-482c-81b9-4be77e95dee2"
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -624,7 +624,7 @@ def test_executor_with_json_body_and_unquoted_uuid_with_newlines():
def test_executor_with_json_body_preserves_numbers_and_strings():
"""Test that numbers are preserved and string values are properly quoted."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -110,12 +110,14 @@ def _build_http_node(
call_depth=0,
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(user_id="user", files=[]), user_inputs={}),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="user", files=[]), user_inputs={}
),
start_at=time.perf_counter(),
)
return HttpRequestNode(
node_id="http-node",
config=HttpRequestNodeData.model_validate(node_data),
data=HttpRequestNodeData.model_validate(node_data),
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
http_request_config=HTTP_REQUEST_CONFIG,
@@ -8,7 +8,7 @@ def test_render_body_template_replaces_variable_values():
subject="Subject",
body="Hello {{#node1.value#}} {{#url#}}",
)
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
variable_pool.add(["node1", "value"], "World")
result = config.render_body_template(body=config.body, url="https://example.com", variable_pool=variable_pool)
@@ -149,7 +149,7 @@ def _build_human_input_node(
)
return HumanInputNode(
node_id=node_id,
config=typed_node_data,
data=typed_node_data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
runtime=runtime,
@@ -250,7 +250,7 @@ class TestUserAction:
("field_name", "value"),
[
("id", "a" * 21),
("title", "b" * 21),
("title", "b" * 101),
],
)
def test_user_action_length_limits(self, field_name: str, value: str):
@@ -427,7 +427,7 @@ class TestHumanInputNodeVariableResolution:
"""Tests for resolving variable-based defaults in HumanInputNode."""
def test_resolves_variable_defaults(self):
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="user",
app_id="app",
@@ -504,7 +504,7 @@ class TestHumanInputNodeVariableResolution:
assert params.resolved_default_values == expected_values
def test_debugger_falls_back_to_recipient_token_when_webapp_disabled(self):
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="user",
app_id="app",
@@ -565,7 +565,7 @@ class TestHumanInputNodeVariableResolution:
assert not hasattr(pause_event.reason, "form_token")
def test_webapp_runtime_keeps_form_visible_in_ui_when_webapp_delivery_is_enabled(self):
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="user",
app_id="app",
@@ -631,7 +631,7 @@ class TestHumanInputNodeVariableResolution:
assert params.display_in_ui is True
def test_debugger_debug_mode_overrides_email_recipients(self):
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="user-123",
app_id="app",
@@ -748,7 +748,7 @@ class TestHumanInputNodeRenderedContent:
"""Tests for rendering submitted content."""
def test_replaces_outputs_placeholders_after_submission(self):
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="user",
app_id="app",
@@ -40,7 +40,7 @@ def _create_human_input_node(
)
return HumanInputNode(
node_id=config["id"],
config=node_data,
data=node_data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
form_repository=repo,
@@ -51,7 +51,9 @@ def _create_human_input_node(
def _build_node(form_content: str = "Please enter your name:\n\n{{#$output.name#}}") -> HumanInputNode:
system_variables = default_system_variables()
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=system_variables, user_inputs={}, environment_variables=[]),
variable_pool=VariablePool.from_bootstrap(
system_variables=system_variables, user_inputs={}, environment_variables=[]
),
start_at=0.0,
)
graph_init_params = GraphInitParams(
@@ -114,7 +116,9 @@ def _build_node(form_content: str = "Please enter your name:\n\n{{#$output.name#
def _build_timeout_node() -> HumanInputNode:
system_variables = default_system_variables()
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=system_variables, user_inputs={}, environment_variables=[]),
variable_pool=VariablePool.from_bootstrap(
system_variables=system_variables, user_inputs={}, environment_variables=[]
),
start_at=0.0,
)
graph_init_params = GraphInitParams(
@@ -32,7 +32,7 @@ class _MissingGraphBuilder:
def _build_runtime_state() -> GraphRuntimeState:
return GraphRuntimeState(
variable_pool=VariablePool(system_variables=default_system_variables(), user_inputs={}),
variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables(), user_inputs={}),
start_at=0.0,
)
@@ -46,7 +46,7 @@ def _build_iteration_node(
init_params = build_test_graph_init_params(graph_config=graph_config)
return IterationNode(
node_id="iteration-node",
config=IterationNodeData(
data=IterationNodeData(
type="iteration",
title="Iteration",
iterator_selector=["start", "items"],
@@ -40,7 +40,7 @@ def mock_graph_init_params():
@pytest.fixture
def mock_graph_runtime_state():
"""Create mock GraphRuntimeState."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id=str(uuid.uuid4()), files=[]),
user_inputs={},
environment_variables=[],
@@ -102,7 +102,7 @@ def _build_node(
) -> KnowledgeIndexNode:
return KnowledgeIndexNode(
node_id=node_id,
config=(
data=(
node_data
if isinstance(node_data, KnowledgeIndexNodeData)
else KnowledgeIndexNodeData.model_validate(node_data)
@@ -46,7 +46,7 @@ def mock_graph_init_params():
@pytest.fixture
def mock_graph_runtime_state():
"""Create mock GraphRuntimeState."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id=str(uuid.uuid4()), files=[]),
user_inputs={},
environment_variables=[],
@@ -117,7 +117,7 @@ class TestKnowledgeRetrievalNode:
# Act
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -146,7 +146,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -205,7 +205,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -249,7 +249,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -285,7 +285,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -320,7 +320,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -361,7 +361,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -400,7 +400,7 @@ class TestKnowledgeRetrievalNode:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -481,7 +481,7 @@ class TestFetchDatasetRetriever:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -518,7 +518,7 @@ class TestFetchDatasetRetriever:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -573,7 +573,7 @@ class TestFetchDatasetRetriever:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -621,7 +621,7 @@ class TestFetchDatasetRetriever:
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -682,7 +682,7 @@ class TestFetchDatasetRetriever:
config = {"id": node_id, "data": node_data.model_dump()}
node = KnowledgeRetrievalNode(
node_id=node_id,
config=KnowledgeRetrievalNodeData.model_validate(config["data"]),
data=KnowledgeRetrievalNodeData.model_validate(config["data"]),
graph_init_params=mock_graph_init_params,
graph_runtime_state=mock_graph_runtime_state,
)
@@ -15,7 +15,7 @@ from core.app.llm.model_access import (
)
from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle
from core.entities.provider_entities import CustomConfiguration, SystemConfiguration
from core.plugin.impl.model_runtime_factory import create_plugin_model_runtime
from core.plugin.impl.model_runtime_factory import create_model_type_instance, create_plugin_model_runtime
from core.prompt.entities.advanced_prompt_entities import MemoryConfig
from core.workflow.system_variables import default_system_variables
from graphon.entities import GraphInitParams
@@ -187,7 +187,7 @@ def graph_init_params() -> GraphInitParams:
@pytest.fixture
def graph_runtime_state() -> GraphRuntimeState:
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -208,7 +208,7 @@ def llm_node(
http_client = mock.MagicMock()
node = LLMNode(
node_id="1",
config=llm_node_data,
data=llm_node_data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=mock_credentials_provider,
@@ -241,9 +241,13 @@ def model_config(monkeypatch):
)
# Create actual provider and model type instances
model_provider_factory = ModelProviderFactory(model_runtime=create_plugin_model_runtime(tenant_id="test"))
model_provider_factory = ModelProviderFactory(runtime=create_plugin_model_runtime(tenant_id="test"))
provider_instance = model_provider_factory.get_model_provider("openai")
model_type_instance = model_provider_factory.get_model_type_instance("openai", ModelType.LLM)
model_type_instance = create_model_type_instance(
factory=model_provider_factory,
provider="openai",
model_type=ModelType.LLM,
)
# Create a ProviderModelBundle
provider_model_bundle = ProviderModelBundle(
@@ -1173,7 +1177,7 @@ def llm_node_for_multimodal(llm_node_data, graph_init_params, graph_runtime_stat
http_client = mock.MagicMock()
node = LLMNode(
node_id="1",
config=llm_node_data,
data=llm_node_data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=mock_credentials_provider,
@@ -28,7 +28,7 @@ def _build_template_transform_node(
)
return TemplateTransformNode(
node_id=node_id,
config=typed_node_data,
data=typed_node_data,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
**kwargs,
@@ -39,7 +39,7 @@ def mock_graph_runtime_state():
def test_node_uses_default_max_output_length_when_not_overridden(graph_init_params, mock_graph_runtime_state):
node = TemplateTransformNode(
node_id="test_node",
config=TemplateTransformNodeData(
data=TemplateTransformNodeData(
title="Template Transform",
type="template-transform",
variables=[],
@@ -35,7 +35,9 @@ def _build_context(graph_config: Mapping[str, object]) -> tuple[GraphInitParams,
invoke_from="debugger",
)
runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(user_id="user", files=[]), user_inputs={}),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="user", files=[]), user_inputs={}
),
start_at=0.0,
)
return init_params, runtime_state
@@ -62,7 +64,7 @@ def test_node_hydrates_data_during_initialization():
node = _SampleNode(
node_id="node-1",
config=_build_node_data(),
data=_build_node_data(),
graph_init_params=init_params,
graph_runtime_state=runtime_state,
)
@@ -82,13 +84,15 @@ def test_node_accepts_invoke_from_enum():
invoke_from=InvokeFrom.DEBUGGER,
)
runtime_state = GraphRuntimeState(
variable_pool=VariablePool(system_variables=build_system_variables(user_id="user", files=[]), user_inputs={}),
variable_pool=VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="user", files=[]), user_inputs={}
),
start_at=0.0,
)
node = _SampleNode(
node_id="node-1",
config=_build_node_data(),
data=_build_node_data(),
graph_init_params=init_params,
graph_runtime_state=runtime_state,
)
@@ -140,7 +144,7 @@ def test_node_hydration_preserves_compatibility_extra_fields():
node = _SampleNode(
node_id="node-1",
config=node_config["data"],
data=node_config["data"],
graph_init_params=init_params,
graph_runtime_state=runtime_state,
)
@@ -49,7 +49,7 @@ def document_extractor_node(graph_init_params):
http_client = Mock()
node = DocumentExtractorNode(
node_id="test_node_id",
config=node_data,
data=node_data,
graph_init_params=graph_init_params,
graph_runtime_state=Mock(),
http_client=http_client,
@@ -188,10 +188,20 @@ def test_run_extract_text(
if mime_type == "application/pdf":
mock_pdf_extract = Mock(return_value=expected_text[0])
monkeypatch.setattr("graphon.nodes.document_extractor.node._extract_text_from_pdf", mock_pdf_extract)
if extension:
monkeypatch.setattr(
"graphon.nodes.document_extractor.node._extract_text_by_file_extension", mock_pdf_extract
)
else:
monkeypatch.setattr("graphon.nodes.document_extractor.node._extract_text_by_mime_type", mock_pdf_extract)
elif mime_type.startswith("application/vnd.openxmlformats"):
mock_docx_extract = Mock(return_value=expected_text[0])
monkeypatch.setattr("graphon.nodes.document_extractor.node._extract_text_from_docx", mock_docx_extract)
if extension:
monkeypatch.setattr(
"graphon.nodes.document_extractor.node._extract_text_by_file_extension", mock_docx_extract
)
else:
monkeypatch.setattr("graphon.nodes.document_extractor.node._extract_text_by_mime_type", mock_docx_extract)
result = document_extractor_node._run()
@@ -439,13 +449,18 @@ def test_extract_text_from_file_routes_excel_inputs(document_extractor_node, ext
file.extension = extension
file.mime_type = mime_type
extract_patch_target = (
"graphon.nodes.document_extractor.node._extract_text_by_file_extension"
if extension
else "graphon.nodes.document_extractor.node._extract_text_by_mime_type"
)
with (
patch(
"graphon.nodes.document_extractor.node._download_file_content",
return_value=b"excel",
),
patch(
"graphon.nodes.document_extractor.node._extract_text_from_excel",
extract_patch_target,
return_value="excel text",
) as mock_extract,
):
@@ -456,7 +471,6 @@ def test_extract_text_from_file_routes_excel_inputs(document_extractor_node, ext
)
assert result == "excel text"
mock_extract.assert_called_once_with(b"excel")
def test_extract_text_from_file_rejects_missing_extension_and_mime_type(document_extractor_node):
@@ -29,7 +29,7 @@ def _build_if_else_node(
node_id=str(uuid.uuid4()),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
config=node_data if isinstance(node_data, IfElseNodeData) else IfElseNodeData.model_validate(node_data),
data=node_data if isinstance(node_data, IfElseNodeData) else IfElseNodeData.model_validate(node_data),
)
@@ -48,7 +48,7 @@ def test_execute_if_else_result_true():
)
# construct variable pool
pool = VariablePool(system_variables=build_system_variables(user_id="aaa", files=[]), user_inputs={})
pool = VariablePool.from_bootstrap(system_variables=build_system_variables(user_id="aaa", files=[]), user_inputs={})
pool.add(["start", "array_contains"], ["ab", "def"])
pool.add(["start", "array_not_contains"], ["ac", "def"])
pool.add(["start", "contains"], "cabcde")
@@ -148,7 +148,7 @@ def test_execute_if_else_result_false():
)
# construct variable pool
pool = VariablePool(
pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="aaa", files=[]),
user_inputs={},
environment_variables=[],
@@ -305,7 +305,7 @@ def test_execute_if_else_boolean_conditions(condition: Condition):
)
# construct variable pool with boolean values
pool = VariablePool(
pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(files=[], user_id="aaa"),
)
pool.add(["start", "bool_true"], True)
@@ -359,7 +359,7 @@ def test_execute_if_else_boolean_false_conditions():
)
# construct variable pool with boolean values
pool = VariablePool(
pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(files=[], user_id="aaa"),
)
pool.add(["start", "bool_true"], True)
@@ -424,7 +424,7 @@ def test_execute_if_else_boolean_cases_structure():
)
# construct variable pool with boolean values
pool = VariablePool(
pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(files=[], user_id="aaa"),
)
pool.add(["start", "bool_true"], True)
@@ -22,7 +22,7 @@ from graphon.variables import ArrayFileSegment
def _build_list_operator_node(node_data: ListOperatorNodeData, graph_init_params) -> ListOperatorNode:
return ListOperatorNode(
node_id="test_node_id",
config=node_data,
data=node_data,
graph_init_params=graph_init_params,
graph_runtime_state=MagicMock(),
)
@@ -31,7 +31,7 @@ def make_start_node(user_inputs, variables):
return StartNode(
node_id="start",
config=node_data,
data=node_data,
graph_init_params=build_test_graph_init_params(
workflow_id="wf",
graph_config={},
@@ -260,7 +260,7 @@ def test_start_node_outputs_full_variable_pool_snapshot():
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
node = StartNode(
node_id="start",
config=node_data,
data=node_data,
graph_init_params=build_test_graph_init_params(
workflow_id="wf",
graph_config={},
@@ -99,7 +99,7 @@ def tool_node(monkeypatch) -> ToolNode:
call_depth=0,
)
variable_pool = VariablePool(system_variables=build_system_variables(user_id="user-id"))
variable_pool = VariablePool.from_bootstrap(system_variables=build_system_variables(user_id="user-id"))
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=0.0)
config = graph_config["nodes"][0]
@@ -110,7 +110,7 @@ def tool_node(monkeypatch) -> ToolNode:
node = ToolNode(
node_id="node-instance",
config=ToolNodeData.model_validate(config["data"]),
data=ToolNodeData.model_validate(config["data"]),
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
tool_file_manager_factory=tool_file_manager_factory,
@@ -44,7 +44,7 @@ def test_trigger_event_node_run_populates_trigger_info_metadata() -> None:
init_params, runtime_state = _build_context(graph_config={})
node = TriggerEventNode(
node_id="node-1",
config=_build_node_data(),
data=_build_node_data(),
graph_init_params=init_params,
graph_runtime_state=runtime_state,
)
@@ -52,7 +52,7 @@ def create_webhook_node(
node = TriggerWebhookNode(
node_id="webhook-node-1",
config=webhook_data,
data=webhook_data,
graph_init_params=graph_init_params,
graph_runtime_state=runtime_state,
)
@@ -44,7 +44,7 @@ def create_webhook_node(webhook_data: WebhookData, variable_pool: VariablePool)
)
node = TriggerWebhookNode(
node_id="1",
config=webhook_data,
data=webhook_data,
graph_init_params=graph_init_params,
graph_runtime_state=runtime_state,
)
@@ -470,7 +470,7 @@ class TestDifyNodeFactoryCreateNode:
matched_node_class.assert_called_once()
kwargs = matched_node_class.call_args.kwargs
assert kwargs["node_id"] == "node-id"
_assert_typed_node_config(kwargs["config"], node_id="node-id", node_type=BuiltinNodeTypes.START, version="9")
_assert_typed_node_config(kwargs["data"], node_id="node-id", node_type=BuiltinNodeTypes.START, version="9")
assert kwargs["graph_init_params"] is sentinel.graph_init_params
assert kwargs["graph_runtime_state"] is factory.graph_runtime_state
latest_node_class.assert_not_called()
@@ -490,7 +490,7 @@ class TestDifyNodeFactoryCreateNode:
latest_node_class.assert_called_once()
kwargs = latest_node_class.call_args.kwargs
assert kwargs["node_id"] == "node-id"
_assert_typed_node_config(kwargs["config"], node_id="node-id", node_type=BuiltinNodeTypes.START, version="9")
_assert_typed_node_config(kwargs["data"], node_id="node-id", node_type=BuiltinNodeTypes.START, version="9")
assert kwargs["graph_init_params"] is sentinel.graph_init_params
assert kwargs["graph_runtime_state"] is factory.graph_runtime_state
@@ -528,7 +528,7 @@ class TestDifyNodeFactoryCreateNode:
assert result is created_node
kwargs = constructor.call_args.kwargs
assert kwargs["node_id"] == "node-id"
_assert_typed_node_config(kwargs["config"], node_id="node-id", node_type=node_type)
_assert_typed_node_config(kwargs["data"], node_id="node-id", node_type=node_type)
assert kwargs["graph_init_params"] is sentinel.graph_init_params
assert kwargs["graph_runtime_state"] is factory.graph_runtime_state
@@ -707,7 +707,7 @@ class TestDifyNodeFactoryCreateNode:
constructor_kwargs = constructor.call_args.kwargs
assert constructor_kwargs["node_id"] == "node-id"
_assert_typed_node_config(constructor_kwargs["config"], node_id="node-id", node_type=node_type)
_assert_typed_node_config(constructor_kwargs["data"], node_id="node-id", node_type=node_type)
assert constructor_kwargs["graph_init_params"] is sentinel.graph_init_params
assert constructor_kwargs["graph_runtime_state"] is factory.graph_runtime_state
assert constructor_kwargs["credentials_provider"] is sentinel.credentials_provider
@@ -41,7 +41,7 @@ from models.utils.file_input_compat import rebuild_serialized_graph_files_withou
@pytest.fixture
def pool():
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(
variable_pool,
build_system_variables(
@@ -82,7 +82,7 @@ def test_get_file_attribute(pool, file):
class TestVariablePool:
def test_constructor(self):
pool = VariablePool()
pool = VariablePool.from_bootstrap()
assert pool.variable_dictionary == defaultdict(dict)
complex_system_vars = build_system_variables(
@@ -110,7 +110,7 @@ class TestVariablePool:
assert pool.get([CONVERSATION_VARIABLE_NODE_ID, "conv_var_1"]) is not None
def test_constructor_loads_legacy_bootstrap_kwargs(self):
pool = VariablePool(
pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(user_id="test_user_id"),
environment_variables=[StringVariable(name="env_var", value="env-value")],
conversation_variables=[StringVariable(name="conv_var", value="conv-value")],
@@ -139,7 +139,7 @@ class TestVariablePool:
conversation_id="test_conv_id",
dialogue_count=5,
)
pool = VariablePool()
pool = VariablePool.from_bootstrap()
add_variables_to_pool(pool, sys_var)
system_values = system_variables_to_mapping(sys_var)
@@ -259,7 +259,7 @@ class TestVariablePoolSerialization:
}
# Create VariablePool
pool = VariablePool()
pool = VariablePool.from_bootstrap()
add_variables_to_pool(pool, system_vars)
add_variables_to_pool(pool, env_vars)
add_variables_to_pool(pool, conv_vars)
@@ -302,7 +302,7 @@ class TestVariablePoolSerialization:
conversation_id="test_conv_id",
dialogue_count=5,
)
pool = VariablePool()
pool = VariablePool.from_bootstrap()
add_variables_to_pool(pool, sys_vars)
json = pool.model_dump_json()
pool2 = VariablePool.model_validate_json(json)
@@ -437,7 +437,7 @@ class TestVariablePoolSerialization:
def test_get_attr():
vp = VariablePool()
vp = VariablePool.from_bootstrap()
value = {"output": StringSegment(value="hello")}
vp.add(["node", "name"], value)
@@ -55,7 +55,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_to_variable_pool_with_system_variables(self):
"""Test mapping system variables from user inputs to variable pool."""
# Initialize variable pool with system variables
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="test_user_id",
app_id="test_app_id",
@@ -128,7 +128,7 @@ class TestWorkflowEntry:
return NodeConfigDictAdapter.validate_python(node_config)
workflow = StubWorkflow()
variable_pool = VariablePool(system_variables=default_system_variables(), user_inputs={})
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables(), user_inputs={})
expected_limits = CodeNodeLimits(
max_string_length=dify_config.CODE_MAX_STRING_LENGTH,
max_number=dify_config.CODE_MAX_NUMBER,
@@ -157,7 +157,7 @@ class TestWorkflowEntry:
"""Test mapping environment variables from user inputs to variable pool."""
# Initialize variable pool with environment variables
env_var = StringVariable(name="API_KEY", value="existing_key")
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
environment_variables=[env_var],
user_inputs={},
@@ -198,7 +198,7 @@ class TestWorkflowEntry:
"""Test mapping conversation variables from user inputs to variable pool."""
# Initialize variable pool with conversation variables
conv_var = StringVariable(name="last_message", value="Hello")
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
conversation_variables=[conv_var],
user_inputs={},
@@ -239,7 +239,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_to_variable_pool_with_regular_variables(self):
"""Test mapping regular node variables from user inputs to variable pool."""
# Initialize empty variable pool
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -281,7 +281,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_with_file_handling(self):
"""Test mapping file inputs from user inputs to variable pool."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -340,7 +340,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_missing_variable_error(self):
"""Test that mapping raises error when required variable is missing."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -366,7 +366,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_with_alternative_key_format(self):
"""Test mapping with alternative key format (without node prefix)."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -396,7 +396,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_with_complex_selectors(self):
"""Test mapping with complex node variable keys."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -432,7 +432,7 @@ class TestWorkflowEntry:
def test_mapping_user_inputs_invalid_node_variable(self):
"""Test that mapping handles invalid node variable format."""
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=default_system_variables(),
user_inputs={},
)
@@ -463,7 +463,7 @@ class TestWorkflowEntry:
env_var = StringVariable(name="API_KEY", value="existing_key")
conv_var = StringVariable(name="session_id", value="session123")
variable_pool = VariablePool(
variable_pool = VariablePool.from_bootstrap(
system_variables=build_system_variables(
user_id="test_user",
app_id="test_app",
@@ -165,7 +165,7 @@ class TestWorkflowChildEngineBuilder:
_ = graph_config
node = node_cls(
node_id=root_node_id,
config=BaseNodeData(
data=BaseNodeData(
type=node_cls.node_type,
title="Child Model",
),
@@ -334,7 +334,7 @@ class TestWorkflowEntrySingleStepRun:
def extract_variable_selector_to_variable_mapping(**_kwargs):
return {}
variable_pool = VariablePool(system_variables=default_system_variables(), user_inputs={})
variable_pool = VariablePool.from_bootstrap(system_variables=default_system_variables(), user_inputs={})
variable_loader = MagicMock()
variable_loader.load_variables.return_value = [
StringVariable(
@@ -644,7 +644,9 @@ class TestWorkflowEntryHelpers:
with (
patch.object(workflow_entry, "default_system_variables", return_value=sentinel.system_variables),
patch.object(workflow_entry, "VariablePool", return_value=sentinel.variable_pool) as variable_pool_cls,
patch(
"graphon.runtime.VariablePool.from_bootstrap", return_value=sentinel.variable_pool
) as variable_pool_cls,
patch.object(workflow_entry, "add_variables_to_pool") as add_variables_to_pool,
patch.object(
workflow_entry, "DifyGraphInitContext", return_value=sentinel.graph_init_context
@@ -732,7 +734,7 @@ class TestWorkflowEntryHelpers:
with (
patch.object(workflow_entry, "default_system_variables", return_value=sentinel.system_variables),
patch.object(workflow_entry, "VariablePool", return_value=sentinel.variable_pool),
patch("graphon.runtime.VariablePool.from_bootstrap", return_value=sentinel.variable_pool),
patch.object(workflow_entry, "add_variables_to_pool"),
patch.object(workflow_entry, "DifyGraphInitContext", return_value=sentinel.graph_init_context),
patch.object(workflow_entry, "GraphRuntimeState", return_value=sentinel.graph_runtime_state),
@@ -2103,7 +2103,7 @@ class TestSetupVariablePool:
# Act
with (
patch("services.workflow_service.VariablePool") as MockPool,
patch("services.workflow_service.VariablePool.from_bootstrap") as mock_pool_from_bootstrap,
patch("services.workflow_service.build_system_variables") as mock_build_system_variables,
patch("services.workflow_service.build_bootstrap_variables") as mock_build_bootstrap_variables,
patch("services.workflow_service.add_variables_to_pool") as mock_add_variables_to_pool,
@@ -2122,14 +2122,14 @@ class TestSetupVariablePool:
)
# Assert — start nodes should build bootstrap variables and attach node inputs.
MockPool.assert_called_once_with()
mock_pool_from_bootstrap.assert_called_once_with()
mock_build_system_variables.assert_called_once()
mock_add_variables_to_pool.assert_called_once_with(
MockPool.return_value,
mock_pool_from_bootstrap.return_value,
mock_build_bootstrap_variables.return_value,
)
mock_add_node_inputs_to_pool.assert_called_once_with(
MockPool.return_value,
mock_pool_from_bootstrap.return_value,
node_id="start-node",
inputs={"k": "v"},
)
@@ -2142,7 +2142,7 @@ class TestSetupVariablePool:
# Act
with (
patch("services.workflow_service.VariablePool") as MockPool,
patch("services.workflow_service.VariablePool.from_bootstrap") as mock_pool_from_bootstrap,
patch("services.workflow_service.default_system_variables") as mock_default_system_variables,
patch("services.workflow_service.build_bootstrap_variables") as mock_build_bootstrap_variables,
patch("services.workflow_service.add_variables_to_pool") as mock_add_variables_to_pool,
@@ -2162,9 +2162,9 @@ class TestSetupVariablePool:
# Assert — default system variables should be used and node inputs should not be added.
mock_default_system_variables.assert_called_once()
MockPool.assert_called_once_with()
mock_pool_from_bootstrap.assert_called_once_with()
mock_add_variables_to_pool.assert_called_once_with(
MockPool.return_value,
mock_pool_from_bootstrap.return_value,
mock_build_bootstrap_variables.return_value,
)
mock_add_node_inputs_to_pool.assert_not_called()
@@ -2180,7 +2180,7 @@ class TestSetupVariablePool:
# Act
with (
patch("services.workflow_service.VariablePool") as MockPool,
patch("services.workflow_service.VariablePool.from_bootstrap") as mock_pool_from_bootstrap,
patch("services.workflow_service.build_system_variables") as mock_build_system_variables,
patch("services.workflow_service.build_bootstrap_variables"),
patch("services.workflow_service.add_variables_to_pool"),
@@ -2199,7 +2199,7 @@ class TestSetupVariablePool:
)
# Assert — chatflow system variables should include query, conversation_id and dialogue_count.
MockPool.assert_called_once_with()
mock_pool_from_bootstrap.assert_called_once_with()
system_variable_values = mock_build_system_variables.call_args.args[0]
assert system_variable_values["query"] == "what is AI?"
assert system_variable_values["conversation_id"] == "conv-abc"
@@ -2513,7 +2513,7 @@ class TestWorkflowServiceDraftExecution:
patch("services.workflow_service.db"),
patch("services.workflow_service.Session"),
patch("services.workflow_service.WorkflowDraftVariableService"),
patch("services.workflow_service.VariablePool") as mock_pool_cls,
patch("services.workflow_service.VariablePool.from_bootstrap") as mock_pool_cls,
patch("services.workflow_service.default_system_variables") as mock_default_system_variables,
patch("services.workflow_service.build_bootstrap_variables") as mock_build_bootstrap_variables,
patch("services.workflow_service.add_variables_to_pool") as mock_add_variables_to_pool,
@@ -2737,7 +2737,7 @@ class TestWorkflowServiceHumanInputOperations:
patch("services.workflow_service.db"),
patch("services.workflow_service.Session"),
patch("services.workflow_service.WorkflowDraftVariableService"),
patch("services.workflow_service.VariablePool") as mock_pool_cls,
patch("services.workflow_service.VariablePool.from_bootstrap") as mock_pool_cls,
patch("services.workflow_service.DraftVarLoader"),
patch("services.workflow_service.HumanInputNode.extract_variable_selector_to_variable_mapping"),
patch("services.workflow_service.load_into_variable_pool"),
@@ -2845,7 +2845,7 @@ class TestWorkflowServiceFreeNodeExecution:
mock_node_cls.validate_node_data.assert_called_once_with(sentinel.adapted_node_data)
mock_node_cls.assert_called_once_with(
node_id="n-1",
config=sentinel.node_data,
data=sentinel.node_data,
graph_init_params=mock_graph_init_context_cls.return_value.to_graph_init_params.return_value,
graph_runtime_state=ANY,
runtime=mock_runtime_cls.return_value,
@@ -124,7 +124,7 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext:
call_depth=0,
workflow_execution_id="run-1",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
runtime_state = GraphRuntimeState(variable_pool=VariablePool.from_bootstrap(), start_at=0.0)
runtime_state.register_paused_node("node-1")
runtime_state.outputs = {"result": "value"}
wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity)
@@ -230,7 +230,7 @@ def _build_resumption_context_additional(task_id: str) -> WorkflowResumptionCont
call_depth=0,
workflow_execution_id="run-1",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
runtime_state = GraphRuntimeState(variable_pool=VariablePool.from_bootstrap(), start_at=0.0)
runtime_state.outputs = {"answer": "ok"}
wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity)
return WorkflowResumptionContext(
@@ -67,7 +67,7 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext:
call_depth=0,
workflow_execution_id="run-1",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
runtime_state = GraphRuntimeState(variable_pool=VariablePool.from_bootstrap(), start_at=0.0)
runtime_state.outputs = {"answer": "ok"}
wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity)
return WorkflowResumptionContext(
@@ -102,7 +102,7 @@ def test_dispatch_human_input_email_task_replaces_body_variables(monkeypatch: py
recipients=[task_module._EmailRecipient(email="user@example.com", token="token-1")],
)
variable_pool = task_module.VariablePool()
variable_pool = task_module.VariablePool.from_bootstrap()
variable_pool.add(["node1", "value"], "OK")
monkeypatch.setattr(task_module, "mail", mail)
+1 -1
View File
@@ -62,7 +62,7 @@ def build_test_variable_pool(
node_id: str | None = None,
inputs: Mapping[str, Any] | None = None,
) -> VariablePool:
variable_pool = VariablePool()
variable_pool = VariablePool.from_bootstrap()
add_variables_to_pool(variable_pool, variables)
if node_id is not None and inputs is not None:
add_node_inputs_to_pool(variable_pool, node_id=node_id, inputs=inputs)
Generated
+4 -4
View File
@@ -1594,7 +1594,7 @@ requires-dist = [
{ name = "gmpy2", specifier = ">=2.3.0" },
{ name = "google-api-python-client", specifier = ">=2.194.0" },
{ name = "google-cloud-aiplatform", specifier = ">=1.148.1,<2.0.0" },
{ name = "graphon", specifier = "~=0.2.2" },
{ name = "graphon", specifier = "~=0.3.0" },
{ name = "gunicorn", specifier = ">=25.3.0" },
{ name = "httpx", extras = ["socks"], specifier = ">=0.28.1,<1.0.0" },
{ name = "httpx-sse", specifier = "~=0.4.0" },
@@ -2937,7 +2937,7 @@ httpx = [
[[package]]
name = "graphon"
version = "0.2.2"
version = "0.3.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "charset-normalizer" },
@@ -2958,9 +2958,9 @@ dependencies = [
{ name = "unstructured", extra = ["docx", "epub", "md", "ppt", "pptx"] },
{ name = "webvtt-py" },
]
sdist = { url = "https://files.pythonhosted.org/packages/08/50/e745a79c5f742f88f6011a1f7c9ba2c2f9cc1beedd982f0b192f1ab8c748/graphon-0.2.2.tar.gz", hash = "sha256:141f0de536171850f1af6f738dc66f0285aadd3c097f1dad2a038636789e0aa5", size = 236360, upload-time = "2026-04-17T08:52:28.047Z" }
sdist = { url = "https://files.pythonhosted.org/packages/bf/62/83593d6e7a139ff124711ea05882cadca7065c11a38763aa9360d7e76804/graphon-0.3.0.tar.gz", hash = "sha256:cd38f842ae3dcfa956428b952efbe2a3ea9c1581446647142accbbdeb638b876", size = 241176, upload-time = "2026-04-21T15:18:48.291Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/de/89/a6340afdaf5169d17a318e00fc685fb67ed99baa602c2cbbbf6af6a76096/graphon-0.2.2-py3-none-any.whl", hash = "sha256:754e544d08779138f99eac6547ab08559463680e2c76488b05e1c978210392b4", size = 340808, upload-time = "2026-04-17T08:52:26.5Z" },
{ url = "https://files.pythonhosted.org/packages/b3/f7/81ee8f0368aa6a2d47f97fecc5d4a12865c987906798cbddd0e3b8387f33/graphon-0.3.0-py3-none-any.whl", hash = "sha256:9cca45ebab2a79fd4d04432f55b5b962e9e4f34fa037cc20fee7f18ec80eaa5d", size = 348486, upload-time = "2026-04-21T15:18:46.737Z" },
]
[[package]]