Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/anyio/itertools.py: 19%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1from __future__ import annotations
3__all__ = (
4 "Chain",
5 "accumulate",
6 "batched",
7 "combinations",
8 "combinations_with_replacement",
9 "compress",
10 "count",
11 "cycle",
12 "dropwhile",
13 "filterfalse",
14 "groupby",
15 "islice",
16 "pairwise",
17 "permutations",
18 "product",
19 "repeat",
20 "starmap",
21 "takewhile",
22 "tee",
23 "zip_longest",
24)
26import itertools
27import operator
28import sys
29from collections.abc import (
30 AsyncGenerator,
31 AsyncIterable,
32 AsyncIterator,
33 Awaitable,
34 Callable,
35 Iterable,
36 Iterator,
37)
38from dataclasses import dataclass, field
39from typing import Any, Generic, TypeVar, cast, overload
41from ._core._synchronization import Lock
42from ._core._tasks import CancelScope
43from .lowlevel import cancel_shielded_checkpoint, checkpoint, checkpoint_if_cancelled
45if sys.version_info < (3, 15):
46 from typing_extensions import sentinel
48T = TypeVar("T")
49R = TypeVar("R")
50_tee_end = sentinel("_tee_end")
53@dataclass(eq=False)
54class _IterableAsyncIterator(AsyncIterator[T]):
55 iterator: Iterator[T]
57 async def __anext__(self) -> T:
58 await checkpoint_if_cancelled()
59 try:
60 result = next(self.iterator)
61 except StopIteration:
62 await cancel_shielded_checkpoint()
63 raise StopAsyncIteration from None
65 await cancel_shielded_checkpoint()
66 return result
69def _iterate(iterable: Iterable[T] | AsyncIterable[T]) -> AsyncIterator[T]:
70 if isinstance(iterable, AsyncIterator):
71 return iterable
73 if isinstance(iterable, AsyncIterable):
74 return iterable.__aiter__()
76 return _IterableAsyncIterator(iter(iterable))
79@dataclass(eq=False)
80class _TeeLink(Generic[T]):
81 value: object | None = None
82 next: _TeeLink[T] | None = None
83 filled: bool = False
86@dataclass(eq=False)
87class _TeeState(Generic[T]):
88 iterator: AsyncIterator[T]
89 lock: Lock = field(default_factory=Lock)
91 async def fill(self, link: _TeeLink[T]) -> bool:
92 if link.filled:
93 return False
95 async with self.lock:
96 if link.filled:
97 return True
99 link.value = await anext(self.iterator, _tee_end)
100 if link.value is not _tee_end:
101 link.next = _TeeLink()
103 link.filled = True
104 return True
107class _TeeAsyncIterator(AsyncIterator[T]):
108 _state: _TeeState[T]
109 _link: _TeeLink[T]
110 _element_yielded: bool
112 def __init__(
113 self, iterable: Iterable[T] | AsyncIterable[T] | _TeeAsyncIterator[T]
114 ) -> None:
115 if isinstance(iterable, _TeeAsyncIterator):
116 self._state = iterable._state
117 self._link = iterable._link
118 else:
119 self._state = _TeeState(_iterate(iterable))
120 self._link = _TeeLink()
122 self._element_yielded = False
124 async def __anext__(self) -> T:
125 had_yieldpoint = await self._state.fill(self._link)
126 if self._link.value is _tee_end:
127 if not self._element_yielded:
128 await checkpoint()
130 raise StopAsyncIteration
132 if not had_yieldpoint:
133 await checkpoint_if_cancelled()
135 self._element_yielded = True
136 value = cast(T, self._link.value)
137 next_link = self._link.next
138 assert next_link is not None
139 self._link = next_link
140 if not had_yieldpoint:
141 await cancel_shielded_checkpoint()
143 return value
146async def _operator_add(x: T, y: T) -> T:
147 return operator.add(x, y)
150async def accumulate(
151 iterable: Iterable[T] | AsyncIterable[T],
152 function: Callable[[T, T], Awaitable[T]] = _operator_add,
153 *,
154 initial: T | None = None,
155) -> AsyncGenerator[T, None]:
156 iterator = _iterate(iterable)
157 if initial is None:
158 try:
159 total = await anext(iterator)
160 except StopAsyncIteration:
161 await checkpoint()
162 return
163 else:
164 await checkpoint_if_cancelled()
165 total = initial
166 await cancel_shielded_checkpoint()
168 yield total
170 async for element in iterator:
171 total = await function(total, element)
172 yield total
175async def batched(
176 iterable: Iterable[T] | AsyncIterable[T], n: int, *, strict: bool = False
177) -> AsyncGenerator[tuple[T, ...], None]:
178 if n < 1:
179 raise ValueError("n must be at least one")
181 iterator = _iterate(iterable)
183 while True:
184 batch: list[T] = []
185 for _ in range(n):
186 try:
187 batch.append(await anext(iterator))
188 except StopAsyncIteration:
189 if not batch:
190 await checkpoint()
191 return
192 if strict:
193 raise ValueError("batched(): incomplete batch") from None
195 yield tuple(batch)
196 return
198 yield tuple(batch)
201class Chain:
202 def __call__(
203 self, *iterables: Iterable[T] | AsyncIterable[T]
204 ) -> AsyncGenerator[T, None]:
205 return self.from_iterable(iterables)
207 async def from_iterable(
208 self,
209 iterables: (
210 Iterable[Iterable[T] | AsyncIterable[T]]
211 | AsyncIterable[Iterable[T] | AsyncIterable[T]]
212 ),
213 ) -> AsyncGenerator[T, None]:
214 element_yielded = False
215 outer_iter = _iterate(iterables)
217 try:
218 async for iterable in outer_iter:
219 async for element in _iterate(iterable):
220 element_yielded = True
221 yield element
222 finally:
223 aclose = getattr(outer_iter, "aclose", None)
224 if aclose is not None:
225 with CancelScope(shield=True):
226 await aclose()
228 if not element_yielded:
229 await checkpoint()
232chain: Chain = Chain()
235async def combinations(
236 iterable: Iterable[T] | AsyncIterable[T], r: int
237) -> AsyncGenerator[tuple[T, ...], None]:
238 pool: list[T] = [element async for element in _iterate(iterable)]
239 async for combination in _iterate(itertools.combinations(pool, r)):
240 yield combination
243async def combinations_with_replacement(
244 iterable: Iterable[T] | AsyncIterable[T], r: int
245) -> AsyncGenerator[tuple[T, ...], None]:
246 pool: list[T] = [element async for element in _iterate(iterable)]
247 async for combination in _iterate(itertools.combinations_with_replacement(pool, r)):
248 yield combination
251async def compress(
252 data: Iterable[T] | AsyncIterable[T],
253 selectors: Iterable[object] | AsyncIterable[object],
254) -> AsyncGenerator[T, None]:
255 data_iterator = _iterate(data)
256 selector_iterator = _iterate(selectors)
257 element_yielded = False
259 while True:
260 try:
261 datum = await anext(data_iterator)
262 selector = await anext(selector_iterator)
263 except StopAsyncIteration:
264 if not element_yielded:
265 await checkpoint()
267 return
269 if selector:
270 element_yielded = True
271 yield datum
274async def count(start: int = 0, step: int = 1) -> AsyncGenerator[int, None]:
275 n = start
276 while True:
277 await checkpoint_if_cancelled()
278 value = n
279 n += step
280 await cancel_shielded_checkpoint()
281 yield value
284async def cycle(
285 iterable: Iterable[T] | AsyncIterable[T],
286) -> AsyncGenerator[T, None]:
287 saved: list[T] = []
288 async for element in _iterate(iterable):
289 saved.append(element)
290 yield element
292 if not saved:
293 await checkpoint()
294 return
296 while True:
297 for element in saved:
298 await checkpoint()
299 yield element
302async def dropwhile(
303 predicate: Callable[[T], Awaitable[object]],
304 iterable: Iterable[T] | AsyncIterable[T],
305) -> AsyncGenerator[T, None]:
306 element_yielded = False
307 dropping = True
309 async for element in _iterate(iterable):
310 if dropping and await predicate(element):
311 continue
313 dropping = False
314 element_yielded = True
315 yield element
317 if not element_yielded:
318 await checkpoint()
321async def filterfalse(
322 predicate: Callable[[T], Awaitable[object]],
323 iterable: Iterable[T] | AsyncIterable[T],
324) -> AsyncGenerator[T, None]:
325 element_yielded = False
327 async for element in _iterate(iterable):
328 if not await predicate(element):
329 element_yielded = True
330 yield element
332 if not element_yielded:
333 await checkpoint()
336@overload
337def groupby(
338 iterable: Iterable[T] | AsyncIterable[T],
339) -> AsyncGenerator[tuple[T, list[T]], None]: ...
342@overload
343def groupby(
344 iterable: Iterable[T] | AsyncIterable[T],
345 key: Callable[[T], Awaitable[R]],
346) -> AsyncGenerator[tuple[R, list[T]], None]: ...
349async def groupby(
350 iterable: Iterable[T] | AsyncIterable[T],
351 key: Callable[[T], Awaitable[object]] | None = None,
352) -> AsyncGenerator[tuple[object, list[T]], None]:
353 iterator = _iterate(iterable)
354 try:
355 element = await anext(iterator)
356 except StopAsyncIteration:
357 await checkpoint()
358 return
360 group_key = element if key is None else await key(element)
361 values = [element]
363 async for element in iterator:
364 next_key = element if key is None else await key(element)
365 if next_key != group_key:
366 completed_group = group_key, values
367 group_key = next_key
368 values = [element]
369 yield completed_group
370 else:
371 values.append(element)
373 yield group_key, values
376@overload
377def islice(
378 iterable: Iterable[T] | AsyncIterable[T],
379 stop: int | None,
380 /,
381) -> AsyncGenerator[T, None]: ...
384@overload
385def islice(
386 iterable: Iterable[T] | AsyncIterable[T],
387 start: int | None,
388 stop: int | None,
389 step: int | None = 1,
390 /,
391) -> AsyncGenerator[T, None]: ...
394async def islice(
395 iterable: Iterable[T] | AsyncIterable[T],
396 *args: int | None,
397) -> AsyncGenerator[T, None]:
398 if not args:
399 raise TypeError("islice expected at least 2 arguments, got 1")
400 if len(args) > 3:
401 raise TypeError(f"islice expected at most 4 arguments, got {len(args) + 1}")
403 slice_args = slice(*args)
405 start_message = (
406 "Indices for islice() must be None or an integer: 0 <= x <= sys.maxsize."
407 )
408 stop_message = (
409 "Stop argument for islice() must be None or an integer: 0 <= x <= sys.maxsize."
410 )
411 step_message = "Step for islice() must be a positive integer or None."
413 def normalize_index(value: object, message: str) -> int:
414 try:
415 index = operator.index(cast(Any, value))
416 except TypeError:
417 raise ValueError(message) from None
419 if index < 0 or index > sys.maxsize:
420 raise ValueError(message)
422 return index
424 start = (
425 0
426 if slice_args.start is None
427 else normalize_index(slice_args.start, start_message)
428 )
429 stop = (
430 None
431 if slice_args.stop is None
432 else normalize_index(slice_args.stop, stop_message)
433 )
434 step = (
435 1 if slice_args.step is None else normalize_index(slice_args.step, step_message)
436 )
438 if step <= 0:
439 raise ValueError(step_message)
441 if stop == 0 or start == stop:
442 await checkpoint()
443 return
445 iterator = _iterate(iterable)
446 index = 0
447 element_yielded = False
449 while stop is None or index < stop:
450 try:
451 element = await anext(iterator)
452 except StopAsyncIteration:
453 if not element_yielded:
454 await checkpoint()
456 return
458 if index >= start and (index - start) % step == 0:
459 index += 1
460 element_yielded = True
461 yield element
462 else:
463 index += 1
465 if not element_yielded:
466 await checkpoint()
469async def pairwise(
470 iterable: Iterable[T] | AsyncIterable[T],
471) -> AsyncGenerator[tuple[T, T], None]:
472 iterator = _iterate(iterable)
473 try:
474 previous = await anext(iterator)
475 except StopAsyncIteration:
476 await checkpoint()
477 return
479 element_yielded = False
480 async for element in iterator:
481 element_yielded = True
482 pair = (previous, element)
483 previous = element
484 yield pair
486 if not element_yielded:
487 await checkpoint()
490async def permutations(
491 iterable: Iterable[T] | AsyncIterable[T], r: int | None = None
492) -> AsyncGenerator[tuple[T, ...], None]:
493 pool: list[T] = [element async for element in _iterate(iterable)]
494 n = len(pool)
495 if r is None:
496 r = n
497 elif not isinstance(r, int):
498 raise TypeError("Expected int as r")
499 elif r < 0:
500 raise ValueError("r must be non-negative")
502 async for permutation in _iterate(itertools.permutations(pool, r)):
503 yield permutation
506async def product(
507 *iterables: Iterable[T] | AsyncIterable[T], repeat: int = 1
508) -> AsyncGenerator[tuple[T, ...], None]:
509 repeat = operator.index(repeat)
510 if repeat < 0:
511 raise ValueError("repeat argument cannot be negative")
513 pools: list[tuple[T, ...]] = []
514 for iterable in iterables:
515 pool: list[T] = [element async for element in _iterate(iterable)]
516 pools.append(tuple(pool))
518 async for value in _iterate(itertools.product(*pools, repeat=repeat)):
519 yield value
522async def repeat(element: T, times: int | None = None) -> AsyncGenerator[T, None]:
523 if times is None:
524 while True:
525 await checkpoint()
526 yield element
528 remaining = operator.index(cast(Any, times))
529 if remaining <= 0:
530 await checkpoint()
531 return
533 while remaining > 0:
534 await checkpoint_if_cancelled()
535 remaining -= 1
536 await cancel_shielded_checkpoint()
537 yield element
540async def starmap(
541 function: Callable[..., Awaitable[R]],
542 iterable: (
543 Iterable[Iterable[object] | AsyncIterable[object]]
544 | AsyncIterable[Iterable[object] | AsyncIterable[object]]
545 ),
546) -> AsyncGenerator[R, None]:
547 result_yielded = False
549 async for args_iterable in _iterate(iterable):
550 args = [element async for element in _iterate(args_iterable)]
551 result_yielded = True
552 yield await function(*args)
554 if not result_yielded:
555 await checkpoint()
558def tee(
559 iterable: Iterable[T] | AsyncIterable[T], n: int = 2
560) -> tuple[AsyncIterator[T], ...]:
561 n = operator.index(cast(Any, n))
562 if n < 0:
563 raise ValueError("n must be >= 0")
564 if n == 0:
565 return ()
567 iterator = _TeeAsyncIterator(iterable)
568 iterators: list[AsyncIterator[T]] = [iterator]
569 iterators.extend(_TeeAsyncIterator(iterator) for _ in range(n - 1))
570 return tuple(iterators)
573async def takewhile(
574 predicate: Callable[[T], Awaitable[object]],
575 iterable: Iterable[T] | AsyncIterable[T],
576) -> AsyncGenerator[T, None]:
577 element_yielded = False
579 async for element in _iterate(iterable):
580 if not await predicate(element):
581 if not element_yielded:
582 await checkpoint()
584 return
586 element_yielded = True
587 yield element
589 if not element_yielded:
590 await checkpoint()
593async def zip_longest(
594 *iterables: Iterable[object] | AsyncIterable[object],
595 fillvalue: object = None,
596) -> AsyncGenerator[tuple[object, ...], None]:
597 iterators = [_iterate(iterable) for iterable in iterables]
598 num_active = len(iterators)
599 if not num_active:
600 await checkpoint()
601 return
603 active = [True] * num_active
604 tuple_yielded = False
606 while True:
607 values: list[object] = []
608 for index, iterator in enumerate(iterators):
609 if not active[index]:
610 values.append(fillvalue)
611 continue
613 try:
614 value = await anext(iterator)
615 except StopAsyncIteration:
616 active[index] = False
617 num_active -= 1
618 if not num_active:
619 if not tuple_yielded:
620 await checkpoint()
622 return
624 value = fillvalue
626 values.append(value)
628 tuple_yielded = True
629 yield tuple(values)