Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions .github/workflows/xparl_security.yml
Original file line number Diff line number Diff line change
@@ -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
14 changes: 6 additions & 8 deletions README.cn.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)。
13 changes: 6 additions & 7 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
27 changes: 27 additions & 0 deletions docs/xparl_security.md
Original file line number Diff line number Diff line change
@@ -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.
15 changes: 10 additions & 5 deletions parl/remote/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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 = []
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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:
Expand All @@ -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)
Expand All @@ -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."
Expand Down Expand Up @@ -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()
Expand Down
12 changes: 5 additions & 7 deletions parl/remote/cluster_monitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
156 changes: 156 additions & 0 deletions parl/remote/control_serialization.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading