Skip to content

Commit 0da14b9

Browse files
Add llama-index rag (#697)
1 parent 9519d71 commit 0da14b9

8 files changed

Lines changed: 477 additions & 41 deletions

File tree

ms_agent/agent/llm_agent.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
from ms_agent.callbacks import Callback, callbacks_mapping
1111
from ms_agent.llm.llm import LLM
1212
from ms_agent.llm.utils import Message, Tool
13-
from ms_agent.rag.base import Rag
13+
from ms_agent.rag.base import RAG
1414
from ms_agent.rag.utils import rag_mapping
1515
from ms_agent.tools import ToolManager
1616
from ms_agent.utils import async_retry
@@ -62,7 +62,7 @@ def __init__(self,
6262
self.tool_manager: Optional[ToolManager] = None
6363
self.memory_tools: List[Memory] = []
6464
self.planer: Optional[Planer] = None
65-
self.rag: Optional[Rag] = None
65+
self.rag: Optional[RAG] = None
6666
self.llm: Optional[LLM] = None
6767
self.runtime: Optional[Runtime] = None
6868
self.max_chat_round: int = 0
@@ -220,7 +220,7 @@ async def _prepare_messages(
220220
Message(role='user', content=inputs or query),
221221
]
222222
if self.rag is not None:
223-
messages = await self.rag.run(messages)
223+
messages = await self.rag.query(messages[1].content)
224224
return messages
225225

226226
async def _prepare_memory(self):
@@ -251,7 +251,7 @@ async def _prepare_rag(self):
251251
assert rag.name in rag_mapping, (
252252
f'{rag.name} not in rag_mapping, '
253253
f'which supports: {list(rag_mapping.keys())}')
254-
self.rag: Rag = rag_mapping(rag.name)(self.config)
254+
self.rag: RAG = rag_mapping(rag.name)(self.config)
255255

256256
async def _refine_memory(self, messages: List[Message]) -> List[Message]:
257257
"""

ms_agent/llm/openai_llm.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@ def __init__(
5050
base_url=base_url,
5151
)
5252
self.args: Dict = OmegaConf.to_container(
53-
getattr(config, 'generation_config', {}))
53+
getattr(config, 'generation_config', DictConfig({})))
5454

5555
def format_tools(self,
5656
tools: Optional[List[Tool]] = None

ms_agent/rag/base.py

Lines changed: 18 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -1,61 +1,52 @@
11
# Copyright (c) Alibaba, Inc. and its affiliates.
2-
from abc import abstractmethod
2+
from abc import ABC, abstractmethod
33
from typing import Any, List
44

5-
from ms_agent.llm import Message
65

7-
8-
class Rag:
6+
class RAG(ABC):
97
"""The base class for rags"""
108

119
def __init__(self, config):
1210
self.config = config
1311

1412
@abstractmethod
15-
async def add_document(self, url: str, content: str, **metadata) -> bool:
13+
async def add_documents(self, documents: List[str]) -> bool:
1614
"""Add document to Rag
1715
1816
Args:
19-
url(`str`): The url of the document
20-
content(`str`): The content of the document
21-
**metadata: Metadata information
17+
documents(`List[str]`): The content of the document
2218
2319
Returns:
2420
success or not
2521
"""
2622
pass
2723

2824
@abstractmethod
29-
async def search_documents(self,
30-
query: str,
31-
limit: int = 5,
32-
score_threshold: float = 0.7,
33-
**filters) -> List[Any]:
34-
"""Search documents in Rag
25+
async def query(self, query: str) -> str:
26+
"""Search documents
3527
3628
Args:
3729
query(`str`): The query to search for
38-
limit(`int`): The number of documents to return
39-
score_threshold(`float`): The score threshold
40-
**filters: Any extra filters
41-
4230
Returns:
43-
List of documents
31+
The query result
4432
"""
4533
pass
4634

4735
@abstractmethod
48-
async def delete_document(self, url: str) -> bool:
49-
"""Delete document from Rag
36+
async def retrieve(self,
37+
query: str,
38+
limit: int = 5,
39+
score_threshold: float = 0.7,
40+
**filters) -> List[Any]:
41+
"""Retrieve documents
5042
5143
Args:
52-
url(`str`): The url of the document
44+
query(`str`): The query to search for
45+
limit(`int`): The number of documents to return
46+
score_threshold(`float`): The score threshold
47+
**filters: Any extra filters
5348
5449
Returns:
55-
bool: True if the document was successfully deleted
50+
List of documents
5651
"""
5752
pass
58-
59-
@abstractmethod
60-
async def run(self, inputs: List[Message]) -> List[Message]:
61-
pass

0 commit comments

Comments
 (0)