itertools.py 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626
  1. from __future__ import annotations
  2. __all__ = (
  3. "accumulate",
  4. "batched",
  5. "Chain",
  6. "combinations",
  7. "combinations_with_replacement",
  8. "compress",
  9. "count",
  10. "cycle",
  11. "dropwhile",
  12. "filterfalse",
  13. "groupby",
  14. "islice",
  15. "pairwise",
  16. "permutations",
  17. "product",
  18. "repeat",
  19. "starmap",
  20. "tee",
  21. "takewhile",
  22. "zip_longest",
  23. )
  24. import itertools
  25. import operator
  26. import sys
  27. from collections.abc import (
  28. AsyncGenerator,
  29. AsyncIterable,
  30. AsyncIterator,
  31. Awaitable,
  32. Callable,
  33. Iterable,
  34. Iterator,
  35. )
  36. from dataclasses import dataclass, field
  37. from typing import Any, Generic, TypeVar, cast, overload
  38. from ._core._synchronization import Lock
  39. from ._core._tasks import CancelScope
  40. from .lowlevel import cancel_shielded_checkpoint, checkpoint, checkpoint_if_cancelled
  41. T = TypeVar("T")
  42. R = TypeVar("R")
  43. _tee_end = object()
  44. @dataclass(eq=False)
  45. class _IterableAsyncIterator(AsyncIterator[T]):
  46. iterator: Iterator[T]
  47. async def __anext__(self) -> T:
  48. await checkpoint_if_cancelled()
  49. try:
  50. result = next(self.iterator)
  51. except StopIteration:
  52. await cancel_shielded_checkpoint()
  53. raise StopAsyncIteration from None
  54. await cancel_shielded_checkpoint()
  55. return result
  56. def _iterate(iterable: Iterable[T] | AsyncIterable[T]) -> AsyncIterator[T]:
  57. if isinstance(iterable, AsyncIterator):
  58. return iterable
  59. if isinstance(iterable, AsyncIterable):
  60. return iterable.__aiter__()
  61. return _IterableAsyncIterator(iter(iterable))
  62. @dataclass(eq=False)
  63. class _TeeLink(Generic[T]):
  64. value: object | None = None
  65. next: _TeeLink[T] | None = None
  66. filled: bool = False
  67. @dataclass(eq=False)
  68. class _TeeState(Generic[T]):
  69. iterator: AsyncIterator[T]
  70. lock: Lock = field(default_factory=Lock)
  71. async def fill(self, link: _TeeLink[T]) -> bool:
  72. if link.filled:
  73. return False
  74. async with self.lock:
  75. if link.filled:
  76. return True
  77. link.value = await anext(self.iterator, _tee_end)
  78. if link.value is not _tee_end:
  79. link.next = _TeeLink()
  80. link.filled = True
  81. return True
  82. class _TeeAsyncIterator(AsyncIterator[T]):
  83. _state: _TeeState[T]
  84. _link: _TeeLink[T]
  85. _element_yielded: bool
  86. def __init__(
  87. self, iterable: Iterable[T] | AsyncIterable[T] | _TeeAsyncIterator[T]
  88. ) -> None:
  89. if isinstance(iterable, _TeeAsyncIterator):
  90. self._state = iterable._state
  91. self._link = iterable._link
  92. else:
  93. self._state = _TeeState(_iterate(iterable))
  94. self._link = _TeeLink()
  95. self._element_yielded = False
  96. async def __anext__(self) -> T:
  97. had_yieldpoint = await self._state.fill(self._link)
  98. if self._link.value is _tee_end:
  99. if not self._element_yielded:
  100. await checkpoint()
  101. raise StopAsyncIteration
  102. if not had_yieldpoint:
  103. await checkpoint_if_cancelled()
  104. self._element_yielded = True
  105. value = cast(T, self._link.value)
  106. next_link = self._link.next
  107. assert next_link is not None
  108. self._link = next_link
  109. if not had_yieldpoint:
  110. await cancel_shielded_checkpoint()
  111. return value
  112. async def _operator_add(x: T, y: T) -> T:
  113. return operator.add(x, y)
  114. async def accumulate(
  115. iterable: Iterable[T] | AsyncIterable[T],
  116. function: Callable[[T, T], Awaitable[T]] = _operator_add,
  117. *,
  118. initial: T | None = None,
  119. ) -> AsyncGenerator[T, None]:
  120. iterator = _iterate(iterable)
  121. if initial is None:
  122. try:
  123. total = await anext(iterator)
  124. except StopAsyncIteration:
  125. await checkpoint()
  126. return
  127. else:
  128. await checkpoint_if_cancelled()
  129. total = initial
  130. await cancel_shielded_checkpoint()
  131. yield total
  132. async for element in iterator:
  133. total = await function(total, element)
  134. yield total
  135. async def batched(
  136. iterable: Iterable[T] | AsyncIterable[T], n: int, *, strict: bool = False
  137. ) -> AsyncGenerator[tuple[T, ...], None]:
  138. if n < 1:
  139. raise ValueError("n must be at least one")
  140. iterator = _iterate(iterable)
  141. while True:
  142. batch: list[T] = []
  143. for _ in range(n):
  144. try:
  145. batch.append(await anext(iterator))
  146. except StopAsyncIteration:
  147. if not batch:
  148. await checkpoint()
  149. return
  150. if strict:
  151. raise ValueError("batched(): incomplete batch") from None
  152. yield tuple(batch)
  153. return
  154. yield tuple(batch)
  155. class Chain:
  156. def __call__(
  157. self, *iterables: Iterable[T] | AsyncIterable[T]
  158. ) -> AsyncGenerator[T, None]:
  159. return self.from_iterable(iterables)
  160. async def from_iterable(
  161. self,
  162. iterables: (
  163. Iterable[Iterable[T] | AsyncIterable[T]]
  164. | AsyncIterable[Iterable[T] | AsyncIterable[T]]
  165. ),
  166. ) -> AsyncGenerator[T, None]:
  167. element_yielded = False
  168. outer_iter = _iterate(iterables)
  169. try:
  170. async for iterable in outer_iter:
  171. async for element in _iterate(iterable):
  172. element_yielded = True
  173. yield element
  174. finally:
  175. aclose = getattr(outer_iter, "aclose", None)
  176. if aclose is not None:
  177. with CancelScope(shield=True):
  178. await aclose()
  179. if not element_yielded:
  180. await checkpoint()
  181. chain: Chain = Chain()
  182. async def combinations(
  183. iterable: Iterable[T] | AsyncIterable[T], r: int
  184. ) -> AsyncGenerator[tuple[T, ...], None]:
  185. pool: list[T] = [element async for element in _iterate(iterable)]
  186. async for combination in _iterate(itertools.combinations(pool, r)):
  187. yield combination
  188. async def combinations_with_replacement(
  189. iterable: Iterable[T] | AsyncIterable[T], r: int
  190. ) -> AsyncGenerator[tuple[T, ...], None]:
  191. pool: list[T] = [element async for element in _iterate(iterable)]
  192. async for combination in _iterate(itertools.combinations_with_replacement(pool, r)):
  193. yield combination
  194. async def compress(
  195. data: Iterable[T] | AsyncIterable[T],
  196. selectors: Iterable[object] | AsyncIterable[object],
  197. ) -> AsyncGenerator[T, None]:
  198. data_iterator = _iterate(data)
  199. selector_iterator = _iterate(selectors)
  200. element_yielded = False
  201. while True:
  202. try:
  203. datum = await anext(data_iterator)
  204. selector = await anext(selector_iterator)
  205. except StopAsyncIteration:
  206. if not element_yielded:
  207. await checkpoint()
  208. return
  209. if selector:
  210. element_yielded = True
  211. yield datum
  212. async def count(start: int = 0, step: int = 1) -> AsyncGenerator[int, None]:
  213. n = start
  214. while True:
  215. await checkpoint_if_cancelled()
  216. value = n
  217. n += step
  218. await cancel_shielded_checkpoint()
  219. yield value
  220. async def cycle(
  221. iterable: Iterable[T] | AsyncIterable[T],
  222. ) -> AsyncGenerator[T, None]:
  223. saved: list[T] = []
  224. async for element in _iterate(iterable):
  225. saved.append(element)
  226. yield element
  227. if not saved:
  228. await checkpoint()
  229. return
  230. while True:
  231. for element in saved:
  232. await checkpoint()
  233. yield element
  234. async def dropwhile(
  235. predicate: Callable[[T], Awaitable[object]],
  236. iterable: Iterable[T] | AsyncIterable[T],
  237. ) -> AsyncGenerator[T, None]:
  238. element_yielded = False
  239. dropping = True
  240. async for element in _iterate(iterable):
  241. if dropping and await predicate(element):
  242. continue
  243. dropping = False
  244. element_yielded = True
  245. yield element
  246. if not element_yielded:
  247. await checkpoint()
  248. async def filterfalse(
  249. predicate: Callable[[T], Awaitable[object]],
  250. iterable: Iterable[T] | AsyncIterable[T],
  251. ) -> AsyncGenerator[T, None]:
  252. element_yielded = False
  253. async for element in _iterate(iterable):
  254. if not await predicate(element):
  255. element_yielded = True
  256. yield element
  257. if not element_yielded:
  258. await checkpoint()
  259. @overload
  260. def groupby(
  261. iterable: Iterable[T] | AsyncIterable[T],
  262. ) -> AsyncGenerator[tuple[T, list[T]], None]: ...
  263. @overload
  264. def groupby(
  265. iterable: Iterable[T] | AsyncIterable[T],
  266. key: Callable[[T], Awaitable[R]],
  267. ) -> AsyncGenerator[tuple[R, list[T]], None]: ...
  268. async def groupby(
  269. iterable: Iterable[T] | AsyncIterable[T],
  270. key: Callable[[T], Awaitable[object]] | None = None,
  271. ) -> AsyncGenerator[tuple[object, list[T]], None]:
  272. iterator = _iterate(iterable)
  273. try:
  274. element = await anext(iterator)
  275. except StopAsyncIteration:
  276. await checkpoint()
  277. return
  278. group_key = element if key is None else await key(element)
  279. values = [element]
  280. async for element in iterator:
  281. next_key = element if key is None else await key(element)
  282. if next_key != group_key:
  283. completed_group = group_key, values
  284. group_key = next_key
  285. values = [element]
  286. yield completed_group
  287. else:
  288. values.append(element)
  289. yield group_key, values
  290. @overload
  291. def islice(
  292. iterable: Iterable[T] | AsyncIterable[T],
  293. stop: int | None,
  294. /,
  295. ) -> AsyncGenerator[T, None]: ...
  296. @overload
  297. def islice(
  298. iterable: Iterable[T] | AsyncIterable[T],
  299. start: int | None,
  300. stop: int | None,
  301. step: int | None = 1,
  302. /,
  303. ) -> AsyncGenerator[T, None]: ...
  304. async def islice(
  305. iterable: Iterable[T] | AsyncIterable[T],
  306. *args: int | None,
  307. ) -> AsyncGenerator[T, None]:
  308. if not args:
  309. raise TypeError("islice expected at least 2 arguments, got 1")
  310. if len(args) > 3:
  311. raise TypeError(f"islice expected at most 4 arguments, got {len(args) + 1}")
  312. slice_args = slice(*args)
  313. start_message = (
  314. "Indices for islice() must be None or an integer: 0 <= x <= sys.maxsize."
  315. )
  316. stop_message = (
  317. "Stop argument for islice() must be None or an integer: 0 <= x <= sys.maxsize."
  318. )
  319. step_message = "Step for islice() must be a positive integer or None."
  320. def normalize_index(value: object, message: str) -> int:
  321. try:
  322. index = operator.index(cast(Any, value))
  323. except TypeError:
  324. raise ValueError(message) from None
  325. if index < 0 or index > sys.maxsize:
  326. raise ValueError(message)
  327. return index
  328. start = (
  329. 0
  330. if slice_args.start is None
  331. else normalize_index(slice_args.start, start_message)
  332. )
  333. stop = (
  334. None
  335. if slice_args.stop is None
  336. else normalize_index(slice_args.stop, stop_message)
  337. )
  338. step = (
  339. 1 if slice_args.step is None else normalize_index(slice_args.step, step_message)
  340. )
  341. if step <= 0:
  342. raise ValueError(step_message)
  343. if stop == 0 or start == stop:
  344. await checkpoint()
  345. return
  346. iterator = _iterate(iterable)
  347. index = 0
  348. element_yielded = False
  349. while stop is None or index < stop:
  350. try:
  351. element = await anext(iterator)
  352. except StopAsyncIteration:
  353. if not element_yielded:
  354. await checkpoint()
  355. return
  356. if index >= start and (index - start) % step == 0:
  357. index += 1
  358. element_yielded = True
  359. yield element
  360. else:
  361. index += 1
  362. if not element_yielded:
  363. await checkpoint()
  364. async def pairwise(
  365. iterable: Iterable[T] | AsyncIterable[T],
  366. ) -> AsyncGenerator[tuple[T, T], None]:
  367. iterator = _iterate(iterable)
  368. try:
  369. previous = await anext(iterator)
  370. except StopAsyncIteration:
  371. await checkpoint()
  372. return
  373. element_yielded = False
  374. async for element in iterator:
  375. element_yielded = True
  376. pair = (previous, element)
  377. previous = element
  378. yield pair
  379. if not element_yielded:
  380. await checkpoint()
  381. async def permutations(
  382. iterable: Iterable[T] | AsyncIterable[T], r: int | None = None
  383. ) -> AsyncGenerator[tuple[T, ...], None]:
  384. pool: list[T] = [element async for element in _iterate(iterable)]
  385. n = len(pool)
  386. if r is None:
  387. r = n
  388. elif not isinstance(r, int):
  389. raise TypeError("Expected int as r")
  390. elif r < 0:
  391. raise ValueError("r must be non-negative")
  392. async for permutation in _iterate(itertools.permutations(pool, r)):
  393. yield permutation
  394. async def product(
  395. *iterables: Iterable[T] | AsyncIterable[T], repeat: int = 1
  396. ) -> AsyncGenerator[tuple[T, ...], None]:
  397. repeat = operator.index(repeat)
  398. if repeat < 0:
  399. raise ValueError("repeat argument cannot be negative")
  400. pools: list[tuple[T, ...]] = []
  401. for iterable in iterables:
  402. pool: list[T] = [element async for element in _iterate(iterable)]
  403. pools.append(tuple(pool))
  404. async for value in _iterate(itertools.product(*pools, repeat=repeat)):
  405. yield value
  406. async def repeat(element: T, times: int | None = None) -> AsyncGenerator[T, None]:
  407. if times is None:
  408. while True:
  409. await checkpoint()
  410. yield element
  411. remaining = operator.index(cast(Any, times))
  412. if remaining <= 0:
  413. await checkpoint()
  414. return
  415. while remaining > 0:
  416. await checkpoint_if_cancelled()
  417. remaining -= 1
  418. await cancel_shielded_checkpoint()
  419. yield element
  420. async def starmap(
  421. function: Callable[..., Awaitable[R]],
  422. iterable: (
  423. Iterable[Iterable[object] | AsyncIterable[object]]
  424. | AsyncIterable[Iterable[object] | AsyncIterable[object]]
  425. ),
  426. ) -> AsyncGenerator[R, None]:
  427. result_yielded = False
  428. async for args_iterable in _iterate(iterable):
  429. args = [element async for element in _iterate(args_iterable)]
  430. result_yielded = True
  431. yield await function(*args)
  432. if not result_yielded:
  433. await checkpoint()
  434. def tee(
  435. iterable: Iterable[T] | AsyncIterable[T], n: int = 2
  436. ) -> tuple[AsyncIterator[T], ...]:
  437. n = operator.index(cast(Any, n))
  438. if n < 0:
  439. raise ValueError("n must be >= 0")
  440. if n == 0:
  441. return ()
  442. iterator = _TeeAsyncIterator(iterable)
  443. iterators: list[AsyncIterator[T]] = [iterator]
  444. iterators.extend(_TeeAsyncIterator(iterator) for _ in range(n - 1))
  445. return tuple(iterators)
  446. async def takewhile(
  447. predicate: Callable[[T], Awaitable[object]],
  448. iterable: Iterable[T] | AsyncIterable[T],
  449. ) -> AsyncGenerator[T, None]:
  450. element_yielded = False
  451. async for element in _iterate(iterable):
  452. if not await predicate(element):
  453. if not element_yielded:
  454. await checkpoint()
  455. return
  456. element_yielded = True
  457. yield element
  458. if not element_yielded:
  459. await checkpoint()
  460. async def zip_longest(
  461. *iterables: Iterable[object] | AsyncIterable[object],
  462. fillvalue: object = None,
  463. ) -> AsyncGenerator[tuple[object, ...], None]:
  464. iterators = [_iterate(iterable) for iterable in iterables]
  465. num_active = len(iterators)
  466. if not num_active:
  467. await checkpoint()
  468. return
  469. active = [True] * num_active
  470. tuple_yielded = False
  471. while True:
  472. values: list[object] = []
  473. for index, iterator in enumerate(iterators):
  474. if not active[index]:
  475. values.append(fillvalue)
  476. continue
  477. try:
  478. value = await anext(iterator)
  479. except StopAsyncIteration:
  480. active[index] = False
  481. num_active -= 1
  482. if not num_active:
  483. if not tuple_yielded:
  484. await checkpoint()
  485. return
  486. value = fillvalue
  487. values.append(value)
  488. tuple_yielded = True
  489. yield tuple(values)