diff --git a/fastembed/text/custom_text_embedding.py b/fastembed/text/custom_text_embedding.py index 118b7313..04715860 100644 --- a/fastembed/text/custom_text_embedding.py +++ b/fastembed/text/custom_text_embedding.py @@ -129,6 +129,8 @@ def add_model( normalization: bool, output_name: str | None = None, ) -> None: + if output_name is not None and (not isinstance(output_name, str) or not output_name): + raise ValueError("output_name must be a non-empty string or None") cls.SUPPORTED_MODELS.append(model_description) cls.POSTPROCESSING_MAPPING[model_description.model] = PostprocessingConfig( pooling=pooling, normalization=normalization, output_name=output_name diff --git a/tests/test_custom_models.py b/tests/test_custom_models.py index 38751bd6..05b1f7f4 100644 --- a/tests/test_custom_models.py +++ b/tests/test_custom_models.py @@ -138,6 +138,58 @@ def test_text_custom_model_output_name(): delete_model_cache(model.model._model_dir) +@pytest.mark.parametrize("output_name", ["", 123, [], False]) +def test_invalid_output_name_preserves_registry_and_allows_retry(output_name): + model_name = "custom/output-validation" + TextEmbedding.add_custom_model( + "custom/existing-output", + pooling=PoolingType.DISABLED, + normalization=False, + sources=ModelSource(hf="local/existing-output"), + dim=1, + ) + existing_models = list(CustomTextEmbedding.SUPPORTED_MODELS) + existing_mapping = dict(CustomTextEmbedding.POSTPROCESSING_MAPPING) + + with pytest.raises(ValueError, match="output_name"): + TextEmbedding.add_custom_model( + model_name, + pooling=PoolingType.DISABLED, + normalization=False, + sources=ModelSource(hf="local/output-validation"), + dim=1, + output_name=output_name, + ) + + assert CustomTextEmbedding.SUPPORTED_MODELS == existing_models + assert CustomTextEmbedding.POSTPROCESSING_MAPPING == existing_mapping + TextEmbedding.add_custom_model( + model_name, + pooling=PoolingType.DISABLED, + normalization=False, + sources=ModelSource(hf="local/output-validation"), + dim=1, + output_name="sentence_embedding", + ) + assert CustomTextEmbedding.POSTPROCESSING_MAPPING[model_name].output_name == ( + "sentence_embedding" + ) + + +@pytest.mark.parametrize("output_name", [None, "sentence_embedding", " output "]) +def test_valid_output_name_preserves_value(output_name): + model_name = "custom/valid-output" + TextEmbedding.add_custom_model( + model_name, + pooling=PoolingType.DISABLED, + normalization=False, + sources=ModelSource(hf="local/valid-output"), + dim=1, + output_name=output_name, + ) + assert CustomTextEmbedding.POSTPROCESSING_MAPPING[model_name].output_name == output_name + + def test_cross_encoder_custom_model(): is_ci = os.getenv("CI") custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"