mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 14:12:58 +00:00
89 lines
3.1 KiB
Python
89 lines
3.1 KiB
Python
"""Drive a streaming ASGI response and disconnect the client mid-stream.
|
|
|
|
Starlette's ``TestClient`` buffers a response until the app returns, so it
|
|
can't exercise what happens when a browser tab closes on an endless SSE
|
|
stream. This harness speaks raw ASGI: it hands the app one GET, collects the
|
|
response, and delivers ``http.disconnect`` once enough body has arrived.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Mapping, Optional
|
|
|
|
import anyio
|
|
|
|
|
|
async def stream_then_disconnect(
|
|
app,
|
|
path: str,
|
|
*,
|
|
headers: Optional[Mapping[str, str]] = None,
|
|
query_string: bytes = b"",
|
|
root_path: str = "",
|
|
chunks_before_disconnect: int = 1,
|
|
timeout: float = 5.0,
|
|
) -> tuple[Optional[int], dict[str, str], list[bytes]]:
|
|
"""Run ``app`` for one GET, disconnecting after ``chunks_before_disconnect`` body chunks.
|
|
|
|
A response that finishes on its own (a 401, a stream that ends) returns
|
|
without any disconnect. ``chunks_before_disconnect=0`` disconnects as soon
|
|
as the request is read, before any body. ``root_path`` mounts the app under
|
|
a prefix, which ASGI servers include in ``path``. ``timeout`` fails the
|
|
test instead of hanging it when a stream never ends.
|
|
|
|
Returns:
|
|
tuple: ``(status, headers, body_chunks)`` with lower-cased header names.
|
|
"""
|
|
disconnected = anyio.Event()
|
|
if chunks_before_disconnect <= 0:
|
|
disconnected.set()
|
|
request_delivered = False
|
|
status: Optional[int] = None
|
|
response_headers: dict[str, str] = {}
|
|
chunks: list[bytes] = []
|
|
|
|
scope = {
|
|
"type": "http",
|
|
"asgi": {"version": "3.0", "spec_version": "2.3"},
|
|
"http_version": "1.1",
|
|
"method": "GET",
|
|
"scheme": "http",
|
|
"path": root_path + path,
|
|
"raw_path": (root_path + path).encode("utf-8"),
|
|
"root_path": root_path,
|
|
"query_string": query_string,
|
|
"headers": [
|
|
(name.lower().encode("latin-1"), value.encode("latin-1"))
|
|
for name, value in (headers or {}).items()
|
|
],
|
|
"client": ("testclient", 50000),
|
|
"server": ("testserver", 80),
|
|
}
|
|
|
|
async def receive():
|
|
nonlocal request_delivered
|
|
if not request_delivered:
|
|
request_delivered = True
|
|
return {"type": "http.request", "body": b"", "more_body": False}
|
|
await disconnected.wait()
|
|
return {"type": "http.disconnect"}
|
|
|
|
async def send(message):
|
|
nonlocal status
|
|
if message["type"] == "http.response.start":
|
|
status = message["status"]
|
|
for name, value in message.get("headers", []):
|
|
response_headers[name.decode("latin-1")] = value.decode("latin-1")
|
|
elif message["type"] == "http.response.body":
|
|
body = message.get("body", b"")
|
|
if body:
|
|
chunks.append(body)
|
|
if len(chunks) >= chunks_before_disconnect:
|
|
disconnected.set()
|
|
if not message.get("more_body", False):
|
|
disconnected.set()
|
|
|
|
with anyio.fail_after(timeout):
|
|
await app(scope, receive, send)
|
|
return status, response_headers, chunks
|