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

370 statements  

1from __future__ import annotations 

2 

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) 

25 

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 

40 

41from ._core._synchronization import Lock 

42from ._core._tasks import CancelScope 

43from .lowlevel import cancel_shielded_checkpoint, checkpoint, checkpoint_if_cancelled 

44 

45if sys.version_info < (3, 15): 

46 from typing_extensions import sentinel 

47 

48T = TypeVar("T") 

49R = TypeVar("R") 

50_tee_end = sentinel("_tee_end") 

51 

52 

53@dataclass(eq=False) 

54class _IterableAsyncIterator(AsyncIterator[T]): 

55 iterator: Iterator[T] 

56 

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 

64 

65 await cancel_shielded_checkpoint() 

66 return result 

67 

68 

69def _iterate(iterable: Iterable[T] | AsyncIterable[T]) -> AsyncIterator[T]: 

70 if isinstance(iterable, AsyncIterator): 

71 return iterable 

72 

73 if isinstance(iterable, AsyncIterable): 

74 return iterable.__aiter__() 

75 

76 return _IterableAsyncIterator(iter(iterable)) 

77 

78 

79@dataclass(eq=False) 

80class _TeeLink(Generic[T]): 

81 value: object | None = None 

82 next: _TeeLink[T] | None = None 

83 filled: bool = False 

84 

85 

86@dataclass(eq=False) 

87class _TeeState(Generic[T]): 

88 iterator: AsyncIterator[T] 

89 lock: Lock = field(default_factory=Lock) 

90 

91 async def fill(self, link: _TeeLink[T]) -> bool: 

92 if link.filled: 

93 return False 

94 

95 async with self.lock: 

96 if link.filled: 

97 return True 

98 

99 link.value = await anext(self.iterator, _tee_end) 

100 if link.value is not _tee_end: 

101 link.next = _TeeLink() 

102 

103 link.filled = True 

104 return True 

105 

106 

107class _TeeAsyncIterator(AsyncIterator[T]): 

108 _state: _TeeState[T] 

109 _link: _TeeLink[T] 

110 _element_yielded: bool 

111 

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() 

121 

122 self._element_yielded = False 

123 

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() 

129 

130 raise StopAsyncIteration 

131 

132 if not had_yieldpoint: 

133 await checkpoint_if_cancelled() 

134 

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() 

142 

143 return value 

144 

145 

146async def _operator_add(x: T, y: T) -> T: 

147 return operator.add(x, y) 

148 

149 

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() 

167 

168 yield total 

169 

170 async for element in iterator: 

171 total = await function(total, element) 

172 yield total 

173 

174 

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") 

180 

181 iterator = _iterate(iterable) 

182 

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 

194 

195 yield tuple(batch) 

196 return 

197 

198 yield tuple(batch) 

199 

200 

201class Chain: 

202 def __call__( 

203 self, *iterables: Iterable[T] | AsyncIterable[T] 

204 ) -> AsyncGenerator[T, None]: 

205 return self.from_iterable(iterables) 

206 

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) 

216 

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() 

227 

228 if not element_yielded: 

229 await checkpoint() 

230 

231 

232chain: Chain = Chain() 

233 

234 

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 

241 

242 

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 

249 

250 

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 

258 

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() 

266 

267 return 

268 

269 if selector: 

270 element_yielded = True 

271 yield datum 

272 

273 

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 

282 

283 

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 

291 

292 if not saved: 

293 await checkpoint() 

294 return 

295 

296 while True: 

297 for element in saved: 

298 await checkpoint() 

299 yield element 

300 

301 

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 

308 

309 async for element in _iterate(iterable): 

310 if dropping and await predicate(element): 

311 continue 

312 

313 dropping = False 

314 element_yielded = True 

315 yield element 

316 

317 if not element_yielded: 

318 await checkpoint() 

319 

320 

321async def filterfalse( 

322 predicate: Callable[[T], Awaitable[object]], 

323 iterable: Iterable[T] | AsyncIterable[T], 

324) -> AsyncGenerator[T, None]: 

325 element_yielded = False 

326 

327 async for element in _iterate(iterable): 

328 if not await predicate(element): 

329 element_yielded = True 

330 yield element 

331 

332 if not element_yielded: 

333 await checkpoint() 

334 

335 

336@overload 

337def groupby( 

338 iterable: Iterable[T] | AsyncIterable[T], 

339) -> AsyncGenerator[tuple[T, list[T]], None]: ... 

340 

341 

342@overload 

343def groupby( 

344 iterable: Iterable[T] | AsyncIterable[T], 

345 key: Callable[[T], Awaitable[R]], 

346) -> AsyncGenerator[tuple[R, list[T]], None]: ... 

347 

348 

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 

359 

360 group_key = element if key is None else await key(element) 

361 values = [element] 

362 

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) 

372 

373 yield group_key, values 

374 

375 

376@overload 

377def islice( 

378 iterable: Iterable[T] | AsyncIterable[T], 

379 stop: int | None, 

380 /, 

381) -> AsyncGenerator[T, None]: ... 

382 

383 

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]: ... 

392 

393 

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}") 

402 

403 slice_args = slice(*args) 

404 

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." 

412 

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 

418 

419 if index < 0 or index > sys.maxsize: 

420 raise ValueError(message) 

421 

422 return index 

423 

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 ) 

437 

438 if step <= 0: 

439 raise ValueError(step_message) 

440 

441 if stop == 0 or start == stop: 

442 await checkpoint() 

443 return 

444 

445 iterator = _iterate(iterable) 

446 index = 0 

447 element_yielded = False 

448 

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() 

455 

456 return 

457 

458 if index >= start and (index - start) % step == 0: 

459 index += 1 

460 element_yielded = True 

461 yield element 

462 else: 

463 index += 1 

464 

465 if not element_yielded: 

466 await checkpoint() 

467 

468 

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 

478 

479 element_yielded = False 

480 async for element in iterator: 

481 element_yielded = True 

482 pair = (previous, element) 

483 previous = element 

484 yield pair 

485 

486 if not element_yielded: 

487 await checkpoint() 

488 

489 

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") 

501 

502 async for permutation in _iterate(itertools.permutations(pool, r)): 

503 yield permutation 

504 

505 

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") 

512 

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)) 

517 

518 async for value in _iterate(itertools.product(*pools, repeat=repeat)): 

519 yield value 

520 

521 

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 

527 

528 remaining = operator.index(cast(Any, times)) 

529 if remaining <= 0: 

530 await checkpoint() 

531 return 

532 

533 while remaining > 0: 

534 await checkpoint_if_cancelled() 

535 remaining -= 1 

536 await cancel_shielded_checkpoint() 

537 yield element 

538 

539 

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 

548 

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) 

553 

554 if not result_yielded: 

555 await checkpoint() 

556 

557 

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 () 

566 

567 iterator = _TeeAsyncIterator(iterable) 

568 iterators: list[AsyncIterator[T]] = [iterator] 

569 iterators.extend(_TeeAsyncIterator(iterator) for _ in range(n - 1)) 

570 return tuple(iterators) 

571 

572 

573async def takewhile( 

574 predicate: Callable[[T], Awaitable[object]], 

575 iterable: Iterable[T] | AsyncIterable[T], 

576) -> AsyncGenerator[T, None]: 

577 element_yielded = False 

578 

579 async for element in _iterate(iterable): 

580 if not await predicate(element): 

581 if not element_yielded: 

582 await checkpoint() 

583 

584 return 

585 

586 element_yielded = True 

587 yield element 

588 

589 if not element_yielded: 

590 await checkpoint() 

591 

592 

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 

602 

603 active = [True] * num_active 

604 tuple_yielded = False 

605 

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 

612 

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() 

621 

622 return 

623 

624 value = fillvalue 

625 

626 values.append(value) 

627 

628 tuple_yielded = True 

629 yield tuple(values)