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

1# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai> 

2# 

3# SPDX-License-Identifier: Apache-2.0 

4 

5from typing import Any 

6 

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 

12 

13 

14@component 

15class TextEmbeddingRetriever: 

16 """ 

17 A component that retrieves documents using a query with an embedding-based retriever. 

18 

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. 

22 

23 ### Usage example 

24 

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 

32 

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 ] 

41 

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) 

48 

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") 

54 

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 """ 

60 

61 def __init__(self, *, retriever: EmbeddingRetriever, text_embedder: TextEmbedder) -> None: 

62 """ 

63 Initialize TextEmbeddingRetriever. 

64 

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 

70 

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() 

78 

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() 

88 

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() 

96 

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() 

106 

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. 

113 

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() 

122 

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"] 

126 

127 # sort 

128 docs.sort(key=lambda x: x.score or 0.0, reverse=True) 

129 return {"documents": docs} 

130 

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. 

137 

138 Uses `run_async` on the text embedder and retriever if available, otherwise falls back to 

139 running `run` in a thread executor. 

140 

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() 

149 

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 ) 

154 

155 docs: list[Document] = result["documents"] 

156 docs.sort(key=lambda x: x.score or 0.0, reverse=True) 

157 return {"documents": docs} 

158 

159 def to_dict(self) -> dict[str, Any]: 

160 """ 

161 Serializes the component to a dictionary. 

162 

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 ) 

171 

172 @classmethod 

173 def from_dict(cls, data: dict[str, Any]) -> "TextEmbeddingRetriever": 

174 """ 

175 Deserializes the component from a dictionary. 

176 

177 :param data: The dictionary to deserialize from. 

178 :returns: 

179 The deserialized component. 

180 """ 

181 return default_from_dict(cls, data)