From e2a42816efebbf7ce5e5298021b56f2722a037f9 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 3 Sep 2026 09:29:45 +0100 Subject: [PATCH 01/10] def 2957 --- sdv/metadata/_single_table.py | 156 +++++++++++++++++++++++++++------- 1 file changed, 125 insertions(+), 31 deletions(-) diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 2d15bf792..1df99b6aa 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -51,6 +51,7 @@ 'is stored as an int but the Regex allows it to start with "0". Please remove the Regex ' 'or update it to correspond to valid ints.' ) +MAX_RANGE_VALUES = 500 class _SingleTableMetadata: @@ -573,6 +574,36 @@ def _detect_id_column(self, column_name): return None + def _detect_ordinal_sdtype(self, data): + """Detect whether a numerical column should have the ordinal sdtype. + + A numerical column is considered ordinal when it contains whole numbers + and has low cardinality. + + Args: + data (pandas.Series): + The data to be analyzed. + + Returns: + str or None: + 'ordinal' if the column is ordinal, otherwise ``None``. + """ + if len(data) <= self._MIN_ROWS_FOR_PREDICTION: + return None + + clean_data = data.dropna() + if clean_data.empty: + return None + + whole_values = (clean_data == clean_data.round()).all() + unique_values = clean_data.nunique() + categorical_threshold = min(round(len(data) / 10), 10) + low_cardinality = unique_values <= categorical_threshold + if whole_values and low_cardinality: + return 'ordinal' + + return None + def _determine_sdtype_for_numbers(self, data, valid_potential_primary_key): """Determine the sdtype for a numerical column. @@ -582,22 +613,16 @@ def _determine_sdtype_for_numbers(self, data, valid_potential_primary_key): valid_potential_primary_key(bool): If the column is unique and doesn't have NaNs. """ - sdtype = 'numerical' + sdtype = self._detect_ordinal_sdtype(data) or 'numerical' pk_candidate = False + if len(data) > self._MIN_ROWS_FOR_PREDICTION: - is_not_null = ~data.isna() - clean_data = (data == data.round()).loc[is_not_null] + clean_data = data.dropna() if clean_data.empty: return sdtype, pk_candidate - whole_values = clean_data.all() - positive_values = (data >= 0).loc[is_not_null].all() - - unique_values = data.nunique() - unique_lt_categorical_threshold = unique_values <= min(round(len(data) / 10), 10) - - if whole_values and positive_values and unique_lt_categorical_threshold: - sdtype = 'categorical' + whole_values = (clean_data == clean_data.round()).all() + positive_values = (clean_data >= 0).all() pk_candidate = valid_potential_primary_key and whole_values and positive_values @@ -733,15 +758,7 @@ def _select_primary_key( A list of primary key candidates that are pii. table_name (str): The name of the table to be analyzed. Defaults to ``None``. - verbose (bool): - A boolean that determines if information should be printed regarding detection. - If True, it prints out information about what is detected. - If False, it does not print out any information about what is detected. - Defaults to False. """ - if verbose: - table_str = f" for table '{table_name}'" if table_name else '' - sys.stdout.write(f'\nDetecting primary key{table_str}:\n') chosen_pk = None sdtype_updated = False pii_removed = False @@ -767,7 +784,83 @@ def _select_primary_key( del self.columns[self.primary_key]['pii'] pii_removed = True - if verbose: + return chosen_pk, sdtype_updated, pii_removed + + def _detect_range_values(self, data): + """Detect the range values for a column. + + This method detects the unique values in a column if there are fewer than + `MAX_RANGE_VALUES` unique values. + + Args: + data (pandas.Series): + The data to be analyzed. + """ + range_values = data.dropna().unique() + if len(range_values) < MAX_RANGE_VALUES: + return range_values.tolist() + + return None + + def _detect_ranges(self, data): + """Detect the range information for all columns. + + Args: + data (pandas.DataFrame): + The data to be analyzed. + """ + for column_name, column_metadata in self.columns.items(): + if column_name == self.primary_key: + continue + + column_data = data[column_name] + sdtype = column_metadata['sdtype'] + + column_metadata['range_is_nullable'] = bool(column_data.isna().any()) + if sdtype == 'numerical': + clean_data = column_data.dropna() + if not clean_data.empty: + column_metadata['range_min'] = clean_data.min().item() + column_metadata['range_max'] = clean_data.max().item() + + column_metadata['decimal_places'] = learn_rounding_digits(column_data) + + elif sdtype == 'datetime': + clean_data = column_data.dropna() + if not clean_data.empty: + datetime_format = column_metadata.get('datetime_format') + clean_data = pd.to_datetime(clean_data, format=datetime_format) + + range_min = clean_data.min() + range_max = clean_data.max() + if datetime_format: + range_min = range_min.strftime(datetime_format) + range_max = range_max.strftime(datetime_format) + else: + range_min = str(range_min) + range_max = str(range_max) + + column_metadata['range_min'] = range_min + column_metadata['range_max'] = range_max + + elif sdtype in {'categorical', 'ordinal'}: + range_values = self._detect_range_values(column_data) + if range_values is not None: + column_metadata['range_values'] = range_values + + def _print_detection( + self, table_name, data, infer_sdtypes, infer_keys, chosen_pk, sdtype_updated, pii_removed + ): + if infer_sdtypes: + table_str = f"table '{table_name}'" if table_name else 'table' + sys.stdout.write(f'\nDetecting {table_str}:\n') + for field in data: + column_metadata = _format_column_metadata(self.columns[field]) + sys.stdout.write(f"- Column '{field}': {column_metadata}\n") + + if infer_keys == 'primary_only': + table_str = f" for table '{table_name}'" if table_name else '' + sys.stdout.write(f'\nDetecting primary key{table_str}:\n') _print_primary_key_detection(chosen_pk, sdtype_updated, pii_removed) def _detect_columns( @@ -796,10 +889,6 @@ def _detect_columns( If False, it does not print out any information about what is detected. Defaults to False. """ - if verbose and infer_sdtypes: - table_str = f"table '{table_name}'" if table_name else 'table' - sys.stdout.write(f'\nDetecting {table_str}:\n') - old_columns = data.columns data.columns = data.columns.astype(str) pk_candidates = [] @@ -823,24 +912,29 @@ def _detect_columns( if sdtype == 'datetime' and dtype == 'O': datetime_format = _get_datetime_format(column_data.iloc[:100]) column_dict['datetime_format'] = datetime_format + else: sdtype = 'unknown' column_dict['pii'] = True column_dict['sdtype'] = sdtype - - if verbose and infer_sdtypes: - column_metadata = _format_column_metadata(column_dict) - sys.stdout.write(f"- Column '{field}': {column_metadata}\n") - self.columns[field] = deepcopy(column_dict) + + chosen_pk = None + sdtype_updated = False + pii_removed = False if infer_keys == 'primary_only': - self._select_primary_key( + chosen_pk, sdtype_updated, pii_removed = self._select_primary_key( infer_sdtypes=infer_sdtypes, pk_candidates=pk_candidates, pii_pk_candidates=pii_pk_candidates, table_name=table_name, - verbose=verbose, + ) + + self._detect_ranges(data) + if verbose: + self._print_detection( + table_name, data, infer_sdtypes, infer_keys, chosen_pk, sdtype_updated, pii_removed ) self._updated = True From f7ec8deeef273277f26acef1f5ef6243bdccc8ae Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 3 Sep 2026 09:29:58 +0100 Subject: [PATCH 02/10] unit tests --- tests/unit/io/local/test_local.py | 38 +- tests/unit/metadata/test_metadata.py | 56 +- tests/unit/metadata/test_single_table.py | 643 ++++++++++++++++++----- 3 files changed, 574 insertions(+), 163 deletions(-) diff --git a/tests/unit/io/local/test_local.py b/tests/unit/io/local/test_local.py index b5fbd7ccd..11ebb0325 100644 --- a/tests/unit/io/local/test_local.py +++ b/tests/unit/io/local/test_local.py @@ -38,31 +38,37 @@ def test_create_metadata(self): # Assert assert isinstance(metadata, Metadata) assert metadata.to_dict() == { - 'METADATA_SPEC_VERSION': 'V2', - 'relationships': [ - { - 'child_foreign_key': 'hotel_id', - 'child_table_name': 'guests', - 'parent_primary_key': 'hotel_id', - 'parent_table_name': 'hotel', - }, - ], 'tables': { - 'guests': { + 'hotel': { 'columns': { - 'guest_id': {'sdtype': 'id'}, 'hotel_id': {'sdtype': 'id'}, + 'stars': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 3, + 'range_max': 5, + 'decimal_places': 0, + }, }, - 'primary_key': 'guest_id', + 'primary_key': 'hotel_id', }, - 'hotel': { + 'guests': { 'columns': { - 'hotel_id': {'sdtype': 'id'}, - 'stars': {'sdtype': 'numerical'}, + 'guest_id': {'sdtype': 'id'}, + 'hotel_id': {'sdtype': 'id', 'range_is_nullable': False}, }, - 'primary_key': 'hotel_id', + 'primary_key': 'guest_id', }, }, + 'relationships': [ + { + 'parent_table_name': 'hotel', + 'child_table_name': 'guests', + 'parent_primary_key': 'hotel_id', + 'child_foreign_key': 'hotel_id', + } + ], + 'METADATA_SPEC_VERSION': 'V2', } diff --git a/tests/unit/metadata/test_metadata.py b/tests/unit/metadata/test_metadata.py index 3012b8456..95d1ca02a 100644 --- a/tests/unit/metadata/test_metadata.py +++ b/tests/unit/metadata/test_metadata.py @@ -807,16 +807,36 @@ def test_add_relationship_child_key_is_primary_key(self): 'primary_key': 'pk', 'columns': { 'pk': {'sdtype': 'id'}, - 'col1': {'sdtype': 'numerical'}, - 'col2': {'sdtype': 'categorical'}, + 'col1': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0.1, + 'range_max': 0.2, + 'decimal_places': 1, + }, + 'col2': { + 'sdtype': 'categorical', + 'range_is_nullable': False, + 'range_values': ['a', 'b', 'c'], + }, }, }, 'table2': { 'primary_key': 'pk', 'columns': { 'pk': {'sdtype': 'id'}, - 'col1': {'sdtype': 'numerical'}, - 'col2': {'sdtype': 'categorical'}, + 'col1': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0.1, + 'range_max': 0.2, + 'decimal_places': 1, + }, + 'col2': { + 'sdtype': 'categorical', + 'range_is_nullable': False, + 'range_values': ['a', 'b', 'c'], + }, }, }, }, @@ -1428,16 +1448,36 @@ def test_validate_child_key_is_primary_key(self): 'table': { 'columns': { 'pk': {'sdtype': 'id'}, - 'col1': {'sdtype': 'numerical'}, - 'col2': {'sdtype': 'categorical'}, + 'col1': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0.1, + 'range_max': 0.2, + 'decimal_places': 1, + }, + 'col2': { + 'sdtype': 'categorical', + 'range_is_nullable': False, + 'range_values': ['a', 'b', 'c'], + }, }, 'primary_key': 'pk', }, 'table2': { 'columns': { 'pk': {'sdtype': 'id'}, - 'col1': {'sdtype': 'numerical'}, - 'col2': {'sdtype': 'categorical'}, + 'col1': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0.1, + 'range_max': 0.2, + 'decimal_places': 1, + }, + 'col2': { + 'sdtype': 'categorical', + 'range_is_nullable': False, + 'range_values': ['a', 'b', 'c'], + }, }, 'primary_key': 'pk', }, diff --git a/tests/unit/metadata/test_single_table.py b/tests/unit/metadata/test_single_table.py index cafd7f320..30d106784 100644 --- a/tests/unit/metadata/test_single_table.py +++ b/tests/unit/metadata/test_single_table.py @@ -1067,6 +1067,36 @@ def test__detect_pii_columns(self): assert metadata._detect_pii_column('StateDepartment') == 'administrative_unit' assert metadata._detect_pii_column('STATEDEPARTMENT') is None + @pytest.mark.parametrize( + ('data', 'expected'), + [ + # Not enough rows + (pd.Series([1, 2, 3]), None), + # Enough rows but all null + (pd.Series([None] * 20), None), + # Low-cardinality integers + (pd.Series([1, 2] * 10), 'ordinal'), + # Low-cardinality whole-number floats + (pd.Series([1.0, 2.0] * 10), 'ordinal'), + # Low-cardinality numerical values that are not whole numbers + (pd.Series([1.1, 2.2] * 10), None), + # Whole numbers with too many unique values + (pd.Series(range(20)), None), + # Null values should be ignored when checking whole values + (pd.Series([1.0, 2.0, None, 1.0, 2.0] * 4), 'ordinal'), + ], + ) + def test__detect_ordinal_sdtype(self, data, expected): + """Test the ``_detect_ordinal_sdtype`` method.""" + # Setup + metadata = _SingleTableMetadata() + + # Run + result = metadata._detect_ordinal_sdtype(data) + + # Assert + assert result == expected + def test__determine_sdtype_for_numbers(self): """Test the ``determine_sdtype_for_numbers`` method. @@ -1074,7 +1104,7 @@ def test__determine_sdtype_for_numbers(self): - Instance of ``_SingleTableMetadata``. - A series of numbers with less than 5 rows. Should be detected as numerical sdtype - A series of numbers with less than 10% unique values. Should be detected as - categorical sdtype + ordinal sdtype - A series of numbers with all unique values. Should be detected as numerical sdtypes - A series of integers. Should be detected as numerical sdtype - A series of floats. Should be detected as numerical sdtype @@ -1138,12 +1168,12 @@ def test__determine_sdtype_for_numbers(self): # Assert assert sdtype_less_than_5_rows == 'numerical' assert candidate is False - assert sdtype_less_than_10_percent_unique_values == 'categorical' + assert sdtype_less_than_10_percent_unique_values == 'ordinal' assert sdtype_all_unique == 'numerical' assert sdtype_numerical_int == 'numerical' assert sdtype_numerical_float == 'numerical' assert sdtype_large_numerical_series == 'numerical' - assert sdtype_large_categorical_series == 'categorical' + assert sdtype_large_categorical_series == 'ordinal' def test__determine_sdtype_for_objects(self): """Test the ``_determine_sdtype_for_objects`` method.""" @@ -1216,10 +1246,121 @@ def test__determine_sdtype_for_objects_with_none(self): assert sdtype == 'categorical' assert candidate is False + @pytest.mark.parametrize( + ('data', 'expected'), + [ + (pd.Series(['a', 'b', 'a', None]), ['a', 'b']), + (pd.Series(range(499)), list(range(499))), + (pd.Series(range(500)), None), + ], + ) + def test__detect_range_values(self, data, expected): + """Test the ``_detect_range_values`` method.""" + # Setup + instance = _SingleTableMetadata() + + # Run + result = instance._detect_range_values(data) + + # Assert + assert result == expected + + @patch('sdv.metadata._single_table.learn_rounding_digits') + def test__detect_ranges(self, mock_learn_rounding_digits): + """Test the ``_detect_ranges`` method.""" + # Setup + instance = _SingleTableMetadata() + instance.columns = { + 'numerical': {'sdtype': 'numerical'}, + 'datetime': { + 'sdtype': 'datetime', + 'datetime_format': '%Y-%m-%d', + }, + 'categorical': {'sdtype': 'categorical'}, + 'ordinal': {'sdtype': 'ordinal'}, + 'boolean': {'sdtype': 'boolean'}, + 'id': {'sdtype': 'id'}, + 'unknown': {'sdtype': 'unknown'}, + } + data = pd.DataFrame({ + 'numerical': [1.1, 2.2, np.nan], + 'datetime': ['2024-01-01', None, '2024-01-03'], + 'categorical': ['a', 'b', None], + 'ordinal': [1, 2, None], + 'boolean': [True, False, None], + 'id': ['id_1', 'id_2', 'id_3'], + 'unknown': ['a', None, 'c'], + }) + mock_learn_rounding_digits.return_value = 1 + + # Run + instance._detect_ranges(data) + + # Assert + assert instance.columns['numerical'] == { + 'sdtype': 'numerical', + 'range_min': 1.1, + 'range_max': 2.2, + 'range_is_nullable': True, + 'decimal_places': 1, + } + assert instance.columns['datetime'] == { + 'sdtype': 'datetime', + 'datetime_format': '%Y-%m-%d', + 'range_min': '2024-01-01', + 'range_max': '2024-01-03', + 'range_is_nullable': True, + } + assert instance.columns['categorical'] == { + 'sdtype': 'categorical', + 'range_values': ['a', 'b'], + 'range_is_nullable': True, + } + assert instance.columns['ordinal'] == { + 'sdtype': 'ordinal', + 'range_values': [1, 2], + 'range_is_nullable': True, + } + assert instance.columns['boolean'] == { + 'sdtype': 'boolean', + 'range_is_nullable': True, + } + assert instance.columns['id'] == { + 'sdtype': 'id', + 'range_is_nullable': False, + } + assert instance.columns['unknown'] == { + 'sdtype': 'unknown', + 'range_is_nullable': True, + } + + mock_learn_rounding_digits.assert_called_once_with(data['numerical']) + + def test__detect_ranges_does_not_add_range_values_with_500_unique_values(self): + """Test that ``range_values`` is not added when there are 500+ unique values.""" + # Setup + instance = _SingleTableMetadata() + instance.columns = { + 'categorical': {'sdtype': 'categorical'}, + } + data = pd.DataFrame({ + 'categorical': [str(value) for value in range(500)], + }) + + # Run + instance._detect_ranges(data) + + # Assert + assert instance.columns['categorical'] == { + 'sdtype': 'categorical', + 'range_is_nullable': False, + } + def test__detect_columns(self, data): """Test the ``_detect_columns`` method.""" # Setup instance = _SingleTableMetadata() + instance._detect_ranges = Mock() expected_datetime_format = '%Y-%m-%d' data['categorical_pk_candidate'] = ['a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k'] @@ -1227,6 +1368,7 @@ def test__detect_columns(self, data): instance._detect_columns(data) # Assert + instance._detect_ranges.assert_called_once_with(data) assert instance.columns['id']['sdtype'] == 'id' assert instance.columns['numerical']['sdtype'] == 'numerical' assert instance.columns['datetime']['sdtype'] == 'datetime' @@ -1340,8 +1482,9 @@ def test__detect_columns_with_error(self, mock__get_datetime_format): 'numerical': [1, 2, 3], }) - instance._determine_sdtype_for_numbers = Mock(return_value=('numerical', 'False')) - instance._determine_sdtype_for_objects = Mock(return_value=('datetime', 'False')) + instance._determine_sdtype_for_numbers = Mock(return_value=('numerical', False)) + instance._determine_sdtype_for_objects = Mock(return_value=('datetime', False)) + mock__get_datetime_format.return_value = '%Y-%m-%d' # Run instance._detect_columns(data) @@ -1514,7 +1657,7 @@ def test_detect_from_dataframe(self, mock_log): 'categorical': ['cat', 'dog', 'cat', np.nan], 'date': pd.to_datetime(['2021-02-02', np.nan, '2021-03-05', '2022-12-09']), 'int': [1, 2, 3, 4], - 'float': [1.0, 2.0, 3.0, 4], + 'float': [1.0, 2.0, 3.0, 4.2], 'bool': [np.nan, True, False, True], }) @@ -1523,11 +1666,36 @@ def test_detect_from_dataframe(self, mock_log): # Assert assert instance.columns == { - 'categorical': {'sdtype': 'categorical'}, - 'date': {'sdtype': 'datetime'}, - 'int': {'sdtype': 'numerical'}, - 'float': {'sdtype': 'numerical'}, - 'bool': {'sdtype': 'categorical'}, + 'categorical': { + 'sdtype': 'categorical', + 'range_is_nullable': True, + 'range_values': ['cat', 'dog'], + }, + 'date': { + 'sdtype': 'datetime', + 'range_is_nullable': True, + 'range_min': '2021-02-02 00:00:00', + 'range_max': '2022-12-09 00:00:00', + }, + 'int': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 4, + 'decimal_places': 0, + }, + 'float': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1.0, + 'range_max': 4.2, + 'decimal_places': 1, + }, + 'bool': { + 'sdtype': 'categorical', + 'range_is_nullable': True, + 'range_values': [True, False], + }, } expected_log_calls = [ @@ -1540,34 +1708,155 @@ def test_detect_from_dataframe(self, mock_log): def test_detect_from_dataframe_numerical_columns(self, mock_log): """Test the detect from dataframe with columns that are integers""" # Setup + np.random.seed(0) num_rows = 100 num_cols = 20 values = {i + 1: np.random.randint(0, 100, size=num_rows) for i in range(num_cols)} data = pd.DataFrame(values) correct_metadata = { + 'METADATA_SPEC_VERSION': 'SINGLE_TABLE_V2', 'columns': { - '1': {'sdtype': 'numerical'}, - '2': {'sdtype': 'numerical'}, - '3': {'sdtype': 'numerical'}, - '4': {'sdtype': 'numerical'}, - '5': {'sdtype': 'numerical'}, - '6': {'sdtype': 'numerical'}, - '7': {'sdtype': 'numerical'}, - '8': {'sdtype': 'numerical'}, - '9': {'sdtype': 'numerical'}, - '10': {'sdtype': 'numerical'}, - '11': {'sdtype': 'numerical'}, - '12': {'sdtype': 'numerical'}, - '13': {'sdtype': 'numerical'}, - '14': {'sdtype': 'numerical'}, - '15': {'sdtype': 'numerical'}, - '16': {'sdtype': 'numerical'}, - '17': {'sdtype': 'numerical'}, - '18': {'sdtype': 'numerical'}, - '19': {'sdtype': 'numerical'}, - '20': {'sdtype': 'numerical'}, + '1': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '2': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '3': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 97, + 'decimal_places': 0, + }, + '4': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '5': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '6': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '7': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 96, + 'decimal_places': 0, + }, + '8': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '9': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '10': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '11': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 2, + 'range_max': 98, + 'decimal_places': 0, + }, + '12': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 96, + 'decimal_places': 0, + }, + '13': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '14': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '15': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 99, + 'decimal_places': 0, + }, + '16': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 2, + 'range_max': 98, + 'decimal_places': 0, + }, + '17': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 99, + 'decimal_places': 0, + }, + '18': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 98, + 'decimal_places': 0, + }, + '19': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 99, + 'decimal_places': 0, + }, + '20': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 0, + 'range_max': 98, + 'decimal_places': 0, + }, }, - 'METADATA_SPEC_VERSION': 'SINGLE_TABLE_V2', } # Run @@ -1628,7 +1917,7 @@ def test_detect_from_csv(self, mock_log, tmp_path): 'categorical': ['cat', 'dog', 'tiger', np.nan], 'date': pd.to_datetime(['2021-02-02', np.nan, '2021-03-05', '2022-12-09']), 'int': [1, 2, 3, 4], - 'float': [1.0, 2.0, 3.0, 4], + 'float': [1.0, 2.0, 3.0, 4.333], 'bool': [np.nan, True, False, True], }) @@ -1639,11 +1928,37 @@ def test_detect_from_csv(self, mock_log, tmp_path): # Assert assert instance.columns == { - 'categorical': {'sdtype': 'categorical'}, - 'date': {'sdtype': 'datetime', 'datetime_format': '%Y-%m-%d'}, - 'int': {'sdtype': 'numerical'}, - 'float': {'sdtype': 'numerical'}, - 'bool': {'sdtype': 'categorical'}, + 'categorical': { + 'sdtype': 'categorical', + 'range_is_nullable': True, + 'range_values': ['cat', 'dog', 'tiger'], + }, + 'date': { + 'datetime_format': '%Y-%m-%d', + 'sdtype': 'datetime', + 'range_is_nullable': True, + 'range_min': '2021-02-02', + 'range_max': '2022-12-09', + }, + 'int': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 4, + 'decimal_places': 0, + }, + 'float': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1.0, + 'range_max': 4.333, + 'decimal_places': 3, + }, + 'bool': { + 'sdtype': 'categorical', + 'range_is_nullable': True, + 'range_values': [True, False], + }, } expected_log_calls = [ @@ -1689,11 +2004,36 @@ def test_detect_from_csv_with_kwargs(self, mock_log, tmp_path): # Assert assert instance.columns == { - 'categorical': {'sdtype': 'categorical'}, - 'date': {'sdtype': 'datetime'}, - 'int': {'sdtype': 'numerical'}, - 'float': {'sdtype': 'numerical'}, - 'bool': {'sdtype': 'categorical'}, + 'categorical': { + 'sdtype': 'categorical', + 'range_is_nullable': True, + 'range_values': ['cat', 'dog', 'tiger'], + }, + 'date': { + 'sdtype': 'datetime', + 'range_is_nullable': True, + 'range_min': '2021-02-02 00:00:00', + 'range_max': '2022-12-09 00:00:00', + }, + 'int': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1, + 'range_max': 4, + 'decimal_places': 0, + }, + 'float': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 1.0, + 'range_max': 4.0, + 'decimal_places': 0, + }, + 'bool': { + 'sdtype': 'categorical', + 'range_is_nullable': True, + 'range_values': [True, False], + }, } expected_log_calls = [ @@ -4052,14 +4392,19 @@ def test__detect_columns_verbose(self, data, capsys): expected_output = ( '\nDetecting table:\n' "- Column 'id': sdtype='id'\n" - "- Column 'numerical': sdtype='numerical'\n" - "- Column 'datetime': sdtype='datetime', datetime_format='%Y-%m-%d'\n" - "- Column 'alternate_id': sdtype='id'\n" - "- Column 'alternate_id_string': sdtype='id'\n" - "- Column 'categorical': sdtype='categorical'\n" - "- Column 'bool': sdtype='categorical'\n" - "- Column 'unknown': sdtype='categorical'\n" - "- Column 'first_name': sdtype='first_name', pii=True\n" + "- Column 'numerical': sdtype='numerical', range_is_nullable=False, range_min=1, " + 'range_max=11, decimal_places=0\n' + "- Column 'datetime': sdtype='datetime', datetime_format='%Y-%m-%d', " + "range_is_nullable=False, range_min='2022-01-01', range_max='2022-11-01'\n" + "- Column 'alternate_id': sdtype='id', range_is_nullable=False\n" + "- Column 'alternate_id_string': sdtype='id', range_is_nullable=False\n" + "- Column 'categorical': sdtype='categorical', range_is_nullable=False, " + "range_values=['a', 'b']\n" + "- Column 'bool': sdtype='categorical', range_is_nullable=False, " + 'range_values=[True, False]\n' + "- Column 'unknown': sdtype='categorical', range_is_nullable=True, range_values=[" + "'a', 'b', 'c', 1, 2.2, 'd', 'e', 'f']\n" + "- Column 'first_name': sdtype='first_name', pii=True, range_is_nullable=False\n" '\nDetecting primary key:\n' "- primary_key='id'\n" ) @@ -4093,15 +4438,20 @@ def test__detect_columns_verbose_infer_keys_none(self, data, capsys): instance = _SingleTableMetadata() expected_output = ( '\nDetecting table:\n' - "- Column 'id': sdtype='id'\n" - "- Column 'numerical': sdtype='numerical'\n" - "- Column 'datetime': sdtype='datetime', datetime_format='%Y-%m-%d'\n" - "- Column 'alternate_id': sdtype='id'\n" - "- Column 'alternate_id_string': sdtype='id'\n" - "- Column 'categorical': sdtype='categorical'\n" - "- Column 'bool': sdtype='categorical'\n" - "- Column 'unknown': sdtype='categorical'\n" - "- Column 'first_name': sdtype='first_name', pii=True\n" + "- Column 'id': sdtype='id', range_is_nullable=False\n" + "- Column 'numerical': sdtype='numerical', range_is_nullable=False, " + 'range_min=1, range_max=11, decimal_places=0\n' + "- Column 'datetime': sdtype='datetime', datetime_format='%Y-%m-%d', " + "range_is_nullable=False, range_min='2022-01-01', range_max='2022-11-01'\n" + "- Column 'alternate_id': sdtype='id', range_is_nullable=False\n" + "- Column 'alternate_id_string': sdtype='id', range_is_nullable=False\n" + "- Column 'categorical': sdtype='categorical', range_is_nullable=False, " + "range_values=['a', 'b']\n" + "- Column 'bool': sdtype='categorical', range_is_nullable=False, " + 'range_values=[True, False]\n' + "- Column 'unknown': sdtype='categorical', range_is_nullable=True, " + "range_values=['a', 'b', 'c', 1, 2.2, 'd', 'e', 'f']\n" + "- Column 'first_name': sdtype='first_name', pii=True, range_is_nullable=False\n" ) # Run @@ -4112,109 +4462,124 @@ def test__detect_columns_verbose_infer_keys_none(self, data, capsys): assert captured == expected_output @pytest.mark.parametrize( - 'table_name,table_str', - [(None, ''), ('users', " for table 'users'")], + ( + 'columns', + 'infer_sdtypes', + 'pk_candidates', + 'pii_pk_candidates', + 'expected_primary_key', + 'expected_sdtype_updated', + 'expected_pii_removed', + ), + [ + pytest.param( + {'email': {'sdtype': 'unknown', 'pii': True}}, + False, + [], + ['email'], + 'email', + True, + True, + id='updates-sdtype-and-removes-pii', + ), + pytest.param( + {'email': {'sdtype': 'id', 'pii': False}}, + True, + ['email'], + [], + 'email', + False, + True, + id='removes-pii-only', + ), + pytest.param( + {'email': {'sdtype': 'unknown'}}, + False, + [], + ['email'], + 'email', + True, + False, + id='updates-sdtype-only', + ), + pytest.param( + {'email': {'sdtype': 'unknown', 'pii': True}}, + True, + [], + [], + None, + False, + False, + id='no-candidates', + ), + ], ) - def test__select_primary_key_verbose(self, capsys, table_name, table_str): - """Test the ``_select_primary_key`` method with verbose .""" - # Setup - instance = _SingleTableMetadata() - instance.columns = { - 'email': {'sdtype': 'unknown', 'pii': True}, - } - expected_output = ( - f'\nDetecting primary key{table_str}:\n' - f"- primary_key='email' (updating sdtype to 'id', removing 'pii' field)\n" - ) - - # Run - instance._select_primary_key( - infer_sdtypes=False, - pk_candidates=[], - pii_pk_candidates=['email'], - table_name=table_name, - verbose=True, - ) - - # Assert - captured = capsys.readouterr().out - assert captured == expected_output - assert instance.primary_key == 'email' - assert instance.columns['email']['sdtype'] == 'id' - assert 'pii' not in instance.columns['email'] - - def test__select_primary_key_verbose_removes_pii_only(self, capsys): - """Test ``_select_primary_key`` verbose output when only the ``pii`` field is removed.""" + def test__select_primary_key_returns_detection_info( + self, + columns, + infer_sdtypes, + pk_candidates, + pii_pk_candidates, + expected_primary_key, + expected_sdtype_updated, + expected_pii_removed, + ): + """Test that ``_select_primary_key`` returns information about the detected primary key.""" # Setup instance = _SingleTableMetadata() - instance.columns = {'email': {'sdtype': 'id', 'pii': False}} - expected_output = ( - "\nDetecting primary key for table 'users':\n" - "- primary_key='email' (removing 'pii' field)\n" - ) + instance.columns = columns # Run - instance._select_primary_key( - infer_sdtypes=True, - pk_candidates=['email'], - pii_pk_candidates=[], - table_name='users', - verbose=True, + primary_key, sdtype_updated, pii_removed = instance._select_primary_key( + infer_sdtypes=infer_sdtypes, + pk_candidates=pk_candidates, + pii_pk_candidates=pii_pk_candidates, ) # Assert - captured = capsys.readouterr().out - assert captured == expected_output - assert instance.primary_key == 'email' - assert instance.columns['email']['sdtype'] == 'id' - assert 'pii' not in instance.columns['email'] + assert primary_key == expected_primary_key + assert sdtype_updated == expected_sdtype_updated + assert pii_removed == expected_pii_removed + assert instance.primary_key == expected_primary_key - def test__select_primary_key_verbose_updates_sdtype_only(self, capsys): - """Test ``_select_primary_key`` verbose output when only the sdtype is updated to 'id'.""" + def test__detect_columns_verbose_with_table_name(self, data, capsys): + """Test the ``_detect_columns`` verbose output when a table name is provided.""" # Setup instance = _SingleTableMetadata() - instance.columns = {'email': {'sdtype': 'unknown'}} - expected_output = ( - "\nDetecting primary key for table 'users':\n" - "- primary_key='email' (updating sdtype to 'id')\n" - ) # Run - instance._select_primary_key( - infer_sdtypes=False, - pk_candidates=[], - pii_pk_candidates=['email'], - table_name='users', - verbose=True, - ) + instance._detect_columns(data, table_name='users', verbose=True) # Assert captured = capsys.readouterr().out - assert captured == expected_output - assert instance.primary_key == 'email' - assert instance.columns['email']['sdtype'] == 'id' - assert 'pii' not in instance.columns['email'] + assert "\nDetecting table 'users':\n" in captured + assert "\nDetecting primary key for table 'users':\n" in captured + assert "- primary_key='id'\n" in captured - def test__select_primary_key_verbose_no_candidates(self, capsys): - """Test the ``_select_primary_key`` method with verbose and no PK candidates.""" + def test__print_detection(self, capsys): + """Test the ``_print_detection`` method.""" # Setup instance = _SingleTableMetadata() - instance.columns = { - 'email': {'sdtype': 'unknown', 'pii': True}, - } - expected_output = "\nDetecting primary key for table 'table':\n- No primary key found\n" + instance.columns = {'id': {'sdtype': 'id'}} + data = pd.DataFrame({'id': [1, 2, 3]}) # Run - instance._select_primary_key( + instance._print_detection( + 'users', + data, infer_sdtypes=True, - pk_candidates=[], - pii_pk_candidates=[], - table_name='table', - verbose=True, + infer_keys='primary_only', + chosen_pk='id', + sdtype_updated=False, + pii_removed=False, ) # Assert captured = capsys.readouterr().out + expected_output = ( + "\nDetecting table 'users':\n" + "- Column 'id': sdtype='id'\n" + "\nDetecting primary key for table 'users':\n" + "- primary_key='id'\n" + ) assert captured == expected_output - assert instance.primary_key is None - assert instance.columns['email'] == {'sdtype': 'unknown', 'pii': True} From 1ef996ec0ac0ec4d089892e4e6cdbac381caf91a Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 3 Sep 2026 14:58:53 +0100 Subject: [PATCH 03/10] integration tests --- sdv/_utils.py | 6 +- sdv/data_processing/data_processor.py | 2 +- sdv/metadata/_single_table.py | 48 ++-- sdv/metadata/metadata.py | 16 +- .../integration/cag/test_one_hot_encoding.py | 10 +- tests/integration/metadata/test_metadata.py | 210 +++++++++--------- .../integration/metadata/test_single_table.py | 48 ++-- tests/unit/metadata/test_metadata.py | 10 +- tests/unit/metadata/test_single_table.py | 2 - tests/unit/single_table/test_dayz.py | 12 +- tests/utils.py | 62 ++++++ 11 files changed, 240 insertions(+), 186 deletions(-) diff --git a/sdv/_utils.py b/sdv/_utils.py index 7dcb1076e..9a1576fba 100644 --- a/sdv/_utils.py +++ b/sdv/_utils.py @@ -678,13 +678,17 @@ def _metadata_range_exceeds_real(data, metadata): return False -def _validate_data_single_table(data): +def _validate_data_single_table(data, table_name=None): """Validate that the data is a dictionary with a single table.""" if len(data) != 1: raise InvalidDataTypeError( 'The `data` parameter must be a dictionary containing exactly one table name ' 'mapped to a pandas DataFrame.' ) + if table_name is not None and table_name not in data: + raise InvalidDataTypeError( + f"The specified table name '{table_name}' is not present in the data." + ) def _get_single_table_data(data): diff --git a/sdv/data_processing/data_processor.py b/sdv/data_processing/data_processor.py index e8ca62624..cdf6c5630 100644 --- a/sdv/data_processing/data_processor.py +++ b/sdv/data_processing/data_processor.py @@ -218,7 +218,7 @@ def create_anonymized_transformer(sdtype, column_metadata, cardinality_rule, loc """ kwargs = {'locales': locales, 'cardinality_rule': cardinality_rule} for key, value in column_metadata.items(): - if key not in ['pii', 'sdtype']: + if key not in ['pii', 'sdtype', 'range_is_nullable']: kwargs[key] = value try: diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 1df99b6aa..963b20198 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -12,7 +12,7 @@ import pandas as pd from rdt.transformers._validators import AddressValidator, GPSValidator from rdt.transformers.pii.anonymization import SDTYPE_ANONYMIZERS, is_faker_function -from rdt.transformers.utils import learn_rounding_digits +from rdt.transformers.utils import MAX_DECIMALS, learn_rounding_digits from sdv._utils import ( _cast_to_datetime64, @@ -815,33 +815,35 @@ def _detect_ranges(self, data): column_data = data[column_name] sdtype = column_metadata['sdtype'] + if sdtype == 'unknown': + continue column_metadata['range_is_nullable'] = bool(column_data.isna().any()) - if sdtype == 'numerical': - clean_data = column_data.dropna() - if not clean_data.empty: - column_metadata['range_min'] = clean_data.min().item() - column_metadata['range_max'] = clean_data.max().item() + clean_data = column_data.dropna() + if clean_data.empty: + continue - column_metadata['decimal_places'] = learn_rounding_digits(column_data) + if sdtype == 'numerical': + column_metadata['range_min'] = clean_data.min().item() + column_metadata['range_max'] = clean_data.max().item() + digits = learn_rounding_digits(column_data) + column_metadata['decimal_places'] = digits if digits is not None else MAX_DECIMALS elif sdtype == 'datetime': - clean_data = column_data.dropna() - if not clean_data.empty: - datetime_format = column_metadata.get('datetime_format') - clean_data = pd.to_datetime(clean_data, format=datetime_format) - - range_min = clean_data.min() - range_max = clean_data.max() - if datetime_format: - range_min = range_min.strftime(datetime_format) - range_max = range_max.strftime(datetime_format) - else: - range_min = str(range_min) - range_max = str(range_max) - - column_metadata['range_min'] = range_min - column_metadata['range_max'] = range_max + datetime_format = column_metadata.get('datetime_format') + clean_data = pd.to_datetime(clean_data, format=datetime_format) + + range_min = clean_data.min() + range_max = clean_data.max() + if datetime_format: + range_min = range_min.strftime(datetime_format) + range_max = range_max.strftime(datetime_format) + else: + range_min = str(range_min) + range_max = str(range_max) + + column_metadata['range_min'] = range_min + column_metadata['range_max'] = range_max elif sdtype in {'categorical', 'ordinal'}: range_values = self._detect_range_values(column_data) diff --git a/sdv/metadata/metadata.py b/sdv/metadata/metadata.py index 6d95ac038..6da504383 100644 --- a/sdv/metadata/metadata.py +++ b/sdv/metadata/metadata.py @@ -21,6 +21,7 @@ _is_numerical, _load_data_from_csv, _validate_boolean_parameter, + _validate_data_single_table, ) from sdv.errors import InvalidDataError from sdv.logging import get_sdv_logger @@ -728,10 +729,16 @@ def _detect_foreign_keys_by_column_name(self, data, verbose=False): try: sdtype_updated = False if pk_sdtype == 'id' and original_fk_sdtype != 'id': + update_kwargs = {'sdtype': 'id'} + if 'range_is_nullable' in original_fk_meta: + update_kwargs['range_is_nullable'] = original_fk_meta[ + 'range_is_nullable' + ] + self.update_column( table_name=child_candidate, column_name=primary_key, - sdtype='id', + **update_kwargs, ) sdtype_updated = True self.add_relationship( @@ -955,7 +962,7 @@ def _detect_from_dataframe( def detect_from_dataframe( cls, data, - table_name=DEFAULT_SINGLE_TABLE_NAME, + table_name, infer_sdtypes=True, infer_keys='primary_only', verbose=False, @@ -966,7 +973,7 @@ def detect_from_dataframe( All data column names are converted to strings. Args: - data (pandas.DataFrame): + data (dict[str, pd.DataFrame]): The data to detect metadata from. table_name (str): The name of the table to detect. If None, a default name will be used. @@ -991,8 +998,9 @@ def detect_from_dataframe( Metadata: A new metadata object with the sdtypes detected from the data. """ + _validate_data_single_table(data, table_name) return cls._detect_from_dataframe( - data=data, + data=data[table_name], table_name=table_name, infer_sdtypes=infer_sdtypes, infer_keys=infer_keys, diff --git a/tests/integration/cag/test_one_hot_encoding.py b/tests/integration/cag/test_one_hot_encoding.py index 3e4e367af..11d5275c5 100644 --- a/tests/integration/cag/test_one_hot_encoding.py +++ b/tests/integration/cag/test_one_hot_encoding.py @@ -162,7 +162,8 @@ def test_end_to_end_numerical_and_categorical(): df = pd.DataFrame(data, columns=columns) # Setup metadata - metadata = Metadata.detect_from_dataframe(df, table_name='one_hot') + data = {'one_hot': df} + metadata = Metadata.detect_from_dataframe(data, table_name='one_hot') for sdtype in ['numerical', 'categorical']: metadata.update_columns(columns, sdtype=sdtype) synthesizer = GaussianCopulaSynthesizer(metadata) @@ -170,7 +171,7 @@ def test_end_to_end_numerical_and_categorical(): # Run synthesizer.add_constraints([constraint]) - synthesizer.fit({'one_hot': df}) + synthesizer.fit(data) samples = synthesizer.sample('one_hot', 100)['one_hot'] # Assert @@ -193,14 +194,15 @@ def test_end_to_end_boolean(): df = pd.DataFrame(data, columns=columns) # Setup metadata - metadata = Metadata.detect_from_dataframe(df, table_name='one_hot') + data = {'one_hot': df} + metadata = Metadata.detect_from_dataframe(data, table_name='one_hot') metadata.update_columns(columns, sdtype='boolean') synthesizer = GaussianCopulaSynthesizer(metadata) constraint = OneHotEncoding(column_names=columns) # Run synthesizer.add_constraints([constraint]) - synthesizer.fit({'one_hot': df}) + synthesizer.fit(data) samples = synthesizer.sample('one_hot', 100)['one_hot'] # Assert diff --git a/tests/integration/metadata/test_metadata.py b/tests/integration/metadata/test_metadata.py index 0c520c34f..d2db5dfac 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -11,7 +11,12 @@ from sdv.metadata.errors import InvalidMetadataError from sdv.metadata.metadata import Metadata from sdv.single_table.copulas import GaussianCopulaSynthesizer -from tests.utils import download_test_demo, get_multi_table_metadata +from tests.utils import ( + compare_metadata, + compare_ranges, + download_test_demo, + get_multi_table_metadata, +) DEFAULT_TABLE_NAME = 'table' @@ -74,12 +79,6 @@ def test_detect_from_dataframes_multi_table(): metadata = Metadata.detect_from_dataframes(real_data) # Assert - metadata.update_column( - table_name='hotels', - column_name='classification', - sdtype='categorical', - ) - expected_metadata = { 'tables': { 'hotels': { @@ -118,7 +117,8 @@ def test_detect_from_dataframes_multi_table(): ], 'METADATA_SPEC_VERSION': 'V2', } - assert metadata.to_dict() == expected_metadata + compare_metadata(metadata, expected_metadata) + compare_ranges(metadata, real_data) def test_detect_from_dataframes_multi_table_without_infer_sdtypes(): @@ -130,12 +130,6 @@ def test_detect_from_dataframes_multi_table_without_infer_sdtypes(): metadata = Metadata.detect_from_dataframes(real_data, infer_sdtypes=False) # Assert - metadata.update_column( - table_name='hotels', - column_name='classification', - sdtype='categorical', - ) - expected_metadata = { 'tables': { 'hotels': { @@ -144,7 +138,7 @@ def test_detect_from_dataframes_multi_table_without_infer_sdtypes(): 'city': {'sdtype': 'unknown', 'pii': True}, 'state': {'sdtype': 'unknown', 'pii': True}, 'rating': {'sdtype': 'unknown', 'pii': True}, - 'classification': {'sdtype': 'categorical'}, + 'classification': {'sdtype': 'unknown', 'pii': True}, }, 'primary_key': 'hotel_id', }, @@ -174,7 +168,8 @@ def test_detect_from_dataframes_multi_table_without_infer_sdtypes(): ], 'METADATA_SPEC_VERSION': 'V2', } - assert metadata.to_dict() == expected_metadata + compare_metadata(metadata, expected_metadata) + compare_ranges(metadata, real_data) def test_detect_from_dataframes_multi_table_with_infer_keys_primary_only(): @@ -186,12 +181,6 @@ def test_detect_from_dataframes_multi_table_with_infer_keys_primary_only(): metadata = Metadata.detect_from_dataframes(real_data, infer_keys='primary_only') # Assert - metadata.update_column( - table_name='hotels', - column_name='classification', - sdtype='categorical', - ) - expected_metadata = { 'tables': { 'hotels': { @@ -223,7 +212,8 @@ def test_detect_from_dataframes_multi_table_with_infer_keys_primary_only(): 'relationships': [], 'METADATA_SPEC_VERSION': 'V2', } - assert metadata.to_dict() == expected_metadata + compare_metadata(metadata, expected_metadata) + compare_ranges(metadata, real_data) def test_detect_from_dataframes_multi_table_with_infer_keys_none(): @@ -235,12 +225,6 @@ def test_detect_from_dataframes_multi_table_with_infer_keys_none(): metadata = Metadata.detect_from_dataframes(real_data, infer_keys=None) # Assert - metadata.update_column( - table_name='hotels', - column_name='classification', - sdtype='categorical', - ) - expected_metadata = { 'tables': { 'hotels': { @@ -270,14 +254,16 @@ def test_detect_from_dataframes_multi_table_with_infer_keys_none(): 'relationships': [], 'METADATA_SPEC_VERSION': 'V2', } - assert metadata.to_dict() == expected_metadata + compare_metadata(metadata, expected_metadata) + compare_ranges(metadata, real_data) def test_detect_from_dataframes_single_table(): """Test the ``detect_from_dataframes`` method works with a single table.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframes({'table_1': data['hotels']}) + data = {'table_1': data['hotels']} + metadata = Metadata.detect_from_dataframes(data) # Run metadata.validate() @@ -299,14 +285,16 @@ def test_detect_from_dataframes_single_table(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_dataframes_single_table_infer_sdtypes_false(): """Test it for a single table when infer_sdtypes is False.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframes({'table_1': data['hotels']}, infer_sdtypes=False) + data = {'table_1': data['hotels']} + metadata = Metadata.detect_from_dataframes(data, infer_sdtypes=False) # Run metadata.validate() @@ -328,16 +316,16 @@ def test_detect_from_dataframes_single_table_infer_sdtypes_false(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_metadata(metadata, expected_metadata) + compare_ranges(metadata, data) def test_detect_from_dataframes_single_table_infer_keys_primary_only(): """Test it for a single table when infer_keys is 'primary_only'.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframes( - {'table_1': data['hotels']}, infer_keys='primary_only' - ) + data = {'table_1': data['hotels']} + metadata = Metadata.detect_from_dataframes(data, infer_keys='primary_only') # Run metadata.validate() @@ -359,14 +347,16 @@ def test_detect_from_dataframes_single_table_infer_keys_primary_only(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_dataframes_single_table_infer_keys_none(): """Test it for a single table when infer_keys is None.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframes({'table_1': data['hotels']}, infer_keys=None) + data = {'table_1': data['hotels']} + metadata = Metadata.detect_from_dataframes(data, infer_keys=None) # Run metadata.validate() @@ -387,15 +377,17 @@ def test_detect_from_dataframes_single_table_infer_keys_none(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_dataframe(): """Test that a single table can be detected as a DataFrame.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') + data = {'table': data['hotels']} - metadata = Metadata.detect_from_dataframe(data['hotels']) + metadata = Metadata.detect_from_dataframe(data, 'table') # Run metadata.validate() @@ -417,14 +409,16 @@ def test_detect_from_dataframe(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_dataframe_infer_sdtypes_false(): """Test it when infer_sdtypes is False.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframe(data['hotels'], infer_sdtypes=False) + data = {'table': data['hotels']} + metadata = Metadata.detect_from_dataframe(data, 'table', infer_sdtypes=False) # Run metadata.validate() @@ -446,14 +440,16 @@ def test_detect_from_dataframe_infer_sdtypes_false(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_dataframe_infer_keys_none(): """Test it when infer_keys is None.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframe(data['hotels'], infer_keys=None) + data = {'table': data['hotels']} + metadata = Metadata.detect_from_dataframe(data, 'table', infer_keys=None) # Run metadata.validate() @@ -474,14 +470,16 @@ def test_detect_from_dataframe_infer_keys_none(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_dataframe_infer_keys_none_infer_sdtypes_false(): """Test it when infer_keys is None and infer_sdtypes is False.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - metadata = Metadata.detect_from_dataframe(data['hotels'], infer_keys=None, infer_sdtypes=False) + data = {'table': data['hotels']} + metadata = Metadata.detect_from_dataframe(data, 'table', infer_keys=None, infer_sdtypes=False) # Run metadata.validate() @@ -502,7 +500,8 @@ def test_detect_from_dataframe_infer_keys_none_infer_sdtypes_false(): }, 'relationships': [], } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, data) + compare_metadata(metadata, expected_metadata) def test_detect_from_csvs(tmp_path): @@ -520,12 +519,6 @@ def test_detect_from_csvs(tmp_path): metadata.detect_from_csvs(folder_name=tmp_path) # Assert - metadata.update_column( - table_name='hotels', - column_name='classification', - sdtype='categorical', - ) - expected_metadata = { 'tables': { 'hotels': { @@ -565,7 +558,8 @@ def test_detect_from_csvs(tmp_path): 'METADATA_SPEC_VERSION': 'V2', } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, real_data) + compare_metadata(metadata, expected_metadata) params = [ @@ -1576,28 +1570,32 @@ def test_detect_from_dataframe_verbose_single(capsys): """Test 'detect_from_dataframe' with verbose True with single table.""" # Setup data, _ = download_test_demo(modality='single_table', dataset_name='fake_hotel_guests') - data_table = data['fake_hotel_guests'] - expected_print = ( - "\nDetecting table 'table':\n" - "- Column 'guest_email': sdtype='email', pii=True\n" - "- Column 'has_rewards': sdtype='categorical'\n" - "- Column 'room_type': sdtype='categorical'\n" - "- Column 'amenities_fee': sdtype='numerical'\n" - "- Column 'checkin_date': sdtype='datetime', datetime_format='%d %b %Y'\n" - "- Column 'checkout_date': sdtype='datetime', datetime_format='%d %b %Y'\n" - "- Column 'room_rate': sdtype='numerical'\n" - "- Column 'billing_address': sdtype='categorical'\n" - "- Column 'credit_card_number': sdtype='credit_card_number', pii=True\n" - "\nDetecting primary key for table 'table':\n" - "- primary_key='guest_email'\n" - ) + data = {'table': data['fake_hotel_guests']} # Run - metadata = Metadata.detect_from_dataframe(data_table, verbose=True) + metadata = Metadata.detect_from_dataframe(data, 'table', verbose=True) # Assert captured = capsys.readouterr().out - assert captured == expected_print + expected_output = [ + "\nDetecting table 'table':\n", + "- Column 'guest_email': sdtype='email', pii=True\n", + "- Column 'has_rewards': sdtype='categorical', range_is_nullable=", + "- Column 'room_type': sdtype='categorical', range_is_nullable=", + "- Column 'amenities_fee': sdtype='numerical', range_is_nullable=", + "- Column 'checkin_date': sdtype='datetime', datetime_format='%d %b %Y', " + 'range_is_nullable=', + "- Column 'checkout_date': sdtype='datetime', datetime_format='%d %b %Y', " + 'range_is_nullable=', + "- Column 'room_rate': sdtype='numerical', range_is_nullable=", + "- Column 'billing_address': sdtype='categorical', range_is_nullable=", + "- Column 'credit_card_number': sdtype='credit_card_number', pii=True, range_is_nullable=", + "\nDetecting primary key for table 'table':\n", + "- primary_key='guest_email'\n", + ] + for line in expected_output: + assert line in captured + assert list(metadata.tables.keys()) == ['table'] assert list(metadata.tables['table'].columns.keys()) == [ 'guest_email', @@ -1616,38 +1614,36 @@ def test_detect_from_dataframes_verbose(capsys): """Test 'detect_from_dataframe' with verbose True with multi table.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') - expected_print = ( - "\nDetecting table 'guests':\n" - "- Column 'guest_email': sdtype='email', pii=True\n" - "- Column 'hotel_id': sdtype='id'\n" - "- Column 'has_rewards': sdtype='categorical'\n" - "- Column 'room_type': sdtype='categorical'\n" - "- Column 'amenities_fee': sdtype='numerical'\n" - "- Column 'checkin_date': sdtype='datetime', datetime_format='%d %b %Y'\n" - "- Column 'checkout_date': sdtype='datetime', datetime_format='%d %b %Y'\n" - "- Column 'room_rate': sdtype='numerical'\n" - "- Column 'billing_address': sdtype='categorical'\n" - "- Column 'credit_card_number': sdtype='credit_card_number', pii=True\n" - "\nDetecting primary key for table 'guests':\n" - "- primary_key='guest_email'\n" - "\nDetecting table 'hotels':\n" - "- Column 'hotel_id': sdtype='id'\n" - "- Column 'city': sdtype='city', pii=True\n" - "- Column 'state': sdtype='administrative_unit', pii=True\n" - "- Column 'rating': sdtype='numerical'\n" - "- Column 'classification': sdtype='categorical'\n" - "\nDetecting primary key for table 'hotels':\n" - "- primary_key='hotel_id'\n" - '\nDetecting foreign keys:\n' - "- Column 'guests.hotel_id' refers to column 'hotels.hotel_id'\n" - ) # Run metadata = Metadata.detect_from_dataframes(data, verbose=True) # Assert captured = capsys.readouterr().out - assert captured == expected_print + expected_output = [ + "\nDetecting table 'guests':\n", + "- Column 'guest_email': sdtype='email', pii=True\n", + "- Column 'hotel_id': sdtype='id', range_is_nullable=", + "- Column 'has_rewards': sdtype='categorical', range_is_nullable=", + "- Column 'amenities_fee': sdtype='numerical', range_is_nullable=", + "- Column 'checkin_date': sdtype='datetime', datetime_format='%d %b %Y', " + 'range_is_nullable=', + "\nDetecting primary key for table 'guests':\n", + "- primary_key='guest_email'\n", + "\nDetecting table 'hotels':\n", + "- Column 'hotel_id': sdtype='id'\n", + "- Column 'city': sdtype='city', pii=True, range_is_nullable=", + "- Column 'state': sdtype='administrative_unit', pii=True, range_is_nullable=", + "- Column 'rating': sdtype='numerical', range_is_nullable=", + "- Column 'classification': sdtype='categorical', range_is_nullable=", + "\nDetecting primary key for table 'hotels':\n", + "- primary_key='hotel_id'\n", + '\nDetecting foreign keys:\n', + "- Column 'guests.hotel_id' refers to column 'hotels.hotel_id'\n", + ] + for line in expected_output: + assert line in captured + assert list(metadata.tables.keys()) == ['guests', 'hotels'] @@ -1665,12 +1661,14 @@ def test_detect_from_dataframes_verbose_updates_fk_sdtype(capsys): } expected_output = ( "\nDetecting table 'users':\n" - "- Column 'account': sdtype='categorical'\n\n" + "- Column 'account': sdtype='id'\n\n" "Detecting primary key for table 'users':\n" "- primary_key='account' (updating sdtype to 'id')\n\n" "Detecting table 'transactions':\n" "- Column 'transaction_id': sdtype='id'\n" - "- Column 'account': sdtype='categorical'\n\n" + "- Column 'account': sdtype='categorical', range_is_nullable=False, " + "range_values=['acct_0', 'acct_1', 'acct_2', 'acct_3', 'acct_4', 'acct_5', " + "'acct_6', 'acct_7', 'acct_8', 'acct_9']\n\n" "Detecting primary key for table 'transactions':\n" "- primary_key='transaction_id'\n\n" 'Detecting foreign keys:\n' @@ -1698,7 +1696,8 @@ def test_detect_from_dataframes_verbose_no_pk_found(capsys): } expected_output = ( "\nDetecting table 'users':\n" - "- Column 'date': sdtype='datetime'\n" + "- Column 'date': sdtype='datetime', range_is_nullable=False, " + "range_min='2023-01-01 00:00:00', range_max='2023-01-10 00:00:00'\n" "\nDetecting primary key for table 'users':\n" '- No primary key found\n' '\nDetecting foreign keys:\n' @@ -1924,12 +1923,6 @@ def test_detect_from_dataframes(): metadata = Metadata.detect_from_dataframes(real_data) # Assert - metadata.update_column( - table_name='hotels', - column_name='classification', - sdtype='categorical', - ) - expected_metadata = { 'tables': { 'hotels': { @@ -1968,7 +1961,8 @@ def test_detect_from_dataframes(): ], 'METADATA_SPEC_VERSION': 'V2', } - assert metadata.to_dict() == expected_metadata + compare_ranges(metadata, real_data) + compare_metadata(metadata, expected_metadata) def test_get_column_names(): diff --git a/tests/integration/metadata/test_single_table.py b/tests/integration/metadata/test_single_table.py index 8b72f88df..8379f8801 100644 --- a/tests/integration/metadata/test_single_table.py +++ b/tests/integration/metadata/test_single_table.py @@ -267,32 +267,6 @@ def test_upgrade_metadata(tmp_path): assert new_metadata == expected_metadata -def test_validate_unknown_sdtype(): - """Test ``validate`` method works with ``unknown`` sdtype.""" - # Setup - data, _ = download_demo(modality='multi_table', dataset_name='fake_hotels') - - metadata = _SingleTableMetadata() - metadata.detect_from_dataframe(data['hotels']) - - # Run - metadata.validate() - - # Assert - expected_metadata = { - 'METADATA_SPEC_VERSION': 'SINGLE_TABLE_V2', - 'columns': { - 'hotel_id': {'sdtype': 'id'}, - 'city': {'sdtype': 'city', 'pii': True}, - 'state': {'sdtype': 'administrative_unit', 'pii': True}, - 'rating': {'sdtype': 'numerical'}, - 'classification': {'sdtype': 'categorical'}, - }, - 'primary_key': 'hotel_id', - } - assert metadata.to_dict() == expected_metadata - - def test_detect_from_dataframe_with_none_nan_and_nat(): """Test ``detect_from_dataframe`` with ``None``, ``np.nan`` and ``pd.NaT``.""" # Setup @@ -332,16 +306,15 @@ def test_detect_from_dataframe_with_pii_names(): # Assert expected_metadata = { - 'METADATA_SPEC_VERSION': 'SINGLE_TABLE_V2', 'primary_key': 'USER PHONE NUMBER', 'columns': { - 'USER PHONE NUMBER': {'sdtype': 'phone_number', 'pii': True}, - 'addr_line_1': {'sdtype': 'street_address', 'pii': True}, - 'First Name': {'sdtype': 'first_name', 'pii': True}, - 'guest_email': {'sdtype': 'email', 'pii': True}, + 'USER PHONE NUMBER': {'pii': True, 'sdtype': 'phone_number'}, + 'addr_line_1': {'pii': True, 'sdtype': 'street_address', 'range_is_nullable': False}, + 'First Name': {'pii': True, 'sdtype': 'first_name', 'range_is_nullable': False}, + 'guest_email': {'pii': True, 'sdtype': 'email', 'range_is_nullable': False}, }, + 'METADATA_SPEC_VERSION': 'SINGLE_TABLE_V2', } - assert metadata.to_dict() == expected_metadata @@ -654,7 +627,16 @@ def test_metadata_detection_numerical_dtypes(): # Assert expected_metadata = { - 'columns': {column: {'sdtype': 'numerical'} for column in data.columns}, + 'columns': { + column: { + 'sdtype': 'numerical', + 'range_is_nullable': data[column].isna().any(), + 'range_min': data[column].min().item(), + 'range_max': data[column].max().item(), + 'decimal_places': not (data[column] == data[column].round()).all(), + } + for column in data.columns + }, } assert metadata.to_dict()['columns'] == expected_metadata['columns'] diff --git a/tests/unit/metadata/test_metadata.py b/tests/unit/metadata/test_metadata.py index 95d1ca02a..d7909bd94 100644 --- a/tests/unit/metadata/test_metadata.py +++ b/tests/unit/metadata/test_metadata.py @@ -4900,22 +4900,24 @@ def test__detect_from_dataframe_bad_input_infer_keys(self): Metadata._detect_from_dataframe(data, infer_keys=infer_keys) @patch.object(Metadata, '_detect_from_dataframe') - def test_detect_from_dataframe(self, mock_detect): + @patch('sdv.metadata.metadata._validate_data_single_table') + def test_detect_from_dataframe(self, mock_validate, mock_detect): """Test the `detect_from_dataframe` method.""" # Setup - data = pd.DataFrame() + data = {'table': pd.DataFrame()} # Run - metadata = Metadata.detect_from_dataframe(data) + metadata = Metadata.detect_from_dataframe(data, 'table') # Assert mock_detect.assert_called_once_with( - data=data, + data=data['table'], table_name='table', infer_sdtypes=True, infer_keys='primary_only', verbose=False, ) + mock_validate.assert_called_once_with(data, 'table') assert metadata == mock_detect.return_value def test__handle_table_name(self): diff --git a/tests/unit/metadata/test_single_table.py b/tests/unit/metadata/test_single_table.py index 30d106784..94af6cc78 100644 --- a/tests/unit/metadata/test_single_table.py +++ b/tests/unit/metadata/test_single_table.py @@ -1331,9 +1331,7 @@ def test__detect_ranges(self, mock_learn_rounding_digits): } assert instance.columns['unknown'] == { 'sdtype': 'unknown', - 'range_is_nullable': True, } - mock_learn_rounding_digits.assert_called_once_with(data['numerical']) def test__detect_ranges_does_not_add_range_values_with_500_unique_values(self): diff --git a/tests/unit/single_table/test_dayz.py b/tests/unit/single_table/test_dayz.py index 4a1904c3d..d53b4c886 100644 --- a/tests/unit/single_table/test_dayz.py +++ b/tests/unit/single_table/test_dayz.py @@ -561,8 +561,8 @@ def test__validate_parameters_errors_with_multi_table_metadata(self): def test_create_parameters_returns_valid_defaults(self): """Test create_parameters returns valid defaults.""" # Setup - data = pd.DataFrame({'col': [np.nan]}) - metadata = Metadata.detect_from_dataframe(data) + data = {'table': pd.DataFrame({'col': [np.nan]})} + metadata = Metadata.detect_from_dataframe(data, 'table') # Run params = DayZSynthesizer.create_parameters(data, metadata) @@ -583,8 +583,8 @@ def test_create_parameters_returns_valid_defaults(self): def test_create_parameters_all_null_categorical_column(self): """Categorical column with all nulls should not have the category_values key parameter.""" # Setup - data = pd.DataFrame({'col': [None, None, np.nan, pd.NA]}) - metadata = Metadata.detect_from_dataframe(data) + data = {'table': pd.DataFrame({'col': [None, None, np.nan, pd.NA]})} + metadata = Metadata.detect_from_dataframe(data, 'table') # Run params = DayZSynthesizer.create_parameters(data, metadata) @@ -629,8 +629,8 @@ def test_create_parameters_all_null_numerical_column(self): def test_create_parameters_all_null_datetime_column(self): """Datetime column with all nulls should omit start/end timestamps.""" # Setup - data = pd.DataFrame({'col': pd.to_datetime([None, None])}) - metadata = Metadata.detect_from_dataframe(data) + data = {'table': pd.DataFrame({'col': pd.to_datetime([None, None])})} + metadata = Metadata.detect_from_dataframe(data, 'table') # Run params = DayZSynthesizer.create_parameters(data, metadata) diff --git a/tests/utils.py b/tests/utils.py index 70f5a73b5..8739aebaf 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -5,6 +5,7 @@ from functools import lru_cache import pandas as pd +from rdt.transformers.utils import learn_rounding_digits from sdv.datasets.demo import download_demo from sdv.logging import get_sdv_logger @@ -12,6 +13,14 @@ from sdv.multi_table import HMASynthesizer from sdv.single_table import GaussianCopulaSynthesizer +RANGE_KEYS = { + 'range_is_nullable', + 'range_min', + 'range_max', + 'range_values', + 'decimal_places', +} + class DataFrameMatcher: """Match a given Pandas DataFrame in a mock function call.""" @@ -231,3 +240,56 @@ def download_test_demo(modality, dataset_name): """ data, metadata = _download_demo(modality, dataset_name) return deepcopy(data), deepcopy(metadata) + + +def compare_metadata(metadata, expected_metadata): + """Compare metadata, allowing detected range fields to be omitted from expected metadata.""" + actual = metadata.to_dict() if isinstance(metadata, Metadata) else deepcopy(metadata) + expected = ( + expected_metadata.to_dict() + if isinstance(expected_metadata, Metadata) + else deepcopy(expected_metadata) + ) + + for table_name, table in actual['tables'].items(): + for column_name, column in table['columns'].items(): + expected_column = expected['tables'][table_name]['columns'][column_name] + for key in RANGE_KEYS: + if key not in expected_column: + column.pop(key, None) + + assert actual == expected + + +def compare_ranges(metadata, data): + """Check that detected ranges are consistent with the source data.""" + metadata = metadata.to_dict() if isinstance(metadata, Metadata) else metadata + for table_name, table in metadata['tables'].items(): + primary_key = table.get('primary_key') + primary_keys = {primary_key} if isinstance(primary_key, str) else set(primary_key or []) + for column_name, column in table['columns'].items(): + sdtype = column.get('sdtype') + range_keys = set(column) & RANGE_KEYS + if column_name in primary_keys or sdtype == 'unknown': + assert not range_keys + continue + + column_data = data[table_name][column_name] + clean_data = column_data.dropna() + + if 'range_is_nullable' in column: + assert column['range_is_nullable'] == column_data.isna().any() + + if 'range_values' in column: + assert set(column['range_values']) == set(clean_data) + + if 'range_min' in column: + if column['sdtype'] == 'datetime': + assert pd.to_datetime(column['range_min']) == pd.to_datetime(clean_data).min() + assert pd.to_datetime(column['range_max']) == pd.to_datetime(clean_data).max() + else: + assert column['range_min'] == clean_data.min() + assert column['range_max'] == clean_data.max() + + if 'decimal_places' in column: + assert column['decimal_places'] == learn_rounding_digits(column_data) From 1821d4bab2c0ce60eb545320e04e8ff10ba522e2 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 3 Sep 2026 15:46:13 +0100 Subject: [PATCH 04/10] cleaning --- sdv/_utils.py | 1 + sdv/metadata/_single_table.py | 4 -- ..._single_table.py => test__single_table.py} | 0 ..._single_table.py => test__single_table.py} | 0 tests/unit/test__utils.py | 42 ++++++++++++++----- 5 files changed, 32 insertions(+), 15 deletions(-) rename tests/integration/metadata/{test_single_table.py => test__single_table.py} (100%) rename tests/unit/metadata/{test_single_table.py => test__single_table.py} (100%) diff --git a/sdv/_utils.py b/sdv/_utils.py index 9a1576fba..963e8ed31 100644 --- a/sdv/_utils.py +++ b/sdv/_utils.py @@ -685,6 +685,7 @@ def _validate_data_single_table(data, table_name=None): 'The `data` parameter must be a dictionary containing exactly one table name ' 'mapped to a pandas DataFrame.' ) + if table_name is not None and table_name not in data: raise InvalidDataTypeError( f"The specified table name '{table_name}' is not present in the data." diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 963b20198..c587d586e 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -583,10 +583,6 @@ def _detect_ordinal_sdtype(self, data): Args: data (pandas.Series): The data to be analyzed. - - Returns: - str or None: - 'ordinal' if the column is ordinal, otherwise ``None``. """ if len(data) <= self._MIN_ROWS_FOR_PREDICTION: return None diff --git a/tests/integration/metadata/test_single_table.py b/tests/integration/metadata/test__single_table.py similarity index 100% rename from tests/integration/metadata/test_single_table.py rename to tests/integration/metadata/test__single_table.py diff --git a/tests/unit/metadata/test_single_table.py b/tests/unit/metadata/test__single_table.py similarity index 100% rename from tests/unit/metadata/test_single_table.py rename to tests/unit/metadata/test__single_table.py diff --git a/tests/unit/test__utils.py b/tests/unit/test__utils.py index 755b73760..706456f0d 100644 --- a/tests/unit/test__utils.py +++ b/tests/unit/test__utils.py @@ -1491,23 +1491,43 @@ def test__metadata_range_exceeds_real(mock__column_range_exceeds_real): ]) +@pytest.mark.parametrize( + 'data, table_name, expected_message', + [ + pytest.param( + { + 'table_1': pd.DataFrame({'col1': ['a', 'b', 'c']}), + 'table_2': pd.DataFrame({'col1': [1, 2, 3]}), + }, + None, + ( + 'The `data` parameter must be a dictionary containing exactly one table name ' + 'mapped to a pandas DataFrame.' + ), + id='multiple-tables', + ), + pytest.param( + {'table': pd.DataFrame({'col1': [1, 2, 3]})}, + 'missing_table', + "The specified table name 'missing_table' is not present in the data.", + id='missing-table-name', + ), + ], +) +def test__validate_data_single_table_invalid(data, table_name, expected_message): + """Test the `_validate_data_single_table` method with invalid inputs.""" + # Run and Assert + with pytest.raises(InvalidDataTypeError, match=re.escape(expected_message)): + _validate_data_single_table(data, table_name=table_name) + + def test__validate_data_single_table(): - """Test the `_validate_data_single_table` method.""" + """Test the `_validate_data_single_table` method with valid input.""" # Setup data = {'table': pd.DataFrame({'col1': [1, 2, 3]})} - invalid_data = { - 'table_1': pd.DataFrame({'col1': ['a', 'b', 'c']}), - 'table_2': pd.DataFrame({'col1': [1, 2, 3]}), - } - expected_message = re.escape( - 'The `data` parameter must be a dictionary containing exactly one table name ' - 'mapped to a pandas DataFrame.' - ) # Run and Assert _validate_data_single_table(data) - with pytest.raises(InvalidDataTypeError, match=expected_message): - _validate_data_single_table(invalid_data) @patch('sdv._utils._validate_data_single_table') From 6718f7c21ece327b6c68814f0cc4287ad4a9f418 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 3 Sep 2026 17:26:36 +0100 Subject: [PATCH 05/10] add integration test small data --- sdv/metadata/_single_table.py | 2 +- tests/integration/metadata/test_metadata.py | 100 ++++++++++++++++++++ 2 files changed, 101 insertions(+), 1 deletion(-) diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index c587d586e..781f700f2 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -827,7 +827,7 @@ def _detect_ranges(self, data): elif sdtype == 'datetime': datetime_format = column_metadata.get('datetime_format') - clean_data = pd.to_datetime(clean_data, format=datetime_format) + clean_data = pd.to_datetime(clean_data, format=datetime_format, errors='coerce') range_min = clean_data.min() range_max = clean_data.max() diff --git a/tests/integration/metadata/test_metadata.py b/tests/integration/metadata/test_metadata.py index d2db5dfac..f89a65594 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -1688,6 +1688,106 @@ def test_detect_from_dataframes_verbose_updates_fk_sdtype(capsys): assert metadata.tables['transactions'].columns['account']['sdtype'] == 'id' +def test_detect_from_dataframes_small_dataset(): + """Test `detect_from_dataframes` by comparing with the expected metadata.""" + # Setup + num_rows = 50 + data = { + 'users': pd.DataFrame({ + 'user_id': [f'user_{i}' for i in range(num_rows)], + 'age': range(20, 20 + num_rows), + 'signup_date': [str(d) for d in pd.date_range('2026-01-01', periods=num_rows)], + 'is_active': [True, False] * 25, + }), + 'transactions': pd.DataFrame({ + 'transaction_id': [f'transaction_{i}' for i in range(num_rows)], + 'user_id': [f'user_{i}' for i in range(num_rows)], + 'category': ['food', 'travel'] * 25, + 'rating': [1, 2, 3, 4, 5] * 10, + 'amount': [10.5 + i for i in range(num_rows)], + }), + } + + expected_metadata = { + 'tables': { + 'users': { + 'primary_key': 'user_id', + 'columns': { + 'user_id': { + 'sdtype': 'id', + }, + 'age': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 20, + 'range_max': 69, + 'decimal_places': 0, + }, + 'signup_date': { + 'sdtype': 'datetime', + 'datetime_format': '%Y-%m-%d %H:%M:%S', + 'range_is_nullable': False, + 'range_min': '2026-01-01 00:00:00', + 'range_max': '2026-02-19 00:00:00', + }, + 'is_active': { + 'sdtype': 'categorical', + 'range_is_nullable': False, + 'range_values': [True, False], + }, + }, + }, + 'transactions': { + 'primary_key': 'transaction_id', + 'columns': { + 'transaction_id': { + 'sdtype': 'id', + }, + 'user_id': { + 'sdtype': 'id', + 'range_is_nullable': False, + }, + 'category': { + 'sdtype': 'categorical', + 'range_is_nullable': False, + 'range_values': ['food', 'travel'], + }, + 'rating': { + 'sdtype': 'ordinal', + 'range_is_nullable': False, + 'range_values': [1, 2, 3, 4, 5], + }, + 'amount': { + 'sdtype': 'numerical', + 'range_is_nullable': False, + 'range_min': 10.5, + 'range_max': 59.5, + 'decimal_places': 1, + }, + }, + }, + }, + 'relationships': [ + { + 'parent_table_name': 'users', + 'child_table_name': 'transactions', + 'parent_primary_key': 'user_id', + 'child_foreign_key': 'user_id', + }, + ], + 'METADATA_SPEC_VERSION': 'V2', + } + + # Run + metadata = Metadata.detect_from_dataframes( + data, + foreign_key_inference_algorithm='column_name_match', + ) + + # Assert + assert metadata.to_dict() == expected_metadata + + def test_detect_from_dataframes_verbose_no_pk_found(capsys): """Test 'detect_from_dataframes' verbose output when no PK found.""" # Setup From 59bbc150a2afc6de24f988a1cc142c6e103944c9 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 10 Sep 2026 14:50:37 +0100 Subject: [PATCH 06/10] use .agg to compute min and max --- sdv/metadata/_single_table.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 781f700f2..27f5a763d 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -820,8 +820,9 @@ def _detect_ranges(self, data): continue if sdtype == 'numerical': - column_metadata['range_min'] = clean_data.min().item() - column_metadata['range_max'] = clean_data.max().item() + ranges = clean_data.agg(['min', 'max']).to_dict() + column_metadata['range_min'] = ranges['min'] + column_metadata['range_max'] = ranges['max'] digits = learn_rounding_digits(column_data) column_metadata['decimal_places'] = digits if digits is not None else MAX_DECIMALS From b2719a567e00187baca2fd798cf923608acef16f Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Mon, 14 Sep 2026 10:56:59 +0100 Subject: [PATCH 07/10] remove Metadata.detect_from_dataframe() --- sdv/metadata/metadata.py | 50 ------------------- .../integration/cag/test_one_hot_encoding.py | 4 +- tests/integration/metadata/test_metadata.py | 14 +++--- tests/unit/metadata/test_metadata.py | 24 +++------ tests/unit/single_table/test_dayz.py | 6 +-- 5 files changed, 18 insertions(+), 80 deletions(-) diff --git a/sdv/metadata/metadata.py b/sdv/metadata/metadata.py index 6da504383..039a53bf2 100644 --- a/sdv/metadata/metadata.py +++ b/sdv/metadata/metadata.py @@ -21,7 +21,6 @@ _is_numerical, _load_data_from_csv, _validate_boolean_parameter, - _validate_data_single_table, ) from sdv.errors import InvalidDataError from sdv.logging import get_sdv_logger @@ -958,55 +957,6 @@ def _detect_from_dataframe( metadata.detect_table_from_dataframe(table_name, data, infer_sdtypes, infer_keys, verbose) return metadata - @classmethod - def detect_from_dataframe( - cls, - data, - table_name, - infer_sdtypes=True, - infer_keys='primary_only', - verbose=False, - ): - """Detect the metadata for a DataFrame. - - This method automatically detects the ``sdtypes`` for the given ``pandas.DataFrame``. - All data column names are converted to strings. - - Args: - data (dict[str, pd.DataFrame]): - The data to detect metadata from. - table_name (str): - The name of the table to detect. If None, a default name will be used. - Defaults to None. - infer_sdtypes (bool): - A boolean describing whether to infer the sdtypes of each column. - If True it infers the sdtypes based on the data. - If False it does not infer the sdtypes and all columns are marked as unknown. - Defaults to True. - infer_keys (str): - A string describing whether to infer the primary keys. Options are: - - 'primary_only': Infer only the primary keys of each table - - None: Do not infer any keys - Defaults to 'primary_only'. - verbose (bool): - A boolean that determines if information should be printed regarding detection. - If True, it prints out information about what is detected. - If False, it does not print out any information about what is detected. - Defaults to False. - - Returns: - Metadata: - A new metadata object with the sdtypes detected from the data. - """ - _validate_data_single_table(data, table_name) - return cls._detect_from_dataframe( - data=data[table_name], - table_name=table_name, - infer_sdtypes=infer_sdtypes, - infer_keys=infer_keys, - verbose=verbose, - ) - def set_primary_key(self, column_name, table_name=None): """Set the primary key of a table. diff --git a/tests/integration/cag/test_one_hot_encoding.py b/tests/integration/cag/test_one_hot_encoding.py index 11d5275c5..742f0fedc 100644 --- a/tests/integration/cag/test_one_hot_encoding.py +++ b/tests/integration/cag/test_one_hot_encoding.py @@ -163,7 +163,7 @@ def test_end_to_end_numerical_and_categorical(): # Setup metadata data = {'one_hot': df} - metadata = Metadata.detect_from_dataframe(data, table_name='one_hot') + metadata = Metadata.detect_from_dataframes(data) for sdtype in ['numerical', 'categorical']: metadata.update_columns(columns, sdtype=sdtype) synthesizer = GaussianCopulaSynthesizer(metadata) @@ -195,7 +195,7 @@ def test_end_to_end_boolean(): # Setup metadata data = {'one_hot': df} - metadata = Metadata.detect_from_dataframe(data, table_name='one_hot') + metadata = Metadata.detect_from_dataframes(data) metadata.update_columns(columns, sdtype='boolean') synthesizer = GaussianCopulaSynthesizer(metadata) constraint = OneHotEncoding(column_names=columns) diff --git a/tests/integration/metadata/test_metadata.py b/tests/integration/metadata/test_metadata.py index f89a65594..1bb10536f 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -387,7 +387,7 @@ def test_detect_from_dataframe(): data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') data = {'table': data['hotels']} - metadata = Metadata.detect_from_dataframe(data, 'table') + metadata = Metadata.detect_from_dataframes(data) # Run metadata.validate() @@ -418,7 +418,7 @@ def test_detect_from_dataframe_infer_sdtypes_false(): # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') data = {'table': data['hotels']} - metadata = Metadata.detect_from_dataframe(data, 'table', infer_sdtypes=False) + metadata = Metadata.detect_from_dataframes(data, infer_sdtypes=False) # Run metadata.validate() @@ -449,7 +449,7 @@ def test_detect_from_dataframe_infer_keys_none(): # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') data = {'table': data['hotels']} - metadata = Metadata.detect_from_dataframe(data, 'table', infer_keys=None) + metadata = Metadata.detect_from_dataframes(data, infer_keys=None) # Run metadata.validate() @@ -479,7 +479,7 @@ def test_detect_from_dataframe_infer_keys_none_infer_sdtypes_false(): # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') data = {'table': data['hotels']} - metadata = Metadata.detect_from_dataframe(data, 'table', infer_keys=None, infer_sdtypes=False) + metadata = Metadata.detect_from_dataframes(data, infer_keys=None, infer_sdtypes=False) # Run metadata.validate() @@ -1567,13 +1567,13 @@ def test_metadata_fails_with_proper_message_when_setting_primary_key(): def test_detect_from_dataframe_verbose_single(capsys): - """Test 'detect_from_dataframe' with verbose True with single table.""" + """Test 'detect_from_dataframes' with verbose True with single table.""" # Setup data, _ = download_test_demo(modality='single_table', dataset_name='fake_hotel_guests') data = {'table': data['fake_hotel_guests']} # Run - metadata = Metadata.detect_from_dataframe(data, 'table', verbose=True) + metadata = Metadata.detect_from_dataframes(data, verbose=True) # Assert captured = capsys.readouterr().out @@ -1611,7 +1611,7 @@ def test_detect_from_dataframe_verbose_single(capsys): def test_detect_from_dataframes_verbose(capsys): - """Test 'detect_from_dataframe' with verbose True with multi table.""" + """Test 'detect_from_dataframes' with verbose True with multi table.""" # Setup data, _ = download_test_demo(modality='multi_table', dataset_name='fake_hotels') diff --git a/tests/unit/metadata/test_metadata.py b/tests/unit/metadata/test_metadata.py index d7909bd94..9bc80a496 100644 --- a/tests/unit/metadata/test_metadata.py +++ b/tests/unit/metadata/test_metadata.py @@ -4899,26 +4899,14 @@ def test__detect_from_dataframe_bad_input_infer_keys(self): with pytest.raises(ValueError, match=expected_message): Metadata._detect_from_dataframe(data, infer_keys=infer_keys) - @patch.object(Metadata, '_detect_from_dataframe') - @patch('sdv.metadata.metadata._validate_data_single_table') - def test_detect_from_dataframe(self, mock_validate, mock_detect): - """Test the `detect_from_dataframe` method.""" + def test_detect_from_dataframe_raise_error(self): + """Test the `detect_from_dataframe` method raises an AttributeError.""" # Setup - data = {'table': pd.DataFrame()} - - # Run - metadata = Metadata.detect_from_dataframe(data, 'table') + data = pd.DataFrame() - # Assert - mock_detect.assert_called_once_with( - data=data['table'], - table_name='table', - infer_sdtypes=True, - infer_keys='primary_only', - verbose=False, - ) - mock_validate.assert_called_once_with(data, 'table') - assert metadata == mock_detect.return_value + # Run and Assert + with pytest.raises(AttributeError): + Metadata.detect_from_dataframe(data) def test__handle_table_name(self): """Test the ``_handle_table_name`` method.""" diff --git a/tests/unit/single_table/test_dayz.py b/tests/unit/single_table/test_dayz.py index d53b4c886..74b0b5948 100644 --- a/tests/unit/single_table/test_dayz.py +++ b/tests/unit/single_table/test_dayz.py @@ -562,7 +562,7 @@ def test_create_parameters_returns_valid_defaults(self): """Test create_parameters returns valid defaults.""" # Setup data = {'table': pd.DataFrame({'col': [np.nan]})} - metadata = Metadata.detect_from_dataframe(data, 'table') + metadata = Metadata.detect_from_dataframes(data) # Run params = DayZSynthesizer.create_parameters(data, metadata) @@ -584,7 +584,7 @@ def test_create_parameters_all_null_categorical_column(self): """Categorical column with all nulls should not have the category_values key parameter.""" # Setup data = {'table': pd.DataFrame({'col': [None, None, np.nan, pd.NA]})} - metadata = Metadata.detect_from_dataframe(data, 'table') + metadata = Metadata.detect_from_dataframes(data) # Run params = DayZSynthesizer.create_parameters(data, metadata) @@ -630,7 +630,7 @@ def test_create_parameters_all_null_datetime_column(self): """Datetime column with all nulls should omit start/end timestamps.""" # Setup data = {'table': pd.DataFrame({'col': pd.to_datetime([None, None])})} - metadata = Metadata.detect_from_dataframe(data, 'table') + metadata = Metadata.detect_from_dataframes(data) # Run params = DayZSynthesizer.create_parameters(data, metadata) From b63931bbb2650e203162a74c69ae57c2fc0eabc8 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Mon, 14 Sep 2026 11:19:08 +0100 Subject: [PATCH 08/10] rename detect_table_from_dataframe -> _detect_table_from_dataframe --- sdv/metadata/metadata.py | 8 +-- tests/integration/multi_table/test_hma.py | 18 +++--- tests/unit/evaluation/test__visualization.py | 8 +-- tests/unit/metadata/test_metadata.py | 64 ++++++++++---------- tests/unit/multi_table/test_base.py | 2 +- 5 files changed, 51 insertions(+), 49 deletions(-) diff --git a/sdv/metadata/metadata.py b/sdv/metadata/metadata.py index 039a53bf2..da7b3de01 100644 --- a/sdv/metadata/metadata.py +++ b/sdv/metadata/metadata.py @@ -785,7 +785,7 @@ def _detect_relationships( if foreign_key_inference_algorithm == 'column_name_match': self._detect_foreign_keys_by_column_name(data, verbose) - def detect_table_from_dataframe( + def _detect_table_from_dataframe( self, table_name, data, @@ -850,7 +850,7 @@ def _detect_from_dataframes( metadata = Metadata() for table_name, dataframe in data.items(): - metadata.detect_table_from_dataframe( + metadata._detect_table_from_dataframe( table_name, dataframe, infer_sdtypes, @@ -937,7 +937,7 @@ def detect_from_csvs(self, folder_name, read_csv_parameters=None): for csv_file in csv_files: table_name = csv_file.stem data[table_name] = _load_data_from_csv(csv_file, read_csv_parameters) - self.detect_table_from_dataframe(table_name, data[table_name]) + self._detect_table_from_dataframe(table_name, data[table_name]) self._detect_relationships(data) @@ -954,7 +954,7 @@ def _detect_from_dataframe( _validate_boolean_parameter(infer_sdtypes, 'infer_sdtypes') metadata = Metadata() - metadata.detect_table_from_dataframe(table_name, data, infer_sdtypes, infer_keys, verbose) + metadata._detect_table_from_dataframe(table_name, data, infer_sdtypes, infer_keys, verbose) return metadata def set_primary_key(self, column_name, table_name=None): diff --git a/tests/integration/multi_table/test_hma.py b/tests/integration/multi_table/test_hma.py index c9a1c1ddc..51dc94426 100644 --- a/tests/integration/multi_table/test_hma.py +++ b/tests/integration/multi_table/test_hma.py @@ -210,9 +210,9 @@ def get_custom_constraint_data_and_metadata(self): }) metadata = Metadata() - metadata.detect_table_from_dataframe('parent', parent_data) + metadata._detect_table_from_dataframe('parent', parent_data) metadata.update_column('primary_key', 'parent', sdtype='id') - metadata.detect_table_from_dataframe('child', child_data) + metadata._detect_table_from_dataframe('child', child_data) metadata.update_column('user_id', 'child', sdtype='id') metadata.update_column('id', 'child', sdtype='id') metadata.set_primary_key('primary_key', 'parent') @@ -304,9 +304,9 @@ def test_hma_with_inequality_constraint(self): data = {'parent_table': parent_table, 'child_table': child_table} metadata = Metadata() - metadata.detect_table_from_dataframe(table_name='parent_table', data=parent_table) + metadata._detect_table_from_dataframe(table_name='parent_table', data=parent_table) metadata.update_column('id', 'parent_table', sdtype='id') - metadata.detect_table_from_dataframe(table_name='child_table', data=child_table) + metadata._detect_table_from_dataframe(table_name='child_table', data=child_table) metadata.update_column('id', 'child_table', sdtype='id') metadata.update_column('parent_id', 'child_table', sdtype='id') @@ -391,7 +391,7 @@ def test_hma_primary_key_and_foreign_key_only(self): metadata = Metadata() for table_name, table in data.items(): - metadata.detect_table_from_dataframe(table_name, table) + metadata._detect_table_from_dataframe(table_name, table) metadata.update_column('user_id', 'users', sdtype='id') metadata.update_column('session_id', 'sessions', sdtype='id') @@ -527,8 +527,8 @@ def test_use_own_data_using_hma(self, tmp_path): # Metadata metadata = Metadata() - metadata.detect_table_from_dataframe(table_name='guests', data=datasets['guests']) - metadata.detect_table_from_dataframe(table_name='hotels', data=datasets['hotels']) + metadata._detect_table_from_dataframe(table_name='guests', data=datasets['guests']) + metadata._detect_table_from_dataframe(table_name='hotels', data=datasets['hotels']) # Assert - detected metadata correctly for table in metadata.tables: @@ -2312,8 +2312,8 @@ def test_detect_from_dataframe_numerical_col(): 'child_data': child_data, } metadata = Metadata() - metadata.detect_table_from_dataframe('parent_data', parent_data) - metadata.detect_table_from_dataframe('child_data', child_data) + metadata._detect_table_from_dataframe('parent_data', parent_data) + metadata._detect_table_from_dataframe('child_data', child_data) metadata.update_column('1', 'parent_data', sdtype='id') metadata.update_column('3', 'child_data', sdtype='id') metadata.update_column('4', 'child_data', sdtype='id') diff --git a/tests/unit/evaluation/test__visualization.py b/tests/unit/evaluation/test__visualization.py index bc6707462..2a12ae30a 100644 --- a/tests/unit/evaluation/test__visualization.py +++ b/tests/unit/evaluation/test__visualization.py @@ -712,7 +712,7 @@ def test_get_column_plot(mock_plot): data1 = {'table': table1} data2 = {'table': table2} metadata = Metadata() - metadata.detect_table_from_dataframe('table', table1) + metadata._detect_table_from_dataframe('table', table1) mock_plot.return_value = 'plot' # Run @@ -731,7 +731,7 @@ def test_get_column_plot_only_real_or_synthetic(mock_plot): table1 = pd.DataFrame({'col': [1, 2, 3]}) data1 = {'table': table1} metadata = Metadata() - metadata.detect_table_from_dataframe('table', table1) + metadata._detect_table_from_dataframe('table', table1) mock_plot.return_value = 'plot' # Run @@ -755,7 +755,7 @@ def test_get_column_pair_plot(mock_plot): data1 = {'table': table1} data2 = {'table': table2} metadata = Metadata() - metadata.detect_table_from_dataframe('table', table1) + metadata._detect_table_from_dataframe('table', table1) mock_plot.return_value = 'plot' # Run @@ -781,7 +781,7 @@ def test_get_column_pair_plot_only_real_or_synthetic(mock_plot): table1 = pd.DataFrame({'col1': [1, 2, 3], 'col2': [3, 2, 1]}) data1 = {'table': table1} metadata = Metadata() - metadata.detect_table_from_dataframe('table', table1) + metadata._detect_table_from_dataframe('table', table1) mock_plot.return_value = 'plot' # Run diff --git a/tests/unit/metadata/test_metadata.py b/tests/unit/metadata/test_metadata.py index 9bc80a496..171fe9cc8 100644 --- a/tests/unit/metadata/test_metadata.py +++ b/tests/unit/metadata/test_metadata.py @@ -790,10 +790,10 @@ def test_add_relationship_child_key_is_primary_key(self): # Setup table = pd.DataFrame({'pk': [1, 2, 3], 'col1': [0.1, 0.1, 0.2], 'col2': ['a', 'b', 'c']}) metadata = Metadata() - metadata.detect_table_from_dataframe('table', table) + metadata._detect_table_from_dataframe('table', table) metadata.update_column(column_name='pk', table_name='table', sdtype='id') metadata.set_primary_key(column_name='pk', table_name='table') - metadata.detect_table_from_dataframe('table2', table) + metadata._detect_table_from_dataframe('table2', table) metadata.update_column(column_name='pk', table_name='table2', sdtype='id') metadata.set_primary_key(column_name='pk', table_name='table2') @@ -1424,10 +1424,10 @@ def test_validate_child_key_is_primary_key(self): # Setup table = pd.DataFrame({'pk': [1, 2, 3], 'col1': [0.1, 0.1, 0.2], 'col2': ['a', 'b', 'c']}) metadata = Metadata() - metadata.detect_table_from_dataframe('table', table) + metadata._detect_table_from_dataframe('table', table) metadata.update_column(column_name='pk', table_name='table', sdtype='id') metadata.set_primary_key(column_name='pk', table_name='table') - metadata.detect_table_from_dataframe('table2', table) + metadata._detect_table_from_dataframe('table2', table) metadata.update_column(column_name='pk', table_name='table2', sdtype='id') metadata.set_primary_key(column_name='pk', table_name='table2') metadata.relationships = [ @@ -2997,7 +2997,7 @@ def test_detect_from_csvs(self, load_data_mock, tmp_path): """Test the ``detect_from_csvs`` method.""" # Setup instance = Metadata() - instance.detect_table_from_dataframe = Mock() + instance._detect_table_from_dataframe = Mock() instance._detect_relationships = Mock() data1 = pd.DataFrame({'col1': [1, 2], 'col2': [3, 4]}) @@ -3034,8 +3034,10 @@ def load_data_side_effect(filepath, _): call('table1', data1), call('table2', data2), ] - instance.detect_table_from_dataframe.assert_has_calls(expected_detect_calls, any_order=True) - assert instance.detect_table_from_dataframe.call_count == 2 + instance._detect_table_from_dataframe.assert_has_calls( + expected_detect_calls, any_order=True + ) + assert instance._detect_table_from_dataframe.call_count == 2 instance._detect_relationships.assert_called_once() table1 = instance._detect_relationships.call_args[0][0]['table1'] @@ -3064,7 +3066,7 @@ def test_detect_from_csvs_no_csv(self, tmp_path): @patch('sdv.metadata.metadata.LOGGER') @patch('sdv.metadata.metadata._SingleTableMetadata') def test_detect_table_from_dataframe(self, single_table_mock, log_mock): - """Test the ``detect_table_from_dataframe`` method. + """Test the ``_detect_table_from_dataframe`` method. If the table does not already exist, a ``_SingleTableMetadata`` instance should be created and call the ``detect_from_dataframe`` method. @@ -3083,7 +3085,7 @@ def test_detect_table_from_dataframe(self, single_table_mock, log_mock): } # Run - metadata.detect_table_from_dataframe('table', data) + metadata._detect_table_from_dataframe('table', data) # Assert single_table_mock.return_value._detect_columns.assert_called_once_with( @@ -3104,7 +3106,7 @@ def test_detect_table_from_dataframe(self, single_table_mock, log_mock): log_mock.info.assert_has_calls([expected_log_calls]) def test_detect_table_from_dataframe_table_already_exists(self): - """Test the ``detect_table_from_dataframe`` method. + """Test the ``_detect_table_from_dataframe`` method. If the table already exists, an error should be raised. @@ -3128,17 +3130,17 @@ def test_detect_table_from_dataframe_table_already_exists(self): 'create a new Metadata object for other data sources.' ) with pytest.raises(InvalidMetadataError, match=error_message): - metadata.detect_table_from_dataframe('table', pd.DataFrame()) + metadata._detect_table_from_dataframe('table', pd.DataFrame()) @patch('sdv.metadata.metadata.Metadata') def test_detect_from_dataframes(self, mock_metadata): """Test ``detect_from_dataframes``. - Expected to call ``detect_table_from_dataframe`` for each table name and dataframe + Expected to call ``_detect_table_from_dataframe`` for each table name and dataframe in the input. """ # Setup - mock_metadata.detect_table_from_dataframe = Mock() + mock_metadata._detect_table_from_dataframe = Mock() mock_metadata._detect_relationships = Mock() guests_table = pd.DataFrame() hotels_table = pd.DataFrame() @@ -3148,10 +3150,10 @@ def test_detect_from_dataframes(self, mock_metadata): metadata = Metadata.detect_from_dataframes(data) # Assert - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( 'guests', guests_table, True, 'primary_only', False ) - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( 'hotels', hotels_table, True, 'primary_only', False ) mock_metadata.return_value._detect_relationships.assert_called_once_with( @@ -4166,7 +4168,7 @@ def test__detect_relationships_verbose(self): @patch('sdv.metadata.metadata._SingleTableMetadata') def test_detect_table_from_dataframe_with_kwargs(self, single_table_mock): - """Test `detect_table_from_dataframe` fsets verbose on `_detect_columns`.""" + """Test `_detect_table_from_dataframe` fsets verbose on `_detect_columns`.""" # Setup metadata = Metadata() data = pd.DataFrame() @@ -4175,7 +4177,7 @@ def test_detect_table_from_dataframe_with_kwargs(self, single_table_mock): } # Run - metadata.detect_table_from_dataframe( + metadata._detect_table_from_dataframe( 'table', data, infer_sdtypes=False, infer_keys=None, verbose=True ) @@ -4684,7 +4686,7 @@ def test_validate_table(self): def test_detect_from_dataframes_infer_keys_none(self, mock_metadata): """Test ``detect_from_dataframes`` with infer_keys set to None.""" # Setup - mock_metadata.detect_table_from_dataframe = Mock() + mock_metadata._detect_table_from_dataframe = Mock() mock_metadata._detect_relationships = Mock() guests_table = pd.DataFrame() hotels_table = pd.DataFrame() @@ -4694,10 +4696,10 @@ def test_detect_from_dataframes_infer_keys_none(self, mock_metadata): metadata = Metadata.detect_from_dataframes(data, infer_sdtypes=False, infer_keys=None) # Assert - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( 'guests', guests_table, False, None, False ) - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( 'hotels', hotels_table, False, None, False ) mock_metadata.return_value._detect_relationships.assert_not_called() @@ -4724,7 +4726,7 @@ def test_detect_from_dataframes_bad_foreign_key_inference_algorithm(self): def test_detect_from_dataframes_infer_keys_primary_only(self, mock_metadata): """Test ``detect_from_dataframes`` with infer_keys set to 'primary_only'.""" # Setup - mock_metadata.detect_table_from_dataframe = Mock() + mock_metadata._detect_table_from_dataframe = Mock() mock_metadata._detect_relationships = Mock() guests_table = pd.DataFrame() hotels_table = pd.DataFrame() @@ -4736,10 +4738,10 @@ def test_detect_from_dataframes_infer_keys_primary_only(self, mock_metadata): ) # Assert - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( 'guests', guests_table, False, 'primary_only', False ) - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( 'hotels', hotels_table, False, 'primary_only', False ) mock_metadata.return_value._detect_relationships.assert_not_called() @@ -4794,8 +4796,8 @@ def test_detect_from_dataframe_primary_key_to_primary_key(self): }), } instance = Metadata() - instance.detect_table_from_dataframe('table1', data['table1']) - instance.detect_table_from_dataframe('table2', data['table2']) + instance._detect_table_from_dataframe('table1', data['table1']) + instance._detect_table_from_dataframe('table2', data['table2']) # Run instance._detect_foreign_keys_by_column_name(data) @@ -4825,9 +4827,9 @@ def test_detect_from_dataframe_primary_key_to_primary_key_no_cycles(self): }), } instance = Metadata() - instance.detect_table_from_dataframe('table1', data['table1']) - instance.detect_table_from_dataframe('table2', data['table2']) - instance.detect_table_from_dataframe('table3', data['table3']) + instance._detect_table_from_dataframe('table1', data['table1']) + instance._detect_table_from_dataframe('table2', data['table2']) + instance._detect_table_from_dataframe('table3', data['table3']) # Run instance._detect_foreign_keys_by_column_name(data) @@ -4855,17 +4857,17 @@ def test_detect_from_dataframe_primary_key_to_primary_key_no_cycles(self): def test__detect_from_dataframe(self, mock_metadata): """Test that the method calls the detection method and returns the metadata. - Expected to call ``detect_table_from_dataframe`` for the dataframe. + Expected to call ``_detect_table_from_dataframe`` for the dataframe. """ # Setup - mock_metadata.detect_table_from_dataframe = Mock() + mock_metadata._detect_table_from_dataframe = Mock() data = pd.DataFrame() # Run metadata = Metadata._detect_from_dataframe(data) # Assert - mock_metadata.return_value.detect_table_from_dataframe.assert_any_call( + mock_metadata.return_value._detect_table_from_dataframe.assert_any_call( Metadata.DEFAULT_SINGLE_TABLE_NAME, DataFrameMatcher(data), True, 'primary_only', False ) assert metadata == mock_metadata.return_value diff --git a/tests/unit/multi_table/test_base.py b/tests/unit/multi_table/test_base.py index fc270d762..80e3a4963 100644 --- a/tests/unit/multi_table/test_base.py +++ b/tests/unit/multi_table/test_base.py @@ -1634,7 +1634,7 @@ def test_add_constraintss_missing_table_name(self): 'col2': [4, 5, 6], }) metadata = Metadata() - metadata.detect_table_from_dataframe('table', data) + metadata._detect_table_from_dataframe('table', data) constraint = Inequality(low_column_name='col1', high_column_name='col2') model = BaseMultiTableSynthesizer(metadata) From fcc9a86dd10abf1ff5296fbe58b2f6a534a54a73 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Tue, 15 Sep 2026 15:10:20 +0100 Subject: [PATCH 09/10] add metadata validation --- tests/integration/metadata/test_metadata.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/tests/integration/metadata/test_metadata.py b/tests/integration/metadata/test_metadata.py index 1bb10536f..fd02cde9f 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -119,6 +119,7 @@ def test_detect_from_dataframes_multi_table(): } compare_metadata(metadata, expected_metadata) compare_ranges(metadata, real_data) + metadata.validate_data(real_data) def test_detect_from_dataframes_multi_table_without_infer_sdtypes(): @@ -170,6 +171,7 @@ def test_detect_from_dataframes_multi_table_without_infer_sdtypes(): } compare_metadata(metadata, expected_metadata) compare_ranges(metadata, real_data) + metadata.validate_data(real_data) def test_detect_from_dataframes_multi_table_with_infer_keys_primary_only(): @@ -214,6 +216,7 @@ def test_detect_from_dataframes_multi_table_with_infer_keys_primary_only(): } compare_metadata(metadata, expected_metadata) compare_ranges(metadata, real_data) + metadata.validate_data(real_data) def test_detect_from_dataframes_multi_table_with_infer_keys_none(): @@ -256,6 +259,7 @@ def test_detect_from_dataframes_multi_table_with_infer_keys_none(): } compare_metadata(metadata, expected_metadata) compare_ranges(metadata, real_data) + metadata.validate_data(real_data) def test_detect_from_dataframes_single_table(): @@ -287,6 +291,7 @@ def test_detect_from_dataframes_single_table(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_dataframes_single_table_infer_sdtypes_false(): @@ -318,6 +323,7 @@ def test_detect_from_dataframes_single_table_infer_sdtypes_false(): } compare_metadata(metadata, expected_metadata) compare_ranges(metadata, data) + metadata.validate_data(data) def test_detect_from_dataframes_single_table_infer_keys_primary_only(): @@ -349,6 +355,7 @@ def test_detect_from_dataframes_single_table_infer_keys_primary_only(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_dataframes_single_table_infer_keys_none(): @@ -379,6 +386,7 @@ def test_detect_from_dataframes_single_table_infer_keys_none(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_dataframe(): @@ -411,6 +419,7 @@ def test_detect_from_dataframe(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_dataframe_infer_sdtypes_false(): @@ -442,6 +451,7 @@ def test_detect_from_dataframe_infer_sdtypes_false(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_dataframe_infer_keys_none(): @@ -472,6 +482,7 @@ def test_detect_from_dataframe_infer_keys_none(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_dataframe_infer_keys_none_infer_sdtypes_false(): @@ -502,6 +513,7 @@ def test_detect_from_dataframe_infer_keys_none_infer_sdtypes_false(): } compare_ranges(metadata, data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(data) def test_detect_from_csvs(tmp_path): @@ -560,6 +572,7 @@ def test_detect_from_csvs(tmp_path): compare_ranges(metadata, real_data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(real_data) params = [ @@ -2063,6 +2076,7 @@ def test_detect_from_dataframes(): } compare_ranges(metadata, real_data) compare_metadata(metadata, expected_metadata) + metadata.validate_data(real_data) def test_get_column_names(): From ee264d41afb44aa7b6407ca88270336be1e16d13 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Wed, 16 Sep 2026 17:40:25 +0100 Subject: [PATCH 10/10] add docstring ordinal logic --- sdv/metadata/_single_table.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 27f5a763d..6a7e2df9d 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -577,8 +577,10 @@ def _detect_id_column(self, column_name): def _detect_ordinal_sdtype(self, data): """Detect whether a numerical column should have the ordinal sdtype. - A numerical column is considered ordinal when it contains whole numbers - and has low cardinality. + A numerical column is considered ordinal if: + - It contains only whole numbers + - It has low cardinality, defined as having at most 10% unique values relative + to the total number of rows, capped at 10 unique values. Args: data (pandas.Series): @@ -593,8 +595,8 @@ def _detect_ordinal_sdtype(self, data): whole_values = (clean_data == clean_data.round()).all() unique_values = clean_data.nunique() - categorical_threshold = min(round(len(data) / 10), 10) - low_cardinality = unique_values <= categorical_threshold + ordinal_threshold = min(round(len(data) / 10), 10) + low_cardinality = unique_values <= ordinal_threshold if whole_values and low_cardinality: return 'ordinal' @@ -1565,7 +1567,7 @@ def _validate_column_data(self, column, sdtype_warnings): if decimal_places is not None: column_values = column.dropna() data_digits = learn_rounding_digits(column_values) - if data_digits > decimal_places: + if data_digits is not None and data_digits > decimal_places: errors += [ f"Values found for numerical column '{column.name}' exceed the allowed " f'decimal places ({decimal_places}).'