1use crate::trilogy_parser::{Rule, TrilogyParser};
12use pest::iterators::Pair;
13use pest::Parser;
14use std::path::{Path, PathBuf};
15use thiserror::Error;
16
17#[derive(Debug, Clone, PartialEq, Eq, Hash)]
18pub struct ImportStatement {
19 pub raw_path: String,
20 pub parent_dirs: usize,
21 pub alias: Option<String>,
22 pub is_stdlib: bool,
23}
24
25impl ImportStatement {
26 pub fn resolve(&self, working_dir: &Path) -> Option<PathBuf> {
27 if self.is_stdlib {
28 return None;
29 }
30
31 let mut base = working_dir.to_path_buf();
32 for _ in 0..self.parent_dirs {
33 base = base.parent()?.to_path_buf();
34 }
35 for part in self.raw_path.split('.') {
36 base.push(part);
37 }
38 base.set_extension("preql");
39 Some(base)
40 }
41
42 pub fn effective_alias(&self) -> &str {
43 self.alias
44 .as_deref()
45 .unwrap_or_else(|| self.raw_path.split('.').last().unwrap_or(&self.raw_path))
46 }
47}
48
49#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
53pub enum AddressKind {
54 Literal,
56 Templated,
58 Query,
60 File,
62}
63
64impl AddressKind {
65 pub fn as_str(&self) -> &'static str {
66 match self {
67 AddressKind::Literal => "literal",
68 AddressKind::Templated => "templated",
69 AddressKind::Query => "query",
70 AddressKind::File => "file",
71 }
72 }
73}
74
75impl std::fmt::Display for AddressKind {
76 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
77 write!(f, "{}", self.as_str())
78 }
79}
80
81#[derive(Debug, Clone, PartialEq, Eq, Hash)]
82pub struct DatasourceDeclaration {
83 pub name: String,
84 pub address: Option<String>,
87 pub address_kind: AddressKind,
88 pub is_root: bool,
90 pub is_partitioned: bool,
92}
93
94#[derive(Debug, Clone, PartialEq, Eq, Hash)]
95pub struct PersistStatement {
96 pub mode: PersistMode,
97 pub target_datasource: String,
98}
99
100#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
101pub enum PersistMode {
102 Append,
103 Overwrite,
104 Persist,
105}
106
107impl std::fmt::Display for PersistMode {
108 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
109 match self {
110 PersistMode::Append => write!(f, "append"),
111 PersistMode::Overwrite => write!(f, "overwrite"),
112 PersistMode::Persist => write!(f, "persist"),
113 }
114 }
115}
116
117#[derive(Debug, Clone, Default)]
118pub struct ParsedFile {
119 pub imports: Vec<ImportStatement>,
120 pub datasources: Vec<DatasourceDeclaration>,
121 pub persists: Vec<PersistStatement>,
122}
123
124#[derive(Error, Debug)]
125pub enum ParseError {
126 #[error("Failed to parse file: {0}")]
127 PestError(#[from] pest::error::Error<Rule>),
128
129 #[error("Invalid import statement structure")]
130 InvalidImportStructure,
131
132 #[error("Invalid datasource statement structure")]
133 InvalidDatasourceStructure,
134
135 #[error("Invalid persist statement structure")]
136 InvalidPersistStructure,
137}
138
139pub fn parse_file(content: &str) -> Result<ParsedFile, ParseError> {
140 let mut pairs = TrilogyParser::parse(Rule::start, content)?;
141 let start = pairs
142 .next()
143 .ok_or(ParseError::InvalidImportStructure)?;
144
145 let mut result = ParsedFile::default();
146 for top in start.into_inner() {
147 if top.as_rule() != Rule::block {
148 continue;
149 }
150 for stmt in top.into_inner() {
154 match stmt.as_rule() {
155 Rule::import_statement | Rule::selective_import_statement => {
161 result.imports.push(extract_import(stmt)?);
162 }
163 Rule::datasource => {
164 result.datasources.push(extract_datasource(stmt)?);
165 }
166 Rule::persist_statement => {
167 result.persists.push(extract_persist(stmt)?);
168 }
169 _ => {}
170 }
171 }
172 }
173
174 Ok(result)
175}
176
177pub fn parse_imports(content: &str) -> Result<Vec<ImportStatement>, ParseError> {
178 Ok(parse_file(content)?.imports)
179}
180
181fn extract_import(pair: Pair<Rule>) -> Result<ImportStatement, ParseError> {
185 let full_text = pair.as_str();
186 let mut n_dots = 0usize;
187 let mut idents: Vec<String> = Vec::new();
188 for child in pair.into_inner() {
189 match child.as_rule() {
190 Rule::IMPORT_DOT => n_dots += 1,
191 Rule::IDENTIFIER => idents.push(child.as_str().to_string()),
192 _ => {}
193 }
194 }
195 if idents.is_empty() {
196 return Err(ParseError::InvalidImportStructure);
197 }
198
199 let has_alias = full_text
202 .split_ascii_whitespace()
203 .any(|tok| tok.eq_ignore_ascii_case("as"));
204 let alias = if has_alias && idents.len() >= 2 {
205 Some(idents.pop().unwrap())
206 } else {
207 None
208 };
209
210 let raw_path = idents.join(".");
211 let is_stdlib = raw_path == "std" || raw_path.starts_with("std.");
212 let parent_dirs = n_dots.saturating_sub(1);
216
217 Ok(ImportStatement {
218 raw_path,
219 parent_dirs,
220 alias,
221 is_stdlib,
222 })
223}
224
225fn extract_datasource(pair: Pair<Rule>) -> Result<DatasourceDeclaration, ParseError> {
230 let mut name: Option<String> = None;
231 let mut address: Option<String> = None;
232 let mut address_kind: Option<AddressKind> = None;
233 let mut is_root = false;
234 let mut is_partitioned = false;
235
236 for child in pair.into_inner() {
237 match child.as_rule() {
238 Rule::DATASOURCE_ROOT => is_root = true,
239 Rule::IDENTIFIER if name.is_none() => name = Some(child.as_str().to_string()),
240 Rule::address => {
242 let tok = child
243 .into_inner()
244 .next()
245 .ok_or(ParseError::InvalidDatasourceStructure)?;
246 match tok.as_rule() {
247 Rule::F_QUOTED_ADDRESS => {
248 address_kind = Some(AddressKind::Templated);
251 address = Some(
252 tok.as_str()
253 .trim_start_matches(['f', 'F'])
254 .trim_matches('`')
255 .to_string(),
256 );
257 }
258 Rule::QUOTED_ADDRESS => {
259 address_kind = Some(AddressKind::Literal);
261 address = Some(
262 tok.as_str().trim_matches('`').trim_matches('\'').to_string(),
263 );
264 }
265 _ => {
266 address_kind = Some(AddressKind::Literal);
267 address = Some(tok.as_str().to_string());
268 }
269 }
270 }
271 Rule::query => address_kind = Some(AddressKind::Query),
272 Rule::file => {
273 address_kind = Some(AddressKind::File);
274 address = Some(child.as_str()[4..].trim().to_string());
276 }
277 Rule::datasource_partition_clause => is_partitioned = true,
278 _ => {}
279 }
280 }
281
282 match (name, address_kind) {
283 (Some(name), Some(address_kind)) => Ok(DatasourceDeclaration {
284 name,
285 address,
286 address_kind,
287 is_root,
288 is_partitioned,
289 }),
290 _ => Err(ParseError::InvalidDatasourceStructure),
291 }
292}
293
294fn extract_persist(pair: Pair<Rule>) -> Result<PersistStatement, ParseError> {
295 let inner = pair
297 .into_inner()
298 .next()
299 .ok_or(ParseError::InvalidPersistStructure)?;
300 match inner.as_rule() {
301 Rule::auto_persist => extract_auto_persist(inner),
302 Rule::full_persist => extract_full_persist(inner),
303 _ => Err(ParseError::InvalidPersistStructure),
304 }
305}
306
307fn extract_auto_persist(pair: Pair<Rule>) -> Result<PersistStatement, ParseError> {
309 let mut mode: Option<PersistMode> = None;
310 let mut target: Option<String> = None;
311 for child in pair.into_inner() {
312 match child.as_rule() {
313 Rule::PERSIST_MODE => mode = Some(parse_persist_mode(child.as_str())),
314 Rule::IDENTIFIER if target.is_none() => {
315 target = Some(child.as_str().to_string());
316 }
317 _ => {}
318 }
319 }
320 match (mode, target) {
321 (Some(mode), Some(target_datasource)) => Ok(PersistStatement {
322 mode,
323 target_datasource,
324 }),
325 _ => Err(ParseError::InvalidPersistStructure),
326 }
327}
328
329fn extract_full_persist(pair: Pair<Rule>) -> Result<PersistStatement, ParseError> {
334 let mut mode: Option<PersistMode> = None;
335 let mut last_ident: Option<String> = None;
336 for child in pair.into_inner() {
337 match child.as_rule() {
338 Rule::PERSIST_MODE => mode = Some(parse_persist_mode(child.as_str())),
339 Rule::IDENTIFIER => last_ident = Some(child.as_str().to_string()),
340 _ => {}
341 }
342 }
343 match (mode, last_ident) {
344 (Some(mode), Some(target_datasource)) => Ok(PersistStatement {
345 mode,
346 target_datasource,
347 }),
348 _ => Err(ParseError::InvalidPersistStructure),
349 }
350}
351
352fn parse_persist_mode(s: &str) -> PersistMode {
353 match s.to_ascii_lowercase().as_str() {
354 "append" => PersistMode::Append,
355 "overwrite" => PersistMode::Overwrite,
356 _ => PersistMode::Persist,
357 }
358}
359
360#[cfg(test)]
361mod tests {
362 use super::*;
363
364 #[test]
365 fn test_simple_import() {
366 let parsed = parse_file("import models.customer;").unwrap();
367 assert_eq!(parsed.imports.len(), 1);
368 assert_eq!(parsed.imports[0].raw_path, "models.customer");
369 assert_eq!(parsed.imports[0].parent_dirs, 0);
370 assert!(parsed.imports[0].alias.is_none());
371 }
372
373 #[test]
374 fn test_import_with_alias() {
375 let parsed = parse_file("import models.customer as cust;").unwrap();
376 assert_eq!(parsed.imports.len(), 1);
377 assert_eq!(parsed.imports[0].raw_path, "models.customer");
378 assert_eq!(parsed.imports[0].alias, Some("cust".to_string()));
379 }
380
381 #[test]
382 fn test_relative_import() {
383 let parsed = parse_file("import ..models.customer;").unwrap();
384 assert_eq!(parsed.imports.len(), 1);
385 assert_eq!(parsed.imports[0].raw_path, "models.customer");
386 assert_eq!(parsed.imports[0].parent_dirs, 1);
387 }
388
389 #[test]
390 fn test_sibling_relative_import() {
391 let parsed = parse_file("import .customer;").unwrap();
392 assert_eq!(parsed.imports[0].raw_path, "customer");
393 assert_eq!(parsed.imports[0].parent_dirs, 0);
394 }
395
396 #[test]
397 fn test_stdlib_import() {
398 let parsed = parse_file("import std.aggregates;").unwrap();
399 assert!(parsed.imports[0].is_stdlib);
400 }
401
402 #[test]
403 fn test_datasource_simple() {
404 let content = r#"
405 key order_id int;
406 datasource orders (
407 order_id: order_id,
408 amount: amount
409 )
410 grain (order_id)
411 address my_database.orders;
412 "#;
413 let parsed = parse_file(content).unwrap();
414 assert_eq!(parsed.datasources.len(), 1);
415 let ds = &parsed.datasources[0];
416 assert_eq!(ds.name, "orders");
417 assert_eq!(ds.address.as_deref(), Some("my_database.orders"));
418 assert_eq!(ds.address_kind, AddressKind::Literal);
419 assert!(!ds.is_root);
420 assert!(!ds.is_partitioned);
421 }
422
423 #[test]
424 fn test_datasource_with_quoted_address() {
425 let content = r#"
426 key customer_id int;
427 datasource customers (
428 id: customer_id,
429 name: customer_name
430 )
431 grain (customer_id)
432 address `my_db.customers`;
433 "#;
434 let parsed = parse_file(content).unwrap();
435 assert_eq!(parsed.datasources.len(), 1);
436 let ds = &parsed.datasources[0];
437 assert_eq!(ds.name, "customers");
438 assert_eq!(ds.address.as_deref(), Some("my_db.customers"));
439 assert_eq!(ds.address_kind, AddressKind::Literal);
440 }
441
442 #[test]
443 fn test_root_partitioned_datasource() {
444 let content = r#"
445 key event_id int;
446 root datasource events (
447 event_id: event_id
448 )
449 grain (event_id)
450 address analytics.events
451 partition by event_id;
452 "#;
453 let parsed = parse_file(content).unwrap();
454 let ds = &parsed.datasources[0];
455 assert!(ds.is_root);
456 assert!(ds.is_partitioned);
457 assert_eq!(ds.address.as_deref(), Some("analytics.events"));
458 }
459
460 #[test]
461 fn test_templated_address_datasource() {
462 let content = r#"
463 key order_id int;
464 datasource orders (
465 order_id: order_id
466 )
467 grain (order_id)
468 address f`{{env}}.orders`;
469 "#;
470 let parsed = parse_file(content).unwrap();
471 let ds = &parsed.datasources[0];
472 assert_eq!(ds.address_kind, AddressKind::Templated);
473 assert_eq!(ds.address.as_deref(), Some("{{env}}.orders"));
475 }
476
477 #[test]
478 fn test_query_datasource_has_no_address() {
479 let content = r#"
480 key order_id int;
481 datasource order_view (
482 order_id: order_id
483 )
484 grain (order_id)
485 query '''select 1 as order_id''';
486 "#;
487 let parsed = parse_file(content).unwrap();
488 let ds = &parsed.datasources[0];
489 assert_eq!(ds.address_kind, AddressKind::Query);
490 assert!(ds.address.is_none());
491 }
492
493 #[test]
494 fn test_selective_import_creates_edge() {
495 let parsed = parse_file("from models.customer import customer_id;").unwrap();
496 assert_eq!(parsed.imports.len(), 1);
497 assert_eq!(parsed.imports[0].raw_path, "models.customer");
498 assert!(parsed.imports[0].alias.is_none());
499 }
500
501 #[test]
502 fn test_selective_import_with_alias() {
503 let parsed = parse_file("from ..models.customer as cust import customer_id, name;").unwrap();
504 assert_eq!(parsed.imports.len(), 1);
505 assert_eq!(parsed.imports[0].raw_path, "models.customer");
506 assert_eq!(parsed.imports[0].alias, Some("cust".to_string()));
507 assert_eq!(parsed.imports[0].parent_dirs, 1);
508 }
509
510 #[test]
511 fn test_self_import_is_not_an_edge() {
512 let parsed = parse_file("self import as me;").unwrap();
513 assert!(parsed.imports.is_empty());
514 }
515
516 #[test]
517 fn test_auto_persist() {
518 let parsed = parse_file("persist orders;").unwrap();
519 assert_eq!(parsed.persists.len(), 1);
520 assert_eq!(parsed.persists[0].target_datasource, "orders");
521 assert_eq!(parsed.persists[0].mode, PersistMode::Persist);
522 }
523
524 #[test]
525 fn test_append_auto_persist() {
526 let parsed = parse_file("append orders;").unwrap();
527 assert_eq!(parsed.persists.len(), 1);
528 assert_eq!(parsed.persists[0].target_datasource, "orders");
529 assert_eq!(parsed.persists[0].mode, PersistMode::Append);
530 }
531
532 #[test]
533 fn test_full_persist() {
534 let content = r#"
535 key order_id int;
536 overwrite into target_orders from select order_id;
537 "#;
538 let parsed = parse_file(content).unwrap();
539 assert_eq!(parsed.persists.len(), 1);
540 assert_eq!(parsed.persists[0].target_datasource, "target_orders");
541 assert_eq!(parsed.persists[0].mode, PersistMode::Overwrite);
542 }
543
544 #[test]
545 fn test_multiple_imports() {
546 let content = r#"
547 import models.customer;
548 import models.orders as ord;
549 // comment
550 import ..shared.utils;
551 "#;
552 let parsed = parse_file(content).unwrap();
553 assert_eq!(parsed.imports.len(), 3);
554 assert_eq!(parsed.imports[1].alias, Some("ord".to_string()));
555 assert_eq!(parsed.imports[2].parent_dirs, 1);
556 }
557
558 #[test]
559 fn test_mixed_file() {
560 let content = r#"
561 import models.customer;
562
563 key order_id int;
564 datasource local_orders (
565 order_id: order_id
566 )
567 grain (order_id)
568 address local.orders;
569
570 persist local_orders;
571 "#;
572 let parsed = parse_file(content).unwrap();
573 assert_eq!(parsed.imports.len(), 1);
574 assert_eq!(parsed.datasources.len(), 1);
575 assert_eq!(parsed.persists.len(), 1);
576 assert_eq!(parsed.datasources[0].name, "local_orders");
577 assert_eq!(parsed.persists[0].target_datasource, "local_orders");
578 }
579}