1use sqlparser::ast::Statement;
2use sqlparser::dialect::{PostgreSqlDialect, SQLiteDialect};
3use sqlparser::parser::Parser;
4
5use super::error::ParseError;
6use super::segments::{SqlObjectName, SqlSegment, SqlStatementKind, segment_sql};
7use super::table_recovery::recover_table_sql;
8use super::{postgres, sqlite};
9use crate::dialects::Dialect;
10use crate::states::types::EntityKind;
11use crate::states::{
12 ExtensionDef, FunctionDef, Index, Schema, Table, TriggerDef, ViewDef, schema_qualified_key,
13};
14
15pub(crate) struct ParseContext {
16 pub(crate) schema: Schema,
17 pending_indexes: Vec<(String, Index)>,
18}
19
20impl ParseContext {
21 pub(crate) fn new() -> Self {
22 Self {
23 schema: Schema::default(),
24 pending_indexes: Vec::new(),
25 }
26 }
27
28 pub(crate) fn insert_table(&mut self, table: (String, Table)) -> Result<(), ParseError> {
29 let (key, table) = table;
30 if self.schema.tables.contains_key(&key) {
31 return Err(ParseError::DuplicateTable(key));
32 }
33 self.schema.tables.insert(key, table);
34 Ok(())
35 }
36
37 pub(crate) fn push_index(&mut self, index: (String, Index)) {
38 self.pending_indexes.push(index);
39 }
40
41 fn preserve_opaque_index_source(&mut self, raw: &str) {
43 let Some((_, index)) = self.pending_indexes.last_mut() else {
44 return;
45 };
46 if index.is_opaque() {
47 *index = Index::from_raw(index.name.clone(), raw.to_string());
48 }
49 }
50
51 fn finish(self) -> Result<Schema, ParseError> {
52 self.finish_raw()
53 }
54
55 fn finish_raw(mut self) -> Result<Schema, ParseError> {
56 for (table_name, index) in self.pending_indexes {
57 match self.schema.tables.get_mut(&table_name) {
58 Some(table) => table.indexes.push(index),
59 None => return Err(ParseError::UnknownTable { table: table_name }),
60 }
61 }
62 Ok(self.schema)
63 }
64}
65
66pub fn parse_sql(sql: &str, dialect: Dialect) -> Result<Schema, ParseError> {
68 Ok(parse_sql_raw(sql, dialect)?.prepare(dialect)?)
69}
70
71pub(crate) fn parse_sql_raw(sql: &str, dialect: Dialect) -> Result<Schema, ParseError> {
72 let segments = segment_sql(sql, dialect)?;
73 if matches!(dialect, Dialect::Mysql) {
74 return Err(ParseError::UnsupportedDialect("mysql".to_string()));
75 }
76
77 let mut ctx = ParseContext::new();
78 for segment in segments {
79 ensure_schema_segment(&segment, dialect)?;
80 let statements = match parse_segment(&segment.sql, dialect) {
81 Ok(statements) => statements,
82 Err(error) => {
83 if recover_modeled_table(&segment, &mut ctx, dialect)? {
84 continue;
85 }
86 if lower_raw_segment(&segment, &mut ctx, dialect).is_ok() {
87 continue;
88 }
89 return Err(ParseError::parse_in_segment(dialect, &segment, error));
90 }
91 };
92 for stmt in &statements {
93 ensure_modeled_create_statement(stmt, dialect)?;
94 let lowered = match dialect {
95 Dialect::Postgres => postgres::lower_statement(stmt, &mut ctx),
96 Dialect::Sqlite => sqlite::lower_statement(stmt, &mut ctx),
97 Dialect::Mysql => unreachable!("mysql returned above"),
98 };
99 if let Err(error) = lowered {
100 if matches!(segment.kind, Some(SqlStatementKind::Ddl(ref ddl)) if ddl.entity == EntityKind::Table)
101 {
102 return Err(error);
103 }
104 lower_raw_segment(&segment, &mut ctx, dialect)?;
105 } else if matches!(
106 segment.kind,
107 Some(SqlStatementKind::Ddl(ref ddl)) if ddl.entity == EntityKind::Index
108 ) {
109 ctx.preserve_opaque_index_source(&segment.sql);
110 }
111 }
112 }
113 ctx.finish()
114}
115
116fn ensure_schema_segment(segment: &SqlSegment, dialect: Dialect) -> Result<(), ParseError> {
118 if matches!(segment.kind, Some(SqlStatementKind::Ddl(_))) {
119 return Ok(());
120 }
121 Err(ParseError::unsupported(
122 dialect,
123 segment.sql.clone(),
124 "schema SQL must be a CREATE statement for a known Gaman entity kind",
125 ))
126}
127
128fn parse_segment(
130 sql: &str,
131 dialect: Dialect,
132) -> Result<Vec<Statement>, sqlparser::parser::ParserError> {
133 match dialect {
134 Dialect::Postgres => Parser::parse_sql(&PostgreSqlDialect {}, sql),
135 Dialect::Sqlite => Parser::parse_sql(&SQLiteDialect {}, sql),
136 Dialect::Mysql => unreachable!("mysql returned before segment parsing"),
137 }
138}
139
140fn recover_modeled_table(
142 segment: &SqlSegment,
143 ctx: &mut ParseContext,
144 dialect: Dialect,
145) -> Result<bool, ParseError> {
146 let Some(SqlStatementKind::Ddl(ddl)) = &segment.kind else {
147 return Ok(false);
148 };
149 if ddl.entity != EntityKind::Table {
150 return Ok(false);
151 }
152 let Some(recovered) = recover_table_sql(&segment.sql, dialect) else {
153 return Ok(false);
154 };
155 let Ok(statements) = parse_segment(&recovered.core_sql, dialect) else {
156 return Ok(false);
157 };
158 for statement in &statements {
159 ensure_modeled_create_statement(statement, dialect)?;
160 match dialect {
161 Dialect::Postgres => postgres::lower_statement(statement, ctx)?,
162 Dialect::Sqlite => sqlite::lower_statement(statement, ctx)?,
163 Dialect::Mysql => unreachable!("mysql returned before table recovery"),
164 }
165 }
166 attach_recovered_options(ddl.name.as_ref(), recovered, ctx);
167 Ok(true)
168}
169
170fn attach_recovered_options(
172 name: Option<&SqlObjectName>,
173 recovered: super::table_recovery::RecoveredTableSql,
174 ctx: &mut ParseContext,
175) {
176 let Some(name) = name else {
177 return;
178 };
179 let (name, schema) = object_name_parts(name);
180 let key = schema_qualified_key(&name, schema.as_deref());
181 let Some(table) = ctx.schema.tables.get_mut(&key) else {
182 return;
183 };
184 let mut header = table.options.header_raw.clone();
185 let mut tail = table.options.tail_raw.clone();
186 header.extend(recovered.header_options);
187 tail.extend(recovered.tail_options);
188 table.options = crate::states::TableOptionsMeta::from_parts(header, tail);
189}
190
191fn lower_raw_segment(
192 segment: &SqlSegment,
193 ctx: &mut ParseContext,
194 dialect: Dialect,
195) -> Result<(), ParseError> {
196 let Some(SqlStatementKind::Ddl(ddl)) = &segment.kind else {
197 return Err(ParseError::unsupported(
198 dialect,
199 segment.sql.clone(),
200 "schema SQL must be CREATE statements for known Gaman entity kinds",
201 ));
202 };
203 let Some(name) = &ddl.name else {
204 return Err(ParseError::unsupported(
205 dialect,
206 segment.sql.clone(),
207 "raw fallback requires a recoverable object name",
208 ));
209 };
210 match ddl.entity {
211 EntityKind::Table => Err(ParseError::unsupported(
212 dialect,
213 segment.sql.clone(),
214 "CREATE TABLE must parse into a modeled table; opaque tables are not supported",
215 )),
216 EntityKind::Index => lower_raw_index(segment, ctx, dialect, name, ddl.owner.as_ref()),
217 EntityKind::Trigger => lower_raw_trigger(segment, ctx, dialect, name, ddl.owner.as_ref()),
218 EntityKind::Function => {
219 let (name, schema) = object_name_parts(name);
220 let key = schema_qualified_key(&name, schema.as_deref());
221 let mut function = FunctionDef::from_raw(name, segment.sql.clone());
222 function.schema = schema;
223 ctx.schema.functions.insert(key, function);
224 Ok(())
225 }
226 EntityKind::View => {
227 let (name, schema) = object_name_parts(name);
228 let key = schema_qualified_key(&name, schema.as_deref());
229 let mut view = ViewDef::from_raw(name, segment.sql.clone());
230 view.schema = schema;
231 ctx.schema.views.insert(key, view);
232 Ok(())
233 }
234 EntityKind::Extension => {
235 let (name, schema) = object_name_parts(name);
236 let key = schema_qualified_key(&name, schema.as_deref());
237 let mut extension = ExtensionDef::from_raw(name, segment.sql.clone());
238 extension.schema = schema;
239 ctx.schema.extensions.insert(key, extension);
240 Ok(())
241 }
242 EntityKind::Enum | EntityKind::Column | EntityKind::ForeignKey | EntityKind::Constraint => {
243 Err(ParseError::unsupported(
244 dialect,
245 segment.sql.clone(),
246 "this entity kind must be recovered from modeled SQL, not raw fallback",
247 ))
248 }
249 }
250}
251
252fn lower_raw_index(
253 segment: &SqlSegment,
254 ctx: &mut ParseContext,
255 dialect: Dialect,
256 name: &SqlObjectName,
257 owner: Option<&SqlObjectName>,
258) -> Result<(), ParseError> {
259 let Some(owner) = owner else {
260 return Err(ParseError::unsupported(
261 dialect,
262 segment.sql.clone(),
263 "opaque index fallback requires a recoverable target table after ON",
264 ));
265 };
266 let (table, schema) = object_name_parts(owner);
267 let table_name = schema_qualified_key(&table, schema.as_deref());
268 let (index_name, _) = object_name_parts(name);
269 ctx.push_index((table_name, Index::from_raw(index_name, segment.sql.clone())));
270 Ok(())
271}
272
273fn lower_raw_trigger(
274 segment: &SqlSegment,
275 ctx: &mut ParseContext,
276 dialect: Dialect,
277 name: &SqlObjectName,
278 owner: Option<&SqlObjectName>,
279) -> Result<(), ParseError> {
280 let Some(owner) = owner else {
281 return Err(ParseError::unsupported(
282 dialect,
283 segment.sql.clone(),
284 "opaque trigger fallback requires a recoverable target table after ON",
285 ));
286 };
287 let (table, schema) = object_name_parts(owner);
288 let table_name = schema_qualified_key(&table, schema.as_deref());
289 let table =
290 ctx.schema
291 .tables
292 .get_mut(&table_name)
293 .ok_or_else(|| ParseError::UnknownTriggerTable {
294 table: table_name.clone(),
295 })?;
296 let (trigger_name, _) = object_name_parts(name);
297 table
298 .triggers
299 .push(TriggerDef::from_raw(trigger_name, segment.sql.clone()));
300 Ok(())
301}
302
303fn object_name_parts(name: &SqlObjectName) -> (String, Option<String>) {
304 match name.parts.as_slice() {
305 [schema, name] if schema != "public" => (name.clone(), Some(schema.clone())),
306 [_, name] => (name.clone(), None),
307 [name] => (name.clone(), None),
308 _ => (name.raw.clone(), None),
309 }
310}
311
312fn ensure_modeled_create_statement(stmt: &Statement, dialect: Dialect) -> Result<(), ParseError> {
313 if is_modeled_create_statement(stmt) {
314 return Ok(());
315 }
316 Err(ParseError::unsupported(
317 dialect,
318 stmt.to_string(),
319 "only CREATE statements for modeled schema entities are parsed",
320 ))
321}
322
323fn is_modeled_create_statement(stmt: &Statement) -> bool {
324 matches!(
325 stmt,
326 Statement::CreateExtension(_)
327 | Statement::CreateFunction(_)
328 | Statement::CreateIndex(_)
329 | Statement::CreateTable(_)
330 | Statement::CreateTrigger(_)
331 | Statement::CreateType { .. }
332 | Statement::CreateView(_)
333 )
334}
335
336#[cfg(test)]
337mod recovery_tests {
338 use super::*;
339
340 #[test]
342 fn parse_sql_recovers_unlogged_table() {
343 let schema = parse_sql(
344 "CREATE UNLOGGED TABLE events (id integer NOT NULL)",
345 Dialect::Postgres,
346 )
347 .expect("parse table");
348 let table = schema.tables.get("events").expect("events table");
349 assert_eq!(table.columns.len(), 1);
350 assert_eq!(table.options.header_raw, ["UNLOGGED"]);
351 }
352
353 #[test]
355 fn raw_index_uses_classified_owner() {
356 let schema = parse_sql(
357 "CREATE TABLE app.users (email text); CREATE INDEX users_email_idx ON app.users ((lower(email)))",
358 Dialect::Postgres,
359 )
360 .expect("parse schema");
361 let table = schema.tables.get("app.users").expect("qualified table");
362 assert!(
363 table
364 .indexes
365 .iter()
366 .any(|index| index.name == "users_email_idx")
367 );
368 }
369}