thread_runner.py 21 KB

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