From 0545278ec895cc6443a0c8b52df9a24b822935d1 Mon Sep 17 00:00:00 2001 From: hexiaonan-800 Date: Mon, 21 Sep 2026 17:17:35 +0800 Subject: [PATCH] feat: Optimize loop iterations by caching template compilation --- .../workflow/loop_workflow_manage.py | 14 ++++---- apps/application/workflow/workflow_manage.py | 10 +++--- apps/common/utils/prompt_template.py | 34 +++++++++++++++++++ 3 files changed, 44 insertions(+), 14 deletions(-) create mode 100644 apps/common/utils/prompt_template.py diff --git a/apps/application/workflow/loop_workflow_manage.py b/apps/application/workflow/loop_workflow_manage.py index 16d103c6ce7..9a3ef30351d 100644 --- a/apps/application/workflow/loop_workflow_manage.py +++ b/apps/application/workflow/loop_workflow_manage.py @@ -7,11 +7,12 @@ @desc: """ -from typing import Dict, Optional, Callable +from typing import Dict, Callable -from application.workflow.common import Workflow, WorkflowType, Node +from application.workflow.common import Workflow, WorkflowType from application.workflow.i_node import INode from application.workflow.workflow_manage import WorkflowManage, CallBack +from common.utils.prompt_template import render_prompt class LoopWorkFlowManage(WorkflowManage): @@ -34,13 +35,10 @@ def get_parent_context(self, node_id, key): return self.parent_workflow_manage.get_context(node_id, key) def generate_prompt(self, prompt): - prompt = self.workflow.reset_prompt(prompt) - prompt = self.parent_workflow_manage.workflow.reset_prompt(prompt) + input_template = self.workflow.reset_prompt(prompt) + input_template = self.parent_workflow_manage.workflow.reset_prompt(input_template) context = {**self.context, **self.parent_workflow_manage.context} - from langchain_core.prompts import PromptTemplate - - prompt_template = PromptTemplate.from_template(prompt, template_format="jinja2") - return prompt_template.format(context=context) + return render_prompt(input_template, context) def get_reference_field(self, node_id, fields): """ diff --git a/apps/application/workflow/workflow_manage.py b/apps/application/workflow/workflow_manage.py index 11afd733472..5360ccc2480 100644 --- a/apps/application/workflow/workflow_manage.py +++ b/apps/application/workflow/workflow_manage.py @@ -12,13 +12,12 @@ import threading from typing import List, Dict, Optional, Callable -from langchain_core.prompts import PromptTemplate - from application.workflow.common import Workflow, WorkflowType, Node, get_node_parameters from application.workflow.i_node import INode, Signal -from application.workflow.message.struct.content import Content, Position +from application.workflow.message.struct.content import Content from application.workflow.status import Status +from common.utils.prompt_template import render_prompt class CallBack: @@ -223,9 +222,8 @@ def generate_prompt(self, prompt): @param prompt: 提示词 @return: 处理后的提示词 """ - prompt = self.workflow.reset_prompt(prompt) - prompt_template = PromptTemplate.from_template(prompt, template_format="jinja2") - return prompt_template.format(context=self.context) + input_template = self.workflow.reset_prompt(prompt) + return render_prompt(input_template, self.context) def get_reference_field(self, node_id, fields): """ diff --git a/apps/common/utils/prompt_template.py b/apps/common/utils/prompt_template.py new file mode 100644 index 00000000000..30b686d5c45 --- /dev/null +++ b/apps/common/utils/prompt_template.py @@ -0,0 +1,34 @@ +# coding=utf-8 + +import hashlib + +from jinja2.sandbox import SandboxedEnvironment + +from common.cache.mem_cache import MemCache + +# 缓存 reset_prompt 后的模板编译产物(只依赖配置,同一工作流运行期间恒定) +template_cache = MemCache( + "workflow_template_cache", + { + "TIMEOUT": 3600, # 缓存有效期为 1 小时 + "OPTIONS": { + "MAX_ENTRIES": 1000, # 最多缓存 1000 个条目 + "CULL_FREQUENCY": 10, # 达到上限时,删除约 1/10 的缓存 + }, + }, +) + + +def render_prompt(input_template, context): + """ + 渲染提示词,编译一次后缓存编译产物,之后每轮仅渲染 + @param input_template: reset_prompt 处理后的模板字符串 + @param context: 渲染模板所需的上下文数据 + @return: 渲染后的提示词字符串 + """ + key = f"workflow_template::{hashlib.sha256(input_template.encode('utf-8')).hexdigest()}" + template = template_cache.get(key) + if template is None: + template = SandboxedEnvironment().from_string(input_template) + template_cache.set(key, template) + return template.render(context=context)