| 4 | import os |
| 5 | |
| 6 | class Retriever: |
| 7 | def __init__(self, |
| 8 | openai_api_key: str = None, |
| 9 | model: str = "text-embedding-ada-002"): |
| 10 | if openai_api_key is None: |
| 11 | openai_api_key = os.environ.get("OPENAI_API_KEY") |
| 12 | self.embed = OpenAIEmbeddings(openai_api_key=openai_api_key, model=model) |
| 13 | self.documents = dict() |
| 14 | |
| 15 | def add_tool(self, tool_name: str, api_info: Dict) -> None: |
| 16 | if tool_name in self.documents: |
| 17 | return |
| 18 | document = api_info["name_for_model"] + ". " + api_info["description_for_model"] |
| 19 | document_embedding = self.embed.embed_documents([document]) |
| 20 | self.documents[tool_name] = { |
| 21 | "document": document, |
| 22 | "embedding": document_embedding[0] |
| 23 | } |
| 24 | |
| 25 | def query(self, query: str, topk: int = 3) -> List[str]: |
| 26 | query_embedding = self.embed.embed_query(query) |
| 27 | |
| 28 | queue = PriorityQueue() |
| 29 | for tool_name, tool_info in self.documents.items(): |
| 30 | tool_embedding = tool_info["embedding"] |
| 31 | tool_sim = self.similarity(query_embedding, tool_embedding) |
| 32 | queue.put([-tool_sim, tool_name]) |
| 33 | |
| 34 | result = [] |
| 35 | for i in range(min(topk, len(queue.queue))): |
| 36 | tool = queue.get() |
| 37 | result.append(tool[1]) |
| 38 | |
| 39 | return result |
| 40 | |
| 41 | def similarity(self, query: List[float], document: List[float]) -> float: |
| 42 | return sum([i * j for i, j in zip(query, document)]) |
| 43 | |