Coverage for haystack/components/retrievers/text_embedding_retriever.py: 100%
52 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-21 13:53 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-21 13:53 +0000
1# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai>
2#
3# SPDX-License-Identifier: Apache-2.0
5from typing import Any
7from haystack import Document, component, default_from_dict, default_to_dict
8from haystack.components.embedders.types.protocol import TextEmbedder
9from haystack.components.retrievers.types import EmbeddingRetriever
10from haystack.core.serialization import component_to_dict
11from haystack.utils.async_utils import _execute_component_async
14@component
15class TextEmbeddingRetriever:
16 """
17 A component that retrieves documents using a query with an embedding-based retriever.
19 This component takes a text query, converts it to an embedding using a text embedder, and then uses an
20 embedding-based retriever to find relevant documents.
21 The results are sorted by relevance score.
23 ### Usage example
25 ```python
26 from haystack import Document
27 from haystack.document_stores.in_memory import InMemoryDocumentStore
28 from haystack.document_stores.types import DuplicatePolicy
29 from haystack.components.embedders import OpenAITextEmbedder, OpenAIDocumentEmbedder
30 from haystack.components.retrievers import InMemoryEmbeddingRetriever, TextEmbeddingRetriever
31 from haystack.components.writers import DocumentWriter
33 documents = [
34 Document(content="Renewable energy is energy that is collected from renewable resources."),
35 Document(content="Solar energy is a type of green energy that is harnessed from the sun."),
36 Document(content="Wind energy is another type of green energy that is generated by wind turbines."),
37 Document(content="Geothermal energy is heat that comes from the sub-surface of the earth."),
38 Document(content="Biomass energy is produced from organic materials, such as plant and animal waste."),
39 Document(content="Fossil fuels, such as coal, oil, and natural gas, are non-renewable energy sources."),
40 ]
42 # Populate the document store
43 doc_store = InMemoryDocumentStore()
44 doc_embedder = OpenAIDocumentEmbedder()
45 doc_writer = DocumentWriter(document_store=doc_store, policy=DuplicatePolicy.SKIP)
46 documents = doc_embedder.run(documents)["documents"]
47 doc_writer.run(documents=documents)
49 # Run the retriever
50 in_memory_retriever = InMemoryEmbeddingRetriever(document_store=doc_store, top_k=1)
51 text_embedder = OpenAITextEmbedder()
52 retriever = TextEmbeddingRetriever(retriever=in_memory_retriever, text_embedder=text_embedder)
53 result = retriever.run(query="Geothermal energy")
55 for doc in result["documents"]:
56 print(f"Content: {doc.content}, Score: {doc.score}")
57 # >> Content: Geothermal energy is heat that comes from the sub-surface of the earth., Score: 0.8509603046266574
58 ```
59 """
61 def __init__(self, *, retriever: EmbeddingRetriever, text_embedder: TextEmbedder) -> None:
62 """
63 Initialize TextEmbeddingRetriever.
65 :param retriever: The embedding-based retriever to use for document retrieval.
66 :param text_embedder: The text embedder to convert a text query to an embedding.
67 """
68 self.retriever = retriever
69 self.text_embedder = text_embedder
71 def warm_up(self) -> None:
72 """
73 Warm up the text embedder and the retriever.
74 """
75 for inner in (self.text_embedder, self.retriever):
76 if hasattr(inner, "warm_up"):
77 inner.warm_up()
79 async def warm_up_async(self) -> None:
80 """
81 Warm up the text embedder and the retriever on the serving event loop.
82 """
83 for inner in (self.text_embedder, self.retriever):
84 if hasattr(inner, "warm_up_async"):
85 await inner.warm_up_async()
86 elif hasattr(inner, "warm_up"):
87 inner.warm_up()
89 def close(self) -> None:
90 """
91 Release the text embedder's and the retriever's resources.
92 """
93 for inner in (self.text_embedder, self.retriever):
94 if hasattr(inner, "close"):
95 inner.close()
97 async def close_async(self) -> None:
98 """
99 Release the text embedder's and the retriever's async resources.
100 """
101 for inner in (self.text_embedder, self.retriever):
102 if hasattr(inner, "close_async"):
103 await inner.close_async()
104 elif hasattr(inner, "close"):
105 inner.close()
107 @component.output_types(documents=list[Document])
108 def run(
109 self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None
110 ) -> dict[str, list[Document]]:
111 """
112 Retrieve documents using a single query.
114 :param query: The query to retrieve documents for.
115 :param filters: A dictionary of filters to apply when retrieving documents.
116 :param top_k: The maximum number of documents to return.
117 :returns:
118 A dictionary containing:
119 - `documents`: List of retrieved documents sorted by relevance score.
120 """
121 self.warm_up()
123 embedding_result = self.text_embedder.run(text=query)
124 result = self.retriever.run(query_embedding=embedding_result["embedding"], filters=filters, top_k=top_k)
125 docs: list[Document] = result["documents"]
127 # sort
128 docs.sort(key=lambda x: x.score or 0.0, reverse=True)
129 return {"documents": docs}
131 @component.output_types(documents=list[Document])
132 async def run_async(
133 self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None
134 ) -> dict[str, list[Document]]:
135 """
136 Retrieve documents using a single query asynchronously.
138 Uses `run_async` on the text embedder and retriever if available, otherwise falls back to
139 running `run` in a thread executor.
141 :param query: The query to retrieve documents for.
142 :param filters: A dictionary of filters to apply when retrieving documents.
143 :param top_k: The maximum number of documents to return.
144 :returns:
145 A dictionary containing:
146 - `documents`: List of retrieved documents sorted by relevance score.
147 """
148 await self.warm_up_async()
150 embedding_result = await _execute_component_async(self.text_embedder, text=query)
151 result = await _execute_component_async(
152 self.retriever, query_embedding=embedding_result["embedding"], filters=filters, top_k=top_k
153 )
155 docs: list[Document] = result["documents"]
156 docs.sort(key=lambda x: x.score or 0.0, reverse=True)
157 return {"documents": docs}
159 def to_dict(self) -> dict[str, Any]:
160 """
161 Serializes the component to a dictionary.
163 :returns:
164 A dictionary representing the serialized component.
165 """
166 return default_to_dict(
167 self,
168 retriever=component_to_dict(obj=self.retriever, name="retriever"),
169 text_embedder=component_to_dict(obj=self.text_embedder, name="text_embedder"),
170 )
172 @classmethod
173 def from_dict(cls, data: dict[str, Any]) -> "TextEmbeddingRetriever":
174 """
175 Deserializes the component from a dictionary.
177 :param data: The dictionary to deserialize from.
178 :returns:
179 The deserialized component.
180 """
181 return default_from_dict(cls, data)