codekingpro/portable-devtools
114k
1import pytest
2from typing import List, Any, Callable, Dict
3from jsonschema import ValidationError
4from unittest.mock import MagicMock, create_autospec
5from chromadb.utils.embedding_functions.schemas import (
6 validate_config_schema,
7 load_schema,
8 get_available_schemas,
9)
10from chromadb.utils.embedding_functions import (
11 known_embedding_functions,
12 sparse_known_embedding_functions,
13)
14from chromadb.api.types import Documents, Embeddings
15from pytest import MonkeyPatch
16
17# Skip these embedding functions in tests
18SKIP_EMBEDDING_FUNCTIONS = [
19 "chroma_langchain",
20]
21
22
23def get_embedding_function_names() -> List[str]:
24 """Get all embedding function names to test"""
25 return [
26 name
27 for name in known_embedding_functions.keys()
28 if name not in SKIP_EMBEDDING_FUNCTIONS
29 ]
30
31
32class TestEmbeddingFunctionSchemas:
33 """Test class for embedding function schemas"""
34
35 @pytest.mark.parametrize("ef_name", get_embedding_function_names())
36 def test_embedding_function_config_roundtrip(
37 self,
38 ef_name: str,
39 mock_embeddings: Callable[[Documents], Embeddings],
40 mock_common_deps: MonkeyPatch,
41 ) -> None:
42 """Test embedding function configuration roundtrip"""
43 ef_class = known_embedding_functions[ef_name]
44
45 # Create an autospec of the embedding function class
46 mock_ef = create_autospec(ef_class, instance=True)
47
48 # Mock the __call__ method
49 mock_call = MagicMock(return_value=mock_embeddings(["test"]))
50 mock_ef.__call__ = mock_call
51
52 # For chroma-cloud-qwen, mock get_config to return valid data
53 if ef_name == "chroma-cloud-qwen":
54 from chromadb.utils.embedding_functions.chroma_cloud_qwen_embedding_function import (
55 ChromaCloudQwenEmbeddingModel,
56 CHROMA_CLOUD_QWEN_DEFAULT_INSTRUCTIONS,
57 )
58
59 mock_ef.get_config.return_value = {
60 "api_key_env_var": "CHROMA_API_KEY",
61 "model": ChromaCloudQwenEmbeddingModel.QWEN3_EMBEDDING_0p6B.value,
62 "task": "nl_to_code",
63 "instructions": CHROMA_CLOUD_QWEN_DEFAULT_INSTRUCTIONS,
64 }
65
66 # Mock the class constructor to return our mock instance
67 mock_common_deps.setattr(
68 ef_class, "__new__", lambda cls, *args, **kwargs: mock_ef
69 )
70
71 # Create instance with minimal args (constructor will be mocked)
72 ef_instance = ef_class()
73
74 # Get the config (this will use the real method)
75 config = ef_instance.get_config()
76
77 # Test recreation from config
78 new_instance = ef_class.build_from_config(config)
79 new_config = new_instance.get_config()
80
81 # Configs should match
82 assert (
83 config == new_config
84 ), f"Configs don't match after recreation for {ef_name}"
85
86 def test_schema_required_fields(self) -> None:
87 """Test that schemas enforce required fields"""
88 for schema_name in get_available_schemas():
89 schema = load_schema(schema_name)
90 if "required" not in schema:
91 continue
92
93 # Create minimal valid config
94 config = {}
95 for field in schema["required"]:
96 field_schema = schema["properties"][field]
97 field_type = (
98 field_schema["type"][0]
99 if isinstance(field_schema["type"], list)
100 else field_schema["type"]
101 )
102 config[field] = self._get_dummy_value(field_type)
103
104 # Test each required field
105 for field in schema["required"]:
106 test_config = config.copy()
107 del test_config[field]
108 with pytest.raises(ValidationError):
109 validate_config_schema(test_config, schema_name)
110
111 @staticmethod
112 def _get_dummy_value(field_type: str) -> Any:
113 """Get a dummy value for a given field type"""
114 type_map = {
115 "string": "dummy",
116 "integer": 0,
117 "number": 0.0,
118 "boolean": False,
119 "object": {},
120 "array": [],
121 }
122 return type_map.get(field_type, "dummy")
123
124 def test_schema_additional_properties(self) -> None:
125 """Test that schemas reject additional properties"""
126 for schema_name in get_available_schemas():
127 schema = load_schema(schema_name)
128 config = {}
129
130 # Add required fields
131 if "required" in schema:
132 for field in schema["required"]:
133 field_schema = schema["properties"][field]
134 field_type = (
135 field_schema["type"][0]
136 if isinstance(field_schema["type"], list)
137 else field_schema["type"]
138 )
139 config[field] = self._get_dummy_value(field_type)
140
141 # Add additional property
142 test_config = config.copy()
143 test_config["additional_property"] = "value"
144
145 # Test validation
146 if schema.get("additionalProperties", True) is False:
147 with pytest.raises(ValidationError):
148 validate_config_schema(test_config, schema_name)
149
150 def _create_valid_config_from_schema(
151 self, schema: Dict[str, Any]
152 ) -> Dict[str, Any]:
153 """Create a valid config from a schema by filling in required fields"""
154 config: Dict[str, Any] = {}
155
156 if "required" in schema and "properties" in schema:
157 for field in schema["required"]:
158 if field in schema["properties"]:
159 field_schema = schema["properties"][field]
160 config[field] = self._get_value_from_field_schema(field_schema)
161
162 return config
163
164 def _get_value_from_field_schema(self, field_schema: Dict[str, Any]) -> Any:
165 """Get a valid value from a field schema"""
166 # Handle enums - use first enum value
167 if "enum" in field_schema:
168 return field_schema["enum"][0]
169
170 # Handle type (could be a list or single value)
171 field_type = field_schema.get("type")
172 if field_type is None:
173 return "dummy" # Fallback if no type specified
174
175 if isinstance(field_type, list):
176 # If null is in the type list, prefer non-null type
177 non_null_types = [t for t in field_type if t != "null"]
178 field_type = non_null_types[0] if non_null_types else field_type[0]
179
180 if field_type == "object":
181 # Handle nested objects
182 nested_config = {}
183 if "properties" in field_schema:
184 nested_required = field_schema.get("required", [])
185 for prop in nested_required:
186 if prop in field_schema["properties"]:
187 nested_config[prop] = self._get_value_from_field_schema(
188 field_schema["properties"][prop]
189 )
190 return nested_config if nested_config else {}
191
192 if field_type == "array":
193 # Return empty array for arrays
194 return []
195
196 # Use the existing dummy value method for primitive types
197 return self._get_dummy_value(field_type)
198
199 def _has_custom_validation(self, ef_class: Any) -> bool:
200 """Check if validate_config actually validates (not just base implementation)"""
201 try:
202 # Try with an obviously invalid config - if it doesn't raise, it's base implementation
203 invalid_config = {"__invalid_test_config__": True}
204 try:
205 ef_class.validate_config(invalid_config)
206 # If we get here without exception, it's using base implementation
207 return False
208 except (ValidationError, ValueError, FileNotFoundError):
209 # If it raises any validation-related error, it's actually validating
210 return True
211 except Exception:
212 # Any other exception means it's trying to validate (e.g., schema not found)
213 return True
214
215 def _setup_env_vars_for_ef(
216 self, ef_name: str, mock_common_deps: MonkeyPatch
217 ) -> None:
218 """Set up environment variables needed for embedding function instantiation"""
219 # Map of embedding function names to their default API key environment variable names
220 api_key_env_vars = {
221 "cohere": "CHROMA_COHERE_API_KEY",
222 "openai": "CHROMA_OPENAI_API_KEY",
223 "huggingface": "CHROMA_HUGGINGFACE_API_KEY",
224 "huggingface_server": "CHROMA_HUGGINGFACE_API_KEY",
225 "google_palm": "CHROMA_GOOGLE_PALM_API_KEY",
226 "google_genai": "GEMINI_API_KEY",
227 "google_generative_ai": "GEMINI_API_KEY",
228 "google_vertex": "CHROMA_GOOGLE_VERTEX_API_KEY",
229 "jina": "CHROMA_JINA_API_KEY",
230 "mistral": "MISTRAL_API_KEY",
231 "morph": "MORPH_API_KEY",
232 "voyageai": "CHROMA_VOYAGE_API_KEY",
233 "cloudflare_workers_ai": "CHROMA_CLOUDFLARE_API_KEY",
234 "together_ai": "CHROMA_TOGETHER_AI_API_KEY",
235 "baseten": "CHROMA_BASETEN_API_KEY",
236 "roboflow": "CHROMA_ROBOFLOW_API_KEY",
237 "amazon_bedrock": "AWS_ACCESS_KEY_ID", # AWS uses different env vars
238 "chroma-cloud-qwen": "CHROMA_API_KEY",
239 # Sparse embedding functions
240 "chroma-cloud-splade": "CHROMA_API_KEY",
241 }
242
243 # Set API key environment variable if needed
244 if ef_name in api_key_env_vars:
245 mock_common_deps.setenv(api_key_env_vars[ef_name], "test-api-key")
246
247 # Special cases that need additional environment variables
248 if ef_name == "amazon_bedrock":
249 mock_common_deps.setenv("AWS_SECRET_ACCESS_KEY", "test-secret-key")
250 mock_common_deps.setenv("AWS_REGION", "us-east-1")
251
252 def _create_ef_instance(
253 self, ef_name: str, ef_class: Any, mock_common_deps: MonkeyPatch
254 ) -> Any:
255 """Create an embedding function instance, handling special cases"""
256 # Set up environment variables first
257 self._setup_env_vars_for_ef(ef_name, mock_common_deps)
258
259 # Mock missing modules that are imported inside __init__ methods
260 import sys
261
262 # Create mock modules
263 mock_pil = MagicMock()
264 mock_pil_image = MagicMock()
265 mock_google_genai = MagicMock()
266 mock_vertexai = MagicMock()
267 mock_vertexai_lm = MagicMock()
268 mock_boto3 = MagicMock()
269 mock_jina = MagicMock()
270 mock_mistralai = MagicMock()
271
272 # Mock boto3.Session for amazon_bedrock
273 mock_boto3_session = MagicMock()
274 mock_session_instance = MagicMock()
275 mock_session_instance.region_name = "us-east-1"
276 mock_session_instance.profile_name = None
277 mock_session_instance.client.return_value = MagicMock()
278 mock_boto3_session.return_value = mock_session_instance
279 mock_boto3.Session = mock_boto3_session
280
281 # Mock vertexai.init and TextEmbeddingModel
282 mock_text_embedding_model = MagicMock()
283 mock_text_embedding_model.from_pretrained.return_value = MagicMock()
284 mock_vertexai_lm.TextEmbeddingModel = mock_text_embedding_model
285 mock_vertexai.language_models = mock_vertexai_lm
286 mock_vertexai.init = MagicMock()
287
288 # Mock google.generativeai and google.genai - need to set up google module first
289 mock_google = MagicMock()
290 mock_google_genai.configure = MagicMock() # For palm.configure()
291 mock_google_genai.GenerativeModel = MagicMock(return_value=MagicMock())
292 mock_google.generativeai = mock_google_genai
293 mock_google_genai_new = MagicMock()
294 mock_google.genai = mock_google_genai_new
295
296 # Mock jina Client
297 mock_jina.Client = MagicMock()
298
299 # Mock mistralai
300 mock_mistral_client = MagicMock()
301 mock_mistral_client.return_value.embeddings.create.return_value.data = [
302 MagicMock(embedding=[0.1, 0.2, 0.3])
303 ]
304 mock_mistralai.Mistral = mock_mistral_client
305
306 # Add missing modules to sys.modules using monkeypatch
307 modules_to_mock = {
308 "PIL": mock_pil,
309 "PIL.Image": mock_pil_image,
310 "google": mock_google,
311 "google.generativeai": mock_google_genai,
312 "google.genai": mock_google_genai_new,
313 "google.genai.types": MagicMock(),
314 "vertexai": mock_vertexai,
315 "vertexai.language_models": mock_vertexai_lm,
316 "boto3": mock_boto3,
317 "jina": mock_jina,
318 "mistralai": mock_mistralai,
319 }
320
321 for module_name, mock_module in modules_to_mock.items():
322 mock_common_deps.setitem(sys.modules, module_name, mock_module)
323
324 # Special cases that need additional arguments
325 if ef_name == "cloudflare_workers_ai":
326 return ef_class(
327 model_name="test-model",
328 account_id="test-account-id",
329 )
330 elif ef_name == "baseten":
331 # Baseten needs api_key explicitly passed even with env var
332 return ef_class(
333 api_key="test-api-key",
334 api_base="https://test.api.baseten.co",
335 )
336 elif ef_name == "amazon_bedrock":
337 # Amazon Bedrock needs a boto3 session - create a mock session
338 # boto3 is already mocked in sys.modules above
339 mock_session = mock_boto3.Session(region_name="us-east-1")
340 return ef_class(
341 session=mock_session,
342 model_name="amazon.titan-embed-text-v1",
343 )
344 elif ef_name == "huggingface_server":
345 return ef_class(url="http://localhost:8080")
346 elif ef_name == "google_vertex":
347 return ef_class(project_id="test-project", region="us-central1")
348 elif ef_name == "mistral":
349 return ef_class(model="mistral-embed")
350 elif ef_name == "roboflow":
351 return ef_class() # No model_name needed
352 elif ef_name == "chroma-cloud-qwen":
353 from chromadb.utils.embedding_functions.chroma_cloud_qwen_embedding_function import (
354 ChromaCloudQwenEmbeddingModel,
355 )
356
357 return ef_class(
358 model=ChromaCloudQwenEmbeddingModel.QWEN3_EMBEDDING_0p6B,
359 task="nl_to_code",
360 )
361 else:
362 # Try with no args first
363 try:
364 return ef_class()
365 except Exception:
366 # If that fails, try with common minimal args
367 return ef_class(model_name="test-model")
368
369 @pytest.mark.parametrize("ef_name", get_embedding_function_names())
370 def test_validate_config_with_schema(
371 self,
372 ef_name: str,
373 mock_embeddings: Callable[[Documents], Embeddings],
374 mock_common_deps: MonkeyPatch,
375 ) -> None:
376 """Test that validate_config works correctly with actual configs from embedding functions"""
377 ef_class = known_embedding_functions[ef_name]
378
379 # Skip if the embedding function doesn't have a validate_config method
380 if not hasattr(ef_class, "validate_config"):
381 pytest.skip(f"{ef_name} does not have validate_config method")
382
383 # Check if it's callable (static methods are callable on the class)
384 if not callable(getattr(ef_class, "validate_config", None)):
385 pytest.skip(f"{ef_name} validate_config is not callable")
386
387 # Skip if using base implementation (doesn't actually validate)
388 if not self._has_custom_validation(ef_class):
389 pytest.skip(
390 f"{ef_name} uses base validate_config implementation (no validation)"
391 )
392
393 # Create a real instance to get the actual config
394 # We'll mock __call__ to avoid needing to actually generate embeddings
395 try:
396 ef_instance = self._create_ef_instance(ef_name, ef_class, mock_common_deps)
397 except Exception as e:
398 pytest.skip(
399 f"{ef_name} requires arguments that we cannot provide without external deps: {e}"
400 )
401
402 # Mock only __call__ to avoid needing to actually generate embeddings
403 mock_call = MagicMock(return_value=mock_embeddings(["test"]))
404 mock_common_deps.setattr(ef_instance, "__call__", mock_call)
405
406 # Get the actual config from the embedding function (this uses the real get_config method)
407 config = ef_instance.get_config()
408
409 # Filter out None values - optional fields with None shouldn't be included in validation
410 # This matches common JSON schema practice where optional fields are omitted rather than null
411 config = {k: v for k, v in config.items() if v is not None}
412
413 # Validate the actual config using the embedding function's validate_config method
414 ef_class.validate_config(config)
415
416 def test_validate_config_sparse_embedding_functions(
417 self,
418 mock_embeddings: Callable[[Documents], Embeddings],
419 mock_common_deps: MonkeyPatch,
420 ) -> None:
421 """Test validate_config for sparse embedding functions with actual configs"""
422 for ef_name, ef_class in sparse_known_embedding_functions.items():
423 # Skip if the embedding function doesn't have a validate_config method
424 if not hasattr(ef_class, "validate_config"):
425 continue
426
427 # Check if it's callable (static methods are callable on the class)
428 if not callable(getattr(ef_class, "validate_config", None)):
429 continue
430
431 # Skip if using base implementation (doesn't actually validate)
432 if not self._has_custom_validation(ef_class):
433 continue
434
435 # Create a real instance to get the actual config
436 # We'll mock __call__ to avoid needing to actually generate embeddings
437 try:
438 ef_instance = self._create_ef_instance(
439 ef_name, ef_class, mock_common_deps
440 )
441 except Exception:
442 continue # Skip if we can't create instance
443
444 # Mock only __call__ to avoid needing to actually generate embeddings
445 mock_call = MagicMock(return_value=mock_embeddings(["test"]))
446 mock_common_deps.setattr(ef_instance, "__call__", mock_call)
447
448 # Get the actual config from the embedding function (this uses the real get_config method)
449 config = ef_instance.get_config()
450
451 # Filter out None values - optional fields with None shouldn't be included in validation
452 # This matches common JSON schema practice where optional fields are omitted rather than null
453 config = {k: v for k, v in config.items() if v is not None}
454
455 # Validate the actual config using the embedding function's validate_config method
456 ef_class.validate_config(config)
457 