aphilippov/python-server-api
0
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 