1use crate::comprehension::cardinality::{Interval, MeasureName, ProductMeasure};
36use crate::comprehension::source::{LiteralValue, Source};
37
38pub fn parse_source(text: &str) -> Result<Source, SourceParseError> {
40 let trimmed = text.trim();
41
42 if let Some(name) = strip_curly(trimmed) {
48 return Ok(Source::WorkloadParamList {
49 name,
50 len_hint: None,
51 });
52 }
53 if let Some(dyn_text) = strip_dynamic_curly(trimmed) {
54 return Ok(Source::WorkloadParamList {
55 name: dyn_text,
56 len_hint: None,
57 });
58 }
59
60 if trimmed.len() >= 2 && trimmed.starts_with('"') && trimmed.ends_with('"') {
69 let inner = &trimmed[1..trimmed.len() - 1];
70 let values = super::super::source::split_string_comprehension(inner)
71 .into_iter()
72 .map(parse_literal_value)
73 .collect();
74 return Ok(Source::Literal { values });
75 }
76 if trimmed.len() >= 2 && trimmed.starts_with('\'') && trimmed.ends_with('\'') {
77 let inner = &trimmed[1..trimmed.len() - 1];
78 return Ok(Source::Literal {
79 values: vec![LiteralValue::String(inner.to_string())],
80 });
81 }
82
83 if trimmed.starts_with('[') && trimmed.ends_with(']') {
94 let inner = &trimmed[1..trimmed.len() - 1];
95 if bracket_is_pure_literal(inner) {
96 return parse_literal_list(inner);
97 }
98 return Ok(Source::Generator {
99 expr: trimmed.to_string(),
100 cardinality_hint: None,
101 });
102 }
103
104 if let Some(result) = parse_distribution(trimmed) {
109 return result;
110 }
111
112 if let Some(idx) = find_top_level(trimmed, "..") {
114 return parse_range(trimmed, idx);
115 }
116
117 if looks_like_function_call(trimmed) {
119 return Ok(Source::Generator {
120 expr: trimmed.to_string(),
121 cardinality_hint: None,
122 });
123 }
124
125 if let Some(value) = try_parse_bare_scalar(trimmed) {
131 return Ok(Source::Literal {
132 values: vec![value],
133 });
134 }
135
136 if trimmed.contains(',') && looks_like_bare_value_list(trimmed) {
143 return parse_literal_list(trimmed);
144 }
145
146 Ok(Source::Generator {
154 expr: trimmed.to_string(),
155 cardinality_hint: None,
156 })
157}
158
159fn parse_distribution(text: &str) -> Option<Result<Source, SourceParseError>> {
170 let (call, support_text) = match split_on_keyword(text, " on ") {
171 Some((call, rest)) => (call, Some(rest)),
172 None => (text, None),
173 };
174 let (name, args) = split_call(call)?;
175 let measure = MeasureName::from_text(name)?;
176 Some(build_distribution(text, measure, args, support_text))
177}
178
179fn build_distribution(
182 text: &str,
183 measure: MeasureName,
184 args: &str,
185 support_text: Option<&str>,
186) -> Result<Source, SourceParseError> {
187 let invalid = || SourceParseError::InvalidRange(text.to_string());
188 let mut params = Vec::new();
189 for arg in args.split(',') {
190 let arg = arg.trim();
191 if arg.is_empty() {
192 if params.is_empty() && args.trim().is_empty() {
193 break;
194 }
195 return Err(invalid());
196 }
197 params.push(arg.parse::<f64>().map_err(|_| invalid())?);
198 }
199 let params = measure.resolve_params(¶ms).map_err(|_| invalid())?;
200 let support = match support_text {
201 None => measure.support(¶ms),
202 Some(interval_text) => match parse_range_text(interval_text)? {
203 Source::ContinuousInterval { interval, .. } => interval,
204 _ => return Err(invalid()),
205 },
206 };
207 Ok(Source::Distribution {
208 distribution: measure,
209 support,
210 params,
211 })
212}
213
214fn split_call(text: &str) -> Option<(&str, &str)> {
217 let text = text.trim();
218 let open = text.find('(')?;
219 let inner = text.strip_suffix(')')?;
220 let name = text[..open].trim();
221 if name.is_empty() || !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
222 return None;
223 }
224 Some((name, &inner[open + 1..]))
225}
226
227fn split_on_keyword<'a>(text: &'a str, keyword: &str) -> Option<(&'a str, &'a str)> {
230 let bytes = text.as_bytes();
231 let mut depth = 0i32;
232 let mut quote: Option<u8> = None;
233 for i in 0..bytes.len() {
234 let c = bytes[i];
235 match quote {
236 Some(q) => {
237 if c == q {
238 quote = None;
239 }
240 }
241 None => match c {
242 b'"' | b'\'' => quote = Some(c),
243 b'(' | b'[' | b'{' => depth += 1,
244 b')' | b']' | b'}' => depth -= 1,
245 _ => {
246 if depth == 0 && text[i..].starts_with(keyword) {
247 return Some((&text[..i], &text[i + keyword.len()..]));
248 }
249 }
250 },
251 }
252 }
253 None
254}
255
256fn looks_like_bare_value_list(text: &str) -> bool {
262 !text.chars().any(|c| {
263 matches!(
264 c,
265 '(' | ')'
266 | '['
267 | ']'
268 | '{'
269 | '}'
270 | '\''
271 | '"'
272 | '+'
273 | '*'
274 | '/'
275 | '%'
276 | '='
277 | '<'
278 | '>'
279 | '!'
280 | '&'
281 | '|'
282 | '~'
283 | '^'
284 | '?'
285 )
286 })
287}
288
289fn strip_dynamic_curly(s: &str) -> Option<String> {
293 let s = s.trim();
294 if !s.starts_with('{') || !s.ends_with('}') {
295 return None;
296 }
297 let inner = &s[1..s.len() - 1];
298 if !inner.contains('{') {
302 return None;
303 }
304 Some(inner.to_string())
305}
306
307fn parse_literal_list(inner: &str) -> Result<Source, SourceParseError> {
311 let parts: Vec<&str> = inner
312 .split(',')
313 .map(|s| s.trim())
314 .filter(|s| !s.is_empty())
315 .collect();
316
317 if parts.is_empty() {
318 return Ok(Source::Literal { values: Vec::new() });
319 }
320
321 let values: Vec<LiteralValue> = parts.iter().map(|s| parse_literal_value(s)).collect();
322
323 Ok(Source::Literal { values })
324}
325
326fn bracket_is_pure_literal(inner: &str) -> bool {
333 let elems: Vec<&str> = inner
334 .split(',')
335 .map(str::trim)
336 .filter(|s| !s.is_empty())
337 .collect();
338 if elems.is_empty() {
339 return true; }
341 elems.iter().all(|e| {
342 if e.ends_with('…') || e.ends_with("...") {
343 return false; }
345 e.eq_ignore_ascii_case("true")
346 || e.eq_ignore_ascii_case("false")
347 || ((e.starts_with('"') && e.ends_with('"'))
348 || (e.starts_with('\'') && e.ends_with('\'')))
349 || e.parse::<i64>().is_ok()
350 || e.parse::<f64>().is_ok()
351 })
352}
353
354fn parse_literal_value(s: &str) -> LiteralValue {
355 let s = s.trim();
356 if s.eq_ignore_ascii_case("true") {
357 return LiteralValue::Bool(true);
358 }
359 if s.eq_ignore_ascii_case("false") {
360 return LiteralValue::Bool(false);
361 }
362 if (s.starts_with('"') && s.ends_with('"')) || (s.starts_with('\'') && s.ends_with('\'')) {
364 let inner = &s[1..s.len() - 1];
365 return LiteralValue::String(inner.to_string());
366 }
367 if let Ok(n) = s.parse::<i64>() {
369 return LiteralValue::Int(n);
370 }
371 if let Ok(f) = s.parse::<f64>() {
373 return LiteralValue::Float(f);
374 }
375 LiteralValue::String(s.to_string())
377}
378
379fn parse_range_text(text: &str) -> Result<Source, SourceParseError> {
382 let text = text.trim();
383 match find_top_level(text, "..") {
384 Some(idx) => parse_range(text, idx),
385 None => Err(SourceParseError::InvalidRange(text.to_string())),
386 }
387}
388
389fn parse_range(text: &str, dotdot_idx: usize) -> Result<Source, SourceParseError> {
390 let lo_str = text[..dotdot_idx].trim();
391 let after = &text[dotdot_idx + 2..];
392
393 let (inclusive_end, after) = if let Some(rest) = after.strip_prefix('=') {
395 (true, rest)
396 } else {
397 (false, after)
398 };
399
400 let (rhs, step) = if let Some(step_pos) = after.find(" step ") {
405 let rhs = after[..step_pos].trim();
406 let step_str = after[step_pos + 6..].trim();
407 let step: i64 = step_str
408 .parse()
409 .map_err(|_| SourceParseError::InvalidRange(text.to_string()))?;
410 (rhs, step)
411 } else if let Some(step_pos) = after.find("..") {
412 let rhs = after[..step_pos].trim();
415 let step_str = after[step_pos + 2..].trim();
416 let step: i64 = step_str
417 .parse()
418 .map_err(|_| SourceParseError::InvalidRange(text.to_string()))?;
419 (rhs, step)
420 } else {
421 (after.trim(), 1)
422 };
423
424 if let (Ok(lo_i), Ok(hi_i)) = (lo_str.parse::<i64>(), rhs.parse::<i64>()) {
426 let hi = if inclusive_end { hi_i + 1 } else { hi_i };
427 return Ok(Source::IntRange { lo: lo_i, hi, step });
428 }
429 if let (Ok(lo_f), Ok(hi_f)) = (lo_str.parse::<f64>(), rhs.parse::<f64>()) {
431 let interval = Interval {
432 lo: lo_f,
433 hi: hi_f,
434 lo_open: false,
435 hi_open: !inclusive_end,
436 };
437 return Ok(Source::ContinuousInterval {
438 interval,
439 measure: ProductMeasure::Uniform,
440 });
441 }
442
443 Err(SourceParseError::InvalidRange(text.to_string()))
444}
445
446fn strip_curly(s: &str) -> Option<String> {
447 let s = s.trim();
448 if s.starts_with('{') && s.ends_with('}') {
449 let inner = &s[1..s.len() - 1];
450 let trimmed = inner.trim();
451 if !trimmed.is_empty() && trimmed.chars().all(|c| c.is_alphanumeric() || c == '_') {
452 return Some(trimmed.to_string());
453 }
454 }
455 None
456}
457
458fn try_parse_bare_scalar(s: &str) -> Option<LiteralValue> {
464 if s.eq_ignore_ascii_case("true") {
465 return Some(LiteralValue::Bool(true));
466 }
467 if s.eq_ignore_ascii_case("false") {
468 return Some(LiteralValue::Bool(false));
469 }
470 if (s.starts_with('"') && s.ends_with('"')) || (s.starts_with('\'') && s.ends_with('\'')) {
471 let inner = &s[1..s.len() - 1];
472 return Some(LiteralValue::String(inner.to_string()));
473 }
474 if let Ok(n) = s.parse::<i64>() {
475 return Some(LiteralValue::Int(n));
476 }
477 if let Ok(f) = s.parse::<f64>() {
478 return Some(LiteralValue::Float(f));
479 }
480 None
481}
482
483fn looks_like_function_call(s: &str) -> bool {
484 let Some(open) = s.find('(') else {
485 return false;
486 };
487 if !s.ends_with(')') {
488 return false;
489 }
490 let name = &s[..open];
491 !name.is_empty() && name.chars().all(|c| c.is_alphanumeric() || c == '_')
492}
493
494fn find_top_level(s: &str, needle: &str) -> Option<usize> {
497 let bytes = s.as_bytes();
498 let needle_bytes = needle.as_bytes();
499 let mut depth = 0i64;
500 let mut i = 0;
501 while i + needle_bytes.len() <= bytes.len() {
502 match bytes[i] {
503 b'(' | b'[' | b'{' => depth += 1,
504 b')' | b']' | b'}' => depth -= 1,
505 _ => {}
506 }
507 if depth == 0 && &bytes[i..i + needle_bytes.len()] == needle_bytes {
508 return Some(i);
509 }
510 i += 1;
511 }
512 None
513}
514
515#[derive(Debug, Clone, PartialEq)]
517pub enum SourceParseError {
518 Unrecognized(String),
520 InvalidRange(String),
523}
524
525impl std::fmt::Display for SourceParseError {
526 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
527 match self {
528 SourceParseError::Unrecognized(s) => {
529 write!(f, "unrecognized source expression: {s:?}")
530 }
531 SourceParseError::InvalidRange(s) => {
532 write!(f, "invalid range expression: {s:?}")
533 }
534 }
535 }
536}
537
538impl std::error::Error for SourceParseError {}
539
540#[cfg(test)]
541mod tests {
542 use super::*;
543
544 #[test]
545 fn int_range_exclusive() {
546 let s = parse_source("1..10").unwrap();
547 assert!(matches!(
548 s,
549 Source::IntRange {
550 lo: 1,
551 hi: 10,
552 step: 1
553 }
554 ));
555 }
556
557 #[test]
558 fn double_quoted_source_is_string_comprehension_striped() {
559 let s = parse_source(r#""rerank_def, rerank_1x, rerank_2x""#).unwrap();
561 match s {
562 Source::Literal { values } => {
563 assert_eq!(
564 values,
565 vec![
566 LiteralValue::String("rerank_def".into()),
567 LiteralValue::String("rerank_1x".into()),
568 LiteralValue::String("rerank_2x".into()),
569 ]
570 );
571 }
572 other => panic!("expected striped Literal, got {other:?}"),
573 }
574 }
575
576 #[test]
577 fn single_quoted_source_is_atomic() {
578 let s = parse_source("'rerank_def, rerank_1x'").unwrap();
580 match s {
581 Source::Literal { values } => {
582 assert_eq!(
583 values,
584 vec![LiteralValue::String("rerank_def, rerank_1x".into())]
585 );
586 }
587 other => panic!("expected atomic Literal, got {other:?}"),
588 }
589 }
590
591 #[test]
592 fn int_range_inclusive() {
593 let s = parse_source("1..=10").unwrap();
594 assert!(matches!(
595 s,
596 Source::IntRange {
597 lo: 1,
598 hi: 11,
599 step: 1
600 }
601 ));
602 }
603
604 #[test]
605 fn int_range_with_step() {
606 let s = parse_source("0..100 step 10").unwrap();
607 assert!(matches!(
608 s,
609 Source::IntRange {
610 lo: 0,
611 hi: 100,
612 step: 10
613 }
614 ));
615 }
616
617 #[test]
618 fn literal_int_list() {
619 let s = parse_source("[1, 2, 3]").unwrap();
620 match s {
621 Source::Literal { values } => {
622 assert_eq!(values.len(), 3);
623 assert_eq!(values[0], LiteralValue::Int(1));
624 assert_eq!(values[2], LiteralValue::Int(3));
625 }
626 other => panic!("expected Literal, got {other:?}"),
627 }
628 }
629
630 #[test]
631 fn bracket_bare_words_are_references_not_strings() {
632 let s = parse_source("[a, b, c]").unwrap();
638 match s {
639 Source::Generator { expr, .. } => assert_eq!(expr, "[a, b, c]"),
640 other => panic!("expected deferred Generator, got {other:?}"),
641 }
642 }
643
644 #[test]
645 fn bracket_with_spread_defers_to_generator() {
646 let s = parse_source("[xs…]").unwrap();
647 assert!(
648 matches!(s, Source::Generator { .. }),
649 "spread list must defer: {s:?}"
650 );
651 }
652
653 #[test]
654 fn literal_quoted_strings() {
655 let s = parse_source(r#"["hello", "world"]"#).unwrap();
656 match s {
657 Source::Literal { values } => {
658 assert_eq!(values[0], LiteralValue::String("hello".into()));
659 assert_eq!(values[1], LiteralValue::String("world".into()));
660 }
661 other => panic!("expected Literal, got {other:?}"),
662 }
663 }
664
665 #[test]
666 fn literal_float_list() {
667 let s = parse_source("[1.5, 2.5, 3.5]").unwrap();
668 match s {
669 Source::Literal { values } => {
670 assert_eq!(values[0], LiteralValue::Float(1.5));
671 }
672 other => panic!("expected Literal, got {other:?}"),
673 }
674 }
675
676 #[test]
677 fn workload_param_ref() {
678 let s = parse_source("{profiles}").unwrap();
679 match s {
680 Source::WorkloadParamList { name, .. } => assert_eq!(name, "profiles"),
681 other => panic!("expected WorkloadParamList, got {other:?}"),
682 }
683 }
684
685 #[test]
686 fn generator_function_call() {
687 let s = parse_source("fib(8)").unwrap();
688 match s {
689 Source::Generator { expr, .. } => assert_eq!(expr, "fib(8)"),
690 other => panic!("expected Generator, got {other:?}"),
691 }
692 }
693
694 #[test]
695 fn continuous_interval_via_floats() {
696 let s = parse_source("0.0..1.0").unwrap();
697 match s {
698 Source::ContinuousInterval { interval, measure } => {
699 assert_eq!(interval.lo, 0.0);
700 assert_eq!(interval.hi, 1.0);
701 assert!(matches!(measure, ProductMeasure::Uniform));
702 }
703 other => panic!("expected ContinuousInterval, got {other:?}"),
704 }
705 }
706
707 #[test]
708 fn continuous_interval_inclusive() {
709 let s = parse_source("0.0..=1.0").unwrap();
710 match s {
711 Source::ContinuousInterval { interval, .. } => {
712 assert!(!interval.hi_open);
713 }
714 other => panic!("expected ContinuousInterval, got {other:?}"),
715 }
716 }
717
718 #[test]
719 fn unrecognized_source_falls_back_to_generator() {
720 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}