LIANJie-Jason Claude Opus 4.6 commited on
Commit
863e8aa
·
1 Parent(s): 9649eb0

feat: add SQL ingestion — load tabular files into SQLite with schema registry

Browse files
Files changed (2) hide show
  1. src/sql_ingest.py +217 -0
  2. tests/test_sql_ingest.py +102 -0
src/sql_ingest.py CHANGED
@@ -1,6 +1,11 @@
1
  """SQL ingestion: load tabular files into SQLite for structured queries."""
2
 
 
 
 
3
  import re
 
 
4
 
5
  # Extensions that trigger SQL ingestion (tabular formats)
6
  SQL_EXTENSIONS = {".csv", ".tab", ".tsv", ".xlsx", ".xls", ".dta", ".sav", ".rds", ".rda"}
@@ -59,3 +64,215 @@ def _get_sample_values(values: list, n: int = 3) -> list:
59
  if len(samples) >= n:
60
  break
61
  return samples
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  """SQL ingestion: load tabular files into SQLite for structured queries."""
2
 
3
+ import csv
4
+ import json
5
+ import os
6
  import re
7
+ import sqlite3
8
+ from pathlib import Path
9
 
10
  # Extensions that trigger SQL ingestion (tabular formats)
11
  SQL_EXTENSIONS = {".csv", ".tab", ".tsv", ".xlsx", ".xls", ".dta", ".sav", ".rds", ".rda"}
 
64
  if len(samples) >= n:
65
  break
66
  return samples
67
+
68
+
69
+ def _load_rows_from_csv(file_path: str, ext: str) -> list[tuple]:
70
+ """Load CSV/TSV/TAB into (sheet_or_name, headers, rows) tuples."""
71
+ delimiter = "\t" if ext in (".tab", ".tsv") else ","
72
+ with open(file_path, "r", encoding="utf-8", errors="replace") as f:
73
+ reader = csv.reader(f, delimiter=delimiter)
74
+ all_rows = list(reader)
75
+ if len(all_rows) < 2:
76
+ return []
77
+ headers = [h.strip().strip('"') for h in all_rows[0]]
78
+ return [(None, headers, all_rows[1:])]
79
+
80
+
81
+ def _load_rows_from_excel(file_path: str, ext: str) -> list[tuple]:
82
+ """Load Excel into (sheet_name, headers, rows) tuples."""
83
+ if ext == ".xlsx":
84
+ import openpyxl
85
+ wb = openpyxl.load_workbook(file_path, read_only=True, data_only=True)
86
+ results = []
87
+ for sheet_name in wb.sheetnames:
88
+ ws = wb[sheet_name]
89
+ all_rows = list(ws.iter_rows(values_only=True))
90
+ if len(all_rows) < 2:
91
+ continue
92
+ headers = [str(h) if h is not None else "" for h in all_rows[0]]
93
+ rows = [[str(c) if c is not None else "" for c in row] for row in all_rows[1:]]
94
+ results.append((sheet_name, headers, rows))
95
+ wb.close()
96
+ return results
97
+ else: # .xls
98
+ import xlrd
99
+ wb = xlrd.open_workbook(file_path)
100
+ results = []
101
+ for sheet_idx in range(wb.nsheets):
102
+ ws = wb.sheet_by_index(sheet_idx)
103
+ if ws.nrows < 2:
104
+ continue
105
+ headers = [str(ws.cell_value(0, c)) for c in range(ws.ncols)]
106
+ rows = [[str(ws.cell_value(r, c)) for c in range(ws.ncols)] for r in range(1, ws.nrows)]
107
+ results.append((ws.name, headers, rows))
108
+ return results
109
+
110
+
111
+ def _load_rows_from_stata(file_path: str) -> list[tuple]:
112
+ """Load Stata .dta into (None, headers, rows) tuples."""
113
+ import pyreadstat
114
+ df, _meta = pyreadstat.read_dta(file_path)
115
+ headers = list(df.columns)
116
+ rows = [[str(v) for v in row] for row in df.values.tolist()]
117
+ return [(None, headers, rows)]
118
+
119
+
120
+ def _load_rows_from_spss(file_path: str) -> list[tuple]:
121
+ """Load SPSS .sav into (None, headers, rows) tuples."""
122
+ import pyreadstat
123
+ df, _meta = pyreadstat.read_sav(file_path)
124
+ headers = list(df.columns)
125
+ rows = [[str(v) for v in row] for row in df.values.tolist()]
126
+ return [(None, headers, rows)]
127
+
128
+
129
+ def _load_rows_from_rdata(file_path: str) -> list[tuple]:
130
+ """Load R data files into (object_name, headers, rows) tuples."""
131
+ import pyreadr
132
+ result = pyreadr.read_r(file_path)
133
+ tables = []
134
+ for name, df in result.items():
135
+ headers = list(df.columns)
136
+ rows = [[str(v) for v in row] for row in df.values.tolist()]
137
+ tables.append((name, headers, rows))
138
+ return tables
139
+
140
+
141
+ def _load_tabular_file(file_path: str, ext: str) -> list[tuple]:
142
+ """Dispatch to the correct loader based on extension.
143
+
144
+ Returns list of (sheet_or_name, headers, rows) tuples.
145
+ """
146
+ if ext in (".csv", ".tab", ".tsv"):
147
+ return _load_rows_from_csv(file_path, ext)
148
+ if ext in (".xlsx", ".xls"):
149
+ return _load_rows_from_excel(file_path, ext)
150
+ if ext == ".dta":
151
+ return _load_rows_from_stata(file_path)
152
+ if ext == ".sav":
153
+ return _load_rows_from_spss(file_path)
154
+ if ext in (".rds", ".rda"):
155
+ return _load_rows_from_rdata(file_path)
156
+ return []
157
+
158
+
159
+ def _sanitize_column_name(name: str) -> str:
160
+ """Sanitize a column name into a valid SQL identifier."""
161
+ safe = re.sub(r"[^a-zA-Z0-9_]", "_", name).strip("_")
162
+ if not safe or not safe[0].isalpha():
163
+ safe = "col_" + safe
164
+ return safe
165
+
166
+
167
+ def ingest_to_sql(files: list[tuple], documents_dir: str, cfg: dict) -> dict:
168
+ """Ingest tabular files into SQLite and generate schema registry.
169
+
170
+ Args:
171
+ files: List of (file_path, dataset_name) tuples.
172
+ documents_dir: Root documents directory.
173
+ cfg: App config dict.
174
+
175
+ Returns:
176
+ Schema registry dict (also saved to sql_schemas.json).
177
+ """
178
+ sql_db_dir = cfg.get("paths", {}).get("sql_db", "sql_db")
179
+ if not os.path.isabs(sql_db_dir):
180
+ project_root = Path(__file__).resolve().parent.parent
181
+ sql_db_dir = os.path.join(str(project_root), sql_db_dir)
182
+
183
+ os.makedirs(sql_db_dir, exist_ok=True)
184
+ db_path = os.path.join(sql_db_dir, "knowledge_base.db")
185
+ schema_path = os.path.join(sql_db_dir, "sql_schemas.json")
186
+
187
+ # Filter to tabular files only
188
+ tabular_files = [(fp, ds) for fp, ds in files if fp.suffix.lower() in SQL_EXTENSIONS]
189
+
190
+ if not tabular_files:
191
+ if os.path.exists(schema_path):
192
+ os.remove(schema_path)
193
+ return {}
194
+
195
+ # Clear existing DB
196
+ if os.path.exists(db_path):
197
+ os.remove(db_path)
198
+
199
+ conn = sqlite3.connect(db_path)
200
+ schema_registry = {}
201
+
202
+ for file_path, dataset_name in tabular_files:
203
+ ext = file_path.suffix.lower()
204
+ rel_path = file_path.relative_to(documents_dir)
205
+ source_file = str(rel_path)
206
+
207
+ try:
208
+ tables = _load_tabular_file(str(file_path), ext)
209
+ except Exception as e:
210
+ print(f" SQL ingest error for {file_path.name}: {e}")
211
+ continue
212
+
213
+ for sheet_or_name, headers, rows in tables:
214
+ table_name = _sanitize_table_name(dataset_name, file_path.stem, ext, sheet_or_name)
215
+
216
+ # Infer column types and collect samples
217
+ safe_headers = [_sanitize_column_name(h) for h in headers]
218
+ col_types = []
219
+ col_samples = []
220
+ for col_idx in range(len(headers)):
221
+ col_values = [row[col_idx] if col_idx < len(row) else None for row in rows]
222
+ col_types.append(_infer_column_type(col_values))
223
+ col_samples.append(_get_sample_values(col_values))
224
+
225
+ # Create table
226
+ col_defs = ", ".join(f'"{h}" {t}' for h, t in zip(safe_headers, col_types))
227
+ conn.execute(f'DROP TABLE IF EXISTS "{table_name}"')
228
+ conn.execute(f'CREATE TABLE "{table_name}" ({col_defs})')
229
+
230
+ # Insert rows
231
+ placeholders = ", ".join(["?"] * len(safe_headers))
232
+ insert_sql = f'INSERT INTO "{table_name}" VALUES ({placeholders})'
233
+
234
+ for row in rows:
235
+ values = []
236
+ for col_idx, col_type in enumerate(col_types):
237
+ raw = row[col_idx] if col_idx < len(row) else None
238
+ if raw is None or str(raw).strip() == "" or str(raw).lower() == "nan":
239
+ values.append(None)
240
+ elif col_type == "INTEGER":
241
+ try:
242
+ values.append(int(float(str(raw).strip())))
243
+ except (ValueError, TypeError):
244
+ values.append(None)
245
+ elif col_type == "REAL":
246
+ try:
247
+ values.append(float(str(raw).strip()))
248
+ except (ValueError, TypeError):
249
+ values.append(None)
250
+ else:
251
+ values.append(str(raw).strip())
252
+ conn.execute(insert_sql, values)
253
+
254
+ conn.commit()
255
+
256
+ row_count = conn.execute(f'SELECT COUNT(*) FROM "{table_name}"').fetchone()[0]
257
+
258
+ schema_registry[table_name] = {
259
+ "source_file": source_file,
260
+ "columns": [
261
+ {
262
+ "name": safe_headers[i],
263
+ "original_name": headers[i],
264
+ "type": col_types[i],
265
+ "sample": col_samples[i],
266
+ }
267
+ for i in range(len(headers))
268
+ ],
269
+ "row_count": row_count,
270
+ }
271
+ print(f" SQL: {table_name} ({row_count} rows, {len(headers)} columns)")
272
+
273
+ conn.close()
274
+
275
+ with open(schema_path, "w", encoding="utf-8") as f:
276
+ json.dump(schema_registry, f, indent=2, default=str)
277
+
278
+ return schema_registry
tests/test_sql_ingest.py CHANGED
@@ -61,3 +61,105 @@ def test_get_sample_values_fewer_than_n():
61
  from src.sql_ingest import _get_sample_values
62
  samples = _get_sample_values(["a", None, "a"], n=3)
63
  assert samples == ["a"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
61
  from src.sql_ingest import _get_sample_values
62
  samples = _get_sample_values(["a", None, "a"], n=3)
63
  assert samples == ["a"]
64
+
65
+
66
+ import csv
67
+ import json
68
+ import sqlite3
69
+
70
+
71
+ def test_ingest_to_sql_creates_db_and_schema(tmp_path):
72
+ from pathlib import Path
73
+ from src.sql_ingest import ingest_to_sql
74
+
75
+ # Create a CSV file
76
+ kb_dir = tmp_path / "knowledge_base"
77
+ ds_dir = kb_dir / "testds"
78
+ ds_dir.mkdir(parents=True)
79
+ csv_file = ds_dir / "data.csv"
80
+ csv_file.write_text("Country,Year,Score\nChina,2005,4.0\nIndia,2005,3.0\n")
81
+
82
+ sql_db_dir = tmp_path / "sql_db"
83
+ cfg = {"paths": {"sql_db": str(sql_db_dir)}}
84
+ files = [(Path(csv_file), "testds")]
85
+
86
+ schema = ingest_to_sql(files, str(kb_dir), cfg)
87
+
88
+ # Check DB was created
89
+ db_path = sql_db_dir / "knowledge_base.db"
90
+ assert db_path.exists()
91
+
92
+ # Check schema registry
93
+ schema_path = sql_db_dir / "sql_schemas.json"
94
+ assert schema_path.exists()
95
+ with open(schema_path) as f:
96
+ saved = json.load(f)
97
+ assert len(saved) == 1
98
+ table_name = list(saved.keys())[0]
99
+ assert "testds" in table_name
100
+ assert saved[table_name]["row_count"] == 2
101
+ assert len(saved[table_name]["columns"]) == 3
102
+
103
+ # Check data in SQLite
104
+ conn = sqlite3.connect(str(db_path))
105
+ rows = conn.execute(f'SELECT * FROM "{table_name}"').fetchall()
106
+ conn.close()
107
+ assert len(rows) == 2
108
+ assert rows[0][0] == "China"
109
+
110
+
111
+ def test_ingest_to_sql_type_inference(tmp_path):
112
+ from pathlib import Path
113
+ from src.sql_ingest import ingest_to_sql
114
+
115
+ kb_dir = tmp_path / "knowledge_base"
116
+ ds_dir = kb_dir / "ds"
117
+ ds_dir.mkdir(parents=True)
118
+ csv_file = ds_dir / "typed.csv"
119
+ csv_file.write_text("Name,Year,Score\nChina,2005,4.5\nIndia,2006,3.0\n")
120
+
121
+ sql_db_dir = tmp_path / "sql_db"
122
+ cfg = {"paths": {"sql_db": str(sql_db_dir)}}
123
+ files = [(Path(csv_file), "ds")]
124
+
125
+ schema = ingest_to_sql(files, str(kb_dir), cfg)
126
+ table_name = list(schema.keys())[0]
127
+
128
+ # Name=TEXT, Year=INTEGER, Score=REAL
129
+ col_types = {c["name"]: c["type"] for c in schema[table_name]["columns"]}
130
+ assert col_types["Name"] == "TEXT"
131
+ assert col_types["Year"] == "INTEGER"
132
+ assert col_types["Score"] == "REAL"
133
+
134
+
135
+ def test_ingest_to_sql_no_tabular_files(tmp_path):
136
+ from pathlib import Path
137
+ from src.sql_ingest import ingest_to_sql
138
+
139
+ sql_db_dir = tmp_path / "sql_db"
140
+ cfg = {"paths": {"sql_db": str(sql_db_dir)}}
141
+ # Pass empty list (no tabular files)
142
+ schema = ingest_to_sql([], str(tmp_path), cfg)
143
+ assert schema == {}
144
+
145
+
146
+ def test_ingest_to_sql_clears_on_rerun(tmp_path):
147
+ from pathlib import Path
148
+ from src.sql_ingest import ingest_to_sql
149
+
150
+ kb_dir = tmp_path / "knowledge_base"
151
+ ds_dir = kb_dir / "ds"
152
+ ds_dir.mkdir(parents=True)
153
+ csv_file = ds_dir / "data.csv"
154
+ csv_file.write_text("A,B\n1,2\n3,4\n")
155
+
156
+ sql_db_dir = tmp_path / "sql_db"
157
+ cfg = {"paths": {"sql_db": str(sql_db_dir)}}
158
+ files = [(Path(csv_file), "ds")]
159
+
160
+ # First run
161
+ ingest_to_sql(files, str(kb_dir), cfg)
162
+ # Second run — should not duplicate
163
+ schema = ingest_to_sql(files, str(kb_dir), cfg)
164
+ table_name = list(schema.keys())[0]
165
+ assert schema[table_name]["row_count"] == 2