Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- #!/usr/bin/env python3
- import json
- import sys
- from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
- from urllib import error as urlerror
- from urllib import request as urlrequest
- from urllib.parse import urlsplit, urlunsplit
- LISTEN_HOST = "0.0.0.0"
- LISTEN_PORT = 1234
- TARGET = "http://127.0.0.1:1235"
- TEMPERATURE = 0.6
- TOP_P = 0.95
- TOP_K = 20
- MIN_P = 0.0
- REASONING_EFFORT = "low"
- CHAT_TEMPLATE_KWARGS = '{"preserve_thinking": true, "reasoning_effort": "low"}'
- UPSTREAM_TIMEOUT = None
- HOP_BY_HOP_HEADERS = {
- "connection",
- "keep-alive",
- "proxy-authenticate",
- "proxy-authorization",
- "proxy-connection",
- "te",
- "trailer",
- "transfer-encoding",
- "upgrade",
- }
- class ProxyHandler(BaseHTTPRequestHandler):
- protocol_version = "HTTP/1.1"
- server_version = "vllm-override-proxy/0.1"
- def do_GET(self):
- self._proxy()
- def do_POST(self):
- self._proxy()
- def do_PUT(self):
- self._proxy()
- def do_PATCH(self):
- self._proxy()
- def do_DELETE(self):
- self._proxy()
- def do_OPTIONS(self):
- self._proxy()
- def _proxy(self):
- body = self._read_request_body()
- body = self._override_json_body(body)
- upstream_url = self._upstream_url()
- headers = self._upstream_headers(body)
- req = urlrequest.Request(
- upstream_url,
- data=body,
- headers=headers,
- method=self.command,
- )
- try:
- with urlrequest.urlopen(req, timeout=self.server.upstream_timeout) as resp:
- self._send_upstream_response(resp)
- except urlerror.HTTPError as exc:
- self._send_upstream_response(exc)
- except urlerror.URLError as exc:
- self.send_error(502, "bad gateway: %s" % exc.reason)
- def _log_request(self, body):
- if body:
- try:
- payload = json.dumps(json.loads(body.decode("utf-8")), indent=2)
- except (UnicodeDecodeError, json.JSONDecodeError):
- payload = body.decode("utf-8", "replace")
- print(">>> %s %s\n%s" % (self.command, self.path, payload), flush=True)
- else:
- print(">>> %s %s" % (self.command, self.path), flush=True)
- def _read_request_body(self):
- length = self.headers.get("Content-Length")
- if length is None:
- return None
- try:
- length_int = int(length)
- except ValueError:
- return None
- if length_int <= 0:
- return b""
- return self.rfile.read(length_int)
- def _override_json_body(self, body):
- if not body:
- return body
- content_type = self.headers.get("Content-Type", "").lower()
- looks_like_json = body.lstrip().startswith((b"{", b"["))
- if "json" not in content_type and not looks_like_json:
- return body
- try:
- payload = json.loads(body.decode("utf-8"))
- except (UnicodeDecodeError, json.JSONDecodeError):
- return body
- if not isinstance(payload, dict):
- return body
- payload.pop("temp", None)
- payload.pop("top-k", None)
- payload.pop("top-p", None)
- payload.pop("min-p", None)
- payload["temperature"] = self.server.temperature
- payload["top_k"] = self.server.top_k
- payload["top_p"] = self.server.top_p
- payload["min_p"] = self.server.min_p
- if self.server.reasoning_effort:
- payload.pop("reasoning_effort", None)
- payload["reasoning_effort"] = self.server.reasoning_effort
- if self.server.chat_template_kwargs:
- payload.pop("chat_template_kwargs", None)
- value = self.server.chat_template_kwargs
- try:
- value = json.loads(value)
- except ValueError:
- pass
- payload["chat_template_kwargs"] = value
- return json.dumps(payload, separators=(",", ":")).encode("utf-8")
- def _upstream_url(self):
- path = self.path
- if path.startswith(("http://", "https://")):
- parts = urlsplit(path)
- path = urlunsplit(("", "", parts.path or "/", parts.query, ""))
- if not path.startswith("/"):
- path = "/" + path
- return self.server.target.rstrip("/") + path
- def _upstream_headers(self, body):
- headers = {}
- for key, value in self.headers.items():
- lower = key.lower()
- if lower in HOP_BY_HOP_HEADERS or lower in {"host", "content-length"}:
- continue
- headers[key] = value
- headers["Host"] = urlsplit(self.server.target).netloc
- headers["Accept-Encoding"] = "identity"
- if body is not None:
- headers["Content-Length"] = str(len(body))
- return headers
- def _send_upstream_response(self, resp):
- status = getattr(resp, "status", getattr(resp, "code", 502))
- reason = getattr(resp, "reason", None)
- self.send_response(status, reason)
- for key, value in resp.headers.items():
- lower = key.lower()
- if lower in HOP_BY_HOP_HEADERS or lower == "content-length":
- continue
- self.send_header(key, value)
- self.send_header("Connection", "close")
- self.end_headers()
- reader = getattr(resp, "read1", resp.read)
- while True:
- chunk = reader(64 * 1024)
- if not chunk:
- break
- self.wfile.write(chunk)
- self.wfile.flush()
- self.close_connection = True
- def log_message(self, fmt, *args):
- sys.stderr.write("%s - - [%s] %s\n" % (self.address_string(), self.log_date_time_string(), fmt % args))
- class OverrideProxyServer(ThreadingHTTPServer):
- daemon_threads = True
- def __init__(self, server_address, handler_class):
- super().__init__(server_address, handler_class)
- self.target = TARGET
- self.temperature = TEMPERATURE
- self.top_k = TOP_K
- self.top_p = TOP_P
- self.min_p = MIN_P
- self.reasoning_effort = REASONING_EFFORT or ""
- self.chat_template_kwargs = CHAT_TEMPLATE_KWARGS or ""
- self.upstream_timeout = UPSTREAM_TIMEOUT
- def main():
- server = OverrideProxyServer((LISTEN_HOST, LISTEN_PORT), ProxyHandler)
- print(
- "proxy listening on http://%s:%d -> %s, temperature=%s, top_k=%s, top_p=%s, min_p=%s, reasoning_effort=%s, chat_template_kwargs=%s"
- % (LISTEN_HOST, LISTEN_PORT, TARGET, TEMPERATURE, TOP_K, TOP_P, MIN_P, REASONING_EFFORT or "-", CHAT_TEMPLATE_KWARGS or "-"),
- flush=True,
- )
- server.serve_forever()
- if __name__ == "__main__":
- main()
Advertisement
Add Comment
Please, Sign In to add comment