Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 6 additions & 8 deletions apps/application/workflow/loop_workflow_manage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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):
"""
Expand Down
10 changes: 4 additions & 6 deletions apps/application/workflow/workflow_manage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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):
"""
Expand Down
34 changes: 34 additions & 0 deletions apps/common/utils/prompt_template.py
Original file line number Diff line number Diff line change
@@ -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)
Loading