ast_grep_core/tree_sitter/
mod.rs1pub 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#[derive(Debug, Error)]
21pub enum TSParseError {
22 #[error("incompatible `Language` is assigned to a `Parser`.")]
23 Language(#[from] LanguageError),
24 #[error("general error when tree-sitter fails to parse.")]
30 TreeUnavailable,
31}
32
33thread_local! {
34 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 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 .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 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 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 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 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 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
285pub trait LanguageExt: Language {
287 fn ast_grep<S: AsRef<str>>(&self, source: S) -> AstGrep<StrDoc<Self>> {
289 AstGrep::new(source, self.clone())
290 }
291
292 fn get_ts_language(&self) -> TSLanguage;
294
295 fn injectable_languages(&self) -> Option<&'static [&'static str]> {
296 None
297 }
298
299 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 pub matched: Cow<'r, str>,
403 pub leading: &'r str,
405 pub trailing: &'r str,
407 pub start_line: usize,
409}
410
411impl<'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 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 let offset = if lines_before == 0 {
444 before
445 } else {
446 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 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 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 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}