refac
This commit is contained in:
@@ -23,10 +23,8 @@ log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
STREAMING_CONTENT_TYPES = ("application/octet-stream", "image/", "application/pdf")
|
||||
STRIPPED_RESPONSE_HEADERS = frozenset(
|
||||
("transfer-encoding", "connection", "content-encoding", "content-length")
|
||||
)
|
||||
STREAMING_CONTENT_TYPES = ('application/octet-stream', 'image/', 'application/pdf')
|
||||
STRIPPED_RESPONSE_HEADERS = frozenset(('transfer-encoding', 'connection', 'content-encoding', 'content-length'))
|
||||
|
||||
|
||||
def _sanitize_proxy_path(path: str) -> str | None:
|
||||
@@ -37,14 +35,14 @@ def _sanitize_proxy_path(path: str) -> str | None:
|
||||
decoded = unquote(path)
|
||||
normalized = posixpath.normpath(decoded)
|
||||
# Remove any leading slashes that would reset the base
|
||||
cleaned = normalized.lstrip("/")
|
||||
cleaned = normalized.lstrip('/')
|
||||
# Reject if normpath resolved to parent traversal or current-dir only
|
||||
if cleaned.startswith("..") or cleaned == ".":
|
||||
if cleaned.startswith('..') or cleaned == '.':
|
||||
return None
|
||||
return cleaned
|
||||
|
||||
|
||||
@router.get("/")
|
||||
@router.get('/')
|
||||
async def list_terminal_servers(request: Request, user=Depends(get_verified_user)):
|
||||
"""Return terminal servers the authenticated user has access to."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
@@ -52,20 +50,19 @@ async def list_terminal_servers(request: Request, user=Depends(get_verified_user
|
||||
|
||||
return [
|
||||
{
|
||||
"id": connection.get("id", ""),
|
||||
"url": connection.get("url", ""),
|
||||
"name": connection.get("name", ""),
|
||||
'id': connection.get('id', ''),
|
||||
'url': connection.get('url', ''),
|
||||
'name': connection.get('name', ''),
|
||||
}
|
||||
for connection in connections
|
||||
if connection.get("enabled", True)
|
||||
and has_connection_access(user, connection, user_group_ids)
|
||||
if connection.get('enabled', True) and has_connection_access(user, connection, user_group_ids)
|
||||
]
|
||||
|
||||
|
||||
PROXY_METHODS = ["GET", "POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"]
|
||||
PROXY_METHODS = ['GET', 'POST', 'PUT', 'PATCH', 'DELETE', 'HEAD', 'OPTIONS']
|
||||
|
||||
|
||||
@router.api_route("/{server_id}/{path:path}", methods=PROXY_METHODS)
|
||||
@router.api_route('/{server_id}/{path:path}', methods=PROXY_METHODS)
|
||||
async def proxy_terminal(
|
||||
server_id: str,
|
||||
path: str,
|
||||
@@ -74,56 +71,52 @@ async def proxy_terminal(
|
||||
):
|
||||
"""Proxy a request to the admin terminal server identified by *server_id*."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connection = next((c for c in connections if c.get("id") == server_id), None)
|
||||
connection = next((c for c in connections if c.get('id') == server_id), None)
|
||||
|
||||
if connection is None:
|
||||
return JSONResponse(
|
||||
{"error": f"Terminal server '{server_id}' not found"}, status_code=404
|
||||
)
|
||||
return JSONResponse({'error': f"Terminal server '{server_id}' not found"}, status_code=404)
|
||||
|
||||
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
|
||||
if not has_connection_access(user, connection, user_group_ids):
|
||||
return JSONResponse({"error": "Access denied"}, status_code=403)
|
||||
return JSONResponse({'error': 'Access denied'}, status_code=403)
|
||||
|
||||
base_url = (connection.get("url") or "").rstrip("/")
|
||||
base_url = (connection.get('url') or '').rstrip('/')
|
||||
if not base_url:
|
||||
return JSONResponse(
|
||||
{"error": "Terminal server URL not configured"}, status_code=503
|
||||
)
|
||||
return JSONResponse({'error': 'Terminal server URL not configured'}, status_code=503)
|
||||
|
||||
safe_path = _sanitize_proxy_path(path)
|
||||
if safe_path is None:
|
||||
return JSONResponse({"error": "Invalid path"}, status_code=400)
|
||||
return JSONResponse({'error': 'Invalid path'}, status_code=400)
|
||||
|
||||
target_url = f"{base_url}/{safe_path}"
|
||||
target_url = f'{base_url}/{safe_path}'
|
||||
|
||||
# Route through orchestrator policy endpoint if policy_id is set
|
||||
policy_id = connection.get("policy_id")
|
||||
policy_id = connection.get('policy_id')
|
||||
if policy_id:
|
||||
target_url = f"{base_url}/p/{policy_id}/{safe_path}"
|
||||
target_url = f'{base_url}/p/{policy_id}/{safe_path}'
|
||||
|
||||
if request.query_params:
|
||||
target_url += f"?{request.query_params}"
|
||||
target_url += f'?{request.query_params}'
|
||||
|
||||
headers = {"X-User-Id": user.id}
|
||||
headers = {'X-User-Id': user.id}
|
||||
cookies = {}
|
||||
auth_type = connection.get("auth_type", "bearer")
|
||||
auth_type = connection.get('auth_type', 'bearer')
|
||||
|
||||
if auth_type == "bearer":
|
||||
headers["Authorization"] = f"Bearer {connection.get('key', '')}"
|
||||
elif auth_type == "session":
|
||||
if auth_type == 'bearer':
|
||||
headers['Authorization'] = f'Bearer {connection.get("key", "")}'
|
||||
elif auth_type == 'session':
|
||||
cookies = request.cookies
|
||||
headers["Authorization"] = f"Bearer {request.state.token.credentials}"
|
||||
elif auth_type == "system_oauth":
|
||||
headers['Authorization'] = f'Bearer {request.state.token.credentials}'
|
||||
elif auth_type == 'system_oauth':
|
||||
cookies = request.cookies
|
||||
oauth_token = request.headers.get("x-oauth-access-token", "")
|
||||
oauth_token = request.headers.get('x-oauth-access-token', '')
|
||||
if oauth_token:
|
||||
headers["Authorization"] = f"Bearer {oauth_token}"
|
||||
headers['Authorization'] = f'Bearer {oauth_token}'
|
||||
# auth_type == "none": no Authorization header
|
||||
|
||||
content_type = request.headers.get("content-type")
|
||||
content_type = request.headers.get('content-type')
|
||||
if content_type:
|
||||
headers["Content-Type"] = content_type
|
||||
headers['Content-Type'] = content_type
|
||||
|
||||
body = await request.body()
|
||||
session = aiohttp.ClientSession(
|
||||
@@ -140,7 +133,7 @@ async def proxy_terminal(
|
||||
data=body or None,
|
||||
)
|
||||
|
||||
upstream_content_type = upstream_response.headers.get("content-type", "")
|
||||
upstream_content_type = upstream_response.headers.get('content-type', '')
|
||||
filtered_headers = {
|
||||
key: value
|
||||
for key, value in upstream_response.headers.items()
|
||||
@@ -167,16 +160,12 @@ async def proxy_terminal(
|
||||
await upstream_response.release()
|
||||
await session.close()
|
||||
|
||||
return Response(
|
||||
content=response_body, status_code=status_code, headers=filtered_headers
|
||||
)
|
||||
return Response(content=response_body, status_code=status_code, headers=filtered_headers)
|
||||
|
||||
except Exception as error:
|
||||
await session.close()
|
||||
log.exception("Terminal proxy error: %s", error)
|
||||
return JSONResponse(
|
||||
{"error": f"Terminal proxy error: {error}"}, status_code=502
|
||||
)
|
||||
log.exception('Terminal proxy error: %s', error)
|
||||
return JSONResponse({'error': f'Terminal proxy error: {error}'}, status_code=502)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -201,42 +190,42 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str):
|
||||
try:
|
||||
raw = await asyncio.wait_for(ws.receive_text(), timeout=10.0)
|
||||
payload = json.loads(raw)
|
||||
if payload.get("type") != "auth":
|
||||
await ws.close(code=4001, reason="Expected auth message")
|
||||
if payload.get('type') != 'auth':
|
||||
await ws.close(code=4001, reason='Expected auth message')
|
||||
return None
|
||||
token = payload.get("token", "")
|
||||
token = payload.get('token', '')
|
||||
data = decode_token(token)
|
||||
if data is None or "id" not in data:
|
||||
await ws.close(code=4001, reason="Invalid token")
|
||||
if data is None or 'id' not in data:
|
||||
await ws.close(code=4001, reason='Invalid token')
|
||||
return None
|
||||
user = Users.get_user_by_id(data["id"])
|
||||
user = Users.get_user_by_id(data['id'])
|
||||
if user is None:
|
||||
await ws.close(code=4001, reason="User not found")
|
||||
await ws.close(code=4001, reason='User not found')
|
||||
return None
|
||||
except (asyncio.TimeoutError, json.JSONDecodeError):
|
||||
await ws.close(code=4001, reason="Auth timeout or invalid payload")
|
||||
await ws.close(code=4001, reason='Auth timeout or invalid payload')
|
||||
return None
|
||||
except Exception:
|
||||
await ws.close(code=4001, reason="Invalid token")
|
||||
await ws.close(code=4001, reason='Invalid token')
|
||||
return None
|
||||
|
||||
# Resolve terminal server
|
||||
connections = ws.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connection = next((c for c in connections if c.get("id") == server_id), None)
|
||||
connection = next((c for c in connections if c.get('id') == server_id), None)
|
||||
|
||||
if connection is None:
|
||||
await ws.close(code=4004, reason="Terminal server not found")
|
||||
await ws.close(code=4004, reason='Terminal server not found')
|
||||
return None
|
||||
|
||||
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
|
||||
if not has_connection_access(user, connection, user_group_ids):
|
||||
await ws.close(code=4003, reason="Access denied")
|
||||
await ws.close(code=4003, reason='Access denied')
|
||||
return None
|
||||
|
||||
return user, connection
|
||||
|
||||
|
||||
@router.websocket("/{server_id}/api/terminals/{session_id}")
|
||||
@router.websocket('/{server_id}/api/terminals/{session_id}')
|
||||
async def ws_terminal(
|
||||
ws: WebSocket,
|
||||
server_id: str,
|
||||
@@ -255,28 +244,28 @@ async def ws_terminal(
|
||||
return
|
||||
user, connection = result
|
||||
|
||||
base_url = (connection.get("url") or "").rstrip("/")
|
||||
base_url = (connection.get('url') or '').rstrip('/')
|
||||
if not base_url:
|
||||
await ws.close(code=4003, reason="Terminal server URL not configured")
|
||||
await ws.close(code=4003, reason='Terminal server URL not configured')
|
||||
return
|
||||
|
||||
# Build upstream WebSocket URL (no token in URL)
|
||||
ws_base = base_url.replace("https://", "wss://").replace("http://", "ws://")
|
||||
ws_base = base_url.replace('https://', 'wss://').replace('http://', 'ws://')
|
||||
|
||||
# Route through orchestrator policy endpoint if policy_id is set
|
||||
policy_id = connection.get("policy_id")
|
||||
policy_id = connection.get('policy_id')
|
||||
upstream_params = {}
|
||||
# For orchestrator-backed servers, pass user_id
|
||||
upstream_params["user_id"] = user.id
|
||||
upstream_params['user_id'] = user.id
|
||||
|
||||
import urllib.parse
|
||||
|
||||
if policy_id:
|
||||
upstream_url = f"{ws_base}/p/{policy_id}/api/terminals/{session_id}"
|
||||
upstream_url = f'{ws_base}/p/{policy_id}/api/terminals/{session_id}'
|
||||
else:
|
||||
upstream_url = f"{ws_base}/api/terminals/{session_id}"
|
||||
upstream_url = f'{ws_base}/api/terminals/{session_id}'
|
||||
if upstream_params:
|
||||
upstream_url += f"?{urllib.parse.urlencode(upstream_params)}"
|
||||
upstream_url += f'?{urllib.parse.urlencode(upstream_params)}'
|
||||
|
||||
session = aiohttp.ClientSession()
|
||||
try:
|
||||
@@ -285,22 +274,22 @@ async def ws_terminal(
|
||||
import json as _json
|
||||
|
||||
# First-message auth to upstream terminal server
|
||||
auth_type = connection.get("auth_type", "bearer")
|
||||
if auth_type == "bearer":
|
||||
key = connection.get("key", "")
|
||||
await upstream.send_str(_json.dumps({"type": "auth", "token": key}))
|
||||
auth_type = connection.get('auth_type', 'bearer')
|
||||
if auth_type == 'bearer':
|
||||
key = connection.get('key', '')
|
||||
await upstream.send_str(_json.dumps({'type': 'auth', 'token': key}))
|
||||
|
||||
async def _client_to_upstream():
|
||||
"""Forward client → upstream."""
|
||||
try:
|
||||
while True:
|
||||
msg = await ws.receive()
|
||||
if msg["type"] == "websocket.disconnect":
|
||||
if msg['type'] == 'websocket.disconnect':
|
||||
break
|
||||
elif "bytes" in msg and msg["bytes"]:
|
||||
await upstream.send_bytes(msg["bytes"])
|
||||
elif "text" in msg and msg["text"]:
|
||||
await upstream.send_str(msg["text"])
|
||||
elif 'bytes' in msg and msg['bytes']:
|
||||
await upstream.send_bytes(msg['bytes'])
|
||||
elif 'text' in msg and msg['text']:
|
||||
await upstream.send_str(msg['text'])
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -326,7 +315,7 @@ async def ws_terminal(
|
||||
return_exceptions=True,
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception("Terminal WebSocket proxy error: %s", e)
|
||||
log.exception('Terminal WebSocket proxy error: %s', e)
|
||||
finally:
|
||||
await session.close()
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user