Team Ai
Apppublic

aphilippov/python-server-api

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
retrieve_user_proxy_agent.py457 linesDownload Raw Back to contrib
1import re2 3try:4    import chromadb5except ImportError:6    raise ImportError("Please install dependencies first. `pip install pyautogen[retrievechat]`")7from autogen.agentchat.agent import Agent8from autogen.agentchat import UserProxyAgent9from autogen.retrieve_utils import create_vector_db_from_dir, query_vector_db, TEXT_FORMATS10from autogen.token_count_utils import count_token11from autogen.code_utils import extract_code12from autogen import logger13 14from typing import Callable, Dict, Optional, Union, List, Tuple, Any15from IPython import get_ipython16 17try:18    from termcolor import colored19except ImportError:20 21    def colored(x, *args, **kwargs):22        return x23 24 25PROMPT_DEFAULT = """You're a retrieve augmented chatbot. You answer user's questions based on your own knowledge and the26context provided by the user. You should follow the following steps to answer a question:27Step 1, you estimate the user's intent based on the question and context. The intent can be a code generation task or28a question answering task.29Step 2, you reply based on the intent.30If you can't answer the question with or without the current context, you should reply exactly `UPDATE CONTEXT`.31If user's intent is code generation, you must obey the following rules:32Rule 1. You MUST NOT install any packages because all the packages needed are already installed.33Rule 2. You must follow the formats below to write your code:34```language35# your code36```37 38If user's intent is question answering, you must give as short an answer as possible.39 40User's question is: {input_question}41 42Context is: {input_context}43"""44 45PROMPT_CODE = """You're a retrieve augmented coding assistant. You answer user's questions based on your own knowledge and the46context provided by the user.47If you can't answer the question with or without the current context, you should reply exactly `UPDATE CONTEXT`.48For code generation, you must obey the following rules:49Rule 1. You MUST NOT install any packages because all the packages needed are already installed.50Rule 2. You must follow the formats below to write your code:51```language52# your code53```54 55User's question is: {input_question}56 57Context is: {input_context}58"""59 60PROMPT_QA = """You're a retrieve augmented chatbot. You answer user's questions based on your own knowledge and the61context provided by the user.62If you can't answer the question with or without the current context, you should reply exactly `UPDATE CONTEXT`.63You must give as short an answer as possible.64 65User's question is: {input_question}66 67Context is: {input_context}68"""69 70 71class RetrieveUserProxyAgent(UserProxyAgent):72    def __init__(73        self,74        name="RetrieveChatAgent",  # default set to RetrieveChatAgent75        human_input_mode: Optional[str] = "ALWAYS",76        is_termination_msg: Optional[Callable[[Dict], bool]] = None,77        retrieve_config: Optional[Dict] = None,  # config for the retrieve agent78        **kwargs,79    ):80        """81        Args:82            name (str): name of the agent.83            human_input_mode (str): whether to ask for human inputs every time a message is received.84                Possible values are "ALWAYS", "TERMINATE", "NEVER".85                (1) When "ALWAYS", the agent prompts for human input every time a message is received.86                    Under this mode, the conversation stops when the human input is "exit",87                    or when is_termination_msg is True and there is no human input.88                (2) When "TERMINATE", the agent only prompts for human input only when a termination message is received or89                    the number of auto reply reaches the max_consecutive_auto_reply.90                (3) When "NEVER", the agent will never prompt for human input. Under this mode, the conversation stops91                    when the number of auto reply reaches the max_consecutive_auto_reply or when is_termination_msg is True.92            is_termination_msg (function): a function that takes a message in the form of a dictionary93                and returns a boolean value indicating if this received message is a termination message.94                The dict can contain the following keys: "content", "role", "name", "function_call".95            retrieve_config (dict or None): config for the retrieve agent.96                To use default config, set to None. Otherwise, set to a dictionary with the following keys:97                - task (Optional, str): the task of the retrieve chat. Possible values are "code", "qa" and "default". System98                    prompt will be different for different tasks. The default value is `default`, which supports both code and qa.99                - client (Optional, chromadb.Client): the chromadb client. If key not provided, a default client `chromadb.Client()`100                    will be used. If you want to use other vector db, extend this class and override the `retrieve_docs` function.101                - docs_path (Optional, Union[str, List[str]]): the path to the docs directory. It can also be the path to a single file,102                    the url to a single file or a list of directories, files and urls. Default is None, which works only if the collection is already created.103                - extra_docs (Optional, bool): when true, allows adding documents with unique IDs without overwriting existing ones; when false, it replaces existing documents using default IDs, risking collection overwrite.,104                    when set to true it enables the system to assign unique IDs starting from "length+i" for new document chunks, preventing the replacement of existing documents and facilitating the addition of more content to the collection..105                    By default, "extra_docs" is set to false, starting document IDs from zero. This poses a risk as new documents might overwrite existing ones, potentially causing unintended loss or alteration of data in the collection.106                - collection_name (Optional, str): the name of the collection.107                    If key not provided, a default name `autogen-docs` will be used.108                - model (Optional, str): the model to use for the retrieve chat.109                    If key not provided, a default model `gpt-4` will be used.110                - chunk_token_size (Optional, int): the chunk token size for the retrieve chat.111                    If key not provided, a default size `max_tokens * 0.4` will be used.112                - context_max_tokens (Optional, int): the context max token size for the retrieve chat.113                    If key not provided, a default size `max_tokens * 0.8` will be used.114                - chunk_mode (Optional, str): the chunk mode for the retrieve chat. Possible values are115                    "multi_lines" and "one_line". If key not provided, a default mode `multi_lines` will be used.116                - must_break_at_empty_line (Optional, bool): chunk will only break at empty line if True. Default is True.117                    If chunk_mode is "one_line", this parameter will be ignored.118                - embedding_model (Optional, str): the embedding model to use for the retrieve chat.119                    If key not provided, a default model `all-MiniLM-L6-v2` will be used. All available models120                    can be found at `https://www.sbert.net/docs/pretrained_models.html`. The default model is a121                    fast model. If you want to use a high performance model, `all-mpnet-base-v2` is recommended.122                - embedding_function (Optional, Callable): the embedding function for creating the vector db. Default is None,123                    SentenceTransformer with the given `embedding_model` will be used. If you want to use OpenAI, Cohere, HuggingFace or124                    other embedding functions, you can pass it here, follow the examples in `https://docs.trychroma.com/embeddings`.125                - customized_prompt (Optional, str): the customized prompt for the retrieve chat. Default is None.126                - customized_answer_prefix (Optional, str): the customized answer prefix for the retrieve chat. Default is "".127                    If not "" and the customized_answer_prefix is not in the answer, `Update Context` will be triggered.128                - update_context (Optional, bool): if False, will not apply `Update Context` for interactive retrieval. Default is True.129                - get_or_create (Optional, bool): if True, will create/return a collection for the retrieve chat. This is the same as that used in chromadb.130                    Default is False. Will raise ValueError if the collection already exists and get_or_create is False. Will be set to True if docs_path is None.131                - custom_token_count_function (Optional, Callable): a custom function to count the number of tokens in a string.132                    The function should take (text:str, model:str) as input and return the token_count(int). the retrieve_config["model"] will be passed in the function.133                    Default is autogen.token_count_utils.count_token that uses tiktoken, which may not be accurate for non-OpenAI models.134                - custom_text_split_function (Optional, Callable): a custom function to split a string into a list of strings.135                    Default is None, will use the default function in `autogen.retrieve_utils.split_text_to_chunks`.136                - custom_text_types (Optional, List[str]): a list of file types to be processed. Default is `autogen.retrieve_utils.TEXT_FORMATS`.137                    This only applies to files under the directories in `docs_path`. Explicitly included files and urls will be chunked regardless of their types.138                - recursive (Optional, bool): whether to search documents recursively in the docs_path. Default is True.139            **kwargs (dict): other kwargs in [UserProxyAgent](../user_proxy_agent#__init__).140 141        Example of overriding retrieve_docs:142        If you have set up a customized vector db, and it's not compatible with chromadb, you can easily plug in it with below code.143        ```python144        class MyRetrieveUserProxyAgent(RetrieveUserProxyAgent):145            def query_vector_db(146                self,147                query_texts: List[str],148                n_results: int = 10,149                search_string: str = "",150                **kwargs,151            ) -> Dict[str, Union[List[str], List[List[str]]]]:152                # define your own query function here153                pass154 155            def retrieve_docs(self, problem: str, n_results: int = 20, search_string: str = "", **kwargs):156                results = self.query_vector_db(157                    query_texts=[problem],158                    n_results=n_results,159                    search_string=search_string,160                    **kwargs,161                )162 163                self._results = results164                print("doc_ids: ", results["ids"])165        ```166        """167        super().__init__(168            name=name,169            human_input_mode=human_input_mode,170            **kwargs,171        )172 173        self._retrieve_config = {} if retrieve_config is None else retrieve_config174        self._task = self._retrieve_config.get("task", "default")175        self._client = self._retrieve_config.get("client", chromadb.Client())176        self._docs_path = self._retrieve_config.get("docs_path", None)177        self._extra_docs = self._retrieve_config.get("extra_docs", False)178        self._collection_name = self._retrieve_config.get("collection_name", "autogen-docs")179        if "docs_path" not in self._retrieve_config:180            logger.warning(181                "docs_path is not provided in retrieve_config. "182                f"Will raise ValueError if the collection `{self._collection_name}` doesn't exist. "183                "Set docs_path to None to suppress this warning."184            )185        self._model = self._retrieve_config.get("model", "gpt-4")186        self._max_tokens = self.get_max_tokens(self._model)187        self._chunk_token_size = int(self._retrieve_config.get("chunk_token_size", self._max_tokens * 0.4))188        self._chunk_mode = self._retrieve_config.get("chunk_mode", "multi_lines")189        self._must_break_at_empty_line = self._retrieve_config.get("must_break_at_empty_line", True)190        self._embedding_model = self._retrieve_config.get("embedding_model", "all-MiniLM-L6-v2")191        self._embedding_function = self._retrieve_config.get("embedding_function", None)192        self.customized_prompt = self._retrieve_config.get("customized_prompt", None)193        self.customized_answer_prefix = self._retrieve_config.get("customized_answer_prefix", "").upper()194        self.update_context = self._retrieve_config.get("update_context", True)195        self._get_or_create = self._retrieve_config.get("get_or_create", False) if self._docs_path is not None else True196        self.custom_token_count_function = self._retrieve_config.get("custom_token_count_function", count_token)197        self.custom_text_split_function = self._retrieve_config.get("custom_text_split_function", None)198        self._custom_text_types = self._retrieve_config.get("custom_text_types", TEXT_FORMATS)199        self._recursive = self._retrieve_config.get("recursive", True)200        self._context_max_tokens = self._max_tokens * 0.8201        self._collection = True if self._docs_path is None else False  # whether the collection is created202        self._ipython = get_ipython()203        self._doc_idx = -1  # the index of the current used doc204        self._results = {}  # the results of the current query205        self._intermediate_answers = set()  # the intermediate answers206        self._doc_contents = []  # the contents of the current used doc207        self._doc_ids = []  # the ids of the current used doc208        self._search_string = ""  # the search string used in the current query209        # update the termination message function210        self._is_termination_msg = (211            self._is_termination_msg_retrievechat if is_termination_msg is None else is_termination_msg212        )213        self.register_reply(Agent, RetrieveUserProxyAgent._generate_retrieve_user_reply, position=2)214 215    def _is_termination_msg_retrievechat(self, message):216        """Check if a message is a termination message.217        For code generation, terminate when no code block is detected. Currently only detect python code blocks.218        For question answering, terminate when don't update context, i.e., answer is given.219        """220        if isinstance(message, dict):221            message = message.get("content")222            if message is None:223                return False224        cb = extract_code(message)225        contain_code = False226        for c in cb:227            # todo: support more languages228            if c[0] == "python":229                contain_code = True230                break231        update_context_case1, update_context_case2 = self._check_update_context(message)232        return not (contain_code or update_context_case1 or update_context_case2)233 234    @staticmethod235    def get_max_tokens(model="gpt-3.5-turbo"):236        if "32k" in model:237            return 32000238        elif "16k" in model:239            return 16000240        elif "gpt-4" in model:241            return 8000242        else:243            return 4000244 245    def _reset(self, intermediate=False):246        self._doc_idx = -1  # the index of the current used doc247        self._results = {}  # the results of the current query248        if not intermediate:249            self._intermediate_answers = set()  # the intermediate answers250            self._doc_contents = []  # the contents of the current used doc251            self._doc_ids = []  # the ids of the current used doc252 253    def _get_context(self, results: Dict[str, Union[List[str], List[List[str]]]]):254        doc_contents = ""255        current_tokens = 0256        _doc_idx = self._doc_idx257        _tmp_retrieve_count = 0258        for idx, doc in enumerate(results["documents"][0]):259            if idx <= _doc_idx:260                continue261            if results["ids"][0][idx] in self._doc_ids:262                continue263            _doc_tokens = self.custom_token_count_function(doc, self._model)264            if _doc_tokens > self._context_max_tokens:265                func_print = f"Skip doc_id {results['ids'][0][idx]} as it is too long to fit in the context."266                print(colored(func_print, "green"), flush=True)267                self._doc_idx = idx268                continue269            if current_tokens + _doc_tokens > self._context_max_tokens:270                break271            func_print = f"Adding doc_id {results['ids'][0][idx]} to context."272            print(colored(func_print, "green"), flush=True)273            current_tokens += _doc_tokens274            doc_contents += doc + "\n"275            self._doc_idx = idx276            self._doc_ids.append(results["ids"][0][idx])277            self._doc_contents.append(doc)278            _tmp_retrieve_count += 1279            if _tmp_retrieve_count >= self.n_results:280                break281        return doc_contents282 283    def _generate_message(self, doc_contents, task="default"):284        if not doc_contents:285            print(colored("No more context, will terminate.", "green"), flush=True)286            return "TERMINATE"287        if self.customized_prompt:288            message = self.customized_prompt.format(input_question=self.problem, input_context=doc_contents)289        elif task.upper() == "CODE":290            message = PROMPT_CODE.format(input_question=self.problem, input_context=doc_contents)291        elif task.upper() == "QA":292            message = PROMPT_QA.format(input_question=self.problem, input_context=doc_contents)293        elif task.upper() == "DEFAULT":294            message = PROMPT_DEFAULT.format(input_question=self.problem, input_context=doc_contents)295        else:296            raise NotImplementedError(f"task {task} is not implemented.")297        return message298 299    def _check_update_context(self, message):300        if isinstance(message, dict):301            message = message.get("content", "")302        elif not isinstance(message, str):303            message = ""304        update_context_case1 = "UPDATE CONTEXT" in message[-20:].upper() or "UPDATE CONTEXT" in message[:20].upper()305        update_context_case2 = self.customized_answer_prefix and self.customized_answer_prefix not in message.upper()306        return update_context_case1, update_context_case2307 308    def _generate_retrieve_user_reply(309        self,310        messages: Optional[List[Dict]] = None,311        sender: Optional[Agent] = None,312        config: Optional[Any] = None,313    ) -> Tuple[bool, Union[str, Dict, None]]:314        """In this function, we will update the context and reset the conversation based on different conditions.315        We'll update the context and reset the conversation if update_context is True and either of the following:316        (1) the last message contains "UPDATE CONTEXT",317        (2) the last message doesn't contain "UPDATE CONTEXT" and the customized_answer_prefix is not in the message.318        """319        if config is None:320            config = self321        if messages is None:322            messages = self._oai_messages[sender]323        message = messages[-1]324        update_context_case1, update_context_case2 = self._check_update_context(message)325        if (update_context_case1 or update_context_case2) and self.update_context:326            print(colored("Updating context and resetting conversation.", "green"), flush=True)327            # extract the first sentence in the response as the intermediate answer328            _message = message.get("content", "").split("\n")[0].strip()329            _intermediate_info = re.split(r"(?<=[.!?])\s+", _message)330            self._intermediate_answers.add(_intermediate_info[0])331 332            if update_context_case1:333                # try to get more context from the current retrieved doc results because the results may be too long to fit334                # in the LLM context.335                doc_contents = self._get_context(self._results)336 337                # Always use self.problem as the query text to retrieve docs, but each time we replace the context with the338                # next similar docs in the retrieved doc results.339                if not doc_contents:340                    for _tmp_retrieve_count in range(1, 5):341                        self._reset(intermediate=True)342                        self.retrieve_docs(343                            self.problem, self.n_results * (2 * _tmp_retrieve_count + 1), self._search_string344                        )345                        doc_contents = self._get_context(self._results)346                        if doc_contents:347                            break348            elif update_context_case2:349                # Use the current intermediate info as the query text to retrieve docs, and each time we append the top similar350                # docs in the retrieved doc results to the context.351                for _tmp_retrieve_count in range(5):352                    self._reset(intermediate=True)353                    self.retrieve_docs(354                        _intermediate_info[0], self.n_results * (2 * _tmp_retrieve_count + 1), self._search_string355                    )356                    self._get_context(self._results)357                    doc_contents = "\n".join(self._doc_contents)  # + "\n" + "\n".join(self._intermediate_answers)358                    if doc_contents:359                        break360 361            self.clear_history()362            sender.clear_history()363            return True, self._generate_message(doc_contents, task=self._task)364        else:365            return False, None366 367    def retrieve_docs(self, problem: str, n_results: int = 20, search_string: str = ""):368        """Retrieve docs based on the given problem and assign the results to the class property `_results`.369        In case you want to customize the retrieval process, such as using a different vector db whose APIs are not370        compatible with chromadb or filter results with metadata, you can override this function. Just keep the current371        parameters and add your own parameters with default values, and keep the results in below type.372 373        Type of the results: Dict[str, List[List[Any]]], should have keys "ids" and "documents", "ids" for the ids of374        the retrieved docs and "documents" for the contents of the retrieved docs. Any other keys are optional. Refer375        to `chromadb.api.types.QueryResult` as an example.376            ids: List[string]377            documents: List[List[string]]378 379        Args:380            problem (str): the problem to be solved.381            n_results (int): the number of results to be retrieved. Default is 20.382            search_string (str): only docs that contain an exact match of this string will be retrieved. Default is "".383        """384        if not self._collection or not self._get_or_create:385            print("Trying to create collection.")386            self._client = create_vector_db_from_dir(387                dir_path=self._docs_path,388                max_tokens=self._chunk_token_size,389                client=self._client,390                collection_name=self._collection_name,391                chunk_mode=self._chunk_mode,392                must_break_at_empty_line=self._must_break_at_empty_line,393                embedding_model=self._embedding_model,394                get_or_create=self._get_or_create,395                embedding_function=self._embedding_function,396                custom_text_split_function=self.custom_text_split_function,397                custom_text_types=self._custom_text_types,398                recursive=self._recursive,399                extra_docs=self._extra_docs,400            )401            self._collection = True402            self._get_or_create = True403 404        results = query_vector_db(405            query_texts=[problem],406            n_results=n_results,407            search_string=search_string,408            client=self._client,409            collection_name=self._collection_name,410            embedding_model=self._embedding_model,411            embedding_function=self._embedding_function,412        )413        self._search_string = search_string414        self._results = results415        print("doc_ids: ", results["ids"])416 417    def generate_init_message(self, problem: str, n_results: int = 20, search_string: str = ""):418        """Generate an initial message with the given problem and prompt.419 420        Args:421            problem (str): the problem to be solved.422            n_results (int): the number of results to be retrieved.423            search_string (str): only docs containing this string will be retrieved.424 425        Returns:426            str: the generated prompt ready to be sent to the assistant agent.427        """428        self._reset()429        self.retrieve_docs(problem, n_results, search_string)430        self.problem = problem431        self.n_results = n_results432        doc_contents = self._get_context(self._results)433        message = self._generate_message(doc_contents, self._task)434        return message435 436    def run_code(self, code, **kwargs):437        lang = kwargs.get("lang", None)438        if code.startswith("!") or code.startswith("pip") or lang in ["bash", "shell", "sh"]:439            return (440                0,441                "You MUST NOT install any packages because all the packages needed are already installed.",442                None,443            )444        if self._ipython is None or lang != "python":445            return super().run_code(code, **kwargs)446        else:447            result = self._ipython.run_cell(code)448            log = str(result.result)449            exitcode = 0 if result.success else 1450            if result.error_before_exec is not None:451                log += f"\n{result.error_before_exec}"452                exitcode = 1453            if result.error_in_exec is not None:454                log += f"\n{result.error_in_exec}"455                exitcode = 1456            return exitcode, log, None457