Skip to content
File

Blob: src/pyodide/internal/workers-api/src/asgi.py

python322 lines
1import logging
2from asyncio import Event, Future, Queue, create_task, ensure_future, sleep
3from collections.abc import Awaitable
4from contextlib import contextmanager
5from inspect import isawaitable
6from typing import Any
7 
8import js
9from workers import Context, Request
10 
11ASGI = {"spec_version": "2.0", "version": "3.0"}
12logger = logging.getLogger("asgi")
13background_tasks = set()
14 
15 
16def run_in_background(coro: Awaitable[Any]) -> None:
17 fut = ensure_future(coro)
18 background_tasks.add(fut)
19 
20 def _on_done(f):
21 background_tasks.discard(f)
22 exc = f.exception() if not f.cancelled() else None
23 if exc is not None:
24 logger.error("Unhandled exception in background task", exc_info=exc)
25 
26 fut.add_done_callback(_on_done)
27 
28 
29@contextmanager
30def acquire_js_buffer(pybuffer):
31 from pyodide.ffi import create_proxy
32 
33 px = create_proxy(pybuffer)
34 buf = px.getBuffer()
35 px.destroy()
36 try:
37 yield buf.data
38 finally:
39 buf.release()
40 
41 
42def request_to_scope(req, env, ws=False):
43 from js import URL
44 
45 # @app.get("/example")
46 # async def example(request: Request):
47 # request.headers.get("content-type")
48 # - this will error if header is not "bytes" as in ASGI spec.
49 
50 # Support both JS and Python http.client.HTTPMessage headers.
51 req_headers = req.headers.items() if isinstance(req, Request) else req.headers
52 
53 headers = [(k.lower().encode(), v.encode()) for k, v in req_headers]
54 url = URL.new(req.url)
55 assert url.protocol[-1] == ":"
56 scheme = url.protocol[:-1]
57 path = url.pathname
58 assert "?".startswith(url.search[0:1])
59 query_string = url.search[1:].encode()
60 if ws:
61 ty = "websocket"
62 else:
63 ty = "http"
64 return {
65 "asgi": ASGI,
66 "headers": headers,
67 "http_version": "1.1",
68 "method": req.method,
69 "scheme": scheme,
70 "path": path,
71 "query_string": query_string,
72 "type": ty,
73 "env": env,
74 }
75 
76 
77async def start_application(app):
78 shutdown_future = Future()
79 
80 async def shutdown():
81 shutdown_future.set_result(None)
82 await sleep(0)
83 
84 it = iter([{"type": "lifespan.startup"}, Future()])
85 
86 async def receive():
87 res = next(it)
88 if isawaitable(res):
89 await res
90 return res
91 
92 ready = Future()
93 
94 async def send(got):
95 if got["type"] == "lifespan.startup.complete":
96 ready.set_result(None)
97 return
98 if got["type"] == "lifespan.shutdown.complete":
99 return
100 raise RuntimeError(f"Unexpected lifespan event {got['type']}")
101 
102 run_in_background(
103 app(
104 {
105 "asgi": ASGI,
106 "state": {},
107 "type": "lifespan",
108 },
109 receive,
110 send,
111 )
112 )
113 await ready
114 return shutdown
115 
116 
117async def process_request(
118 app: Any, req: "Request | js.Request", env: Any, ctx: Context
119) -> js.Response:
120 from js import Object, Response, TransformStream
121 
122 from pyodide.ffi import create_proxy
123 
124 status = None
125 headers = None
126 result = Future()
127 is_sse = False
128 finished_response = Event()
129 
130 receive_queue = Queue()
131 if req.body:
132 async for data in req.body:
133 await receive_queue.put(
134 {
135 "body": data.to_bytes(),
136 "more_body": True,
137 "type": "http.request",
138 }
139 )
140 await receive_queue.put({"body": b"", "more_body": False, "type": "http.request"})
141 
142 async def receive():
143 message = None
144 if not receive_queue.empty():
145 message = await receive_queue.get()
146 else:
147 await finished_response.wait()
148 message = {"type": "http.disconnect"}
149 return message
150 
151 # Create a transform stream for handling streaming responses
152 transform_stream = TransformStream.new()
153 readable = transform_stream.readable
154 writable = transform_stream.writable
155 writer = writable.getWriter()
156 
157 async def send(got):
158 nonlocal status
159 nonlocal headers
160 nonlocal is_sse
161 
162 if got["type"] == "http.response.start":
163 status = got["status"]
164 # Like above, we need to convert byte-pairs into string explicitly.
165 headers = [(k.decode(), v.decode()) for k, v in got["headers"]]
166 # Check if this is a server-sent events response
167 for k, v in headers:
168 if k.lower() == "content-type" and v.lower().startswith(
169 "text/event-stream"
170 ):
171 is_sse = True
172 break
173 if is_sse:
174 # For SSE, create and return the response immediately after http.response.start
175 resp = Response.new(
176 readable, headers=Object.fromEntries(headers), status=status
177 )
178 result.set_result(resp)
179 
180 elif got["type"] == "http.response.body":
181 body = got["body"]
182 more_body = got.get("more_body", False)
183 
184 # Convert body to JS buffer
185 px = create_proxy(body)
186 buf = px.getBuffer()
187 px.destroy()
188 
189 if is_sse:
190 # For SSE, write chunk to the stream
191 await writer.write(buf.data)
192 # If this is the last chunk, close the writer
193 if not more_body:
194 await writer.close()
195 finished_response.set()
196 else:
197 resp = Response.new(
198 buf.data, headers=Object.fromEntries(headers), status=status
199 )
200 result.set_result(resp)
201 await writer.close()
202 finished_response.set()
203 
204 # Run the application in the background to handle SSE
205 async def run_app():
206 try:
207 await app(request_to_scope(req, env), receive, send)
208 
209 # If we get here and no response has been set yet, the app didn't generate a response
210 if not result.done():
211 raise RuntimeError("The application did not generate a response") # noqa: TRY301
212 except Exception as e:
213 if not result.done():
214 result.set_exception(e)
215 await writer.close() # Close the writer
216 finished_response.set()
217 else:
218 # Response already sent — exception can't be propagated to the
219 # client, so log it to avoid silently swallowing errors.
220 logger.exception("Exception in ASGI application after response started")
221 
222 # Create task to run the application in the background
223 app_task = create_task(run_app())
224 
225 # Wait for the result (the response)
226 response = await result
227 
228 # For non-SSE responses, we need to wait for the application to complete
229 if not is_sse:
230 await app_task
231 else: # noqa: PLR5501
232 if ctx is not None:
233 ctx.waitUntil(create_proxy(app_task))
234 else:
235 raise RuntimeError(
236 "Server-Side-Events require ctx to be passed to asgi.fetch"
237 )
238 return response
239 
240 
241async def process_websocket(app: Any, req: "Request | js.Request") -> js.Response:
242 from js import Response, WebSocketPair
243 
244 client, server = WebSocketPair.new().object_values()
245 server.accept()
246 queue = Queue()
247 
248 def onopen(evt):
249 msg = {"type": "websocket.connect"}
250 queue.put_nowait(msg)
251 
252 # onopen doesn't seem to get called. WS lifecycle events are a bit messed up
253 # here.
254 onopen(1)
255 
256 def onclose(evt):
257 msg = {"type": "websocket.close", "code": evt.code, "reason": evt.reason}
258 queue.put_nowait(msg)
259 
260 def onmessage(evt):
261 msg = {"type": "websocket.receive", "text": evt.data}
262 queue.put_nowait(msg)
263 
264 server.onopen = onopen
265 server.onopen = onclose
266 server.onmessage = onmessage
267 
268 async def ws_send(got):
269 if got["type"] == "websocket.send":
270 b = got.get("bytes", None)
271 s = got.get("text", None)
272 if b:
273 with acquire_js_buffer(b) as jsbytes:
274 # Unlike the `Response` constructor, server.send seems to
275 # eagerly copy the source buffer
276 server.send(jsbytes)
277 if s:
278 server.send(s)
279 
280 else:
281 logger.warning(" == Not implemented %s", got["type"])
282 
283 async def ws_receive():
284 received = await queue.get()
285 return received
286 
287 env = {}
288 run_in_background(app(request_to_scope(req, env, ws=True), ws_receive, ws_send))
289 
290 return Response.new(None, status=101, webSocket=client)
291 
292 
293async def fetch(
294 app: Any, req: "Request | js.Request", env: Any, ctx: Context | None = None
295) -> js.Response:
296 logger.debug("ASGI request: %s %s", req.method, req.url)
297 shutdown = await start_application(app)
298 try:
299 result = await process_request(app, req, env, ctx)
300 except Exception:
301 logger.exception("ASGI request failed")
302 raise
303 await shutdown()
304 return result
305 
306 
307async def websocket(app: Any, req: "Request | js.Request") -> js.Response:
308 return await process_websocket(app, req)
309 
310 
311def __getattr__(name):
312 if name == "env":
313 from fastapi import Depends, Request
314 
315 @Depends
316 async def env(request: Request):
317 return request.scope["env"]
318 
319 return env
320 
321 raise AttributeError(f"module {__name__!r} has no attribute {name!r}")