| 3 | from finagent.query import QUERY_TYPES |
| 4 | from typing import Dict, Any, List |
| 5 | class DiverseQuery(): |
| 6 | def __init__(self, |
| 7 | memory: MemoryInterface, |
| 8 | provider: EmbeddingProvider, |
| 9 | top_k: int = 5): |
| 10 | self.memory = memory |
| 11 | self.provider = provider |
| 12 | self.top_k = top_k |
| 13 | |
| 14 | def query(self, |
| 15 | params: Dict= None, |
| 16 | query_types: List[str] = ["plain", "short_term", "long_term"], |
| 17 | top_k: int = None): |
| 18 | |
| 19 | return self.diverse_query(params, query_types=query_types, top_k=top_k) |
| 20 | |
| 21 | def diverse_query(self, |
| 22 | params: Dict, |
| 23 | query_types: List[str] = ["plain", "short_term", "long_term"], |
| 24 | top_k: int = None): |
| 25 | |
| 26 | top_k = top_k if top_k is not None else self.top_k |
| 27 | |
| 28 | type = params["type"] |
| 29 | symbol = params["symbol"] |
| 30 | |
| 31 | res = {} |
| 32 | |
| 33 | for query_type in query_types: |
| 34 | |
| 35 | query_text = QUERY_TYPES[query_type](params) |
| 36 | embedding = self.provider.embed_query(query_text) |
| 37 | query_items, _ = self.memory.query_memory(type=type, |
| 38 | symbol=symbol, |
| 39 | data={"embedding": embedding}, |
| 40 | embedding_query="embedding", |
| 41 | top_k=top_k) |
| 42 | |
| 43 | pre_query_items = query_items |
| 44 | |
| 45 | if len(pre_query_items) == 0: |
| 46 | post_query_items = [] |
| 47 | else: |
| 48 | post_query_items = pre_query_items |
| 49 | |
| 50 | res[query_type] = { |
| 51 | "query_text": query_text, |
| 52 | "query_items": post_query_items |
| 53 | } |
| 54 | |
| 55 | return res |