diff --git a/.github/workflows/xparl_security.yml b/.github/workflows/xparl_security.yml new file mode 100644 index 000000000..9bf527975 --- /dev/null +++ b/.github/workflows/xparl_security.yml @@ -0,0 +1,29 @@ +name: xparl security +on: + pull_request: + branches: [develop] + push: + branches: [develop] + +permissions: + contents: read + +jobs: + control-security: + runs-on: ubuntu-22.04 + strategy: + matrix: + python-version: ['3.9', '3.10'] + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + - name: Install xparl test dependencies + run: | + python -m pip install 'setuptools<81' 'numpy<2' scipy 'pyzmq==22.3.0' 'cloudpickle==1.6.0' 'grpcio>=1.48' 'protobuf==3.20.0' termcolor psutil flask flask-cors click requests pynvml six + python -m pip install --no-deps . + - name: Test control authentication and JSON metadata + env: + XPARL_igonre_core: '1' + run: python -m unittest discover -s parl/remote/tests -p control_security_test.py -v diff --git a/README.cn.md b/README.cn.md index c3cb7f2b3..8533d2509 100644 --- a/README.cn.md +++ b/README.cn.md @@ -128,15 +128,13 @@ pip install parl # xparl 安全说明 -`xparl` 提供了跨多机集群的多进程并行功能,类似于 Python 自带的单机多进程。这意味着在某个客户端上编写代码后,可以在集群内的任意机器上执行任意代码,例如获取其他机器上的数据、增删文件等。 - -这是设计的初衷,因为强化学习环境多种多样,`env_wrapper` 需要具备执行各种可能操作的能力。`xparl` 使用了 `pickle` 实现这一功能(类似于 `ray`)。与大多数情况下将 `pickle` 视为可注入代码的漏洞不同,这里使用 `pickle` 是一种特性。 +`xparl` 在可信集群机器上执行 Python 代码。启动前,所有 master、worker 和 client 必须配置相同的随机密钥 `XPARL_AUTH_TOKEN`,长度至少 32 字节。所有 ZeroMQ 连接强制使用 CURVE 认证和加密;控制消息使用经过校验的 JSON。master 和 HTTP 服务默认仅监听本机回环地址,监控和日志接口需要身份认证。 ## 安全性注意事项 -由于支持任意代码执行,用户需要确保集群环境是安全的: - -- **不要允许不信任的机器加入集群。** -- **不要让不信任的用户访问集群,例如不要将 `xparl` 的端口暴露在公网。** -- **不要在集群上执行不信任的代码。** +- 仅向可信用户和节点分发集群密钥,持有密钥即拥有在 worker 上执行代码的权限。 +- 使用隔离内网和防火墙限制所有集群端口及心跳端口的访问,禁止公网暴露。 +- 所有节点和客户端必须同时升级并重启,旧进程无法使用新的通信协议。 +- 已发布的历史安装包不会因源码合入自动修复,需升级至修复后的源码或包含补丁的新版本。 +配置、监控访问和升级步骤见[安全和升级说明](docs/xparl_security.md)。 diff --git a/README.md b/README.md index d9ccdc006..c962c0919 100644 --- a/README.md +++ b/README.md @@ -129,14 +129,13 @@ For beginners who know little about reinforcement learning, we also provide an i # xparl Security -`xparl` provides multi-process parallelism across a multi-machine cluster, similar to Python's built-in single-machine multiprocessing. This means that after writing code on a client, you can execute arbitrary code on any machine within the cluster, such as retrieving data from other machines, adding or deleting files, etc. - -This behavior is by design, as reinforcement learning environments are diverse, and `env_wrapper` needs the ability to perform any possible operation. `xparl` achieves this functionality using `pickle` (similar to `ray`). Unlike in most cases where `pickle` may be considered a vulnerability, here it is an essential feature. +`xparl` runs Python code on trusted cluster machines. Set a shared, randomly generated `XPARL_AUTH_TOKEN` of at least 32 bytes on every master, worker, and client before startup. All ZeroMQ connections require CURVE authentication and encryption; control metadata uses validated JSON. The master and HTTP services bind to loopback by default, and HTTP monitoring/log routes require credentials. ## Security Considerations -Since arbitrary code execution is possible, users must ensure the cluster is secure: +- Only trusted users and machines may hold the cluster secret or submit code. +- Keep all cluster and heartbeat ports on a private network behind a firewall. +- Upgrade all nodes and clients together; older processes cannot use the new protocol. +- Existing published packages require an upgrade to patched source or a release containing the fix. -- **Do not allow untrusted machines to join the cluster.** -- **Do not expose the `xparl` ports to the public internet or allow untrusted users to access the cluster.** -- **Do not execute untrusted code on the cluster.** +See the [security and upgrade guide](docs/xparl_security.md) for configuration, monitoring access, and migration steps. diff --git a/docs/xparl_security.md b/docs/xparl_security.md new file mode 100644 index 000000000..c9f6a3d43 --- /dev/null +++ b/docs/xparl_security.md @@ -0,0 +1,27 @@ +# xparl security and upgrade guide + +xparl executes Python code supplied by trusted cluster members. Possession of the cluster secret grants permission to execute code on workers. Use dedicated, least-privileged accounts and a private network; never expose cluster ports to the public Internet. + +## Required authentication + +Set `XPARL_AUTH_TOKEN` to the same cryptographically random secret on the master, workers, and clients **before starting their processes**. The secret must contain at least 32 bytes. Generate a fresh secret locally with `python -c 'import secrets; print(secrets.token_hex(32))'`, then distribute it through your secret management system. Do not put the secret in source control, command-line arguments, screenshots, or logs. + +Every xparl ZeroMQ connection now requires CURVE authentication and encryption. The server accepts only the client public key derived from the cluster secret; knowing the server public key does not authorize a client. Startup fails when the secret is absent, too short, or CURVE is unavailable. There is no unauthenticated compatibility mode. Restart all nodes and clients together after upgrading or rotating the secret. + +Master control messages, worker/job metadata, client/worker status, and monitor status use validated, data-only JSON. Pickle metadata from old versions is rejected. Python classes, arguments, return values, and source files on the authenticated execution channel continue to support Python serialization; only trusted code and trusted secret holders belong in a cluster. + +## Bind addresses and monitoring + +The master control socket, HTTP monitor, and HTTP log server bind to `127.0.0.1` by default. For a multi-machine private cluster, explicitly set `XPARL_BIND_HOST` to the appropriate private interface address on each node, or `0.0.0.0` behind a firewall that permits only trusted cluster nodes. Worker and job ZeroMQ sockets also require authentication, including dynamically allocated ports. Heartbeat gRPC ports still require network isolation. + +All HTTP monitor and log routes require HTTP Basic authentication: username `xparl`, password the cluster secret. The browser prompts for credentials. Use an SSH tunnel or HTTPS reverse proxy for remote HTTP access because HTTP Basic authentication does not encrypt the password. Do not share the cluster secret with users who should only view monitoring data: it also authorizes execution on workers. + +## Upgrading affected installations + +The fix is in the `develop` source branch. Existing PyPI packages and earlier tags are not changed by a source merge. Install the patched source commit on all nodes, or upgrade to a subsequently published release containing it. This JSON/CURVE protocol is incompatible with older xparl processes. Stop the old cluster, configure the secret and private bind addresses, upgrade every node/client, and then restart. + +For an installation that cannot yet upgrade, restrict the master, monitor, log, worker/job, and heartbeat ports to trusted hosts through firewalls or security groups. The reported ports 8010/8137/8200 are examples; check actual configuration and dynamic ports. Remove public exposure immediately. If an old installation was exposed, investigate its host and logs; applying a source patch does not establish whether it was previously compromised. + +## Regression checks + +Run `python -m unittest discover -s parl/remote/tests -p control_security_test.py -v` in a PARL environment. The suite checks plaintext clients, wrong cluster secrets, unrelated CURVE keys, malicious pickle metadata in all four master branches, invalid JSON and message frames, valid status/monitor messages, HTTP authentication, and fail-closed configuration. The `xparl security` GitHub Actions workflow runs these checks with Python 3.9 and 3.10 and the existing pyzmq 22.3.0 pin. diff --git a/parl/remote/client.py b/parl/remote/client.py index b5c15872b..5fb3523e3 100644 --- a/parl/remote/client.py +++ b/parl/remote/client.py @@ -13,12 +13,14 @@ # limitations under the License. import cloudpickle +from parl.remote import control_serialization import datetime import os import socket import sys import threading import zmq +from parl.remote.security import SecureContext import parl import time import glob @@ -27,6 +29,7 @@ from parl.utils import to_str, to_byte, get_ip_address, logger, isnotebook from parl.remote.utils import get_subfiles_recursively from parl.remote import remote_constants +from parl.remote.message import InitializedJob from parl.remote.grpc_heartbeat import HeartbeatServerThread, HeartbeatServerProcess from parl.remote.utils import get_version @@ -65,7 +68,7 @@ def __init__(self, master_address, process_id, distributed_files=[]): th.start() self.master_address = master_address self.process_id = process_id - self.ctx = zmq.Context() + self.ctx = SecureContext() self.lock = threading.Lock() self.log_monitor_url = None self.threads = [] @@ -175,6 +178,7 @@ def _create_sockets(self, master_address): # submit_job_socket: submits job to master self.submit_job_socket = self.ctx.socket(zmq.REQ) + self.ctx.authenticate_client(self.submit_job_socket) self.submit_job_socket.linger = 0 self.submit_job_socket.setsockopt(zmq.RCVTIMEO, remote_constants.HEARTBEAT_TIMEOUT_S * 1000) self.submit_job_socket.connect("tcp://{}".format(master_address)) @@ -259,7 +263,7 @@ def _update_client_status_to_master(self): self.submit_job_socket.send_multipart([ remote_constants.CLIENT_STATUS_UPDATE_TAG, to_byte(self.reply_master_heartbeat_address), - cloudpickle.dumps(client_status) + control_serialization.dumps(client_status) ]) message = self.submit_job_socket.recv_multipart() except zmq.error.Again as e: @@ -276,6 +280,7 @@ def _check_job(self, job_ping_address, max_memory, gpu): """ # job_ping_socket: sends ping signal to job job_ping_socket = self.ctx.socket(zmq.REQ) + self.ctx.authenticate_client(job_ping_socket) job_ping_socket.linger = 0 job_ping_socket.setsockopt(zmq.RCVTIMEO, int(0.9 * 1000)) job_ping_socket.connect("tcp://" + job_ping_address) @@ -302,8 +307,8 @@ def _create_heartbeat_server(self): """ job_heartbeat_port = mp.Value('i', 0) self.actor_num = mp.Value('i', 0) - self.job_heartbeat_process = HeartbeatServerProcess(job_heartbeat_port, self.actor_num, - self.client_is_alive, self.dead_job_queue) + self.job_heartbeat_process = HeartbeatServerProcess(job_heartbeat_port, self.actor_num, self.client_is_alive, + self.dead_job_queue) self.job_heartbeat_process.daemon = True self.job_heartbeat_process.start() assert job_heartbeat_port.value != 0, "fail to initialize heartbeat server for jobs." @@ -346,7 +351,7 @@ def submit_job(self, max_memory, n_gpu, job_is_alive): self.lock.release() tag = message[0] if tag == remote_constants.NORMAL_TAG: - job_info = cloudpickle.loads(message[1]) + job_info = control_serialization.loads(message[1], InitializedJob) job_ping_address = job_info.ping_heartbeat_address self.lock.acquire() diff --git a/parl/remote/cluster_monitor.py b/parl/remote/cluster_monitor.py index 4953b79ce..ef6c3e184 100644 --- a/parl/remote/cluster_monitor.py +++ b/parl/remote/cluster_monitor.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -import cloudpickle +from parl.remote import control_serialization import threading from collections import defaultdict, deque from parl.utils import to_str @@ -135,12 +135,10 @@ def get_status_info(self): vacant_gpus += worker.get('vacant_gpus', 0) self.lock.release() status_info = "has {} used cpus, {} vacant cpus, {} used_gpus, {} vacant_gpus.".format( - used_cpus, vacant_cpus, used_gpus, vacant_gpus) + used_cpus, vacant_cpus, used_gpus, vacant_gpus) return status_info def get_status(self): - """Return a cloudpickled status.""" - self.lock.acquire() - status = cloudpickle.dumps(self.status) - self.lock.release() - return status + """Return data-only JSON status.""" + with self.lock: + return control_serialization.dumps(self.status) diff --git a/parl/remote/control_serialization.py b/parl/remote/control_serialization.py new file mode 100644 index 000000000..0e2ea62ed --- /dev/null +++ b/parl/remote/control_serialization.py @@ -0,0 +1,156 @@ +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Data-only JSON for cluster metadata. Never fall back to pickle.""" + +import json +import math +from collections import deque + +from parl.remote.message import AllocatedCpu, AllocatedGpu, InitializedJob, InitializedWorker + +MAX_CONTROL_BYTES = 16 * 1024 * 1024 +MAX_DEPTH = 32 +_RECORD_TYPES = {cls.__name__: cls for cls in (AllocatedCpu, AllocatedGpu, InitializedJob, InitializedWorker)} +_FIELDS = { + AllocatedCpu: { + 'worker_address': str, + 'n_cpu': int + }, + AllocatedGpu: { + 'worker_address': str, + 'gpu': str + }, + InitializedJob: { + 'job_address': str, + 'worker_heartbeat_address': str, + 'ping_heartbeat_address': str, + 'worker_address': (str, type(None)), + 'pid': int, + 'is_alive': bool, + 'job_id': (str, type(None)), + 'log_server_address': (str, type(None)), + 'allocated_cpu': (AllocatedCpu, type(None)), + 'allocated_gpu': (AllocatedGpu, type(None)), + 'instance_id': (str, int, type(None)) + }, + InitializedWorker: { + 'worker_address': str, + 'initialized_jobs': list, + 'allocated_cpu': AllocatedCpu, + 'allocated_gpu': AllocatedGpu, + 'hostname': str + } +} + + +def _validate_record(value): + schema = _FIELDS[type(value)] + if set(vars(value)) != set(schema): + raise ValueError('Invalid control record fields.') + for name, allowed in schema.items(): + allowed = allowed if isinstance(allowed, tuple) else (allowed, ) + if type(getattr(value, name)) not in allowed: + raise ValueError('Invalid control record field type.') + if type(value) is InitializedWorker and any(type(job) is not InitializedJob for job in value.initialized_jobs): + raise ValueError('Invalid worker job list.') + if type(value) is AllocatedCpu and value.n_cpu < 0: + raise ValueError('Negative CPU allocation.') + + +def _encode(value, depth=0): + if depth > MAX_DEPTH: + raise ValueError('Control metadata is too deeply nested.') + if value is None or type(value) in (bool, int, str): + return value + if type(value) is float and math.isfinite(value): + return value + if isinstance(value, (list, tuple, deque)): + return [_encode(item, depth + 1) for item in value] + if isinstance(value, dict): + if any(type(key) is not str for key in value) or '__xparl_type__' in value: + raise ValueError('Invalid control dictionary keys.') + return {key: _encode(item, depth + 1) for key, item in value.items()} + if type(value) in _FIELDS: + _validate_record(value) + return {'__xparl_type__': type(value).__name__, 'fields': _encode(vars(value), depth + 1)} + raise ValueError('Unsupported control metadata type.') + + +def _decode(value, depth=0): + if depth > MAX_DEPTH: + raise ValueError('Control metadata is too deeply nested.') + if type(value) is list: + return [_decode(item, depth + 1) for item in value] + if type(value) is dict: + if '__xparl_type__' in value: + name = value['__xparl_type__'] + if type(name) is not str or name not in _RECORD_TYPES or set(value) != {'__xparl_type__', 'fields'}: + raise ValueError('Unknown control record.') + fields = _decode(value['fields'], depth + 1) + if type(fields) is not dict: + raise ValueError('Invalid control record.') + record = object.__new__(_RECORD_TYPES[name]) + record.__dict__.update(fields) + _validate_record(record) + return record + return {key: _decode(item, depth + 1) for key, item in value.items()} + if value is None or type(value) in (str, int, bool) or (type(value) is float and math.isfinite(value)): + return value + raise ValueError('Invalid control metadata value.') + + +def _unique_object(pairs): + value = {} + for key, item in pairs: + if key in value: + raise ValueError('Duplicate JSON key.') + value[key] = item + return value + + +def dumps(value): + data = json.dumps(_encode(value), allow_nan=False, separators=(',', ':')).encode('utf-8') + if len(data) > MAX_CONTROL_BYTES: + raise ValueError('Control metadata is too large.') + return data + + +def loads(data, expected_type=dict): + if type(data) is not bytes or len(data) > MAX_CONTROL_BYTES: + raise ValueError('Invalid control metadata size or type.') + try: + value = _decode(json.loads(data.decode('utf-8'), object_pairs_hook=_unique_object)) + except (RecursionError, UnicodeError) as error: + raise ValueError('Invalid control JSON.') from error + if type(value) is not expected_type: + raise ValueError('Unexpected control metadata type.') + return value + + +def loads_status(data, worker=False): + value = loads(data) + if worker: + schema = { + 'vacant_memory': (int, float), + 'used_memory': (int, float), + 'vacant_gpu_memory': (int, float), + 'used_gpu_memory': (int, float), + 'load_time': (str, ), + 'load_value': (int, float) + } + else: + schema = {'file_path': (str, ), 'actor_num': (int, ), 'time': (str, ), 'log_monitor_url': (str, )} + if set(value) != set(schema) or any(type(value[key]) not in types for key, types in schema.items()): + raise ValueError('Invalid status fields.') + return value diff --git a/parl/remote/job.py b/parl/remote/job.py index 4761139f2..dc5500da0 100644 --- a/parl/remote/job.py +++ b/parl/remote/job.py @@ -23,6 +23,7 @@ import argparse import cloudpickle +from parl.remote import control_serialization import pickle import psutil import re @@ -33,6 +34,7 @@ import time import traceback import zmq +from parl.remote.security import SecureContext import importlib import parl from multiprocessing import Process, Pipe @@ -123,7 +125,7 @@ def _create_sockets(self): self.job_address = self.job_address_receiver.recv() self.job_id = self.job_id_receiver.recv() - self.ctx = zmq.Context() + self.ctx = SecureContext() # create the job_socket self.job_socket = create_client_socket(self.ctx, self.worker_address, heartbeat_timeout=True) @@ -150,7 +152,7 @@ def worker_heartbeat_exit_callback_func(): self.log_server_address) try: - self.job_socket.send_multipart([remote_constants.NORMAL_TAG, cloudpickle.dumps(initialized_job)]) + self.job_socket.send_multipart([remote_constants.NORMAL_TAG, control_serialization.dumps(initialized_job)]) message = self.job_socket.recv_multipart() except zmq.error.Again as e: logger.warning("[Job] Cannot connect to the worker {}. ".format(self.worker_address) + "Job will quit.") @@ -164,6 +166,7 @@ def worker_heartbeat_exit_callback_func(): self.worker_pid = int(to_str(message[2])) worker_heartbeat_server_thread.set_host_pid(self.worker_pid) self.remove_job_socket = self.ctx.socket(zmq.REQ) + self.ctx.authenticate_client(self.remove_job_socket) self.remove_job_socket.setsockopt(zmq.RCVTIMEO, remote_constants.HEARTBEAT_TIMEOUT_S * 1000) self.remove_job_socket.connect("tcp://{}".format(remove_job_address)) @@ -332,10 +335,11 @@ def run(self, job_address_sender, job_id_sender): Args: job_address_sender(sending end of multiprocessing.Pipe): send job address of reply_socket to main process. """ - ctx = zmq.Context() + ctx = SecureContext() # create the reply_socket reply_socket = ctx.socket(zmq.REP) + ctx.authenticate_server(reply_socket) job_port = reply_socket.bind_to_random_port(addr="tcp://*") reply_socket.linger = 0 job_ip = get_ip_address() diff --git a/parl/remote/log_server.py b/parl/remote/log_server.py index 032e32d8c..2710fe448 100644 --- a/parl/remote/log_server.py +++ b/parl/remote/log_server.py @@ -23,8 +23,10 @@ from parl.utils import to_byte, logger from parl.remote.grpc_heartbeat import HeartbeatServerThread from parl.remote.zmq_utils import create_client_socket +from parl.remote.security import SecureContext, get_control_bind_host, require_http_auth app = Flask(__name__) +app.before_request(require_http_auth) CORS(app) @@ -42,23 +44,18 @@ def get_log(): try: job_id = request.args['job_id'] except: - return make_response( - jsonify(message="No job_id provided, please check your request."), - 400) + return make_response(jsonify(message="No job_id provided, please check your request."), 400) log_dir = current_app.config.get('LOG_DIR') log_dir = os.path.expanduser(log_dir) log_file_path = os.path.join(log_dir, job_id, 'stdout.log') if not os.path.isfile(log_file_path): - return make_response( - jsonify(message="Log not exsits, please check your job_id"), 400) + return make_response(jsonify(message="Log not exsits, please check your job_id"), 400) else: line_num = current_app.config.get('LINE_NUM') linecache.checkcache(log_file_path) log_content = ''.join(linecache.getlines(log_file_path)[-line_num:]) - return make_response( - jsonify(message="Log exsits, content in log", log=log_content), - 200) + return make_response(jsonify(message="Log exsits, content in log", log=log_content), 200) @app.route( @@ -75,21 +72,18 @@ def download_log(): try: job_id = request.args['job_id'] except: - return make_response( - jsonify(message="No job_id provided, please check your request."), - 400) + return make_response(jsonify(message="No job_id provided, please check your request."), 400) log_dir = current_app.config.get('LOG_DIR') log_dir = os.path.expanduser(log_dir) log_file_path = os.path.join(log_dir, job_id, 'stdout.log') if not os.path.isfile(log_file_path): - return make_response( - jsonify(message="Log not exsits, please check your job_id"), 400) + return make_response(jsonify(message="Log not exsits, please check your job_id"), 400) else: return send_file(log_file_path, as_attachment=True) def send_heartbeat_addr_to_worker(worker_addr, heartbeat_server_addr): - ctx = zmq.Context() + ctx = SecureContext() socket = create_client_socket(ctx, worker_addr, heartbeat_timeout=True) try: @@ -123,17 +117,13 @@ def send_heartbeat_addr_to_worker(worker_addr, heartbeat_server_addr): ) def heartbeat_exit_callback_func(): - logger.warning( - "[log_server] lost connnect with the worker. Please check if it is still alive." - ) + logger.warning("[log_server] lost connnect with the worker. Please check if it is still alive.") os._exit(1) - heartbeat_server_thread = HeartbeatServerThread( - heartbeat_exit_callback_func=heartbeat_exit_callback_func) + heartbeat_server_thread = HeartbeatServerThread(heartbeat_exit_callback_func=heartbeat_exit_callback_func) heartbeat_server_thread.setDaemon(True) heartbeat_server_thread.start() - send_heartbeat_addr_to_worker(args.worker_address, - heartbeat_server_thread.get_address()) + send_heartbeat_addr_to_worker(args.worker_address, heartbeat_server_thread.get_address()) - app.run(host="0.0.0.0", port=args.port) + app.run(host=get_control_bind_host(), port=args.port) diff --git a/parl/remote/master.py b/parl/remote/master.py index 70d37b101..ba9f5f96a 100644 --- a/parl/remote/master.py +++ b/parl/remote/master.py @@ -13,7 +13,6 @@ # limitations under the License. import os -import pickle import threading import time import zmq @@ -25,9 +24,11 @@ from parl.remote.worker_manager import WorkerManager from parl.remote.cluster_monitor import ClusterMonitor from parl.remote.grpc_heartbeat import HeartbeatClientThread -import cloudpickle import time from parl.remote.utils import get_version +from parl.remote import control_serialization +from parl.remote.message import InitializedWorker, InitializedJob +from parl.remote.security import SecureContext, get_control_bind_host class Master(object): @@ -65,14 +66,19 @@ class Master(object): """ def __init__(self, port, monitor_port=None, device=remote_constants.CPU): - self.ctx = zmq.Context() + self.ctx = SecureContext() self.master_ip = get_ip_address() self.all_client_heartbeat_threads = [] self.all_worker_heartbeat_threads = [] - self.monitor_url = "http://{}:{}".format(self.master_ip, monitor_port) + monitor_host = get_control_bind_host() + if monitor_host in ('0.0.0.0', '*'): + monitor_host = self.master_ip + self.monitor_url = "http://{}:{}".format(monitor_host, monitor_port) logger.set_dir(os.path.expanduser('~/.parl_data/master/{}_{}'.format(self.master_ip, port))) self.client_socket = self.ctx.socket(zmq.REP) - self.client_socket.bind("tcp://*:{}".format(port)) + self.ctx.authenticate_server(self.client_socket) + self.client_socket.setsockopt(zmq.MAXMSGSIZE, control_serialization.MAX_CONTROL_BYTES) + self.client_socket.bind("tcp://{}:{}".format(get_control_bind_host(), port)) self.client_socket.linger = 0 self.port = port self.device = device @@ -112,6 +118,11 @@ def _receive_message(self): submittion; (5) reset job. """ message = self.client_socket.recv_multipart() + try: + self._validate_message(message) + except (ValueError, TypeError, KeyError): + self.client_socket.send_multipart([remote_constants.INVALID_MESSAGE_TAG]) + return tag = message[0] # a new worker connects to the master @@ -128,7 +139,7 @@ def _receive_message(self): self.client_socket.send_multipart([remote_constants.NORMAL_TAG, to_byte(status_info)]) elif tag == remote_constants.WORKER_INITIALIZED_TAG: - initialized_worker = cloudpickle.loads(message[1]) + initialized_worker = control_serialization.loads(message[1], InitializedWorker) worker_address = initialized_worker.worker_address success = self.worker_manager.add_worker(initialized_worker) if not success: @@ -218,7 +229,9 @@ def heartbeat_exit_callback_func(client_heartbeat_address): logger.info("Submitting job...") job_info = self.worker_manager.request_job(n_cpu=n_cpu, n_gpu=n_gpu) if job_info: - self.client_socket.send_multipart([remote_constants.NORMAL_TAG, cloudpickle.dumps(job_info)]) + self.client_socket.send_multipart( + [remote_constants.NORMAL_TAG, + control_serialization.dumps(job_info)]) client_id = to_str(message[2]) self.cluster_monitor.add_client_job(client_id, {job_info.job_id: job_info.log_server_address}) self._print_workers() @@ -230,7 +243,7 @@ def heartbeat_exit_callback_func(client_heartbeat_address): # a worker updates elif tag == remote_constants.NEW_JOB_TAG: - initialized_job = cloudpickle.loads(message[1]) + initialized_job = control_serialization.loads(message[1], InitializedJob) last_job_address = to_str(message[2]) self.client_socket.send_multipart([remote_constants.NORMAL_TAG]) @@ -245,7 +258,7 @@ def heartbeat_exit_callback_func(client_heartbeat_address): # client update status periodically elif tag == remote_constants.CLIENT_STATUS_UPDATE_TAG: client_heartbeat_address = to_str(message[1]) - client_status = cloudpickle.loads(message[2]) + client_status = control_serialization.loads_status(message[2]) client_status['client_hostname'] = self.client_hostname[client_heartbeat_address] self.cluster_monitor.update_client_status(client_heartbeat_address, client_status) @@ -254,7 +267,7 @@ def heartbeat_exit_callback_func(client_heartbeat_address): # worker update status periodically elif tag == remote_constants.WORKER_STATUS_UPDATE_TAG: worker_address = to_str(message[1]) - worker_status = cloudpickle.loads(message[2]) + worker_status = control_serialization.loads_status(message[2], worker=True) vacant_cpus = self.worker_manager.get_vacant_cpu(worker_address) total_cpus = self.worker_manager.get_total_cpu(worker_address) @@ -270,7 +283,51 @@ def heartbeat_exit_callback_func(client_heartbeat_address): self.client_socket.send_multipart([remote_constants.NORMAL_TAG]) else: - raise NotImplementedError() + self.client_socket.send_multipart([remote_constants.INVALID_MESSAGE_TAG]) + + def _validate_message(self, message): + counts = { + remote_constants.WORKER_CONNECT_TAG: 1, + remote_constants.MONITOR_TAG: 1, + remote_constants.STATUS_TAG: 1, + remote_constants.WORKER_INITIALIZED_TAG: 2, + remote_constants.CLIENT_CONNECT_TAG: 4, + remote_constants.CHECK_VERSION_TAG: 1, + remote_constants.CLIENT_SUBMIT_TAG: 5, + remote_constants.NEW_JOB_TAG: 3, + remote_constants.CLIENT_STATUS_UPDATE_TAG: 3, + remote_constants.WORKER_STATUS_UPDATE_TAG: 3, + remote_constants.NORMAL_TAG: 1 + } + if not message or message[0] not in counts or len(message) != counts[message[0]]: + raise ValueError('Invalid control message frames.') + tag = message[0] + payload_frame = { + remote_constants.WORKER_INITIALIZED_TAG: 1, + remote_constants.NEW_JOB_TAG: 1, + remote_constants.CLIENT_STATUS_UPDATE_TAG: 2, + remote_constants.WORKER_STATUS_UPDATE_TAG: 2 + }.get(tag) + for index, frame in enumerate(message[1:], 1): + if index != payload_frame: + if len(frame) > 4096: + raise ValueError('Control field is too large.') + frame.decode('utf-8') + if tag == remote_constants.WORKER_INITIALIZED_TAG: + control_serialization.loads(message[1], InitializedWorker) + elif tag == remote_constants.NEW_JOB_TAG: + job = control_serialization.loads(message[1], InitializedJob) + if job.worker_address not in self.worker_manager.worker_hostname: + raise ValueError('Unknown worker.') + elif tag == remote_constants.CLIENT_STATUS_UPDATE_TAG: + control_serialization.loads_status(message[2]) + elif tag == remote_constants.WORKER_STATUS_UPDATE_TAG: + control_serialization.loads_status(message[2], worker=True) + if to_str(message[1]) not in self.worker_manager.worker_hostname: + raise ValueError('Unknown worker.') + elif tag == remote_constants.CLIENT_SUBMIT_TAG: + if int(message[3]) < 0 or int(message[4]) < 0: + raise ValueError('Negative resource request.') def exit(self): """ Close the master. diff --git a/parl/remote/monitor.py b/parl/remote/monitor.py index 295bcccfb..b4c13cf5b 100644 --- a/parl/remote/monitor.py +++ b/parl/remote/monitor.py @@ -13,15 +13,17 @@ # limitations under the License. import argparse -import pickle import random import time import zmq import threading from flask import Flask, render_template, jsonify, request +from parl.remote import control_serialization +from parl.remote.security import SecureContext, get_control_bind_host, require_http_auth app = Flask(__name__) +app.before_request(require_http_auth) @app.route('/') @@ -42,8 +44,9 @@ class ClusterMonitor(object): """ def __init__(self, master_address, gpu_cluster=False): - ctx = zmq.Context() + ctx = SecureContext() self.socket = ctx.socket(zmq.REQ) + ctx.authenticate_client(self.socket) self.socket.setsockopt(zmq.RCVTIMEO, 30000) self.socket.connect('tcp://{}'.format(master_address)) self.data = None @@ -60,7 +63,7 @@ def run(self): self.socket.send_multipart([b'[MONITOR]']) msg = self.socket.recv_multipart() - status = pickle.loads(msg[1]) + status = control_serialization.loads(msg[1]) data = {'workers': [], 'clients': []} total_vacant_cpus = 0 total_used_cpus = 0 @@ -152,4 +155,4 @@ def get_jobs(): args = parser.parse_args() CLUSTER_MONITOR = ClusterMonitor(args.address, args.gpu_cluster) - app.run(host="0.0.0.0", port=args.monitor_port) + app.run(host=get_control_bind_host(), port=args.monitor_port) diff --git a/parl/remote/remote_constants.py b/parl/remote/remote_constants.py index 6e071f227..d3bb442ea 100644 --- a/parl/remote/remote_constants.py +++ b/parl/remote/remote_constants.py @@ -44,6 +44,7 @@ DESERIALIZE_EXCEPTION_TAG = b'[DESERIALIZE_EXCEPTION]' NORMAL_TAG = b'[NORMAL]' +INVALID_MESSAGE_TAG = b'[INVALID_MESSAGE]' REJECT_GPU_JOB_TAG = b'[REJECT_GPU_JOB]' REJECT_CPU_JOB_TAG = b'[REJECT_CPU_JOB]' REJECT_GPU_WORKER_TAG = b'[REJECT_GPU_WORKER]' diff --git a/parl/remote/remote_wrapper.py b/parl/remote/remote_wrapper.py index 6da07faff..63f01327b 100644 --- a/parl/remote/remote_wrapper.py +++ b/parl/remote/remote_wrapper.py @@ -73,6 +73,7 @@ def __init__(self, *args, **kwargs): # Send actor commands like `init` and `call` to the job. self.job_socket = self.ctx.socket(zmq.REQ) + self.ctx.authenticate_client(self.job_socket) self.job_socket.linger = 0 self.job_socket.connect("tcp://{}".format(job_address)) # check the result every 20s to detect the job is still alive. @@ -160,7 +161,7 @@ def set_remote_attr(self, attr, value): self.job_is_alive.value = False raise NotImplementedError() return - + def _receive_from_remote_instance(self, attr): """Receive message from remote instance while checking the job status every 20 seconds. """ diff --git a/parl/remote/scripts.py b/parl/remote/scripts.py index a1e8ebcf5..8c85dcf80 100644 --- a/parl/remote/scripts.py +++ b/parl/remote/scripts.py @@ -26,6 +26,7 @@ import tempfile import warnings import zmq +from parl.remote.security import SecureContext, get_control_bind_host from multiprocessing import Process from parl.utils import (_IS_WINDOWS, get_free_tcp_port, get_ip_address, get_port_from_range, is_port_available, kill_process, to_str) @@ -50,8 +51,9 @@ def is_master_started(address): - ctx = zmq.Context() + ctx = SecureContext() socket = ctx.socket(zmq.REQ) + ctx.authenticate_client(socket) socket.linger = 0 socket.setsockopt(zmq.RCVTIMEO, 500) socket.connect("tcp://{}".format(address)) @@ -226,7 +228,9 @@ def start_master(port, gpu_cluster, cpu_num, gpu, monitor_port, debug, log_serve break time.sleep(3) - master_ip = get_ip_address() + master_ip = get_control_bind_host() + if master_ip in ('0.0.0.0', '*'): + master_ip = get_ip_address() if monitor_is_started: start_info = """ ## If you want to check cluster status, please view: @@ -324,7 +328,7 @@ def status(): if len(clusters) == 0: click.echo('No active cluster is found.') else: - ctx = zmq.Context() + ctx = SecureContext() status = [] for cluster in clusters: if _IS_WINDOWS: @@ -341,6 +345,7 @@ def status(): master_address = monitors[0].split(' ')[2] monitor_address = "{}:{}".format(get_ip_address(), monitor_port) socket = ctx.socket(zmq.REQ) + ctx.authenticate_client(socket) socket.setsockopt(zmq.RCVTIMEO, 10000) socket.connect('tcp://{}'.format(master_address)) try: diff --git a/parl/remote/security.py b/parl/remote/security.py new file mode 100644 index 000000000..417add44b --- /dev/null +++ b/parl/remote/security.py @@ -0,0 +1,115 @@ +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Authentication for xparl's trusted cluster boundary.""" + +import hashlib +import hmac +import os +import threading + +import zmq +from zmq.auth.thread import ThreadAuthenticator +from zmq.utils import z85 + + +def get_auth_token(): + token = os.environ.get('XPARL_AUTH_TOKEN', '') + if len(token.encode('utf-8')) < 32: + raise ValueError('Set XPARL_AUTH_TOKEN to a random secret of at least 32 bytes on every xparl node.') + return token + + +def _curve_keys(token, role): + secret = z85.encode(hashlib.sha256(('xparl-curve-v1:' + role + ':' + token).encode('utf-8')).digest()) + return zmq.curve_public(secret), secret + + +class _ClusterCredentials(object): + def __init__(self, public_key): + self.public_key = public_key + + def callback(self, domain, key): + return domain == 'xparl' and hmac.compare_digest(key, self.public_key) + + +class SecureContext(zmq.Context): + """A context that owns one CURVE authenticator and closes it with its sockets.""" + + _xparl_token = None + _xparl_authenticator = None + _xparl_auth_lock = None + + def __init__(self, *args, **kwargs): + token = get_auth_token() + if not zmq.has('curve'): + raise RuntimeError('xparl requires a libzmq build with CURVE support.') + super(SecureContext, self).__init__(*args, **kwargs) + self._xparl_token = token + self._xparl_authenticator = None + self._xparl_auth_lock = threading.Lock() + + def authenticate_server(self, socket): + server_public, server_secret = _curve_keys(self._xparl_token, 'server') + client_public, _ = _curve_keys(self._xparl_token, 'client') + with self._xparl_auth_lock: + if self._xparl_authenticator is None: + authenticator = ThreadAuthenticator(self) + authenticator.start() + try: + authenticator.configure_curve_callback('xparl', _ClusterCredentials(client_public)) + except Exception: + authenticator.stop() + raise + self._xparl_authenticator = authenticator + socket.curve_publickey = server_public + socket.curve_secretkey = server_secret + socket.curve_server = True + socket.zap_domain = b'xparl' + + def authenticate_client(self, socket): + client_public, client_secret = _curve_keys(self._xparl_token, 'client') + server_public, _ = _curve_keys(self._xparl_token, 'server') + socket.curve_publickey = client_public + socket.curve_secretkey = client_secret + socket.curve_serverkey = server_public + + def _stop_authenticator(self): + authenticator = getattr(self, '_xparl_authenticator', None) + if authenticator is not None: + authenticator.stop() + self._xparl_authenticator = None + + def term(self): + self._stop_authenticator() + super(SecureContext, self).term() + + def destroy(self, linger=None): + self._stop_authenticator() + super(SecureContext, self).destroy(linger=linger) + + +def get_control_bind_host(): + """Remote control and HTTP exposure require an explicit bind address.""" + return os.environ.get('XPARL_BIND_HOST', '127.0.0.1') + + +def require_http_auth(): + """Protect monitor and log endpoints, including static resources.""" + from flask import request, Response + token = get_auth_token() + auth = request.authorization + if auth is not None and auth.type == 'basic' and auth.username == 'xparl': + if hmac.compare_digest((auth.password or '').encode('utf-8'), token.encode('utf-8')): + return None + return Response('Authentication required.', 401, {'WWW-Authenticate': 'Basic realm="xparl"'}) diff --git a/parl/remote/tests/control_security_test.py b/parl/remote/tests/control_security_test.py new file mode 100644 index 000000000..788ee9477 --- /dev/null +++ b/parl/remote/tests/control_security_test.py @@ -0,0 +1,240 @@ +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import base64 +import os +import pickle +import secrets +import unittest +from unittest.mock import patch + +import zmq + +from parl.remote import control_serialization as codec, remote_constants as tags +from parl.remote.master import Master +from parl.remote.message import AllocatedCpu, AllocatedGpu, InitializedJob, InitializedWorker +from parl.remote.monitor import app as monitor_app +from parl.remote.log_server import app as log_app +from parl.remote.security import SecureContext, _curve_keys, get_control_bind_host +from parl.remote.zmq_utils import create_client_socket + +MARKER = 'XPARL_SECURITY_REGRESSION_MARKER' + + +class MaliciousPickle(object): + def __reduce__(self): + return eval, ("__import__('os').environ.__setitem__('" + MARKER + "', 'executed')", ) + + +def make_job(): + return InitializedJob('127.0.0.1:9001', '127.0.0.1:9002', '127.0.0.1:9003', '127.0.0.1:9004', 123, 'job-1', + '127.0.0.1:9005') + + +class ControlSerializationTest(unittest.TestCase): + def test_worker_and_allocated_job_round_trip(self): + job = make_job() + job.worker_address = None + self.assertIsNone(codec.loads(codec.dumps(job), InitializedJob).worker_address) + job.worker_address = '127.0.0.1:9004' + job.allocated_cpu = AllocatedCpu(job.worker_address, 1) + job.allocated_gpu = AllocatedGpu(job.worker_address, '') + worker = InitializedWorker(job.worker_address, [job], job.allocated_cpu, job.allocated_gpu, 'worker') + restored = codec.loads(codec.dumps(worker), InitializedWorker) + self.assertEqual(vars(restored.allocated_cpu), vars(worker.allocated_cpu)) + self.assertEqual(vars(restored.allocated_gpu), vars(worker.allocated_gpu)) + self.assertEqual(restored.initialized_jobs[0].job_id, job.job_id) + self.assertEqual(restored.initialized_jobs[0].allocated_cpu.n_cpu, 1) + + def test_pickle_payload_never_executes(self): + with patch.dict(os.environ): + os.environ.pop(MARKER, None) + with self.assertRaises(ValueError): + codec.loads(pickle.dumps(MaliciousPickle())) + self.assertNotIn(MARKER, os.environ) + + def test_invalid_json_and_record_types(self): + invalid = [ + b'{}junk', b'{"a":1,"a":2}', b'NaN', b'Infinity', b'[]', b'{"__xparl_type__":"os.system","fields":{}}', + b'{"__xparl_type__":"AllocatedCpu","fields":{"worker_address":"x","n_cpu":true}}', + b'{"__xparl_type__":"AllocatedCpu","fields":{"worker_address":"x","n_cpu":-1}}' + ] + for data in invalid: + with self.subTest(data=data), self.assertRaises(ValueError): + codec.loads(data) + + def test_limits_and_unknown_python_types(self): + with self.assertRaises(ValueError): + codec.loads(b' ' * (codec.MAX_CONTROL_BYTES + 1)) + with self.assertRaises(ValueError): + codec.loads(b'[' * 40 + b'0' + b']' * 40) + with self.assertRaises(ValueError): + codec.dumps(object()) + with self.assertRaises(ValueError): + codec.dumps({'value': float('inf')}) + + def test_status_schemas(self): + client = {'file_path': 'train.py', 'actor_num': 1, 'time': '0:00:01', 'log_monitor_url': 'http://localhost/'} + self.assertEqual(codec.loads_status(codec.dumps(client)), client) + worker = { + 'vacant_memory': 1.0, + 'used_memory': 2.0, + 'vacant_gpu_memory': 0, + 'used_gpu_memory': 0, + 'load_time': '12:00', + 'load_value': 0.1 + } + self.assertEqual(codec.loads_status(codec.dumps(worker), worker=True), worker) + for data in ({}, dict(client, actor_num=True), dict(worker, load_time=[])): + with self.assertRaises(ValueError): + codec.loads_status(codec.dumps(data), worker='load_value' in data) + + +class ControlTransportTest(unittest.TestCase): + def setUp(self): + self.token = secrets.token_hex(32) + self.environment = patch.dict(os.environ, {'XPARL_AUTH_TOKEN': self.token, 'XPARL_BIND_HOST': '127.0.0.1'}) + self.environment.start() + with patch('parl.remote.master.logger.set_dir'), patch( + 'parl.remote.master.get_ip_address', return_value='127.0.0.1'): + self.master = Master(0) + self.master.client_socket.setsockopt(zmq.RCVTIMEO, 2000) + self.endpoint = self.master.client_socket.getsockopt(zmq.LAST_ENDPOINT).decode() + self.address = self.endpoint[len('tcp://'):] + self.contexts = [self.master.ctx] + self.client_ctx = SecureContext() + self.contexts.append(self.client_ctx) + self.client = create_client_socket(self.client_ctx, self.address) + self.client.setsockopt(zmq.RCVTIMEO, 2000) + + def tearDown(self): + self.master.exit() + for ctx in reversed(self.contexts): + ctx.destroy(linger=0) + self.environment.stop() + + def exchange(self, frames): + self.client.send_multipart(frames) + self.master._receive_message() + return self.client.recv_multipart() + + def assert_healthy(self): + self.assertEqual(self.exchange([tags.CHECK_VERSION_TAG])[0], tags.NORMAL_TAG) + + def test_plaintext_client_cannot_reach_master(self): + ctx = zmq.Context() + self.contexts.append(ctx) + client = ctx.socket(zmq.REQ) + client.linger = 0 + self.addCleanup(client.close, 0) + client.setsockopt(zmq.RCVTIMEO, 200) + client.connect(self.endpoint) + client.send_multipart([tags.WORKER_INITIALIZED_TAG, pickle.dumps(MaliciousPickle())]) + with self.assertRaises(zmq.Again): + client.recv_multipart() + self.assertEqual(self.master.client_socket.poll(100), 0) + self.assert_healthy() + + def test_wrong_token_cannot_reach_master(self): + with patch.dict(os.environ, {'XPARL_AUTH_TOKEN': secrets.token_hex(32)}): + ctx = SecureContext() + self.contexts.append(ctx) + client = create_client_socket(ctx, self.address) + self.addCleanup(client.close, 0) + client.setsockopt(zmq.RCVTIMEO, 200) + client.send_multipart([tags.NORMAL_TAG]) + with self.assertRaises(zmq.Again): + client.recv_multipart() + self.assertEqual(self.master.client_socket.poll(100), 0) + self.assert_healthy() + + def test_unknown_curve_client_key_is_rejected(self): + ctx = zmq.Context() + self.contexts.append(ctx) + client = ctx.socket(zmq.REQ) + client.linger = 0 + self.addCleanup(client.close, 0) + client.curve_publickey, client.curve_secretkey = zmq.curve_keypair() + client.curve_serverkey = _curve_keys(self.token, 'server')[0] + client.setsockopt(zmq.RCVTIMEO, 200) + client.connect(self.endpoint) + client.send_multipart([tags.NORMAL_TAG]) + with self.assertRaises(zmq.Again): + client.recv_multipart() + self.assertEqual(self.master.client_socket.poll(100), 0) + self.assert_healthy() + + def test_all_deserialization_branches_reject_pickle_and_recover(self): + payload = pickle.dumps(MaliciousPickle()) + messages = [[tags.WORKER_INITIALIZED_TAG, payload], [tags.NEW_JOB_TAG, payload, b'old-job'], + [tags.CLIENT_STATUS_UPDATE_TAG, b'client', payload], + [tags.WORKER_STATUS_UPDATE_TAG, b'worker', payload]] + with patch.dict(os.environ): + os.environ.pop(MARKER, None) + for frames in messages: + with self.subTest(tag=frames[0]): + self.assertEqual(self.exchange(frames), [tags.INVALID_MESSAGE_TAG]) + self.assertNotIn(MARKER, os.environ) + self.assert_healthy() + + def test_bad_frames_and_unknown_workers_do_not_stop_master(self): + messages = [[b'unknown'], [tags.NEW_JOB_TAG], [tags.CLIENT_CONNECT_TAG, b'\xff', b'host', b'id'], + [tags.CLIENT_SUBMIT_TAG, b'client', b'id', b'not-a-number', b'0'], + [tags.CLIENT_SUBMIT_TAG, b'client', b'id', b'-1', b'0'], + [tags.WORKER_STATUS_UPDATE_TAG, b'unknown', b'{}']] + for frames in messages: + with self.subTest(frames=frames): + self.assertEqual(self.exchange(frames), [tags.INVALID_MESSAGE_TAG]) + self.assert_healthy() + + def test_client_status_and_monitor_json(self): + status = {'file_path': 'train.py', 'actor_num': 2, 'time': '0:00:01', 'log_monitor_url': 'http://localhost/'} + self.master.client_hostname['client'] = 'test-client' + self.assertEqual( + self.exchange([tags.CLIENT_STATUS_UPDATE_TAG, b'client', + codec.dumps(status)]), [tags.NORMAL_TAG]) + response = self.exchange([tags.MONITOR_TAG]) + self.assertEqual(codec.loads(response[1])['clients']['client']['client_hostname'], 'test-client') + self.assert_healthy() + + +class ConfigurationAndHttpTest(unittest.TestCase): + def test_missing_or_short_token_fails_closed(self): + for token in ('', 'short'): + with patch.dict(os.environ, {'XPARL_AUTH_TOKEN': token}), self.assertRaises(ValueError): + SecureContext() + + def test_bind_is_loopback_by_default(self): + with patch.dict(os.environ): + os.environ.pop('XPARL_BIND_HOST', None) + self.assertEqual(get_control_bind_host(), '127.0.0.1') + + def test_monitor_and_logs_require_credentials(self): + token = secrets.token_hex(32) + with patch.dict(os.environ, {'XPARL_AUTH_TOKEN': token}): + credentials = base64.b64encode(('xparl:' + token).encode()).decode() + for app, route in ((monitor_app, '/cluster'), (log_app, '/get-log')): + client = app.test_client() + self.assertEqual(client.get(route).status_code, 401) + self.assertEqual( + client.get(route, headers={ + 'Authorization': 'Basic eHBhcmw6d3Jvbmc=' + }).status_code, 401) + # A protected static asset avoids dependence on a running cluster. + with client.get('/static/favicon.ico', headers={'Authorization': 'Basic ' + credentials}) as response: + self.assertEqual(response.status_code, 200) + + +if __name__ == '__main__': + unittest.main() diff --git a/parl/remote/tests/log_server_test.py b/parl/remote/tests/log_server_test.py index bcc31f27c..3a9b84423 100644 --- a/parl/remote/tests/log_server_test.py +++ b/parl/remote/tests/log_server_test.py @@ -15,7 +15,6 @@ import json import multiprocessing import os -import pickle import subprocess import sys import tempfile @@ -32,6 +31,12 @@ from parl.utils.test_utils import XparlTestCase from parl.utils import get_free_tcp_port from parl.remote.master import Master +from parl.remote import control_serialization +from parl.remote.security import get_auth_token + + +def authenticated_get(url, **kwargs): + return requests.get(url, auth=('xparl', get_auth_token()), **kwargs) @parl.remote_class @@ -83,7 +88,7 @@ def test_log_server(self): # Get status status = master._get_status() - client_jobs = pickle.loads(status).get('client_jobs') + client_jobs = control_serialization.loads(status).get('client_jobs') self.assertIsNotNone(client_jobs) # Get job id @@ -94,10 +99,10 @@ def test_log_server(self): for job_id, log_server_addr in jobs.items(): log_url = "http://{}/get-log".format(log_server_addr) # Test response without job_id - r = requests.get(log_url) + r = authenticated_get(log_url) self.assertEqual(r.status_code, 400) # Test normal response - r = requests.get(log_url, params={'job_id': job_id}) + r = authenticated_get(log_url, params={'job_id': job_id}) self.assertEqual(r.status_code, 200) log_content = json.loads(r.text).get('log') self.assertIsNotNone(log_content) @@ -106,7 +111,7 @@ def test_log_server(self): # Test download download_url = "http://{}/download-log".format(log_server_addr) - r = requests.get(download_url, params={'job_id': job_id}) + r = authenticated_get(download_url, params={'job_id': job_id}) self.assertEqual(r.status_code, 200) log_content = r.text.replace('\r\n', '\n') self.assertIn(log_content, outputs) @@ -122,8 +127,7 @@ def test_monitor_query_log_server(self): time.sleep(1) # start the cluster monitor monitor_file = __file__.replace('log_server_test.pyc', '../monitor.py') - monitor_file = monitor_file.replace('log_server_test.py', - '../monitor.py') + monitor_file = monitor_file.replace('log_server_test.py', '../monitor.py') command = [ sys.executable, monitor_file, "--monitor_port", str(monitor_port), "--address", "localhost:" + str(master_port) @@ -132,8 +136,7 @@ def test_monitor_query_log_server(self): FNULL = tempfile.TemporaryFile() else: FNULL = open(os.devnull, 'w') - monitor_proc = subprocess.Popen( - command, stdout=FNULL, stderr=subprocess.STDOUT, close_fds=True) + monitor_proc = subprocess.Popen(command, stdout=FNULL, stderr=subprocess.STDOUT, close_fds=True) # Start worker cluster_addr = 'localhost:{}'.format(master_port) @@ -143,15 +146,14 @@ def test_monitor_query_log_server(self): outputs = self._connect_and_create_actor(cluster_addr) time.sleep(5) # Wait for the status update client = get_global_client() - jobs_url = "{}/get-jobs?client_id={}".format(master.monitor_url, - client.client_id) - r = requests.get(jobs_url) + jobs_url = "{}/get-jobs?client_id={}".format(master.monitor_url, client.client_id) + r = authenticated_get(jobs_url) self.assertEqual(r.status_code, 200) data = json.loads(r.text) for job in data: log_url = job.get('log_url') self.assertIsNotNone(log_url) - r = requests.get(log_url) + r = authenticated_get(log_url) self.assertEqual(r.status_code, 200) log_content = json.loads(r.text).get('log') self.assertIsNotNone(log_content) @@ -160,7 +162,7 @@ def test_monitor_query_log_server(self): # Test download download_url = job.get('download_url') - r = requests.get(download_url) + r = authenticated_get(download_url) self.assertEqual(r.status_code, 200) log_content = r.text.replace('\r\n', '\n') self.assertIn(log_content, outputs) diff --git a/parl/remote/worker.py b/parl/remote/worker.py index 8eae5d1a0..577fe6884 100644 --- a/parl/remote/worker.py +++ b/parl/remote/worker.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -import cloudpickle +from parl.remote import control_serialization import multiprocessing as mp import os import psutil @@ -25,12 +25,13 @@ import threading import warnings import zmq +from parl.remote.security import SecureContext, get_control_bind_host from datetime import datetime import pynvml import parl from parl.utils import get_ip_address, to_byte, to_str, logger, _IS_WINDOWS from parl.remote import remote_constants -from parl.remote.message import InitializedWorker, AllocatedCpu, AllocatedGpu +from parl.remote.message import InitializedWorker, InitializedJob, AllocatedCpu, AllocatedGpu from parl.remote.status import WorkerStatus from parl.remote.zmq_utils import create_server_socket, create_client_socket from parl.remote.grpc_heartbeat import HeartbeatServerThread, HeartbeatClientThread @@ -81,7 +82,7 @@ def __init__(self, master_address, cpu_num=None, log_server_port=None, gpu=''): # initialzation self.pid = str(os.getpid()) self.lock = threading.Lock() - self.ctx = zmq.Context.instance() + self.ctx = SecureContext.instance() self.master_address = master_address self.master_is_alive = True self.worker_is_alive = True @@ -171,6 +172,7 @@ def _create_sockets(self): # request_master_socket: sends job address to master self.request_master_socket = self.ctx.socket(zmq.REQ) + self.ctx.authenticate_client(self.request_master_socket) self.request_master_socket.linger = 0 # wait for 0.5 second to check whether master is started @@ -179,12 +181,14 @@ def _create_sockets(self): # reply_job_socket: receives job_address from subprocess self.reply_job_socket = self.ctx.socket(zmq.REP) + self.ctx.authenticate_server(self.reply_job_socket) self.reply_job_socket.linger = 0 reply_job_port = self.reply_job_socket.bind_to_random_port("tcp://*") self.reply_job_address = "{}:{}".format(self.worker_ip, reply_job_port) # remove_job_socket self.remove_job_socket = self.ctx.socket(zmq.REP) + self.ctx.authenticate_server(self.remove_job_socket) self.remove_job_socket.linger = 0 remove_job_port = self.remove_job_socket.bind_to_random_port("tcp://*") self.remove_job_address = "{}:{}".format(self.worker_ip, remove_job_port) @@ -237,7 +241,7 @@ def master_heartbeat_exit_callback_func(): allocated_gpu, socket.gethostname()) self.request_master_socket.send_multipart( [remote_constants.WORKER_INITIALIZED_TAG, - cloudpickle.dumps(initialized_worker)]) + control_serialization.dumps(initialized_worker)]) message = self.request_master_socket.recv_multipart() if message[0] == remote_constants.REJECT_CPU_WORKER_TAG: @@ -295,8 +299,11 @@ def _init_jobs(self, job_num): new_jobs = [] for _ in range(job_num): job_init_message = self.reply_job_socket.recv_multipart() - self.reply_job_socket.send_multipart([remote_constants.NORMAL_TAG, to_byte(self.remove_job_address), to_byte(self.pid)]) - initialized_job = cloudpickle.loads(job_init_message[1]) + self.reply_job_socket.send_multipart( + [remote_constants.NORMAL_TAG, + to_byte(self.remove_job_address), + to_byte(self.pid)]) + initialized_job = control_serialization.loads(job_init_message[1], InitializedJob) new_jobs.append(initialized_job) def heartbeat_exit_callback_func(job): @@ -337,7 +344,7 @@ def _remove_job(self, job_address): self.lock.acquire() self.request_master_socket.send_multipart( [remote_constants.NEW_JOB_TAG, - cloudpickle.dumps(initialized_job), + control_serialization.dumps(initialized_job), to_byte(job_address)]) _ = self.request_master_socket.recv_multipart() self.lock.release() @@ -403,7 +410,7 @@ def _update_worker_status_to_master(self): self.request_master_socket.send_multipart([ remote_constants.WORKER_STATUS_UPDATE_TAG, to_byte(self.master_heartbeat_address), - cloudpickle.dumps(worker_status) + control_serialization.dumps(worker_status) ]) message = self.request_master_socket.recv_multipart() except zmq.error.Again as e: @@ -441,7 +448,10 @@ def _create_log_server(self, port): log_server_proc = subprocess.Popen(command, stdout=FNULL, close_fds=True) FNULL.close() - log_server_address = "{}:{}".format(self.worker_ip, port) + log_host = get_control_bind_host() + if log_host in ('0.0.0.0', '*'): + log_host = self.worker_ip + log_server_address = "{}:{}".format(log_host, port) message = self.reply_log_server_socket.recv_multipart() log_server_heartbeat_addr = to_str(message[1]) diff --git a/parl/remote/zmq_utils.py b/parl/remote/zmq_utils.py index 96b3a6c26..eaec27b36 100644 --- a/parl/remote/zmq_utils.py +++ b/parl/remote/zmq_utils.py @@ -29,9 +29,9 @@ def create_server_socket(ctx, heartbeat_timeout=False): port(int): port of the server socket. """ socket = ctx.socket(zmq.REP) + ctx.authenticate_server(socket) if heartbeat_timeout: - socket.setsockopt(zmq.RCVTIMEO, - remote_constants.HEARTBEAT_RCVTIMEO_S * 1000) + socket.setsockopt(zmq.RCVTIMEO, remote_constants.HEARTBEAT_RCVTIMEO_S * 1000) socket.linger = 0 port = socket.bind_to_random_port(addr="tcp://*") return socket, port @@ -51,9 +51,9 @@ def create_client_socket(ctx, server_socket_address, heartbeat_timeout=False): socket(zmq.Context().socket): socket of the client. """ socket = ctx.socket(zmq.REQ) + ctx.authenticate_client(socket) if heartbeat_timeout: - socket.setsockopt(zmq.RCVTIMEO, - remote_constants.HEARTBEAT_RCVTIMEO_S * 1000) + socket.setsockopt(zmq.RCVTIMEO, remote_constants.HEARTBEAT_RCVTIMEO_S * 1000) socket.linger = 0 socket.connect("tcp://{}".format(server_socket_address))