OpenAI Assistants API:优化响应等待逻辑与实现增量输出
OpenAI Assistants API 代码优化:轮询逻辑改进与增量式回复实现
问题背景
我编写了一套与OpenAI Assistants API交互的Python代码,但对其中_wait_for_run_completion方法的while循环实现不满意,希望找到更优的处理方式;同时想了解如何实现助手回答的增量式显示(类似ChatGPT实时打字的效果)。
核心类原代码
import os import openai from dotenv import load_dotenv import time class OpenAIChatAssistant: def __init__(self, assistant_id, model="gpt-4o"): self.assistant_id = assistant_id self.model = model if self.model != "just_copy": load_dotenv() openai.api_key = os.environ.get("OPENAI_API_KEY") self.client = openai.OpenAI() self._create_new_thread() print('new instance started') def _create_new_thread(self): self.thread = self.client.beta.threads.create() self.thread_id = self.thread.id print(self.thread_id) def reset_thread(self): if self.model != "just_copy": self._create_new_thread() def set_model(self, model_name): self.model = model_name if self.model != "just_copy" and not hasattr(self, 'client'): load_dotenv() openai.api_key = os.environ.get("OPENAI_API_KEY") self.client = openai.OpenAI() self._create_new_thread() def send_message(self, message): if self.model == "just_copy": return message self.client.beta.threads.messages.create( thread_id=self.thread_id, role="user", content=message ) run = self.client.beta.threads.runs.create( thread_id=self.thread_id, assistant_id=self.assistant_id, model=self.model ) return self._wait_for_run_completion(run.id) def _wait_for_run_completion(self, run_id, sleep_interval=1): counter = 1 while True: try: run = self.client.beta.threads.runs.retrieve(thread_id=self.thread_id, run_id=run_id) if run.completed_at: messages = self.client.beta.threads.messages.list(thread_id=self.thread_id) last_message = messages.data[0] response = last_message.content[0].text.value print(f'hello {counter}') return response except Exception as e: raise RuntimeError(f"An error occurred while retrieving answer: {e}") counter += 1 time.sleep(sleep_interval)
控制台原使用示例
import os from openai_chat_assistant import OpenAIChatAssistant def main(): assistant_id = "asst_..." chat_assistant = OpenAIChatAssistant(assistant_id) while True: question = input("Enter your question (or 'exit' to quit, 'clean' to reset): ") if question.lower() == 'exit': break elif question.lower() == 'clean': os.system('cls' if os.name == 'nt' else 'clear') chat_assistant.reset_thread() print("Console cleared and thread reset.") else: response = chat_assistant.send_message(question) print(f"Assistant Response: {response}") if __name__ == "__main__": main()
环境配置
.env文件内容:
OPENAI_API_KEY=sk-proj-...
优化方案
1. 轮询逻辑改进
原轮询采用固定间隔查询,效率较低且未处理异常状态。推荐使用指数退避(间隔逐渐增大),同时增加对run多种状态的处理(如失败、需要人工干预等),避免无限循环。
优化后的_wait_for_run_completion方法:
def _wait_for_run_completion(self, run_id, initial_sleep=1, max_sleep=30): sleep_interval = initial_sleep while True: try: run = self.client.beta.threads.runs.retrieve(thread_id=self.thread_id, run_id=run_id) # 处理不同运行状态 if run.status == "completed": messages = self.client.beta.threads.messages.list(thread_id=self.thread_id) last_message = messages.data[0] return last_message.content[0].text.value elif run.status == "failed": raise RuntimeError(f"Run failed: {run.last_error.message}") elif run.status in ["requires_action", "cancelling", "cancelled"]: raise RuntimeError(f"Run entered unexpected state: {run.status}") # 指数退避,避免频繁请求 time.sleep(sleep_interval) sleep_interval = min(sleep_interval * 1.5, max_sleep) except Exception as e: raise RuntimeError(f"Error retrieving run status: {e}")
2. 增量式回复实现
要实现实时打字效果,需使用OpenAI的流式响应(stream=True)。在创建run时开启流式模式,迭代返回的内容片段并实时输出。
新增流式处理方法send_message_stream:
def send_message_stream(self, message): if self.model == "just_copy": for char in message: print(char, end="", flush=True) time.sleep(0.05) print() return message # 添加用户消息到线程 self.client.beta.threads.messages.create( thread_id=self.thread_id, role="user", content=message ) # 开启流式创建run with self.client.beta.threads.runs.stream( thread_id=self.thread_id, assistant_id=self.assistant_id, model=self.model ) as stream: full_response = "" print("Assistant Response: ", end="", flush=True) for event in stream: # 处理流式返回的内容片段 if event.event == "thread.message.delta": delta = event.data.delta if delta.content: for content_part in delta.content: if content_part.type == "text": text_chunk = content_part.text.value full_response += text_chunk print(text_chunk, end="", flush=True) print() return full_response
优化后的完整类代码
import os import openai from dotenv import load_dotenv import time class OpenAIChatAssistant: def __init__(self, assistant_id, model="gpt-4o"): self.assistant_id = assistant_id self.model = model if self.model != "just_copy": load_dotenv() openai.api_key = os.environ.get("OPENAI_API_KEY") self.client = openai.OpenAI() self._create_new_thread() print('new instance started') def _create_new_thread(self): self.thread = self.client.beta.threads.create() self.thread_id = self.thread.id print(self.thread_id) def reset_thread(self): if self.model != "just_copy": self._create_new_thread() def set_model(self, model_name): self.model = model_name if self.model != "just_copy" and not hasattr(self, 'client'): load_dotenv() openai.api_key = os.environ.get("OPENAI_API_KEY") self.client = openai.OpenAI() self._create_new_thread() def send_message(self, message): if self.model == "just_copy": return message self.client.beta.threads.messages.create( thread_id=self.thread_id, role="user", content=message ) run = self.client.beta.threads.runs.create( thread_id=self.thread_id, assistant_id=self.assistant_id, model=self.model ) return self._wait_for_run_completion(run.id) def _wait_for_run_completion(self, run_id, initial_sleep=1, max_sleep=30): sleep_interval = initial_sleep while True: try: run = self.client.beta.threads.runs.retrieve(thread_id=self.thread_id, run_id=run_id) if run.status == "completed": messages = self.client.beta.threads.messages.list(thread_id=self.thread_id) last_message = messages.data[0] return last_message.content[0].text.value elif run.status == "failed": raise RuntimeError(f"Run failed: {run.last_error.message}") elif run.status in ["requires_action", "cancelling", "cancelled"]: raise RuntimeError(f"Run entered unexpected state: {run.status}") time.sleep(sleep_interval) sleep_interval = min(sleep_interval * 1.5, max_sleep) except Exception as e: raise RuntimeError(f"Error retrieving run status: {e}") def send_message_stream(self, message): if self.model == "just_copy": for char in message: print(char, end="", flush=True) time.sleep(0.05) print() return message self.client.beta.threads.messages.create( thread_id=self.thread_id, role="user", content=message ) with self.client.beta.threads.runs.stream( thread_id=self.thread_id, assistant_id=self.assistant_id, model=self.model ) as stream: full_response = "" print("Assistant Response: ", end="", flush=True) for event in stream: if event.event == "thread.message.delta": delta = event.data.delta if delta.content: for content_part in delta.content: if content_part.type == "text": text_chunk = content_part.text.value full_response += text_chunk print(text_chunk, end="", flush=True) print() return full_response
控制台使用示例(适配流式方法)
修改主函数,调用send_message_stream实现增量显示:
import os from openai_chat_assistant import OpenAIChatAssistant def main(): assistant_id = "asst_..." chat_assistant = OpenAIChatAssistant(assistant_id) while True: question = input("\nEnter your question (or 'exit' to quit, 'clean' to reset): ") if question.lower() == 'exit': break elif question.lower() == 'clean': os.system('cls' if os.name == 'nt' else 'clear') chat_assistant.reset_thread() print("Console cleared and thread reset.") else: chat_assistant.send_message_stream(question) if __name__ == "__main__": main()
内容的提问来源于stack exchange,提问作者Krzysztof Krysztofczyk
相关产品推荐
相关产品推荐

