Team Ai
Apppublic

aphilippov/python-server-api

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
gpt_assistant_agent.py425 linesDownload Raw Back to contrib
1from collections import defaultdict2import openai3import json4import time5import logging6 7from autogen import OpenAIWrapper8from autogen.oai.openai_utils import retrieve_assistants_by_name9from autogen.agentchat.agent import Agent10from autogen.agentchat.assistant_agent import ConversableAgent11from autogen.agentchat.assistant_agent import AssistantAgent12from typing import Dict, Optional, Union, List, Tuple, Any13 14logger = logging.getLogger(__name__)15 16 17class GPTAssistantAgent(ConversableAgent):18    """19    An experimental AutoGen agent class that leverages the OpenAI Assistant API for conversational capabilities.20    This agent is unique in its reliance on the OpenAI Assistant for state management, differing from other agents like ConversableAgent.21    """22 23    def __init__(24        self,25        name="GPT Assistant",26        instructions: Optional[str] = None,27        llm_config: Optional[Union[Dict, bool]] = None,28        overwrite_instructions: bool = False,29        **kwargs,30    ):31        """32        Args:33            name (str): name of the agent. It will be used to find the existing assistant by name. Please remember to delete an old assistant with the same name if you intend to create a new assistant with the same name.34            instructions (str): instructions for the OpenAI assistant configuration.35            When instructions is not None, the system message of the agent will be36            set to the provided instructions and used in the assistant run, irrespective37            of the overwrite_instructions flag. But when instructions is None,38            and the assistant does not exist, the system message will be set to39            AssistantAgent.DEFAULT_SYSTEM_MESSAGE. If the assistant exists, the40            system message will be set to the existing assistant instructions.41            llm_config (dict or False): llm inference configuration.42                - assistant_id: ID of the assistant to use. If None, a new assistant will be created.43                - model: Model to use for the assistant (gpt-4-1106-preview, gpt-3.5-turbo-1106).44                - check_every_ms: check thread run status interval45                - tools: Give Assistants access to OpenAI-hosted tools like Code Interpreter and Knowledge Retrieval,46                        or build your own tools using Function calling. ref https://platform.openai.com/docs/assistants/tools47                - file_ids: files used by retrieval in run48            overwrite_instructions (bool): whether to overwrite the instructions of an existing assistant. This parameter is in effect only when assistant_id is specified in llm_config.49            kwargs (dict): Additional configuration options for the agent.50                - verbose (bool): If set to True, enables more detailed output from the assistant thread.51                - Other kwargs: Except verbose, others are passed directly to ConversableAgent.52        """53        # Use AutoGen OpenAIWrapper to create a client54        oai_wrapper = OpenAIWrapper(**llm_config)55        if len(oai_wrapper._clients) > 1:56            logger.warning("GPT Assistant only supports one OpenAI client. Using the first client in the list.")57        self._openai_client = oai_wrapper._clients[0]58        openai_assistant_id = llm_config.get("assistant_id", None)59        if openai_assistant_id is None:60            # try to find assistant by name first61            candidate_assistants = retrieve_assistants_by_name(self._openai_client, name)62            if len(candidate_assistants) > 0:63                # Filter out candidates with the same name but different instructions, file IDs, and function names.64                candidate_assistants = self.find_matching_assistant(65                    candidate_assistants, instructions, llm_config.get("tools", []), llm_config.get("file_ids", [])66                )67 68            if len(candidate_assistants) == 0:69                logger.warning("No matching assistant found, creating a new assistant")70                # create a new assistant71                if instructions is None:72                    logger.warning(73                        "No instructions were provided for new assistant. Using default instructions from AssistantAgent.DEFAULT_SYSTEM_MESSAGE."74                    )75                    instructions = AssistantAgent.DEFAULT_SYSTEM_MESSAGE76                self._openai_assistant = self._openai_client.beta.assistants.create(77                    name=name,78                    instructions=instructions,79                    tools=llm_config.get("tools", []),80                    model=llm_config.get("model", "gpt-4-1106-preview"),81                    file_ids=llm_config.get("file_ids", []),82                )83            else:84                logger.warning(85                    "Matching assistant found, using the first matching assistant: %s",86                    candidate_assistants[0].__dict__,87                )88                self._openai_assistant = candidate_assistants[0]89        else:90            # retrieve an existing assistant91            self._openai_assistant = self._openai_client.beta.assistants.retrieve(openai_assistant_id)92            # if no instructions are provided, set the instructions to the existing instructions93            if instructions is None:94                logger.warning(95                    "No instructions were provided for given assistant. Using existing instructions from assistant API."96                )97                instructions = self.get_assistant_instructions()98            elif overwrite_instructions is True:99                logger.warning(100                    "overwrite_instructions is True. Provided instructions will be used and will modify the assistant in the API"101                )102                self._openai_assistant = self._openai_client.beta.assistants.update(103                    assistant_id=openai_assistant_id,104                    instructions=instructions,105                )106            else:107                logger.warning(108                    "overwrite_instructions is False. Provided instructions will be used without permanently modifying the assistant in the API."109                )110 111        self._verbose = kwargs.pop("verbose", False)112        super().__init__(113            name=name, system_message=instructions, human_input_mode="NEVER", llm_config=llm_config, **kwargs114        )115 116        # lazily create threads117        self._openai_threads = {}118        self._unread_index = defaultdict(int)119        self.register_reply(Agent, GPTAssistantAgent._invoke_assistant)120 121    def _invoke_assistant(122        self,123        messages: Optional[List[Dict]] = None,124        sender: Optional[Agent] = None,125        config: Optional[Any] = None,126    ) -> Tuple[bool, Union[str, Dict, None]]:127        """128        Invokes the OpenAI assistant to generate a reply based on the given messages.129 130        Args:131            messages: A list of messages in the conversation history with the sender.132            sender: The agent instance that sent the message.133            config: Optional configuration for message processing.134 135        Returns:136            A tuple containing a boolean indicating success and the assistant's reply.137        """138 139        if messages is None:140            messages = self._oai_messages[sender]141        unread_index = self._unread_index[sender] or 0142        pending_messages = messages[unread_index:]143 144        # Check and initiate a new thread if necessary145        if self._openai_threads.get(sender, None) is None:146            self._openai_threads[sender] = self._openai_client.beta.threads.create(147                messages=[],148            )149        assistant_thread = self._openai_threads[sender]150        # Process each unread message151        for message in pending_messages:152            self._openai_client.beta.threads.messages.create(153                thread_id=assistant_thread.id,154                content=message["content"],155                role=message["role"],156            )157 158        # Create a new run to get responses from the assistant159        run = self._openai_client.beta.threads.runs.create(160            thread_id=assistant_thread.id,161            assistant_id=self._openai_assistant.id,162            # pass the latest system message as instructions163            instructions=self.system_message,164        )165 166        run_response_messages = self._get_run_response(assistant_thread, run)167        assert len(run_response_messages) > 0, "No response from the assistant."168 169        response = {170            "role": run_response_messages[-1]["role"],171            "content": "",172        }173        for message in run_response_messages:174            # just logging or do something with the intermediate messages?175            # if current response is not empty and there is more, append new lines176            if len(response["content"]) > 0:177                response["content"] += "\n\n"178            response["content"] += message["content"]179 180        self._unread_index[sender] = len(self._oai_messages[sender]) + 1181        return True, response182 183    def _get_run_response(self, thread, run):184        """185        Waits for and processes the response of a run from the OpenAI assistant.186 187        Args:188            run: The run object initiated with the OpenAI assistant.189 190        Returns:191            Updated run object, status of the run, and response messages.192        """193        while True:194            run = self._wait_for_run(run.id, thread.id)195            if run.status == "completed":196                response_messages = self._openai_client.beta.threads.messages.list(thread.id, order="asc")197 198                new_messages = []199                for msg in response_messages:200                    if msg.run_id == run.id:201                        for content in msg.content:202                            if content.type == "text":203                                new_messages.append(204                                    {"role": msg.role, "content": self._format_assistant_message(content.text)}205                                )206                            elif content.type == "image_file":207                                new_messages.append(208                                    {209                                        "role": msg.role,210                                        "content": f"Received file id={content.image_file.file_id}",211                                    }212                                )213                return new_messages214            elif run.status == "requires_action":215                actions = []216                for tool_call in run.required_action.submit_tool_outputs.tool_calls:217                    function = tool_call.function218                    is_exec_success, tool_response = self.execute_function(function.dict(), self._verbose)219                    tool_response["metadata"] = {220                        "tool_call_id": tool_call.id,221                        "run_id": run.id,222                        "thread_id": thread.id,223                    }224 225                    logger.info(226                        "Intermediate executing(%s, Success: %s) : %s",227                        tool_response["name"],228                        is_exec_success,229                        tool_response["content"],230                    )231                    actions.append(tool_response)232 233                submit_tool_outputs = {234                    "tool_outputs": [235                        {"output": action["content"], "tool_call_id": action["metadata"]["tool_call_id"]}236                        for action in actions237                    ],238                    "run_id": run.id,239                    "thread_id": thread.id,240                }241 242                run = self._openai_client.beta.threads.runs.submit_tool_outputs(**submit_tool_outputs)243            else:244                run_info = json.dumps(run.dict(), indent=2)245                raise ValueError(f"Unexpected run status: {run.status}. Full run info:\n\n{run_info})")246 247    def _wait_for_run(self, run_id: str, thread_id: str) -> Any:248        """249        Waits for a run to complete or reach a final state.250 251        Args:252            run_id: The ID of the run.253            thread_id: The ID of the thread associated with the run.254 255        Returns:256            The updated run object after completion or reaching a final state.257        """258        in_progress = True259        while in_progress:260            run = self._openai_client.beta.threads.runs.retrieve(run_id, thread_id=thread_id)261            in_progress = run.status in ("in_progress", "queued")262            if in_progress:263                time.sleep(self.llm_config.get("check_every_ms", 1000) / 1000)264        return run265 266    def _format_assistant_message(self, message_content):267        """268        Formats the assistant's message to include annotations and citations.269        """270 271        annotations = message_content.annotations272        citations = []273 274        # Iterate over the annotations and add footnotes275        for index, annotation in enumerate(annotations):276            # Replace the text with a footnote277            message_content.value = message_content.value.replace(annotation.text, f" [{index}]")278 279            # Gather citations based on annotation attributes280            if file_citation := getattr(annotation, "file_citation", None):281                try:282                    cited_file = self._openai_client.files.retrieve(file_citation.file_id)283                    citations.append(f"[{index}] {cited_file.filename}: {file_citation.quote}")284                except Exception as e:285                    logger.error(f"Error retrieving file citation: {e}")286            elif file_path := getattr(annotation, "file_path", None):287                try:288                    cited_file = self._openai_client.files.retrieve(file_path.file_id)289                    citations.append(f"[{index}] Click <here> to download {cited_file.filename}")290                except Exception as e:291                    logger.error(f"Error retrieving file citation: {e}")292                # Note: File download functionality not implemented above for brevity293 294        # Add footnotes to the end of the message before displaying to user295        message_content.value += "\n" + "\n".join(citations)296        return message_content.value297 298    def can_execute_function(self, name: str) -> bool:299        """Whether the agent can execute the function."""300        return False301 302    def reset(self):303        """304        Resets the agent, clearing any existing conversation thread and unread message indices.305        """306        super().reset()307        for thread in self._openai_threads.values():308            # Delete the existing thread to start fresh in the next conversation309            self._openai_client.beta.threads.delete(thread.id)310        self._openai_threads = {}311        # Clear the record of unread messages312        self._unread_index.clear()313 314    def clear_history(self, agent: Optional[Agent] = None):315        """Clear the chat history of the agent.316 317        Args:318            agent: the agent with whom the chat history to clear. If None, clear the chat history with all agents.319        """320        super().clear_history(agent)321        if self._openai_threads.get(agent, None) is not None:322            # Delete the existing thread to start fresh in the next conversation323            thread = self._openai_threads[agent]324            logger.info("Clearing thread %s", thread.id)325            self._openai_client.beta.threads.delete(thread.id)326            self._openai_threads.pop(agent)327            self._unread_index[agent] = 0328 329    def pretty_print_thread(self, thread):330        """Pretty print the thread."""331        if thread is None:332            print("No thread to print")333            return334        # NOTE: that list may not be in order, sorting by created_at is important335        messages = self._openai_client.beta.threads.messages.list(336            thread_id=thread.id,337        )338        messages = sorted(messages.data, key=lambda x: x.created_at)339        print("~~~~~~~THREAD CONTENTS~~~~~~~")340        for message in messages:341            content_types = [content.type for content in message.content]342            print(f"[{message.created_at}]", message.role, ": [", ", ".join(content_types), "]")343            for content in message.content:344                content_type = content.type345                if content_type == "text":346                    print(content.type, ": ", content.text.value)347                elif content_type == "image_file":348                    print(content.type, ": ", content.image_file.file_id)349                else:350                    print(content.type, ": ", content)351        print("~~~~~~~~~~~~~~~~~~~~~~~~~~~~~")352 353    @property354    def oai_threads(self) -> Dict[Agent, Any]:355        """Return the threads of the agent."""356        return self._openai_threads357 358    @property359    def assistant_id(self):360        """Return the assistant id"""361        return self._openai_assistant.id362 363    @property364    def openai_client(self):365        return self._openai_client366 367    def get_assistant_instructions(self):368        """Return the assistant instructions from OAI assistant API"""369        return self._openai_assistant.instructions370 371    def delete_assistant(self):372        """Delete the assistant from OAI assistant API"""373        logger.warning("Permanently deleting assistant...")374        self._openai_client.beta.assistants.delete(self.assistant_id)375 376    def find_matching_assistant(self, candidate_assistants, instructions, tools, file_ids):377        """378        Find the matching assistant from a list of candidate assistants.379        Filter out candidates with the same name but different instructions, file IDs, and function names.380        TODO: implement accurate match based on assistant metadata fields.381        """382        matching_assistants = []383 384        # Preprocess the required tools for faster comparison385        required_tool_types = set(tool.get("type") for tool in tools)386        required_function_names = set(387            tool.get("function", {}).get("name")388            for tool in tools389            if tool.get("type") not in ["code_interpreter", "retrieval"]390        )391        required_file_ids = set(file_ids)  # Convert file_ids to a set for unordered comparison392 393        for assistant in candidate_assistants:394            # Check if instructions are similar395            if instructions and instructions != getattr(assistant, "instructions", None):396                logger.warning(397                    "instructions not match, skip assistant(%s): %s",398                    assistant.id,399                    getattr(assistant, "instructions", None),400                )401                continue402 403            # Preprocess the assistant's tools404            assistant_tool_types = set(tool.type for tool in assistant.tools)405            assistant_function_names = set(tool.function.name for tool in assistant.tools if hasattr(tool, "function"))406            assistant_file_ids = set(getattr(assistant, "file_ids", []))  # Convert to set for comparison407 408            # Check if the tool types, function names, and file IDs match409            if required_tool_types != assistant_tool_types or required_function_names != assistant_function_names:410                logger.warning(411                    "tools not match, skip assistant(%s): tools %s, functions %s",412                    assistant.id,413                    assistant_tool_types,414                    assistant_function_names,415                )416                continue417            if required_file_ids != assistant_file_ids:418                logger.warning("file_ids not match, skip assistant(%s): %s", assistant.id, assistant_file_ids)419                continue420 421            # Append assistant to matching list if all conditions are met422            matching_assistants.append(assistant)423 424        return matching_assistants425