Skip to main content

core/iter/adapters/
take.rs

1use crate::cmp;
2use crate::iter::adapters::SourceIter;
3use crate::iter::{FusedIterator, InPlaceIterable, TrustedFused, TrustedLen, TrustedRandomAccess};
4use crate::num::NonZero;
5use crate::ops::{ControlFlow, Try};
6
7/// An iterator that only iterates over the first `n` iterations of `iter`.
8///
9/// This `struct` is created by the [`take`] method on [`Iterator`]. See its
10/// documentation for more.
11///
12/// [`take`]: Iterator::take
13/// [`Iterator`]: trait.Iterator.html
14#[derive(Clone, Debug)]
15#[must_use = "iterators are lazy and do nothing unless consumed"]
16#[stable(feature = "rust1", since = "1.0.0")]
17#[ferrocene::prevalidated]
18pub struct Take<I> {
19    iter: I,
20    n: usize,
21}
22
23impl<I> Take<I> {
24    #[ferrocene::prevalidated]
25    pub(in crate::iter) const fn new(iter: I, n: usize) -> Take<I> {
26        Take { iter, n }
27    }
28}
29
30#[stable(feature = "rust1", since = "1.0.0")]
31impl<I> Iterator for Take<I>
32where
33    I: Iterator,
34{
35    type Item = <I as Iterator>::Item;
36
37    #[inline]
38    #[ferrocene::prevalidated]
39    fn next(&mut self) -> Option<<I as Iterator>::Item> {
40        if self.n != 0 {
41            self.n -= 1;
42            self.iter.next()
43        } else {
44            None
45        }
46    }
47
48    #[inline]
49    #[ferrocene::prevalidated]
50    fn nth(&mut self, n: usize) -> Option<I::Item> {
51        if self.n > n {
52            self.n -= n + 1;
53            self.iter.nth(n)
54        } else {
55            if self.n > 0 {
56                self.iter.nth(self.n - 1);
57                self.n = 0;
58            }
59            None
60        }
61    }
62
63    #[inline]
64    fn count(mut self) -> usize {
65        if self.n == 0 {
66            return 0;
67        }
68        // Advancing consumes the same elements `next` would have yielded,
69        // while benefiting from the inner iterator's `advance_by` fast path.
70        match self.iter.advance_by(self.n) {
71            Ok(()) => self.n,
72            Err(remaining) => self.n - remaining.get(),
73        }
74    }
75
76    #[inline]
77    #[ferrocene::prevalidated]
78    fn size_hint(&self) -> (usize, Option<usize>) {
79        if self.n == 0 {
80            return (0, Some(0));
81        }
82
83        let (lower, upper) = self.iter.size_hint();
84
85        let lower = cmp::min(lower, self.n);
86
87        let upper = match upper {
88            Some(x) if x < self.n => Some(x),
89            _ => Some(self.n),
90        };
91
92        (lower, upper)
93    }
94
95    #[inline]
96    #[ferrocene::prevalidated]
97    fn try_fold<Acc, Fold, R>(&mut self, init: Acc, fold: Fold) -> R
98    where
99        Fold: FnMut(Acc, Self::Item) -> R,
100        R: Try<Output = Acc>,
101    {
102        #[ferrocene::prevalidated]
103        fn check<'a, T, Acc, R: Try<Output = Acc>>(
104            n: &'a mut usize,
105            mut fold: impl FnMut(Acc, T) -> R + 'a,
106        ) -> impl FnMut(Acc, T) -> ControlFlow<R, Acc> + 'a {
107            move |acc, x| {
108                *n -= 1;
109                let r = fold(acc, x);
110                if *n == 0 { ControlFlow::Break(r) } else { ControlFlow::from_try(r) }
111            }
112        }
113
114        if self.n == 0 {
115            try { init }
116        } else {
117            let n = &mut self.n;
118            self.iter.try_fold(init, check(n, fold)).into_try()
119        }
120    }
121
122    #[inline]
123    #[ferrocene::prevalidated]
124    fn fold<B, F>(self, init: B, f: F) -> B
125    where
126        Self: Sized,
127        F: FnMut(B, Self::Item) -> B,
128    {
129        Self::spec_fold(self, init, f)
130    }
131
132    #[inline]
133    #[ferrocene::prevalidated]
134    fn for_each<F: FnMut(Self::Item)>(self, f: F) {
135        Self::spec_for_each(self, f)
136    }
137
138    #[inline]
139    #[rustc_inherit_overflow_checks]
140    #[ferrocene::prevalidated]
141    fn advance_by(&mut self, n: usize) -> Result<(), NonZero<usize>> {
142        let min = self.n.min(n);
143        let rem = match self.iter.advance_by(min) {
144            Ok(()) => 0,
145            Err(rem) => rem.get(),
146        };
147        let advanced = min - rem;
148        self.n -= advanced;
149        NonZero::new(n - advanced).map_or(Ok(()), Err)
150    }
151}
152
153#[unstable(issue = "none", feature = "inplace_iteration")]
154unsafe impl<I> SourceIter for Take<I>
155where
156    I: SourceIter,
157{
158    type Source = I::Source;
159
160    #[inline]
161    unsafe fn as_inner(&mut self) -> &mut I::Source {
162        // SAFETY: unsafe function forwarding to unsafe function with the same requirements
163        unsafe { SourceIter::as_inner(&mut self.iter) }
164    }
165}
166
167#[unstable(issue = "none", feature = "inplace_iteration")]
168unsafe impl<I: InPlaceIterable> InPlaceIterable for Take<I> {
169    const EXPAND_BY: Option<NonZero<usize>> = I::EXPAND_BY;
170    const MERGE_BY: Option<NonZero<usize>> = I::MERGE_BY;
171}
172
173#[stable(feature = "double_ended_take_iterator", since = "1.38.0")]
174impl<I> DoubleEndedIterator for Take<I>
175where
176    I: DoubleEndedIterator + ExactSizeIterator,
177{
178    #[inline]
179    fn next_back(&mut self) -> Option<Self::Item> {
180        if self.n == 0 {
181            None
182        } else {
183            let n = self.n;
184            self.n -= 1;
185            self.iter.nth_back(self.iter.len().saturating_sub(n))
186        }
187    }
188
189    #[inline]
190    fn nth_back(&mut self, n: usize) -> Option<Self::Item> {
191        let len = self.iter.len();
192        if self.n > n {
193            let m = len.saturating_sub(self.n) + n;
194            self.n -= n + 1;
195            self.iter.nth_back(m)
196        } else {
197            if len > 0 {
198                self.iter.nth_back(len - 1);
199            }
200            None
201        }
202    }
203
204    #[inline]
205    fn try_rfold<Acc, Fold, R>(&mut self, init: Acc, fold: Fold) -> R
206    where
207        Self: Sized,
208        Fold: FnMut(Acc, Self::Item) -> R,
209        R: Try<Output = Acc>,
210    {
211        if self.n == 0 {
212            try { init }
213        } else {
214            let len = self.iter.len();
215            if len > self.n && self.iter.nth_back(len - self.n - 1).is_none() {
216                try { init }
217            } else {
218                self.iter.try_rfold(init, fold)
219            }
220        }
221    }
222
223    #[inline]
224    fn rfold<Acc, Fold>(mut self, init: Acc, fold: Fold) -> Acc
225    where
226        Self: Sized,
227        Fold: FnMut(Acc, Self::Item) -> Acc,
228    {
229        if self.n == 0 {
230            init
231        } else {
232            let len = self.iter.len();
233            if len > self.n && self.iter.nth_back(len - self.n - 1).is_none() {
234                init
235            } else {
236                self.iter.rfold(init, fold)
237            }
238        }
239    }
240
241    #[inline]
242    #[rustc_inherit_overflow_checks]
243    fn advance_back_by(&mut self, n: usize) -> Result<(), NonZero<usize>> {
244        // The amount by which the inner iterator needs to be shortened for it to be
245        // at most as long as the take() amount.
246        let trim_inner = self.iter.len().saturating_sub(self.n);
247        // The amount we need to advance inner to fulfill the caller's request.
248        // take(), advance_by() and len() all can be at most usize, so we don't have to worry
249        // about having to advance more than usize::MAX here.
250        let advance_by = trim_inner.saturating_add(n);
251
252        let remainder = match self.iter.advance_back_by(advance_by) {
253            Ok(()) => 0,
254            Err(rem) => rem.get(),
255        };
256        let advanced_by_inner = advance_by - remainder;
257        let advanced_by = advanced_by_inner - trim_inner;
258        self.n -= advanced_by;
259        NonZero::new(n - advanced_by).map_or(Ok(()), Err)
260    }
261}
262
263#[stable(feature = "rust1", since = "1.0.0")]
264impl<I> ExactSizeIterator for Take<I> where I: ExactSizeIterator {}
265
266#[stable(feature = "fused", since = "1.26.0")]
267impl<I> FusedIterator for Take<I> where I: FusedIterator {}
268
269#[unstable(issue = "none", feature = "trusted_fused")]
270unsafe impl<I: TrustedFused> TrustedFused for Take<I> {}
271
272#[unstable(feature = "trusted_len", issue = "37572")]
273unsafe impl<I: TrustedLen> TrustedLen for Take<I> {}
274
275trait SpecTake: Iterator {
276    fn spec_fold<B, F>(self, init: B, f: F) -> B
277    where
278        Self: Sized,
279        F: FnMut(B, Self::Item) -> B;
280
281    fn spec_for_each<F: FnMut(Self::Item)>(self, f: F);
282}
283
284impl<I: Iterator> SpecTake for Take<I> {
285    #[inline]
286    #[ferrocene::prevalidated]
287    default fn spec_fold<B, F>(mut self, init: B, f: F) -> B
288    where
289        Self: Sized,
290        F: FnMut(B, Self::Item) -> B,
291    {
292        use crate::ops::NeverShortCircuit;
293        self.try_fold(init, NeverShortCircuit::wrap_mut_2(f)).0
294    }
295
296    #[inline]
297    #[ferrocene::prevalidated]
298    default fn spec_for_each<F: FnMut(Self::Item)>(mut self, f: F) {
299        // The default implementation would use a unit accumulator, so we can
300        // avoid a stateful closure by folding over the remaining number
301        // of items we wish to return instead.
302        #[ferrocene::prevalidated]
303        fn check<'a, Item>(
304            mut action: impl FnMut(Item) + 'a,
305        ) -> impl FnMut(usize, Item) -> Option<usize> + 'a {
306            move |more, x| {
307                action(x);
308                more.checked_sub(1)
309            }
310        }
311
312        let remaining = self.n;
313        if remaining > 0 {
314            self.iter.try_fold(remaining - 1, check(f));
315        }
316    }
317}
318
319impl<I: Iterator + TrustedRandomAccess> SpecTake for Take<I> {
320    #[inline]
321    fn spec_fold<B, F>(mut self, init: B, mut f: F) -> B
322    where
323        Self: Sized,
324        F: FnMut(B, Self::Item) -> B,
325    {
326        let mut acc = init;
327        let end = self.n.min(self.iter.size());
328        for i in 0..end {
329            // SAFETY: i < end <= self.iter.size() and we discard the iterator at the end
330            let val = unsafe { self.iter.__iterator_get_unchecked(i) };
331            acc = f(acc, val);
332        }
333        acc
334    }
335
336    #[inline]
337    fn spec_for_each<F: FnMut(Self::Item)>(mut self, mut f: F) {
338        let end = self.n.min(self.iter.size());
339        for i in 0..end {
340            // SAFETY: i < end <= self.iter.size() and we discard the iterator at the end
341            let val = unsafe { self.iter.__iterator_get_unchecked(i) };
342            f(val);
343        }
344    }
345}
346
347#[stable(feature = "exact_size_take_repeat", since = "1.82.0")]
348impl<T: Clone> DoubleEndedIterator for Take<crate::iter::Repeat<T>> {
349    #[inline]
350    fn next_back(&mut self) -> Option<Self::Item> {
351        self.next()
352    }
353
354    #[inline]
355    fn nth_back(&mut self, n: usize) -> Option<Self::Item> {
356        self.nth(n)
357    }
358
359    #[inline]
360    fn try_rfold<Acc, Fold, R>(&mut self, init: Acc, fold: Fold) -> R
361    where
362        Self: Sized,
363        Fold: FnMut(Acc, Self::Item) -> R,
364        R: Try<Output = Acc>,
365    {
366        self.try_fold(init, fold)
367    }
368
369    #[inline]
370    fn rfold<Acc, Fold>(self, init: Acc, fold: Fold) -> Acc
371    where
372        Self: Sized,
373        Fold: FnMut(Acc, Self::Item) -> Acc,
374    {
375        self.fold(init, fold)
376    }
377
378    #[inline]
379    #[rustc_inherit_overflow_checks]
380    fn advance_back_by(&mut self, n: usize) -> Result<(), NonZero<usize>> {
381        self.advance_by(n)
382    }
383}
384
385// Note: It may be tempting to impl DoubleEndedIterator for Take<RepeatWith>.
386// One must fight that temptation since such implementation wouldn’t be correct
387// because we have no way to return value of nth invocation of repeater followed
388// by n-1st without remembering all results.
389
390#[stable(feature = "exact_size_take_repeat", since = "1.82.0")]
391impl<T: Clone> ExactSizeIterator for Take<crate::iter::Repeat<T>> {
392    fn len(&self) -> usize {
393        self.n
394    }
395}
396
397#[stable(feature = "exact_size_take_repeat", since = "1.82.0")]
398impl<F: FnMut() -> A, A> ExactSizeIterator for Take<crate::iter::RepeatWith<F>> {
399    fn len(&self) -> usize {
400        self.n
401    }
402}