Coverage for src/utils/dataframe_helpers.py: 14%

134 statements  

« prev     ^ index     » next       coverage.py v7.8.2, created at 2025-05-25 13:51 -0400

1# Copyright (c) CloudZero - ALL RIGHTS RESERVED - PROPRIETARY AND CONFIDENTIAL 

2# Unauthorized copying of this file and/or project, via any medium is strictly prohibited. 

3# Direct all questions to legal@cloudzero.com 

4 

5""" 

6DataFrame utilities for common data processing operations. 

7Provides standardized DataFrame manipulation patterns. 

8""" 

9 

10import polars as pl 

11from typing import List, Dict, Any, Optional, Tuple 

12from .console_helpers import print_progress 

13 

14 

15def get_column_sample_values(df: pl.DataFrame, column: str, sample_size: int = 5) -> List[Any]: 

16 """Get sample values from a column, excluding nulls.""" 

17 non_null_data = df[column].drop_nulls() 

18 if len(non_null_data) == 0: 

19 return [] 

20 return non_null_data.head(sample_size).to_list() 

21 

22 

23def calculate_cardinality_stats(df: pl.DataFrame, column: str) -> Dict[str, Any]: 

24 """Calculate comprehensive cardinality statistics for a column.""" 

25 col_data = df[column] 

26 total_count = len(col_data) 

27 null_count = col_data.null_count() 

28 non_null_count = total_count - null_count 

29 

30 if non_null_count > 0: 

31 non_null_data = col_data.drop_nulls() 

32 unique_count = non_null_data.n_unique() 

33 cardinality_ratio = unique_count / non_null_count 

34 else: 

35 unique_count = 0 

36 cardinality_ratio = 0.0 

37 

38 # Classify cardinality level 

39 if cardinality_ratio >= 0.95: 

40 cardinality_level = "very_high" 

41 elif cardinality_ratio >= 0.7: 

42 cardinality_level = "high" 

43 elif cardinality_ratio >= 0.1: 

44 cardinality_level = "medium" 

45 else: 

46 cardinality_level = "low" 

47 

48 return { 

49 "total_count": total_count, 

50 "null_count": null_count, 

51 "non_null_count": non_null_count, 

52 "unique_count": unique_count, 

53 "cardinality_ratio": cardinality_ratio, 

54 "cardinality": cardinality_level 

55 } 

56 

57 

58def get_value_counts(df: pl.DataFrame, column: str, top_n: int = 10) -> Optional[Dict[str, int]]: 

59 """Get value counts for a column, limited to top N values.""" 

60 try: 

61 col_data = df[column].drop_nulls() 

62 if len(col_data) == 0: 

63 return None 

64 

65 value_counts = col_data.value_counts().head(top_n) 

66 return {str(row[0]): row[1] for row in value_counts.iter_rows()} 

67 except Exception: 

68 return None 

69 

70 

71def filter_valid_rows(df: pl.DataFrame, required_columns: List[str]) -> pl.DataFrame: 

72 """Filter DataFrame to only include rows with non-null values in required columns.""" 

73 if not required_columns: 

74 return df 

75 

76 # Create filter expressions for each required column 

77 filter_expressions = [pl.col(col).is_not_null() for col in required_columns] 

78 

79 # Combine all filters with AND logic 

80 combined_filter = filter_expressions[0] 

81 for expr in filter_expressions[1:]: 

82 combined_filter = combined_filter & expr 

83 

84 return df.filter(combined_filter) 

85 

86 

87def create_composite_column( 

88 df: pl.DataFrame, 

89 columns: List[str], 

90 separator: str = "|", 

91 new_column_name: Optional[str] = None 

92) -> Tuple[pl.DataFrame, str]: 

93 """Create a composite column by concatenating multiple columns.""" 

94 if new_column_name is None: 

95 new_column_name = f"composite_{'_'.join(columns)}" 

96 

97 # Create composite expression 

98 composite_expr = pl.concat_str( 

99 [pl.col(col).cast(pl.Utf8) for col in columns], 

100 separator=separator, 

101 ignore_nulls=False 

102 ) 

103 

104 # Add composite column to DataFrame 

105 df_with_composite = df.with_columns(composite_expr.alias(new_column_name)) 

106 

107 return df_with_composite, new_column_name 

108 

109 

110def analyze_null_patterns(df: pl.DataFrame, columns: Optional[List[str]] = None) -> Dict[str, int]: 

111 """Analyze null patterns across specified columns.""" 

112 if columns is None: 

113 columns = df.columns 

114 

115 null_counts = {} 

116 for col in columns: 

117 if col in df.columns: 

118 null_count = df[col].null_count() 

119 if null_count > 0: 

120 null_counts[col] = null_count 

121 

122 return null_counts 

123 

124 

125def get_data_quality_summary(df: pl.DataFrame) -> Dict[str, Any]: 

126 """Get comprehensive data quality summary.""" 

127 total_rows = len(df) 

128 total_cols = len(df.columns) 

129 

130 # Calculate null statistics 

131 null_stats = {} 

132 for col in df.columns: 

133 null_count = df[col].null_count() 

134 null_stats[col] = { 

135 "null_count": null_count, 

136 "null_percentage": (null_count / total_rows) * 100 if total_rows > 0 else 0 

137 } 

138 

139 return { 

140 "total_rows": total_rows, 

141 "total_columns": total_cols, 

142 "null_statistics": null_stats, 

143 "columns_with_nulls": sum(1 for stats in null_stats.values() if stats["null_count"] > 0), 

144 "total_null_cells": sum(stats["null_count"] for stats in null_stats.values()) 

145 } 

146 

147 

148def sample_dataframe(df: pl.DataFrame, sample_size: int, seed: int = 42) -> pl.DataFrame: 

149 """Sample DataFrame with validation.""" 

150 if sample_size >= len(df): 

151 return df 

152 

153 return df.sample(n=sample_size, seed=seed) 

154 

155 

156def detect_datetime_columns(df: pl.DataFrame) -> List[str]: 

157 """Detect columns that contain datetime data.""" 

158 datetime_columns = [] 

159 

160 for col in df.columns: 

161 col_type = df[col].dtype 

162 # Check if already datetime/date type 

163 if col_type in [pl.Datetime, pl.Date]: 

164 datetime_columns.append(col) 

165 # Check string columns for datetime patterns 

166 elif col_type == pl.Utf8: 

167 sample_values = get_column_sample_values(df, col, 10) 

168 if _contains_datetime_patterns(sample_values): 

169 datetime_columns.append(col) 

170 

171 return datetime_columns 

172 

173 

174def detect_numeric_columns(df: pl.DataFrame) -> List[str]: 

175 """Detect columns that contain numeric data.""" 

176 numeric_types = { 

177 pl.Int8, pl.Int16, pl.Int32, pl.Int64, 

178 pl.UInt8, pl.UInt16, pl.UInt32, pl.UInt64, 

179 pl.Float32, pl.Float64 

180 } 

181 

182 return [col for col in df.columns if df[col].dtype in numeric_types] 

183 

184 

185def detect_string_columns(df: pl.DataFrame) -> List[str]: 

186 """Detect columns that contain string data.""" 

187 return [col for col in df.columns if df[col].dtype == pl.Utf8] 

188 

189 

190def classify_column_cardinality(df: pl.DataFrame, columns: Optional[List[str]] = None) -> Dict[str, str]: 

191 """Classify cardinality level for multiple columns.""" 

192 if columns is None: 

193 columns = df.columns 

194 

195 cardinality_classification = {} 

196 for col in columns: 

197 if col in df.columns: 

198 stats = calculate_cardinality_stats(df, col) 

199 cardinality_classification[col] = stats["cardinality"] 

200 

201 return cardinality_classification 

202 

203 

204def get_column_summary(df: pl.DataFrame, column: str) -> Dict[str, Any]: 

205 """Get comprehensive summary for a single column.""" 

206 if column not in df.columns: 

207 raise ValueError(f"Column '{column}' not found in DataFrame") 

208 

209 col_data = df[column] 

210 col_type = col_data.dtype 

211 

212 # Basic stats 

213 stats = calculate_cardinality_stats(df, column) 

214 

215 # Sample values 

216 sample_values = get_column_sample_values(df, column, 5) 

217 

218 # Value counts for low cardinality columns 

219 value_counts = None 

220 if stats["cardinality"] in ["low", "medium"]: 

221 value_counts = get_value_counts(df, column, 10) 

222 

223 return { 

224 "column": column, 

225 "dtype": str(col_type), 

226 "sample_values": sample_values, 

227 "value_counts": value_counts, 

228 **stats 

229 } 

230 

231 

232def _contains_datetime_patterns(values: List[Any]) -> bool: 

233 """Check if values contain datetime patterns.""" 

234 import re 

235 

236 datetime_patterns = [ 

237 r'\d{4}-\d{2}-\d{2}', # YYYY-MM-DD 

238 r'\d{2}/\d{2}/\d{4}', # MM/DD/YYYY 

239 r'\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2}', # YYYY-MM-DD HH:MM:SS 

240 ] 

241 

242 for value in values[:5]: # Check first 5 values 

243 if value is None: 

244 continue 

245 value_str = str(value) 

246 for pattern in datetime_patterns: 

247 if re.search(pattern, value_str): 

248 return True 

249 

250 return False 

251 

252 

253def print_dataframe_info(df: pl.DataFrame, filename: str = "DataFrame") -> None: 

254 """Print standardized DataFrame information.""" 

255 rows, cols = df.shape 

256 print_progress(f"{filename}: {rows:,} rows, {cols} columns") 

257 

258 # Print null summary if there are nulls 

259 null_counts = analyze_null_patterns(df) 

260 if null_counts: 

261 total_nulls = sum(null_counts.values()) 

262 print_progress(f"Null values found: {total_nulls:,} across {len(null_counts)} columns") 

263 

264 

265def ensure_column_types(df: pl.DataFrame, type_mapping: Dict[str, pl.DataType]) -> pl.DataFrame: 

266 """Ensure columns have the specified types.""" 

267 cast_expressions = [] 

268 

269 for col, target_type in type_mapping.items(): 

270 if col in df.columns and df[col].dtype != target_type: 

271 cast_expressions.append(pl.col(col).cast(target_type)) 

272 

273 if cast_expressions: 

274 return df.with_columns(cast_expressions) 

275 

276 return df