stapled.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150
  1. from __future__ import annotations
  2. __all__ = (
  3. "MultiListener",
  4. "StapledByteStream",
  5. "StapledObjectStream",
  6. )
  7. from collections.abc import Callable, Mapping, Sequence
  8. from dataclasses import dataclass
  9. from typing import Any, Generic, TypeVar
  10. from ..abc import (
  11. ByteReceiveStream,
  12. ByteSendStream,
  13. ByteStream,
  14. Listener,
  15. ObjectReceiveStream,
  16. ObjectSendStream,
  17. ObjectStream,
  18. TaskGroup,
  19. )
  20. T_Item = TypeVar("T_Item")
  21. T_Stream = TypeVar("T_Stream")
  22. @dataclass(eq=False)
  23. class StapledByteStream(ByteStream):
  24. """
  25. Combines two byte streams into a single, bidirectional byte stream.
  26. Extra attributes will be provided from both streams, with the receive stream
  27. providing the values in case of a conflict.
  28. :param ByteSendStream send_stream: the sending byte stream
  29. :param ByteReceiveStream receive_stream: the receiving byte stream
  30. """
  31. send_stream: ByteSendStream
  32. receive_stream: ByteReceiveStream
  33. async def receive(self, max_bytes: int = 65536) -> bytes:
  34. if max_bytes < 1:
  35. raise ValueError("max_bytes must be a positive integer")
  36. return await self.receive_stream.receive(max_bytes)
  37. async def send(self, item: bytes) -> None:
  38. await self.send_stream.send(item)
  39. async def send_eof(self) -> None:
  40. await self.send_stream.aclose()
  41. async def aclose(self) -> None:
  42. await self.send_stream.aclose()
  43. await self.receive_stream.aclose()
  44. @property
  45. def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
  46. return {
  47. **self.send_stream.extra_attributes,
  48. **self.receive_stream.extra_attributes,
  49. }
  50. @dataclass(eq=False)
  51. class StapledObjectStream(Generic[T_Item], ObjectStream[T_Item]):
  52. """
  53. Combines two object streams into a single, bidirectional object stream.
  54. Extra attributes will be provided from both streams, with the receive stream
  55. providing the values in case of a conflict.
  56. :param ObjectSendStream send_stream: the sending object stream
  57. :param ObjectReceiveStream receive_stream: the receiving object stream
  58. """
  59. send_stream: ObjectSendStream[T_Item]
  60. receive_stream: ObjectReceiveStream[T_Item]
  61. async def receive(self) -> T_Item:
  62. return await self.receive_stream.receive()
  63. async def send(self, item: T_Item) -> None:
  64. await self.send_stream.send(item)
  65. async def send_eof(self) -> None:
  66. await self.send_stream.aclose()
  67. async def aclose(self) -> None:
  68. await self.send_stream.aclose()
  69. await self.receive_stream.aclose()
  70. @property
  71. def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
  72. return {
  73. **self.send_stream.extra_attributes,
  74. **self.receive_stream.extra_attributes,
  75. }
  76. @dataclass(eq=False)
  77. class MultiListener(Generic[T_Stream], Listener[T_Stream]):
  78. """
  79. Combines multiple listeners into one, serving connections from all of them at once.
  80. Any MultiListeners in the given collection of listeners will have their listeners
  81. moved into this one.
  82. Extra attributes are provided from each listener, with each successive listener
  83. overriding any conflicting attributes from the previous one.
  84. :param listeners: listeners to serve
  85. :type listeners: Sequence[Listener[T_Stream]]
  86. """
  87. listeners: Sequence[Listener[T_Stream]]
  88. def __post_init__(self) -> None:
  89. listeners: list[Listener[T_Stream]] = []
  90. for listener in self.listeners:
  91. if isinstance(listener, MultiListener):
  92. listeners.extend(listener.listeners)
  93. del listener.listeners[:] # type: ignore[attr-defined]
  94. else:
  95. listeners.append(listener)
  96. self.listeners = listeners
  97. async def serve(
  98. self, handler: Callable[[T_Stream], Any], task_group: TaskGroup | None = None
  99. ) -> None:
  100. from .. import create_task_group
  101. async with create_task_group() as tg:
  102. for listener in self.listeners:
  103. tg.start_soon(listener.serve, handler, task_group)
  104. async def aclose(self) -> None:
  105. for listener in self.listeners:
  106. await listener.aclose()
  107. @property
  108. def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
  109. attributes: dict = {}
  110. for listener in self.listeners:
  111. attributes.update(listener.extra_attributes)
  112. return attributes