Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_embedding_function_schemas.py457 linesDownload Raw Back to utils
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 
codekingpro/portable-devtools · Team Ai