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}