from __future__ import annotations import json from dataclasses import dataclass, field from pathlib import Path from typing import Any import duckdb import numpy as np import pandas as pd import sqlglot from sqlglot import exp # Convert pandas/NumPy values into JSON-serializable types def _json_default(value: Any): if isinstance(value, np.integer): return int(value) if isinstance(value, np.floating): return float(value) if isinstance(value, np.bool_): return bool(value) if isinstance(value, pd.Timestamp): return value.isoformat() if pd.isna(value): return None return str(value) # Load supported data files into a pandas DataFrame def _read_dataframe(path: Path, max_rows: int | None = None) -> pd.DataFrame: suffix = path.suffix.lower() if suffix == ".csv": return pd.read_csv(path, nrows=max_rows, low_memory=False) if suffix == ".parquet": df = pd.read_parquet(path); return df.head(max_rows) if max_rows else df if suffix == ".json": try: df = pd.read_json(path) except ValueError: df = pd.read_json(path, lines=True) return df.head(max_rows) if max_rows else df if suffix == ".jsonl": df = pd.read_json(path, lines=True); return df.head(max_rows) if max_rows else df if suffix in {".xlsx", ".xls"}: df = pd.read_excel(path); return df.head(max_rows) if max_rows else df raise ValueError(f"Unsupported file type: {suffix}. Use CSV, Parquet, JSON/JSONL, XLSX, or XLS.") # Map pandas dtypes to simplified logical data types def _logical_type(series: pd.Series) -> str: if pd.api.types.is_bool_dtype(series): return "boolean" if pd.api.types.is_integer_dtype(series): return "integer" if pd.api.types.is_float_dtype(series): return "float" if pd.api.types.is_datetime64_any_dtype(series): return "datetime" if pd.api.types.is_numeric_dtype(series): return "number" return "string" # Return a small set of JSON-safe example values from a column def _safe_examples(series: pd.Series, limit: int = 3) -> list[Any]: values = series.dropna().head(limit).tolist() return [json.loads(json.dumps(v, default=_json_default)) for v in values] # Holds the loaded dataset and cached analysis results @dataclass class DataContext: dataframe: pd.DataFrame source_name: str source_path: str | None = None profile_cache: dict[str, Any] | None = field(default=None, repr=False) quality_cache: dict[str, Any] | None = field(default=None, repr=False) # Create a data context directly from a supported file @classmethod def from_path(cls, path: str | Path, max_rows: int | None = None) -> "DataContext": path = Path(path) return cls(_read_dataframe(path, max_rows=max_rows), path.name, str(path)) # Generate dataset-level and column-level profile information def profile(self) -> dict[str, Any]: if self.profile_cache is not None: return self.profile_cache df = self.dataframe columns = [] for column in df.columns: s = df[column]; nulls = int(s.isna().sum()); rows = int(len(df)) columns.append({ "name": str(column), "pandas_dtype": str(s.dtype), "logical_type": _logical_type(s), "nullable": bool(nulls > 0), "null_count": nulls, "null_pct": round(nulls / rows * 100, 2) if rows else 0.0, "unique_count": int(s.nunique(dropna=True)), "examples": _safe_examples(s), }) self.profile_cache = { "source": self.source_name, "rows": int(len(df)), "columns": int(len(df.columns)), "memory_mb": round(float(df.memory_usage(deep=True).sum()/1024/1024), 3), "duplicate_rows": int(df.duplicated().sum()), "column_schema": columns, } return self.profile_cache # Run built-in checks for normal data-quality issues def quality_report(self) -> dict[str, Any]: if self.quality_cache is not None: return self.quality_cache df = self.dataframe; rows = max(len(df), 1); issues = [] # Check for duplicate rows dup = int(df.duplicated().sum()) if dup: issues.append({"severity":"high" if dup/rows>=.05 else "medium","check":"duplicate_rows","column":None, "evidence":f"{dup:,} duplicate rows ({dup/rows*100:.2f}%).", "recommended_fix":"Define a business key and deduplicate with a deterministic ordering rule."}) # Run column-level quality checks for column in df.columns: s = df[column]; nulls = int(s.isna().sum()); pct = nulls/rows*100; non_null = s.dropna() if nulls: sev = "high" if pct>=20 else "medium" if pct>=5 else "low" issues.append({"severity":sev,"check":"missing_values","column":str(column),"evidence":f"{nulls:,} nulls ({pct:.2f}%).", "recommended_fix":"Confirm nullability. Impute, backfill, quarantine, or reject rows based on business meaning."}) if len(non_null) and non_null.nunique(dropna=True)==1: issues.append({"severity":"low","check":"constant_column","column":str(column),"evidence":"Only one non-null value is present.", "recommended_fix":"Remove it if it carries no signal, unless intentionally constant."}) # Additional checks for string-like columns if pd.api.types.is_object_dtype(s) or pd.api.types.is_string_dtype(s): ss = non_null.astype(str); empty = int(ss.str.strip().eq("").sum()) if empty: issues.append({"severity":"medium","check":"blank_strings","column":str(column),"evidence":f"{empty:,} blank or whitespace-only values.", "recommended_fix":"Trim strings and normalize blanks to NULL."}) types = {type(v).__name__ for v in non_null.head(500)} if len(types)>1: issues.append({"severity":"medium","check":"mixed_python_types","column":str(column),"evidence":f"Mixed value types: {sorted(types)}.", "recommended_fix":"Cast to one canonical type and quarantine failed casts."}) if len(ss): ratio = float(pd.to_numeric(ss, errors="coerce").notna().mean()) if .9 <= ratio < 1: issues.append({"severity":"medium","check":"numeric_cast_failures","column":str(column), "evidence":f"{ratio*100:.2f}% of non-null values parse as numeric.", "recommended_fix":"Normalize formatting, safe-cast, and quarantine failures."}) # Boolean columns can be reported as numeric-like by pandas, but # NumPy quantile interpolation does not support boolean subtraction. # Skip booleans for IQR-based numeric outlier detection. if pd.api.types.is_numeric_dtype(s) and not pd.api.types.is_bool_dtype(s) and len(non_null)>=8: q1,q3=float(non_null.quantile(.25)),float(non_null.quantile(.75)); iqr=q3-q1 if iqr>0: lo,hi=q1-3*iqr,q3+3*iqr; n=int(((non_nullhi)).sum()) if n: issues.append({"severity":"low","check":"extreme_numeric_values","column":str(column), "evidence":f"{n:,} values outside 3×IQR bounds [{lo:.3g}, {hi:.3g}].", "recommended_fix":"Verify business validity before clipping or excluding."}) # Sort issues by severity before caching the report rank={"high":0,"medium":1,"low":2}; issues.sort(key=lambda x:(rank.get(x["severity"],9),x.get("column") or "")) self.quality_cache={"source":self.source_name,"issue_count":len(issues),"high":sum(i["severity"]=="high" for i in issues), "medium":sum(i["severity"]=="medium" for i in issues),"low":sum(i["severity"]=="low" for i in issues),"issues":issues[:150]} return self.quality_cache # Format the quality report as a Markdown table def quality_markdown(self) -> str: q=self.quality_report(); lines=["### Data-quality audit",f"- **Issues:** {q['issue_count']}",f"- **High:** {q['high']} | **Medium:** {q['medium']} | **Low:** {q['low']}",""] if not q["issues"]: return "\n".join(lines+["No issues were detected by the built-in checks."]) lines += ["| Severity | Check | Column | Evidence |","|---|---|---|---|"] for i in q["issues"][:30]: col=f"`{i['column']}`" if i["column"] else "dataset"; ev=i["evidence"].replace("|","\\|") lines.append(f"| **{i['severity'].upper()}** | {i['check']} | {col} | {ev} |") return "\n".join(lines) # Validate the dataset against an optional expected JSON schema def validate_schema(self, schema_json: str | None) -> dict[str, Any]: if not schema_json or not schema_json.strip(): return {"valid":True,"message":"No expected schema supplied. Inferred schema only.","errors":[],"warnings":[],"inferred_schema":self.profile()["column_schema"]} try: raw=json.loads(schema_json) except json.JSONDecodeError as exc: return {"valid":False,"message":f"Invalid schema JSON: {exc}","errors":[str(exc)],"warnings":[]} if isinstance(raw,dict) and "columns" in raw: raw=raw["columns"] # Normalize supported schema formats into a name-to-spec mapping expected={} if isinstance(raw,dict): for name,spec in raw.items(): expected[str(name)] = {"type":spec} if isinstance(spec,str) else spec if isinstance(spec,dict) else {"type":str(spec)} elif isinstance(raw,list): for item in raw: if isinstance(item,dict) and "name" in item: expected[str(item["name"])]=item else: return {"valid":False,"message":"Expected a JSON object or list describing columns.","errors":["Schema JSON must be object/list."],"warnings":[]} actual={str(c):self.dataframe[c] for c in self.dataframe.columns}; errors=[]; warnings=[] missing=sorted(set(expected)-set(actual)); unexpected=sorted(set(actual)-set(expected)) if missing: errors.append(f"Missing expected columns: {missing}") if unexpected: warnings.append(f"Unexpected columns present: {unexpected}") # Normalize common type names before comparing schemas aliases={"int":"integer","integer":"integer","long":"integer","bigint":"integer","float":"float","double":"float","decimal":"number","number":"number","numeric":"number","str":"string","string":"string","text":"string","bool":"boolean","boolean":"boolean","date":"datetime","datetime":"datetime","timestamp":"datetime"} for name,spec in expected.items(): if name not in actual: continue et=aliases.get(str(spec.get("type","")).lower(),str(spec.get("type","")).lower()); at=_logical_type(actual[name]) if not (et==at or (et=="number" and at in {"integer","float","number"}) or et==""): errors.append(f"{name}: expected type '{et}', inferred '{at}'.") if spec.get("nullable") is False and actual[name].isna().any(): errors.append(f"{name}: nullable=false but found {int(actual[name].isna().sum())} nulls.") if spec.get("unique") is True and actual[name].dropna().duplicated().any(): errors.append(f"{name}: unique=true but duplicate values were found.") return {"valid":len(errors)==0,"message":"Schema matches." if not errors else "Schema validation failed.","errors":errors,"warnings":warnings, "expected_columns":list(expected),"actual_columns":list(actual)} # Parse SQL and return a normalized form without executing it def validate_sql(self, sql: str, dialect: str = "duckdb") -> dict[str, Any]: try: parsed=sqlglot.parse_one((sql or "").strip(), read=dialect) return {"valid":True,"dialect":dialect,"normalized_sql":parsed.sql(dialect=dialect,pretty=True)} except Exception as exc: return {"valid":False,"dialect":dialect,"error":str(exc)} # Execute read-only SQL against the current DataFrame using DuckDB def execute_sql(self, sql: str, limit: int = 100) -> dict[str, Any]: sql=(sql or "").strip() try: tree=sqlglot.parse_one(sql, read="duckdb") except Exception as exc: return {"ok":False,"error":f"SQL parse failed: {exc}"} # Block statements that could modify data or database state forbidden=(exp.Insert,exp.Update,exp.Delete,exp.Create,exp.Drop,exp.Alter,exp.Command,exp.Copy) if any(tree.find(t) is not None for t in forbidden): return {"ok":False,"error":"Only read-only SELECT/CTE queries are allowed."} try: con=duckdb.connect(database=":memory:"); con.register("dataset",self.dataframe); out=con.execute(sql).df().head(limit); con.close() return {"ok":True,"rows_returned":int(len(out)),"columns":[str(c) for c in out.columns],"preview":json.loads(out.to_json(orient="records",date_format="iso"))} except Exception as exc: return {"ok":False,"error":str(exc)} # Generate a basic PySpark cleaning pipeline using the dataset columns def baseline_pyspark_pipeline(context: DataContext) -> str: cols=", ".join(repr(str(c)) for c in context.dataframe.columns[:20]) return f'''from pyspark.sql import SparkSession, functions as F spark = SparkSession.builder.appName("data-engineering-pipeline").getOrCreate() df = spark.read.option("header", True).option("inferSchema", True).csv("input.csv") source_columns = [{cols}] df = df.select(*[c for c in source_columns if c in df.columns]) for field in df.schema.fields: if field.dataType.simpleString() == "string": df = df.withColumn(field.name, F.when(F.trim(F.col(field.name)) == "", F.lit(None)).otherwise(F.trim(F.col(field.name)))) df = df.dropDuplicates() df.write.mode("overwrite").format("parquet").save("output/clean") ''' # Generate a basic SQL projection and deduplication pipeline def baseline_sql_pipeline(context: DataContext) -> str: cols=[f'"{str(c).replace(chr(34), chr(34)*2)}"' for c in context.dataframe.columns[:30]] projection=",\n ".join(cols) if cols else "*" return f'''WITH source AS ( SELECT {projection} FROM dataset ), deduplicated AS ( SELECT DISTINCT * FROM source ) SELECT * FROM deduplicated; '''