Skip to main content

tree_house/
query_iter.rs

1use core::slice;
2use std::iter::Peekable;
3use std::mem::replace;
4use std::ops::RangeBounds;
5
6use hashbrown::{HashMap, HashSet};
7use ropey::RopeSlice;
8
9use crate::{
10    locals::{Scope, ScopeCursor},
11    Injection, Language, Layer, Range, Syntax, TREE_SITTER_MATCH_LIMIT,
12};
13use tree_sitter::{
14    Capture, InactiveQueryCursor, Node, Pattern, Query, QueryCursor, QueryMatch, RopeInput,
15};
16
17/// A single capture produced by the capture-based iterator.
18#[derive(Debug, Clone)]
19pub struct MatchedNode<'tree> {
20    pub match_id: u32,
21    pub pattern: Pattern,
22    pub node: Node<'tree>,
23    pub capture: Capture,
24    pub scope: Scope,
25}
26
27/// A complete match produced by the match-based iterator.
28/// Contains all captures for one tree-sitter pattern match.
29#[derive(Debug, Clone)]
30pub struct CapturedMatch<'tree> {
31    pub pattern: Pattern,
32    pub nodes: Box<[tree_sitter::MatchedNode<'tree>]>,
33}
34
35impl<'tree> CapturedMatch<'tree> {
36    pub fn nodes_for_capture(&self, capture: Capture) -> impl Iterator<Item = &Node<'tree>> {
37        self.nodes
38            .iter()
39            .filter(move |n| n.capture == capture)
40            .map(|n| &n.node)
41    }
42}
43
44/// Abstracts over capture-based vs match-based cursor advancement.
45pub(crate) trait IterStrategy<'a, 'tree>: Sized {
46    type Peeked: 'tree;
47
48    /// Advance the cursor to produce the next peeked item.
49    /// Returns `None` and leaves `cursor` as `None` when exhausted.
50    fn next_item<Loader: QueryLoader<'a>>(
51        cursor: &mut Option<QueryCursor<'a, 'tree, RopeInput<'a>>>,
52        source: RopeSlice<'_>,
53        scope_cursor: &mut ScopeCursor<'tree>,
54        language: Language,
55        loader: &Loader,
56    ) -> Option<Self::Peeked>;
57
58    fn start_byte(peeked: &Self::Peeked) -> u32;
59    fn end_byte(peeked: &Self::Peeked) -> u32;
60    fn is_empty_range(peeked: &Self::Peeked) -> bool {
61        Self::start_byte(peeked) == Self::end_byte(peeked)
62    }
63}
64
65pub(crate) struct CaptureStrategy;
66
67impl<'a, 'tree> IterStrategy<'a, 'tree> for CaptureStrategy {
68    type Peeked = MatchedNode<'tree>;
69
70    fn next_item<Loader: QueryLoader<'a>>(
71        cursor: &mut Option<QueryCursor<'a, 'tree, RopeInput<'a>>>,
72        source: RopeSlice<'_>,
73        scope_cursor: &mut ScopeCursor<'tree>,
74        language: Language,
75        loader: &Loader,
76    ) -> Option<Self::Peeked> {
77        loop {
78            let mut cur = cursor.take()?;
79            let (query_match, node_idx) = cur.next_matched_node()?;
80            let node = query_match.matched_node(node_idx);
81            let match_id = query_match.id();
82            let pattern = query_match.pattern();
83            let range = node.node.byte_range();
84            let scope = scope_cursor.advance(range.start);
85
86            if !loader.are_predicates_satisfied(language, &query_match, source, scope_cursor) {
87                query_match.remove();
88                *cursor = Some(cur);
89                continue;
90            }
91
92            let result = MatchedNode {
93                match_id,
94                pattern,
95                node: node.node.clone(),
96                capture: node.capture,
97                scope,
98            };
99            *cursor = Some(cur);
100            return Some(result);
101        }
102    }
103
104    fn start_byte(peeked: &Self::Peeked) -> u32 {
105        peeked.node.start_byte()
106    }
107
108    fn end_byte(peeked: &Self::Peeked) -> u32 {
109        peeked.node.end_byte()
110    }
111
112    fn is_empty_range(peeked: &Self::Peeked) -> bool {
113        peeked.node.byte_range().is_empty()
114    }
115}
116
117pub(crate) struct MatchStrategy;
118
119impl<'a, 'tree> IterStrategy<'a, 'tree> for MatchStrategy {
120    type Peeked = CapturedMatch<'tree>;
121
122    fn next_item<Loader: QueryLoader<'a>>(
123        cursor: &mut Option<QueryCursor<'a, 'tree, RopeInput<'a>>>,
124        source: RopeSlice<'_>,
125        scope_cursor: &mut ScopeCursor<'tree>,
126        language: Language,
127        loader: &Loader,
128    ) -> Option<Self::Peeked> {
129        loop {
130            let mut cur = cursor.take()?;
131            let query_match = cur.next_match()?;
132
133            let start = query_match
134                .matched_nodes()
135                .map(|n| n.node.start_byte())
136                .min()
137                .unwrap_or(0);
138            scope_cursor.advance(start);
139
140            if !loader.are_predicates_satisfied(language, &query_match, source, scope_cursor) {
141                query_match.remove();
142                *cursor = Some(cur);
143                continue;
144            }
145
146            let pattern = query_match.pattern();
147            let nodes: Box<[_]> = query_match.matched_nodes().cloned().collect();
148            *cursor = Some(cur);
149            return Some(CapturedMatch { pattern, nodes });
150        }
151    }
152
153    fn start_byte(peeked: &Self::Peeked) -> u32 {
154        peeked
155            .nodes
156            .iter()
157            .map(|n| n.node.start_byte())
158            .min()
159            .unwrap_or(0)
160    }
161
162    fn end_byte(peeked: &Self::Peeked) -> u32 {
163        peeked
164            .nodes
165            .iter()
166            .map(|n| n.node.end_byte())
167            .max()
168            .unwrap_or(0)
169    }
170}
171
172struct LayerIter<'a, 'tree, S: IterStrategy<'a, 'tree>> {
173    cursor: Option<QueryCursor<'a, 'tree, RopeInput<'a>>>,
174    peeked: Option<S::Peeked>,
175    language: Language,
176    scope_cursor: ScopeCursor<'tree>,
177}
178
179impl<'a, 'tree, S: IterStrategy<'a, 'tree>> LayerIter<'a, 'tree, S> {
180    fn peek<Loader: QueryLoader<'a>>(
181        &mut self,
182        source: RopeSlice<'_>,
183        loader: &Loader,
184    ) -> Option<&S::Peeked> {
185        if self.peeked.is_none() {
186            self.peeked = S::next_item(
187                &mut self.cursor,
188                source,
189                &mut self.scope_cursor,
190                self.language,
191                loader,
192            );
193        }
194        self.peeked.as_ref()
195    }
196
197    fn consume(&mut self) -> S::Peeked {
198        self.peeked.take().unwrap()
199    }
200
201    fn has_peeked(&self) -> bool {
202        self.peeked.is_some()
203    }
204}
205
206struct ActiveLayer<'a, 'tree, S: IterStrategy<'a, 'tree>, LayerState> {
207    state: LayerState,
208    layer_iter: LayerIter<'a, 'tree, S>,
209    injections: Peekable<slice::Iter<'a, Injection>>,
210}
211
212struct QueryIterLayerManager<'a, 'tree, Loader, S: IterStrategy<'a, 'tree>, LayerState> {
213    range: Range,
214    loader: Loader,
215    src: RopeSlice<'a>,
216    syntax: &'tree Syntax,
217    active_layers: HashMap<Layer, Box<ActiveLayer<'a, 'tree, S, LayerState>>>,
218    active_injections: Vec<Injection>,
219    /// Layers which are known to have no more captures.
220    finished_layers: HashSet<Layer>,
221}
222
223impl<'a, 'tree: 'a, Loader, S, LayerState> QueryIterLayerManager<'a, 'tree, Loader, S, LayerState>
224where
225    Loader: QueryLoader<'a>,
226    S: IterStrategy<'a, 'tree>,
227    LayerState: Default,
228{
229    fn init_layer(&mut self, injection: Injection) -> Box<ActiveLayer<'a, 'tree, S, LayerState>> {
230        self.active_layers
231            .remove(&injection.layer)
232            .unwrap_or_else(|| {
233                let layer = self.syntax.layer(injection.layer);
234                let start_point = injection.range.start.max(self.range.start);
235                let injection_start = layer
236                    .injections
237                    .partition_point(|child| child.range.end < start_point);
238                let cursor = if self.finished_layers.contains(&injection.layer) {
239                    None
240                } else {
241                    self.loader
242                        .get_query(layer.language)
243                        .and_then(|query| Some((query, layer.tree()?.root_node())))
244                        .map(|(query, node)| {
245                            InactiveQueryCursor::new(self.range.clone(), TREE_SITTER_MATCH_LIMIT)
246                                .execute_query(query, &node, RopeInput::new(self.src))
247                        })
248                };
249                Box::new(ActiveLayer {
250                    state: LayerState::default(),
251                    layer_iter: LayerIter {
252                        language: layer.language,
253                        cursor,
254                        peeked: None,
255                        scope_cursor: layer.locals.scope_cursor(self.range.start),
256                    },
257                    injections: layer.injections[injection_start..].iter().peekable(),
258                })
259            })
260    }
261}
262
263pub(crate) struct BaseIter<
264    'a,
265    'tree,
266    Loader: QueryLoader<'a>,
267    S: IterStrategy<'a, 'tree>,
268    LayerState = (),
269> {
270    layer_manager: Box<QueryIterLayerManager<'a, 'tree, Loader, S, LayerState>>,
271    current_layer: Box<ActiveLayer<'a, 'tree, S, LayerState>>,
272    current_injection: Injection,
273}
274
275impl<'a, 'tree: 'a, Loader, S, LayerState> BaseIter<'a, 'tree, Loader, S, LayerState>
276where
277    Loader: QueryLoader<'a>,
278    S: IterStrategy<'a, 'tree>,
279    LayerState: Default,
280{
281    pub fn new(
282        syntax: &'tree Syntax,
283        src: RopeSlice<'a>,
284        loader: Loader,
285        range: impl RangeBounds<u32>,
286    ) -> Self {
287        let start = match range.start_bound() {
288            std::ops::Bound::Included(&i) => i,
289            std::ops::Bound::Excluded(&i) => i + 1,
290            std::ops::Bound::Unbounded => 0,
291        };
292        let end = match range.end_bound() {
293            std::ops::Bound::Included(&i) => i + 1,
294            std::ops::Bound::Excluded(&i) => i,
295            std::ops::Bound::Unbounded => src.len_bytes() as u32,
296        };
297        let range = start..end;
298        let node = syntax.tree().root_node();
299        let injection = Injection {
300            range: node.byte_range(),
301            layer: syntax.root,
302            matched_node_range: node.byte_range(),
303        };
304        let mut layer_manager = Box::new(QueryIterLayerManager {
305            range,
306            loader,
307            src,
308            syntax,
309            active_layers: HashMap::with_capacity(8),
310            active_injections: Vec::with_capacity(8),
311            finished_layers: HashSet::with_capacity(8),
312        });
313        Self {
314            current_layer: layer_manager.init_layer(injection.clone()),
315            current_injection: injection,
316            layer_manager,
317        }
318    }
319
320    #[inline]
321    pub fn source(&self) -> RopeSlice<'a> {
322        self.layer_manager.src
323    }
324
325    #[inline]
326    pub fn syntax(&self) -> &'tree Syntax {
327        self.layer_manager.syntax
328    }
329
330    #[inline]
331    pub fn loader(&mut self) -> &mut Loader {
332        &mut self.layer_manager.loader
333    }
334
335    #[inline]
336    pub fn current_layer(&self) -> Layer {
337        self.current_injection.layer
338    }
339
340    #[inline]
341    pub fn current_injection(&mut self) -> (Injection, &mut LayerState) {
342        (
343            self.current_injection.clone(),
344            &mut self.current_layer.state,
345        )
346    }
347
348    #[inline]
349    pub fn current_language(&self) -> Language {
350        self.layer_manager
351            .syntax
352            .layer(self.current_injection.layer)
353            .language
354    }
355
356    pub fn layer_state(&mut self, layer: Layer) -> &mut LayerState {
357        if layer == self.current_injection.layer {
358            &mut self.current_layer.state
359        } else {
360            &mut self
361                .layer_manager
362                .active_layers
363                .get_mut(&layer)
364                .unwrap()
365                .state
366        }
367    }
368
369    fn enter_injection(&mut self, injection: Injection) {
370        let active_layer = self.layer_manager.init_layer(injection.clone());
371        let old_injection = replace(&mut self.current_injection, injection);
372        let old_layer = replace(&mut self.current_layer, active_layer);
373        self.layer_manager
374            .active_layers
375            .insert(old_injection.layer, old_layer);
376        self.layer_manager.active_injections.push(old_injection);
377    }
378
379    fn exit_injection(&mut self) -> Option<(Injection, Option<LayerState>)> {
380        let injection = replace(
381            &mut self.current_injection,
382            self.layer_manager.active_injections.pop()?,
383        );
384        let mut layer = replace(
385            &mut self.current_layer,
386            self.layer_manager
387                .active_layers
388                .remove(&self.current_injection.layer)?,
389        );
390        let layer_unfinished = layer.layer_iter.has_peeked() || layer.injections.peek().is_some();
391        if layer_unfinished {
392            self.layer_manager
393                .active_layers
394                .insert(injection.layer, layer);
395            Some((injection, None))
396        } else {
397            self.layer_manager.finished_layers.insert(injection.layer);
398            Some((injection, Some(layer.state)))
399        }
400    }
401}
402
403impl<'a, 'tree: 'a, Loader, S, LayerState> Iterator for BaseIter<'a, 'tree, Loader, S, LayerState>
404where
405    Loader: QueryLoader<'a>,
406    S: IterStrategy<'a, 'tree>,
407    LayerState: Default,
408{
409    type Item = BaseIterEvent<S::Peeked, LayerState>;
410
411    fn next(&mut self) -> Option<Self::Item> {
412        loop {
413            let next_injection = self
414                .current_layer
415                .injections
416                .peek()
417                .filter(|inj| inj.range.start <= self.current_injection.range.end);
418            let next_item = self
419                .current_layer
420                .layer_iter
421                .peek(self.layer_manager.src, &self.layer_manager.loader)
422                .filter(|item| S::start_byte(item) <= self.current_injection.range.end);
423
424            match (next_item, next_injection) {
425                (None, None) => {
426                    return self.exit_injection().map(|(injection, state)| {
427                        BaseIterEvent::ExitInjection { injection, state }
428                    });
429                }
430                (Some(item), _) if S::is_empty_range(item) => {
431                    self.current_layer.layer_iter.consume();
432                    continue;
433                }
434                (Some(_), None) => {
435                    let item = self.current_layer.layer_iter.consume();
436                    return Some(BaseIterEvent::Match(item));
437                }
438                (Some(item), Some(injection)) if S::start_byte(item) < injection.range.end => {
439                    let item = self.current_layer.layer_iter.consume();
440                    if S::start_byte(&item) <= injection.range.start
441                        || injection.range.end < S::end_byte(&item)
442                    {
443                        return Some(BaseIterEvent::Match(item));
444                    }
445                }
446                (Some(_), Some(_)) | (None, Some(_)) => {
447                    let injection = self.current_layer.injections.next().unwrap();
448                    self.enter_injection(injection.clone());
449                    return Some(BaseIterEvent::EnterInjection(injection.clone()));
450                }
451            }
452        }
453    }
454}
455
456pub(crate) enum BaseIterEvent<Item, State = ()> {
457    EnterInjection(Injection),
458    Match(Item),
459    ExitInjection {
460        injection: Injection,
461        state: Option<State>,
462    },
463}
464
465pub struct QueryIter<'a, 'tree, Loader: QueryLoader<'a>, LayerState = ()>(
466    BaseIter<'a, 'tree, Loader, CaptureStrategy, LayerState>,
467);
468
469impl<'a, 'tree: 'a, Loader, LayerState> QueryIter<'a, 'tree, Loader, LayerState>
470where
471    Loader: QueryLoader<'a>,
472    LayerState: Default,
473{
474    pub fn new(
475        syntax: &'tree Syntax,
476        src: RopeSlice<'a>,
477        loader: Loader,
478        range: impl RangeBounds<u32>,
479    ) -> Self {
480        Self(BaseIter::new(syntax, src, loader, range))
481    }
482
483    #[inline]
484    pub fn source(&self) -> RopeSlice<'a> {
485        self.0.source()
486    }
487
488    #[inline]
489    pub fn syntax(&self) -> &'tree Syntax {
490        self.0.syntax()
491    }
492
493    #[inline]
494    pub fn loader(&mut self) -> &mut Loader {
495        self.0.loader()
496    }
497
498    #[inline]
499    pub fn current_layer(&self) -> Layer {
500        self.0.current_layer()
501    }
502
503    #[inline]
504    pub fn current_injection(&mut self) -> (Injection, &mut LayerState) {
505        self.0.current_injection()
506    }
507
508    #[inline]
509    pub fn current_language(&self) -> Language {
510        self.0.current_language()
511    }
512
513    pub fn layer_state(&mut self, layer: Layer) -> &mut LayerState {
514        self.0.layer_state(layer)
515    }
516}
517
518impl<'a, 'tree: 'a, Loader, LayerState> Iterator for QueryIter<'a, 'tree, Loader, LayerState>
519where
520    Loader: QueryLoader<'a>,
521    LayerState: Default,
522{
523    type Item = QueryIterEvent<'tree, LayerState>;
524
525    fn next(&mut self) -> Option<Self::Item> {
526        self.0.next().map(|event| match event {
527            BaseIterEvent::EnterInjection(i) => QueryIterEvent::EnterInjection(i),
528            BaseIterEvent::Match(m) => QueryIterEvent::Match(m),
529            BaseIterEvent::ExitInjection { injection, state } => {
530                QueryIterEvent::ExitInjection { injection, state }
531            }
532        })
533    }
534}
535
536#[derive(Debug)]
537pub enum QueryIterEvent<'tree, State = ()> {
538    EnterInjection(Injection),
539    Match(MatchedNode<'tree>),
540    ExitInjection {
541        injection: Injection,
542        state: Option<State>,
543    },
544}
545
546impl<S> QueryIterEvent<'_, S> {
547    pub fn start_byte(&self) -> u32 {
548        match self {
549            QueryIterEvent::EnterInjection(injection) => injection.range.start,
550            QueryIterEvent::Match(mat) => mat.node.start_byte(),
551            QueryIterEvent::ExitInjection { injection, .. } => injection.range.end,
552        }
553    }
554}
555
556pub struct QueryMatchIter<'a, 'tree, Loader: QueryLoader<'a>, LayerState = ()>(
557    BaseIter<'a, 'tree, Loader, MatchStrategy, LayerState>,
558);
559
560impl<'a, 'tree: 'a, Loader, LayerState> QueryMatchIter<'a, 'tree, Loader, LayerState>
561where
562    Loader: QueryLoader<'a>,
563    LayerState: Default,
564{
565    pub fn new(
566        syntax: &'tree Syntax,
567        src: RopeSlice<'a>,
568        loader: Loader,
569        range: impl RangeBounds<u32>,
570    ) -> Self {
571        Self(BaseIter::new(syntax, src, loader, range))
572    }
573
574    #[inline]
575    pub fn current_language(&self) -> Language {
576        self.0.current_language()
577    }
578
579    #[inline]
580    pub fn current_layer(&self) -> Layer {
581        self.0.current_layer()
582    }
583}
584
585impl<'a, 'tree: 'a, Loader, LayerState> Iterator for QueryMatchIter<'a, 'tree, Loader, LayerState>
586where
587    Loader: QueryLoader<'a>,
588    LayerState: Default,
589{
590    type Item = QueryMatchIterEvent<'tree, LayerState>;
591
592    fn next(&mut self) -> Option<Self::Item> {
593        self.0.next().map(|event| match event {
594            BaseIterEvent::EnterInjection(i) => QueryMatchIterEvent::EnterInjection(i),
595            BaseIterEvent::Match(m) => QueryMatchIterEvent::Match(m),
596            BaseIterEvent::ExitInjection { injection, state } => {
597                QueryMatchIterEvent::ExitInjection { injection, state }
598            }
599        })
600    }
601}
602
603#[derive(Debug)]
604pub enum QueryMatchIterEvent<'tree, State = ()> {
605    EnterInjection(Injection),
606    Match(CapturedMatch<'tree>),
607    ExitInjection {
608        injection: Injection,
609        state: Option<State>,
610    },
611}
612
613pub trait QueryLoader<'a> {
614    fn get_query(&mut self, lang: Language) -> Option<&'a Query>;
615
616    fn are_predicates_satisfied(
617        &self,
618        _lang: Language,
619        _match: &QueryMatch<'_, '_>,
620        _source: RopeSlice<'_>,
621        _locals_cursor: &ScopeCursor<'_>,
622    ) -> bool {
623        true
624    }
625}
626
627impl<'a, F> QueryLoader<'a> for F
628where
629    F: FnMut(Language) -> Option<&'a Query>,
630{
631    fn get_query(&mut self, lang: Language) -> Option<&'a Query> {
632        (self)(lang)
633    }
634}