diff --git a/tests/test_models.py b/tests/test_models.py index 554c807a0..fe325cdf5 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -28,6 +28,22 @@ def test_performance_case_config_applies_top_level_payload_without_mutating_cust assert custom_case == {} +def test_custom_performance_case_config_applies_top_level_payload(): + case_config = CaseConfig( + case_id=CaseType.PerformanceCustomDataset, + custom_case={ + "name": "custom", + "description": "", + "load_timeout": 1, + "optimize_timeout": 1, + "dataset_config": {"size": 1, "dim": 1}, + }, + payload_profile=PayloadProfile.VECTOR, + ) + + assert case_config.case.payload_profile == PayloadProfile.VECTOR + + def test_performance_case_config_payload_round_trip_and_hash_identity(): ids_only = CaseConfig( case_id=CaseType.Performance768D100M, diff --git a/vectordb_bench/backend/cases.py b/vectordb_bench/backend/cases.py index f007aa06b..c92abcd97 100644 --- a/vectordb_bench/backend/cases.py +++ b/vectordb_bench/backend/cases.py @@ -448,6 +448,7 @@ def __init__( dataset=DatasetManager(data=dataset), use_filter=use_filter, label_percentage=label_percentage, + **kwargs, ) @property