Skip to content

Commit 3c295e0

Browse files
style: apply ruff format to src/schemaforge/parsers/sql_parser.py
1 parent cd87229 commit 3c295e0

1 file changed

Lines changed: 89 additions & 46 deletions

File tree

‎src/schemaforge/parsers/sql_parser.py‎

Lines changed: 89 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
"""Parser for SQL DDL into SchemaForge IR."""
2+
23
from __future__ import annotations
34

45
import contextlib
@@ -13,11 +14,22 @@ class SQLParser:
1314

1415
# SQL function keywords that should be stored as fn: prefixed defaults
1516
_SQL_FN_KEYWORDS: set[str] = {
16-
"CURRENT_TIMESTAMP", "CURRENT_DATE", "CURRENT_TIME",
17-
"LOCALTIMESTAMP", "LOCALTIME",
18-
"NOW", "RANDOM", "GEN_RANDOM_UUID", "UUID",
19-
"RAND", "CURDATE", "CURTIME", "SYSDATE",
20-
"UTC_TIMESTAMP", "UTC_DATE", "UTC_TIME",
17+
"CURRENT_TIMESTAMP",
18+
"CURRENT_DATE",
19+
"CURRENT_TIME",
20+
"LOCALTIMESTAMP",
21+
"LOCALTIME",
22+
"NOW",
23+
"RANDOM",
24+
"GEN_RANDOM_UUID",
25+
"UUID",
26+
"RAND",
27+
"CURDATE",
28+
"CURTIME",
29+
"SYSDATE",
30+
"UTC_TIMESTAMP",
31+
"UTC_DATE",
32+
"UTC_TIME",
2133
}
2234

2335
def parse(self, text: str) -> Schema:
@@ -58,13 +70,13 @@ def _split_statements(self, text: str) -> list[str]:
5870
ch = text[i]
5971
if in_string:
6072
current += ch
61-
if ch == string_char and (i == 0 or text[i-1] != '\\'):
73+
if ch == string_char and (i == 0 or text[i - 1] != "\\"):
6274
in_string = False
6375
elif ch in ("'", '"'):
6476
in_string = True
6577
string_char = ch
6678
current += ch
67-
elif ch == ';':
79+
elif ch == ";":
6880
statements.append(current.strip())
6981
current = ""
7082
else:
@@ -79,11 +91,12 @@ def _parse_create_table(self, stmt: str) -> Table | None:
7991
"""Parse a CREATE TABLE statement."""
8092
# Extract table name — support quoted/backtick identifiers
8193
m = re.match(
82-
r'CREATE\s+(?:TEMPORARY\s+)?(?:OR\s+REPLACE\s+)?'
83-
r'TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?'
94+
r"CREATE\s+(?:TEMPORARY\s+)?(?:OR\s+REPLACE\s+)?"
95+
r"TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?"
8496
r'(?:`(\w+)`\.|"(\w+)"\.|(\w+)\.)?'
8597
r'`?"?([\w-]+)"?`?',
86-
stmt, re.IGNORECASE
98+
stmt,
99+
re.IGNORECASE,
87100
)
88101
if not m:
89102
return None
@@ -97,11 +110,11 @@ def _parse_create_table(self, stmt: str) -> Table | None:
97110
paren_depth = 0
98111
start = -1
99112
for i, ch in enumerate(stmt):
100-
if ch == '(':
113+
if ch == "(":
101114
if paren_depth == 0:
102115
start = i + 1
103116
paren_depth += 1
104-
elif ch == ')':
117+
elif ch == ")":
105118
paren_depth -= 1
106119
if paren_depth == 0 and start >= 0:
107120
body = stmt[start:i]
@@ -111,7 +124,7 @@ def _parse_create_table(self, stmt: str) -> Table | None:
111124
return None
112125

113126
# Extract MySQL table options after the closing paren (e.g. ENGINE=InnoDB AUTO_INCREMENT=1)
114-
table_options = self._parse_table_options(stmt[i+1:])
127+
table_options = self._parse_table_options(stmt[i + 1 :])
115128

116129
table = Table(name=table_name, options=table_options)
117130
# Split body into lines/definitions
@@ -150,13 +163,13 @@ def _split_definitions(self, body: str) -> list[str]:
150163
current = ""
151164
paren_depth = 0
152165
for ch in body:
153-
if ch == '(':
166+
if ch == "(":
154167
paren_depth += 1
155168
current += ch
156-
elif ch == ')':
169+
elif ch == ")":
157170
paren_depth -= 1
158171
current += ch
159-
elif ch == ',' and paren_depth == 0:
172+
elif ch == "," and paren_depth == 0:
160173
defs.append(current.strip())
161174
current = ""
162175
else:
@@ -200,17 +213,28 @@ def _parse_column_def(self, defn: str) -> Column | None:
200213
"""Parse a column definition."""
201214
# Skip table constraints
202215
upper = defn.upper()
203-
if (
204-
any(kw in upper for kw in [
205-
"PRIMARY KEY", "FOREIGN KEY", "INDEX",
206-
"KEY", "CHECK", "UNIQUE", "CONSTRAINT",
207-
])
208-
and (
209-
not defn.split()[0].isidentifier()
210-
or defn.split()[0].upper() in (
211-
"PRIMARY", "FOREIGN", "INDEX",
212-
"KEY", "CHECK", "UNIQUE", "CONSTRAINT",
213-
)
216+
if any(
217+
kw in upper
218+
for kw in [
219+
"PRIMARY KEY",
220+
"FOREIGN KEY",
221+
"INDEX",
222+
"KEY",
223+
"CHECK",
224+
"UNIQUE",
225+
"CONSTRAINT",
226+
]
227+
) and (
228+
not defn.split()[0].isidentifier()
229+
or defn.split()[0].upper()
230+
in (
231+
"PRIMARY",
232+
"FOREIGN",
233+
"INDEX",
234+
"KEY",
235+
"CHECK",
236+
"UNIQUE",
237+
"CONSTRAINT",
214238
)
215239
):
216240
return None
@@ -221,37 +245,47 @@ def _parse_column_def(self, defn: str) -> Column | None:
221245
return None
222246

223247
col_name = tokens[0]
224-
if col_name.startswith('"') or col_name.startswith('`') or col_name.startswith('['):
248+
if (
249+
col_name.startswith('"')
250+
or col_name.startswith("`")
251+
or col_name.startswith("[")
252+
):
225253
col_name = col_name.strip('"`[]')
226254

227255
# Find the type (skip quoted name)
228256
type_start = 1
229257
type_raw = tokens[type_start] if type_start < len(tokens) else "TEXT"
230258

231259
# Handle type with parameters that may span multiple tokens (e.g. ENUM('a', 'b', 'c'))
232-
if type_raw.startswith(('ENUM(', 'enum(')) or (
233-
'(' in type_raw and ')' not in type_raw and type_start + 1 < len(tokens)
260+
if type_raw.startswith(("ENUM(", "enum(")) or (
261+
"(" in type_raw and ")" not in type_raw and type_start + 1 < len(tokens)
234262
):
235-
while ')' not in type_raw and type_start + 1 < len(tokens):
263+
while ")" not in type_raw and type_start + 1 < len(tokens):
236264
type_start += 1
237265
type_raw += " " + tokens[type_start]
238266

239-
type_name = type_raw.split('(')[0].upper() if '(' in type_raw else type_raw.upper()
267+
type_name = (
268+
type_raw.split("(")[0].upper() if "(" in type_raw else type_raw.upper()
269+
)
240270

241271
col_type = self._TYPE_MAP.get(type_name, ColumnType.CUSTOM)
242272

243273
# Extract type args
244274
type_args: dict[str, Any] = {}
245-
if '(' in type_raw:
246-
args_str = type_raw[type_raw.index('(')+1:type_raw.index(')')] if ')' in type_raw else ""
275+
if "(" in type_raw:
276+
args_str = (
277+
type_raw[type_raw.index("(") + 1 : type_raw.index(")")]
278+
if ")" in type_raw
279+
else ""
280+
)
247281
if type_name == "ENUM":
248282
# Inline ENUM('a','b','c') — extract values via regex to handle multi-token args
249283
type_args["values"] = re.findall(r"'([^']*)'", args_str)
250284
if type_name in ("VARCHAR", "CHAR"):
251285
with contextlib.suppress(ValueError):
252286
type_args["length"] = int(args_str)
253287
elif type_name in ("DECIMAL", "NUMERIC"):
254-
parts = args_str.split(',')
288+
parts = args_str.split(",")
255289
if len(parts) >= 1:
256290
with contextlib.suppress(ValueError):
257291
type_args["precision"] = int(parts[0])
@@ -260,7 +294,9 @@ def _parse_column_def(self, defn: str) -> Column | None:
260294
type_args["scale"] = int(parts[1])
261295

262296
# Parse constraints
263-
constraints = " ".join(tokens[type_start+1:]) if type_start + 1 < len(tokens) else ""
297+
constraints = (
298+
" ".join(tokens[type_start + 1 :]) if type_start + 1 < len(tokens) else ""
299+
)
264300

265301
is_pk = "PRIMARY KEY" in constraints.upper()
266302
is_not_null = "NOT NULL" in constraints.upper()
@@ -293,7 +329,9 @@ def _parse_column_def(self, defn: str) -> Column | None:
293329
is_fn = (
294330
upper_val in self._SQL_FN_KEYWORDS
295331
or upper_val.rstrip("()") in self._SQL_FN_KEYWORDS
296-
or re.match(r'^\w+\(', val) # Any function call: nextval(), now(), etc.
332+
or re.match(
333+
r"^\w+\(", val
334+
) # Any function call: nextval(), now(), etc.
297335
)
298336
if is_fn:
299337
col.default = f"fn:{val}"
@@ -318,18 +356,17 @@ def _parse_column_def(self, defn: str) -> Column | None:
318356

319357
def _parse_index_definition(self, defn: str) -> Index | None:
320358
"""Parse an index definition."""
321-
m = re.search(r'(?:INDEX|KEY)\s+(\w+)\s*\(([^)]+)\)', defn, re.IGNORECASE)
359+
m = re.search(r"(?:INDEX|KEY)\s+(\w+)\s*\(([^)]+)\)", defn, re.IGNORECASE)
322360
if m:
323361
name = m.group(1)
324-
columns = [c.strip().strip('"`[]') for c in m.group(2).split(',')]
362+
columns = [c.strip().strip('"`[]') for c in m.group(2).split(",")]
325363
return Index(name=name, columns=columns, unique="UNIQUE" in defn.upper())
326364
return None
327365

328366
def _parse_create_enum(self, stmt: str) -> EnumType | None:
329367
"""Parse a CREATE TYPE ... AS ENUM statement."""
330368
m = re.match(
331-
r"CREATE\s+TYPE\s+(\w+)\s+AS\s+ENUM\s*\(([^)]+)\)",
332-
stmt, re.IGNORECASE
369+
r"CREATE\s+TYPE\s+(\w+)\s+AS\s+ENUM\s*\(([^)]+)\)", stmt, re.IGNORECASE
333370
)
334371
if m:
335372
name = m.group(1)
@@ -353,23 +390,29 @@ def _parse_table_options(self, after_paren: str) -> dict[str, str]:
353390
i = 0
354391
while i < len(text):
355392
# Skip whitespace
356-
while i < len(text) and text[i] in (' ', '\t', '\n', '\r'):
393+
while i < len(text) and text[i] in (" ", "\t", "\n", "\r"):
357394
i += 1
358395
if i >= len(text):
359396
break
360397

361398
# Check for "DEFAULT CHARSET" or "DEFAULT COLLATE" two-word keys
362399
remaining = text[i:]
363-
default_m = re.match(r'DEFAULT\s+(CHARSET|CHARACTER\s+SET|COLLATE)\s*=\s*', remaining, re.IGNORECASE)
400+
default_m = re.match(
401+
r"DEFAULT\s+(CHARSET|CHARACTER\s+SET|COLLATE)\s*=\s*",
402+
remaining,
403+
re.IGNORECASE,
404+
)
364405
if default_m:
365406
key = f"DEFAULT {default_m.group(1).upper()}"
366407
val_start = i + default_m.end()
367408
else:
368409
# Single-word key: ENGINE, AUTO_INCREMENT, ROW_FORMAT, etc.
369-
key_m = re.match(r'(\w+)\s*=\s*', remaining)
410+
key_m = re.match(r"(\w+)\s*=\s*", remaining)
370411
if not key_m:
371412
# Try COMMENT 'xxx' (no equals)
372-
comment_m = re.match(r"COMMENT\s+'([^']*)'", remaining, re.IGNORECASE)
413+
comment_m = re.match(
414+
r"COMMENT\s+'([^']*)'", remaining, re.IGNORECASE
415+
)
373416
if comment_m:
374417
options["COMMENT"] = comment_m.group(1)
375418
i += comment_m.end()
@@ -390,7 +433,7 @@ def _parse_table_options(self, after_paren: str) -> dict[str, str]:
390433
break
391434
else:
392435
# Unquoted value — grab until next space, semicolon, or end
393-
val_end = re.match(r'([^\s;]+)', val_rest)
436+
val_end = re.match(r"([^\s;]+)", val_rest)
394437
if val_end:
395438
options[key] = val_end.group(1)
396439
i = val_start + val_end.end()

0 commit comments

Comments
 (0)