From 7b5e375869beb847e2675d1ca7d73a04285e7f97 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 22 Apr 2026 14:55:15 +0000 Subject: [PATCH 01/32] feat(queue): multiprocessing support --- gradio/queueing.py | 86 +++++++++++++++++++++++++++++++--------------- 1 file changed, 58 insertions(+), 28 deletions(-) diff --git a/gradio/queueing.py b/gradio/queueing.py index cd914cbced3..7f7c12187ba 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -3,14 +3,17 @@ import asyncio import copy import inspect +import multiprocessing import os import platform import random import time import traceback import uuid +import warnings from asyncio import Queue as AsyncQueue from collections import defaultdict +from threading import Thread from typing import TYPE_CHECKING, Any, Literal, cast import fastapi @@ -136,6 +139,8 @@ def __init__( self.active_jobs: list[None | list[Event]] = [] self.delete_lock = safe_get_lock() self.server_app = None + self.server_pid = os.getpid() + self.rpc_queue: multiprocessing.Queue[tuple[str, EventMessage]] | None = None self.process_time_per_fn: defaultdict[BlockFunction, ProcessTime] = defaultdict( ProcessTime ) @@ -212,6 +217,19 @@ def start(self): run_coro_in_background(self.start_progress_updates) if not self.live_updates: run_coro_in_background(self.notify_clients) + if os.getenv("GRADIO_QUEUE_MULTIPROCESSING_ENABLED", "").lower() in ("1", "true"): + Thread(target=self.start_rpc, daemon=True).start() + + def start_rpc(self): + try: + ctx = multiprocessing.get_context('fork') + except ValueError: + warnings.warn("GRADIO_QUEUE_MULTIPROCESSING_ENABLED but fork context not available") + return + self.rpc_queue = ctx.Queue() + while True: + event_id, message = self.rpc_queue.get() + self._send_message_rpc(event_id, message) def create_event_queue_for_fn(self, block_fn: BlockFunction): concurrency_id = block_fn.concurrency_id @@ -571,6 +589,27 @@ async def start_progress_updates(self) -> None: await asyncio.sleep(self.progress_update_sleep_when_free) + def _send_message_rpc( + self, + event_id: str, + message: EventMessage, + ): + if os.getpid() != self.server_pid: + if self.rpc_queue is None: + warnings.warn("Sending queue event from child process without GRADIO_QUEUE_MULTIPROCESSING_ENABLED") + else: + self.rpc_queue.put_nowait((event_id, message)) + return + events = [evt for job in self.active_jobs if job is not None for evt in job] + for event in events: + if event._id == event_id: + match message: + case ProgressMessage(): + event.progress = message + event.progress_pending = True + case _: + self.send_message(event, message) + def set_progress( self, event_id: str, @@ -578,23 +617,17 @@ def set_progress( ): if iterables is None: return - for job in self.active_jobs: - if job is None: - continue - for evt in job: - if evt._id == event_id: - progress_data: list[ProgressUnit] = [] - for iterable in iterables: - progress_unit = ProgressUnit( - index=iterable.index, - length=iterable.length, - unit=iterable.unit, - progress=iterable.progress, - desc=iterable.desc, - ) - progress_data.append(progress_unit) - evt.progress = ProgressMessage(progress_data=progress_data) - evt.progress_pending = True + progress_data: list[ProgressUnit] = [] + for iterable in iterables: + progress_unit = ProgressUnit( + index=iterable.index, + length=iterable.length, + unit=iterable.unit, + progress=iterable.progress, + desc=iterable.desc, + ) + progress_data.append(progress_unit) + self._send_message_rpc(event_id, ProgressMessage(progress_data=progress_data)) def log_message( self, @@ -605,17 +638,14 @@ def log_message( duration: float | None = 10, visible: bool = True, ): - events = [evt for job in self.active_jobs if job is not None for evt in job] - for event in events: - if event._id == event_id: - log_message = LogMessage( - log=log, - level=level, - duration=duration, - visible=visible, - title=title, - ) - self.send_message(event, log_message) + log_message = LogMessage( + log=log, + level=level, + duration=duration, + visible=visible, + title=title, + ) + self._send_message_rpc(event_id, log_message) async def clean_events( self, *, session_hash: str | None = None, event_id: str | None = None From 3389ab8248351e7a41b409889ac82dfeb1127ce4 Mon Sep 17 00:00:00 2001 From: gradio-pr-bot Date: Wed, 22 Apr 2026 14:57:36 +0000 Subject: [PATCH 02/32] add changeset --- .changeset/beige-jobs-switch.md | 5 +++++ 1 file changed, 5 insertions(+) create mode 100644 .changeset/beige-jobs-switch.md diff --git a/.changeset/beige-jobs-switch.md b/.changeset/beige-jobs-switch.md new file mode 100644 index 00000000000..d4d9798af0f --- /dev/null +++ b/.changeset/beige-jobs-switch.md @@ -0,0 +1,5 @@ +--- +"gradio": minor +--- + +feat:feat(queue): multiprocessing support From 43e5fd45f65c1245f377b832fb5c217471ed7139 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 22 Apr 2026 15:01:00 +0000 Subject: [PATCH 03/32] ruff format --- gradio/queueing.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/gradio/queueing.py b/gradio/queueing.py index 7f7c12187ba..a0963be1bad 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -217,14 +217,19 @@ def start(self): run_coro_in_background(self.start_progress_updates) if not self.live_updates: run_coro_in_background(self.notify_clients) - if os.getenv("GRADIO_QUEUE_MULTIPROCESSING_ENABLED", "").lower() in ("1", "true"): + if os.getenv("GRADIO_QUEUE_MULTIPROCESSING_ENABLED", "").lower() in ( + "1", + "true", + ): Thread(target=self.start_rpc, daemon=True).start() def start_rpc(self): try: - ctx = multiprocessing.get_context('fork') + ctx = multiprocessing.get_context("fork") except ValueError: - warnings.warn("GRADIO_QUEUE_MULTIPROCESSING_ENABLED but fork context not available") + warnings.warn( + "GRADIO_QUEUE_MULTIPROCESSING_ENABLED but fork context not available" + ) return self.rpc_queue = ctx.Queue() while True: @@ -596,7 +601,9 @@ def _send_message_rpc( ): if os.getpid() != self.server_pid: if self.rpc_queue is None: - warnings.warn("Sending queue event from child process without GRADIO_QUEUE_MULTIPROCESSING_ENABLED") + warnings.warn( + "Sending queue event from child process without GRADIO_QUEUE_MULTIPROCESSING_ENABLED" + ) else: self.rpc_queue.put_nowait((event_id, message)) return From e39fbe29823e9ee330f59990231d50bdb69f60d4 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 11:04:04 +0000 Subject: [PATCH 04/32] spaces in test requirements --- scripts/create_test_requirements.sh | 1 + test/requirements.in | 1 + test/requirements.txt | 8 +++++++- 3 files changed, 9 insertions(+), 1 deletion(-) diff --git a/scripts/create_test_requirements.sh b/scripts/create_test_requirements.sh index 188e67c2f7c..80623b2fc3b 100755 --- a/scripts/create_test_requirements.sh +++ b/scripts/create_test_requirements.sh @@ -10,5 +10,6 @@ To match the CI environment, this script should be run from a Unix-like system i uv pip compile \ --exclude-newer "${UV_EXCLUDE_NEWER:-7 days}" \ + --exclude-newer-package "spaces=0 days" \ test/requirements.in \ -o test/requirements.txt diff --git a/test/requirements.in b/test/requirements.in index 6efc98da0a7..96355bf7322 100644 --- a/test/requirements.in +++ b/test/requirements.in @@ -28,3 +28,4 @@ tqdm transformers vega_datasets diffusers +spaces>=0.50.dev0 diff --git a/test/requirements.txt b/test/requirements.txt index 2ebf74bdd47..48135ef5d9d 100644 --- a/test/requirements.txt +++ b/test/requirements.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile --python-version 3.10 --exclude-newer 7 days test/requirements.in -o test/requirements.txt +# uv pip compile --python-version 3.10 --exclude-newer 7 days --exclude-newer-package spaces=0 days test/requirements.in -o test/requirements.txt aiofiles==23.2.1 # via gradio altair==5.5.0 @@ -114,6 +114,7 @@ httpx==0.28.1 # openai # respx # safehttpx + # spaces huggingface-hub==1.4.1 # via # -r test/requirements.in @@ -215,6 +216,7 @@ packaging==24.2 # pytest # pytest-rerunfailures # scikit-image + # spaces # transformers pandas==2.2.3 # via @@ -249,6 +251,7 @@ pydantic==2.10.6 # fastapi # gradio # openai + # spaces pydantic-core==2.27.2 # via pydantic pydub==0.25.1 @@ -335,6 +338,8 @@ sniffio==1.3.1 # openai sortedcontainers==2.4.0 # via hypothesis +spaces==0.50.dev0 + # via -r test/requirements.in stack-data==0.6.3 # via ipython starlette==0.45.3 @@ -391,6 +396,7 @@ typing-extensions==4.12.2 # pydantic-core # referencing # rich + # spaces # torch # typer # uvicorn From f5b17a30d55a28afccc062ae31a90635ab651ad4 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 11:04:32 +0000 Subject: [PATCH 05/32] No more auto-wrap --- gradio/block_function.py | 9 --------- 1 file changed, 9 deletions(-) diff --git a/gradio/block_function.py b/gradio/block_function.py index d6a4178c900..94f725f7138 100644 --- a/gradio/block_function.py +++ b/gradio/block_function.py @@ -114,15 +114,6 @@ def __init__( self.component_prop_inputs = component_prop_inputs or [] self.key = key - self.spaces_auto_wrap() - - def spaces_auto_wrap(self): - if spaces is None: - return - if utils.get_space() is None: - return - self.fn = spaces.gradio_auto_wrap(self.fn) - def __str__(self): return str( { From 09bb3bd8c3e911bbcb55b1b03167fa15104baf11 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 11:32:20 +0000 Subject: [PATCH 06/32] ruff format --- gradio/context.py | 12 ++++++++++++ gradio/routes.py | 28 ++++++++++++++++++++++++++-- pyproject.toml | 2 +- 3 files changed, 39 insertions(+), 3 deletions(-) diff --git a/gradio/context.py b/gradio/context.py index c992b20638e..48a335842f9 100644 --- a/gradio/context.py +++ b/gradio/context.py @@ -33,6 +33,18 @@ class LocalContext: ) +class MultiprocessWorkerContextualizer: + def __init__(self): + self.event_id = LocalContext.event_id.get(None) + self.in_event_listener = LocalContext.in_event_listener.get(False) + self.progress = LocalContext.progress.get(None) + + def __call__(self): + LocalContext.event_id.set(self.event_id) + LocalContext.in_event_listener.set(self.in_event_listener) + LocalContext.progress.set(self.progress) + + def get_render_context() -> BlockContext | None: if LocalContext.renderable.get(None): return LocalContext.render_block.get(None) diff --git a/gradio/routes.py b/gradio/routes.py index 23cc66e27d3..305b41164b6 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -71,7 +71,7 @@ utils, ) from gradio.brotli_middleware import BrotliMiddleware -from gradio.context import Context +from gradio.context import Context, MultiprocessWorkerContextualizer from gradio.data_classes import ( CancelBody, ComponentServerBlobBody, @@ -88,7 +88,7 @@ VibeEditBody, ) from gradio.exceptions import Error, InvalidPathError -from gradio.helpers import special_args +from gradio.helpers import log_message, special_args from gradio.i18n import I18n from gradio.node_server import ( start_node_server, @@ -466,6 +466,30 @@ def create_app( excluded_handlers=[mcp_subpath], ) + if utils.is_zero_gpu_space(): + try: + from spaces.zero import ZeroGPUMiddleware + except ImportError: + pass + else: + app.add_middleware( + ZeroGPUMiddleware, + exception_mapper=lambda err, exc: ( + setattr(exc, "print_exception", False) or exc + if isinstance(exc, Error) + else Error( + title=err["detail"]["title"], + message=err["detail"]["message"], + ) + ), + log_emitter=lambda log: log_message( + title=log["title"], + message=log["message"], + level=log["level"], + ), + worker_contextualizer=MultiprocessWorkerContextualizer, + ) + if ssr_mode: @app.middleware("http") diff --git a/pyproject.toml b/pyproject.toml index 2ac636ac3d4..99585cc96bf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -108,7 +108,7 @@ exclude = [ [tool.uv] exclude-newer = "7 days" -exclude-newer-package = {hf-gradio = false, gradio-client = false} +exclude-newer-package = {hf-gradio = false, gradio-client = false, spaces = false} [tool.ruff] exclude = ["gradio/node/*.py", ".venv/*", "gradio/_frontend_code/*.py", "gradio/_vendor/*"] From c4d924485bab85c56d15c3096e8bf971eea05ba4 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 11:51:42 +0000 Subject: [PATCH 07/32] Not needed anymore --- gradio/block_function.py | 5 ----- gradio/blocks.py | 5 ----- 2 files changed, 10 deletions(-) diff --git a/gradio/block_function.py b/gradio/block_function.py index 94f725f7138..5719255f5c8 100644 --- a/gradio/block_function.py +++ b/gradio/block_function.py @@ -6,11 +6,6 @@ from . import utils -try: - import spaces # type: ignore -except Exception: - spaces = None - if TYPE_CHECKING: # Only import for type checking (is False at runtime). from gradio.components.base import Component diff --git a/gradio/blocks.py b/gradio/blocks.py index d638a1ead88..4747bff8b86 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -95,11 +95,6 @@ get_upload_folder, ) -try: - import spaces # type: ignore -except Exception: - spaces = None - if TYPE_CHECKING: # Only import for type checking (is False at runtime). from gradio.components.base import Component From d31eabb3e3b3d1d14f66f2959e5c74e4cda697bc Mon Sep 17 00:00:00 2001 From: gradio-pr-bot Date: Mon, 4 May 2026 11:54:55 +0000 Subject: [PATCH 08/32] add changeset --- .changeset/beige-jobs-switch.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.changeset/beige-jobs-switch.md b/.changeset/beige-jobs-switch.md index d4d9798af0f..62f69e9868b 100644 --- a/.changeset/beige-jobs-switch.md +++ b/.changeset/beige-jobs-switch.md @@ -2,4 +2,4 @@ "gradio": minor --- -feat:feat(queue): multiprocessing support +feat:ZeroGPU native support From f67f2d30b6ea1563830501b7b1ec873fc47b5a6b Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 12:04:06 +0000 Subject: [PATCH 09/32] ruff check --- gradio/block_function.py | 1 - gradio/blocks.py | 1 - 2 files changed, 2 deletions(-) diff --git a/gradio/block_function.py b/gradio/block_function.py index 5719255f5c8..abdf370e8fa 100644 --- a/gradio/block_function.py +++ b/gradio/block_function.py @@ -6,7 +6,6 @@ from . import utils - if TYPE_CHECKING: # Only import for type checking (is False at runtime). from gradio.components.base import Component from gradio.renderable import Renderable diff --git a/gradio/blocks.py b/gradio/blocks.py index 1bf03c59092..9354ba41c5b 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -95,7 +95,6 @@ get_upload_folder, ) - if TYPE_CHECKING: # Only import for type checking (is False at runtime). from gradio.components.base import Component from gradio.mcp import GradioMCPServer From d6bf21d0c8a9c35c9c7ebeab0214c8013309991a Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 12:24:39 +0000 Subject: [PATCH 10/32] ty is not ready for serious typing --- gradio/routes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gradio/routes.py b/gradio/routes.py index 2f380492420..c51ee9ec3fc 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -474,7 +474,7 @@ def create_app( pass else: app.add_middleware( - ZeroGPUMiddleware, + ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] exception_mapper=lambda err, exc: ( setattr(exc, "print_exception", False) or exc if isinstance(exc, Error) From 9dc97e330eaa9f03b07a3268cad99953a7639120 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 14:43:30 +0000 Subject: [PATCH 11/32] ruff format --- gradio/routes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gradio/routes.py b/gradio/routes.py index c51ee9ec3fc..c009791ff11 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -474,7 +474,7 @@ def create_app( pass else: app.add_middleware( - ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] + ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] exception_mapper=lambda err, exc: ( setattr(exc, "print_exception", False) or exc if isinstance(exc, Error) From 097681089fbd40d51ee58075d0ce9a68f474a1f4 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 16:45:30 +0000 Subject: [PATCH 12/32] Shorter --- gradio/queueing.py | 9 ++------- 1 file changed, 2 insertions(+), 7 deletions(-) diff --git a/gradio/queueing.py b/gradio/queueing.py index 885c21208ac..40f1a2ed455 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -217,19 +217,14 @@ def start(self): run_coro_in_background(self.start_progress_updates) if not self.live_updates: run_coro_in_background(self.notify_clients) - if os.getenv("GRADIO_QUEUE_MULTIPROCESSING_ENABLED", "").lower() in ( - "1", - "true", - ): + if os.getenv("GRADIO_QUEUE_MULTIPROCESSING_ENABLED") == "true": Thread(target=self.start_rpc, daemon=True).start() def start_rpc(self): try: ctx = multiprocessing.get_context("fork") except ValueError: - warnings.warn( - "GRADIO_QUEUE_MULTIPROCESSING_ENABLED but fork context not available" - ) + warnings.warn("GRADIO_QUEUE_MULTIPROCESSING_ENABLED but fork not available") return self.rpc_queue = ctx.Queue() while True: From 4809ec0ff7bd6bfc1cdb3c700981ac7b153dd3cb Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 16:47:27 +0000 Subject: [PATCH 13/32] SimpleQueue --- gradio/queueing.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/gradio/queueing.py b/gradio/queueing.py index 40f1a2ed455..fe2b1da3865 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -140,7 +140,7 @@ def __init__( self.delete_lock = safe_get_lock() self.server_app = None self.server_pid = os.getpid() - self.rpc_queue: multiprocessing.Queue[tuple[str, EventMessage]] | None = None + self.rpc_queue: multiprocessing.SimpleQueue[tuple[str, EventMessage]] | None = None self.process_time_per_fn: defaultdict[BlockFunction, ProcessTime] = defaultdict( ProcessTime ) @@ -226,7 +226,7 @@ def start_rpc(self): except ValueError: warnings.warn("GRADIO_QUEUE_MULTIPROCESSING_ENABLED but fork not available") return - self.rpc_queue = ctx.Queue() + self.rpc_queue = ctx.SimpleQueue() while True: event_id, message = self.rpc_queue.get() self._send_message_rpc(event_id, message) @@ -611,7 +611,7 @@ def _send_message_rpc( "Sending queue event from child process without GRADIO_QUEUE_MULTIPROCESSING_ENABLED" ) else: - self.rpc_queue.put_nowait((event_id, message)) + self.rpc_queue.put((event_id, message)) return events = [evt for job in self.active_jobs if job is not None for evt in job] for event in events: From ef07b23e28869c42a71bbb736db66122ec08bcee Mon Sep 17 00:00:00 2001 From: cbensimon Date: Mon, 4 May 2026 17:01:47 +0000 Subject: [PATCH 14/32] handle RPC thread exceptions --- gradio/queueing.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/gradio/queueing.py b/gradio/queueing.py index fe2b1da3865..1503c8b7d39 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -13,6 +13,7 @@ import warnings from asyncio import Queue as AsyncQueue from collections import defaultdict +from multiprocessing import SimpleQueue from threading import Thread from typing import TYPE_CHECKING, Any, Literal, cast @@ -140,7 +141,7 @@ def __init__( self.delete_lock = safe_get_lock() self.server_app = None self.server_pid = os.getpid() - self.rpc_queue: multiprocessing.SimpleQueue[tuple[str, EventMessage]] | None = None + self.rpc_queue: SimpleQueue[tuple[str, EventMessage]] | None = None self.process_time_per_fn: defaultdict[BlockFunction, ProcessTime] = defaultdict( ProcessTime ) @@ -229,7 +230,11 @@ def start_rpc(self): self.rpc_queue = ctx.SimpleQueue() while True: event_id, message = self.rpc_queue.get() - self._send_message_rpc(event_id, message) + try: + self._send_message_rpc(event_id, message) + except Exception: + print("Exception while calling _send_message_rpc from Queue RPC thread") + traceback.print_exc() def create_event_queue_for_fn(self, block_fn: BlockFunction): concurrency_id = block_fn.concurrency_id From 0cd13c1f0665d1b7fbbfa37a83e46bdc71addaf0 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Tue, 5 May 2026 11:20:43 +0000 Subject: [PATCH 15/32] Already in pyproject.toml --- scripts/create_test_requirements.sh | 1 - test/requirements.txt | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/scripts/create_test_requirements.sh b/scripts/create_test_requirements.sh index 80623b2fc3b..188e67c2f7c 100755 --- a/scripts/create_test_requirements.sh +++ b/scripts/create_test_requirements.sh @@ -10,6 +10,5 @@ To match the CI environment, this script should be run from a Unix-like system i uv pip compile \ --exclude-newer "${UV_EXCLUDE_NEWER:-7 days}" \ - --exclude-newer-package "spaces=0 days" \ test/requirements.in \ -o test/requirements.txt diff --git a/test/requirements.txt b/test/requirements.txt index 48135ef5d9d..08bed143f0c 100644 --- a/test/requirements.txt +++ b/test/requirements.txt @@ -1,5 +1,5 @@ # This file was autogenerated by uv via the following command: -# uv pip compile --python-version 3.10 --exclude-newer 7 days --exclude-newer-package spaces=0 days test/requirements.in -o test/requirements.txt +# uv pip compile --exclude-newer 7 days test/requirements.in -o test/requirements.txt aiofiles==23.2.1 # via gradio altair==5.5.0 From e6c84b8d5e1e6aeaa6508235d37426cf8c4860d1 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Tue, 5 May 2026 11:47:34 +0000 Subject: [PATCH 16/32] Add docstring for MultiprocessWorkerContextualizer --- gradio/context.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/gradio/context.py b/gradio/context.py index 48a335842f9..6ec6c73fbf9 100644 --- a/gradio/context.py +++ b/gradio/context.py @@ -34,6 +34,25 @@ class LocalContext: class MultiprocessWorkerContextualizer: + """ + Refreshes LocalContext for persistent multiprocessing workers processing consecutive requests. + + Example usage: + ``` + pool = ProcessPoolExecutor() + + def handler(value): + contextualize = MultiprocessWorkerContextualizer() + return e.submit(process_wrapper, contextualize, value).result() + + def process_wrapper(contextualize, value): + contextualize() + return process(value) + + demo = gr.Interface(handler, gr.Text(), gr.Text()) + ``` + """ + def __init__(self): self.event_id = LocalContext.event_id.get(None) self.in_event_listener = LocalContext.in_event_listener.get(False) From 6964ddbeea9cf0f91447ee0fbee31afe6f915630 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Tue, 5 May 2026 16:24:58 +0000 Subject: [PATCH 17/32] Mount Middleware in launch instead of create_app + START_TIMEOUT --- gradio/blocks.py | 4 ++++ gradio/http_server.py | 3 ++- gradio/routes.py | 28 ++-------------------------- gradio/utils.py | 28 +++++++++++++++++++++++++++- 4 files changed, 35 insertions(+), 28 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index 9354ba41c5b..287ef8c73f8 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -2747,6 +2747,10 @@ def reverse(text): mcp_server=mcp_server, debug=debug, ) + + if utils.is_zero_gpu_space(): + utils.setup_zerogpu_middleware(self.app) + if self.mcp_error and not quiet: print(self.mcp_error) diff --git a/gradio/http_server.py b/gradio/http_server.py index 84924f13738..9226c72d94d 100644 --- a/gradio/http_server.py +++ b/gradio/http_server.py @@ -31,6 +31,7 @@ INITIAL_PORT_VALUE = int(os.getenv("GRADIO_SERVER_PORT", "7860")) TRY_NUM_PORTS = int(os.getenv("GRADIO_NUM_PORTS", "100")) LOCALHOST_NAME = os.getenv("GRADIO_SERVER_NAME", "127.0.0.1") +START_TIMEOUT = int(os.getenv("GRADIO_START_TIMEOUT", "5")) GRADIO_HOT_RELOAD = os.getenv("GRADIO_HOT_RELOAD", "false").lower() @@ -69,7 +70,7 @@ def run_in_thread(self): start = time.time() while not self.started: time.sleep(1e-3) - if time.time() - start > 5: + if time.time() - start > START_TIMEOUT: raise ServerFailedToStartError( "Server failed to start. Please check that the port is available." ) diff --git a/gradio/routes.py b/gradio/routes.py index c009791ff11..0b7307d3fc4 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -71,7 +71,7 @@ utils, ) from gradio.brotli_middleware import BrotliMiddleware -from gradio.context import Context, MultiprocessWorkerContextualizer +from gradio.context import Context from gradio.data_classes import ( CancelBody, ComponentServerBlobBody, @@ -88,7 +88,7 @@ VibeEditBody, ) from gradio.exceptions import Error, InvalidPathError -from gradio.helpers import log_message, special_args +from gradio.helpers import special_args from gradio.i18n import I18n from gradio.node_server import ( start_node_server, @@ -467,30 +467,6 @@ def create_app( excluded_handlers=[mcp_subpath], ) - if utils.is_zero_gpu_space(): - try: - from spaces.zero import ZeroGPUMiddleware - except ImportError: - pass - else: - app.add_middleware( - ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] - exception_mapper=lambda err, exc: ( - setattr(exc, "print_exception", False) or exc - if isinstance(exc, Error) - else Error( - title=err["detail"]["title"], - message=err["detail"]["message"], - ) - ), - log_emitter=lambda log: log_message( - title=log["title"], - message=log["message"], - level=log["level"], - ), - worker_contextualizer=MultiprocessWorkerContextualizer, - ) - if ssr_mode: @app.middleware("http") diff --git a/gradio/utils.py b/gradio/utils.py index 3bc78c57eb8..d5c679d4053 100644 --- a/gradio/utils.py +++ b/gradio/utils.py @@ -62,7 +62,7 @@ import gradio from gradio import themes -from gradio.context import get_blocks_context +from gradio.context import MultiprocessWorkerContextualizer, get_blocks_context from gradio.data_classes import ( BlocksConfigDict, DeveloperPath, @@ -70,6 +70,7 @@ UserProvidedPath, ) from gradio.exceptions import Error, InvalidPathError +from gradio.helpers import log_message from gradio.themes import Default as DefaultTheme from gradio.themes import ThemeClass as Theme @@ -570,6 +571,31 @@ def is_zero_gpu_space() -> bool: return os.getenv("SPACES_ZERO_GPU") == "true" +def setup_zerogpu_middleware(app: App): + try: + from spaces.zero import ZeroGPUMiddleware + except ImportError: + return + + app.add_middleware( + ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] + exception_mapper=lambda err, exc: ( + setattr(exc, "print_exception", False) or exc + if isinstance(exc, Error) + else Error( + title=err["detail"]["title"], + message=err["detail"]["message"], + ) + ), + log_emitter=lambda log: log_message( + title=log["title"], + message=log["message"], + level=log["level"], + ), + worker_contextualizer=MultiprocessWorkerContextualizer, + ) + + def get_theme(theme: Theme | str | None) -> Theme: if theme is None: theme = DefaultTheme() From 3a4dd7b5b1b7e5ba55d47c316d63b11b12904b05 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Tue, 5 May 2026 16:41:45 +0000 Subject: [PATCH 18/32] Move setup inside Blocks method --- gradio/blocks.py | 30 ++++++++++++++++++++++++++++-- gradio/utils.py | 28 +--------------------------- 2 files changed, 29 insertions(+), 29 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index 287ef8c73f8..98e4d41c0cf 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -45,6 +45,7 @@ from gradio.context import ( Context, LocalContext, + MultiprocessWorkerContextualizer, get_blocks_context, get_render_context, set_render_context, @@ -67,11 +68,12 @@ from gradio.exceptions import ( ChecksumMismatchError, DuplicateBlockError, + Error, InvalidApiNameError, InvalidComponentError, ShareCertificateWriteError, ) -from gradio.helpers import create_tracker, skip, special_args +from gradio.helpers import create_tracker, log_message, skip, special_args from gradio.i18n import I18n, I18nData from gradio.node_server import start_node_server from gradio.route_utils import API_PREFIX, MediaStream, slugify @@ -2749,7 +2751,7 @@ def reverse(text): ) if utils.is_zero_gpu_space(): - utils.setup_zerogpu_middleware(self.app) + self._setup_zerogpu_middleware() if self.mcp_error and not quiet: print(self.mcp_error) @@ -3329,3 +3331,27 @@ def route( self.pages.append((path, name, show_in_navbar)) self.current_page = path return self + + def _setup_zerogpu_middleware(self): + try: + from spaces.zero import ZeroGPUMiddleware + except ImportError: + return + + self.app.add_middleware( + ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] + exception_mapper=lambda err, exc: ( + setattr(exc, "print_exception", False) or exc + if isinstance(exc, Error) + else Error( + title=err["detail"]["title"], + message=err["detail"]["message"], + ) + ), + log_emitter=lambda log: log_message( + title=log["title"], + message=log["message"], + level=log["level"], + ), + worker_contextualizer=MultiprocessWorkerContextualizer, + ) diff --git a/gradio/utils.py b/gradio/utils.py index d5c679d4053..3bc78c57eb8 100644 --- a/gradio/utils.py +++ b/gradio/utils.py @@ -62,7 +62,7 @@ import gradio from gradio import themes -from gradio.context import MultiprocessWorkerContextualizer, get_blocks_context +from gradio.context import get_blocks_context from gradio.data_classes import ( BlocksConfigDict, DeveloperPath, @@ -70,7 +70,6 @@ UserProvidedPath, ) from gradio.exceptions import Error, InvalidPathError -from gradio.helpers import log_message from gradio.themes import Default as DefaultTheme from gradio.themes import ThemeClass as Theme @@ -571,31 +570,6 @@ def is_zero_gpu_space() -> bool: return os.getenv("SPACES_ZERO_GPU") == "true" -def setup_zerogpu_middleware(app: App): - try: - from spaces.zero import ZeroGPUMiddleware - except ImportError: - return - - app.add_middleware( - ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] - exception_mapper=lambda err, exc: ( - setattr(exc, "print_exception", False) or exc - if isinstance(exc, Error) - else Error( - title=err["detail"]["title"], - message=err["detail"]["message"], - ) - ), - log_emitter=lambda log: log_message( - title=log["title"], - message=log["message"], - level=log["level"], - ), - worker_contextualizer=MultiprocessWorkerContextualizer, - ) - - def get_theme(theme: Theme | str | None) -> Theme: if theme is None: theme = DefaultTheme() From df84193abab51fbb60d9e563177a9abd4b197f7e Mon Sep 17 00:00:00 2001 From: cbensimon Date: Tue, 5 May 2026 17:42:11 +0000 Subject: [PATCH 19/32] Also setup for mount_gradio_app --- gradio/blocks.py | 8 +++++--- gradio/routes.py | 1 + 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index 98e4d41c0cf..a292048e991 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -2750,8 +2750,7 @@ def reverse(text): debug=debug, ) - if utils.is_zero_gpu_space(): - self._setup_zerogpu_middleware() + self.maybe_setup_zerogpu_middleware() if self.mcp_error and not quiet: print(self.mcp_error) @@ -3332,7 +3331,10 @@ def route( self.current_page = path return self - def _setup_zerogpu_middleware(self): + def maybe_setup_zerogpu_middleware(self): + if not utils.is_zero_gpu_space(): + return + try: from spaces.zero import ZeroGPUMiddleware except ImportError: diff --git a/gradio/routes.py b/gradio/routes.py index 0b7307d3fc4..316d0a91120 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -2818,6 +2818,7 @@ def read_main(): ssr_mode=blocks.ssr_mode, mcp_server=mcp_server, ) + blocks.maybe_setup_zerogpu_middleware() old_lifespan = app.router.lifespan_context @contextlib.asynccontextmanager From 6bb3b3b5e236ffa8e9db8cc15028491d05dbae77 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Tue, 5 May 2026 17:58:06 +0000 Subject: [PATCH 20/32] Move in route_utils.py ... --- gradio/blocks.py | 40 ++++++++-------------------------------- gradio/route_utils.py | 31 +++++++++++++++++++++++++++++++ gradio/routes.py | 3 ++- 3 files changed, 41 insertions(+), 33 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index a292048e991..34878e861e8 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -45,7 +45,6 @@ from gradio.context import ( Context, LocalContext, - MultiprocessWorkerContextualizer, get_blocks_context, get_render_context, set_render_context, @@ -68,15 +67,19 @@ from gradio.exceptions import ( ChecksumMismatchError, DuplicateBlockError, - Error, InvalidApiNameError, InvalidComponentError, ShareCertificateWriteError, ) -from gradio.helpers import create_tracker, log_message, skip, special_args +from gradio.helpers import create_tracker, skip, special_args from gradio.i18n import I18n, I18nData from gradio.node_server import start_node_server -from gradio.route_utils import API_PREFIX, MediaStream, slugify +from gradio.route_utils import ( + API_PREFIX, + MediaStream, + maybe_setup_zerogpu_middleware, + slugify, +) from gradio.routes import INTERNAL_ROUTES, VERSION, App, Request from gradio.state_holder import SessionState, StateHolder from gradio.themes import ThemeClass as Theme @@ -2750,7 +2753,7 @@ def reverse(text): debug=debug, ) - self.maybe_setup_zerogpu_middleware() + maybe_setup_zerogpu_middleware(self.app) if self.mcp_error and not quiet: print(self.mcp_error) @@ -3330,30 +3333,3 @@ def route( self.pages.append((path, name, show_in_navbar)) self.current_page = path return self - - def maybe_setup_zerogpu_middleware(self): - if not utils.is_zero_gpu_space(): - return - - try: - from spaces.zero import ZeroGPUMiddleware - except ImportError: - return - - self.app.add_middleware( - ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] - exception_mapper=lambda err, exc: ( - setattr(exc, "print_exception", False) or exc - if isinstance(exc, Error) - else Error( - title=err["detail"]["title"], - message=err["detail"]["message"], - ) - ), - log_emitter=lambda log: log_message( - title=log["title"], - message=log["message"], - level=log["level"], - ), - worker_contextualizer=MultiprocessWorkerContextualizer, - ) diff --git a/gradio/route_utils.py b/gradio/route_utils.py index 50c040b6ee9..8012561bda4 100644 --- a/gradio/route_utils.py +++ b/gradio/route_utils.py @@ -46,6 +46,7 @@ from starlette.types import ASGIApp, Message, Receive, Scope, Send from gradio import processing_utils, utils +from gradio.context import MultiprocessWorkerContextualizer from gradio.data_classes import ( BlocksConfigDict, MediaStreamChunk, @@ -1160,3 +1161,33 @@ async def iter_body(head: bytes, queue: asyncio.Queue[bytes | None]): yield head while (chunk := await queue.get()) is not None: yield chunk + + +def maybe_setup_zerogpu_middleware(app: App | fastapi.FastAPI): + if not utils.is_zero_gpu_space(): + return + + try: + from spaces.zero import ZeroGPUMiddleware + except ImportError: + return + + from gradio.helpers import log_message + + app.add_middleware( + ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] + exception_mapper=lambda err, exc: ( + setattr(exc, "print_exception", False) or exc + if isinstance(exc, Error) + else Error( + title=err["detail"]["title"], + message=err["detail"]["message"], + ) + ), + log_emitter=lambda log: log_message( + title=log["title"], + message=log["message"], + level=log["level"], + ), + worker_contextualizer=MultiprocessWorkerContextualizer, + ) diff --git a/gradio/routes.py b/gradio/routes.py index 316d0a91120..3aba77cdf75 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -107,6 +107,7 @@ Request, compare_passwords_securely, create_lifespan_handler, + maybe_setup_zerogpu_middleware, move_uploaded_files_to_cache, ) from gradio.screen_recording_utils import process_video_with_ffmpeg @@ -2818,7 +2819,6 @@ def read_main(): ssr_mode=blocks.ssr_mode, mcp_server=mcp_server, ) - blocks.maybe_setup_zerogpu_middleware() old_lifespan = app.router.lifespan_context @contextlib.asynccontextmanager @@ -2834,6 +2834,7 @@ async def new_lifespan(app: FastAPI): app.router.lifespan_context = new_lifespan # type: ignore app.mount(path, gradio_app) + maybe_setup_zerogpu_middleware(app) return app From 594033877277b67f4ecbe9ffb30f653da391eb40 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 09:51:33 +0000 Subject: [PATCH 21/32] Fix utils.is_zero_gpu_space --- gradio/utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gradio/utils.py b/gradio/utils.py index 3bc78c57eb8..45d788246da 100644 --- a/gradio/utils.py +++ b/gradio/utils.py @@ -567,7 +567,7 @@ def get_space() -> str | None: def is_zero_gpu_space() -> bool: - return os.getenv("SPACES_ZERO_GPU") == "true" + return os.getenv("SPACES_ZERO_GPU") in ("1", "true") def get_theme(theme: Theme | str | None) -> Theme: From e7959f2ed06277467dd34a6640980bf450c318d8 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 09:52:21 +0000 Subject: [PATCH 22/32] Remove max_size override since it has never been used until here --- gradio/blocks.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index 34878e861e8..f6fcdff216b 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -2427,8 +2427,6 @@ def queue( """ if api_open is not None: self.api_open = api_open - if utils.is_zero_gpu_space(): - max_size = 1 if max_size is None else max_size self._queue = queueing.Queue( live_updates=status_update_rate == "auto", concurrency_count=self.max_threads, From ac328281c3832746e79a18f141113849f58d60ca Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 18:28:54 +0000 Subject: [PATCH 23/32] Propagate context to handler --- gradio/blocks.py | 10 ++++++++++ gradio/helpers.py | 1 + gradio/mcp.py | 3 +++ gradio/queueing.py | 12 ++++++++++++ gradio/route_utils.py | 3 +++ gradio/routes.py | 4 ++++ 6 files changed, 33 insertions(+) diff --git a/gradio/blocks.py b/gradio/blocks.py index f6fcdff216b..ee43baf994f 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextvars import copy import dataclasses import hashlib @@ -18,6 +19,7 @@ import webbrowser from collections import defaultdict from collections.abc import AsyncIterator, Callable, Coroutine, Sequence, Set +from functools import partial from pathlib import Path from types import ModuleType, SimpleNamespace from typing import TYPE_CHECKING, Any, Literal, Union, cast @@ -1546,6 +1548,7 @@ def __call__(self, *inputs, fn_index: int = 0, api_name: str | None = None): self.process_api, block_fn=fn, inputs=processed_inputs, + context=None, request=None, state={}, explicit_call=True, @@ -1565,6 +1568,7 @@ async def call_function( block_fn: BlockFunction | int, processed_input: list[Any], iterator: AsyncIterator[Any] | None = None, + context: contextvars.Context | None = None, requests: Request | list[Request] | None = None, event_id: str | None = None, event_data: EventData | None = None, @@ -1598,6 +1602,9 @@ async def call_function( state=state, ) + if context is not None: + fn = partial(context.copy().run, fn) + if iterator is None: # If not a generator function that has already run if block_fn.inputs_as_dict: processed_input = [ @@ -2077,6 +2084,7 @@ async def process_api( block_fn: BlockFunction | int, inputs: list[Any], state: SessionState | None = None, + context: contextvars.Context | None = None, request: Request | list[Request] | None = None, iterator: AsyncIterator | None = None, session_hash: str | None = None, @@ -2144,6 +2152,7 @@ async def process_api( block_fn, list(zip(*inputs, strict=False)), None, + context, request, event_id, event_data, @@ -2179,6 +2188,7 @@ async def process_api( block_fn, inputs, old_iterator, + context, request, event_id, event_data, diff --git a/gradio/helpers.py b/gradio/helpers.py index 3420352638d..ad2344bb369 100644 --- a/gradio/helpers.py +++ b/gradio/helpers.py @@ -560,6 +560,7 @@ async def cache(self, example_id: int | None = None) -> None: prediction = await self.root_block.process_api( block_fn=self.root_block.default_config.fns[fn_index], inputs=processed_input, + context=None, request=None, in_event_listener=self.cache_examples != "lazy", ) diff --git a/gradio/mcp.py b/gradio/mcp.py index 849cc4a5547..621e33a0f57 100644 --- a/gradio/mcp.py +++ b/gradio/mcp.py @@ -1,5 +1,6 @@ import base64 import contextlib +import contextvars import copy import os import re @@ -398,6 +399,7 @@ async def call_tool( endpoint_name, processed_args, request_headers, block_fn = ( self._prepare_tool_call_args(name, arguments) ) + context = contextvars.copy_context() processed_args = self.insert_empty_state(block_fn.inputs, processed_args) if not block_fn.queue: @@ -410,6 +412,7 @@ async def call_tool( block_fn=block_fn, inputs=processed_args, state=session_state, + context=context, request=self.mcp_server.request_context.request, ) output_data = raw_output["data"] diff --git a/gradio/queueing.py b/gradio/queueing.py index 1503c8b7d39..413fead8530 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import contextvars import copy import inspect import multiprocessing @@ -60,12 +61,14 @@ def __init__( self, session_hash: str | None, fn: BlockFunction, + context: contextvars.Context | None, request: fastapi.Request, username: str | None, ): self._id = uuid.uuid4().hex self.session_hash: str = session_hash or self._id self.fn = fn + self.context = context self.request = request self.username = username self.concurrency_id = fn.concurrency_id @@ -322,6 +325,7 @@ async def push( fn = self.blocks.fns[body.fn_index] fn = route_utils.get_fn(self.blocks, None, body) + context = contextvars.copy_context() self.create_event_queue_for_fn(fn) if fn.validator is not None: gr_request = route_utils.compile_gr_request( @@ -363,6 +367,7 @@ async def push( event = Event( body.session_hash, validator_fn, + context, request, username, ) @@ -372,6 +377,7 @@ async def push( body=body, gr_request=gr_request, fn=validator_fn, + context=context, root_path=root_path, ) @@ -396,6 +402,7 @@ async def push( event = Event( body.session_hash, fn, + context, request, username, ) @@ -433,6 +440,7 @@ async def push( body=body, gr_request=gr_request, fn=fn, + context=context, root_path=root_path, ) while response and response.get("is_generating", False): @@ -448,6 +456,7 @@ async def push( body=body, gr_request=gr_request, fn=fn, + context=context, root_path=root_path, ) cache_duration = time.time() - cache_start @@ -826,6 +835,7 @@ async def process_events( ) -> None: awake_events: list[Event] = [] fn = events[0].fn + context = events[0].context success = False try: for event in events: @@ -906,6 +916,7 @@ async def process_events( body=body, gr_request=gr_request, fn=fn, + context=context, root_path=root_path, ) end = time.monotonic() @@ -996,6 +1007,7 @@ async def process_events( body=body, gr_request=gr_request, fn=fn, + context=context, root_path=root_path, ) end = time.monotonic() diff --git a/gradio/route_utils.py b/gradio/route_utils.py index 8012561bda4..da6d090320d 100644 --- a/gradio/route_utils.py +++ b/gradio/route_utils.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import contextvars import functools import hashlib import hmac @@ -353,6 +354,7 @@ async def call_process_api( body: PredictBodyInternal, gr_request: Union[Request, list[Request]], fn: BlockFunction, + context: contextvars.Context | None, root_path: str, ): session_state, iterator = restore_session_state(app=app, body=body) @@ -375,6 +377,7 @@ async def call_process_api( output = await app.get_blocks().process_api( block_fn=fn, inputs=inputs, + context=context, request=gr_request, state=session_state, iterator=iterator, diff --git a/gradio/routes.py b/gradio/routes.py index 3aba77cdf75..0b2e66f6496 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -5,6 +5,7 @@ import asyncio import contextlib +import contextvars import hashlib import importlib.resources import inspect @@ -1419,6 +1420,7 @@ async def iterator(): body=body, gr_request=req, fn=app.get_blocks().fns[fn_index], + context=None, root_path=root_path, ) # This will mark the state to be deleted in an hour @@ -1454,6 +1456,7 @@ async def predict( fn = route_utils.get_fn( blocks=app.get_blocks(), api_name=api_name, body=body ) + context = contextvars.copy_context() if not app.get_blocks().api_open and fn.queue: raise HTTPException( @@ -1477,6 +1480,7 @@ async def predict( body=body, gr_request=gr_request, fn=fn, + context=context, root_path=root_path, ) except BaseException as error: From 607023ea42268434866a0c8b6dc0171f298a4080 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 18:47:07 +0000 Subject: [PATCH 24/32] Do not wrap fn with partial --- gradio/blocks.py | 19 ++++++++++++------- 1 file changed, 12 insertions(+), 7 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index ee43baf994f..2a0d4552abf 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -1602,9 +1602,6 @@ async def call_function( state=state, ) - if context is not None: - fn = partial(context.copy().run, fn) - if iterator is None: # If not a generator function that has already run if block_fn.inputs_as_dict: processed_input = [ @@ -1637,11 +1634,19 @@ async def call_function( processed_input[progress_index] = progress_tracker if inspect.iscoroutinefunction(fn): - prediction = await fn(*processed_input) + if context is not None: + prediction = await context.copy().run(fn, *processed_input) + else: + prediction = await fn(*processed_input) else: - prediction = await anyio.to_thread.run_sync( # type: ignore - fn, *processed_input, limiter=self.limiter - ) + if context is not None: + prediction = await anyio.to_thread.run_sync( # type: ignore + context.copy().run, fn, *processed_input, limiter=self.limiter + ) + else: + prediction = await anyio.to_thread.run_sync( # type: ignore + fn, *processed_input, limiter=self.limiter + ) else: prediction = None From e89c90973a8f37a4eecfc0cdebfe338400cd6db8 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 18:51:12 +0000 Subject: [PATCH 25/32] Add a test for context propagation --- test/test_queueing.py | 75 +++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 75 insertions(+) diff --git a/test/test_queueing.py b/test/test_queueing.py index 37ee3b04c6b..cef220201f9 100644 --- a/test/test_queueing.py +++ b/test/test_queueing.py @@ -1,4 +1,5 @@ import asyncio +import contextvars import json import time from unittest.mock import patch @@ -9,6 +10,27 @@ import gradio as gr from gradio.route_utils import API_PREFIX +from gradio.routes import App + + +request_context = contextvars.ContextVar("request_context", default="unset") + + +class ContextHeaderMiddleware: + def __init__(self, app): + self.app = app + + async def __call__(self, scope, receive, send): + if scope["type"] != "http": + await self.app(scope, receive, send) + return + headers = dict(scope.get("headers", [])) + value = headers.get(b"x-test-context", b"missing").decode() + token = request_context.set(value) + try: + await self.app(scope, receive, send) + finally: + request_context.reset(token) class TestQueueing: @@ -224,6 +246,59 @@ def tracking_create_task(coro, **kwargs): demo.close() +def test_queue_event_propagates_context_from_join_request(): + with gr.Blocks() as demo: + start = gr.Button() + output = gr.Textbox() + + def read_context(): + return request_context.get() + + start.click(read_context, None, output) + + demo.queue() + app = App.create_app(demo) + app.add_middleware(ContextHeaderMiddleware) + + try: + with TestClient(app) as test_client: + startup = test_client.get( + f"{API_PREFIX}/startup-events", + headers={"x-test-context": "startup"}, + ) + assert startup.status_code == 200 + + join = test_client.post( + f"{API_PREFIX}/queue/join", + headers={"x-test-context": "join"}, + json={ + "data": [], + "fn_index": 0, + "event_data": None, + "session_hash": "context_session", + "trigger_id": None, + }, + ) + assert join.status_code == 200 + + stream = test_client.get( + f"{API_PREFIX}/queue/data?session_hash=context_session" + ) + + output_data = None + for line in stream.iter_lines(): + if not line.startswith("data: "): + continue + message = json.loads(line[6:]) + if message["msg"] == "process_completed": + output_data = message["output"]["data"] + break + + assert output_data == ["join"] + finally: + demo.close() + + def test_cancel_removes_pending_event_from_queue(): """Cancelling a queued (not yet running) event should remove it from the queue.""" with gr.Blocks() as demo: From 666130994f0e844e318e9202f94468f5a4b31802 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 18:52:53 +0000 Subject: [PATCH 26/32] Ruff --- gradio/blocks.py | 16 +++++++--------- test/test_queueing.py | 1 - 2 files changed, 7 insertions(+), 10 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index 2a0d4552abf..dc9248963b3 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -19,7 +19,6 @@ import webbrowser from collections import defaultdict from collections.abc import AsyncIterator, Callable, Coroutine, Sequence, Set -from functools import partial from pathlib import Path from types import ModuleType, SimpleNamespace from typing import TYPE_CHECKING, Any, Literal, Union, cast @@ -1638,15 +1637,14 @@ async def call_function( prediction = await context.copy().run(fn, *processed_input) else: prediction = await fn(*processed_input) + elif context is not None: + prediction = await anyio.to_thread.run_sync( # type: ignore + context.copy().run, fn, *processed_input, limiter=self.limiter + ) else: - if context is not None: - prediction = await anyio.to_thread.run_sync( # type: ignore - context.copy().run, fn, *processed_input, limiter=self.limiter - ) - else: - prediction = await anyio.to_thread.run_sync( # type: ignore - fn, *processed_input, limiter=self.limiter - ) + prediction = await anyio.to_thread.run_sync( # type: ignore + fn, *processed_input, limiter=self.limiter + ) else: prediction = None diff --git a/test/test_queueing.py b/test/test_queueing.py index cef220201f9..ceebfc52e55 100644 --- a/test/test_queueing.py +++ b/test/test_queueing.py @@ -12,7 +12,6 @@ from gradio.route_utils import API_PREFIX from gradio.routes import App - request_context = contextvars.ContextVar("request_context", default="unset") From 424dbafa9c9e50340fd60a558010ea5783f58361 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Wed, 6 May 2026 20:31:35 +0000 Subject: [PATCH 27/32] Only pass context where needed --- gradio/blocks.py | 9 ++++----- gradio/helpers.py | 1 - gradio/mcp.py | 3 --- gradio/queueing.py | 2 +- gradio/route_utils.py | 2 +- gradio/routes.py | 4 ---- 6 files changed, 6 insertions(+), 15 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index dc9248963b3..43459594a1e 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -1547,7 +1547,6 @@ def __call__(self, *inputs, fn_index: int = 0, api_name: str | None = None): self.process_api, block_fn=fn, inputs=processed_inputs, - context=None, request=None, state={}, explicit_call=True, @@ -1567,12 +1566,12 @@ async def call_function( block_fn: BlockFunction | int, processed_input: list[Any], iterator: AsyncIterator[Any] | None = None, - context: contextvars.Context | None = None, requests: Request | list[Request] | None = None, event_id: str | None = None, event_data: EventData | None = None, in_event_listener: bool = False, state: SessionState | None = None, + context: contextvars.Context | None = None, ): """ Calls function with given index and preprocessed input, and measures process time. @@ -2087,7 +2086,6 @@ async def process_api( block_fn: BlockFunction | int, inputs: list[Any], state: SessionState | None = None, - context: contextvars.Context | None = None, request: Request | list[Request] | None = None, iterator: AsyncIterator | None = None, session_hash: str | None = None, @@ -2097,6 +2095,7 @@ async def process_api( simple_format: bool = False, explicit_call: bool = False, root_path: str | None = None, + context: contextvars.Context | None = None, ) -> dict[str, Any]: """ Processes API calls from the frontend. First preprocesses the data, @@ -2155,12 +2154,12 @@ async def process_api( block_fn, list(zip(*inputs, strict=False)), None, - context, request, event_id, event_data, in_event_listener, state, + context=context, ) manual_cache_used = used_manual_cache() preds = result["prediction"] @@ -2191,12 +2190,12 @@ async def process_api( block_fn, inputs, old_iterator, - context, request, event_id, event_data, in_event_listener, state, + context=context, ) manual_cache_used = used_manual_cache() diff --git a/gradio/helpers.py b/gradio/helpers.py index ad2344bb369..3420352638d 100644 --- a/gradio/helpers.py +++ b/gradio/helpers.py @@ -560,7 +560,6 @@ async def cache(self, example_id: int | None = None) -> None: prediction = await self.root_block.process_api( block_fn=self.root_block.default_config.fns[fn_index], inputs=processed_input, - context=None, request=None, in_event_listener=self.cache_examples != "lazy", ) diff --git a/gradio/mcp.py b/gradio/mcp.py index 621e33a0f57..849cc4a5547 100644 --- a/gradio/mcp.py +++ b/gradio/mcp.py @@ -1,6 +1,5 @@ import base64 import contextlib -import contextvars import copy import os import re @@ -399,7 +398,6 @@ async def call_tool( endpoint_name, processed_args, request_headers, block_fn = ( self._prepare_tool_call_args(name, arguments) ) - context = contextvars.copy_context() processed_args = self.insert_empty_state(block_fn.inputs, processed_args) if not block_fn.queue: @@ -412,7 +410,6 @@ async def call_tool( block_fn=block_fn, inputs=processed_args, state=session_state, - context=context, request=self.mcp_server.request_context.request, ) output_data = raw_output["data"] diff --git a/gradio/queueing.py b/gradio/queueing.py index 413fead8530..7b62c499cc1 100644 --- a/gradio/queueing.py +++ b/gradio/queueing.py @@ -61,7 +61,7 @@ def __init__( self, session_hash: str | None, fn: BlockFunction, - context: contextvars.Context | None, + context: contextvars.Context, request: fastapi.Request, username: str | None, ): diff --git a/gradio/route_utils.py b/gradio/route_utils.py index da6d090320d..b35bef0c77b 100644 --- a/gradio/route_utils.py +++ b/gradio/route_utils.py @@ -354,8 +354,8 @@ async def call_process_api( body: PredictBodyInternal, gr_request: Union[Request, list[Request]], fn: BlockFunction, - context: contextvars.Context | None, root_path: str, + context: contextvars.Context | None = None, ): session_state, iterator = restore_session_state(app=app, body=body) diff --git a/gradio/routes.py b/gradio/routes.py index 0b2e66f6496..3aba77cdf75 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -5,7 +5,6 @@ import asyncio import contextlib -import contextvars import hashlib import importlib.resources import inspect @@ -1420,7 +1419,6 @@ async def iterator(): body=body, gr_request=req, fn=app.get_blocks().fns[fn_index], - context=None, root_path=root_path, ) # This will mark the state to be deleted in an hour @@ -1456,7 +1454,6 @@ async def predict( fn = route_utils.get_fn( blocks=app.get_blocks(), api_name=api_name, body=body ) - context = contextvars.copy_context() if not app.get_blocks().api_open and fn.queue: raise HTTPException( @@ -1480,7 +1477,6 @@ async def predict( body=body, gr_request=gr_request, fn=fn, - context=context, root_path=root_path, ) except BaseException as error: From 27d267e4540be9da34d576cfd2d6c9194e991add Mon Sep 17 00:00:00 2001 From: cbensimon Date: Thu, 7 May 2026 10:24:38 +0000 Subject: [PATCH 28/32] Also for generators x async --- gradio/blocks.py | 11 ++++++++--- gradio/utils.py | 30 +++++++++++++++++++++++++++--- test/test_queueing.py | 27 +++++++++++++++++++++++++-- 3 files changed, 60 insertions(+), 8 deletions(-) diff --git a/gradio/blocks.py b/gradio/blocks.py index 43459594a1e..2006f513313 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import contextvars import copy import dataclasses @@ -1633,7 +1634,9 @@ async def call_function( if inspect.iscoroutinefunction(fn): if context is not None: - prediction = await context.copy().run(fn, *processed_input) + prediction = await context.copy().run( + asyncio.create_task, fn(*processed_input) + ) else: prediction = await fn(*processed_input) elif context is not None: @@ -1652,8 +1655,10 @@ async def call_function( if iterator is None: iterator = cast(AsyncIterator[Any], prediction) if inspect.isgenerator(iterator): - iterator = utils.SyncToAsyncIterator(iterator, self.limiter) - prediction = await utils.async_iteration(iterator) + iterator = utils.SyncToAsyncIterator( + iterator, self.limiter, context + ) + prediction = await utils.async_iteration(iterator, context=context) is_generating = True except StopAsyncIteration: n_outputs = len(block_fn.outputs) diff --git a/gradio/utils.py b/gradio/utils.py index 45d788246da..5478935d7aa 100644 --- a/gradio/utils.py +++ b/gradio/utils.py @@ -1,6 +1,7 @@ """Handy utility functions.""" import asyncio +import contextvars import copy import functools import hashlib @@ -52,6 +53,7 @@ ) import anyio +import anyio.to_thread import gradio_client.utils as client_utils import httpx import orjson @@ -857,14 +859,27 @@ def run_sync_iterator_async(iterator): class SyncToAsyncIterator: """Treat a synchronous iterator as async one.""" - def __init__(self, iterator, limiter) -> None: + def __init__( + self, + iterator, + limiter, + context: contextvars.Context | None = None, + ) -> None: self.iterator = iterator self.limiter = limiter + self.context = context def __aiter__(self): return self async def __anext__(self): + if self.context is not None: + return await anyio.to_thread.run_sync( + self.context.copy().run, + run_sync_iterator_async, + self.iterator, + limiter=self.limiter, + ) return await anyio.to_thread.run_sync( run_sync_iterator_async, self.iterator, limiter=self.limiter ) @@ -882,8 +897,17 @@ async def aclose(self, timeout=60.0, retry_interval=0.05): raise -async def async_iteration(iterator): - return await anext(iterator) +async def async_iteration( + iterator, + context: contextvars.Context | None = None, +): + if context is None: + return await anext(iterator) + + def create_task(): + return asyncio.create_task(anext(iterator)) + + return await context.copy().run(create_task) @contextmanager diff --git a/test/test_queueing.py b/test/test_queueing.py index ceebfc52e55..86bbf086b46 100644 --- a/test/test_queueing.py +++ b/test/test_queueing.py @@ -245,7 +245,11 @@ def tracking_create_task(coro, **kwargs): demo.close() -def test_queue_event_propagates_context_from_join_request(): +@pytest.mark.parametrize("generator", [False, True]) +@pytest.mark.parametrize("asynchronous", [False, True]) +def test_queue_event_propagates_context_from_join_request( + asynchronous: bool, generator: bool +): with gr.Blocks() as demo: start = gr.Button() output = gr.Textbox() @@ -253,7 +257,26 @@ def test_queue_event_propagates_context_from_join_request(): def read_context(): return request_context.get() - start.click(read_context, None, output) + def read_context_gen(): + yield request_context.get() + + async def read_context_async(): + return request_context.get() + + async def read_context_asyncgen(): + yield request_context.get() + + match asynchronous, generator: + case False, False: + fn = read_context + case False, True: + fn = read_context_gen + case True, False: + fn = read_context_async + case True, True: + fn = read_context_asyncgen + + start.click(fn, None, output) demo.queue() app = App.create_app(demo) From 7f70042b1c43db0c038728c067937c8d070a2bac Mon Sep 17 00:00:00 2001 From: cbensimon Date: Thu, 7 May 2026 10:26:46 +0000 Subject: [PATCH 29/32] ty --- gradio/route_utils.py | 2 +- test/test_queueing.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/gradio/route_utils.py b/gradio/route_utils.py index b35bef0c77b..0ff727a1cdb 100644 --- a/gradio/route_utils.py +++ b/gradio/route_utils.py @@ -1166,7 +1166,7 @@ async def iter_body(head: bytes, queue: asyncio.Queue[bytes | None]): yield chunk -def maybe_setup_zerogpu_middleware(app: App | fastapi.FastAPI): +def maybe_setup_zerogpu_middleware(app: fastapi.FastAPI): if not utils.is_zero_gpu_space(): return diff --git a/test/test_queueing.py b/test/test_queueing.py index 86bbf086b46..cc140fcf727 100644 --- a/test/test_queueing.py +++ b/test/test_queueing.py @@ -280,7 +280,7 @@ async def read_context_asyncgen(): demo.queue() app = App.create_app(demo) - app.add_middleware(ContextHeaderMiddleware) + app.add_middleware(ContextHeaderMiddleware) # ty: ignore[invalid-argument-type] try: with TestClient(app) as test_client: From 43dde934095a9706204daab7ab52a50e61c217a4 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Thu, 7 May 2026 17:22:29 +0000 Subject: [PATCH 30/32] Move middleware setup in dedicated zerogpu module + add disconnect detector --- gradio/blocks.py | 2 +- gradio/route_utils.py | 31 ---------------------- gradio/routes.py | 2 +- gradio/zerogpu.py | 61 +++++++++++++++++++++++++++++++++++++++++++ test/requirements.txt | 2 +- 5 files changed, 64 insertions(+), 34 deletions(-) create mode 100644 gradio/zerogpu.py diff --git a/gradio/blocks.py b/gradio/blocks.py index 2006f513313..20219205258 100644 --- a/gradio/blocks.py +++ b/gradio/blocks.py @@ -79,7 +79,6 @@ from gradio.route_utils import ( API_PREFIX, MediaStream, - maybe_setup_zerogpu_middleware, slugify, ) from gradio.routes import INTERNAL_ROUTES, VERSION, App, Request @@ -101,6 +100,7 @@ get_package_version, get_upload_folder, ) +from gradio.zerogpu import maybe_setup_zerogpu_middleware if TYPE_CHECKING: # Only import for type checking (is False at runtime). from gradio.components.base import Component diff --git a/gradio/route_utils.py b/gradio/route_utils.py index 0ff727a1cdb..bc7ac59fae7 100644 --- a/gradio/route_utils.py +++ b/gradio/route_utils.py @@ -47,7 +47,6 @@ from starlette.types import ASGIApp, Message, Receive, Scope, Send from gradio import processing_utils, utils -from gradio.context import MultiprocessWorkerContextualizer from gradio.data_classes import ( BlocksConfigDict, MediaStreamChunk, @@ -1164,33 +1163,3 @@ async def iter_body(head: bytes, queue: asyncio.Queue[bytes | None]): yield head while (chunk := await queue.get()) is not None: yield chunk - - -def maybe_setup_zerogpu_middleware(app: fastapi.FastAPI): - if not utils.is_zero_gpu_space(): - return - - try: - from spaces.zero import ZeroGPUMiddleware - except ImportError: - return - - from gradio.helpers import log_message - - app.add_middleware( - ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] - exception_mapper=lambda err, exc: ( - setattr(exc, "print_exception", False) or exc - if isinstance(exc, Error) - else Error( - title=err["detail"]["title"], - message=err["detail"]["message"], - ) - ), - log_emitter=lambda log: log_message( - title=log["title"], - message=log["message"], - level=log["level"], - ), - worker_contextualizer=MultiprocessWorkerContextualizer, - ) diff --git a/gradio/routes.py b/gradio/routes.py index 3aba77cdf75..ef66a0ccf71 100644 --- a/gradio/routes.py +++ b/gradio/routes.py @@ -107,7 +107,6 @@ Request, compare_passwords_securely, create_lifespan_handler, - maybe_setup_zerogpu_middleware, move_uploaded_files_to_cache, ) from gradio.screen_recording_utils import process_video_with_ffmpeg @@ -129,6 +128,7 @@ get_upload_folder, safe_aclose_iterator, ) +from gradio.zerogpu import maybe_setup_zerogpu_middleware if TYPE_CHECKING: from gradio.blocks import Block diff --git a/gradio/zerogpu.py b/gradio/zerogpu.py new file mode 100644 index 00000000000..ea2dc03749c --- /dev/null +++ b/gradio/zerogpu.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import fastapi + +from gradio import utils +from gradio.context import LocalContext, MultiprocessWorkerContextualizer +from gradio.exceptions import Error +from gradio.helpers import log_message + +if TYPE_CHECKING: + from spaces.zero import ZeroGPUErrorResponse, ZeroGPULog + + +def exception_mapper(err: ZeroGPUErrorResponse, exc: Exception | None): + if isinstance(exc, Error): + exc.print_exception = False + return exc + detail = err["detail"] + return Error( + title=detail["title"], + message=detail["message"], + ) + + +def log_emitter(log: ZeroGPULog): + log_message( + title=log["title"], + message=log["message"], + level=log["level"], + ) + + +def disconnect_detector(): + blocks = LocalContext.blocks.get(None) + event_id = LocalContext.event_id.get(None) + if blocks is not None and event_id is not None: + jobs = blocks._queue.active_jobs + for event in [evt for job in jobs if job is not None for evt in job]: + if event._id == event_id: + return not event.alive + return False + + +def maybe_setup_zerogpu_middleware(app: fastapi.FastAPI): + if not utils.is_zero_gpu_space(): + return + + try: + from spaces.zero import ZeroGPUMiddleware + except ImportError: + return + + app.add_middleware( + ZeroGPUMiddleware, # ty: ignore[invalid-argument-type] + exception_mapper=exception_mapper, + log_emitter=log_emitter, + worker_contextualizer=MultiprocessWorkerContextualizer, + disconnect_detector=disconnect_detector, + ) diff --git a/test/requirements.txt b/test/requirements.txt index 08bed143f0c..50d0949e085 100644 --- a/test/requirements.txt +++ b/test/requirements.txt @@ -338,7 +338,7 @@ sniffio==1.3.1 # openai sortedcontainers==2.4.0 # via hypothesis -spaces==0.50.dev0 +spaces==0.50.dev2 # via -r test/requirements.in stack-data==0.6.3 # via ipython From 99f0df68c0b6aebbb87cdc22986ff05ae8f2b3b3 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Thu, 7 May 2026 17:34:17 +0000 Subject: [PATCH 31/32] spaces 0.50.dev3 --- test/requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/requirements.txt b/test/requirements.txt index 50d0949e085..61a44ca5059 100644 --- a/test/requirements.txt +++ b/test/requirements.txt @@ -338,7 +338,7 @@ sniffio==1.3.1 # openai sortedcontainers==2.4.0 # via hypothesis -spaces==0.50.dev2 +spaces==0.50.dev3 # via -r test/requirements.in stack-data==0.6.3 # via ipython From 76cc5666f9e71e8706331172ba7de34fec1e5592 Mon Sep 17 00:00:00 2001 From: cbensimon Date: Thu, 7 May 2026 18:11:46 +0000 Subject: [PATCH 32/32] Add a test for cursor bot --- test/test_queueing.py | 65 +++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 65 insertions(+) diff --git a/test/test_queueing.py b/test/test_queueing.py index cc140fcf727..79796ce1114 100644 --- a/test/test_queueing.py +++ b/test/test_queueing.py @@ -1,6 +1,7 @@ import asyncio import contextvars import json +import threading import time from unittest.mock import patch @@ -321,6 +322,70 @@ async def read_context_asyncgen(): demo.close() +def test_queue_context_task_is_cancelled_with_event(): + started = threading.Event() + cancelled = threading.Event() + completed = threading.Event() + + with gr.Blocks() as demo: + start = gr.Button() + output = gr.Textbox() + + async def wait_forever(): + try: + assert request_context.get() == "join" + started.set() + await asyncio.Event().wait() + completed.set() + return "done" + finally: + cancelled.set() + + start.click(wait_forever, None, output) + + demo.queue() + app = App.create_app(demo) + app.add_middleware(ContextHeaderMiddleware) # ty: ignore[invalid-argument-type] + + try: + with TestClient(app) as test_client: + startup = test_client.get( + f"{API_PREFIX}/startup-events", + headers={"x-test-context": "startup"}, + ) + assert startup.status_code == 200 + + join = test_client.post( + f"{API_PREFIX}/queue/join", + headers={"x-test-context": "join"}, + json={ + "data": [], + "fn_index": 0, + "event_data": None, + "session_hash": "cancel_context_session", + "trigger_id": None, + }, + ) + assert join.status_code == 200 + event_id = join.json()["event_id"] + + assert started.wait(timeout=2) + + cancel = test_client.post( + f"{API_PREFIX}/cancel", + json={ + "session_hash": "cancel_context_session", + "fn_index": 0, + "event_id": event_id, + }, + ) + assert cancel.status_code == 200 + assert cancelled.wait(timeout=2) + assert not completed.is_set() + finally: + demo.close() + + def test_cancel_removes_pending_event_from_queue(): """Cancelling a queued (not yet running) event should remove it from the queue.""" with gr.Blocks() as demo: