q6 commited on
Commit
4f1440f
·
1 Parent(s): 4122d52

Decrypt-proxied-gradio-payloads

Browse files
Files changed (1) hide show
  1. app.py +141 -0
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=["*"],