Skip to main content

ast_grep_core/tree_sitter/
mod.rs

1pub mod traversal;
2
3use crate::node::Root;
4use crate::replacer::Replacer;
5use crate::source::{Content, Doc, Edit, SgNode};
6use crate::{AstGrep, Matcher};
7use crate::{Language, Position, node::KindId};
8use std::borrow::Cow;
9use std::cell::RefCell;
10use std::collections::HashMap;
11use std::collections::hash_map::Entry;
12use std::num::NonZero;
13use thiserror::Error;
14pub use traversal::{TsPre, Visitor};
15pub use tree_sitter::Language as TSLanguage;
16use tree_sitter::{InputEdit, LanguageError, Node, Parser, Point, Tree};
17pub use tree_sitter::{Point as TSPoint, Range as TSRange};
18
19/// Represents tree-sitter related error
20#[derive(Debug, Error)]
21pub enum TSParseError {
22  #[error("incompatible `Language` is assigned to a `Parser`.")]
23  Language(#[from] LanguageError),
24  /// A general error when tree sitter fails to parse in time. It can be caused by
25  /// the following reasons but tree-sitter does not provide error detail.
26  /// * The timeout set with [Parser::set_timeout_micros] expired
27  /// * The cancellation flag set with [Parser::set_cancellation_flag] was flipped
28  /// * The parser has not yet had a language assigned with [Parser::set_language]
29  #[error("general error when tree-sitter fails to parse.")]
30  TreeUnavailable,
31}
32
33thread_local! {
34  /// Keep one parser for each language on every worker thread. A parser is not
35  /// shared between workers, so parsing remains lock-free, while repeated files
36  /// avoid recreating the parser and rebuilding its language-specific tables.
37  static PARSER_CACHE: RefCell<HashMap<TSLanguage, Parser>> = RefCell::new(HashMap::new());
38}
39
40#[inline]
41fn parse_lang(
42  parse_fn: impl Fn(&mut Parser) -> Option<Tree>,
43  ts_lang: TSLanguage,
44) -> Result<Tree, TSParseError> {
45  PARSER_CACHE.with(|cache| {
46    let mut cache = cache.borrow_mut();
47    let parser = match cache.entry(ts_lang) {
48      Entry::Occupied(entry) => entry.into_mut(),
49      Entry::Vacant(entry) => {
50        let mut parser = Parser::new();
51        parser.set_language(entry.key())?;
52        entry.insert(parser)
53      }
54    };
55    // A failed or cancelled parse can retain resumable parser state. Every
56    // parse_lang call represents a complete document, so always start clean.
57    parser.reset();
58    parse_fn(parser).ok_or(TSParseError::TreeUnavailable)
59  })
60}
61
62#[derive(Clone)]
63pub struct StrDoc<L: LanguageExt> {
64  pub src: String,
65  pub lang: L,
66  pub tree: Tree,
67}
68
69impl<L: LanguageExt> StrDoc<L> {
70  pub fn try_new(src: &str, lang: L) -> Result<Self, String> {
71    let src = src.to_string();
72    let ts_lang = lang.get_ts_language();
73    let tree = parse_lang(|p| p.parse(src.as_bytes(), None), ts_lang).map_err(|e| e.to_string())?;
74    Ok(Self { src, lang, tree })
75  }
76  pub fn new(src: &str, lang: L) -> Self {
77    Self::try_new(src, lang).expect("Parser tree error")
78  }
79  fn parse(&self, old_tree: Option<&Tree>) -> Result<Tree, TSParseError> {
80    let source = self.get_source();
81    let lang = self.get_lang().get_ts_language();
82    parse_lang(|p| p.parse(source.as_bytes(), old_tree), lang)
83  }
84}
85
86impl<L: LanguageExt> Doc for StrDoc<L> {
87  type Source = String;
88  type Lang = L;
89  type Node<'r> = Node<'r>;
90  fn get_lang(&self) -> &Self::Lang {
91    &self.lang
92  }
93  fn get_source(&self) -> &Self::Source {
94    &self.src
95  }
96  fn do_edit(&mut self, edit: &Edit<Self::Source>) -> Result<(), String> {
97    let source = &mut self.src;
98    perform_edit(&mut self.tree, source, edit);
99    self.tree = self.parse(Some(&self.tree)).map_err(|e| e.to_string())?;
100    Ok(())
101  }
102  fn root_node(&self) -> Node<'_> {
103    self.tree.root_node()
104  }
105  fn get_node_text<'a>(&'a self, node: &Self::Node<'a>) -> Cow<'a, str> {
106    Cow::Borrowed(
107      node
108        .utf8_text(self.src.as_bytes())
109        .expect("invalid source text encoding"),
110    )
111  }
112}
113
114struct NodeWalker<'tree> {
115  cursor: tree_sitter::TreeCursor<'tree>,
116  count: usize,
117}
118
119impl<'tree> Iterator for NodeWalker<'tree> {
120  type Item = Node<'tree>;
121  fn next(&mut self) -> Option<Self::Item> {
122    if self.count == 0 {
123      return None;
124    }
125    let ret = Some(self.cursor.node());
126    self.cursor.goto_next_sibling();
127    self.count -= 1;
128    ret
129  }
130}
131
132impl ExactSizeIterator for NodeWalker<'_> {
133  fn len(&self) -> usize {
134    self.count
135  }
136}
137
138impl<'r> SgNode<'r> for Node<'r> {
139  fn parent(&self) -> Option<Self> {
140    Node::parent(self)
141  }
142  fn ancestors(&self, root: Self) -> impl Iterator<Item = Self> {
143    let mut ancestor = Some(root);
144    let self_id = self.id();
145    std::iter::from_fn(move || {
146      let inner = ancestor.take()?;
147      if inner.id() == self_id {
148        return None;
149      }
150      ancestor = inner.child_with_descendant(*self);
151      Some(inner)
152    })
153    // We must iterate up the tree to preserve backwards compatibility
154    .collect::<Vec<_>>()
155    .into_iter()
156    .rev()
157  }
158  fn dfs(&self) -> impl Iterator<Item = Self> {
159    TsPre::new(self)
160  }
161  fn child(&self, nth: usize) -> Option<Self> {
162    Node::child(self, nth as u32)
163  }
164  fn children(&self) -> impl ExactSizeIterator<Item = Self> {
165    let mut cursor = self.walk();
166    cursor.goto_first_child();
167    NodeWalker {
168      cursor,
169      count: self.child_count(),
170    }
171  }
172  fn child_by_field_id(&self, field_id: u16) -> Option<Self> {
173    Node::child_by_field_id(self, field_id)
174  }
175  fn next(&self) -> Option<Self> {
176    self.next_sibling()
177  }
178  fn prev(&self) -> Option<Self> {
179    self.prev_sibling()
180  }
181  fn next_all(&self) -> impl Iterator<Item = Self> {
182    // if root is none, use self as fallback to return a type-stable Iterator
183    let node = self.parent().unwrap_or(*self);
184    let mut cursor = node.walk();
185    cursor.goto_first_child_for_byte(self.start_byte());
186    std::iter::from_fn(move || {
187      if cursor.goto_next_sibling() {
188        Some(cursor.node())
189      } else {
190        None
191      }
192    })
193  }
194  fn prev_all(&self) -> impl Iterator<Item = Self> {
195    // if root is none, use self as fallback to return a type-stable Iterator
196    let node = self.parent().unwrap_or(*self);
197    let mut cursor = node.walk();
198    cursor.goto_first_child_for_byte(self.start_byte());
199    std::iter::from_fn(move || {
200      if cursor.goto_previous_sibling() {
201        Some(cursor.node())
202      } else {
203        None
204      }
205    })
206  }
207  fn is_named(&self) -> bool {
208    Node::is_named(self)
209  }
210  /// N.B. it is different from is_named && is_leaf
211  /// if a node has no named children.
212  fn is_named_leaf(&self) -> bool {
213    self.named_child_count() == 0
214  }
215  fn is_leaf(&self) -> bool {
216    self.child_count() == 0
217  }
218  fn kind(&self) -> Cow<'_, str> {
219    Cow::Borrowed(Node::kind(self))
220  }
221  fn kind_id(&self) -> KindId {
222    Node::kind_id(self)
223  }
224  fn node_id(&self) -> usize {
225    self.id()
226  }
227  fn range(&self) -> std::ops::Range<usize> {
228    self.start_byte()..self.end_byte()
229  }
230  fn start_pos(&self) -> Position {
231    let pos = self.start_position();
232    let byte = self.start_byte();
233    Position::new(pos.row, pos.column, byte)
234  }
235  fn end_pos(&self) -> Position {
236    let pos = self.end_position();
237    let byte = self.end_byte();
238    Position::new(pos.row, pos.column, byte)
239  }
240  // missing node is a tree-sitter specific concept
241  fn is_missing(&self) -> bool {
242    Node::is_missing(self)
243  }
244  fn is_error(&self) -> bool {
245    Node::is_error(self)
246  }
247  fn is_extra(&self) -> bool {
248    Node::is_extra(self)
249  }
250
251  fn field(&self, name: &str) -> Option<Self> {
252    self.child_by_field_name(name)
253  }
254  fn field_children(&self, field_id: Option<u16>) -> impl Iterator<Item = Self> {
255    let field_id = field_id.and_then(NonZero::new);
256    let mut cursor = self.walk();
257    let has_children = cursor.goto_first_child();
258    // if field_id is not found, iteration is done
259    let mut done = field_id.is_none() || !has_children;
260
261    std::iter::from_fn(move || {
262      if done {
263        return None;
264      }
265      while cursor.field_id() != field_id {
266        if !cursor.goto_next_sibling() {
267          return None;
268        }
269      }
270      let ret = cursor.node();
271      if !cursor.goto_next_sibling() {
272        done = true;
273      }
274      Some(ret)
275    })
276  }
277}
278
279pub fn perform_edit<S: ContentExt>(tree: &mut Tree, input: &mut S, edit: &Edit<S>) -> InputEdit {
280  let edit = input.accept_edit(edit);
281  tree.edit(&edit);
282  edit
283}
284
285/// tree-sitter specific language trait
286pub trait LanguageExt: Language {
287  /// Create an [`AstGrep`] instance for the language
288  fn ast_grep<S: AsRef<str>>(&self, source: S) -> AstGrep<StrDoc<Self>> {
289    AstGrep::new(source, self.clone())
290  }
291
292  /// tree sitter language to parse the source
293  fn get_ts_language(&self) -> TSLanguage;
294
295  fn injectable_languages(&self) -> Option<&'static [&'static str]> {
296    None
297  }
298
299  /// Get injected language regions in the root document. e.g. get JavaScripts in HTML.
300  /// Each entry is parsed as an **independent** tree-sitter document.
301  /// Multiple entries for the same language produce separate parse trees.
302  /// Also see <https://tree-sitter.github.io/tree-sitter/using-parsers#multi-language-documents>
303  fn extract_injections<L: LanguageExt>(
304    &self,
305    _root: crate::Node<StrDoc<L>>,
306  ) -> Vec<(String, Vec<TSRange>)> {
307    Vec::new()
308  }
309}
310
311fn position_for_offset(input: &[u8], offset: usize) -> Point {
312  debug_assert!(offset <= input.len());
313  let (mut row, mut col) = (0, 0);
314  for c in &input[0..offset] {
315    if *c as char == '\n' {
316      row += 1;
317      col = 0;
318    } else {
319      col += 1;
320    }
321  }
322  Point::new(row, col)
323}
324
325impl<L: LanguageExt> AstGrep<StrDoc<L>> {
326  pub fn new<S: AsRef<str>>(src: S, lang: L) -> Self {
327    Root::str(src.as_ref(), lang)
328  }
329
330  pub fn source(&self) -> &str {
331    self.doc.get_source().as_str()
332  }
333
334  pub fn generate(self) -> String {
335    self.doc.src
336  }
337}
338
339pub trait ContentExt: Content {
340  fn accept_edit(&mut self, edit: &Edit<Self>) -> InputEdit;
341}
342impl ContentExt for String {
343  fn accept_edit(&mut self, edit: &Edit<Self>) -> InputEdit {
344    let start_byte = edit.position;
345    let old_end_byte = edit.position + edit.deleted_length;
346    let new_end_byte = edit.position + edit.inserted_text.len();
347    let input = unsafe { self.as_mut_vec() };
348    let start_position = position_for_offset(input, start_byte);
349    let old_end_position = position_for_offset(input, old_end_byte);
350    input.splice(start_byte..old_end_byte, edit.inserted_text.clone());
351    let new_end_position = position_for_offset(input, new_end_byte);
352    InputEdit {
353      start_byte,
354      old_end_byte,
355      new_end_byte,
356      start_position,
357      old_end_position,
358      new_end_position,
359    }
360  }
361}
362
363impl<L: LanguageExt> Root<StrDoc<L>> {
364  pub fn str(src: &str, lang: L) -> Self {
365    Self::try_new(src, lang).expect("should parse")
366  }
367  pub fn try_new(src: &str, lang: L) -> Result<Self, String> {
368    let doc = StrDoc::try_new(src, lang)?;
369    Ok(Self { doc })
370  }
371  pub fn get_text(&self) -> &str {
372    &self.doc.src
373  }
374
375  pub fn get_injections<F: Fn(&str) -> Option<L>>(&self, get_lang: F) -> Vec<Self> {
376    let root = self.root();
377    let range = self.lang().extract_injections(root);
378
379    range
380      .into_iter()
381      .filter_map(|(lang_str, ranges)| {
382        let lang = get_lang(&lang_str)?;
383        let source = self.doc.get_source();
384        let mut parser = Parser::new();
385        parser.set_included_ranges(&ranges).ok()?;
386        parser.set_language(&lang.get_ts_language()).ok()?;
387        let tree = parser.parse(source, None)?;
388        Some(Self {
389          doc: StrDoc {
390            src: self.doc.src.clone(),
391            lang,
392            tree,
393          },
394        })
395      })
396      .collect()
397  }
398}
399
400pub struct DisplayContext<'r> {
401  /// content for the matched node
402  pub matched: Cow<'r, str>,
403  /// content before the matched node
404  pub leading: &'r str,
405  /// content after the matched node
406  pub trailing: &'r str,
407  /// zero-based start line of the context
408  pub start_line: usize,
409}
410
411/// these methods are only for `StrDoc`
412impl<'r, L: LanguageExt> crate::Node<'r, StrDoc<L>> {
413  #[doc(hidden)]
414  pub fn display_context(&self, before: usize, after: usize) -> DisplayContext<'r> {
415    let source = self.root.doc.get_source().as_str();
416    let bytes = source.as_bytes();
417    let start = self.inner.start_byte();
418    let end = self.inner.end_byte();
419    let (mut leading, mut trailing) = (start, end);
420    let mut lines_before = before + 1;
421    while leading > 0 {
422      if bytes[leading - 1] == b'\n' {
423        lines_before -= 1;
424        if lines_before == 0 {
425          break;
426        }
427      }
428      leading -= 1;
429    }
430    let mut lines_after = after + 1;
431    // tree-sitter will append line ending to source so trailing can be out of bound
432    trailing = trailing.min(bytes.len());
433    while trailing < bytes.len() {
434      if bytes[trailing] == b'\n' {
435        lines_after -= 1;
436        if lines_after == 0 {
437          break;
438        }
439      }
440      trailing += 1;
441    }
442    // lines_before means we matched all context, offset is `before` itself
443    let offset = if lines_before == 0 {
444      before
445    } else {
446      // otherwise, there are fewer than `before` line in src, compute the actual line
447      before + 1 - lines_before
448    };
449    DisplayContext {
450      matched: self.text(),
451      leading: &source[leading..start],
452      trailing: &source[end..trailing],
453      start_line: self.start_pos().line() - offset,
454    }
455  }
456
457  pub fn replace_all<M: Matcher, R: Replacer<StrDoc<L>>>(
458    &self,
459    matcher: M,
460    replacer: R,
461  ) -> Vec<Edit<String>> {
462    // TODO: support nested matches like Some(Some(1)) with pattern Some($A)
463    Visitor::new(&matcher)
464      .reentrant(false)
465      .visit(self.clone())
466      .map(|matched| matched.make_edit(&matcher, &replacer))
467      .collect()
468  }
469}
470
471#[cfg(test)]
472mod test {
473  use super::*;
474  use crate::language::Tsx;
475  use tree_sitter::Point;
476
477  fn parse(src: &str) -> Result<Tree, TSParseError> {
478    parse_lang(|p| p.parse(src, None), Tsx.get_ts_language())
479  }
480
481  #[test]
482  fn test_tree_sitter() -> Result<(), TSParseError> {
483    let tree = parse("var a = 1234")?;
484    let root_node = tree.root_node();
485    assert_eq!(root_node.kind(), "program");
486    assert_eq!(root_node.start_position().column, 0);
487    assert_eq!(root_node.end_position().column, 12);
488    assert_eq!(
489      root_node.to_sexp(),
490      "(program (variable_declaration (variable_declarator name: (identifier) value: (number))))"
491    );
492    Ok(())
493  }
494
495  #[test]
496  fn test_parser_cache_reuses_parser_per_language() -> Result<(), TSParseError> {
497    PARSER_CACHE.with(|cache| cache.borrow_mut().clear());
498
499    parse("let one = 1")?;
500    parse("let two = 2")?;
501    assert_eq!(PARSER_CACHE.with(|cache| cache.borrow().len()), 1);
502
503    let ts_lang = tree_sitter_typescript::LANGUAGE_TYPESCRIPT.into();
504    parse_lang(|p| p.parse("let three = 3", None), ts_lang)?;
505    assert_eq!(PARSER_CACHE.with(|cache| cache.borrow().len()), 2);
506    Ok(())
507  }
508
509  #[test]
510  fn test_object_literal() -> Result<(), TSParseError> {
511    let tree = parse("{a: $X}")?;
512    let root_node = tree.root_node();
513    // wow this is not label. technically it is wrong but practically it is better LOL
514    assert_eq!(
515      root_node.to_sexp(),
516      "(program (expression_statement (object (pair key: (property_identifier) value: (identifier)))))"
517    );
518    Ok(())
519  }
520
521  #[test]
522  fn test_string() -> Result<(), TSParseError> {
523    let tree = parse("'$A'")?;
524    let root_node = tree.root_node();
525    assert_eq!(
526      root_node.to_sexp(),
527      "(program (expression_statement (string (string_fragment))))"
528    );
529    Ok(())
530  }
531
532  #[test]
533  fn test_row_col() -> Result<(), TSParseError> {
534    let tree = parse("😄")?;
535    let root = tree.root_node();
536    assert_eq!(root.start_position(), Point::new(0, 0));
537    // NOTE: Point in tree-sitter is counted in bytes instead of char
538    assert_eq!(root.end_position(), Point::new(0, 4));
539    Ok(())
540  }
541
542  #[test]
543  fn test_edit() -> Result<(), TSParseError> {
544    let mut src = "a + b".to_string();
545    let mut tree = parse(&src)?;
546    let _ = perform_edit(
547      &mut tree,
548      &mut src,
549      &Edit {
550        position: 1,
551        deleted_length: 0,
552        inserted_text: " * b".into(),
553      },
554    );
555    let tree2 = parse_lang(|p| p.parse(&src, Some(&tree)), Tsx.get_ts_language())?;
556    assert_eq!(
557      tree.root_node().to_sexp(),
558      "(program (expression_statement (binary_expression left: (identifier) right: (identifier))))"
559    );
560    assert_eq!(
561      tree2.root_node().to_sexp(),
562      "(program (expression_statement (binary_expression left: (binary_expression left: (identifier) right: (identifier)) right: (identifier))))"
563    );
564    Ok(())
565  }
566}