# retoor <retoor@molodetz.nl>
from __future__ import annotations
import asyncio
import logging
import httpx
import websockets
from fastapi import Request, WebSocket
from starlette.responses import Response
logger = logging.getLogger(__name__)
HOP_HEADERS = {
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailers",
"transfer-encoding",
"upgrade",
"host",
"content-length",
"content-encoding",
}
METHODS = ["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"]
DEFAULT_TIMEOUT = 300.0
def forward_headers(request: Request, prefix: str = "") -> dict:
headers = {k: v for k, v in request.headers.items() if k.lower() not in HOP_HEADERS}
if prefix:
headers["X-Forwarded-Prefix"] = prefix
headers["X-Script-Name"] = prefix
headers["X-Forwarded-Host"] = request.headers.get(
"host", request.url.hostname or ""
)
headers["X-Forwarded-Proto"] = request.headers.get(
"x-forwarded-proto", request.url.scheme
)
headers["Accept-Encoding"] = "identity"
return headers
def inject_base(body: bytes, prefix: str) -> bytes:
lowered = body.lower()
if b"<base" in lowered:
return body
tag = f'<base href="{prefix}/">'.encode()
head = lowered.find(b"<head")
anchor = (
lowered.find(b">", head)
if head != -1
else lowered.find(b">", lowered.find(b"<html"))
)
if anchor == -1:
return tag + body
return body[: anchor + 1] + tag + body[anchor + 1 :]
def rewrite_location(value: str, prefix: str) -> str:
if not prefix:
return value
if value.startswith("/") and not value.startswith("//"):
return prefix + value
return value
async def proxy_http(
request: Request,
host: str,
port: int,
path: str,
*,
prefix: str = "",
timeout: float = DEFAULT_TIMEOUT,
rewrite_html: bool = True,
) -> Response:
url = f"http://{host}:{port}/{path}"
headers = forward_headers(request, prefix)
body = await request.body()
try:
async with httpx.AsyncClient(
timeout=timeout, follow_redirects=False
) as client:
upstream = await client.request(
request.method,
url,
params=request.query_params,
headers=headers,
content=body,
)
except httpx.HTTPError as error:
return Response(f"upstream error: {error}", status_code=502)
out_headers = {
k: v
for k, v in upstream.headers.items()
if k.lower() not in HOP_HEADERS and k.lower() != "set-cookie"
}
if "location" in out_headers:
out_headers["location"] = rewrite_location(out_headers["location"], prefix)
content_type = upstream.headers.get("content-type", "")
content = upstream.content
if prefix and rewrite_html and "text/html" in content_type.lower():
content = inject_base(content, prefix)
response = Response(
content=content,
status_code=upstream.status_code,
headers=out_headers,
media_type=content_type or None,
)
for cookie in upstream.headers.get_list("set-cookie"):
response.headers.append("set-cookie", cookie)
return response
async def proxy_ws(
websocket: WebSocket, host: str, port: int, path: str, *, accepted: bool = False
) -> None:
upstream_url = f"ws://{host}:{port}/{path}"
if websocket.url.query:
upstream_url += f"?{websocket.url.query}"
if not accepted:
await websocket.accept()
try:
async with websockets.connect(
upstream_url, open_timeout=10, max_size=None
) as upstream:
await pump(websocket, upstream)
except Exception as error:
logger.debug("ws proxy to %s:%s failed: %s", host, port, error)
try:
await websocket.close(code=1011)
except Exception:
pass
async def pump(client_ws: WebSocket, upstream) -> None:
async def client_to_upstream():
try:
while True:
message = await client_ws.receive()
if message["type"] == "websocket.disconnect":
break
if message.get("text") is not None:
await upstream.send(message["text"])
elif message.get("bytes") is not None:
await upstream.send(message["bytes"])
except Exception:
pass
finally:
await upstream.close()
async def upstream_to_client():
try:
async for message in upstream:
if isinstance(message, (bytes, bytearray)):
await client_ws.send_bytes(bytes(message))
else:
await client_ws.send_text(message)
except Exception:
pass
finally:
try:
await client_ws.close()
except Exception:
pass
await asyncio.gather(client_to_upstream(), upstream_to_client())