Skip to main content

scah_query_ir/query/compiler/
query.rs

1use std::ops::Range;
2
3use super::builder::{QueryBuilder, Save, SelectionKind};
4use super::error::SelectorParseError;
5use super::transition::Transition;
6use crate::query::selector::Combinator;
7
8#[derive(PartialEq, Eq, PartialOrd, Ord, Debug, Clone, Copy)]
9pub struct TransitionId(pub usize);
10
11impl TransitionId {
12    #[inline(always)]
13    pub fn index(self) -> usize {
14        self.0
15    }
16}
17
18impl From<usize> for TransitionId {
19    fn from(value: usize) -> Self {
20        Self(value)
21    }
22}
23
24#[derive(PartialEq, Eq, PartialOrd, Ord, Debug, Clone, Copy)]
25pub struct QuerySectionId(pub usize);
26
27impl QuerySectionId {
28    #[inline(always)]
29    pub fn index(self) -> usize {
30        self.0
31    }
32}
33
34impl From<usize> for QuerySectionId {
35    fn from(value: usize) -> Self {
36        Self(value)
37    }
38}
39
40struct PositionIterator<'query, Q: QuerySpec<'query>> {
41    arena: &'query Q,
42    current: Option<Position>,
43}
44impl<'query, Q: QuerySpec<'query>> Iterator for PositionIterator<'query, Q> {
45    type Item = Position;
46    fn next(&mut self) -> Option<Self::Item> {
47        self.current
48            .inspect(|position| self.current = position.next_sibling(self.arena))
49    }
50}
51
52pub trait QuerySpec<'query> {
53    fn states(&self) -> &[Transition<'query>];
54    fn queries(&self) -> &[QuerySection<'query>];
55    fn exit_at_section_end(&self) -> Option<QuerySectionId>;
56
57    fn requires_text_content(&self) -> bool {
58        self.queries()
59            .iter()
60            .any(|section| section.save.text_content)
61    }
62
63    fn get_transition(&self, state: TransitionId) -> &Transition<'query> {
64        &self.states()[state.index()]
65    }
66
67    fn get_section_selection_kind(&self, section_index: QuerySectionId) -> SelectionKind {
68        self.queries()[section_index.index()].kind
69    }
70
71    fn get_selection(&self, section_index: QuerySectionId) -> &QuerySection<'query> {
72        &self.queries()[section_index.index()]
73    }
74
75    fn is_descendant(&self, state: TransitionId) -> bool {
76        self.get_transition(state).guard == Combinator::Descendant
77    }
78
79    fn is_save_point(&self, position: &Position) -> bool {
80        debug_assert!(
81            self.get_selection(position.selection)
82                .range
83                .contains(&position.state)
84        );
85        self.get_selection(position.selection).range.end.index() - 1 == position.state.index()
86    }
87
88    fn is_last_save_point(&self, position: &Position) -> bool {
89        debug_assert!(position.selection.index() < self.queries().len());
90        let is_last_query = self.queries().len() - 1 == position.selection.index();
91        let is_last_state =
92            self.get_selection(position.selection).range.end.index() - 1 == position.state.index();
93        is_last_query && is_last_state
94    }
95
96    fn children(&'query self, position: &Position) -> Option<impl Iterator<Item = Position>>
97    where
98        Self: Sized,
99    {
100        position.next_child(self).map(|child| PositionIterator {
101            arena: self,
102            current: Some(child),
103        })
104    }
105}
106
107#[derive(PartialEq, Debug, Clone, Copy)]
108pub struct Position {
109    pub selection: QuerySectionId,
110    pub state: TransitionId,
111}
112
113impl Position {
114    pub fn next_transition<'query, Q: QuerySpec<'query> + ?Sized>(
115        &self,
116        query: &Q,
117    ) -> Option<TransitionId> {
118        debug_assert!(self.selection.index() < query.queries().len());
119        debug_assert!(
120            query
121                .get_selection(self.selection)
122                .range
123                .contains(&self.state)
124        );
125
126        let selection_range = &query.get_selection(self.selection).range;
127        if self.state.index() + 1 < selection_range.end.index() {
128            Some(TransitionId(self.state.index() + 1))
129        } else {
130            None
131        }
132    }
133
134    pub fn next_child<'query, Q: QuerySpec<'query> + ?Sized>(&self, query: &Q) -> Option<Self> {
135        debug_assert!(self.selection.index() < query.queries().len());
136        debug_assert!(
137            query
138                .get_selection(self.selection)
139                .range
140                .contains(&self.state)
141        );
142
143        if self.selection.index() == query.queries().len() - 1 {
144            return None;
145        }
146
147        let next_selection_index = QuerySectionId(self.selection.index() + 1);
148        let next_selection = query.get_selection(next_selection_index);
149        if next_selection.parent.is_some_and(|p| p == self.selection) {
150            return Some(Self {
151                selection: next_selection_index,
152                state: next_selection.range.start,
153            });
154        }
155
156        None
157    }
158
159    pub fn next_sibling<'query, Q: QuerySpec<'query> + ?Sized>(&self, query: &Q) -> Option<Self> {
160        debug_assert!(self.selection.index() < query.queries().len());
161        debug_assert!(
162            query
163                .get_selection(self.selection)
164                .range
165                .contains(&self.state)
166        );
167
168        query
169            .get_selection(self.selection)
170            .next_sibling
171            .map(|sibling| Self {
172                selection: sibling,
173                state: query.get_selection(sibling).range.start,
174            })
175    }
176
177    pub fn back<'query, Q: QuerySpec<'query> + ?Sized>(&mut self, query: &Q) {
178        debug_assert!(self.selection.index() < query.queries().len());
179        debug_assert!(self.state < query.get_selection(self.selection).range.end);
180
181        let selection = query.get_selection(self.selection);
182        if self.state.index() > selection.range.start.index() {
183            self.state = TransitionId(self.state.index() - 1);
184        } else if let Some(parent) = selection.parent {
185            self.selection = parent;
186            self.state = TransitionId(query.get_selection(self.selection).range.end.index() - 1);
187        }
188    }
189}
190
191#[derive(Debug, Clone, PartialEq)]
192pub struct QuerySection<'query> {
193    pub source: &'query str,
194    pub range: Range<TransitionId>,
195    pub parent: Option<QuerySectionId>,
196    pub next_sibling: Option<QuerySectionId>,
197    pub save: Save,
198    pub kind: SelectionKind,
199}
200
201impl<'query> QuerySection<'query> {
202    pub fn new(
203        source: &'query str,
204        save: Save,
205        kind: SelectionKind,
206        range: Range<TransitionId>,
207        parent: Option<QuerySectionId>,
208    ) -> Self {
209        Self {
210            source,
211            save,
212            kind,
213            range,
214            parent,
215            next_sibling: None,
216        }
217    }
218
219    pub const fn new_const(
220        source: &'query str,
221        save: Save,
222        kind: SelectionKind,
223        range: Range<TransitionId>,
224        parent: Option<QuerySectionId>,
225        next_sibling: Option<QuerySectionId>,
226    ) -> Self {
227        Self {
228            source,
229            save,
230            kind,
231            range,
232            parent,
233            next_sibling,
234        }
235    }
236}
237
238#[derive(Debug, PartialEq, Clone)]
239pub struct Query<'query> {
240    pub states: Box<[Transition<'query>]>,
241    pub queries: Box<[QuerySection<'query>]>,
242    pub exit_at_section_end: Option<QuerySectionId>,
243}
244
245impl<'query> QuerySpec<'query> for Query<'query> {
246    fn states(&self) -> &[Transition<'query>] {
247        &self.states
248    }
249
250    fn queries(&self) -> &[QuerySection<'query>] {
251        &self.queries
252    }
253
254    fn exit_at_section_end(&self) -> Option<QuerySectionId> {
255        self.exit_at_section_end
256    }
257}
258
259#[derive(Debug, PartialEq, Clone)]
260pub struct StaticQuery<'query, const N_STATES: usize, const N_SECTIONS: usize> {
261    pub states: [Transition<'query>; N_STATES],
262    pub queries: [QuerySection<'query>; N_SECTIONS],
263    pub exit_at_section_end: Option<QuerySectionId>,
264}
265
266impl<'query, const N_STATES: usize, const N_SECTIONS: usize>
267    StaticQuery<'query, N_STATES, N_SECTIONS>
268{
269    pub const fn new(
270        states: [Transition<'query>; N_STATES],
271        queries: [QuerySection<'query>; N_SECTIONS],
272        exit_at_section_end: Option<QuerySectionId>,
273    ) -> Self {
274        Self {
275            states,
276            queries,
277            exit_at_section_end,
278        }
279    }
280}
281
282impl<'query, const N_STATES: usize, const N_SECTIONS: usize> QuerySpec<'query>
283    for StaticQuery<'query, N_STATES, N_SECTIONS>
284{
285    fn states(&self) -> &[Transition<'query>] {
286        &self.states
287    }
288
289    fn queries(&self) -> &[QuerySection<'query>] {
290        &self.queries
291    }
292
293    fn exit_at_section_end(&self) -> Option<QuerySectionId> {
294        self.exit_at_section_end
295    }
296}
297
298impl<'query> Query<'query> {
299    pub fn first(
300        query: &'query str,
301        save: Save,
302    ) -> Result<QueryBuilder<'query>, SelectorParseError> {
303        let states = Transition::generate_transitions_from_string(query)?;
304        let queries = vec![QuerySection::new(
305            query,
306            save,
307            SelectionKind::First,
308            TransitionId(0)..TransitionId(states.len()),
309            None,
310        )];
311
312        Ok(QueryBuilder {
313            states,
314            selection: queries,
315        })
316    }
317
318    pub fn all(query: &'query str, save: Save) -> Result<QueryBuilder<'query>, SelectorParseError> {
319        let states = Transition::generate_transitions_from_string(query)?;
320        let queries = vec![QuerySection::new(
321            query,
322            save,
323            SelectionKind::All,
324            TransitionId(0)..TransitionId(states.len()),
325            None,
326        )];
327
328        Ok(QueryBuilder {
329            states,
330            selection: queries,
331        })
332    }
333}
334
335#[cfg(test)]
336mod tests {
337    use crate::query::compiler::transition::Transition;
338    use crate::query::selector::AttributeSelection;
339    use crate::query::selector::AttributeSelectionKind;
340    use crate::query::selector::AttributeSelections;
341    use crate::query::selector::ClassSelections;
342    use crate::query::selector::Combinator;
343    use crate::query::selector::ElementPredicate;
344    use crate::{Query, QuerySection, QuerySectionId, Save, SelectionKind, TransitionId};
345
346    #[test]
347    fn test_query_builder_one_selection() {
348        let query = Query::all("a", Save::all()).unwrap().build();
349
350        assert_eq!(
351            query.states.iter().as_slice(),
352            [Transition {
353                predicate: ElementPredicate {
354                    name: Some("a"),
355                    id: None,
356                    classes: ClassSelections::from_static(&[]),
357                    attributes: AttributeSelections::from_static(&[])
358                },
359                guard: Combinator::Descendant,
360            }]
361        );
362
363        assert_eq!(
364            query.queries.iter().as_slice(),
365            [QuerySection {
366                source: "a",
367                save: Save::all(),
368                kind: SelectionKind::All,
369                parent: None,
370                range: TransitionId(0)..TransitionId(1),
371                next_sibling: None,
372            }]
373        );
374    }
375
376    #[test]
377    fn test_query_builder_chainned_selection() {
378        let query = Query::first("span", Save::all())
379            .unwrap()
380            .all("a", Save::all())
381            .unwrap()
382            .build();
383
384        assert_eq!(
385            query.states.iter().as_slice(),
386            [
387                Transition {
388                    predicate: ElementPredicate {
389                        name: Some("span"),
390                        id: None,
391                        classes: ClassSelections::from_static(&[]),
392                        attributes: AttributeSelections::from_static(&[])
393                    },
394                    guard: Combinator::Descendant,
395                },
396                Transition {
397                    predicate: ElementPredicate {
398                        name: Some("a"),
399                        id: None,
400                        classes: ClassSelections::from_static(&[]),
401                        attributes: AttributeSelections::from_static(&[])
402                    },
403                    guard: Combinator::Descendant,
404                }
405            ]
406        );
407    }
408
409    #[test]
410    fn test_query_builder_chainned_multi_element_selection() {
411        let query = Query::first("span#top.inner", Save::all())
412            .unwrap()
413            .all("a#link1.foo[href^=\"https\"]", Save::all())
414            .unwrap()
415            .build();
416
417        assert_eq!(query.states.len(), 2);
418        assert_eq!(query.queries.len(), 2);
419        assert_eq!(
420            query.states[1].predicate,
421            ElementPredicate {
422                name: Some("a"),
423                id: Some("link1"),
424                classes: ClassSelections::from_static(&["foo"]),
425                attributes: AttributeSelections::from(vec![AttributeSelection {
426                    name: "href",
427                    value: Some("https"),
428                    kind: AttributeSelectionKind::Prefix,
429                }]),
430            }
431        );
432    }
433
434    #[test]
435    fn test_query_builder_chainned_multi_element_selection_with_branching() {
436        let query = Query::first("div", Save::all())
437            .unwrap()
438            .then(|ctx| {
439                Ok([
440                    ctx.all("a", Save::all())?,
441                    ctx.first("p.note", Save::none())?,
442                ])
443            })
444            .unwrap()
445            .build();
446
447        assert_eq!(query.queries.len(), 3);
448        assert_eq!(query.queries[1].next_sibling, Some(QuerySectionId(2)));
449        assert_eq!(query.queries[2].next_sibling, None);
450    }
451}