thread_runner.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450
  1. from functools import partial
  2. import logging
  3. from typing import List
  4. from concurrent.futures import Executor
  5. from sqlalchemy.orm import Session
  6. from app.models.token_relation import RelationType
  7. from config.config import settings
  8. from config.llm import llm_settings, tool_settings
  9. from app.core.runner.llm_backend import LLMBackend
  10. from app.core.runner.llm_callback_handler import LLMCallbackHandler
  11. from app.core.runner.memory import Memory, find_memory
  12. from app.core.runner.pub_handler import StreamEventHandler
  13. from app.core.runner.utils import message_util as msg_util
  14. from app.core.runner.utils.tool_call_util import (
  15. tool_call_recognize,
  16. internal_tool_call_invoke,
  17. tool_call_request,
  18. tool_call_id,
  19. tool_call_output,
  20. )
  21. from app.core.tools import find_tools, BaseTool
  22. from app.libs.thread_executor import get_executor_for_config, run_with_executor
  23. from app.models.message import Message, MessageUpdate
  24. from app.models.run import Run
  25. from app.models.run_step import RunStep
  26. from app.models.token_relation import RelationType
  27. from app.services.assistant.assistant import AssistantService
  28. from app.services.file.file import FileService
  29. from app.services.message.message import MessageService
  30. from app.services.run.run import RunService
  31. from app.services.run.run_step import RunStepService
  32. from app.services.token.token import TokenService
  33. from app.services.token.token_relation import TokenRelationService
  34. class ThreadRunner:
  35. """
  36. ThreadRunner 封装 run 的执行逻辑
  37. """
  38. tool_executor: Executor = get_executor_for_config(
  39. tool_settings.TOOL_WORKER_NUM, "tool_worker_"
  40. )
  41. def __init__(
  42. self, run_id: str, token_id: str, session: Session, stream: bool = False
  43. ):
  44. self.run_id = run_id
  45. self.token_id = token_id
  46. self.session = session
  47. self.stream = stream
  48. self.max_step = llm_settings.LLM_MAX_STEP
  49. self.event_handler: StreamEventHandler = None
  50. def run(self):
  51. """
  52. 完成一次 run 的执行,基本步骤
  53. 1. 初始化,获取 run 以及相关 tools, 构造 system instructions;
  54. 2. 开始循环,查询已有 run step, 进行 chat message 生成;
  55. 3. 调用 llm 并解析返回结果;
  56. 4. 根据返回结果,生成新的 run step(tool calls 处理) 或者 message
  57. """
  58. # TODO: 重构,将 run 的状态变更逻辑放到 RunService 中
  59. run = RunService.get_run_sync(session=self.session, run_id=self.run_id)
  60. self.event_handler = StreamEventHandler(
  61. run_id=self.run_id, is_stream=self.stream
  62. )
  63. run = RunService.to_in_progress(session=self.session, run_id=self.run_id)
  64. self.event_handler.pub_run_in_progress(run)
  65. logging.info("processing ThreadRunner task, run_id: %s", self.run_id)
  66. # get memory from assistant metadata
  67. # format likes {"memory": {"type": "window", "window_size": 20, "max_token_size": 4000}}
  68. ast = AssistantService.get_assistant_sync(
  69. session=self.session, assistant_id=run.assistant_id
  70. )
  71. metadata = ast.metadata_ or {}
  72. memory = find_memory(metadata.get("memory", {}))
  73. instructions = (
  74. [run.instructions or ""] if run.instructions else [ast.instructions or ""]
  75. )
  76. asst_ids = []
  77. ids = []
  78. if ast.tool_resources and "file_search" in ast.tool_resources:
  79. ids = (
  80. ast.tool_resources.get("file_search")
  81. .get("vector_stores")[0]
  82. .get("folder_ids")
  83. )
  84. if ids:
  85. asst_ids += ids
  86. ids = (
  87. ast.tool_resources.get("file_search")
  88. .get("vector_stores")[0]
  89. .get("file_ids")
  90. )
  91. if ids:
  92. asst_ids += ids
  93. if len(asst_ids) > 0:
  94. if len(run.file_ids) > 0:
  95. run.tools.append({"type": "knowledge_search"})
  96. else:
  97. for tool in run.tools:
  98. if tool.get("type") == "file_search":
  99. tool["type"] = "knowledge_search"
  100. tools = find_tools(run, self.session)
  101. for tool in tools:
  102. tool.configure(session=self.session, run=run)
  103. instruction_supplement = tool.instruction_supplement()
  104. if instruction_supplement:
  105. instructions += [instruction_supplement or ""]
  106. instruction = "\n".join(instructions)
  107. llm = self.__init_llm_backend(run.assistant_id)
  108. loop = True
  109. while loop:
  110. print(
  111. "looplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooplooploop"
  112. )
  113. run_steps = RunStepService.get_run_step_list(
  114. session=self.session, run_id=self.run_id, thread_id=run.thread_id
  115. )
  116. loop = self.__run_step(llm, run, run_steps, instruction, tools, memory)
  117. # 任务结束
  118. self.event_handler.pub_run_completed(run)
  119. self.event_handler.pub_done()
  120. def __run_step(
  121. self,
  122. llm: LLMBackend,
  123. run: Run,
  124. run_steps: List[RunStep],
  125. instruction: str,
  126. tools: List[BaseTool],
  127. memory: Memory,
  128. ):
  129. """
  130. 执行 run step
  131. """
  132. logging.info("step %d is running", len(run_steps) + 1)
  133. assistant_system_message = [msg_util.system_message(instruction)]
  134. # 获取已有 message 上下文记录
  135. chat_messages = self.__generate_chat_messages(
  136. MessageService.get_message_list(
  137. session=self.session, thread_id=run.thread_id
  138. ),
  139. run,
  140. )
  141. tool_call_messages = []
  142. for step in run_steps:
  143. if step.type == "tool_calls" and step.status == "completed":
  144. tool_call_messages += (
  145. self.__convert_assistant_tool_calls_to_chat_messages(step)
  146. )
  147. # tool_call_messages = tool_call_messages
  148. # memory
  149. messages = (
  150. assistant_system_message
  151. + memory.integrate_context(chat_messages)
  152. + tool_call_messages
  153. )
  154. logging.info("messages: run %s", run)
  155. logging.info(messages)
  156. logging.info(tools)
  157. response_stream = llm.run(
  158. messages=messages,
  159. model=run.model,
  160. tools=[tool.openai_function for tool in tools],
  161. tool_choice="auto" if len(run_steps) < self.max_step else "none",
  162. stream=self.stream,
  163. stream_options=run.stream_options,
  164. extra_body=run.extra_body,
  165. temperature=run.temperature,
  166. top_p=run.top_p,
  167. response_format=run.response_format,
  168. parallel_tool_calls=run.parallel_tool_calls,
  169. audio=run.audio,
  170. modalities=run.modalities,
  171. )
  172. # create message callback
  173. create_message_callback = partial(
  174. MessageService.new_message,
  175. session=self.session,
  176. assistant_id=run.assistant_id,
  177. thread_id=run.thread_id,
  178. run_id=run.id,
  179. role="assistant",
  180. )
  181. # create 'message creation' run step callback
  182. def _create_message_creation_run_step(message_id):
  183. return RunStepService.new_run_step(
  184. session=self.session,
  185. type="message_creation",
  186. assistant_id=run.assistant_id,
  187. thread_id=run.thread_id,
  188. run_id=run.id,
  189. step_details={
  190. "type": "message_creation",
  191. "message_creation": {"message_id": message_id},
  192. },
  193. )
  194. llm_callback_handler = LLMCallbackHandler(
  195. run_id=run.id,
  196. on_step_create_func=_create_message_creation_run_step,
  197. on_message_create_func=create_message_callback,
  198. event_handler=self.event_handler,
  199. )
  200. if self.stream == False and hasattr(response_stream, "choices"):
  201. response_stream = [response_stream]
  202. response_msg = llm_callback_handler.handle_llm_response(response_stream)
  203. message_creation_run_step = llm_callback_handler.step
  204. print("444444444444444444444444455555555577777777777777777777777")
  205. logging.info("chat_response_message: %s", response_msg)
  206. if msg_util.is_tool_call(response_msg):
  207. # tool & tool_call definition dict
  208. tool_calls = [
  209. tool_call_recognize(tool_call, tools)
  210. for tool_call in response_msg.tool_calls
  211. ]
  212. # new run step for tool calls
  213. new_run_step = RunStepService.new_run_step(
  214. session=self.session,
  215. type="tool_calls",
  216. assistant_id=run.assistant_id,
  217. thread_id=run.thread_id,
  218. run_id=run.id,
  219. step_details={
  220. "type": "tool_calls",
  221. "tool_calls": [tool_call_dict for _, tool_call_dict in tool_calls],
  222. },
  223. )
  224. self.event_handler.pub_run_step_created(new_run_step)
  225. self.event_handler.pub_run_step_in_progress(new_run_step)
  226. internal_tool_calls = list(
  227. filter(lambda _tool_calls: _tool_calls[0] is not None, tool_calls)
  228. )
  229. external_tool_call_dict = [
  230. tool_call_dict for tool, tool_call_dict in tool_calls if tool is None
  231. ]
  232. # 为减少线程同步逻辑,依次处理内/外 tool_call 调用
  233. if internal_tool_calls:
  234. try:
  235. print(
  236. "==========================internal_tool_callsinternal_tool_callsinternal_tool_calls"
  237. )
  238. print(internal_tool_calls)
  239. ## 线程执行有问题 可以改成异步, 这里如果是filesearch要确定只执行一次
  240. tool_calls_with_outputs = run_with_executor(
  241. executor=ThreadRunner.tool_executor,
  242. func=internal_tool_call_invoke,
  243. tasks=internal_tool_calls,
  244. timeout=tool_settings.TOOL_WORKER_EXECUTION_TIMEOUT,
  245. )
  246. new_run_step = RunStepService.update_step_details(
  247. session=self.session,
  248. run_step_id=new_run_step.id,
  249. step_details={
  250. "type": "tool_calls",
  251. "tool_calls": tool_calls_with_outputs,
  252. },
  253. completed=not external_tool_call_dict,
  254. )
  255. except Exception as e:
  256. RunStepService.to_failed(
  257. session=self.session, run_step_id=new_run_step.id, last_error=e
  258. )
  259. raise e
  260. print(
  261. "aaaaaaaaaaaaaaa===============================================================8888888888888888888888888"
  262. )
  263. print(external_tool_call_dict)
  264. if external_tool_call_dict:
  265. # run 设置为 action required,等待业务完成更新并再次拉起
  266. run = RunService.to_requires_action(
  267. session=self.session,
  268. run_id=run.id,
  269. required_action={
  270. "type": "submit_tool_outputs",
  271. "submit_tool_outputs": {"tool_calls": external_tool_call_dict},
  272. },
  273. )
  274. self.event_handler.pub_run_step_delta(
  275. step_id=new_run_step.id,
  276. step_details={
  277. "type": "tool_calls",
  278. "tool_calls": external_tool_call_dict,
  279. },
  280. )
  281. print(run)
  282. self.event_handler.pub_run_requires_action(run)
  283. else:
  284. self.event_handler.pub_run_step_completed(new_run_step)
  285. return True
  286. else:
  287. if response_msg.content == "":
  288. response_msg.content = (
  289. '[{"text": {"value": "", "annotations": []}, "type": "text"}]'
  290. )
  291. # 无 tool call 信息,message 生成结束,更新状态
  292. new_message = MessageService.modify_message_sync(
  293. session=self.session,
  294. thread_id=run.thread_id,
  295. message_id=llm_callback_handler.message.id,
  296. body=MessageUpdate(content=response_msg.content),
  297. )
  298. self.event_handler.pub_message_completed(new_message)
  299. new_step = RunStepService.update_step_details(
  300. session=self.session,
  301. run_step_id=message_creation_run_step.id,
  302. step_details={
  303. "type": "message_creation",
  304. "message_creation": {"message_id": new_message.id},
  305. },
  306. completed=True,
  307. )
  308. RunService.to_completed(session=self.session, run_id=run.id)
  309. self.event_handler.pub_run_step_completed(new_step)
  310. return False
  311. def __init_llm_backend(self, assistant_id):
  312. if settings.AUTH_ENABLE:
  313. # init llm backend with token id
  314. if self.token_id:
  315. token_id = self.token_id
  316. else:
  317. token_id = TokenRelationService.get_token_id_by_relation(
  318. session=self.session,
  319. relation_type=RelationType.Assistant,
  320. relation_id=assistant_id,
  321. )
  322. print(
  323. "token_idtoken_idtoken_idtoken_idtoken_idtoken_idtoken_idtoken_idtoken_idtoken_idtoken_idtoken_id"
  324. )
  325. print(self.token_id)
  326. print(token_id)
  327. try:
  328. if token_id is not None and len(token_id) > 0:
  329. token = TokenService.get_token_by_id(self.session, token_id)
  330. print(token)
  331. return LLMBackend(
  332. base_url=token.llm_base_url, api_key=token.llm_api_key
  333. )
  334. except Exception as e:
  335. print(e)
  336. token = {
  337. "llm_base_url": "http://172.16.12.13:3000/v1",
  338. "llm_api_key": "sk-vTqeBKDC2j6osbGt89A2202dAd1c4fE8B1D294388b569e54",
  339. }
  340. return LLMBackend(
  341. base_url=token.get("llm_base_url"), api_key=token.get("llm_api_key")
  342. )
  343. else:
  344. # init llm backend with llm settings
  345. return LLMBackend(
  346. base_url=llm_settings.OPENAI_API_BASE,
  347. api_key=llm_settings.OPENAI_API_KEY,
  348. )
  349. def __generate_chat_messages(self, messages: List[Message], run: Run):
  350. """
  351. 根据历史信息生成 chat message
  352. """
  353. chat_messages = []
  354. is_audio_num = 0
  355. for message in messages:
  356. role = message.role
  357. if role == "user":
  358. message_content = []
  359. """
  360. if message.file_ids:
  361. files = FileService.get_file_list_by_ids(
  362. session=self.session, file_ids=message.file_ids
  363. )
  364. for file in files:
  365. chat_messages.append(
  366. msg_util.new_message(
  367. role,
  368. f'The file "{file.filename}" can be used as a reference',
  369. )
  370. )
  371. else:
  372. """
  373. for content in message.content:
  374. if content["type"] == "text":
  375. message_content.append(
  376. {"type": "text", "text": content["text"]["value"]}
  377. )
  378. elif content["type"] == "image_url" and run.audio is None:
  379. message_content.append(content)
  380. elif (
  381. content.get("type") == "input_audio"
  382. and run.audio is not None
  383. and is_audio_num < 5
  384. ):
  385. message_content.append(content)
  386. is_audio_num += 1
  387. chat_messages.append(msg_util.new_message(role, message_content))
  388. elif role == "assistant":
  389. message_content = ""
  390. for content in message.content:
  391. if content["type"] == "text":
  392. message_content += content["text"]["value"]
  393. chat_messages.append(msg_util.new_message(role, message_content))
  394. return chat_messages ### 暂时只支持5条消息,后续正价token上限
  395. def __convert_assistant_tool_calls_to_chat_messages(self, run_step: RunStep):
  396. """
  397. 根据 run step 执行结果生成 message 信息
  398. 每个 tool call run step 包含两部分,调用与结果(结果可能为多个信息)
  399. """
  400. tool_calls = run_step.step_details["tool_calls"]
  401. tool_call_requests = [
  402. msg_util.tool_calls(
  403. [tool_call_request(tool_call) for tool_call in tool_calls]
  404. )
  405. ]
  406. tool_call_outputs = [
  407. msg_util.tool_call_result(
  408. tool_call_id(tool_call), tool_call_output(tool_call)
  409. )
  410. for tool_call in tool_calls
  411. ]
  412. return tool_call_requests + tool_call_outputs