1use crate::comprehension::cardinality::{Interval, MeasureName, ProductMeasure};
34use crate::comprehension::source::{LiteralValue, Source};
35
36pub fn parse_source(text: &str) -> Result<Source, SourceParseError> {
38 let trimmed = text.trim();
39
40 if let Some(name) = strip_curly(trimmed) {
46 return Ok(Source::WorkloadParamList {
47 name,
48 len_hint: None,
49 });
50 }
51 if let Some(dyn_text) = strip_dynamic_curly(trimmed) {
52 return Ok(Source::WorkloadParamList {
53 name: dyn_text,
54 len_hint: None,
55 });
56 }
57
58 if trimmed.len() >= 2 && trimmed.starts_with('"') && trimmed.ends_with('"') {
67 let inner = &trimmed[1..trimmed.len() - 1];
68 let values = super::super::source::split_string_comprehension(inner)
69 .into_iter()
70 .map(parse_literal_value)
71 .collect();
72 return Ok(Source::Literal { values });
73 }
74 if trimmed.len() >= 2 && trimmed.starts_with('\'') && trimmed.ends_with('\'') {
75 let inner = &trimmed[1..trimmed.len() - 1];
76 return Ok(Source::Literal {
77 values: vec![LiteralValue::String(inner.to_string())],
78 });
79 }
80
81 if trimmed.starts_with('[') && trimmed.ends_with(']') {
92 let inner = &trimmed[1..trimmed.len() - 1];
93 if bracket_is_pure_literal(inner) {
94 return parse_literal_list(inner);
95 }
96 return Ok(Source::Generator {
97 expr: trimmed.to_string(),
98 cardinality_hint: None,
99 });
100 }
101
102 if let Some(result) = parse_distribution(trimmed) {
108 return result;
109 }
110
111 if let Some(idx) = find_top_level(trimmed, "..") {
113 return parse_range(trimmed, idx);
114 }
115
116 if looks_like_function_call(trimmed) {
118 return Ok(Source::Generator {
119 expr: trimmed.to_string(),
120 cardinality_hint: None,
121 });
122 }
123
124 if let Some(value) = try_parse_bare_scalar(trimmed) {
128 return Ok(Source::Literal {
129 values: vec![value],
130 });
131 }
132
133 if trimmed.contains(',') && looks_like_bare_value_list(trimmed) {
140 return parse_literal_list(trimmed);
141 }
142
143 Ok(Source::Generator {
149 expr: trimmed.to_string(),
150 cardinality_hint: None,
151 })
152}
153
154fn parse_distribution(text: &str) -> Option<Result<Source, SourceParseError>> {
165 let (call, support_text) = match split_on_keyword(text, " on ") {
166 Some((call, rest)) => (call, Some(rest)),
167 None => (text, None),
168 };
169 let (name, args) = split_call(call)?;
170 let measure = MeasureName::from_text(name)?;
171 Some(build_distribution(text, measure, args, support_text))
172}
173
174fn build_distribution(
177 text: &str,
178 measure: MeasureName,
179 args: &str,
180 support_text: Option<&str>,
181) -> Result<Source, SourceParseError> {
182 let invalid = || SourceParseError::InvalidRange(text.to_string());
183 let mut params = Vec::new();
184 for arg in args.split(',') {
185 let arg = arg.trim();
186 if arg.is_empty() {
187 if params.is_empty() && args.trim().is_empty() {
188 break;
189 }
190 return Err(invalid());
191 }
192 params.push(arg.parse::<f64>().map_err(|_| invalid())?);
193 }
194 let params = measure.resolve_params(¶ms).map_err(|_| invalid())?;
195 let support = match support_text {
196 None => measure.support(¶ms),
197 Some(interval_text) => match parse_range_text(interval_text)? {
198 Source::ContinuousInterval { interval, .. } => interval,
199 _ => return Err(invalid()),
200 },
201 };
202 Ok(Source::Distribution {
203 distribution: measure,
204 support,
205 params,
206 })
207}
208
209fn split_call(text: &str) -> Option<(&str, &str)> {
212 let text = text.trim();
213 let open = text.find('(')?;
214 let inner = text.strip_suffix(')')?;
215 let name = text[..open].trim();
216 if name.is_empty() || !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
217 return None;
218 }
219 Some((name, &inner[open + 1..]))
220}
221
222fn split_on_keyword<'a>(text: &'a str, keyword: &str) -> Option<(&'a str, &'a str)> {
225 let bytes = text.as_bytes();
226 let mut depth = 0i32;
227 let mut quote: Option<u8> = None;
228 for i in 0..bytes.len() {
229 let c = bytes[i];
230 match quote {
231 Some(q) => {
232 if c == q {
233 quote = None;
234 }
235 }
236 None => match c {
237 b'"' | b'\'' => quote = Some(c),
238 b'(' | b'[' | b'{' => depth += 1,
239 b')' | b']' | b'}' => depth -= 1,
240 _ => {
241 if depth == 0 && text[i..].starts_with(keyword) {
242 return Some((&text[..i], &text[i + keyword.len()..]));
243 }
244 }
245 },
246 }
247 }
248 None
249}
250
251fn looks_like_bare_value_list(text: &str) -> bool {
257 !text.chars().any(|c| {
258 matches!(
259 c,
260 '(' | ')'
261 | '['
262 | ']'
263 | '{'
264 | '}'
265 | '\''
266 | '"'
267 | '+'
268 | '*'
269 | '/'
270 | '%'
271 | '='
272 | '<'
273 | '>'
274 | '!'
275 | '&'
276 | '|'
277 | '~'
278 | '^'
279 | '?'
280 )
281 })
282}
283
284fn strip_dynamic_curly(s: &str) -> Option<String> {
288 let s = s.trim();
289 if !s.starts_with('{') || !s.ends_with('}') {
290 return None;
291 }
292 let inner = &s[1..s.len() - 1];
293 if !inner.contains('{') {
297 return None;
298 }
299 Some(inner.to_string())
300}
301
302fn parse_literal_list(inner: &str) -> Result<Source, SourceParseError> {
306 let parts: Vec<&str> = inner
307 .split(',')
308 .map(|s| s.trim())
309 .filter(|s| !s.is_empty())
310 .collect();
311
312 if parts.is_empty() {
313 return Ok(Source::Literal { values: Vec::new() });
314 }
315
316 let values: Vec<LiteralValue> = parts.iter().map(|s| parse_literal_value(s)).collect();
317
318 Ok(Source::Literal { values })
319}
320
321fn bracket_is_pure_literal(inner: &str) -> bool {
328 let elems: Vec<&str> = inner
329 .split(',')
330 .map(str::trim)
331 .filter(|s| !s.is_empty())
332 .collect();
333 if elems.is_empty() {
334 return true; }
336 elems.iter().all(|e| {
337 if e.ends_with('…') || e.ends_with("...") {
338 return false; }
340 e.eq_ignore_ascii_case("true")
341 || e.eq_ignore_ascii_case("false")
342 || ((e.starts_with('"') && e.ends_with('"'))
343 || (e.starts_with('\'') && e.ends_with('\'')))
344 || e.parse::<i64>().is_ok()
345 || e.parse::<f64>().is_ok()
346 })
347}
348
349fn parse_literal_value(s: &str) -> LiteralValue {
350 let s = s.trim();
351 if s.eq_ignore_ascii_case("true") {
352 return LiteralValue::Bool(true);
353 }
354 if s.eq_ignore_ascii_case("false") {
355 return LiteralValue::Bool(false);
356 }
357 if (s.starts_with('"') && s.ends_with('"')) || (s.starts_with('\'') && s.ends_with('\'')) {
359 let inner = &s[1..s.len() - 1];
360 return LiteralValue::String(inner.to_string());
361 }
362 if let Ok(n) = s.parse::<i64>() {
364 return LiteralValue::Int(n);
365 }
366 if let Ok(n) = s.parse::<u64>() {
367 return LiteralValue::UInt(n);
368 }
369 if let Ok(f) = s.parse::<f64>() {
371 return LiteralValue::Float(f);
372 }
373 LiteralValue::String(s.to_string())
375}
376
377fn parse_range_text(text: &str) -> Result<Source, SourceParseError> {
380 let text = text.trim();
381 match find_top_level(text, "..") {
382 Some(idx) => parse_range(text, idx),
383 None => Err(SourceParseError::InvalidRange(text.to_string())),
384 }
385}
386
387fn parse_range(text: &str, dotdot_idx: usize) -> Result<Source, SourceParseError> {
388 let lo_str = text[..dotdot_idx].trim();
389 let after = &text[dotdot_idx + 2..];
390
391 let (inclusive_end, after) = if let Some(rest) = after.strip_prefix('=') {
393 (true, rest)
394 } else {
395 (false, after)
396 };
397
398 let (rhs, step) = if let Some(step_pos) = after.find(" step ") {
403 let rhs = after[..step_pos].trim();
404 let step_str = after[step_pos + 6..].trim();
405 let step: i64 = step_str
406 .parse()
407 .map_err(|_| SourceParseError::InvalidRange(text.to_string()))?;
408 (rhs, step)
409 } else if let Some(step_pos) = after.find("..") {
410 let rhs = after[..step_pos].trim();
413 let step_str = after[step_pos + 2..].trim();
414 let step: i64 = step_str
415 .parse()
416 .map_err(|_| SourceParseError::InvalidRange(text.to_string()))?;
417 (rhs, step)
418 } else {
419 (after.trim(), 1)
420 };
421
422 if let (Ok(lo_i), Ok(hi_i)) = (lo_str.parse::<i64>(), rhs.parse::<i64>()) {
424 let hi = if inclusive_end { hi_i + 1 } else { hi_i };
425 return Ok(Source::IntRange { lo: lo_i, hi, step });
426 }
427 if let (Ok(lo_f), Ok(hi_f)) = (lo_str.parse::<f64>(), rhs.parse::<f64>()) {
429 let interval = Interval {
430 lo: lo_f,
431 hi: hi_f,
432 lo_open: false,
433 hi_open: !inclusive_end,
434 };
435 return Ok(Source::ContinuousInterval {
436 interval,
437 measure: ProductMeasure::Uniform,
438 });
439 }
440
441 Err(SourceParseError::InvalidRange(text.to_string()))
442}
443
444fn strip_curly(s: &str) -> Option<String> {
445 let s = s.trim();
446 if s.starts_with('{') && s.ends_with('}') {
447 let inner = &s[1..s.len() - 1];
448 let trimmed = inner.trim();
449 if !trimmed.is_empty() && trimmed.chars().all(|c| c.is_alphanumeric() || c == '_') {
450 return Some(trimmed.to_string());
451 }
452 }
453 None
454}
455
456fn try_parse_bare_scalar(s: &str) -> Option<LiteralValue> {
462 if s.eq_ignore_ascii_case("true") {
463 return Some(LiteralValue::Bool(true));
464 }
465 if s.eq_ignore_ascii_case("false") {
466 return Some(LiteralValue::Bool(false));
467 }
468 if (s.starts_with('"') && s.ends_with('"')) || (s.starts_with('\'') && s.ends_with('\'')) {
469 let inner = &s[1..s.len() - 1];
470 return Some(LiteralValue::String(inner.to_string()));
471 }
472 if let Ok(n) = s.parse::<i64>() {
473 return Some(LiteralValue::Int(n));
474 }
475 if let Ok(n) = s.parse::<u64>() {
476 return Some(LiteralValue::UInt(n));
477 }
478 if let Ok(f) = s.parse::<f64>() {
479 return Some(LiteralValue::Float(f));
480 }
481 None
482}
483
484fn looks_like_function_call(s: &str) -> bool {
485 let Some(open) = s.find('(') else {
486 return false;
487 };
488 if !s.ends_with(')') {
489 return false;
490 }
491 let name = &s[..open];
492 !name.is_empty() && name.chars().all(|c| c.is_alphanumeric() || c == '_')
493}
494
495fn find_top_level(s: &str, needle: &str) -> Option<usize> {
498 let bytes = s.as_bytes();
499 let needle_bytes = needle.as_bytes();
500 let mut depth = 0i64;
501 let mut i = 0;
502 while i + needle_bytes.len() <= bytes.len() {
503 match bytes[i] {
504 b'(' | b'[' | b'{' => depth += 1,
505 b')' | b']' | b'}' => depth -= 1,
506 _ => {}
507 }
508 if depth == 0 && &bytes[i..i + needle_bytes.len()] == needle_bytes {
509 return Some(i);
510 }
511 i += 1;
512 }
513 None
514}
515
516#[derive(Debug, Clone, PartialEq)]
518pub enum SourceParseError {
519 Unrecognized(String),
521 InvalidRange(String),
524}
525
526impl std::fmt::Display for SourceParseError {
527 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
528 match self {
529 SourceParseError::Unrecognized(s) => {
530 write!(f, "unrecognized source expression: {s:?}")
531 }
532 SourceParseError::InvalidRange(s) => {
533 write!(f, "invalid range expression: {s:?}")
534 }
535 }
536 }
537}
538
539impl std::error::Error for SourceParseError {}
540
541#[cfg(test)]
542mod tests {
543 use super::*;
544
545 #[test]
546 fn int_range_exclusive() {
547 let s = parse_source("1..10").unwrap();
548 assert!(matches!(
549 s,
550 Source::IntRange {
551 lo: 1,
552 hi: 10,
553 step: 1
554 }
555 ));
556 }
557
558 #[test]
559 fn double_quoted_source_is_string_comprehension_striped() {
560 let s = parse_source(r#""rerank_def, rerank_1x, rerank_2x""#).unwrap();
562 match s {
563 Source::Literal { values } => {
564 assert_eq!(
565 values,
566 vec![
567 LiteralValue::String("rerank_def".into()),
568 LiteralValue::String("rerank_1x".into()),
569 LiteralValue::String("rerank_2x".into()),
570 ]
571 );
572 }
573 other => panic!("expected striped Literal, got {other:?}"),
574 }
575 }
576
577 #[test]
578 fn single_quoted_source_is_atomic() {
579 let s = parse_source("'rerank_def, rerank_1x'").unwrap();
581 match s {
582 Source::Literal { values } => {
583 assert_eq!(
584 values,
585 vec![LiteralValue::String("rerank_def, rerank_1x".into())]
586 );
587 }
588 other => panic!("expected atomic Literal, got {other:?}"),
589 }
590 }
591
592 #[test]
593 fn int_range_inclusive() {
594 let s = parse_source("1..=10").unwrap();
595 assert!(matches!(
596 s,
597 Source::IntRange {
598 lo: 1,
599 hi: 11,
600 step: 1
601 }
602 ));
603 }
604
605 #[test]
606 fn int_range_with_step() {
607 let s = parse_source("0..100 step 10").unwrap();
608 assert!(matches!(
609 s,
610 Source::IntRange {
611 lo: 0,
612 hi: 100,
613 step: 10
614 }
615 ));
616 }
617
618 #[test]
619 fn literal_int_list() {
620 let s = parse_source("[1, 2, 3]").unwrap();
621 match s {
622 Source::Literal { values } => {
623 assert_eq!(values.len(), 3);
624 assert_eq!(values[0], LiteralValue::Int(1));
625 assert_eq!(values[2], LiteralValue::Int(3));
626 }
627 other => panic!("expected Literal, got {other:?}"),
628 }
629 }
630
631 #[test]
632 fn bracket_bare_words_are_references_not_strings() {
633 let s = parse_source("[a, b, c]").unwrap();
639 match s {
640 Source::Generator { expr, .. } => assert_eq!(expr, "[a, b, c]"),
641 other => panic!("expected deferred Generator, got {other:?}"),
642 }
643 }
644
645 #[test]
646 fn bracket_with_spread_defers_to_generator() {
647 let s = parse_source("[xs…]").unwrap();
648 assert!(
649 matches!(s, Source::Generator { .. }),
650 "spread list must defer: {s:?}"
651 );
652 }
653
654 #[test]
655 fn literal_quoted_strings() {
656 let s = parse_source(r#"["hello", "world"]"#).unwrap();
657 match s {
658 Source::Literal { values } => {
659 assert_eq!(values[0], LiteralValue::String("hello".into()));
660 assert_eq!(values[1], LiteralValue::String("world".into()));
661 }
662 other => panic!("expected Literal, got {other:?}"),
663 }
664 }
665
666 #[test]
667 fn literal_float_list() {
668 let s = parse_source("[1.5, 2.5, 3.5]").unwrap();
669 match s {
670 Source::Literal { values } => {
671 assert_eq!(values[0], LiteralValue::Float(1.5));
672 }
673 other => panic!("expected Literal, got {other:?}"),
674 }
675 }
676
677 #[test]
678 fn workload_param_ref() {
679 let s = parse_source("{profiles}").unwrap();
680 match s {
681 Source::WorkloadParamList { name, .. } => assert_eq!(name, "profiles"),
682 other => panic!("expected WorkloadParamList, got {other:?}"),
683 }
684 }
685
686 #[test]
687 fn generator_function_call() {
688 let s = parse_source("fib(8)").unwrap();
689 match s {
690 Source::Generator { expr, .. } => assert_eq!(expr, "fib(8)"),
691 other => panic!("expected Generator, got {other:?}"),
692 }
693 }
694
695 #[test]
696 fn continuous_interval_via_floats() {
697 let s = parse_source("0.0..1.0").unwrap();
698 match s {
699 Source::ContinuousInterval { interval, measure } => {
700 assert_eq!(interval.lo, 0.0);
701 assert_eq!(interval.hi, 1.0);
702 assert!(matches!(measure, ProductMeasure::Uniform));
703 }
704 other => panic!("expected ContinuousInterval, got {other:?}"),
705 }
706 }
707
708 #[test]
709 fn continuous_interval_inclusive() {
710 let s = parse_source("0.0..=1.0").unwrap();
711 match s {
712 Source::ContinuousInterval { interval, .. } => {
713 assert!(!interval.hi_open);
714 }
715 other => panic!("expected ContinuousInterval, got {other:?}"),
716 }
717 }
718
719 #[test]
720 fn unrecognized_source_falls_back_to_generator() {
721 let s = parse_source("totally nonsense").unwrap();
726 match s {
727 Source::Generator { expr, .. } => assert_eq!(expr, "totally nonsense"),
728 other => panic!("expected Generator, got {other:?}"),
729 }
730 }
731
732 #[test]
733 fn empty_literal_list() {
734 let s = parse_source("[]").unwrap();
735 match s {
736 Source::Literal { values } => assert!(values.is_empty()),
737 other => panic!("expected empty Literal, got {other:?}"),
738 }
739 }
740
741 #[test]
744 fn a_named_measure_parses_as_a_distribution() {
745 match parse_source("normal(0, 1)").unwrap() {
746 Source::Distribution {
747 distribution: MeasureName::Normal,
748 support,
749 params,
750 } => {
751 assert_eq!(params, vec![0.0, 1.0]);
752 assert!(support.lo.is_infinite() && support.hi.is_infinite());
753 }
754 other => panic!("expected a normal distribution, got {other:?}"),
755 }
756 match parse_source("exponential(2) on 0.0..1.0").unwrap() {
757 Source::Distribution {
758 distribution: MeasureName::Exponential,
759 support,
760 params,
761 } => {
762 assert_eq!(params, vec![2.0]);
763 assert_eq!(support, Interval::half_open(0.0, 1.0));
764 }
765 other => panic!("expected a restricted exponential, got {other:?}"),
766 }
767 match parse_source("pareto(3, 2)").unwrap() {
770 Source::Distribution {
771 support, params, ..
772 } => {
773 assert_eq!(params, vec![3.0, 2.0]);
774 assert_eq!(support.lo, 3.0);
775 }
776 other => panic!("{other:?}"),
777 }
778 assert!(matches!(
779 parse_source("uniform01()").unwrap(),
780 Source::Distribution {
781 distribution: MeasureName::Uniform01,
782 ..
783 }
784 ));
785 }
786
787 #[test]
790 fn a_call_that_names_no_measure_is_still_a_generator() {
791 for text in ["fib(8)", "partitions(\"*/4\", 100)", "range(0, 10)"] {
792 assert!(
793 matches!(parse_source(text).unwrap(), Source::Generator { .. }),
794 "{text}"
795 );
796 }
797 assert!(matches!(
799 parse_source("0.0..1.0").unwrap(),
800 Source::ContinuousInterval { .. }
801 ));
802 }
803
804 #[test]
807 fn a_measure_with_the_wrong_arguments_is_an_error() {
808 for text in ["normal(1)", "normal(0, 1, 2)", "beta(a, b)", "gamma(1,)"] {
809 assert!(parse_source(text).is_err(), "{text}");
810 }
811 }
812}