feat: add vector retrieval and update policy/template (#5218)
* Updated collection_indexing_policy to store the correct json. Added support for graph retrival and other minor imporvements * Added RagGraph template * [autofix.ci] apply automated fixes * Corrected the class name to avoid ut failures * [autofix.ci] apply automated fixes * Updated _map_search_type to be less idiotic * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * Renamed AstraDBGraphVectorStoreComponent back to its original form for convention sake * Unrelated to the graph work * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * Linting * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * [autofix.ci] apply automated fixes * Delete src/backend/base/langflow/initial_setup/starter_projects/RagGraph.json Remove template as per langflow team --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Eric Hare <ericrhare@gmail.com> Co-authored-by: Gabriel Luiz Freitas Almeida <gabriel@langflow.org>
This commit is contained in:
parent
6b21682b33
commit
b98b225037
1 changed files with 42 additions and 14 deletions
|
|
@ -6,11 +6,12 @@ from loguru import logger
|
||||||
|
|
||||||
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
|
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
|
||||||
from langflow.helpers import docs_to_data
|
from langflow.helpers import docs_to_data
|
||||||
from langflow.inputs import DictInput, FloatInput
|
from langflow.inputs import (
|
||||||
from langflow.io import (
|
|
||||||
BoolInput,
|
BoolInput,
|
||||||
DataInput,
|
DataInput,
|
||||||
|
DictInput,
|
||||||
DropdownInput,
|
DropdownInput,
|
||||||
|
FloatInput,
|
||||||
HandleInput,
|
HandleInput,
|
||||||
IntInput,
|
IntInput,
|
||||||
MultilineInput,
|
MultilineInput,
|
||||||
|
|
@ -71,11 +72,10 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
advanced=True,
|
advanced=True,
|
||||||
),
|
),
|
||||||
HandleInput(
|
HandleInput(
|
||||||
name="embedding",
|
name="embedding_model",
|
||||||
display_name="Embedding Model",
|
display_name="Embedding Model",
|
||||||
input_types=["Embeddings"],
|
input_types=["Embeddings"],
|
||||||
info="Embedding model.",
|
info="Allows an embedding model configuration.",
|
||||||
required=True,
|
|
||||||
),
|
),
|
||||||
DropdownInput(
|
DropdownInput(
|
||||||
name="metric",
|
name="metric",
|
||||||
|
|
@ -156,8 +156,14 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
name="search_type",
|
name="search_type",
|
||||||
display_name="Search Type",
|
display_name="Search Type",
|
||||||
info="Search type to use",
|
info="Search type to use",
|
||||||
options=["Similarity", "Similarity with score threshold", "MMR (Max Marginal Relevance)"],
|
options=[
|
||||||
value="Similarity",
|
"Similarity",
|
||||||
|
"Similarity with score threshold",
|
||||||
|
"MMR (Max Marginal Relevance)",
|
||||||
|
"Graph Traversal",
|
||||||
|
"MMR (Max Marginal Relevance) Graph Traversal",
|
||||||
|
],
|
||||||
|
value="MMR (Max Marginal Relevance) Graph Traversal",
|
||||||
advanced=True,
|
advanced=True,
|
||||||
),
|
),
|
||||||
FloatInput(
|
FloatInput(
|
||||||
|
|
@ -199,8 +205,10 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
raise ValueError(msg) from e
|
raise ValueError(msg) from e
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
logger.debug(f"Initializing Graph Vector Store {self.collection_name}")
|
||||||
|
|
||||||
vector_store = AstraDBGraphVectorStore(
|
vector_store = AstraDBGraphVectorStore(
|
||||||
embedding=self.embedding,
|
embedding=self.embedding_model,
|
||||||
collection_name=self.collection_name,
|
collection_name=self.collection_name,
|
||||||
metadata_incoming_links_key=self.metadata_incoming_links_key or "incoming_links",
|
metadata_incoming_links_key=self.metadata_incoming_links_key or "incoming_links",
|
||||||
token=self.token,
|
token=self.token,
|
||||||
|
|
@ -216,7 +224,7 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
pre_delete_collection=self.pre_delete_collection,
|
pre_delete_collection=self.pre_delete_collection,
|
||||||
metadata_indexing_include=[s for s in self.metadata_indexing_include if s] or None,
|
metadata_indexing_include=[s for s in self.metadata_indexing_include if s] or None,
|
||||||
metadata_indexing_exclude=[s for s in self.metadata_indexing_exclude if s] or None,
|
metadata_indexing_exclude=[s for s in self.metadata_indexing_exclude if s] or None,
|
||||||
collection_indexing_policy=orjson.dumps(self.collection_indexing_policy)
|
collection_indexing_policy=orjson.loads(self.collection_indexing_policy.encode("utf-8"))
|
||||||
if self.collection_indexing_policy
|
if self.collection_indexing_policy
|
||||||
else None,
|
else None,
|
||||||
)
|
)
|
||||||
|
|
@ -224,6 +232,7 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
msg = f"Error initializing AstraDBGraphVectorStore: {e}"
|
msg = f"Error initializing AstraDBGraphVectorStore: {e}"
|
||||||
raise ValueError(msg) from e
|
raise ValueError(msg) from e
|
||||||
|
|
||||||
|
logger.debug(f"Vector Store initialized: {vector_store.astra_env.collection_name}")
|
||||||
self._add_documents_to_vector_store(vector_store)
|
self._add_documents_to_vector_store(vector_store)
|
||||||
|
|
||||||
return vector_store
|
return vector_store
|
||||||
|
|
@ -248,10 +257,18 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
logger.debug("No documents to add to the Vector Store.")
|
logger.debug("No documents to add to the Vector Store.")
|
||||||
|
|
||||||
def _map_search_type(self) -> str:
|
def _map_search_type(self) -> str:
|
||||||
if self.search_type == "Similarity with score threshold":
|
match self.search_type:
|
||||||
|
case "Similarity":
|
||||||
|
return "similarity"
|
||||||
|
case "Similarity with score threshold":
|
||||||
return "similarity_score_threshold"
|
return "similarity_score_threshold"
|
||||||
if self.search_type == "MMR (Max Marginal Relevance)":
|
case "MMR (Max Marginal Relevance)":
|
||||||
return "mmr"
|
return "mmr"
|
||||||
|
case "Graph Traversal":
|
||||||
|
return "traversal"
|
||||||
|
case "MMR (Max Marginal Relevance) Graph Traversal":
|
||||||
|
return "mmr_traversal"
|
||||||
|
case _:
|
||||||
return "similarity"
|
return "similarity"
|
||||||
|
|
||||||
def _build_search_args(self):
|
def _build_search_args(self):
|
||||||
|
|
@ -270,6 +287,7 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
if not vector_store:
|
if not vector_store:
|
||||||
vector_store = self.build_vector_store()
|
vector_store = self.build_vector_store()
|
||||||
|
|
||||||
|
logger.debug("Searching for documents in AstraDBGraphVectorStore.")
|
||||||
logger.debug(f"Search input: {self.search_input}")
|
logger.debug(f"Search input: {self.search_input}")
|
||||||
logger.debug(f"Search type: {self.search_type}")
|
logger.debug(f"Search type: {self.search_type}")
|
||||||
logger.debug(f"Number of results: {self.number_of_results}")
|
logger.debug(f"Number of results: {self.number_of_results}")
|
||||||
|
|
@ -280,6 +298,14 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
search_args = self._build_search_args()
|
search_args = self._build_search_args()
|
||||||
|
|
||||||
docs = vector_store.search(query=self.search_input, search_type=search_type, **search_args)
|
docs = vector_store.search(query=self.search_input, search_type=search_type, **search_args)
|
||||||
|
|
||||||
|
# Drop links from the metadata. At this point the links don't add any value for building the
|
||||||
|
# context and haven't been restored to json which causes the conversion to fail.
|
||||||
|
logger.debug("Removing links from metadata.")
|
||||||
|
for doc in docs:
|
||||||
|
if "links" in doc.metadata:
|
||||||
|
doc.metadata.pop("links")
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
msg = f"Error performing search in AstraDBGraphVectorStore: {e}"
|
msg = f"Error performing search in AstraDBGraphVectorStore: {e}"
|
||||||
raise ValueError(msg) from e
|
raise ValueError(msg) from e
|
||||||
|
|
@ -287,7 +313,9 @@ class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
|
||||||
logger.debug(f"Retrieved documents: {len(docs)}")
|
logger.debug(f"Retrieved documents: {len(docs)}")
|
||||||
|
|
||||||
data = docs_to_data(docs)
|
data = docs_to_data(docs)
|
||||||
|
|
||||||
logger.debug(f"Converted documents to data: {len(data)}")
|
logger.debug(f"Converted documents to data: {len(data)}")
|
||||||
|
|
||||||
self.status = data
|
self.status = data
|
||||||
return data
|
return data
|
||||||
logger.debug("No search input provided. Skipping search.")
|
logger.debug("No search input provided. Skipping search.")
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue