from __future__ import annotations
import functools
import hashlib
import hmac
import marshal
import os
import pickle
import socket
import struct
import sys
import threading
import traceback
import types
import cloudpickle as _cloudpickle
__all__ = [
"RemoteError",
"RemoteFunction",
"RemoteTraceback",
"Session",
"connect",
"remote",
"serve",
"serve_forever",
]
_MAX_FRAME = 1 << 34
class RemoteTraceback(Exception):
class RemoteError(Exception):
def _send(sock, header, payload=None, bufs=()):
hb = pickle.dumps(header)
sock.sendall(struct.pack("<II", len(hb), len(bufs) + (payload is not None)))
sock.sendall(hb)
if payload is not None:
m = memoryview(payload).cast("B")
sock.sendall(struct.pack("<Q", m.nbytes))
sock.sendall(m)
for b in bufs:
m = memoryview(b).cast("B")
sock.sendall(struct.pack("<Q", m.nbytes))
sock.sendall(m)
def _recv_exact(sock, n):
buf = bytearray(n)
view = memoryview(buf)
i = 0
while i < n:
k = sock.recv_into(view[i:], n - i)
if not k:
raise ConnectionError("peer closed")
i += k
return buf
def _recv(sock):
hlen, nbufs = struct.unpack("<II", _recv_exact(sock, 8))
header = pickle.loads(_recv_exact(sock, hlen))
bufs = []
for _ in range(nbufs):
(blen,) = struct.unpack("<Q", _recv_exact(sock, 8))
if blen > _MAX_FRAME:
raise ConnectionError(f"oversized frame ({blen} bytes)")
bufs.append(_recv_exact(sock, blen))
return header, bufs
def _dumps_oob(obj):
oob = []
payload = _cloudpickle.dumps(obj, protocol=5, buffer_callback=lambda b: oob.append(b.raw()))
return payload, oob
def _authenticate(sock, authkey, *, server):
if not isinstance(authkey, bytes):
raise TypeError("authkey must be bytes")
def challenge():
nonce = os.urandom(32)
sock.sendall(nonce)
reply = _recv_exact(sock, 32)
if not hmac.compare_digest(hmac.digest(authkey, nonce, "sha256"), reply):
raise ConnectionError("authentication failed")
def respond():
nonce = _recv_exact(sock, 32)
sock.sendall(hmac.digest(authkey, bytes(nonce), "sha256"))
if server:
challenge()
respond()
else:
respond()
challenge()
def _default_ship(fn):
if "." in fn.__module__ or "<locals>" in fn.__qualname__:
return "pickle"
mod = sys.modules.get(fn.__module__)
file = getattr(mod, "__file__", None)
if file and os.path.isfile(file):
return "source"
return "pickle"
def _pack_function(fn, ship):
if ship is None:
ship = _default_ship(fn)
if ship == "pickle":
bundle = {"mode": "pickle", "data": _cloudpickle.dumps(fn)}
elif ship == "source":
mod = sys.modules.get(fn.__module__)
file = getattr(mod, "__file__", None)
if not (file and os.path.isfile(file)) or "<locals>" in fn.__qualname__:
raise RuntimeError(
f"ship='source' needs {fn.__qualname__} at top level of a "
"module with a source file"
)
with open(file, "rb") as fh:
source = fh.read()
bundle = {
"mode": "source",
"source": source,
"modname": fn.__module__,
"qualname": fn.__qualname__,
}
elif ship == "code":
if fn.__closure__:
raise RuntimeError(
f"ship='code' cannot carry closures ({fn.__qualname__}); "
"use the default cloudpickle mode"
)
bundle = {
"mode": "code",
"code": marshal.dumps(fn.__code__),
"name": fn.__name__,
"defaults": fn.__defaults__,
"kwdefaults": fn.__kwdefaults__,
}
else:
raise ValueError(f"unknown ship mode {ship!r}")
payload = pickle.dumps(bundle, protocol=5)
return hashlib.sha256(payload).hexdigest()[:16], payload
def _load_function(payload, code_hash):
bundle = pickle.loads(payload)
mode = bundle["mode"]
if mode == "pickle":
fn = _cloudpickle.loads(bundle["data"])
elif mode == "source":
name = f"_omp_remote_{code_hash}"
mod = sys.modules.get(name)
if mod is None:
mod = types.ModuleType(name)
mod.__dict__["__omp_remote_origin__"] = bundle["modname"]
sys.modules[name] = mod
code = compile(bundle["source"], f"<remote {bundle['modname']}>", "exec")
exec(code, mod.__dict__)
obj = mod
for part in bundle["qualname"].split("."):
obj = getattr(obj, part)
fn = obj
elif mode == "code":
code = marshal.loads(bundle["code"])
namespace = {"__builtins__": __builtins__}
fn = types.FunctionType(code, namespace, bundle["name"], bundle["defaults"])
fn.__kwdefaults__ = bundle["kwdefaults"]
else:
raise ValueError(f"unknown bundle mode {mode!r}")
return fn.fn if isinstance(fn, RemoteFunction) else fn
_default_session = None
class RemoteFunction:
def __init__(self, fn, ship=None):
self.fn = fn
self._ship = ship
self._packed = None functools.update_wrapper(self, fn)
def __call__(self, *args, **kwargs):
return self.fn(*args, **kwargs)
def _pack(self):
if self._packed is None:
self._packed = _pack_function(self.fn, self._ship)
return self._packed
def remote(self, *args, **kwargs):
if _default_session is None:
raise RuntimeError("no default session; call omp_remote.connect() first")
return _default_session.call(self, *args, **kwargs)
def remote(fn=None, *, ship=None):
if fn is None:
return lambda f: RemoteFunction(f, ship)
return RemoteFunction(fn, ship)
class Session:
def __init__(self, sock, authkey=None):
if authkey is not None:
_authenticate(sock, authkey, server=False)
self._sock = sock
self._lock = threading.Lock()
def call(self, rf, /, *args, **kwargs):
if not isinstance(rf, RemoteFunction):
rf = RemoteFunction(rf)
code_hash, code_payload = rf._pack()
payload, oob = _dumps_oob((args, kwargs))
with self._lock:
_send(self._sock, {"op": "call", "hash": code_hash}, payload, oob)
header, frames = _recv(self._sock)
if header["op"] == "need_code":
_send(self._sock, {"op": "register", "hash": code_hash}, code_payload)
header, frames = _recv(self._sock)
if header["op"] == "error":
try:
exc = pickle.loads(frames[0])
except Exception:
exc = RemoteError(header["exc"])
raise exc from RemoteTraceback(header["traceback"])
return pickle.loads(frames[0], buffers=frames[1:])
def close(self):
self._sock.close()
def __enter__(self):
return self
def __exit__(self, *exc):
self.close()
def connect(address, authkey=None):
global _default_session
if isinstance(address, tuple):
sock = socket.create_connection(address)
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
else:
sock = socket.socket(socket.AF_UNIX)
sock.connect(address)
_default_session = Session(sock, authkey)
return _default_session
def serve(sock, authkey=None):
if authkey is not None:
_authenticate(sock, authkey, server=True)
fns = {}
pending = None while True:
try:
header, frames = _recv(sock)
except ConnectionError:
return
op = header["op"]
if op == "register":
try:
fns[header["hash"]] = _load_function(frames[0], header["hash"])
except BaseException as exc: _send_error(sock, exc)
pending = None
continue
if pending and pending[0] == header["hash"]:
_execute(sock, fns[header["hash"]], pending[1])
pending = None
elif op == "call":
fn = fns.get(header["hash"])
if fn is None:
pending = (header["hash"], frames)
_send(sock, {"op": "need_code"})
else:
_execute(sock, fn, frames)
elif op == "shutdown":
return
else:
raise ValueError(f"unknown op {op!r}")
def _execute(sock, fn, frames):
try:
args, kwargs = pickle.loads(frames[0], buffers=frames[1:])
payload, oob = _dumps_oob(fn(*args, **kwargs))
_send(sock, {"op": "result"}, payload, oob)
except BaseException as exc: _send_error(sock, exc)
def _send_error(sock, exc):
summary = f"{type(exc).__name__}: {exc}"
tb = traceback.format_exc()
try:
data = _cloudpickle.dumps(exc)
except Exception:
data = pickle.dumps(RemoteError(summary))
_send(sock, {"op": "error", "exc": summary, "traceback": tb}, data)
def serve_forever(address, authkey=None):
if isinstance(address, tuple):
srv = socket.create_server(address)
else:
if os.path.exists(address):
os.unlink(address)
srv = socket.socket(socket.AF_UNIX)
srv.bind(address)
srv.listen()
while True:
conn, _ = srv.accept()
threading.Thread(target=serve, args=(conn, authkey), daemon=True).start()