Skip to content

Commit f08bdb9

Browse files
cowork-bot: fix generate_schema to infer primitive column types instead of hardcoding TEXT
1 parent 94ee7aa commit f08bdb9

2 files changed

Lines changed: 34 additions & 31 deletions

File tree

src/json2sql/converter.py

Lines changed: 10 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -43,18 +43,12 @@ def convert(self, json_text: str, table_name: str = "data") -> str:
4343
# Add any extra tables from flattening
4444
for name, columns, rows in self._extra_tables:
4545
statements.insert(0, create_table_sql(name, columns, self.dialect))
46-
statements.append(
47-
insert_sql(name, list(columns.keys()), rows, self.dialect)
48-
)
46+
statements.append(insert_sql(name, list(columns.keys()), rows, self.dialect))
4947

5048
result = "\n\n".join(s for s in statements if s)
5149
# An empty object / nested-only root legitimately produces no SQL; say
5250
# so explicitly instead of returning "" (avoids a silent green no-op).
53-
return (
54-
result
55-
if result
56-
else "-- No columns to generate (empty or nested-only object)."
57-
)
51+
return result if result else "-- No columns to generate (empty or nested-only object)."
5852

5953
def generate_schema(self, json_text: str, table_name: str = "data") -> str:
6054
"""Generate only CREATE TABLE statements from JSON data."""
@@ -74,7 +68,9 @@ def generate_schema(self, json_text: str, table_name: str = "data") -> str:
7468
else:
7569
columns = self._infer_columns(objects)
7670
else:
77-
columns = {"value": "TEXT"}
71+
# Primitive array — infer type from first element when available
72+
col_type = sql_type_for(data[0], self.dialect) if isinstance(data, list) and data else "TEXT"
73+
columns = {"value": col_type}
7874

7975
statements = []
8076
if columns:
@@ -96,11 +92,7 @@ def _convert_objects(self, objects: list[dict], table_name: str) -> str:
9692
# Process nested arrays into child tables
9793
for obj in objects:
9894
for key, value in obj.items():
99-
if (
100-
isinstance(value, list)
101-
and value
102-
and all(isinstance(v, dict) for v in value)
103-
):
95+
if isinstance(value, list) and value and all(isinstance(v, dict) for v in value):
10496
self._flatten_nested(table_name, key, value, obj)
10597
else:
10698
columns = self._infer_columns(objects)
@@ -134,9 +126,7 @@ def _convert_objects(self, objects: list[dict], table_name: str) -> str:
134126
return ""
135127
parts = [create_table_sql(table_name, columns, self.dialect)]
136128
if rows:
137-
parts.append(
138-
insert_sql(table_name, list(columns.keys()), rows, self.dialect)
139-
)
129+
parts.append(insert_sql(table_name, list(columns.keys()), rows, self.dialect))
140130
return "\n\n".join(parts)
141131

142132
def _convert_primitives(self, values: list, table_name: str) -> str:
@@ -215,15 +205,8 @@ def _infer_columns_flattened(
215205
columns[flat_key] = inferred
216206
flat_map[flat_key] = (key, sub_key)
217207
elif inferred is not None:
218-
columns[flat_key] = self._merge_type(
219-
columns[flat_key], inferred
220-
)
221-
elif (
222-
isinstance(value, list)
223-
and value
224-
and self.flatten
225-
and all(isinstance(v, dict) for v in value)
226-
):
208+
columns[flat_key] = self._merge_type(columns[flat_key], inferred)
209+
elif isinstance(value, list) and value and self.flatten and all(isinstance(v, dict) for v in value):
227210
# Skip - goes to separate table
228211
pass
229212
else:
@@ -279,9 +262,5 @@ def _process_flatten(self, objects: list, table_name: str) -> None:
279262
return
280263
for obj in objects:
281264
for key, value in obj.items():
282-
if (
283-
isinstance(value, list)
284-
and value
285-
and all(isinstance(v, dict) for v in value)
286-
):
265+
if isinstance(value, list) and value and all(isinstance(v, dict) for v in value):
287266
self._flatten_nested(table_name, key, value, obj)

tests/test_converter.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -535,3 +535,27 @@ def test_version_in_init_matches_pyproject(self):
535535
assert data["project"]["version"] == __version__, (
536536
f"pyproject.toml version ({data['project']['version']}) != __init__.__version__ ({__version__})"
537537
)
538+
539+
540+
class TestGenerateSchemaPrimitiveTypeInference:
541+
"""Regression: generate_schema must infer primitive column types, not always TEXT."""
542+
543+
def test_schema_primitive_int_array(self):
544+
conv = JSONToSQLConverter(dialect=Dialect.POSTGRES)
545+
data = json.dumps([1, 2, 3])
546+
result = conv.generate_schema(data, table_name="nums")
547+
assert "INTEGER" in result
548+
assert "TEXT" not in result.split("CREATE TABLE")[1].split(")")[0]
549+
550+
def test_schema_primitive_float_array_mysql(self):
551+
conv = JSONToSQLConverter(dialect=Dialect.MYSQL)
552+
data = json.dumps([1.5, 2.5])
553+
result = conv.generate_schema(data, table_name="vals")
554+
assert "DOUBLE" in result
555+
556+
def test_schema_primitive_bool_array_sqlite(self):
557+
conv = JSONToSQLConverter(dialect=Dialect.SQLITE)
558+
data = json.dumps([True, False])
559+
result = conv.generate_schema(data, table_name="flags")
560+
# SQLite bool -> INTEGER
561+
assert "INTEGER" in result

0 commit comments

Comments
 (0)