forked from retoor/devplacepy
update
This commit is contained in:
@@ -0,0 +1,216 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
import uuid_utils
|
||||
import websockets
|
||||
from starlette.datastructures import Headers
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.routing import Route
|
||||
from starlette.websockets import WebSocket
|
||||
|
||||
from devplacepy_services.base.config import service_url
|
||||
from devplacepy_services.base.errors import error_response
|
||||
from devplacepy_services.base.manifest import INGRESS_ROUTES as _INGRESS_SPECS
|
||||
from devplacepy_services.base.proxy import HOP_HEADERS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
INGRESS_ROUTES = [(route.prefix, route.service) for route in _INGRESS_SPECS]
|
||||
|
||||
TIMEOUTS = {
|
||||
"xmlrpc": 120.0,
|
||||
}
|
||||
|
||||
|
||||
def _forward_headers_from_scope(scope) -> dict[str, str]:
|
||||
original = Headers(scope=scope)
|
||||
headers = {
|
||||
k: v for k, v in original.items() if k.lower() not in HOP_HEADERS
|
||||
}
|
||||
headers["X-Request-Id"] = headers.get("X-Request-Id") or uuid_utils.uuid7().hex
|
||||
headers["Accept-Encoding"] = "identity"
|
||||
# Mirror nginx's `proxy_set_header Host $host` (Appendix F) so an
|
||||
# upstream-generated absolute URL (redirect, url_for) reflects the
|
||||
# public :10500 endpoint the browser is actually talking to, not the
|
||||
# upstream service's own internal bind address/port.
|
||||
original_host = original.get("host")
|
||||
if original_host:
|
||||
headers["Host"] = original_host
|
||||
return headers
|
||||
|
||||
|
||||
def _forward_headers(request: Request) -> dict[str, str]:
|
||||
return _forward_headers_from_scope(request.scope)
|
||||
|
||||
|
||||
_WS_HANDSHAKE_HEADERS = frozenset(
|
||||
{
|
||||
"sec-websocket-key",
|
||||
"sec-websocket-version",
|
||||
"sec-websocket-extensions",
|
||||
"sec-websocket-protocol",
|
||||
"sec-websocket-accept",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _forward_ws_headers(scope) -> dict[str, str]:
|
||||
headers = _forward_headers_from_scope(scope)
|
||||
for key in list(headers):
|
||||
if key.lower() in _WS_HANDSHAKE_HEADERS:
|
||||
del headers[key]
|
||||
return headers
|
||||
|
||||
|
||||
def _target_path(scope, prefix: str) -> str:
|
||||
# Starlette's Mount rewrites scope["root_path"] to the cumulative mount
|
||||
# prefix but leaves scope["path"] as the FULL original request path (it
|
||||
# does not strip the prefix) - so the full path alone is already the
|
||||
# correct upstream path; concatenating root_path in front double-prefixes it.
|
||||
return scope.get("path", "") or prefix
|
||||
|
||||
|
||||
def _upstream_http_url(service: str, path: str, query: str) -> str:
|
||||
base = service_url(service).rstrip("/")
|
||||
url = f"{base}{path}"
|
||||
if query:
|
||||
url = f"{url}?{query}"
|
||||
return url
|
||||
|
||||
|
||||
def _upstream_ws_url(service: str, path: str, query: str) -> str:
|
||||
parsed = urlparse(service_url(service))
|
||||
scheme = "wss" if parsed.scheme == "https" else "ws"
|
||||
netloc = parsed.netloc
|
||||
url = f"{scheme}://{netloc}{path}"
|
||||
if query:
|
||||
url = f"{url}?{query}"
|
||||
return url
|
||||
|
||||
|
||||
class IngressProxy:
|
||||
def __init__(self, prefix: str, service: str) -> None:
|
||||
self.prefix = prefix
|
||||
self.service = service
|
||||
self.timeout = TIMEOUTS.get(service, 30.0)
|
||||
|
||||
async def __call__(self, scope, receive, send) -> None:
|
||||
if scope["type"] == "http":
|
||||
await self._proxy_http(scope, receive, send)
|
||||
elif scope["type"] == "websocket":
|
||||
await self._proxy_ws(scope, receive, send)
|
||||
|
||||
async def _proxy_http(self, scope, receive, send) -> None:
|
||||
request = Request(scope, receive)
|
||||
path = _target_path(scope, self.prefix)
|
||||
query = scope.get("query_string", b"").decode()
|
||||
url = _upstream_http_url(self.service, path, query)
|
||||
body = await request.body()
|
||||
headers = _forward_headers(request)
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.timeout, follow_redirects=False
|
||||
) as client:
|
||||
upstream = await client.request(
|
||||
request.method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=body,
|
||||
)
|
||||
except httpx.HTTPError as exc:
|
||||
logger.warning("ingress %s upstream error: %s", self.prefix, exc)
|
||||
response = error_response(
|
||||
502, "Upstream service unavailable", "upstream_error"
|
||||
)
|
||||
await response(scope, receive, send)
|
||||
return
|
||||
out_headers = {
|
||||
k: v
|
||||
for k, v in upstream.headers.items()
|
||||
if k.lower() not in HOP_HEADERS and k.lower() != "set-cookie"
|
||||
}
|
||||
response = Response(
|
||||
content=upstream.content,
|
||||
status_code=upstream.status_code,
|
||||
headers=out_headers,
|
||||
media_type=upstream.headers.get("content-type"),
|
||||
)
|
||||
for cookie in upstream.headers.get_list("set-cookie"):
|
||||
response.headers.append("set-cookie", cookie)
|
||||
await response(scope, receive, send)
|
||||
|
||||
async def _proxy_ws(self, scope, receive, send) -> None:
|
||||
client_ws = WebSocket(scope, receive, send)
|
||||
path = _target_path(scope, self.prefix)
|
||||
query = scope.get("query_string", b"").decode()
|
||||
upstream_url = _upstream_ws_url(self.service, path, query)
|
||||
headers = _forward_ws_headers(scope)
|
||||
await client_ws.accept()
|
||||
try:
|
||||
async with websockets.connect(
|
||||
upstream_url,
|
||||
open_timeout=10,
|
||||
max_size=None,
|
||||
additional_headers=headers,
|
||||
) as upstream:
|
||||
await _pump(client_ws, upstream)
|
||||
except Exception as exc:
|
||||
logger.debug("ingress ws %s failed: %s", self.prefix, exc)
|
||||
try:
|
||||
await client_ws.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())
|
||||
|
||||
|
||||
def mount_ingress(app) -> None:
|
||||
for prefix, service in INGRESS_ROUTES:
|
||||
proxy = IngressProxy(prefix, service)
|
||||
# A bare hit on the prefix itself (no trailing slash, nothing after -
|
||||
# e.g. a POST to /xmlrpc or GET /tools) never matches Mount's own
|
||||
# path regex (it requires a "/" plus content after the prefix), so
|
||||
# Starlette's router-level redirect_slashes fallback 307s it to
|
||||
# "<prefix>/" before the Mount ever sees it. If the upstream service's
|
||||
# own router registers an exact route at its mount root (as /tools
|
||||
# and /xmlrpc both do), THAT redirects back to the bare prefix -
|
||||
# an infinite loop between the two opposite trailing-slash
|
||||
# conventions. Registering an explicit Route at the exact prefix
|
||||
# bypasses Mount's regex/redirect fallback entirely for that one path.
|
||||
app.router.routes.append(Route(prefix, endpoint=proxy, methods=None))
|
||||
app.mount(prefix, proxy)
|
||||
Reference in New Issue
Block a user