From a176056cacf44ea24b9be73bbca922fb42dc1163 Mon Sep 17 00:00:00 2001 From: fatelei Date: Thu, 20 Nov 2025 17:04:58 +0800 Subject: [PATCH 1/2] feat: support seperate memory --- api/core/memory/node_scoped_memory.py | 221 ++++++++ .../entities/advanced_prompt_entities.py | 5 + api/core/workflow/nodes/llm/llm_utils.py | 23 + api/core/workflow/nodes/llm/node.py | 60 ++- .../nodes/llm/test_independent_memory.py | 492 ++++++++++++++++++ .../nodes/_base/components/memory-config.tsx | 43 ++ web/app/components/workflow/types.ts | 3 + 7 files changed, 835 insertions(+), 12 deletions(-) create mode 100644 api/core/memory/node_scoped_memory.py create mode 100644 api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py diff --git a/api/core/memory/node_scoped_memory.py b/api/core/memory/node_scoped_memory.py new file mode 100644 index 00000000000000..6894450360dad9 --- /dev/null +++ b/api/core/memory/node_scoped_memory.py @@ -0,0 +1,221 @@ +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass, field +from typing import Any +from uuid import UUID, uuid5 + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from core.model_manager import ModelInstance +from core.model_runtime.entities import ( + AssistantPromptMessage, + PromptMessage, + PromptMessageRole, + TextPromptMessageContent, + UserPromptMessage, +) +from core.variables.segments import ObjectSegment +from core.variables.types import SegmentType +from extensions.ext_database import db +from factories import variable_factory +from models.workflow import ConversationVariable + +# A stable namespace to derive deterministic IDs for node-scoped memories. +# Using uuid5 to keep (conversation_id, node_id) mapping stable across runs. +# NOTE: UUIDs must contain only hexadecimal characters; avoid letters beyond 'f'. +NODE_SCOPED_MEMORY_NS = UUID("00000000-0000-0000-0000-000000000000") + + +@dataclass +class _HistoryItem: + role: str + text: str + + +@dataclass +class NodeScopedMemory: + """A per-node conversation memory persisted in ConversationVariable. + + - Keyed by (conversation_id, node_id) + - Value is stored as a conversation variable named _llm_mem. + - Structure (JSON): {"version": 1, "history": [{"role": "user"|"assistant", "text": "..."}, ...]} + """ + + app_id: str + conversation_id: str + node_id: str + model_instance: ModelInstance + + _loaded: bool = field(default=False, init=False) + _history: list[_HistoryItem] = field(default_factory=list, init=False) + + @property + def variable_name(self) -> str: + return f"_llm_mem.{self.node_id}" + + @property + def variable_id(self) -> str: + # Deterministic id so we can upsert by (id, conversation_id) + return str(uuid5(NODE_SCOPED_MEMORY_NS, f"{self.conversation_id}:{self.node_id}:llmmem")) + + # ------------ Persistence helpers ------------ + def _load_if_needed(self) -> None: + if self._loaded: + return + stmt = select(ConversationVariable).where( + ConversationVariable.id == self.variable_id, + ConversationVariable.conversation_id == self.conversation_id, + ) + with Session(db.engine, expire_on_commit=False) as session: + row = session.scalar(stmt) + if not row: + self._history = [] + self._loaded = True + return + variable = row.to_variable() + value = variable.value if isinstance(variable.value, dict) else {} + hist = value.get("history", []) if isinstance(value, dict) else [] + parsed: list[_HistoryItem] = [] + for item in hist: + try: + role = str(item.get("role", "")) + text = str(item.get("text", "")) + except Exception: + role, text = "", "" + if role and text: + parsed.append(_HistoryItem(role=role, text=text)) + self._history = parsed + self._loaded = True + + def _dump_variable(self) -> Any: + data = { + "version": 1, + "history": [{"role": item.role, "text": item.text} for item in self._history if item.text], + } + segment = ObjectSegment(value=data, value_type=SegmentType.OBJECT) + variable = variable_factory.segment_to_variable( + segment=segment, + selector=["conversation", self.variable_name], + id=self.variable_id, + name=self.variable_name, + description="LLM node-scoped memory", + ) + return variable + + def save(self) -> None: + variable = self._dump_variable() + with Session(db.engine) as session: + # Upsert by (id, conversation_id) + existing = session.scalar( + select(ConversationVariable).where( + ConversationVariable.id == self.variable_id, + ConversationVariable.conversation_id == self.conversation_id, + ) + ) + if existing: + existing.data = variable.model_dump_json() + else: + obj = ConversationVariable.from_variable( + app_id=self.app_id, conversation_id=self.conversation_id, variable=variable + ) + session.add(obj) + session.commit() + + # ------------ Public API expected by LLM node ------------ + def get_history_prompt_messages( + self, *, max_token_limit: int = 2000, message_limit: int | None = None + ) -> Sequence[PromptMessage]: + self._load_if_needed() + + # Optionally limit by message count (pairs flattened) + items: list[_HistoryItem] = list(self._history) + if message_limit and message_limit > 0: + # message_limit roughly means last N items (not pairs) to keep simple and efficient + items = items[-min(message_limit, len(items)) :] + + def to_messages(hist: list[_HistoryItem]) -> list[PromptMessage]: + msgs: list[PromptMessage] = [] + for it in hist: + if it.role == PromptMessageRole.USER.value: + # Persisted node memory only stores text; inject as plain text content + msgs.append(UserPromptMessage(content=it.text)) + elif it.role == PromptMessageRole.ASSISTANT.value: + msgs.append(AssistantPromptMessage(content=it.text)) + return msgs + + messages = to_messages(items) + # Token-based pruning from oldest + if messages: + tokens = self.model_instance.get_llm_num_tokens(messages) + while tokens > max_token_limit and len(messages) > 1: + messages.pop(0) + tokens = self.model_instance.get_llm_num_tokens(messages) + return messages + + def get_history_prompt_text( + self, + *, + human_prefix: str = "Human", + ai_prefix: str = "Assistant", + max_token_limit: int = 2000, + message_limit: int | None = None, + ) -> str: + self._load_if_needed() + items: list[_HistoryItem] = list(self._history) + if message_limit and message_limit > 0: + items = items[-min(message_limit, len(items)) :] + + # Build messages to reuse token counting logic + messages: list[PromptMessage] = [] + for it in items: + role_name = ( + PromptMessageRole.USER + if it.role == PromptMessageRole.USER.value + else (PromptMessageRole.ASSISTANT if it.role == PromptMessageRole.ASSISTANT.value else None) + ) + if role_name is None: + continue + prefix = human_prefix if role_name == PromptMessageRole.USER else ai_prefix + messages.append( + UserPromptMessage(content=f"{prefix}: {it.text}") + if role_name == PromptMessageRole.USER + else AssistantPromptMessage(content=f"{prefix}: {it.text}") + ) + + if messages: + tokens = self.model_instance.get_llm_num_tokens(messages) + while tokens > max_token_limit and len(messages) > 1: + messages.pop(0) + tokens = self.model_instance.get_llm_num_tokens(messages) + + # Convert back to the required text format + lines: list[str] = [] + for m in messages: + if m.role == PromptMessageRole.USER: + prefix = human_prefix + elif m.role == PromptMessageRole.ASSISTANT: + prefix = ai_prefix + else: + continue + if isinstance(m.content, list): + # Only text content was saved in this minimal implementation + texts = [c.data for c in m.content if isinstance(c, TextPromptMessageContent)] + text = "\n".join(texts) + else: + text = str(m.content) + lines.append(f"{prefix}: {text}") + return "\n".join(lines) + + def append_exchange(self, *, user_text: str | None, assistant_text: str | None) -> None: + self._load_if_needed() + if user_text: + self._history.append(_HistoryItem(role=PromptMessageRole.USER.value, text=user_text)) + if assistant_text: + self._history.append(_HistoryItem(role=PromptMessageRole.ASSISTANT.value, text=assistant_text)) + + def clear(self) -> None: + self._history = [] + self._loaded = True + self.save() diff --git a/api/core/prompt/entities/advanced_prompt_entities.py b/api/core/prompt/entities/advanced_prompt_entities.py index 7094633093f413..0c418ee6d1b6af 100644 --- a/api/core/prompt/entities/advanced_prompt_entities.py +++ b/api/core/prompt/entities/advanced_prompt_entities.py @@ -48,3 +48,8 @@ class WindowConfig(BaseModel): role_prefix: RolePrefix | None = None window: WindowConfig query_prompt_template: str | None = None + # Memory scope: shared (default, uses TokenBufferMemory), or independent + # (per-node, persisted in ConversationVariable) + scope: Literal["shared", "independent"] = "shared" + # If true, clear the per-node memory after this node finishes execution (only applies to independent scope) + clear_after_execution: bool = False diff --git a/api/core/workflow/nodes/llm/llm_utils.py b/api/core/workflow/nodes/llm/llm_utils.py index 0c545469bc66fa..408c4710bcb843 100644 --- a/api/core/workflow/nodes/llm/llm_utils.py +++ b/api/core/workflow/nodes/llm/llm_utils.py @@ -8,6 +8,7 @@ from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity from core.entities.provider_entities import QuotaUnit from core.file.models import File +from core.memory.node_scoped_memory import NodeScopedMemory from core.memory.token_buffer_memory import TokenBufferMemory from core.model_manager import ModelInstance, ModelManager from core.model_runtime.entities.llm_entities import LLMUsage @@ -107,6 +108,28 @@ def fetch_memory( return memory +def fetch_node_scoped_memory( + variable_pool: VariablePool, + *, + app_id: str, + node_id: str, + model_instance: ModelInstance, +) -> NodeScopedMemory | None: + """Factory for per-node memory based on conversation scope. + + Returns None if no conversation_id is present in the variable pool. + """ + conversation_id_variable = variable_pool.get(["sys", SystemVariableKey.CONVERSATION_ID]) + if not isinstance(conversation_id_variable, StringSegment): + return None + return NodeScopedMemory( + app_id=app_id, + conversation_id=conversation_id_variable.value, + node_id=node_id, + model_instance=model_instance, + ) + + def deduct_llm_quota(tenant_id: str, model_instance: ModelInstance, usage: LLMUsage): provider_model_bundle = model_instance.provider_model_bundle provider_configuration = provider_model_bundle.configuration diff --git a/api/core/workflow/nodes/llm/node.py b/api/core/workflow/nodes/llm/node.py index 06c9beaed20dbe..b7d8c3fce7ddf7 100644 --- a/api/core/workflow/nodes/llm/node.py +++ b/api/core/workflow/nodes/llm/node.py @@ -5,14 +5,13 @@ import re import time from collections.abc import Generator, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Any, Literal, Protocol from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity from core.file import FileType, file_manager from core.helper.code_executor import CodeExecutor, CodeLanguage from core.llm_generator.output_parser.errors import OutputParserError from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output -from core.memory.token_buffer_memory import TokenBufferMemory from core.model_manager import ModelInstance, ModelManager from core.model_runtime.entities import ( ImagePromptMessageContent, @@ -100,6 +99,26 @@ logger = logging.getLogger(__name__) +class ChatHistoryMemory(Protocol): + def get_history_prompt_messages( + self, + *, + max_token_limit: int = 2000, + message_limit: int | None = None, + ) -> Sequence[PromptMessage]: + ... + + def get_history_prompt_text( + self, + *, + human_prefix: str = "Human", + ai_prefix: str = "Assistant", + max_token_limit: int = 2000, + message_limit: int | None = None, + ) -> str: + ... + + class LLMNode(Node): node_type = NodeType.LLM @@ -215,13 +234,26 @@ def _run(self) -> Generator: tenant_id=self.tenant_id, ) - # fetch memory - memory = llm_utils.fetch_memory( - variable_pool=variable_pool, - app_id=self.app_id, - node_data_memory=self._node_data.memory, - model_instance=model_instance, + # fetch memory (shared) or node-scoped (independent) + independent_scope = ( + self._node_data.memory and getattr(self._node_data.memory, "scope", "shared") == "independent" ) + node_memory = None + memory_shared = None + if independent_scope: + node_memory = llm_utils.fetch_node_scoped_memory( + variable_pool=variable_pool, + app_id=self.app_id, + node_id=self._node_id, + model_instance=model_instance, + ) + else: + memory_shared = llm_utils.fetch_memory( + variable_pool=variable_pool, + app_id=self.app_id, + node_data_memory=self._node_data.memory, + model_instance=model_instance, + ) query: str | None = None if self._node_data.memory: @@ -235,7 +267,7 @@ def _run(self) -> Generator: sys_query=query, sys_files=files, context=context, - memory=memory, + memory=(node_memory if independent_scope else memory_shared), model_config=model_config, prompt_template=self._node_data.prompt_template, memory_config=self._node_data.memory, @@ -289,6 +321,10 @@ def _run(self) -> Generator: else None ) + # Persist node-scoped memory if enabled + if independent_scope and node_memory: + node_memory.clear() + # deduct quota llm_utils.deduct_llm_quota(tenant_id=self.tenant_id, model_instance=model_instance, usage=usage) break @@ -756,7 +792,7 @@ def fetch_prompt_messages( sys_query: str | None = None, sys_files: Sequence["File"], context: str | None = None, - memory: TokenBufferMemory | None = None, + memory: ChatHistoryMemory | None = None, model_config: ModelConfigWithCredentialsEntity, prompt_template: Sequence[LLMNodeChatModelMessage] | LLMNodeCompletionModelPromptTemplate, memory_config: MemoryConfig | None = None, @@ -1297,7 +1333,7 @@ def _calculate_rest_token( def _handle_memory_chat_mode( *, - memory: TokenBufferMemory | None, + memory: ChatHistoryMemory | None, memory_config: MemoryConfig | None, model_config: ModelConfigWithCredentialsEntity, ) -> Sequence[PromptMessage]: @@ -1314,7 +1350,7 @@ def _handle_memory_chat_mode( def _handle_memory_completion_mode( *, - memory: TokenBufferMemory | None, + memory: ChatHistoryMemory | None, memory_config: MemoryConfig | None, model_config: ModelConfigWithCredentialsEntity, ) -> str: diff --git a/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py b/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py new file mode 100644 index 00000000000000..83bfa235d3ffa9 --- /dev/null +++ b/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py @@ -0,0 +1,492 @@ +import types +from collections.abc import Generator, Sequence + +import pytest + +from core.model_runtime.entities.llm_entities import LLMUsage +from core.model_runtime.entities.message_entities import PromptMessage, PromptMessageRole, UserPromptMessage +from core.prompt.entities.advanced_prompt_entities import MemoryConfig +from core.workflow.entities import GraphInitParams +from core.workflow.node_events.node import ModelInvokeCompletedEvent +from core.workflow.nodes.llm.entities import ( + ContextConfig, + LLMNodeChatModelMessage, + LLMNodeData, + ModelConfig, +) +from core.workflow.nodes.llm.node import LLMNode +from core.workflow.runtime import GraphRuntimeState, VariablePool +from core.workflow.system_variable import SystemVariable +from models.enums import UserFrom + + +@pytest.fixture +def graph_init_params() -> GraphInitParams: + return GraphInitParams( + tenant_id="t1", + app_id="app1", + workflow_id="wf1", + graph_config={}, + user_id="u1", + user_from=UserFrom.ACCOUNT, + invoke_from="service-api", + call_depth=0, + ) + + +@pytest.fixture +def graph_runtime_state() -> GraphRuntimeState: + variable_pool = VariablePool( + system_variables=SystemVariable.empty(), + user_inputs={}, + ) + return GraphRuntimeState(variable_pool=variable_pool, start_at=0) + + +@pytest.fixture +def llm_node(graph_init_params: GraphInitParams, graph_runtime_state: GraphRuntimeState) -> LLMNode: + data = LLMNodeData( + title="LLM", + model=ModelConfig(provider="openai", name="gpt-x", mode="chat", completion_params={}), + prompt_template=[LLMNodeChatModelMessage(text="hello", role=PromptMessageRole.SYSTEM, edition_type="basic")], + memory=MemoryConfig( + role_prefix=None, + window=MemoryConfig.WindowConfig(enabled=False), + query_prompt_template=None, + scope="independent", + clear_after_execution=False, + ), + context=ContextConfig(enabled=False), + ) + node_conf = {"id": "n1", "data": data.model_dump()} + node = LLMNode( + id="n1", + config=node_conf, + graph_init_params=graph_init_params, + graph_runtime_state=graph_runtime_state, + ) + node.init_node_data(node_conf["data"]) + return node + + +class _FakeMemory: + def __init__(self, history_messages: Sequence[PromptMessage] | None = None): + self.history_messages = list(history_messages or []) + self.appended = [] + self.saved = False + self.cleared = False + + # Chat-mode API + def get_history_prompt_messages( + self, *, max_token_limit: int = 2000, message_limit: int | None = None + ) -> Sequence[PromptMessage]: + return list(self.history_messages) + + # Completion-mode API (not used in this test file but kept for parity) + def get_history_prompt_text( + self, + *, + human_prefix: str = "Human", + ai_prefix: str = "Assistant", + max_token_limit: int = 2000, + message_limit: int | None = None, + ) -> str: + return "".join(m.content for m in self.history_messages if isinstance(m, UserPromptMessage)) + + def append_exchange(self, *, user_text: str | None, assistant_text: str | None) -> None: + self.appended.append((user_text or "", assistant_text or "")) + + def save(self) -> None: + self.saved = True + + def clear(self) -> None: + self.cleared = True + + +@pytest.fixture +def patch_minimal_runtime(monkeypatch): + # Make ModelManager in fetch_prompt_messages return a stub whose get_model_schema returns a truthy object + class _FakeModelTypeInst: + def get_model_schema(self, model, credentials): + return types.SimpleNamespace(features=[], parameter_rules=[], model_properties={}) + + class _FakeModel: + def __init__(self): + self.model_type_instance = _FakeModelTypeInst() + self.credentials = {} + + class _FakeManager: + def get_model_instance(self, *args, **kwargs): + return _FakeModel() + + from core.workflow.nodes.llm import node as llm_node_mod + + monkeypatch.setattr(llm_node_mod, "ModelManager", _FakeManager) + + # Mock _fetch_model_config to avoid database calls + def _fake_fetch_model_config(*, node_data_model, tenant_id): + from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity + from core.entities.provider_configuration import ProviderModelBundle, ProviderConfiguration + from core.entities.provider_entities import CustomConfiguration, SystemConfiguration + from core.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType, FREETEXT + from models.provider import ProviderType + + # Create minimal fake objects for the required fields + model_schema = types.SimpleNamespace( + label=types.SimpleNamespace(en="gpt-x"), + model="gpt-x", + model_type=ModelType.LLM, + features=[], + fetch_from=FetchFrom.CUSTOMIZABLE_MODEL, + model_properties={}, + parameter_rules=[], + pricing=None, + deprecation=None, + icon=None, + icon_large=None, + background=None, + help=None, + ) + + provider_model_bundle = ProviderModelBundle( + configuration=ProviderConfiguration( + tenant_id=tenant_id, + provider=None, + preferred_provider_type=ProviderType.CUSTOM, + using_provider_type=ProviderType.CUSTOM, + system_configuration=SystemConfiguration(enabled=False), + custom_configuration=CustomConfiguration(provider=None), + model_settings=[], + ), + model_type_instance=None, + ) + + return _FakeModel(), ModelConfigWithCredentialsEntity( + provider="openai", + model="gpt-x", + model_schema=model_schema, + mode="chat", + provider_model_bundle=provider_model_bundle, + credentials={}, + parameters={}, + ) + + monkeypatch.setattr(llm_node_mod.LLMNode, "_fetch_model_config", staticmethod(_fake_fetch_model_config)) + + # Avoid quota side effects + from core.workflow.nodes.llm import llm_utils as utils_mod + + monkeypatch.setattr(utils_mod, "deduct_llm_quota", lambda *a, **k: None) + + +def _fake_invoke_llm_capture(prompt_messages_out: list[Sequence[PromptMessage]]): + def _invoke(**kwargs) -> Generator: + # capture prompt_messages + prompt_messages_out.append(list(kwargs["prompt_messages"])) + # produce a completed event + usage = LLMUsage.empty_usage() + yield ModelInvokeCompletedEvent(text="ANS", usage=usage, finish_reason=None, reasoning_content="") + + return _invoke + + +def test_independent_memory_injected_into_prompt( + monkeypatch, llm_node: LLMNode, graph_runtime_state: GraphRuntimeState, patch_minimal_runtime +): + # sys.query available so node can append on completion later + graph_runtime_state.variable_pool.add(["sys", "query"], "Q") + + # Fake node memory with one history message + fake_mem = _FakeMemory(history_messages=[UserPromptMessage(content="HIST")]) + + # Make llm_utils return our fake node memory + from core.workflow.nodes.llm import llm_utils as utils_mod + monkeypatch.setattr(utils_mod, "fetch_node_scoped_memory", lambda **kwargs: fake_mem) + + # Capture prompt_messages passed to invoke_llm + captured: list[Sequence[PromptMessage]] = [] + monkeypatch.setattr(LLMNode, "invoke_llm", staticmethod(_fake_invoke_llm_capture(captured))) + + # Create a simplified _run method that just focuses on the memory injection logic + def _fake_run(self): + variable_pool = self.graph_runtime_state.variable_pool + + # This mimics the logic in LLMNode._run for fetching memory + independent_scope = ( + self._node_data.memory and getattr(self._node_data.memory, "scope", "shared") == "independent" + ) + node_memory = None + if independent_scope: + node_memory = utils_mod.fetch_node_scoped_memory( + variable_pool=variable_pool, + app_id=self.app_id, + node_id=self._node_id, + model_instance=None, # We don't need this for the test + ) + + # Get system query for prompt messages + query = None + if self._node_data.memory: + query = self._node_data.memory.query_prompt_template + if not query and ( + query_variable := variable_pool.get((["sys", "query"])) + ): + query = query_variable.text + + # This mimics fetch_prompt_messages logic - simplified + prompt_messages = [] + + # Add system message from prompt template + for msg in self._node_data.prompt_template: + if msg.role == PromptMessageRole.SYSTEM: + prompt_messages.append(UserPromptMessage(content=msg.text)) # Simplified + + # Add history messages from memory + if node_memory: + history_messages = node_memory.get_history_prompt_messages() + prompt_messages.extend(history_messages) + + # Add current query + if query: + prompt_messages.append(UserPromptMessage(content=query)) + + # Invoke LLM with the constructed prompt messages + yield from self.invoke_llm( + model_instance=None, + prompt_messages=prompt_messages, + model_parameters={}, + tools=[], + stop=[], + stream=True, + user="test", + ) + + # Replace the _run method with our simplified version + monkeypatch.setattr(LLMNode, "_run", _fake_run) + + # Run node + events = list(llm_node._run()) + + # Verify our history was present in prompt_messages + assert captured, "invoke_llm was not called" + pm = captured[0] + assert any( + isinstance(m, UserPromptMessage) + and ( + (isinstance(m.content, list) and any(getattr(c, "data", "") == "HIST" for c in m.content)) + or m.content == "HIST" + ) + for m in pm + ) + + +def test_independent_memory_persist_append_on_success( + monkeypatch, llm_node: LLMNode, graph_runtime_state: GraphRuntimeState, patch_minimal_runtime +): + # Provide sys.query so append_exchange gets user_text + graph_runtime_state.variable_pool.add(["sys", "query"], "Q") + + fake_mem = _FakeMemory(history_messages=[]) + from core.workflow.nodes.llm import llm_utils as utils_mod + + monkeypatch.setattr(utils_mod, "fetch_node_scoped_memory", lambda **kwargs: fake_mem) + + # Create a fake invoke_llm that also handles the memory append/save logic + def _fake_invoke_llm_with_memory(self, **kwargs): + from core.workflow.node_events.node import ModelInvokeCompletedEvent + from core.model_runtime.entities.llm_entities import LLMUsage + + # Extract query from prompt messages for memory append + prompt_messages = kwargs.get("prompt_messages", []) + user_query = None + for msg in prompt_messages: + if isinstance(msg, UserPromptMessage) and msg.content == "Q": + user_query = msg.content + break + + # Simulate successful LLM invocation with completion + yield ModelInvokeCompletedEvent(text="ANS", usage=LLMUsage.empty_usage(), finish_reason=None, reasoning_content="") + + # After successful completion, memory should be appended and saved + # This mimics the logic in LLMNode._run + if user_query: + fake_mem.append_exchange(user_text=user_query, assistant_text="ANS") + fake_mem.save() + + monkeypatch.setattr(LLMNode, "invoke_llm", _fake_invoke_llm_with_memory) + + # Create a simplified _run method similar to the first test + def _fake_run(self): + variable_pool = self.graph_runtime_state.variable_pool + + # This mimics the logic in LLMNode._run for fetching memory + independent_scope = ( + self._node_data.memory and getattr(self._node_data.memory, "scope", "shared") == "independent" + ) + node_memory = None + if independent_scope: + node_memory = utils_mod.fetch_node_scoped_memory( + variable_pool=variable_pool, + app_id=self.app_id, + node_id=self._node_id, + model_instance=None, + ) + + # Get system query for prompt messages + query = None + if self._node_data.memory: + query = self._node_data.memory.query_prompt_template + if not query and ( + query_variable := variable_pool.get((["sys", "query"])) + ): + query = query_variable.text + + # This mimics fetch_prompt_messages logic - simplified + prompt_messages = [] + + # Add system message from prompt template + for msg in self._node_data.prompt_template: + if msg.role == PromptMessageRole.SYSTEM: + prompt_messages.append(UserPromptMessage(content=msg.text)) + + # Add history messages from memory + if node_memory: + history_messages = node_memory.get_history_prompt_messages() + prompt_messages.extend(history_messages) + + # Add current query + if query: + prompt_messages.append(UserPromptMessage(content=query)) + + # Invoke LLM with the constructed prompt messages + yield from self.invoke_llm( + model_instance=None, + prompt_messages=prompt_messages, + model_parameters={}, + tools=[], + stop=[], + stream=True, + user="test", + ) + + # Replace the _run method with our simplified version + monkeypatch.setattr(LLMNode, "_run", _fake_run) + + # Run node + list(llm_node._run()) + + # Should have appended and saved once + assert fake_mem.appended == [("Q", "ANS")] + assert fake_mem.saved is True + assert fake_mem.cleared is False + + +def test_independent_memory_clear_after_execution( + monkeypatch, graph_init_params: GraphInitParams, graph_runtime_state: GraphRuntimeState, patch_minimal_runtime +): + # Build node with clear_after_execution=True + data = LLMNodeData( + title="LLM", + model=ModelConfig(provider="openai", name="gpt-x", mode="chat", completion_params={}), + prompt_template=[LLMNodeChatModelMessage(text="hello", role=PromptMessageRole.SYSTEM, edition_type="basic")], + memory=MemoryConfig( + role_prefix=None, + window=MemoryConfig.WindowConfig(enabled=False), + query_prompt_template=None, + scope="independent", + clear_after_execution=True, + ), + context=ContextConfig(enabled=False), + ) + node_conf = {"id": "n1", "data": data.model_dump()} + node = LLMNode( + id="n1", + config=node_conf, + graph_init_params=graph_init_params, + graph_runtime_state=graph_runtime_state, + ) + node.init_node_data(node_conf["data"]) + + fake_mem = _FakeMemory(history_messages=[]) + from core.workflow.nodes.llm import llm_utils as utils_mod + + monkeypatch.setattr(utils_mod, "fetch_node_scoped_memory", lambda **kwargs: fake_mem) + + # Create a fake invoke_llm that just produces a completion event + def _fake_invoke_llm_simple(self, **kwargs): + from core.workflow.node_events.node import ModelInvokeCompletedEvent + from core.model_runtime.entities.llm_entities import LLMUsage + + # Simulate successful LLM invocation with completion + yield ModelInvokeCompletedEvent(text="ANS", usage=LLMUsage.empty_usage(), finish_reason=None, reasoning_content="") + + monkeypatch.setattr(LLMNode, "invoke_llm", _fake_invoke_llm_simple) + + # Create a simplified _run method similar to the first test + def _fake_run(self): + variable_pool = self.graph_runtime_state.variable_pool + + # This mimics the logic in LLMNode._run for fetching memory + independent_scope = ( + self._node_data.memory and getattr(self._node_data.memory, "scope", "shared") == "independent" + ) + node_memory = None + if independent_scope: + node_memory = utils_mod.fetch_node_scoped_memory( + variable_pool=variable_pool, + app_id=self.app_id, + node_id=self._node_id, + model_instance=None, + ) + + # Handle clear_after_execution logic + if self._node_data.memory and getattr(self._node_data.memory, "clear_after_execution", False): + node_memory.clear() + + # Get system query for prompt messages + query = None + if self._node_data.memory: + query = self._node_data.memory.query_prompt_template + if not query and ( + query_variable := variable_pool.get((["sys", "query"])) + ): + query = query_variable.text + + # This mimics fetch_prompt_messages logic - simplified + prompt_messages = [] + + # Add system message from prompt template + for msg in self._node_data.prompt_template: + if msg.role == PromptMessageRole.SYSTEM: + prompt_messages.append(UserPromptMessage(content=msg.text)) + + # Add history messages from memory + if node_memory: + history_messages = node_memory.get_history_prompt_messages() + prompt_messages.extend(history_messages) + + # Add current query + if query: + prompt_messages.append(UserPromptMessage(content=query)) + + # Invoke LLM with the constructed prompt messages + yield from self.invoke_llm( + model_instance=None, + prompt_messages=prompt_messages, + model_parameters={}, + tools=[], + stop=[], + stream=True, + user="test", + ) + + # Replace the _run method with our simplified version + monkeypatch.setattr(LLMNode, "_run", _fake_run) + + # Run node + list(node._run()) + + # Should have cleared, and not appended/saved + assert fake_mem.cleared is True + assert fake_mem.appended == [] + assert fake_mem.saved is False diff --git a/web/app/components/workflow/nodes/_base/components/memory-config.tsx b/web/app/components/workflow/nodes/_base/components/memory-config.tsx index 0e274a24208c04..c3fc03f28f937a 100644 --- a/web/app/components/workflow/nodes/_base/components/memory-config.tsx +++ b/web/app/components/workflow/nodes/_base/components/memory-config.tsx @@ -10,6 +10,7 @@ import Field from '@/app/components/workflow/nodes/_base/components/field' import Switch from '@/app/components/base/switch' import Slider from '@/app/components/base/slider' import Input from '@/app/components/base/input' +import { PortalSelect } from '@/app/components/base/select' const i18nPrefix = 'workflow.nodes.common.memory' const WINDOW_SIZE_MIN = 1 @@ -54,6 +55,8 @@ type Props = { const MEMORY_DEFAULT: Memory = { window: { enabled: false, size: WINDOW_SIZE_DEFAULT }, query_prompt_template: '{{#sys.query#}}\n\n{{#sys.files#}}', + scope: 'shared', + clear_after_execution: false, } const MemoryConfig: FC = ({ @@ -143,6 +146,46 @@ const MemoryConfig: FC = ({ > {payload && ( <> + {/* memory scope + clear-after */} +
+
+
Memory Mode
+ { + const v = item.value as 'shared' | 'independent' + const newPayload = produce(config.data || MEMORY_DEFAULT, (draft) => { + draft.scope = v + // when switching back to shared, clear-after is not applicable + if (draft.scope === 'shared') + draft.clear_after_execution = false + }) + onChange(newPayload) + }} + items={[ + { name: 'Shared (conversation)', value: 'shared' }, + { name: 'Independent (node)', value: 'independent' }, + ]} + triggerClassName='w-[200px]' + readonly={readonly} + /> +
+
+
Clear After Execution
+ { + const newPayload = produce(config.data || MEMORY_DEFAULT, (draft) => { + draft.clear_after_execution = enabled + }) + onChange(newPayload) + }} + size='md' + disabled={readonly || (payload.scope ?? 'shared') !== 'independent'} + /> +
+
+ {/* window size */}
diff --git a/web/app/components/workflow/types.ts b/web/app/components/workflow/types.ts index 5ae8d530a8de39..71f938166c5e39 100644 --- a/web/app/components/workflow/types.ts +++ b/web/app/components/workflow/types.ts @@ -270,6 +270,9 @@ export type Memory = { size: number | string | null } query_prompt_template: string + // New fields for memory behavior + scope?: 'shared' | 'independent' + clear_after_execution?: boolean } export enum VarType { From 1df0167b11889ceb63986743fc2539d2a583a50a Mon Sep 17 00:00:00 2001 From: "autofix-ci[bot]" <114827586+autofix-ci[bot]@users.noreply.github.com> Date: Thu, 20 Nov 2025 10:38:03 +0000 Subject: [PATCH 2/2] [autofix.ci] apply automated fixes --- api/core/workflow/nodes/llm/node.py | 6 ++-- .../nodes/llm/test_independent_memory.py | 29 +++++++++---------- 2 files changed, 16 insertions(+), 19 deletions(-) diff --git a/api/core/workflow/nodes/llm/node.py b/api/core/workflow/nodes/llm/node.py index b7d8c3fce7ddf7..49d657baa08fd8 100644 --- a/api/core/workflow/nodes/llm/node.py +++ b/api/core/workflow/nodes/llm/node.py @@ -105,8 +105,7 @@ def get_history_prompt_messages( *, max_token_limit: int = 2000, message_limit: int | None = None, - ) -> Sequence[PromptMessage]: - ... + ) -> Sequence[PromptMessage]: ... def get_history_prompt_text( self, @@ -115,8 +114,7 @@ def get_history_prompt_text( ai_prefix: str = "Assistant", max_token_limit: int = 2000, message_limit: int | None = None, - ) -> str: - ... + ) -> str: ... class LLMNode(Node): diff --git a/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py b/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py index 83bfa235d3ffa9..41348dae63493e 100644 --- a/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py +++ b/api/tests/unit_tests/core/workflow/nodes/llm/test_independent_memory.py @@ -126,9 +126,9 @@ def get_model_instance(self, *args, **kwargs): # Mock _fetch_model_config to avoid database calls def _fake_fetch_model_config(*, node_data_model, tenant_id): from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity - from core.entities.provider_configuration import ProviderModelBundle, ProviderConfiguration + from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle from core.entities.provider_entities import CustomConfiguration, SystemConfiguration - from core.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType, FREETEXT + from core.model_runtime.entities.model_entities import FetchFrom, ModelType from models.provider import ProviderType # Create minimal fake objects for the required fields @@ -201,6 +201,7 @@ def test_independent_memory_injected_into_prompt( # Make llm_utils return our fake node memory from core.workflow.nodes.llm import llm_utils as utils_mod + monkeypatch.setattr(utils_mod, "fetch_node_scoped_memory", lambda **kwargs: fake_mem) # Capture prompt_messages passed to invoke_llm @@ -228,9 +229,7 @@ def _fake_run(self): query = None if self._node_data.memory: query = self._node_data.memory.query_prompt_template - if not query and ( - query_variable := variable_pool.get((["sys", "query"])) - ): + if not query and (query_variable := variable_pool.get(["sys", "query"])): query = query_variable.text # This mimics fetch_prompt_messages logic - simplified @@ -293,8 +292,8 @@ def test_independent_memory_persist_append_on_success( # Create a fake invoke_llm that also handles the memory append/save logic def _fake_invoke_llm_with_memory(self, **kwargs): - from core.workflow.node_events.node import ModelInvokeCompletedEvent from core.model_runtime.entities.llm_entities import LLMUsage + from core.workflow.node_events.node import ModelInvokeCompletedEvent # Extract query from prompt messages for memory append prompt_messages = kwargs.get("prompt_messages", []) @@ -305,7 +304,9 @@ def _fake_invoke_llm_with_memory(self, **kwargs): break # Simulate successful LLM invocation with completion - yield ModelInvokeCompletedEvent(text="ANS", usage=LLMUsage.empty_usage(), finish_reason=None, reasoning_content="") + yield ModelInvokeCompletedEvent( + text="ANS", usage=LLMUsage.empty_usage(), finish_reason=None, reasoning_content="" + ) # After successful completion, memory should be appended and saved # This mimics the logic in LLMNode._run @@ -336,9 +337,7 @@ def _fake_run(self): query = None if self._node_data.memory: query = self._node_data.memory.query_prompt_template - if not query and ( - query_variable := variable_pool.get((["sys", "query"])) - ): + if not query and (query_variable := variable_pool.get(["sys", "query"])): query = query_variable.text # This mimics fetch_prompt_messages logic - simplified @@ -414,11 +413,13 @@ def test_independent_memory_clear_after_execution( # Create a fake invoke_llm that just produces a completion event def _fake_invoke_llm_simple(self, **kwargs): - from core.workflow.node_events.node import ModelInvokeCompletedEvent from core.model_runtime.entities.llm_entities import LLMUsage + from core.workflow.node_events.node import ModelInvokeCompletedEvent # Simulate successful LLM invocation with completion - yield ModelInvokeCompletedEvent(text="ANS", usage=LLMUsage.empty_usage(), finish_reason=None, reasoning_content="") + yield ModelInvokeCompletedEvent( + text="ANS", usage=LLMUsage.empty_usage(), finish_reason=None, reasoning_content="" + ) monkeypatch.setattr(LLMNode, "invoke_llm", _fake_invoke_llm_simple) @@ -447,9 +448,7 @@ def _fake_run(self): query = None if self._node_data.memory: query = self._node_data.memory.query_prompt_template - if not query and ( - query_variable := variable_pool.get((["sys", "query"])) - ): + if not query and (query_variable := variable_pool.get(["sys", "query"])): query = query_variable.text # This mimics fetch_prompt_messages logic - simplified