diff --git a/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py b/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py index ae744850396..4921e01059f 100644 --- a/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py +++ b/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py @@ -657,8 +657,8 @@ def get_details(self, index: int = 0, position: dict = None, old_details: dict = "question": self.data.get("question"), "answer": self.get_context("answer"), "reasoning_content": self.get_context("reasoning_content"), - "message_tokens": self.get_context("message_tokens"), - "answer_tokens": self.get_context("answer_tokens"), + "message_tokens": self.data.get("message_tokens"), + "answer_tokens": self.data.get("answer_tokens"), "history_message": self.get_context("history_message"), "messages": messages, } diff --git a/apps/application/workflow/nodes/loop_node/loop_node.py b/apps/application/workflow/nodes/loop_node/loop_node.py index c33e4449bcf..ea60b45e8e0 100644 --- a/apps/application/workflow/nodes/loop_node/loop_node.py +++ b/apps/application/workflow/nodes/loop_node/loop_node.py @@ -114,6 +114,12 @@ def on_next(wf_manage, content): def on_complete(wf_manage, error): loop_details_list = self.data.setdefault("loop_details_list", []) loop_details_list.append(wf_manage.get_details()) + self.data["message_tokens"] = (self.data.get("message_tokens") or 0) + sum( + (n.data.get("message_tokens") or 0) for n in wf_manage.nodes + ) + self.data["answer_tokens"] = (self.data.get("answer_tokens") or 0) + sum( + (n.data.get("answer_tokens") or 0) for n in wf_manage.nodes + ) self.write_context("index", index) self.write_context("item", item) last_context = self.workflow_manage.get_context(self.node.id, "last_context") @@ -193,7 +199,13 @@ def get_context(): def get_details(self, index: int = 0, position: dict = None, old_details: dict = None, **kwargs): details = super().get_details(index, position, old_details, **kwargs) details.update( - {"params": self.data.get("params"), "index": self.get_context("index"), "item": self.get_context("item")} + { + "params": self.data.get("params"), + "index": self.get_context("index"), + "item": self.get_context("item"), + "message_tokens": self.data.get("message_tokens"), + "answer_tokens": self.data.get("answer_tokens"), + } ) loop_details = [] position_index = 0 diff --git a/apps/application/workflow/workflow_manage.py b/apps/application/workflow/workflow_manage.py index 5360ccc2480..c0e6b7e9df1 100644 --- a/apps/application/workflow/workflow_manage.py +++ b/apps/application/workflow/workflow_manage.py @@ -10,6 +10,7 @@ from __future__ import annotations import threading +import time from typing import List, Dict, Optional, Callable from application.workflow.common import Workflow, WorkflowType, Node, get_node_parameters @@ -64,6 +65,7 @@ def __init__( self.signal = None self.details = {"position": {}, "details": {}} self.start_node = get_start_node(workflow, self) + self.start_time = time.time() def run(self): """ diff --git a/apps/chat/serializers/chat.py b/apps/chat/serializers/chat.py index ffe5d60ef79..d3b08ce854d 100644 --- a/apps/chat/serializers/chat.py +++ b/apps/chat/serializers/chat.py @@ -12,6 +12,7 @@ import queue import queue as thread_queue import threading +import time import uuid_utils import uuid_utils.compat as uuid @@ -221,21 +222,15 @@ def get_defaults_record(self, question): } @staticmethod - def _usage_from_context(workflow_context): - """从 workflow_context 汇总 token 用量:prompt=message_tokens, completion=answer_tokens。""" - prompt_tokens = sum( - v.get("message_tokens", 0) - for v in workflow_context.values() - if isinstance(v, dict) and "message_tokens" in v - ) - completion_tokens = sum( - v.get("answer_tokens", 0) for v in workflow_context.values() if isinstance(v, dict) and "answer_tokens" in v - ) + def _usage_from_details(details): + """从节点详情列表汇总 token 用量:prompt=message_tokens, completion=answer_tokens。""" + prompt_tokens = sum((d.get("message_tokens") or 0) for d in details if isinstance(d, dict)) + completion_tokens = sum((d.get("answer_tokens") or 0) for d in details if isinstance(d, dict)) return {"prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens} @staticmethod - def update_chat_record(chat_user_id, chat_record_id, workflow_context, messages, details): - usage = ChatSerializers._usage_from_context(workflow_context) + def update_chat_record(chat_user_id, chat_record_id, workflow_context, messages, details, run_time): + usage = ChatSerializers._usage_from_details(details) message_tokens = usage["prompt_tokens"] answer_tokens = usage["completion_tokens"] ChatUserTokenQuota.consume(chat_user_id, message_tokens + answer_tokens) @@ -245,6 +240,7 @@ def update_chat_record(chat_user_id, chat_record_id, workflow_context, messages, message_tokens=message_tokens, answer_tokens=answer_tokens, details=details, + run_time=run_time, ) # ---------- 执行 ---------- @@ -341,7 +337,8 @@ def on_complete(wf_manage, error): if position and chat_record.messages: messages = list({m.get("id"): m for m in [*chat_record.messages, *messages]}.values()) details = wf_manage.get_details(position=position, old_details=old_details) - self.update_chat_record(chat_user_id, chat_record_id_str, wf_manage.context, messages, details) + run_time = time.time() - wf_manage.start_time + self.update_chat_record(chat_user_id, chat_record_id_str, wf_manage.context, messages, details, run_time) ChatCountSerializer(data={"chat_id": chat_id}).update_chat() # 表单续跑时 message_dict.content 为空;保留原记录里的用户问题,避免 WORKFLOW 历史丢问题 question = chat_record.question if (chat_record and chat_record.question) else message_dict @@ -395,7 +392,7 @@ def generate(): end_frame = base_to_response.to_stream_end( chat_id, chat_record_id_str, - usage=self._usage_from_context(work_flow_manage.context), + usage=self._usage_from_details(work_flow_manage.get_details()), ) if end_frame is not None: yield "data: " + end_frame + "\n\n" @@ -422,7 +419,7 @@ def generate(): break if msg_type == "error": raise data - usage = self._usage_from_context(work_flow_manage.context) + usage = self._usage_from_details(work_flow_manage.get_details()) return base_to_response.to_block(chat_id, chat_record_id_str, aggregation.get_contents(), usage) def chat(self, instance: dict, base_to_response: BaseToResponse = SystemToResponse()):