Add support for Text and Document inputs in RetrievalQAComponent
This commit is contained in:
parent
155a679a6d
commit
05b088cdfe
1 changed files with 11 additions and 3 deletions
|
|
@ -2,8 +2,9 @@ from typing import Callable, Optional, Union
|
||||||
|
|
||||||
from langchain.chains.combine_documents.base import BaseCombineDocumentsChain
|
from langchain.chains.combine_documents.base import BaseCombineDocumentsChain
|
||||||
from langchain.chains.retrieval_qa.base import BaseRetrievalQA, RetrievalQA
|
from langchain.chains.retrieval_qa.base import BaseRetrievalQA, RetrievalQA
|
||||||
|
from langchain_core.documents import Document
|
||||||
from langflow import CustomComponent
|
from langflow import CustomComponent
|
||||||
from langflow.field_typing import BaseMemory, BaseRetriever
|
from langflow.field_typing import BaseMemory, BaseRetriever, Text
|
||||||
|
|
||||||
|
|
||||||
class RetrievalQAComponent(CustomComponent):
|
class RetrievalQAComponent(CustomComponent):
|
||||||
|
|
@ -18,18 +19,20 @@ class RetrievalQAComponent(CustomComponent):
|
||||||
"input_key": {"display_name": "Input Key", "advanced": True},
|
"input_key": {"display_name": "Input Key", "advanced": True},
|
||||||
"output_key": {"display_name": "Output Key", "advanced": True},
|
"output_key": {"display_name": "Output Key", "advanced": True},
|
||||||
"return_source_documents": {"display_name": "Return Source Documents"},
|
"return_source_documents": {"display_name": "Return Source Documents"},
|
||||||
|
"inputs": {"display_name": "Input", "input_types": ["Text", "Document"]},
|
||||||
}
|
}
|
||||||
|
|
||||||
def build(
|
def build(
|
||||||
self,
|
self,
|
||||||
combine_documents_chain: BaseCombineDocumentsChain,
|
combine_documents_chain: BaseCombineDocumentsChain,
|
||||||
retriever: BaseRetriever,
|
retriever: BaseRetriever,
|
||||||
|
inputs: str = "",
|
||||||
memory: Optional[BaseMemory] = None,
|
memory: Optional[BaseMemory] = None,
|
||||||
input_key: str = "query",
|
input_key: str = "query",
|
||||||
output_key: str = "result",
|
output_key: str = "result",
|
||||||
return_source_documents: bool = True,
|
return_source_documents: bool = True,
|
||||||
) -> Union[BaseRetrievalQA, Callable]:
|
) -> Union[BaseRetrievalQA, Callable, Text]:
|
||||||
return RetrievalQA(
|
runnable = RetrievalQA(
|
||||||
combine_documents_chain=combine_documents_chain,
|
combine_documents_chain=combine_documents_chain,
|
||||||
retriever=retriever,
|
retriever=retriever,
|
||||||
memory=memory,
|
memory=memory,
|
||||||
|
|
@ -37,3 +40,8 @@ class RetrievalQAComponent(CustomComponent):
|
||||||
output_key=output_key,
|
output_key=output_key,
|
||||||
return_source_documents=return_source_documents,
|
return_source_documents=return_source_documents,
|
||||||
)
|
)
|
||||||
|
if isinstance(inputs, Document):
|
||||||
|
inputs = inputs.page_content
|
||||||
|
|
||||||
|
result = runnable.invoke({input_key: inputs})
|
||||||
|
return result.content if hasattr(result, "content") else result
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue