-
Notifications
You must be signed in to change notification settings - Fork 9.9k
Expand file tree
/
Copy pathtest_retrieval.py
More file actions
374 lines (289 loc) · 15.4 KB
/
Copy pathtest_retrieval.py
File metadata and controls
374 lines (289 loc) · 15.4 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
import contextlib
import json
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from langflow.base.knowledge_bases.knowledge_base_utils import get_knowledge_bases
from lfx.components.files_and_knowledge.retrieval import KnowledgeBaseComponent
from pydantic import SecretStr
from tests.base import ComponentTestBaseWithClient
class TestKnowledgeBaseComponent(ComponentTestBaseWithClient):
@pytest.fixture
def component_class(self):
"""Return the component class to test."""
return KnowledgeBaseComponent
@pytest.fixture(autouse=True)
def mock_knowledge_base_path(self, tmp_path):
"""Mock the knowledge base root path directly."""
with patch("langflow.components.knowledge_bases.retrieval._KNOWLEDGE_BASES_ROOT_PATH", tmp_path):
yield
@pytest.fixture
def default_kwargs(self, tmp_path, active_user):
"""Return default kwargs for component instantiation."""
# Create knowledge base directory structure
kb_name = "test_kb"
kb_path = tmp_path / active_user.username / kb_name
kb_path.mkdir(parents=True, exist_ok=True)
# Create embedding metadata file
metadata = {
"embedding_provider": "HuggingFace",
"embedding_model": "sentence-transformers/all-MiniLM-L6-v2",
"api_key": None,
"api_key_used": False,
"chunk_size": 1000,
"created_at": "2024-01-01T00:00:00Z",
}
(kb_path / "embedding_metadata.json").write_text(json.dumps(metadata))
return {
"knowledge_base": kb_name,
"kb_root_path": str(tmp_path),
"api_key": None,
"search_query": "",
"top_k": 5,
"include_embeddings": True,
"_user_id": active_user.id,
}
@pytest.fixture
def file_names_mapping(self):
"""Return file names mapping for version testing."""
# This is a new component, so it doesn't exist in older versions
return []
async def test_get_knowledge_bases(self, tmp_path, active_user):
"""Test getting list of knowledge bases."""
# Create additional test directories
(tmp_path / active_user.username / "kb1").mkdir(parents=True, exist_ok=True)
(tmp_path / active_user.username / "kb2").mkdir(parents=True, exist_ok=True)
(tmp_path / active_user.username / ".hidden").mkdir(parents=True, exist_ok=True) # Should be ignored
kb_list = await get_knowledge_bases(tmp_path, user_id=active_user.id)
assert "test_kb" in kb_list
assert "kb1" in kb_list
assert "kb2" in kb_list
assert ".hidden" not in kb_list
async def test_update_build_config(self, component_class, default_kwargs, tmp_path, active_user):
"""Test updating build configuration."""
component = component_class(**default_kwargs)
# Create additional KB directories
(tmp_path / active_user.username / "kb1").mkdir(parents=True, exist_ok=True)
(tmp_path / active_user.username / "kb2").mkdir(parents=True, exist_ok=True)
build_config = {"knowledge_base": {"value": "test_kb", "options": []}}
result = await component.update_build_config(build_config, None, "knowledge_base")
assert "test_kb" in result["knowledge_base"]["options"]
assert "kb1" in result["knowledge_base"]["options"]
assert "kb2" in result["knowledge_base"]["options"]
async def test_update_build_config_invalid_kb(self, component_class, default_kwargs):
"""Test updating build config when selected KB is not available."""
component = component_class(**default_kwargs)
build_config = {"knowledge_base": {"value": "nonexistent_kb", "options": ["test_kb"]}}
result = await component.update_build_config(build_config, None, "knowledge_base")
assert result["knowledge_base"]["value"] is None
def test_get_kb_metadata_success(self, component_class, default_kwargs, active_user):
"""Test successful metadata loading."""
component = component_class(**default_kwargs)
kb_path = Path(default_kwargs["kb_root_path"]) / active_user.username / default_kwargs["knowledge_base"]
with patch("langflow.components.knowledge_bases.retrieval.decrypt_api_key") as mock_decrypt:
mock_decrypt.return_value = "decrypted_key"
metadata = component._get_kb_metadata(kb_path)
assert metadata["embedding_provider"] == "HuggingFace"
assert metadata["embedding_model"] == "sentence-transformers/all-MiniLM-L6-v2"
assert "chunk_size" in metadata
def test_get_kb_metadata_no_file(self, component_class, default_kwargs, tmp_path, active_user):
"""Test metadata loading when file doesn't exist."""
component = component_class(**default_kwargs)
nonexistent_path = tmp_path / active_user.username / "nonexistent"
nonexistent_path.mkdir(parents=True, exist_ok=True)
metadata = component._get_kb_metadata(nonexistent_path)
assert metadata == {}
def test_get_kb_metadata_json_error(self, component_class, default_kwargs, tmp_path, active_user):
"""Test metadata loading with invalid JSON."""
component = component_class(**default_kwargs)
kb_path = tmp_path / active_user.username / "invalid_json_kb"
kb_path.mkdir(parents=True, exist_ok=True)
# Create invalid JSON file
(kb_path / "embedding_metadata.json").write_text("invalid json content")
metadata = component._get_kb_metadata(kb_path)
assert metadata == {}
def test_get_kb_metadata_decrypt_error(self, component_class, default_kwargs, tmp_path, active_user):
"""Test metadata loading with decryption error."""
component = component_class(**default_kwargs)
kb_path = tmp_path / active_user.username / "decrypt_error_kb"
kb_path.mkdir(parents=True, exist_ok=True)
# Create metadata with encrypted key
metadata = {
"embedding_provider": "OpenAI",
"embedding_model": "text-embedding-ada-002",
"api_key": "encrypted_key", # pragma:allowlist secret
"chunk_size": 1000,
}
(kb_path / "embedding_metadata.json").write_text(json.dumps(metadata))
with patch("langflow.components.knowledge_bases.retrieval.decrypt_api_key") as mock_decrypt:
mock_decrypt.side_effect = ValueError("Decryption failed")
result = component._get_kb_metadata(kb_path)
assert result["api_key"] is None
@patch("langchain_huggingface.HuggingFaceEmbeddings")
def test_build_embeddings_huggingface(self, mock_hf_embeddings, component_class, default_kwargs):
"""Test building HuggingFace embeddings."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "HuggingFace",
"embedding_model": "sentence-transformers/all-MiniLM-L6-v2",
"chunk_size": 1000,
}
mock_embeddings = MagicMock()
mock_hf_embeddings.return_value = mock_embeddings
result = component._build_embeddings(metadata)
mock_hf_embeddings.assert_called_once_with(model="sentence-transformers/all-MiniLM-L6-v2")
assert result == mock_embeddings
@patch("langchain_openai.OpenAIEmbeddings")
def test_build_embeddings_openai(self, mock_openai_embeddings, component_class, default_kwargs):
"""Test building OpenAI embeddings."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "OpenAI",
"embedding_model": "text-embedding-ada-002",
"api_key": "test-api-key", # pragma:allowlist secret
"chunk_size": 1000,
}
mock_embeddings = MagicMock()
mock_openai_embeddings.return_value = mock_embeddings
result = component._build_embeddings(metadata)
mock_openai_embeddings.assert_called_once_with(
model="text-embedding-ada-002",
api_key="test-api-key", # pragma:allowlist secret
chunk_size=1000, # pragma:allowlist secret
)
assert result == mock_embeddings
def test_build_embeddings_openai_no_key(self, component_class, default_kwargs):
"""Test building OpenAI embeddings without API key raises error."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "OpenAI",
"embedding_model": "text-embedding-ada-002",
"api_key": None,
"chunk_size": 1000,
}
with pytest.raises(ValueError, match="OpenAI API key is required"):
component._build_embeddings(metadata)
@patch("langchain_cohere.CohereEmbeddings")
def test_build_embeddings_cohere(self, mock_cohere_embeddings, component_class, default_kwargs):
"""Test building Cohere embeddings."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "Cohere",
"embedding_model": "embed-english-v3.0",
"api_key": "test-api-key", # pragma:allowlist secret
"chunk_size": 1000,
}
mock_embeddings = MagicMock()
mock_cohere_embeddings.return_value = mock_embeddings
result = component._build_embeddings(metadata)
mock_cohere_embeddings.assert_called_once_with(
model="embed-english-v3.0",
cohere_api_key="test-api-key", # pragma:allowlist secret
) # pragma:allowlist secret
assert result == mock_embeddings
def test_build_embeddings_cohere_no_key(self, component_class, default_kwargs):
"""Test building Cohere embeddings without API key raises error."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "Cohere",
"embedding_model": "embed-english-v3.0",
"api_key": None,
"chunk_size": 1000,
}
with pytest.raises(ValueError, match="Cohere API key is required"):
component._build_embeddings(metadata)
def test_build_embeddings_custom_not_supported(self, component_class, default_kwargs):
"""Test building custom embeddings raises NotImplementedError."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "Custom",
"embedding_model": "custom-model",
"api_key": "test-key", # pragma:allowlist secret
} # pragma:allowlist secret
with pytest.raises(NotImplementedError, match="Custom embedding models not yet supported"):
component._build_embeddings(metadata)
def test_build_embeddings_unsupported_provider(self, component_class, default_kwargs):
"""Test building embeddings with unsupported provider raises NotImplementedError."""
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "UnsupportedProvider",
"embedding_model": "some-model",
"api_key": "test-key", # pragma:allowlist secret
} # pragma:allowlist secret
with pytest.raises(NotImplementedError, match="Embedding provider 'UnsupportedProvider' is not supported"):
component._build_embeddings(metadata)
def test_build_embeddings_with_user_api_key(self, component_class, default_kwargs):
"""Test that user-provided API key overrides stored one."""
# Use a real SecretStr object instead of a mock
mock_secret = SecretStr("user-provided-key")
default_kwargs["api_key"] = mock_secret
component = component_class(**default_kwargs)
metadata = {
"embedding_provider": "OpenAI",
"embedding_model": "text-embedding-ada-002",
"api_key": "stored-key", # pragma:allowlist secret
"chunk_size": 1000,
}
with patch("langchain_openai.OpenAIEmbeddings") as mock_openai:
mock_embeddings = MagicMock()
mock_openai.return_value = mock_embeddings
component._build_embeddings(metadata)
# The user-provided key should override the stored key in metadata
mock_openai.assert_called_once_with(
model="text-embedding-ada-002",
api_key="user-provided-key", # pragma:allowlist secret
chunk_size=1000,
)
async def test_retrieve_data_no_metadata(self, component_class, default_kwargs, tmp_path, active_user):
"""Test retrieving data when metadata is missing."""
# Remove metadata file
kb_path = tmp_path / active_user.username / default_kwargs["knowledge_base"]
metadata_file = kb_path / "embedding_metadata.json"
if metadata_file.exists():
metadata_file.unlink()
component = component_class(**default_kwargs)
with pytest.raises(ValueError, match="Metadata not found for knowledge base"):
await component.retrieve_data()
def test_retrieve_data_path_construction(self, component_class, default_kwargs):
"""Test that retrieve_data constructs the correct paths."""
component = component_class(**default_kwargs)
# Test that the component correctly builds the KB path
assert component.kb_root_path == default_kwargs["kb_root_path"]
assert component.knowledge_base == default_kwargs["knowledge_base"]
# Test that paths are correctly expanded
expanded_path = Path(component.kb_root_path).expanduser()
assert expanded_path.exists() # tmp_path should exist
# Verify method exists with correct parameters
assert hasattr(component, "retrieve_data")
assert hasattr(component, "search_query")
assert hasattr(component, "top_k")
assert hasattr(component, "include_embeddings")
async def test_retrieve_data_method_exists(self, component_class, default_kwargs):
"""Test that retrieve_data method exists and can be called."""
component = component_class(**default_kwargs)
# Just verify the method exists and has the right signature
assert hasattr(component, "retrieve_data"), "Component should have retrieve_data method"
# Mock all external calls to avoid integration issues
with (
patch.object(component, "_get_kb_metadata") as mock_get_metadata,
patch.object(component, "_build_embeddings") as mock_build_embeddings,
patch("langchain_chroma.Chroma"),
):
mock_get_metadata.return_value = {"embedding_provider": "HuggingFace", "embedding_model": "test-model"}
mock_build_embeddings.return_value = MagicMock()
# This is a unit test focused on the component's internal logic
with contextlib.suppress(Exception):
await component.retrieve_data()
# Verify internal methods were called
mock_get_metadata.assert_called_once()
mock_build_embeddings.assert_called_once()
def test_include_embeddings_parameter(self, component_class, default_kwargs):
"""Test that include_embeddings parameter is properly set."""
# Test with embeddings enabled
default_kwargs["include_embeddings"] = True
component = component_class(**default_kwargs)
assert component.include_embeddings is True
# Test with embeddings disabled
default_kwargs["include_embeddings"] = False
component = component_class(**default_kwargs)
assert component.include_embeddings is False