concurrency.py 1.7 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859
  1. from __future__ import annotations
  2. import functools
  3. import warnings
  4. from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator
  5. from typing import ParamSpec, TypeVar
  6. import anyio.to_thread
  7. from starlette.exceptions import StarletteDeprecationWarning
  8. P = ParamSpec("P")
  9. T = TypeVar("T")
  10. async def run_until_first_complete(*args: tuple[Callable, dict]) -> None: # type: ignore[type-arg]
  11. warnings.warn(
  12. "run_until_first_complete is deprecated and will be removed in a future version.",
  13. StarletteDeprecationWarning,
  14. )
  15. async with anyio.create_task_group() as task_group:
  16. async def run(func: Callable[[], Coroutine]) -> None: # type: ignore[type-arg]
  17. await func()
  18. task_group.cancel_scope.cancel()
  19. for func, kwargs in args:
  20. task_group.start_soon(run, functools.partial(func, **kwargs))
  21. async def run_in_threadpool(func: Callable[P, T], *args: P.args, **kwargs: P.kwargs) -> T:
  22. func = functools.partial(func, *args, **kwargs)
  23. return await anyio.to_thread.run_sync(func)
  24. class _StopIteration(Exception):
  25. pass
  26. def _next(iterator: Iterator[T]) -> T:
  27. # We can't raise `StopIteration` from within the threadpool iterator
  28. # and catch it outside that context, so we coerce them into a different
  29. # exception type.
  30. try:
  31. return next(iterator)
  32. except StopIteration:
  33. raise _StopIteration
  34. async def iterate_in_threadpool(
  35. iterator: Iterable[T],
  36. ) -> AsyncIterator[T]:
  37. as_iterator = iter(iterator)
  38. while True:
  39. try:
  40. yield await anyio.to_thread.run_sync(_next, as_iterator)
  41. except _StopIteration:
  42. break