1use std::collections::BTreeMap;
12
13use harn_hostlib::ast::Language;
14use serde::Serialize;
15
16use crate::constraint::CompiledConstraint;
17use crate::error::RulesError;
18use crate::evaluator::CompiledRuleTree;
19use crate::fix::{interpolate, splice, AppliedEdit};
20use crate::model::{Applicability, Rule, Safety, Severity};
21use crate::semantic::enrich_harn_matches;
22use crate::transform::CompiledTransform;
23
24#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
27pub struct Span {
28 pub start_byte: usize,
30 pub end_byte: usize,
32 pub start_row: usize,
34 pub start_col: usize,
36 pub end_row: usize,
38 pub end_col: usize,
40}
41
42impl Span {
43 pub(crate) fn of(node: tree_sitter::Node<'_>) -> Self {
44 let start = node.start_position();
45 let end = node.end_position();
46 Span {
47 start_byte: node.start_byte(),
48 end_byte: node.end_byte(),
49 start_row: start.row,
50 start_col: start.column,
51 end_row: end.row,
52 end_col: end.column,
53 }
54 }
55}
56
57#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize)]
61pub struct BindingMetadata {
62 #[serde(skip_serializing_if = "Option::is_none")]
64 pub resolved: Option<ResolvedBinding>,
65 #[serde(rename = "type", skip_serializing_if = "Option::is_none")]
67 pub ty: Option<String>,
68}
69
70impl BindingMetadata {
71 pub fn is_empty(&self) -> bool {
73 self.resolved.is_none() && self.ty.is_none()
74 }
75}
76
77#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
79pub struct ResolvedBinding {
80 pub id: String,
82 pub name: String,
84 pub kind: String,
87 #[serde(flatten)]
89 pub span: Span,
90}
91
92#[derive(Debug, Clone)]
94pub struct Binding {
95 pub text: String,
97 pub span: Span,
99 pub metadata: BindingMetadata,
101}
102
103impl Binding {
104 pub(crate) fn new(text: String, span: Span) -> Self {
105 Binding {
106 text,
107 span,
108 metadata: BindingMetadata::default(),
109 }
110 }
111}
112
113#[derive(Debug, Clone)]
115pub struct RuleMatch {
116 pub rule_id: String,
118 pub span: Span,
120 pub text: String,
122 pub bindings: BTreeMap<String, Binding>,
125}
126
127#[derive(Debug, Clone)]
129pub struct CodemodResult {
130 pub rewritten: String,
132 pub edits: Vec<AppliedEdit>,
134 pub changed: bool,
136 pub safety: Safety,
138 pub applicability: Applicability,
140 pub idempotent: bool,
143}
144
145pub struct CompiledRule {
147 rule_id: String,
148 language: Language,
149 execution: Execution,
150 constraints: Vec<CompiledConstraint>,
152 transforms: Vec<(String, CompiledTransform)>,
154 fix: Option<String>,
156 fix_target: Option<String>,
159 safety: Safety,
161 message: String,
163 severity: Severity,
165}
166
167#[derive(Debug, Clone)]
170pub struct Diagnostic {
171 pub rule_id: String,
173 pub message: String,
175 pub severity: Severity,
177 pub span: Span,
179 pub applicability: Applicability,
181 pub fix: Option<String>,
184}
185
186enum Execution {
187 SourceRegex(regex::Regex),
190 Tree(Box<CompiledRuleTree>),
192}
193
194impl CompiledRule {
195 pub fn compile(rule: &Rule) -> Result<Self, RulesError> {
197 let language =
198 Language::from_name(&rule.language).ok_or_else(|| RulesError::UnknownLanguage {
199 rule: rule.id.clone(),
200 language: rule.language.clone(),
201 })?;
202
203 let execution = if rule.rule.is_pure_regex() {
207 let pattern = rule.rule.regex.as_ref().expect("pure regex");
208 Execution::SourceRegex(regex::Regex::new(pattern).map_err(|err| {
209 RulesError::PatternCompile {
210 rule: rule.id.clone(),
211 message: format!("invalid regex `{pattern}`: {err}"),
212 }
213 })?)
214 } else {
215 Execution::Tree(Box::new(CompiledRuleTree::compile(
216 &rule.id,
217 language,
218 &rule.rule,
219 &rule.utils,
220 )?))
221 };
222
223 let constraints = rule
224 .where_constraints
225 .iter()
226 .map(|c| CompiledConstraint::compile(&rule.id, language, c))
227 .collect::<Result<Vec<_>, _>>()?;
228
229 let transforms = rule
230 .transform
231 .iter()
232 .map(|(name, t)| {
233 CompiledTransform::compile(&rule.id, name, t).map(|c| (name.clone(), c))
234 })
235 .collect::<Result<Vec<_>, _>>()?;
236
237 Ok(CompiledRule {
238 rule_id: rule.id.clone(),
239 language,
240 execution,
241 constraints,
242 transforms,
243 fix: rule.fix.clone(),
244 fix_target: rule.fix_target.clone(),
245 safety: rule.safety,
246 message: rule.message.clone(),
247 severity: rule.severity,
248 })
249 }
250
251 pub fn language(&self) -> Language {
253 self.language
254 }
255
256 pub fn safety(&self) -> Safety {
258 self.safety
259 }
260
261 pub fn applicability(&self) -> Applicability {
264 self.safety.applicability()
265 }
266
267 pub fn id(&self) -> &str {
269 &self.rule_id
270 }
271
272 pub fn severity(&self) -> Severity {
275 self.severity
276 }
277
278 pub fn message(&self) -> &str {
280 &self.message
281 }
282
283 pub fn run(&self, source: &str) -> Result<Vec<RuleMatch>, RulesError> {
286 let mut matches = match &self.execution {
287 Execution::SourceRegex(regex) => self.run_regex(regex, source),
288 Execution::Tree(tree) => tree
289 .find(&self.rule_id, self.language, source)?
290 .into_iter()
291 .map(|m| RuleMatch {
292 rule_id: self.rule_id.clone(),
293 span: m.span,
294 text: m.text,
295 bindings: m.bindings,
296 })
297 .collect(),
298 };
299 if self.language == Language::Harn && !matches.is_empty() {
300 enrich_harn_matches(source, &mut matches).map_err(|message| {
301 RulesError::SourceParse {
302 rule: self.rule_id.clone(),
303 message,
304 }
305 })?;
306 }
307 if !self.constraints.is_empty() {
308 matches.retain(|m| self.satisfies_constraints(m));
309 }
310 Ok(matches)
311 }
312
313 fn satisfies_constraints(&self, m: &RuleMatch) -> bool {
316 self.constraints
317 .iter()
318 .all(|c| m.bindings.get(&c.metavar).is_some_and(|b| c.evaluate(b)))
319 }
320
321 pub fn apply(&self, source: &str) -> Result<CodemodResult, RulesError> {
330 let (rewritten, edits) = self.rewrite(source)?;
331 let changed = rewritten != source;
332 let (twice, _) = self.rewrite(&rewritten)?;
335 let idempotent = twice == rewritten;
336 Ok(CodemodResult {
337 rewritten,
338 edits,
339 changed,
340 safety: self.safety,
341 applicability: self.applicability(),
342 idempotent,
343 })
344 }
345
346 pub fn auto_apply(&self, source: &str) -> Result<CodemodResult, RulesError> {
350 if !self.safety.is_auto_applicable() {
351 return Err(RulesError::NotAutoApplicable {
352 rule: self.rule_id.clone(),
353 safety: format!("{:?}", self.safety),
354 });
355 }
356 self.apply(source)
357 }
358
359 pub fn apply_checked(&self, source: &str) -> Result<CodemodResult, RulesError> {
363 let result = self.apply(source)?;
364 if !result.idempotent {
365 return Err(RulesError::NotIdempotent {
366 rule: self.rule_id.clone(),
367 });
368 }
369 Ok(result)
370 }
371
372 pub fn diagnostics(&self, source: &str) -> Result<Vec<Diagnostic>, RulesError> {
377 let applicability = self.applicability();
378 let matches = self.run(source)?;
379 matches
380 .iter()
381 .map(|m| {
382 let span = if self.fix.is_some() && self.fix_target.is_some() {
386 self.edit_target(m)?.0
387 } else {
388 m.span
389 };
390 Ok(Diagnostic {
391 rule_id: self.rule_id.clone(),
392 message: self.message.clone(),
393 severity: self.severity,
394 span,
395 applicability,
396 fix: self.fix.as_ref().map(|template| {
397 let vars = self.metavars_for(m);
398 interpolate(template, &vars)
399 }),
400 })
401 })
402 .collect()
403 }
404
405 fn rewrite(&self, source: &str) -> Result<(String, Vec<AppliedEdit>), RulesError> {
408 let template = self
409 .fix
410 .as_ref()
411 .ok_or_else(|| RulesError::PatternCompile {
412 rule: self.rule_id.clone(),
413 message: "apply requires a `fix` template; this rule has none".into(),
414 })?;
415
416 let matches = dedupe_overlapping(self.run(source)?);
417 let edits: Vec<AppliedEdit> = matches
418 .iter()
419 .map(|m| {
420 let vars = self.metavars_for(m);
421 let (span, before) = self.edit_target(m)?;
422 Ok(AppliedEdit {
423 span,
424 before,
425 replacement: interpolate(template, &vars),
426 })
427 })
428 .collect::<Result<_, RulesError>>()?;
429 Ok((splice(source, &edits), edits))
430 }
431
432 fn edit_target(&self, m: &RuleMatch) -> Result<(Span, String), RulesError> {
435 match &self.fix_target {
436 None => Ok((m.span, m.text.clone())),
437 Some(name) => {
438 let binding =
439 m.bindings
440 .get(name)
441 .ok_or_else(|| RulesError::PatternCompile {
442 rule: self.rule_id.clone(),
443 message: format!(
444 "fixTarget `{name}` is not bound by this match; name a capture the matcher always binds"
445 ),
446 })?;
447 Ok((binding.span, binding.text.clone()))
448 }
449 }
450 }
451
452 fn metavars_for(&self, m: &RuleMatch) -> BTreeMap<String, String> {
455 let mut vars: BTreeMap<String, String> = m
456 .bindings
457 .iter()
458 .map(|(name, binding)| (name.clone(), binding.text.clone()))
459 .collect();
460 for (name, transform) in &self.transforms {
461 let input = m
462 .bindings
463 .get(&transform.source)
464 .map(|b| b.text.as_str())
465 .unwrap_or("");
466 vars.insert(name.clone(), transform.apply(input));
467 }
468 vars
469 }
470
471 fn run_regex(&self, regex: ®ex::Regex, source: &str) -> Vec<RuleMatch> {
472 let mut matches = Vec::new();
473 let mut cursor = RowColCursor::new(source);
478 for m in regex.find_iter(source) {
479 let (start_row, start_col) = cursor.advance_to(m.start());
480 let (end_row, end_col) = cursor.advance_to(m.end());
481 matches.push(RuleMatch {
482 rule_id: self.rule_id.clone(),
483 span: Span {
484 start_byte: m.start(),
485 end_byte: m.end(),
486 start_row,
487 start_col,
488 end_row,
489 end_col,
490 },
491 text: m.as_str().to_string(),
492 bindings: BTreeMap::new(),
493 });
494 }
495 matches
496 }
497}
498
499fn dedupe_overlapping(mut matches: Vec<RuleMatch>) -> Vec<RuleMatch> {
507 matches.sort_by(|a, b| {
510 a.span
511 .start_byte
512 .cmp(&b.span.start_byte)
513 .then(b.span.end_byte.cmp(&a.span.end_byte))
514 });
515 let mut kept: Vec<RuleMatch> = Vec::with_capacity(matches.len());
516 let mut covered_to = 0usize; for m in matches {
518 if m.span.start_byte >= covered_to {
522 covered_to = m.span.end_byte.max(covered_to);
523 kept.push(m);
524 }
525 }
526 kept
527}
528
529struct RowColCursor<'a> {
534 source: &'a str,
535 byte: usize,
536 row: usize,
537 col: usize,
538}
539
540impl<'a> RowColCursor<'a> {
541 fn new(source: &'a str) -> Self {
542 Self {
543 source,
544 byte: 0,
545 row: 0,
546 col: 0,
547 }
548 }
549
550 #[expect(
554 clippy::string_slice,
555 reason = "byte and target are regex match offsets on source, so both are char boundaries"
556 )]
557 fn advance_to(&mut self, target: usize) -> (usize, usize) {
558 for ch in self.source[self.byte..target].chars() {
559 if ch == '\n' {
560 self.row += 1;
561 self.col = 0;
562 } else {
563 self.col += 1;
564 }
565 }
566 self.byte = target;
567 (self.row, self.col)
568 }
569}
570
571#[cfg(test)]
572mod tests {
573 use super::*;
574 use crate::model::Rule;
575
576 fn rule(toml: &str) -> CompiledRule {
577 let parsed = Rule::from_toml_str(toml).expect("rule parses");
578 CompiledRule::compile(&parsed).expect("rule compiles")
579 }
580
581 #[test]
582 fn pattern_rule_binds_metavars() {
583 let compiled = rule(
584 r#"
585 id = "destructure-default"
586 language = "typescript"
587 fix = "{ $KEY: $SRC }"
588 [rule]
589 pattern = "$SRC?.$KEY ?? $DEFAULT"
590 "#,
591 );
592 let matches = compiled
593 .run("const a = cfg?.timeout ?? 30;\nconst b = opts?.retries ?? 3;\n")
594 .unwrap();
595 assert_eq!(matches.len(), 2);
596 assert_eq!(matches[0].bindings["SRC"].text, "cfg");
597 assert_eq!(matches[0].bindings["KEY"].text, "timeout");
598 assert_eq!(matches[0].bindings["DEFAULT"].text, "30");
599 assert_eq!(matches[1].bindings["SRC"].text, "opts");
600 assert_eq!(matches[0].text, "cfg?.timeout ?? 30");
602 assert_eq!(matches[0].span.start_row, 0);
603 assert_eq!(matches[1].span.start_row, 1);
604 }
605
606 #[test]
607 fn nested_matches_do_not_corrupt_or_panic_on_apply() {
608 let compiled = rule(
613 r#"
614 id = "sum-binop"
615 language = "typescript"
616 fix = "sum($X, $Y)"
617 [rule]
618 pattern = "$X + $Y"
619 "#,
620 );
621 assert!(compiled.run("const z = a + b + c;\n").unwrap().len() >= 2);
623 let result = compiled.apply("const z = a + b + c;\n").unwrap();
624 assert_eq!(result.rewritten, "const z = sum(a + b, c);\n");
626 assert_eq!(result.edits.len(), 1);
627 assert!(result.changed);
628 }
629
630 #[test]
631 fn dedupe_overlapping_keeps_outermost_in_document_order() {
632 let span = |s: usize, e: usize| Span {
633 start_byte: s,
634 end_byte: e,
635 start_row: 0,
636 start_col: s,
637 end_row: 0,
638 end_col: e,
639 };
640 let m = |s: usize, e: usize| RuleMatch {
641 rule_id: "r".into(),
642 span: span(s, e),
643 text: String::new(),
644 bindings: BTreeMap::new(),
645 };
646 let kept = dedupe_overlapping(vec![m(0, 5), m(0, 9), m(10, 14)]);
648 let spans: Vec<_> = kept
649 .iter()
650 .map(|m| (m.span.start_byte, m.span.end_byte))
651 .collect();
652 assert_eq!(spans, vec![(0, 9), (10, 14)]);
653 }
654
655 #[test]
656 fn kind_rule_matches_node_kind() {
657 let compiled = rule(
658 r#"
659 id = "find-calls"
660 language = "python"
661 [rule]
662 kind = "call"
663 "#,
664 );
665 let matches = compiled.run("print(x)\nlog(y)\n").unwrap();
666 assert_eq!(matches.len(), 2);
667 assert_eq!(matches[0].text, "print(x)");
668 assert!(matches[0].bindings.is_empty());
669 }
670
671 #[test]
672 fn regex_rule_matches_text() {
673 let compiled = rule(
674 r#"
675 id = "todo"
676 language = "rust"
677 message = "Found a TODO"
678 [rule]
679 regex = "TODO\\(\\w+\\)"
680 "#,
681 );
682 let matches = compiled
683 .run("fn f() {\n // TODO(ken) fix\n // todo lower\n}\n")
684 .unwrap();
685 assert_eq!(matches.len(), 1);
686 assert_eq!(matches[0].text, "TODO(ken)");
687 assert_eq!(matches[0].span.start_row, 1);
688 }
689
690 #[test]
691 fn raw_query_with_fix_target_rewrites_only_the_capture() {
692 let compiled = rule(
698 r#"
699 id = "let-to-const"
700 language = "harn"
701 fix = "const"
702 fixTarget = "kw"
703 [rule]
704 query = '(let_binding "let" @kw) @__match'
705 "#,
706 );
707 let result = compiled
708 .apply("fn f() {\n let x: Int = 1\n let y = 2\n}\n")
709 .unwrap();
710 assert_eq!(
711 result.rewritten,
712 "fn f() {\n const x: Int = 1\n const y = 2\n}\n"
713 );
714 assert_eq!(result.edits.len(), 2);
715 assert!(result.edits.iter().all(|e| e.before == "let"));
717 assert!(result.changed);
718 assert!(compiled.apply(&result.rewritten).unwrap().rewritten == result.rewritten);
720 }
721
722 #[test]
723 fn raw_query_scales_to_many_bindings() {
724 let compiled = rule(
733 r#"
734 id = "let-to-const"
735 language = "harn"
736 fix = "const"
737 fixTarget = "kw"
738 [rule]
739 query = '(let_binding "let" @kw) @__match'
740 "#,
741 );
742 let n = 200;
743 let mut src = String::from("fn f() {\n");
744 for i in 0..n {
745 src.push_str(&format!(" let v{i} = {i}\n"));
746 }
747 src.push_str("}\n");
748 let result = compiled.apply(&src).unwrap();
749 assert_eq!(result.edits.len(), n, "every binding keyword rewritten");
750 assert!(result.edits.iter().all(|e| e.before == "let"));
751 assert!(!result.rewritten.contains("let v"));
752 assert_eq!(result.rewritten.matches("const v").count(), n);
753 }
754
755 #[test]
756 fn raw_query_without_root_capture_is_a_compile_error() {
757 let parsed = Rule::from_toml_str(
758 r#"
759 id = "no-root"
760 language = "harn"
761 fix = "const"
762 [rule]
763 query = '(let_binding "let" @kw)'
764 "#,
765 )
766 .unwrap();
767 let compiled = CompiledRule::compile(&parsed);
768 assert!(
769 matches!(&compiled, Err(RulesError::PatternCompile { message, .. }) if message.contains("__match")),
770 "expected a missing-@__match compile error"
771 );
772 }
773
774 #[test]
775 fn fix_target_naming_an_unbound_capture_errors_on_apply() {
776 let compiled = rule(
777 r#"
778 id = "bad-target"
779 language = "harn"
780 fix = "const"
781 fixTarget = "missing"
782 [rule]
783 query = '(let_binding "let" @kw) @__match'
784 "#,
785 );
786 let result = compiled.apply("fn f() {\n let x = 1\n}\n");
787 assert!(
788 matches!(&result, Err(RulesError::PatternCompile { message, .. }) if message.contains("missing")),
789 "expected an unbound-fixTarget error"
790 );
791 }
792
793 #[test]
794 fn unknown_language_is_an_error() {
795 let parsed = Rule::from_toml_str(
796 r#"
797 id = "x"
798 language = "cobol"
799 [rule]
800 kind = "foo"
801 "#,
802 )
803 .unwrap();
804 assert!(matches!(
805 CompiledRule::compile(&parsed),
806 Err(RulesError::UnknownLanguage { .. })
807 ));
808 }
809
810 #[test]
811 fn invalid_pattern_surfaces_compile_error() {
812 let parsed = Rule::from_toml_str(
813 r#"
814 id = "x"
815 language = "typescript"
816 [rule]
817 pattern = "foo($$$ARGS)"
818 "#,
819 )
820 .unwrap();
821 assert!(matches!(
822 CompiledRule::compile(&parsed),
823 Err(RulesError::PatternCompile { .. })
824 ));
825 }
826
827 #[test]
828 fn harn_resolves_same_named_call_sites_by_binding_identity() {
829 let compiled = rule(
830 r#"
831 id = "top-level-target"
832 language = "harn"
833 [rule]
834 pattern = "$FN($ARG)"
835
836 [[where]]
837 metavar = "FN"
838 resolvesTo = { name = "target", kind = "fn", line = 1 }
839 "#,
840 );
841 let source = r"fn target(value: int) -> int {
842 return value
843}
844
845fn call_shadowed(target: fn(int) -> int) {
846 target(1)
847}
848
849fn call_global() {
850 target(2)
851}
852";
853 let matches = compiled.run(source).unwrap();
854 assert_eq!(matches.len(), 1);
855 assert_eq!(matches[0].text, "target(2)");
856 let binding = &matches[0].bindings["FN"];
857 let resolved = binding.metadata.resolved.as_ref().unwrap();
858 assert_eq!(resolved.name, "target");
859 assert_eq!(resolved.kind, "fn");
860 assert_eq!(resolved.span.start_row, 0);
861 assert_eq!(binding.metadata.ty.as_deref(), Some("fn(int) -> int"));
862 }
863
864 #[test]
865 fn harn_capture_type_constraint_filters_matches() {
866 let compiled = rule(
867 r#"
868 id = "int-logs"
869 language = "harn"
870 [rule]
871 pattern = "log($VALUE)"
872
873 [[where]]
874 metavar = "VALUE"
875 type = "int"
876 "#,
877 );
878 let source = r#"fn main() {
879 let count: int = 1
880 let label: string = "one"
881 log(count)
882 log(label)
883}
884"#;
885 let matches = compiled.run(source).unwrap();
886 assert_eq!(matches.len(), 1);
887 let value = &matches[0].bindings["VALUE"];
888 assert_eq!(value.text, "count");
889 assert_eq!(value.metadata.ty.as_deref(), Some("int"));
890 assert_eq!(
891 value
892 .metadata
893 .resolved
894 .as_ref()
895 .map(|resolved| resolved.kind.as_str()),
896 Some("let")
897 );
898 }
899
900 #[test]
901 fn harn_initializer_uses_outer_binding_scope() {
902 let compiled = rule(
903 r#"
904 id = "outer-initializer"
905 language = "harn"
906 [rule]
907 pattern = "log($VALUE)"
908
909 [[where]]
910 metavar = "VALUE"
911 resolvesTo = { name = "value", kind = "let", line = 2 }
912 "#,
913 );
914 let source = r"fn main() {
915 let value: int = 1
916 if true {
917 let value: string = log(value)
918 }
919}
920";
921 let matches = compiled.run(source).unwrap();
922 assert_eq!(matches.len(), 1);
923 let value = &matches[0].bindings["VALUE"];
924 assert_eq!(value.text, "value");
925 assert_eq!(
926 value
927 .metadata
928 .resolved
929 .as_ref()
930 .map(|resolved| resolved.span.start_row),
931 Some(1)
932 );
933 }
934}