Decrypt-proxied-gradio-payloads
Browse files
app.py
CHANGED
|
@@ -40,6 +40,7 @@ import numpy as np
|
|
| 40 |
import py7zr
|
| 41 |
import spaces
|
| 42 |
import torch
|
|
|
|
| 43 |
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
| 44 |
from cryptography.hazmat.primitives.kdf.scrypt import Scrypt
|
| 45 |
from fastapi import Body, Depends, HTTPException, Query
|
|
@@ -100,6 +101,10 @@ SALT_SIZE = 16
|
|
| 100 |
NONCE_SIZE = 12
|
| 101 |
SCRYPT_N = 2**14
|
| 102 |
FILE_MAGIC = b"EPNG1"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
IMAGE_SUFFIXES = (".epng",)
|
| 104 |
PREVIEW_CACHE_SIZE = 256
|
| 105 |
PREVIEW_QUALITY = 70
|
|
@@ -1874,6 +1879,42 @@ def image_key(salt):
|
|
| 1874 |
).derive(PASSWORD.encode())
|
| 1875 |
|
| 1876 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1877 |
def save_image(image, from_api):
|
| 1878 |
path = IMAGE_DIR / datetime.now(TIMEZONE).date().isoformat()
|
| 1879 |
path.mkdir(parents=True, exist_ok=True)
|
|
@@ -2759,7 +2800,107 @@ def run_job(job_id, workflow):
|
|
| 2759 |
}
|
| 2760 |
|
| 2761 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2762 |
api = App()
|
|
|
|
| 2763 |
api.add_middleware(
|
| 2764 |
CORSMiddleware,
|
| 2765 |
allow_origins=["*"],
|
|
|
|
| 40 |
import py7zr
|
| 41 |
import spaces
|
| 42 |
import torch
|
| 43 |
+
from cryptography.exceptions import InvalidTag
|
| 44 |
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
| 45 |
from cryptography.hazmat.primitives.kdf.scrypt import Scrypt
|
| 46 |
from fastapi import Body, Depends, HTTPException, Query
|
|
|
|
| 101 |
NONCE_SIZE = 12
|
| 102 |
SCRYPT_N = 2**14
|
| 103 |
FILE_MAGIC = b"EPNG1"
|
| 104 |
+
PROXY_ACCESS_VALUE = b"haip"
|
| 105 |
+
PROXY_ENCRYPTION = b"aes-256-gcm"
|
| 106 |
+
PROXY_ENCRYPTION_HEADER = b"x-gradio-comfy-encryption"
|
| 107 |
+
PROXY_MAGIC = b"GCV1"
|
| 108 |
IMAGE_SUFFIXES = (".epng",)
|
| 109 |
PREVIEW_CACHE_SIZE = 256
|
| 110 |
PREVIEW_QUALITY = 70
|
|
|
|
| 1879 |
).derive(PASSWORD.encode())
|
| 1880 |
|
| 1881 |
|
| 1882 |
+
def proxy_key(salt):
|
| 1883 |
+
return Scrypt(
|
| 1884 |
+
salt=salt,
|
| 1885 |
+
length=32,
|
| 1886 |
+
n=SCRYPT_N,
|
| 1887 |
+
r=8,
|
| 1888 |
+
p=1,
|
| 1889 |
+
).derive(PROXY_ACCESS_VALUE)
|
| 1890 |
+
|
| 1891 |
+
|
| 1892 |
+
def encrypt_proxy_payload(data):
|
| 1893 |
+
salt = os.urandom(SALT_SIZE)
|
| 1894 |
+
nonce = os.urandom(NONCE_SIZE)
|
| 1895 |
+
return (
|
| 1896 |
+
PROXY_MAGIC
|
| 1897 |
+
+ salt
|
| 1898 |
+
+ nonce
|
| 1899 |
+
+ AESGCM(proxy_key(salt)).encrypt(nonce, data, PROXY_MAGIC)
|
| 1900 |
+
)
|
| 1901 |
+
|
| 1902 |
+
|
| 1903 |
+
def decrypt_proxy_payload(data):
|
| 1904 |
+
if len(data) < len(PROXY_MAGIC) + SALT_SIZE + NONCE_SIZE + 16:
|
| 1905 |
+
raise ValueError("invalid encrypted payload")
|
| 1906 |
+
if not data.startswith(PROXY_MAGIC):
|
| 1907 |
+
raise ValueError("invalid encrypted payload")
|
| 1908 |
+
salt_start = len(PROXY_MAGIC)
|
| 1909 |
+
nonce_start = salt_start + SALT_SIZE
|
| 1910 |
+
data_start = nonce_start + NONCE_SIZE
|
| 1911 |
+
return AESGCM(proxy_key(data[salt_start:nonce_start])).decrypt(
|
| 1912 |
+
data[nonce_start:data_start],
|
| 1913 |
+
data[data_start:],
|
| 1914 |
+
PROXY_MAGIC,
|
| 1915 |
+
)
|
| 1916 |
+
|
| 1917 |
+
|
| 1918 |
def save_image(image, from_api):
|
| 1919 |
path = IMAGE_DIR / datetime.now(TIMEZONE).date().isoformat()
|
| 1920 |
path.mkdir(parents=True, exist_ok=True)
|
|
|
|
| 2800 |
}
|
| 2801 |
|
| 2802 |
|
| 2803 |
+
def replace_asgi_headers(headers, replacements):
|
| 2804 |
+
names = {name for name, _ in replacements}
|
| 2805 |
+
return [item for item in headers if item[0].lower() not in names] + replacements
|
| 2806 |
+
|
| 2807 |
+
|
| 2808 |
+
class GradioEncryptionMiddleware:
|
| 2809 |
+
def __init__(self, app):
|
| 2810 |
+
self.app = app
|
| 2811 |
+
|
| 2812 |
+
async def __call__(self, scope, receive, send):
|
| 2813 |
+
if scope["type"] != "http":
|
| 2814 |
+
await self.app(scope, receive, send)
|
| 2815 |
+
return
|
| 2816 |
+
headers = dict(scope.get("headers", []))
|
| 2817 |
+
encrypted = (
|
| 2818 |
+
headers.get(PROXY_ENCRYPTION_HEADER) == PROXY_ENCRYPTION
|
| 2819 |
+
and "/gradio_api/call/health" not in scope.get("path", "")
|
| 2820 |
+
)
|
| 2821 |
+
if not encrypted:
|
| 2822 |
+
await self.app(scope, receive, send)
|
| 2823 |
+
return
|
| 2824 |
+
|
| 2825 |
+
decrypted_receive = receive
|
| 2826 |
+
if scope.get("method") not in {"GET", "HEAD"}:
|
| 2827 |
+
chunks = []
|
| 2828 |
+
while True:
|
| 2829 |
+
message = await receive()
|
| 2830 |
+
if message["type"] == "http.disconnect":
|
| 2831 |
+
return
|
| 2832 |
+
chunks.append(message.get("body", b""))
|
| 2833 |
+
if not message.get("more_body", False):
|
| 2834 |
+
break
|
| 2835 |
+
try:
|
| 2836 |
+
body = decrypt_proxy_payload(b"".join(chunks))
|
| 2837 |
+
except (InvalidTag, ValueError):
|
| 2838 |
+
content = b'{"detail":"Invalid encrypted payload"}'
|
| 2839 |
+
await send({
|
| 2840 |
+
"type": "http.response.start",
|
| 2841 |
+
"status": 400,
|
| 2842 |
+
"headers": [
|
| 2843 |
+
(b"content-type", b"application/json"),
|
| 2844 |
+
(b"content-length", str(len(content)).encode()),
|
| 2845 |
+
],
|
| 2846 |
+
})
|
| 2847 |
+
await send({"type": "http.response.body", "body": content})
|
| 2848 |
+
return
|
| 2849 |
+
|
| 2850 |
+
scope = dict(scope)
|
| 2851 |
+
scope["headers"] = replace_asgi_headers(
|
| 2852 |
+
scope.get("headers", []),
|
| 2853 |
+
[
|
| 2854 |
+
(b"content-type", b"application/json"),
|
| 2855 |
+
(b"content-length", str(len(body)).encode()),
|
| 2856 |
+
],
|
| 2857 |
+
)
|
| 2858 |
+
delivered = False
|
| 2859 |
+
|
| 2860 |
+
async def decrypted_receive():
|
| 2861 |
+
nonlocal delivered
|
| 2862 |
+
if delivered:
|
| 2863 |
+
return {"type": "http.request", "body": b"", "more_body": False}
|
| 2864 |
+
delivered = True
|
| 2865 |
+
return {"type": "http.request", "body": body, "more_body": False}
|
| 2866 |
+
|
| 2867 |
+
start = None
|
| 2868 |
+
response_chunks = []
|
| 2869 |
+
|
| 2870 |
+
async def encrypted_send(message):
|
| 2871 |
+
nonlocal start
|
| 2872 |
+
if message["type"] == "http.response.start":
|
| 2873 |
+
start = message
|
| 2874 |
+
return
|
| 2875 |
+
if message["type"] == "http.response.pathsend":
|
| 2876 |
+
await send(start)
|
| 2877 |
+
await send(message)
|
| 2878 |
+
return
|
| 2879 |
+
if message["type"] != "http.response.body":
|
| 2880 |
+
await send(message)
|
| 2881 |
+
return
|
| 2882 |
+
response_chunks.append(message.get("body", b""))
|
| 2883 |
+
if message.get("more_body", False):
|
| 2884 |
+
return
|
| 2885 |
+
content = b"".join(response_chunks)
|
| 2886 |
+
if not content.startswith(FILE_MAGIC):
|
| 2887 |
+
content = encrypt_proxy_payload(content)
|
| 2888 |
+
start = dict(start)
|
| 2889 |
+
start["headers"] = replace_asgi_headers(
|
| 2890 |
+
start.get("headers", []),
|
| 2891 |
+
[
|
| 2892 |
+
(PROXY_ENCRYPTION_HEADER, PROXY_ENCRYPTION),
|
| 2893 |
+
(b"content-length", str(len(content)).encode()),
|
| 2894 |
+
],
|
| 2895 |
+
)
|
| 2896 |
+
await send(start)
|
| 2897 |
+
await send({"type": "http.response.body", "body": content})
|
| 2898 |
+
|
| 2899 |
+
await self.app(scope, decrypted_receive, encrypted_send)
|
| 2900 |
+
|
| 2901 |
+
|
| 2902 |
api = App()
|
| 2903 |
+
api.add_middleware(GradioEncryptionMiddleware)
|
| 2904 |
api.add_middleware(
|
| 2905 |
CORSMiddleware,
|
| 2906 |
allow_origins=["*"],
|