How can I broadcast asyncio StreamReader to several consumers?

Viewed 439

I'm trying to use aiohttp to make a sort of advanced reverse proxy.

I want to get content of HTTP request and pass it to new HTTP request without pulling it to memory. While there is the only upstream the task is fairly easy: aiohttp server returns request content as StreamReader and aiohttp client can accept StreamReader as request body.

The problem is that I want to send origin request to several upstreams or, for example, simultaneously send content to upstream and write it on disk.

Is there some instruments to broadcast content of StreamReader?

I've tried to make some naive broadcaster but it fails on large objects. What do I do wrong?

class StreamBroadcast:
    async def __do_broadcast(self):
        while True:
            chunk = await self.__source.read(self.__n)
            if not chunk:
                break
            for output in self.__sinks:
                output.feed_data(chunk)
        for output in self.__sinks:
            output.feed_eof()

    def __init__(self, source: StreamReader, sinks_count: int, n: int = -1):
        self.__source = source
        self.__n = n
        self.__sinks = [StreamReader() for i in range(sinks_count)]
        self.__task = asyncio.create_task(self.__do_broadcast())

    @property
    def sinks(self) -> Iterable[StreamReader]:
        return self.__sinks

    @property
    def ready(self) -> Task:
        return self.__task
1 Answers

Well, I've looked through asyncio sources and discovered that I should use Transport to pump data over a stream. Here is my solution.

import asyncio
from asyncio import StreamReader, StreamWriter, ReadTransport, StreamReaderProtocol
from typing import Iterable


class _BroadcastReadTransport(ReadTransport):
    """
    Internal class, is not meant to be instantiated manually
    """

    def __init__(self, source: StreamReader, sinks: Iterable[StreamReader]):
        super().__init__()
        self.__source = source
        self.__sinks = tuple(StreamReaderProtocol(s) for s in sinks)
        for sink in sinks:
            sink.set_transport(self)
        self.__waiting_for_data = len(self.__sinks)

        asyncio.create_task(self.__broadcast_next_chunk(), name='initial-chunk-broadcast')

    def is_reading(self):
        return self.__waiting_for_data == len(self.__sinks)

    def pause_reading(self):
        self.__waiting_for_data -= 1

    async def __broadcast_next_chunk(self):
        data = await self.__source.read()
        if data:
            for sink in self.__sinks:
                sink.data_received(data)
            if self.is_reading():
                asyncio.create_task(self.__broadcast_next_chunk())
        else:
            for sink in self.__sinks:
                sink.eof_received()

    def resume_reading(self):
        self.__waiting_for_data += 1
        if self.__waiting_for_data == len(self.__sinks):
            asyncio.create_task(self.__broadcast_next_chunk(), name='chunk-broadcast')

    @property
    def is_completed(self):
        return self.__source.at_eof()


class StreamBroadcast:
    def __init__(self, source: StreamReader, sinks_count: int):
        self.__source = source
        self.__sinks = tuple(StreamReader() for _ in range(sinks_count))
        self.__transport = _BroadcastReadTransport(self.__source, self.__sinks)

    @property
    def sinks(self) -> Iterable[StreamReader]:
        return self.__sinks

    @property
    def is_completed(self):
        return self.__transport.is_completed

Hope once I'll pack it to pip module.

Related