1919 write_to_app_db ,
2020)
2121from testgen .common .database .column_chars import ColumnChars
22- from testgen .common .database .database_service import ThreadedProgress
22+ from testgen .common .database .database_service import ThreadedProgress , get_flavor_service
2323from testgen .common .job_context import job_context
2424from testgen .common .mixpanel_service import MixpanelService
2525from testgen .common .models import get_current_session , with_database_session
@@ -86,8 +86,9 @@ def run_profiling(
8686 if data_chars :
8787 sql_generator = ProfilingSQL (connection , table_group , profiling_run )
8888
89- _run_column_profiling (sql_generator , data_chars )
90- _run_frequency_analysis (sql_generator )
89+ sampling_params = _compute_sampling_params (sql_generator , data_chars )
90+ _run_column_profiling (sql_generator , data_chars , sampling_params )
91+ _run_frequency_analysis (sql_generator , sampling_params )
9192 _run_hygiene_issue_detection (sql_generator )
9293
9394 # if table_group.profile_do_pair_rules == "Y":
@@ -143,26 +144,40 @@ def _exclude_xde_columns(data_chars: list[ColumnChars], table_group_id: UUID) ->
143144 return filtered
144145
145146
146- def _run_column_profiling (sql_generator : ProfilingSQL , data_chars : list [ColumnChars ]) -> None :
147+ def _compute_sampling_params (
148+ sql_generator : ProfilingSQL , data_chars : list [ColumnChars ]
149+ ) -> dict [str , TableSampling ]:
150+ table_group = sql_generator .table_group
151+ sampling_params : dict [str , TableSampling ] = {}
152+ if not table_group .profile_use_sampling :
153+ return sampling_params
154+
155+ sampleable_types = get_flavor_service (sql_generator .flavor ).sampleable_object_types
156+ for column in data_chars :
157+ if sampling_params .get (column .table_name ):
158+ continue
159+ if sampleable_types is not None and column .object_type not in sampleable_types :
160+ continue
161+ result = calculate_sampling_params (
162+ table_name = column .table_name ,
163+ record_count = column .record_ct ,
164+ sample_percent_raw = table_group .profile_sample_percent ,
165+ min_sample = table_group .profile_sample_min_count ,
166+ )
167+ if result :
168+ sampling_params [column .table_name ] = result
169+ return sampling_params
170+
171+
172+ def _run_column_profiling (
173+ sql_generator : ProfilingSQL , data_chars : list [ColumnChars ], sampling_params : dict [str , TableSampling ]
174+ ) -> None :
147175 profiling_run = sql_generator .profiling_run
148176 profiling_run .set_progress ("col_profiling" , "Running" )
149177 profiling_run .save ()
150178 get_current_session ().commit ()
151179
152180 LOG .info (f"Running column profiling queries: { len (data_chars )} " )
153- table_group = sql_generator .table_group
154- sampling_params : dict [str , TableSampling ] = {}
155- if table_group .profile_use_sampling :
156- for column in data_chars :
157- if not sampling_params .get (column .table_name ):
158- result = calculate_sampling_params (
159- table_name = column .table_name ,
160- record_count = column .record_ct ,
161- sample_percent_raw = table_group .profile_sample_percent ,
162- min_sample = table_group .profile_sample_min_count ,
163- )
164- if result :
165- sampling_params [column .table_name ] = result
166181
167182 def update_column_progress (progress : ThreadedProgress ) -> None :
168183 profiling_run .set_progress (
@@ -218,7 +233,7 @@ def update_column_progress(progress: ThreadedProgress) -> None:
218233 )
219234
220235
221- def _run_frequency_analysis (sql_generator : ProfilingSQL ) -> None :
236+ def _run_frequency_analysis (sql_generator : ProfilingSQL , sampling_params : dict [ str , TableSampling ] ) -> None :
222237 profiling_run = sql_generator .profiling_run
223238 profiling_run .set_progress ("freq_analysis" , "Running" )
224239 profiling_run .save ()
@@ -227,7 +242,7 @@ def _run_frequency_analysis(sql_generator: ProfilingSQL) -> None:
227242 error_data = None
228243 try :
229244 LOG .info ("Selecting columns for frequency analysis" )
230- frequency_columns = fetch_dict_from_db (* sql_generator .get_frequency_analysis_columns ())
245+ frequency_columns = [ ColumnChars ( ** column ) for column in fetch_dict_from_db (* sql_generator .get_frequency_analysis_columns ())]
231246
232247 if frequency_columns :
233248 LOG .info (f"Running frequency analysis queries: { len (frequency_columns )} " )
@@ -240,7 +255,10 @@ def update_frequency_progress(progress: ThreadedProgress) -> None:
240255 get_current_session ().commit ()
241256
242257 frequency_results , result_columns , error_data = fetch_from_db_threaded (
243- [sql_generator .run_frequency_analysis (ColumnChars (** column )) for column in frequency_columns ],
258+ [
259+ sql_generator .run_frequency_analysis (column , sampling_params .get (column .table_name ))
260+ for column in frequency_columns
261+ ],
244262 use_target_db = True ,
245263 max_threads = sql_generator .connection .max_threads ,
246264 progress_callback = update_frequency_progress ,
0 commit comments