11"""Parser for SQL DDL into SchemaForge IR."""
2+
23from __future__ import annotations
34
45import 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