Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/fresh-lizards-knock.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"gradio": minor
---

feat:workflow: forward x-ip-token so zerogpu spaces bill the caller
14 changes: 10 additions & 4 deletions gradio/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@
utils,
)
from gradio.brotli_middleware import BrotliMiddleware
from gradio.context import Context
from gradio.context import Context, LocalContext
from gradio.data_classes import (
CancelBody,
ComponentServerBlobBody,
Expand Down Expand Up @@ -1718,12 +1718,18 @@ async def component_server(
request, # type: ignore
None,
)
if inspect.iscoroutinefunction(fn):
return await fn(*processed_input)
else:
# So `gradio_client.Client` in a server fn can forward `x-ip-token`
# (ZeroGPU quota) — regular event handlers get this via
# `get_function_with_locals`, this path bypasses that wrapping.
LocalContext.request.set(request) # type: ignore
Comment thread
hannahblair marked this conversation as resolved.
Outdated
try:
if inspect.iscoroutinefunction(fn):
return await fn(*processed_input)
return await anyio.to_thread.run_sync(
fn, *processed_input, limiter=app.get_blocks().limiter
)
finally:
LocalContext.request.set(None)

@router.get(
"/queue/status",
Expand Down
35 changes: 35 additions & 0 deletions test/test_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -2553,6 +2553,41 @@ def get_url(self, request: gr.Request):
assert response.json()["_url"].endswith("/gradio_api/component_server")


def test_server_fn_forwards_x_ip_token_via_local_context():
"""/component_server must set LocalContext.request so gradio_client.Client
can forward the caller's x-ip-token to downstream ZeroGPU Spaces."""
import requests

from gradio.components.base import server
from gradio.context import LocalContext

def get_ip_token(self, _data):
req = LocalContext.request.get(None)
return req.headers.get("x-ip-token") if req else None

tb = gr.Textbox()
tb.get_ip_token = server(get_ip_token) # type: ignore
iface = gr.Interface(lambda x: x, inputs=tb, outputs="text")
component_id = next(
c["id"]
for c in iface.config["components"]
if c["type"] == "textbox" # type: ignore
)
_, local_url, _ = iface.launch(prevent_thread_lock=True)
response = requests.post(
f"{local_url}/gradio_api/component_server",
json={
"session_hash": "foo",
"component_id": component_id,
"fn_name": "get_ip_token",
"data": json.dumps({}),
},
headers={"x-ip-token": "test-token"},
)
assert response.status_code == 200
assert response.json() == "test-token"


def test_slugify():
items = (
("Hello, World!", "hello-world"),
Expand Down
Loading