1use std::collections::{BTreeMap, BTreeSet};
8
9use super::advisor::{
10 AdviceDecision, AdviceDecisionKind, AdviceSource, ConceptNamingAdvisor, MdlMergeAdvisor,
11 MergeAdvisor, MergeCandidate, MergeRequest, NamingAdvisor, NamingRequest,
12};
13pub use super::eval::MembershipOracle;
14use super::eval::{sample, GrammarOracle, SampleConfig};
15use super::prior::{build_structural_prior, ByteSpan, Delimiter, LeafKind, PriorOptions, SeedNode};
16use crate::grammar::{Grammar, GrammarExpr, GrammarFormat, GrammarRule};
17
18const ROOT_RULE: &str = "Root";
19const DEFAULT_MAX_ITERATIONS: usize = 64;
20const DEFAULT_SAMPLE_BUDGET: usize = 256;
21
22#[derive(Clone, Copy, Debug, PartialEq, Eq)]
24pub struct InferenceOptions {
25 pub incremental: bool,
28 pub max_iterations: usize,
30 pub sample_budget: usize,
32}
33
34impl Default for InferenceOptions {
35 fn default() -> Self {
36 Self {
37 incremental: false,
38 max_iterations: DEFAULT_MAX_ITERATIONS,
39 sample_budget: DEFAULT_SAMPLE_BUDGET,
40 }
41 }
42}
43
44#[derive(Clone, Debug, Default, PartialEq, Eq)]
46pub struct InferenceReport {
47 pub rules: usize,
49 pub bubbles_proposed: usize,
51 pub merges_accepted: usize,
53 pub merges_rejected: usize,
55 pub advice: Vec<AdviceDecision>,
57}
58
59impl InferenceReport {
60 fn record_advice(
61 &mut self,
62 kind: AdviceDecisionKind,
63 target: impl Into<String>,
64 source: AdviceSource,
65 ) {
66 self.advice.push(AdviceDecision::new(kind, target, source));
67 }
68}
69
70#[derive(Clone, Debug, PartialEq, Eq)]
72pub struct InferenceResult {
73 pub grammar: Grammar,
75 pub report: InferenceReport,
77}
78
79pub trait Oracle {
81 fn accepts_all_positive(&self, grammar: &Grammar, examples: &[String]) -> bool {
83 let grammar_oracle = GrammarOracle::new(grammar);
84 examples
85 .iter()
86 .all(|example| grammar_oracle.accepts(example))
87 }
88
89 fn membership(&self) -> Option<&dyn MembershipOracle> {
91 None
92 }
93}
94
95#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
97pub struct PositiveOnlyOracle;
98
99impl PositiveOnlyOracle {
100 #[must_use]
102 pub const fn new() -> Self {
103 Self
104 }
105}
106
107impl Oracle for PositiveOnlyOracle {}
108
109#[must_use]
111pub fn infer_cfg(
112 examples: &[String],
113 oracle: &dyn Oracle,
114 opts: InferenceOptions,
115) -> InferenceResult {
116 infer_cfg_with_advisors(
117 examples,
118 oracle,
119 opts,
120 &ConceptNamingAdvisor,
121 &MdlMergeAdvisor,
122 )
123}
124
125#[must_use]
127pub fn infer_cfg_with_advisors(
128 examples: &[String],
129 oracle: &dyn Oracle,
130 opts: InferenceOptions,
131 naming_advisor: &dyn NamingAdvisor,
132 merge_advisor: &dyn MergeAdvisor,
133) -> InferenceResult {
134 let positives = sorted_unique_examples(examples);
135 let mut report = InferenceReport::default();
136
137 if positives.is_empty() {
138 let grammar = Grammar::new().with_source_format(GrammarFormat::Inferred);
139 return InferenceResult { grammar, report };
140 }
141
142 let mut candidate =
143 structured_grammar(&positives, opts, naming_advisor, merge_advisor, &mut report);
144 if !oracle.accepts_all_positive(&candidate, examples)
145 || membership_rejects_candidate(&candidate, oracle, opts)
146 {
147 report.merges_rejected = report.merges_rejected.saturating_add(1);
148 candidate = exact_positive_grammar(&positives);
149 }
150
151 report.rules = candidate.rules().len();
152 InferenceResult {
153 grammar: candidate,
154 report,
155 }
156}
157
158fn sorted_unique_examples(examples: &[String]) -> Vec<String> {
159 examples
160 .iter()
161 .cloned()
162 .collect::<BTreeSet<_>>()
163 .into_iter()
164 .collect()
165}
166
167fn structured_grammar(
168 examples: &[String],
169 opts: InferenceOptions,
170 naming_advisor: &dyn NamingAdvisor,
171 merge_advisor: &dyn MergeAdvisor,
172 report: &mut InferenceReport,
173) -> Grammar {
174 let prior = build_structural_prior(examples, PriorOptions::default());
175 let mut draft = Draft::default();
176 let mut root_alternatives = Vec::new();
177
178 for tree in &prior.trees {
179 root_alternatives.push(draft.expr_for_node(&tree.root, &tree.example));
180 }
181
182 report.bubbles_proposed = report
183 .bubbles_proposed
184 .saturating_add(draft.bubbles_proposed);
185
186 let before_root = root_alternatives.len();
187 let root_merge_exprs = root_alternatives.clone();
188 let root_expr = choice_expr(root_alternatives);
189 let root_merges = before_root.saturating_sub(choice_len(&root_expr));
190 report.merges_accepted = report.merges_accepted.saturating_add(root_merges);
191
192 let mut grammar = Grammar::new().with_source_format(GrammarFormat::Inferred);
193 record_naming_advice(
194 report,
195 naming_advisor,
196 &grammar,
197 ROOT_RULE,
198 &root_expr,
199 examples,
200 );
201 grammar.add_rule(GrammarRule::new(ROOT_RULE, root_expr));
202 if root_merges > 0 {
203 record_merge_advice(
204 report,
205 merge_advisor,
206 examples,
207 ROOT_RULE,
208 &root_merge_exprs,
209 );
210 }
211
212 for delimiter in [Delimiter::Paren, Delimiter::Curly, Delimiter::Square] {
213 let Some(alternatives) = draft.group_alternatives.remove(&delimiter) else {
214 continue;
215 };
216 let before = alternatives.len();
217 let merge_exprs = alternatives
218 .iter()
219 .cloned()
220 .map(seq_expr)
221 .collect::<Vec<_>>();
222 let rules = rules_for_group(delimiter, alternatives);
223 let after = rules.first().map_or(0, |rule| choice_len(rule.expr()));
224 let group_merges = before.saturating_sub(after);
225 report.merges_accepted = report.merges_accepted.saturating_add(group_merges);
226 if group_merges > 0 {
227 record_merge_advice(
228 report,
229 merge_advisor,
230 examples,
231 group_rule_name(delimiter),
232 &merge_exprs,
233 );
234 }
235 for rule in rules {
236 record_naming_advice(
237 report,
238 naming_advisor,
239 &grammar,
240 rule.name(),
241 rule.expr(),
242 &[],
243 );
244 grammar.add_rule(rule);
245 }
246 }
247
248 if opts.incremental {
249 report.bubbles_proposed = report.bubbles_proposed.saturating_add(examples.len());
250 }
251
252 grammar.set_start(ROOT_RULE);
253 grammar
254}
255
256fn record_naming_advice(
257 report: &mut InferenceReport,
258 advisor: &dyn NamingAdvisor,
259 grammar: &Grammar,
260 target: impl Into<String>,
261 rule_expr: &GrammarExpr,
262 sample_yields: &[String],
263) {
264 let source = advisor
265 .propose_names(&NamingRequest {
266 grammar,
267 rule_expr,
268 sample_yields,
269 })
270 .first()
271 .map_or(AdviceSource::Deterministic, |candidate| candidate.source);
272
273 report.record_advice(AdviceDecisionKind::Naming, target, source);
274}
275
276fn record_merge_advice(
277 report: &mut InferenceReport,
278 advisor: &dyn MergeAdvisor,
279 examples: &[String],
280 target: impl Into<String>,
281 alternatives: &[GrammarExpr],
282) {
283 let target = target.into();
284 let source = merge_advice_request_grammar(&target, alternatives).map_or(
285 AdviceSource::Deterministic,
286 |(advice_grammar, candidates)| {
287 advisor
288 .rank_merges(&MergeRequest {
289 grammar: &advice_grammar,
290 candidates: &candidates,
291 examples,
292 })
293 .first()
294 .map_or(AdviceSource::Deterministic, |score| score.source)
295 },
296 );
297
298 report.record_advice(AdviceDecisionKind::Merge, target, source);
299}
300
301fn merge_advice_request_grammar(
302 target: &str,
303 alternatives: &[GrammarExpr],
304) -> Option<(Grammar, Vec<MergeCandidate>)> {
305 if alternatives.len() < 2 {
306 return None;
307 }
308
309 let mut grammar = Grammar::new().with_source_format(GrammarFormat::Inferred);
310 let mut names = Vec::with_capacity(alternatives.len());
311 for (index, alternative) in alternatives.iter().enumerate() {
312 let name = format!("{target}_alternative_{}", index + 1);
313 grammar.add_rule(GrammarRule::new(name.clone(), alternative.clone()));
314 names.push(name);
315 }
316
317 grammar.add_rule(GrammarRule::new(
318 target,
319 GrammarExpr::choice(false, names.iter().cloned().map(GrammarExpr::non_terminal)),
320 ));
321 grammar.set_start(target);
322
323 let winner = names.first()?.clone();
324 let candidates = names
325 .iter()
326 .skip(1)
327 .map(|loser| MergeCandidate::new(&winner, loser))
328 .collect::<Vec<_>>();
329 Some((grammar, candidates))
330}
331
332#[derive(Default)]
333struct Draft {
334 group_alternatives: BTreeMap<Delimiter, Vec<Vec<GrammarExpr>>>,
335 bubbles_proposed: usize,
336}
337
338impl Draft {
339 fn expr_for_node(&mut self, node: &SeedNode, example: &str) -> GrammarExpr {
340 match node {
341 SeedNode::Leaf { span, kind } => terminal_for_leaf(example, *span, *kind),
342 SeedNode::Group {
343 delimiter,
344 children,
345 span,
346 } if *delimiter == Delimiter::Root => {
347 seq_expr(self.sequence_for_children(*delimiter, children, *span, example))
348 }
349 SeedNode::Group {
350 delimiter,
351 children,
352 span,
353 } => {
354 let inner = self.sequence_for_children(*delimiter, children, *span, example);
355 self.group_alternatives
356 .entry(*delimiter)
357 .or_default()
358 .push(inner);
359 self.bubbles_proposed = self.bubbles_proposed.saturating_add(1);
360 GrammarExpr::non_terminal(group_rule_name(*delimiter))
361 }
362 }
363 }
364
365 fn sequence_for_children(
366 &mut self,
367 delimiter: Delimiter,
368 children: &[SeedNode],
369 span: ByteSpan,
370 example: &str,
371 ) -> Vec<GrammarExpr> {
372 let (content_start, content_end) = content_span(delimiter, span);
373 let mut cursor = content_start;
374 let mut sequence = Vec::new();
375
376 for child in children {
377 let child_span = node_span(child);
378 push_gap(example, cursor, child_span.start, &mut sequence);
379 sequence.push(self.expr_for_node(child, example));
380 cursor = child_span.end;
381 }
382
383 push_gap(example, cursor, content_end, &mut sequence);
384 sequence
385 }
386}
387
388const fn content_span(delimiter: Delimiter, span: ByteSpan) -> (usize, usize) {
389 match delimiter {
390 Delimiter::Root => (span.start, span.end),
391 Delimiter::Paren | Delimiter::Curly | Delimiter::Square => {
392 (span.start.saturating_add(1), span.end.saturating_sub(1))
393 }
394 }
395}
396
397const fn node_span(node: &SeedNode) -> ByteSpan {
398 match node {
399 SeedNode::Leaf { span, .. } | SeedNode::Group { span, .. } => *span,
400 }
401}
402
403fn push_gap(example: &str, start: usize, end: usize, sequence: &mut Vec<GrammarExpr>) {
404 if start < end {
405 sequence.push(GrammarExpr::terminal(&example[start..end]));
406 }
407}
408
409fn terminal_for_leaf(example: &str, span: ByteSpan, kind: LeafKind) -> GrammarExpr {
410 let value = &example[span.start..span.end];
411 match kind {
412 LeafKind::Backtick => GrammarExpr::terminal_insensitive(value),
413 LeafKind::Text | LeafKind::SingleQuote | LeafKind::DoubleQuote => {
414 GrammarExpr::terminal(value)
415 }
416 }
417}
418
419fn rules_for_group(delimiter: Delimiter, alternatives: Vec<Vec<GrammarExpr>>) -> Vec<GrammarRule> {
420 let name = group_rule_name(delimiter);
421 if let Some(list) = infer_comma_list_group(delimiter, &name, &alternatives) {
422 return list;
423 }
424
425 let inner = choice_expr(alternatives.into_iter().map(seq_expr).collect());
426 vec![GrammarRule::new(name, wrap_delimited(delimiter, inner))]
427}
428
429fn infer_comma_list_group(
430 delimiter: Delimiter,
431 name: &str,
432 alternatives: &[Vec<GrammarExpr>],
433) -> Option<Vec<GrammarRule>> {
434 let mut saw_empty = false;
435 let mut saw_separator = false;
436 let mut item_alternatives = Vec::new();
437 let mut separator_alternatives = Vec::new();
438
439 for alternative in alternatives {
440 if alternative.is_empty() {
441 saw_empty = true;
442 continue;
443 }
444
445 let parts = split_comma_list(alternative)?;
446 saw_separator |= !parts.separators.is_empty();
447 item_alternatives.extend(parts.items);
448 separator_alternatives.extend(parts.separators);
449 }
450
451 if !saw_separator || item_alternatives.is_empty() {
452 return None;
453 }
454
455 let item_name = format!("{name}_item");
456 let items_name = format!("{name}_items");
457 let separator = choice_expr(separator_alternatives);
458 let content = if saw_empty {
459 GrammarExpr::optional(GrammarExpr::non_terminal(&items_name))
460 } else {
461 GrammarExpr::non_terminal(&items_name)
462 };
463
464 let group_rule = GrammarRule::new(name, wrap_delimited(delimiter, content));
465 let item_rule = GrammarRule::new(item_name.clone(), choice_expr(item_alternatives));
466 let items_rule = GrammarRule::new(
467 items_name.clone(),
468 GrammarExpr::choice(
469 false,
470 [
471 GrammarExpr::non_terminal(&item_name),
472 GrammarExpr::sequence([
473 GrammarExpr::non_terminal(&item_name),
474 separator,
475 GrammarExpr::non_terminal(&items_name),
476 ]),
477 ],
478 ),
479 );
480
481 Some(vec![group_rule, item_rule, items_rule])
482}
483
484#[derive(Clone, Debug, PartialEq, Eq)]
485struct ListParts {
486 items: Vec<GrammarExpr>,
487 separators: Vec<GrammarExpr>,
488}
489
490fn split_comma_list(sequence: &[GrammarExpr]) -> Option<ListParts> {
491 let mut cursor = 0usize;
492 let mut items = Vec::new();
493 let mut separators = Vec::new();
494
495 while let Some(comma) = find_comma(sequence, cursor) {
496 let item_start = trim_start(sequence, cursor, comma);
497 let item_end = trim_end(sequence, item_start, comma);
498 if item_start == item_end {
499 return None;
500 }
501 items.push(seq_expr(sequence[item_start..item_end].to_vec()));
502
503 let separator_start = item_end;
504 let separator_end = trim_start(sequence, comma.saturating_add(1), sequence.len());
505 separators.push(seq_expr(sequence[separator_start..separator_end].to_vec()));
506 cursor = separator_end;
507 }
508
509 let item_start = trim_start(sequence, cursor, sequence.len());
510 let item_end = trim_end(sequence, item_start, sequence.len());
511 if item_start == item_end {
512 return None;
513 }
514 items.push(seq_expr(sequence[item_start..item_end].to_vec()));
515
516 Some(ListParts { items, separators })
517}
518
519fn find_comma(sequence: &[GrammarExpr], start: usize) -> Option<usize> {
520 sequence
521 .iter()
522 .enumerate()
523 .skip(start)
524 .find_map(|(index, expr)| is_comma(expr).then_some(index))
525}
526
527fn trim_start(sequence: &[GrammarExpr], mut start: usize, end: usize) -> usize {
528 while start < end && is_whitespace(&sequence[start]) {
529 start += 1;
530 }
531 start
532}
533
534fn trim_end(sequence: &[GrammarExpr], start: usize, mut end: usize) -> usize {
535 while start < end && is_whitespace(&sequence[end - 1]) {
536 end -= 1;
537 }
538 end
539}
540
541fn is_comma(expr: &GrammarExpr) -> bool {
542 matches!(expr, GrammarExpr::Terminal(value) if value == ",")
543}
544
545fn is_whitespace(expr: &GrammarExpr) -> bool {
546 matches!(expr, GrammarExpr::Terminal(value) if !value.is_empty() && value.chars().all(char::is_whitespace))
547}
548
549fn wrap_delimited(delimiter: Delimiter, inner: GrammarExpr) -> GrammarExpr {
550 let (open, close) = delimiter_tokens(delimiter);
551 let mut items = vec![GrammarExpr::terminal(open)];
552 if inner != GrammarExpr::Empty {
553 items.push(inner);
554 }
555 items.push(GrammarExpr::terminal(close));
556 GrammarExpr::sequence(items)
557}
558
559const fn delimiter_tokens(delimiter: Delimiter) -> (&'static str, &'static str) {
560 match delimiter {
561 Delimiter::Paren => ("(", ")"),
562 Delimiter::Curly => ("{", "}"),
563 Delimiter::Square => ("[", "]"),
564 Delimiter::Root => ("", ""),
565 }
566}
567
568fn group_rule_name(delimiter: Delimiter) -> String {
569 match delimiter {
570 Delimiter::Paren => "n0",
571 Delimiter::Curly => "n1",
572 Delimiter::Square => "n2",
573 Delimiter::Root => ROOT_RULE,
574 }
575 .to_string()
576}
577
578fn choice_expr(alternatives: Vec<GrammarExpr>) -> GrammarExpr {
579 let mut unique = BTreeMap::<String, GrammarExpr>::new();
580 for alternative in alternatives {
581 unique.entry(expr_key(&alternative)).or_insert(alternative);
582 }
583
584 match unique.len() {
585 0 => GrammarExpr::Empty,
586 1 => unique
587 .into_values()
588 .next()
589 .expect("one alternative must be present"),
590 _ => GrammarExpr::choice(false, unique.into_values()),
591 }
592}
593
594fn seq_expr(items: Vec<GrammarExpr>) -> GrammarExpr {
595 match items.len() {
596 0 => GrammarExpr::Empty,
597 1 => items
598 .into_iter()
599 .next()
600 .expect("one sequence item must be present"),
601 _ => GrammarExpr::sequence(items),
602 }
603}
604
605fn choice_len(expr: &GrammarExpr) -> usize {
606 match expr {
607 GrammarExpr::Choice { alternatives, .. } => alternatives.len(),
608 GrammarExpr::Empty => 0,
609 _ => 1,
610 }
611}
612
613fn expr_key(expr: &GrammarExpr) -> String {
614 format!("{expr:?}")
615}
616
617fn exact_positive_grammar(examples: &[String]) -> Grammar {
618 let alternatives = examples.iter().map(|example| {
619 if example.is_empty() {
620 GrammarExpr::Empty
621 } else {
622 GrammarExpr::terminal(example)
623 }
624 });
625
626 Grammar::new()
627 .with_source_format(GrammarFormat::Inferred)
628 .with_rule(GrammarRule::new(
629 ROOT_RULE,
630 choice_expr(alternatives.collect()),
631 ))
632 .with_start(ROOT_RULE)
633}
634
635fn membership_rejects_candidate(
636 candidate: &Grammar,
637 oracle: &dyn Oracle,
638 opts: InferenceOptions,
639) -> bool {
640 let Some(membership) = oracle.membership() else {
641 return false;
642 };
643
644 let config = SampleConfig {
645 count: opts.sample_budget.max(1),
646 max_depth: opts.max_iterations.max(1),
647 ..SampleConfig::default()
648 };
649
650 sample(candidate, &config)
651 .is_ok_and(|samples| samples.iter().any(|sample| !membership.accepts(sample)))
652}