feat(ws): implement full duplex bi-directional WebSocket proxy bridge supporting subdomains and path-based routing

This commit is contained in:
bolbol
2026-08-20 15:04:12 +01:00
parent 442e128450
commit 23168a13e8
3 changed files with 134 additions and 4 deletions
+34 -3
View File
@@ -1,4 +1,4 @@
from fastapi import FastAPI, Request
from fastapi import FastAPI, Request, WebSocket, WebSocketDisconnect
from fastapi.responses import HTMLResponse
from fastapi.staticfiles import StaticFiles
from fastapi.templating import Jinja2Templates
@@ -8,7 +8,7 @@ from contextlib import asynccontextmanager
from app.config import settings, UNIVERSES, BASE_DIR
from app.routers import api, proxy
from app.proxy import close_http_client, proxy_request
from app.proxy import close_http_client, proxy_request, proxy_websocket
from app.personas import ensure_dirs
def get_subdomain_target(host: str) -> Optional[Tuple[str, str]]:
@@ -55,7 +55,7 @@ app = FastAPI(
lifespan=lifespan
)
# Subdomain Routing Middleware
# Subdomain Routing Middleware for HTTP
@app.middleware("http")
async def subdomain_routing_middleware(request: Request, call_next):
host = request.headers.get("host", "")
@@ -96,6 +96,37 @@ async def index_view(request: Request):
}
)
# WebSocket Proxy Route (intercepts any WebSocket connection across subdomains)
@app.websocket("/{path:path}")
@app.websocket("")
async def websocket_proxy_endpoint(websocket: WebSocket, path: str = ""):
host = websocket.headers.get("host", "")
target = get_subdomain_target(host)
if target:
backend_url, universe_id = target
await proxy_websocket(
client_ws=websocket,
backend_url=backend_url,
path=path,
universe_id=universe_id
)
elif path.startswith("u/"):
parts = path.split("/", 2)
if len(parts) >= 2 and parts[1] in UNIVERSES:
u_id = parts[1]
subpath = parts[2] if len(parts) > 2 else ""
await proxy_websocket(
client_ws=websocket,
backend_url=UNIVERSES[u_id].backend_url,
path=subpath,
universe_id=u_id
)
else:
await websocket.close(code=1008)
else:
await websocket.close(code=1008)
if __name__ == "__main__":
import uvicorn
uvicorn.run("app.main:app", host=settings.host, port=settings.port, reload=settings.debug)
+99 -1
View File
@@ -1,8 +1,11 @@
import httpx
import re
import asyncio
from typing import AsyncGenerator, Dict, Any, Optional, List, Tuple
from fastapi import Request, Response
from fastapi import Request, Response, WebSocket, WebSocketDisconnect
from fastapi.responses import StreamingResponse
import websockets
from websockets.exceptions import ConnectionClosed
import logging
logger = logging.getLogger("hermes_hub.proxy")
@@ -19,6 +22,16 @@ HOP_BY_HOP_HEADERS = {
"content-length"
}
WS_HOP_BY_HOP_HEADERS = {
"connection",
"upgrade",
"sec-websocket-key",
"sec-websocket-version",
"sec-websocket-extensions",
"sec-websocket-accept",
"host"
}
_http_transport: Optional[httpx.AsyncHTTPTransport] = None
def get_http_transport() -> httpx.AsyncHTTPTransport:
@@ -158,3 +171,88 @@ async def proxy_request(
status_code=500,
media_type="text/plain"
)
async def proxy_websocket(
client_ws: WebSocket,
backend_url: str,
path: str,
universe_id: Optional[str] = None
):
"""
Bi-directional full duplex WebSocket proxy bridge.
Forwards connection from client browser to backend Hermes universe instance.
"""
ws_base = backend_url.replace("https://", "wss://").replace("http://", "ws://").rstrip("/")
sub_path = path.lstrip("/")
upstream_url = f"{ws_base}/{sub_path}" if sub_path else ws_base
if client_ws.url.query:
upstream_url = f"{upstream_url}?{client_ws.url.query}"
upstream_headers = {}
for k, v in client_ws.headers.items():
if k.lower() not in WS_HOP_BY_HOP_HEADERS:
upstream_headers[k] = v
upstream_headers["x-forwarded-host"] = client_ws.headers.get("host", "")
upstream_headers["x-forwarded-proto"] = client_ws.url.scheme or "http"
if client_ws.client:
upstream_headers["x-forwarded-for"] = client_ws.client.host
# Check subprotocols
subprotocols_raw = client_ws.headers.get("sec-websocket-protocol", "")
subprotocols = [s.strip() for s in subprotocols_raw.split(",") if s.strip()] or None
try:
async with websockets.connect(
upstream_url,
extra_headers=upstream_headers,
subprotocols=subprotocols,
ping_interval=20,
ping_timeout=20,
max_size=10 * 1024 * 1024
) as upstream_ws:
# Accept client connection with negotiated subprotocol
await client_ws.accept(subprotocol=upstream_ws.subprotocol)
async def client_to_upstream():
try:
while True:
msg = await client_ws.receive()
if "text" in msg:
await upstream_ws.send(msg["text"])
elif "bytes" in msg:
await upstream_ws.send(msg["bytes"])
elif msg.get("type") == "websocket.disconnect":
break
except (WebSocketDisconnect, ConnectionClosed, asyncio.CancelledError):
pass
async def upstream_to_client():
try:
async for msg in upstream_ws:
if isinstance(msg, str):
await client_ws.send_text(msg)
elif isinstance(msg, bytes):
await client_ws.send_bytes(msg)
except (WebSocketDisconnect, ConnectionClosed, asyncio.CancelledError):
pass
done, pending = await asyncio.wait(
[
asyncio.create_task(client_to_upstream()),
asyncio.create_task(upstream_to_client())
],
return_when=asyncio.FIRST_COMPLETED
)
for task in pending:
task.cancel()
except (WebSocketDisconnect, ConnectionClosed):
pass
except Exception as e:
logger.error(f"WebSocket proxy bridge error for {upstream_url}: {e}")
try:
await client_ws.close(code=1011)
except Exception:
pass