diff --git a/src/json2sql/converter.py b/src/json2sql/converter.py index a326f8c..be8c27b 100644 --- a/src/json2sql/converter.py +++ b/src/json2sql/converter.py @@ -43,18 +43,12 @@ def convert(self, json_text: str, table_name: str = "data") -> str: # Add any extra tables from flattening for name, columns, rows in self._extra_tables: statements.insert(0, create_table_sql(name, columns, self.dialect)) - statements.append( - insert_sql(name, list(columns.keys()), rows, self.dialect) - ) + statements.append(insert_sql(name, list(columns.keys()), rows, self.dialect)) result = "\n\n".join(s for s in statements if s) # An empty object / nested-only root legitimately produces no SQL; say # so explicitly instead of returning "" (avoids a silent green no-op). - return ( - result - if result - else "-- No columns to generate (empty or nested-only object)." - ) + return result if result else "-- No columns to generate (empty or nested-only object)." def generate_schema(self, json_text: str, table_name: str = "data") -> str: """Generate only CREATE TABLE statements from JSON data.""" @@ -74,7 +68,9 @@ def generate_schema(self, json_text: str, table_name: str = "data") -> str: else: columns = self._infer_columns(objects) else: - columns = {"value": "TEXT"} + # Primitive array — infer type from first element when available + col_type = sql_type_for(data[0], self.dialect) if isinstance(data, list) and data else "TEXT" + columns = {"value": col_type} statements = [] if columns: @@ -96,11 +92,7 @@ def _convert_objects(self, objects: list[dict], table_name: str) -> str: # Process nested arrays into child tables for obj in objects: for key, value in obj.items(): - if ( - isinstance(value, list) - and value - and all(isinstance(v, dict) for v in value) - ): + if isinstance(value, list) and value and all(isinstance(v, dict) for v in value): self._flatten_nested(table_name, key, value, obj) else: columns = self._infer_columns(objects) @@ -134,9 +126,7 @@ def _convert_objects(self, objects: list[dict], table_name: str) -> str: return "" parts = [create_table_sql(table_name, columns, self.dialect)] if rows: - parts.append( - insert_sql(table_name, list(columns.keys()), rows, self.dialect) - ) + parts.append(insert_sql(table_name, list(columns.keys()), rows, self.dialect)) return "\n\n".join(parts) def _convert_primitives(self, values: list, table_name: str) -> str: @@ -215,15 +205,8 @@ def _infer_columns_flattened( columns[flat_key] = inferred flat_map[flat_key] = (key, sub_key) elif inferred is not None: - columns[flat_key] = self._merge_type( - columns[flat_key], inferred - ) - elif ( - isinstance(value, list) - and value - and self.flatten - and all(isinstance(v, dict) for v in value) - ): + columns[flat_key] = self._merge_type(columns[flat_key], inferred) + elif isinstance(value, list) and value and self.flatten and all(isinstance(v, dict) for v in value): # Skip - goes to separate table pass else: @@ -279,9 +262,5 @@ def _process_flatten(self, objects: list, table_name: str) -> None: return for obj in objects: for key, value in obj.items(): - if ( - isinstance(value, list) - and value - and all(isinstance(v, dict) for v in value) - ): + if isinstance(value, list) and value and all(isinstance(v, dict) for v in value): self._flatten_nested(table_name, key, value, obj) diff --git a/tests/test_converter.py b/tests/test_converter.py index 402a453..c5e475b 100644 --- a/tests/test_converter.py +++ b/tests/test_converter.py @@ -535,3 +535,27 @@ def test_version_in_init_matches_pyproject(self): assert data["project"]["version"] == __version__, ( f"pyproject.toml version ({data['project']['version']}) != __init__.__version__ ({__version__})" ) + + +class TestGenerateSchemaPrimitiveTypeInference: + """Regression: generate_schema must infer primitive column types, not always TEXT.""" + + def test_schema_primitive_int_array(self): + conv = JSONToSQLConverter(dialect=Dialect.POSTGRES) + data = json.dumps([1, 2, 3]) + result = conv.generate_schema(data, table_name="nums") + assert "INTEGER" in result + assert "TEXT" not in result.split("CREATE TABLE")[1].split(")")[0] + + def test_schema_primitive_float_array_mysql(self): + conv = JSONToSQLConverter(dialect=Dialect.MYSQL) + data = json.dumps([1.5, 2.5]) + result = conv.generate_schema(data, table_name="vals") + assert "DOUBLE" in result + + def test_schema_primitive_bool_array_sqlite(self): + conv = JSONToSQLConverter(dialect=Dialect.SQLITE) + data = json.dumps([True, False]) + result = conv.generate_schema(data, table_name="flags") + # SQLite bool -> INTEGER + assert "INTEGER" in result