itertools.py 16 KB

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