Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
15 changes: 11 additions & 4 deletions gradio/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -1718,12 +1718,19 @@ async def component_server(
request, # type: ignore
None,
)
fn = utils.get_function_with_locals(
fn,
app.get_blocks(),
event_id=None,
in_event_listener=False,
request=request, # type: ignore
state=state,
)
if inspect.iscoroutinefunction(fn):
return await fn(*processed_input)
else:
return await anyio.to_thread.run_sync(
fn, *processed_input, limiter=app.get_blocks().limiter
)
return await anyio.to_thread.run_sync(
fn, *processed_input, limiter=app.get_blocks().limiter
)

@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