"""Own upstream responses across the route-to-ASGI streaming handoff.""" import asyncio import inspect from fastapi.responses import StreamingResponse class OwnedStreamingResponse(StreamingResponse): def __init__(self, content, *, cleanup, **kwargs): super().__init__(content, **kwargs) self.cleanup = cleanup async def __call__(self, scope, receive, send): try: await super().__call__(scope, receive, send) finally: # A disconnect may happen before the generator's first iteration. try: result = self.cleanup() if inspect.isawaitable(result): await result finally: await self.body_iterator.aclose() class UpstreamLease: """Own an upstream response and an optional, already-acquired pool slot.""" def __init__(self, semaphore=None): self.response = None self.semaphore = semaphore self.closed = False self.transferred = False def close(self): if self.closed: return self.closed = True try: if self.response is not None: self.response.release() finally: if self.semaphore is not None: self.semaphore.release() def __enter__(self): return self def __exit__(self, *args): if not self.transferred: self.close() async def chunks(self, error_label): try: async for chunk in self.response.content.iter_chunked(65536): yield chunk except Exception as exc: print(f"{error_label}: {exc}") raise finally: self.close() def streaming_response(self, content, **kwargs): response = OwnedStreamingResponse(content, cleanup=self.close, **kwargs) self.transferred = True return response class SharedRangeFlight: """Stream one bounded response to independently owned subscribers.""" def __init__(self, create, limit, finished): self.create = create self.limit = limit self.finished = finished self.ready = asyncio.get_running_loop().create_future() self.ready.add_done_callback(lambda future: None if future.cancelled() else future.exception()) self.changed = asyncio.Event() self.started = asyncio.Event() self.task = None self.response = None self.headers = {} self.chunks = [] self.size = 0 self.subscribers = 0 self.shared = False self.done = False self.error = None def subscribe(self): self.subscribers += 1 if self.task is None: self.task = asyncio.create_task(self._run()) self.task.add_done_callback(self._settled) return RangeSubscription(self) def _settled(self, task): error = asyncio.CancelledError() if task.cancelled() else task.exception() if error is not None and self.error is None: self.error = error if not self.done: self.done = True if not self.ready.done(): self.ready.set_exception(self.error or RuntimeError("range flight stopped")) self.changed.set() if not self.subscribers: self.finished(self) async def _release_response(self): response, self.response = self.response, None if response is None: return cleanup = getattr(response, "cleanup", None) if cleanup: cleanup() iterator = getattr(response, "body_iterator", None) if iterator is not None: await iterator.aclose() async def _run(self): try: self.response = await self.create() try: length = int(self.response.headers.get("content-length", "")) except (TypeError, ValueError): length = -1 if self.response.status_code != 206 or not 0 <= length <= self.limit: self.ready.set_result(False) return self.shared = True self.headers = dict(self.response.headers) self.ready.set_result(True) await self.started.wait() iterator = getattr(self.response, "body_iterator", None) if iterator is None: self.chunks.append(self.response.body) self.size = len(self.response.body) self.changed.set() else: async for chunk in iterator: if self.size + len(chunk) > length: raise ValueError("shared range body exceeds declared length") self.chunks.append(chunk) self.size += len(chunk) self.changed.set() except BaseException as error: self.error = error if not self.ready.done(): self.ready.set_exception(error) finally: try: if self.shared or self.error or not self.subscribers: await self._release_response() finally: self.done = True self.changed.set() if not self.subscribers: self.finished(self) def take_response(self): response, self.response = self.response, None return response def unsubscribe(self): self.subscribers -= 1 if self.subscribers: return if not self.done: self.task.cancel() self.finished(self) elif self.response is not None: # An unclaimed declined response still owns its upstream lease. self.task = asyncio.create_task(self._close_declined()) else: self.finished(self) async def _close_declined(self): try: await self._release_response() finally: self.finished(self) async def close(self): if self.task and not self.task.done(): self.task.cancel() if self.task: await asyncio.gather(self.task, return_exceptions=True) await self._release_response() self.finished(self) class RangeSubscription: def __init__(self, flight): self.flight = flight self.closed = False def close(self): if not self.closed: self.closed = True self.flight.unsubscribe() async def release(self): flight = self.flight if flight is None: return self.close() if not flight.subscribers and flight.task: await asyncio.gather(flight.task, return_exceptions=True) self.flight = None async def stream(self): cursor = 0 self.flight.started.set() try: while not self.closed: while cursor < len(self.flight.chunks): chunk = self.flight.chunks[cursor] cursor += 1 yield chunk if self.flight.done: if self.flight.error: raise self.flight.error return self.flight.changed.clear() if cursor < len(self.flight.chunks) or self.flight.done: continue await self.flight.changed.wait() finally: await self.release()