gradio-app--gradio
adf0d17497
publish / version_or_publish (push) Has been cancelled
storybook-build / changes (push) Has been cancelled
storybook-build / :storybook-build (push) Has been cancelled
Sync Gradio Skills to Hugging Face / sync-skills (push) Has been cancelled
functional / changes (push) Has been cancelled
functional / build-frontend (push) Has been cancelled
functional / functional-test-SSR=false (push) Has been cancelled
functional / functional-reload (push) Has been cancelled
js / changes (push) Has been cancelled
js / js-test (push) Has been cancelled
docs-build / changes (push) Has been cancelled
docs-build / docs-build (push) Has been cancelled
docs-build / website-build (push) Has been cancelled
functional / functional-test-SSR=true (push) Has been cancelled
hygiene / hygiene-test (push) Has been cancelled
python / changes (push) Has been cancelled
python / build (push) Has been cancelled
python / test-ubuntu-latest-flaky (push) Has been cancelled
python / test-ubuntu-latest-not-flaky (push) Has been cancelled
python / test-windows-latest-flaky (push) Has been cancelled
python / test-windows-latest-not-flaky (push) Has been cancelled
277 行
8.5 KiB
Python
277 行
8.5 KiB
Python
"""Centralized code snippet generation for Gradio API endpoints. Generates Python, JavaScript, and cURL code snippets from API info dicts."""
|
|
|
|
import copy
|
|
import json
|
|
import re
|
|
from typing import Any
|
|
|
|
BLOB_COMPONENTS = {
|
|
"Audio",
|
|
"DownloadButton",
|
|
"File",
|
|
"Image",
|
|
"ImageSlider",
|
|
"Model3D",
|
|
"UploadButton",
|
|
"Video",
|
|
}
|
|
|
|
|
|
def _is_file_data(obj: Any) -> bool:
|
|
return (
|
|
isinstance(obj, dict)
|
|
and "url" in obj
|
|
and obj.get("url")
|
|
and "meta" in obj
|
|
and isinstance(obj.get("meta"), dict)
|
|
and obj["meta"].get("_type") == "gradio.FileData"
|
|
)
|
|
|
|
|
|
def _has_file_data(obj: Any) -> bool:
|
|
if isinstance(obj, dict):
|
|
if _is_file_data(obj):
|
|
return True
|
|
return any(_has_file_data(v) for v in obj.values())
|
|
if isinstance(obj, (list, tuple)):
|
|
return any(_has_file_data(item) for item in obj)
|
|
return False
|
|
|
|
|
|
def _replace_file_data_py(obj: Any) -> Any:
|
|
if isinstance(obj, dict) and _is_file_data(obj):
|
|
return f"handle_file('{obj['url']}')"
|
|
if isinstance(obj, (list, tuple)):
|
|
return [_replace_file_data_py(item) for item in obj]
|
|
if isinstance(obj, dict):
|
|
return {k: _replace_file_data_py(v) for k, v in obj.items()}
|
|
return obj
|
|
|
|
|
|
def _simplify_file_data(obj: Any) -> Any:
|
|
if isinstance(obj, dict) and _is_file_data(obj):
|
|
return {"path": obj["url"], "meta": {"_type": "gradio.FileData"}}
|
|
if isinstance(obj, (list, tuple)):
|
|
return [_simplify_file_data(item) for item in obj]
|
|
if isinstance(obj, dict):
|
|
return {k: _simplify_file_data(v) for k, v in obj.items()}
|
|
return obj
|
|
|
|
|
|
_UNQUOTED = "UNQUOTED_GRADIO_"
|
|
|
|
|
|
def _stringify_py(obj: Any) -> str:
|
|
def _prepare(o: Any) -> Any:
|
|
if o is None:
|
|
return f"{_UNQUOTED}None"
|
|
if isinstance(o, bool):
|
|
return f"{_UNQUOTED}True" if o else f"{_UNQUOTED}False"
|
|
if isinstance(o, str) and o.startswith("handle_file(") and o.endswith(")"):
|
|
return f"{_UNQUOTED}{o}"
|
|
if isinstance(o, (list, tuple)):
|
|
return [_prepare(item) for item in o]
|
|
if isinstance(o, dict):
|
|
return {k: _prepare(v) for k, v in o.items()}
|
|
return o
|
|
|
|
prepared = _prepare(obj)
|
|
result = json.dumps(prepared, default=str)
|
|
result = re.sub(
|
|
rf'"{_UNQUOTED}(handle_file\([^)]*\))"',
|
|
r"\1",
|
|
result,
|
|
)
|
|
result = result.replace(f'"{_UNQUOTED}None"', "None")
|
|
result = result.replace(f'"{_UNQUOTED}True"', "True")
|
|
result = result.replace(f'"{_UNQUOTED}False"', "False")
|
|
return result
|
|
|
|
|
|
def _represent_value(value: Any, python_type: str | None, lang: str) -> str:
|
|
if python_type is None:
|
|
return "None" if lang == "py" else "null"
|
|
if value is None:
|
|
return "None" if lang == "py" else "null"
|
|
if python_type in ("string", "str"):
|
|
return f'"{value}"'
|
|
if python_type == "number":
|
|
return str(value)
|
|
if python_type in ("boolean", "bool"):
|
|
if lang == "py":
|
|
return "True" if value else "False"
|
|
return str(value).lower() if isinstance(value, bool) else str(value)
|
|
if python_type == "List[str]":
|
|
return json.dumps(value)
|
|
if python_type.startswith("Literal['"):
|
|
return f'"{value}"'
|
|
|
|
if isinstance(value, str):
|
|
if value == "":
|
|
return "None" if lang == "py" else "null"
|
|
return value
|
|
|
|
value = copy.deepcopy(value)
|
|
if lang == "bash":
|
|
value = _simplify_file_data(value)
|
|
if lang == "py":
|
|
value = _replace_file_data_py(value)
|
|
return _stringify_py(value)
|
|
|
|
|
|
def _get_param_value(param: dict) -> Any:
|
|
if param.get("parameter_has_default"):
|
|
return param.get("parameter_default")
|
|
return param.get("example_input")
|
|
|
|
|
|
def generate_python_snippet(
|
|
api_name: str,
|
|
params: list[dict],
|
|
src: str,
|
|
) -> str:
|
|
has_file = any(_has_file_data(p.get("example_input")) for p in params)
|
|
imports = "from gradio_client import Client"
|
|
if has_file:
|
|
imports += ", handle_file"
|
|
|
|
lines = [imports, ""]
|
|
lines.append(f'client = Client("{src}")')
|
|
|
|
predict_args = []
|
|
for p in params:
|
|
name = p.get("parameter_name") or p.get("label", "input")
|
|
value = _get_param_value(p)
|
|
ptype = p.get("python_type", {}).get("type")
|
|
formatted = _represent_value(value, ptype, "py")
|
|
predict_args.append(f"\t{name}={formatted},")
|
|
|
|
lines.append("result = client.predict(")
|
|
lines.extend(predict_args)
|
|
lines.append(f'\tapi_name="{api_name}",')
|
|
lines.append(")")
|
|
lines.append("print(result)")
|
|
|
|
return "\n".join(lines)
|
|
|
|
|
|
def generate_js_snippet(
|
|
api_name: str,
|
|
params: list[dict],
|
|
src: str,
|
|
) -> str:
|
|
blob_params = [p for p in params if p.get("component") in BLOB_COMPONENTS]
|
|
|
|
lines = ['import { Client } from "@gradio/client";', ""]
|
|
|
|
for i, bp in enumerate(blob_params):
|
|
example = bp.get("example_input", {})
|
|
url = example.get("url", "") if isinstance(example, dict) else ""
|
|
component = bp.get("component", "")
|
|
lines.append(f'const response_{i} = await fetch("{url}");')
|
|
lines.append(f"const example{component} = await response_{i}.blob();")
|
|
|
|
if blob_params:
|
|
lines.append("")
|
|
|
|
lines.append(f'const client = await Client.connect("{src}");')
|
|
|
|
blob_component_names = {bp.get("component") for bp in blob_params}
|
|
|
|
predict_args = []
|
|
for p in params:
|
|
name = p.get("parameter_name") or p.get("label", "input")
|
|
component = p.get("component", "")
|
|
if component in blob_component_names:
|
|
predict_args.append(f"\t\t{name}: example{component},")
|
|
else:
|
|
value = _get_param_value(p)
|
|
ptype = p.get("python_type", {}).get("type")
|
|
formatted = _represent_value(value, ptype, "js")
|
|
predict_args.append(f"\t\t{name}: {formatted},")
|
|
|
|
lines.append(f'const result = await client.predict("{api_name}", {{')
|
|
lines.extend(predict_args)
|
|
lines.append("});")
|
|
lines.append("")
|
|
lines.append("console.log(result.data);")
|
|
|
|
return "\n".join(lines)
|
|
|
|
|
|
def generate_bash_snippet(
|
|
api_name: str,
|
|
params: list[dict],
|
|
root: str,
|
|
api_prefix: str = "/",
|
|
) -> str:
|
|
normalised_root = root.rstrip("/")
|
|
normalised_prefix = api_prefix if api_prefix else "/"
|
|
endpoint_name = api_name.lstrip("/")
|
|
|
|
has_file = any(_has_file_data(p.get("example_input")) for p in params)
|
|
upload_url = f"{normalised_root}{normalised_prefix}upload"
|
|
|
|
lines: list[str] = []
|
|
|
|
file_param_names: list[str] = []
|
|
if has_file:
|
|
for p in params:
|
|
if _has_file_data(p.get("example_input")):
|
|
name = p.get("parameter_name") or p.get("label", "input")
|
|
file_param_names.append(name)
|
|
lines.append(
|
|
f"FILE_PATH=$(curl -s -X POST {upload_url}"
|
|
" -F 'files=@/path/to/your/file'"
|
|
" | tr -d '[]\" ')"
|
|
)
|
|
lines.append("")
|
|
|
|
data_dict = {}
|
|
for p in params:
|
|
name = p.get("parameter_name") or p.get("label", "input")
|
|
if name in file_param_names:
|
|
data_dict[name] = "FILE_PATH_PLACEHOLDER"
|
|
else:
|
|
value = _get_param_value(p)
|
|
ptype = p.get("python_type", {}).get("type")
|
|
formatted = _represent_value(value, ptype, "bash")
|
|
data_dict[name] = formatted
|
|
|
|
data_entries = ", ".join(f'"{k}": {v}' for k, v in data_dict.items())
|
|
data_str = "{" + data_entries + "}"
|
|
for _ in file_param_names:
|
|
replacement = '{"path": "\'$FILE_PATH\'", "meta": {"_type": "gradio.FileData"}}'
|
|
data_str = data_str.replace("FILE_PATH_PLACEHOLDER", replacement)
|
|
|
|
base_url = f"{normalised_root}{normalised_prefix}call/v2/{endpoint_name}"
|
|
get_url = f"{normalised_root}{normalised_prefix}call/{endpoint_name}"
|
|
|
|
lines.extend(
|
|
[
|
|
f'curl -X POST {base_url} -s -H "Content-Type: application/json" \\',
|
|
f" -d '{data_str}' \\",
|
|
" | awk -F'\"' '{ print $4}' \\",
|
|
f" | read EVENT_ID; curl -N {get_url}/$EVENT_ID",
|
|
]
|
|
)
|
|
|
|
return "\n".join(lines)
|
|
|
|
|
|
def generate_code_snippets(
|
|
api_name: str,
|
|
endpoint_info: dict,
|
|
root: str,
|
|
space_id: str | None = None,
|
|
api_prefix: str = "/",
|
|
) -> dict[str, str]:
|
|
params = endpoint_info.get("parameters", [])
|
|
src = space_id or root
|
|
|
|
return {
|
|
"python": generate_python_snippet(api_name, params, src),
|
|
"javascript": generate_js_snippet(api_name, params, src),
|
|
"bash": generate_bash_snippet(api_name, params, root, api_prefix),
|
|
}
|