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
44 changes: 43 additions & 1 deletion apps/application/flow/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -249,7 +249,8 @@ def is_valid_model_params(self):
node_list = [node for node in self.nodes if (
node.type == 'ai-chat-node' or node.type == 'question-node' or node.type == 'parameter-extraction-node')]
for node in node_list:
if (node.properties.get('node_data', {}).get('model_id_type') or 'custom') == 'reference':
model_id_type = node.properties.get('node_data', {}).get('model_id_type') or 'custom'
if model_id_type in ('reference', 'default'):
continue
model = QuerySet(Model).filter(id=node.properties.get('node_data', {}).get('model_id')).first()
if model is None:
Expand Down Expand Up @@ -282,3 +283,44 @@ def is_valid_base_node(self):
raise AppApiException(500, _('Basic information node is required'))
if len(base_node_list) > 1:
raise AppApiException(500, _('There can only be one basic information node'))


def get_base_node_model(application, model_type):
"""
解析基本信息节点(base-node)实际使用的模型配置。
model_type: 'STT' | 'TTS' | 'LLM'(长期记忆)。
WORK_FLOW 应用读 base-node node_data 的模式字段;default/DEFAULT 时取 default_model_setting;
SIMPLE 应用(无 base-node)回退到顶层模型字段。
返回 {'model_id', 'model_params'}。
"""
mode_field = {'STT': 'stt_model_id_type', 'TTS': 'tts_type', 'LLM': 'long_term_model_id_type'}[model_type]
model_field = {'STT': 'stt_model_id', 'TTS': 'tts_model_id', 'LLM': 'long_term_model_id'}[model_type]
param_field = {
'STT': 'stt_model_params_setting',
'TTS': 'tts_model_params_setting',
'LLM': 'long_term_model_params_setting',
}[model_type]
node_data = None
for node in ((getattr(application, 'work_flow', None) or {}).get('nodes') or []):
if node.get('id') == 'base-node':
node_data = (node.get('properties') or {}).get('node_data') or {}
break
if node_data is None:
return {
'model_id': getattr(application, model_field, None),
'model_params': getattr(application, param_field, None) or {},
}
mode = node_data.get(mode_field)
if mode == 'default' or mode == 'DEFAULT':
default_setting = (application.default_model_setting or {}).get(model_type, {}) or {}
return {
'model_id': default_setting.get('model_id'),
'model_params': default_setting.get('model_params_setting') or {},
}
if model_type == 'TTS' and mode == 'BROWSER':
# 浏览器播放不使用 TTS 模型,残留 model_id 不参与解析
return {'model_id': None, 'model_params': {}}
return {
'model_id': node_data.get(model_field),
'model_params': node_data.get(param_field) or {},
}
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,11 @@ def execute(
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get("model_id", model_id)
model_params_setting = reference_data.get("model_params_setting")
elif model_id_type == "default":
default_setting = self.workflow_manage.get_default_model_setting("LLM")
if default_setting.get("model_id"):
model_id = default_setting.get("model_id")
model_params_setting = default_setting.get("model_params_setting", model_params_setting)
if model_id is None or model_id == "":
raise Exception(_("Model is not allowed to be empty"))

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,11 @@ def execute(self, model_id, prompt, negative_prompt, dialogue_number, dialogue_t
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('TTI')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

if model_id is None or model_id == '':
raise Exception(_('Model is not allowed to be empty'))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,11 @@ def execute(self, model_id, prompt, negative_prompt, dialogue_number, dialogue_t
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('ITV')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)
if model_id is None or model_id == '':
raise Exception(_('Model is not allowed to be empty'))
workspace_id = self.workflow_manage.get_body().get('workspace_id')
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,11 @@ def execute(self, model_id, system, prompt, dialogue_number, dialogue_type, hist
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('IMAGE')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

if model_id is None or model_id == '':
raise Exception(_('Model is not allowed to be empty'))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,11 @@ def execute(self, model_id, dialogue_number, history_chat_record, user_input, br
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('LLM')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)
if not model_id:
raise Exception(_('Model is not allowed to be empty'))

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,11 @@ def _run(self):
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('LLM')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

input_variable = self.workflow_manage.get_reference_field(
self.node_params_serializer.data.get('input_variable')[0],
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,11 @@ def execute(self, model_id, system, prompt, dialogue_number, history_chat_record
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('LLM')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)
if not model_id:
raise Exception(_('Model is not allowed to be empty'))

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,10 @@ def _run(self):
if reference_data and isinstance(reference_data, dict):
reranker_model_id = reference_data.get('reranker_model_id',
reference_data.get('model_id', reranker_model_id))
elif reranker_model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('RERANKER')
if default_setting.get('model_id'):
reranker_model_id = default_setting.get('model_id')
if reranker_model_id is None or reranker_model_id == '':
raise Exception(_('Model is not allowed to be empty'))

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,11 @@ def execute(self, stt_model_id, audio, model_params_setting=None, stt_model_id_t
if reference_data and isinstance(reference_data, dict):
stt_model_id = reference_data.get('stt_model_id', reference_data.get('model_id', stt_model_id))
model_params_setting = reference_data.get('model_params_setting')
elif stt_model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('STT')
if default_setting.get('model_id'):
stt_model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

from django.utils.translation import gettext_lazy as _

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,11 @@ def execute(self, tts_model_id,
if reference_data and isinstance(reference_data, dict):
tts_model_id = reference_data.get('tts_model_id', reference_data.get('model_id', tts_model_id))
model_params_setting = reference_data.get('model_params_setting')
elif tts_model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('TTS')
if default_setting.get('model_id'):
tts_model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

from django.utils.translation import gettext_lazy as _

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,11 @@ def execute(self, model_id, prompt, negative_prompt, dialogue_number, dialogue_t
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('TTV')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

if model_id is None or model_id == '':
raise Exception(_('Model is not allowed to be empty'))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -228,7 +228,8 @@ def workflow_manage_new_instance(start_node_id=None,
'stream': True,
'workspace_id': workspace_id,
'user_id': runtime_user_id,
**parameters},
**parameters,
'default_model_setting': tool_workflow_version.default_model_setting},
ToolWorkflowPostHandler(took_execute, tool_lib_id),
base_to_response=LoopToResponse(),
start_node_id=start_node_id,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -147,6 +147,11 @@ def execute(self, model_id, system, prompt, dialogue_number, dialogue_type, hist
if reference_data and isinstance(reference_data, dict):
model_id = reference_data.get('model_id', model_id)
model_params_setting = reference_data.get('model_params_setting')
elif model_id_type == 'default':
default_setting = self.workflow_manage.get_default_model_setting('IMAGE')
if default_setting.get('model_id'):
model_id = default_setting.get('model_id')
model_params_setting = default_setting.get('model_params_setting', model_params_setting)

from django.utils.translation import gettext_lazy as _

Expand Down
69 changes: 58 additions & 11 deletions apps/application/flow/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -816,6 +816,34 @@ async def anext_async(agen):
return await agen.__anext__()


def _get_node_model_id(node, model_field, mode_field):
"""节点为 default/reference 模式时不返回节点内 model_id(运行时才解析,避免脏映射)。"""
node_data = (node.get("properties") or {}).get("node_data") or {}
if node_data.get(mode_field) in ("default", "reference"):
return None
return node_data.get(model_field)


# base-node 三类模型:mode 判定与 validate_workflow_default_models/get_base_node_model 保持一致
# (stt/长期记忆用 'default'/'reference',tts 用大写 'DEFAULT'/'BROWSER')
_base_node_model_specs = (
("stt_model_id_type", ("default", "reference"), "stt_model_enable", "stt_model_id"),
("tts_type", ("DEFAULT", "BROWSER"), "tts_model_enable", "tts_model_id"),
("long_term_model_id_type", ("default", "reference"), "long_term_enable", "long_term_model_id"),
)


def _get_base_node_model_ids(node):
"""返回 base-node node_data 中实际自定义的 STT/TTS/长期记忆 model_id(default/BROWSER 时运行时解析,不映射)。"""
node_data = (node.get("properties") or {}).get("node_data") or {}
model_ids = []
for mode_field, skip_modes, enable_field, model_field in _base_node_model_specs:
if node_data.get(enable_field) and node_data.get(mode_field) not in skip_modes:
if node_data.get(model_field):
model_ids.append(node_data.get(model_field))
return model_ids


target_source_node_mapping = {
"TOOL": {
"tool-lib-node": lambda n: [n.get("properties").get("node_data").get("tool_lib_id")],
Expand All @@ -828,17 +856,18 @@ async def anext_async(agen):
"tool-workflow-lib-node": lambda n: [n.get("properties").get("node_data").get("tool_lib_id")],
},
"MODEL": {
"ai-chat-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"question-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"speech-to-text-node": lambda n: [n.get("properties").get("node_data").get("stt_model_id")],
"text-to-speech-node": lambda n: [n.get("properties").get("node_data").get("tts_model_id")],
"image-to-video-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"image-generate-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"intent-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"image-understand-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"parameter-extraction-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"video-understand-node": lambda n: [n.get("properties").get("node_data").get("model_id")],
"reranker-node": lambda n: [n.get("properties").get("node_data").get("reranker_model_id")],
"ai-chat-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"question-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"speech-to-text-node": lambda n: [v for v in [_get_node_model_id(n, 'stt_model_id', 'stt_model_id_type')] if v],
"text-to-speech-node": lambda n: [v for v in [_get_node_model_id(n, 'tts_model_id', 'tts_model_id_type')] if v],
"image-to-video-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"image-generate-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"intent-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"image-understand-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"parameter-extraction-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"video-understand-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v],
"reranker-node": lambda n: [v for v in [_get_node_model_id(n, 'reranker_model_id', 'reranker_model_id_type')] if v],
"base-node": _get_base_node_model_ids,
},
"KNOWLEDGE": {
"search-knowledge-node": lambda n: n.get("properties").get("node_data").get("knowledge_id_list"),
Expand Down Expand Up @@ -905,6 +934,7 @@ def get_workflow_resource(workflow, node_handle):
lambda instance: [instance.long_term_model_id] if instance.long_term_model_id else [],
lambda instance: [instance.tts_model_id] if instance.tts_model_id else [],
lambda instance: [instance.stt_model_id] if instance.stt_model_id else [],
lambda instance: [v.get('model_id') for v in (instance.default_model_setting or {}).values() if (v or {}).get('model_id')],
],
}
knowledge_instance_field_call_dict = {
Expand All @@ -928,6 +958,22 @@ def get_instance_resource(instance, source_type, source_id, instance_field_call_
return response


def append_default_model_mapping(instance_mapping, default_model_setting, source_type, source_id):
"""把 default_model_setting 各类别 model_id 追加为 MODEL 资源映射(方案A),返回追加后的列表。"""
from system_manage.models.resource_mapping import ResourceMapping, ResourceType

for value in (default_model_setting or {}).values():
model_id = (value or {}).get('model_id')
if model_id:
instance_mapping.append(
ResourceMapping(
source_type=source_type, target_type=ResourceType.MODEL,
source_id=str(source_id), target_id=model_id,
)
)
return instance_mapping


def save_workflow_mapping(workflow, source_type, source_id, other_resource_mapping=None):
if not other_resource_mapping:
other_resource_mapping = []
Expand Down Expand Up @@ -1058,6 +1104,7 @@ def inner(**kwargs):
"workspace_id": workspace_id,
"user_id": user_id,
**kwargs,
"default_model_setting": qv.default_model_setting,
},
ToolWorkflowPostHandler(took_execute, tool_id),
is_the_task_interrupted=lambda: False,
Expand Down
29 changes: 29 additions & 0 deletions apps/application/flow/workflow_manage.py
Original file line number Diff line number Diff line change
Expand Up @@ -831,3 +831,32 @@ def get_source_type(self):

def get_source_id(self):
return self.params.get('application_id')

def get_application(self):
"""获取当前应用的 Application 实例(仅 APPLICATION 模式有;其余模式返回 None)"""
if self.work_flow_post_handler is None:
return None
chat_info = getattr(self.work_flow_post_handler, 'chat_info', None)
if chat_info is None:
return None
return getattr(chat_info, 'application', None)

def get_default_model_setting(self, model_type):
"""获取指定模型类别的默认模型配置 {model_id, model_params_setting,...};未配置返回 {}"""
setting = self._load_default_model_setting()
return (setting or {}).get(model_type, {}) or {}

def _load_default_model_setting(self):
"""按执行来源读取 default_model_setting(params 透传),每次 run 缓存一次;Application 回退 chat_info.application。"""
if getattr(self, '_default_model_setting', None) is not None:
return self._default_model_setting
setting = self.params.get('default_model_setting')
if setting is None:
application = self.get_application()
setting = getattr(application, 'default_model_setting', {}) or {}
self._default_model_setting = setting
return setting

def get_default_model_id(self, model_type):
"""获取指定模型类别的默认模型 id;未配置返回 None"""
return self.get_default_model_setting(model_type).get('model_id')
Loading