This commit is contained in:
Mark Qvist
2026-07-18 19:23:02 +02:00
parent 88833f170b
commit d93f9798ae
7 changed files with 80 additions and 99 deletions
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.2.0" __version__ = "0.3.0"
-2
View File
@@ -53,8 +53,6 @@ class permit(AbstractContextManager):
""" """
def __init__(self, *exceptions): self._exceptions = exceptions def __init__(self, *exceptions): self._exceptions = exceptions
def __enter__(self): pass def __enter__(self): pass
def __exit__(self, exctype, excinst, exctb): def __exit__(self, exctype, excinst, exctb):
return exctype is not None and not issubclass(exctype, self._exceptions) return exctype is not None and not issubclass(exctype, self._exceptions)
-1
View File
@@ -55,5 +55,4 @@ class SleepRate:
return sleep_for if sleep_for > 0 else 0 return sleep_for if sleep_for > 0 else 0
async def sleep_async(self): await asyncio.sleep(self.next_sleep_time()) async def sleep_async(self): await asyncio.sleep(self.next_sleep_time())
def sleep_block(self): time.sleep(self.next_sleep_time()) def sleep_block(self): time.sleep(self.next_sleep_time())
+2 -10
View File
@@ -239,12 +239,6 @@ async def initiate(configdir: str, rnsconfigdir:str, identitypath: str, verbosit
channel = _link.get_channel() channel = _link.get_channel()
protocol.register_message_types(channel) protocol.register_message_types(channel)
channel.add_message_handler(_client_message_handler) channel.add_message_handler(_client_message_handler)
# Next step after linking and identifying: send version
# if not await _spin(lambda: messenger.is_outlet_ready(outlet), timeout=5, quiet=quietness > 0):
# print("Error bringing up link")
# return 253
channel.send(protocol.VersionInfoMessage()) channel.send(protocol.VersionInfoMessage())
try: try:
vm = _pq.get(timeout=max(outlet.rtt * 20, 5)) vm = _pq.get(timeout=max(outlet.rtt * 20, 5))
@@ -283,10 +277,8 @@ async def initiate(configdir: str, rnsconfigdir:str, identitypath: str, verbosit
return None return None
elif b == "L": elif b == "L":
line_mode = not line_mode line_mode = not line_mode
if line_mode: if line_mode: os.write(1, "\n\rLine-interactive mode enabled\n\r".encode("utf-8"))
os.write(1, "\n\rLine-interactive mode enabled\n\r".encode("utf-8")) else: os.write(1, "\n\rLine-interactive mode disabled\n\r".encode("utf-8"))
else:
os.write(1, "\n\rLine-interactive mode disabled\n\r".encode("utf-8"))
return None return None
return b return b
+4 -10
View File
@@ -88,8 +88,7 @@ class RetryThread(AbstractContextManager):
self._thread = threading.Thread(name=name, target=self._thread_run, daemon=True) self._thread = threading.Thread(name=name, target=self._thread_run, daemon=True)
self._thread.start() self._thread.start()
def is_alive(self): def is_alive(self): return self._thread.is_alive()
return self._thread.is_alive()
def close(self, loop: asyncio.AbstractEventLoop = None) -> asyncio.Future: def close(self, loop: asyncio.AbstractEventLoop = None) -> asyncio.Future:
RNS.log("Stopping timer thread", RNS.LOG_DEBUG) RNS.log("Stopping timer thread", RNS.LOG_DEBUG)
@@ -102,17 +101,12 @@ class RetryThread(AbstractContextManager):
return self._finished return self._finished
def wait(self, timeout: float = None): def wait(self, timeout: float = None):
if timeout: if timeout: timeout = timeout + time.time()
timeout = timeout + time.time()
while timeout is None or time.time() < timeout: while timeout is None or time.time() < timeout:
with self._lock: with self._lock: task_count = len(self._statuses)
task_count = len(self._statuses) if task_count == 0: return
if task_count == 0:
return
time.sleep(0.1) time.sleep(0.1)
def _thread_run(self): def _thread_run(self):
while self._run and self._finished is None: while self._run and self._finished is None:
time.sleep(self._loop_period) time.sleep(self._loop_period)
+2 -4
View File
@@ -105,9 +105,8 @@ def ensure_config_directory():
async def _rnsh_cli_main(): async def _rnsh_cli_main():
global verbose_set global verbose_set
args, parser = parse_arguments() args, parser = parse_arguments()
verbose_set = args.verbose > 0 verbose_set = args.verbose > 0
configdir = ensure_config_directory()
configdir = ensure_config_directory()
if args.print_identity: if args.print_identity:
print_identity(args.config, args.identity, args.service, args.listen) print_identity(args.config, args.identity, args.service, args.listen)
@@ -170,5 +169,4 @@ def main():
if verbose_set and exc: raise exc if verbose_set and exc: raise exc
sys.exit(return_code if return_code is not None else 255) sys.exit(return_code if return_code is not None else 255)
if __name__ == "__main__": main() if __name__ == "__main__": main()
+71 -71
View File
@@ -103,6 +103,8 @@ class ListenerSession:
self.outlet.set_link_closed_callback(self._link_closed) self.outlet.set_link_closed_callback(self._link_closed)
self.loop = loop self.loop = loop
self.state: LSState = None self.state: LSState = None
self.terminated = False
self.authenticated = False
self.remote_identity = None self.remote_identity = None
self.term: str | None = None self.term: str | None = None
self.stdin_is_pipe: bool = False self.stdin_is_pipe: bool = False
@@ -122,8 +124,10 @@ class ListenerSession:
self.return_code_sent = False self.return_code_sent = False
self.process: process.CallbackSubprocess | None = None self.process: process.CallbackSubprocess | None = None
if self.allow_all: self._set_state(LSState.LSSTATE_WAIT_VERS) if not self.allow_all: self._set_state(LSState.LSSTATE_WAIT_IDENT)
else: self._set_state(LSState.LSSTATE_WAIT_IDENT) else:
self._set_state(LSState.LSSTATE_WAIT_VERS)
self.authenticated = True
self.sessions.append(self) self.sessions.append(self)
protocol.register_message_types(self.channel) protocol.register_message_types(self.channel)
@@ -143,37 +147,31 @@ class ListenerSession:
def _call(self, func: callable, delay: float = 0): def _call(self, func: callable, delay: float = 0):
def call_inner(): def call_inner():
if delay == 0: func() if delay == 0: func()
else: self.loop.call_later(delay, func) else: self.loop.call_later(delay, func)
self.loop.call_soon_threadsafe(call_inner) self.loop.call_soon_threadsafe(call_inner)
def send(self, message: RNS.MessageBase): def send(self, message: RNS.MessageBase): self.channel.send(message)
self.channel.send(message) def _protocol_error(self, name: str): self.terminate(f"Protocol error ({name})")
def _protocol_timeout_error(self, name: str): self.terminate(f"Protocol timeout error: {name}")
def _protocol_error(self, name: str):
self.terminate(f"Protocol error ({name})")
def _protocol_timeout_error(self, name: str):
self.terminate(f"Protocol timeout error: {name}")
def terminate(self, error: str = None): def terminate(self, error: str = None):
self.terminated = True
self.state = LSState.LSSTATE_ERROR
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
RNS.log("Terminating session" + (f": {error}" if error else ""), RNS.LOG_DEBUG) RNS.log("Terminating session" + (f": {error}" if error else ""), RNS.LOG_DEBUG)
if error and self.state != LSState.LSSTATE_TEARDOWN: if error and self.state != LSState.LSSTATE_TEARDOWN:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
self.send(protocol.ErrorMessage(error, True)) self.send(protocol.ErrorMessage(error, True))
self.state = LSState.LSSTATE_ERROR
self._terminate_process() self._terminate_process()
self._call(self._prune, max(self.outlet.rtt * 3, process.CallbackSubprocess.PROCESS_PIPE_TIME+5)) self._call(self._prune, max(self.outlet.rtt * 3, process.CallbackSubprocess.PROCESS_PIPE_TIME+5))
def _prune(self): def _prune(self):
self.state = LSState.LSSTATE_TEARDOWN self.state = LSState.LSSTATE_TEARDOWN
RNS.log("Pruning session", RNS.LOG_DEBUG) RNS.log("Pruning session", RNS.LOG_DEBUG)
with contextlib.suppress(ValueError): with contextlib.suppress(ValueError): self.sessions.remove(self)
self.sessions.remove(self) with contextlib.suppress(Exception): self.outlet.teardown()
with contextlib.suppress(Exception):
self.outlet.teardown()
def _check_protocol_timeout(self, fail_condition: Callable[[], bool], name: str): def _check_protocol_timeout(self, fail_condition: Callable[[], bool], name: str):
timeout = True timeout = True
@@ -183,7 +181,6 @@ class ListenerSession:
def _link_closed(self, outlet: LSOutletBase): def _link_closed(self, outlet: LSOutletBase):
outlet.unset_link_closed_callback() outlet.unset_link_closed_callback()
if outlet != self.outlet: if outlet != self.outlet:
RNS.log("Link closed received from incorrect outlet", RNS.LOG_DEBUG) RNS.log("Link closed received from incorrect outlet", RNS.LOG_DEBUG)
return return
@@ -200,11 +197,14 @@ class ListenerSession:
if self.state not in [LSState.LSSTATE_WAIT_IDENT, LSState.LSSTATE_WAIT_VERS]: if self.state not in [LSState.LSSTATE_WAIT_IDENT, LSState.LSSTATE_WAIT_VERS]:
self._protocol_error(LSState.LSSTATE_WAIT_IDENT.name) self._protocol_error(LSState.LSSTATE_WAIT_IDENT.name)
if not self.allow_all and identity.hash not in self.allowed_identity_hashes and identity.hash not in self.allowed_file_identity_hashes: if self.allow_all or identity.hash in self.allowed_identity_hashes or identity.hash in self.allowed_file_identity_hashes:
self.terminate("Identity is not allowed.") self.authenticated = True
self.remote_identity = identity
self._set_state(LSState.LSSTATE_WAIT_VERS)
self.remote_identity = identity else:
self._set_state(LSState.LSSTATE_WAIT_VERS) self.authenticated = False
self.terminate("Identity not allowed")
@classmethod @classmethod
async def pump_all(cls) -> True: async def pump_all(cls) -> True:
@@ -241,8 +241,7 @@ class ListenerSession:
if compressed_length < max_data_len and compressed_length < chunk_segment_length: if compressed_length < max_data_len and compressed_length < chunk_segment_length:
comp_success = True comp_success = True
break break
else: else: comp_try += 1
comp_try += 1
if comp_success: if comp_success:
diff = max_data_len - len(compressed_chunk) diff = max_data_len - len(compressed_chunk)
@@ -255,10 +254,8 @@ class ListenerSession:
return comp_success, processed_length, chunk return comp_success, processed_length, chunk
try: try:
if self.state != LSState.LSSTATE_RUNNING: if self.state != LSState.LSSTATE_RUNNING: return False
return False elif not self.channel.is_ready_to_send(): return False
elif not self.channel.is_ready_to_send():
return False
elif len(self.stderr_buf) > 0: elif len(self.stderr_buf) > 0:
comp_success, processed_length, data = compress_adaptive(self.stderr_buf) comp_success, processed_length, data = compress_adaptive(self.stderr_buf)
self.stderr_buf = self.stderr_buf[processed_length:] self.stderr_buf = self.stderr_buf[processed_length:]
@@ -267,8 +264,7 @@ class ListenerSession:
msg = protocol.StreamDataMessage(protocol.StreamDataMessage.STREAM_ID_STDERR, msg = protocol.StreamDataMessage(protocol.StreamDataMessage.STREAM_ID_STDERR,
data, send_eof, comp_success) data, send_eof, comp_success)
self.send(msg) self.send(msg)
if send_eof: if send_eof: self.stderr_eof_sent = True
self.stderr_eof_sent = True
return True return True
elif len(self.stdout_buf) > 0: elif len(self.stdout_buf) > 0:
comp_success, processed_length, data = compress_adaptive(self.stdout_buf) comp_success, processed_length, data = compress_adaptive(self.stdout_buf)
@@ -278,8 +274,7 @@ class ListenerSession:
msg = protocol.StreamDataMessage(protocol.StreamDataMessage.STREAM_ID_STDOUT, msg = protocol.StreamDataMessage(protocol.StreamDataMessage.STREAM_ID_STDOUT,
data, send_eof, comp_success) data, send_eof, comp_success)
self.send(msg) self.send(msg)
if send_eof: if send_eof: self.stdout_eof_sent = True
self.stdout_eof_sent = True
return True return True
elif self.return_code is not None and not self.return_code_sent: elif self.return_code is not None and not self.return_code_sent:
msg = protocol.CommandExitedMessage(self.return_code) msg = protocol.CommandExitedMessage(self.return_code)
@@ -295,22 +290,32 @@ class ListenerSession:
def _terminate_process(self): def _terminate_process(self):
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
if self.process and self.process.running: if self.process and self.process.running: self.process.terminate()
self.process.terminate()
def _start_cmd(self, cmdline: [str], pipe_stdin: bool, pipe_stdout: bool, pipe_stderr: bool, tcflags: [any], def _start_cmd(self, cmdline: [str], pipe_stdin: bool, pipe_stdout: bool, pipe_stderr: bool, tcflags: [any],
term: str | None, rows: int, cols: int, hpix: int, vpix: int): term: str | None, rows: int, cols: int, hpix: int, vpix: int):
if self.terminated:
RNS.log(f"Attempt to start command on terminated session: {cmdline}", RNS.LOG_WARNING)
return
if not self.authenticated:
RNS.log(f"Attempt to start command on unauthenticated session: {cmdline}", RNS.LOG_WARNING)
self.terminate("Invalid state")
return
if self.state != LSState.LSSTATE_RUNNING:
RNS.log(f"Attempt to start command on session not in running state: {cmdline}", RNS.LOG_WARNING)
self.terminate("Invalid state")
return
self.cmdline = self.default_command self.cmdline = self.default_command
if not self.allow_remote_command and cmdline and len(cmdline) > 0: if not self.allow_remote_command and cmdline and len(cmdline) > 0:
self.terminate("Remote command line not allowed by listener") self.terminate("Remote command line not allowed by listener")
return return
if self.remote_cmd_as_args and cmdline and len(cmdline) > 0: if self.remote_cmd_as_args and cmdline and len(cmdline) > 0: self.cmdline.extend(cmdline)
self.cmdline.extend(cmdline) elif cmdline and len(cmdline) > 0: self.cmdline = cmdline
elif cmdline and len(cmdline) > 0:
self.cmdline = cmdline
self.stdin_is_pipe = pipe_stdin self.stdin_is_pipe = pipe_stdin
self.stdout_is_pipe = pipe_stdout self.stdout_is_pipe = pipe_stdout
@@ -318,11 +323,8 @@ class ListenerSession:
self.tcflags = tcflags self.tcflags = tcflags
self.term = term self.term = term
def stdout(data: bytes): def stdout(data: bytes): self.stdout_buf.extend(data)
self.stdout_buf.extend(data) def stderr(data: bytes): self.stderr_buf.extend(data)
def stderr(data: bytes):
self.stderr_buf.extend(data)
try: try:
self.process = process.CallbackSubprocess(argv=self.cmdline, self.process = process.CallbackSubprocess(argv=self.cmdline,
@@ -347,22 +349,28 @@ class ListenerSession:
self.cols = cols self.cols = cols
self.hpix = hpix self.hpix = hpix
self.vpix = vpix self.vpix = vpix
with contextlib.suppress(Exception): with contextlib.suppress(Exception): self.process.set_winsize(rows, cols, hpix, vpix)
self.process.set_winsize(rows, cols, hpix, vpix)
def _received_stdin(self, data: bytes, eof: bool): def _received_stdin(self, data: bytes, eof: bool):
if data and len(data) > 0: if data and len(data) > 0: self.process.write(data)
self.process.write(data) if eof: self.process.close_stdin()
if eof:
self.process.close_stdin()
def _handle_message(self, message: RNS.MessageBase): def _handle_message(self, message: RNS.MessageBase):
if self.terminated:
RNS.log(f"Received packet, but session was terminated", RNS.LOG_ERROR)
return
if self.state > LSState.LSSTATE_RUNNING:
RNS.log(f"Received packet, but in state {self.state.name}", RNS.LOG_ERROR)
return
if self.state == LSState.LSSTATE_WAIT_IDENT: if self.state == LSState.LSSTATE_WAIT_IDENT:
# Ignore any messages until the initiator has identified to avoid race conditions # Ignore any messages until the initiator has identified to avoid race conditions
# between identity announcement and early protocol messages. # between identity announcement and early protocol messages.
RNS.log("Ignoring message while waiting for identification", RNS.LOG_DEBUG) RNS.log("Ignoring message while waiting for identification", RNS.LOG_DEBUG)
return return
if self.state == LSState.LSSTATE_WAIT_VERS: elif self.state == LSState.LSSTATE_WAIT_VERS:
if not self.authenticated:
self.terminate("Invalid state")
return
if not isinstance(message, protocol.VersionInfoMessage): if not isinstance(message, protocol.VersionInfoMessage):
self._protocol_error(self.state.name) self._protocol_error(self.state.name)
return return
@@ -374,6 +382,9 @@ class ListenerSession:
self._set_state(LSState.LSSTATE_WAIT_CMD) self._set_state(LSState.LSSTATE_WAIT_CMD)
return return
elif self.state == LSState.LSSTATE_WAIT_CMD: elif self.state == LSState.LSSTATE_WAIT_CMD:
if not self.authenticated:
self.terminate("Invalid state")
return
if not isinstance(message, protocol.ExecuteCommandMesssage): if not isinstance(message, protocol.ExecuteCommandMesssage):
return self._protocol_error(self.state.name) return self._protocol_error(self.state.name)
RNS.log(f"Execute command message on link {self.outlet}: {message.cmdline}", RNS.LOG_VERBOSE) RNS.log(f"Execute command message on link {self.outlet}: {message.cmdline}", RNS.LOG_VERBOSE)
@@ -382,6 +393,9 @@ class ListenerSession:
message.tcflags, message.term, message.rows, message.cols, message.hpix, message.vpix) message.tcflags, message.term, message.rows, message.cols, message.hpix, message.vpix)
return return
elif self.state == LSState.LSSTATE_RUNNING: elif self.state == LSState.LSSTATE_RUNNING:
if not self.authenticated:
self.terminate("Invalid state")
return
if isinstance(message, protocol.WindowSizeMessage): if isinstance(message, protocol.WindowSizeMessage):
self._set_window_size(message.rows, message.cols, message.hpix, message.vpix) self._set_window_size(message.rows, message.cols, message.hpix, message.vpix)
elif isinstance(message, protocol.StreamDataMessage): elif isinstance(message, protocol.StreamDataMessage):
@@ -394,9 +408,6 @@ class ListenerSession:
# echo noop only on listener--used for keepalive/connectivity check # echo noop only on listener--used for keepalive/connectivity check
self.send(message) self.send(message)
return return
elif self.state in [LSState.LSSTATE_ERROR, LSState.LSSTATE_TEARDOWN]:
RNS.log(f"Received packet, but in state {self.state.name}", RNS.LOG_ERROR)
return
else: else:
self._protocol_error("unexpected message") self._protocol_error("unexpected message")
return return
@@ -405,29 +416,20 @@ class ListenerSession:
class RNSOutlet(LSOutletBase): class RNSOutlet(LSOutletBase):
def set_initiator_identified_callback(self, cb: Callable[[LSOutletBase, _TIdentity], None]): def set_initiator_identified_callback(self, cb: Callable[[LSOutletBase, _TIdentity], None]):
def inner_cb(link, identity: _TIdentity): def inner_cb(link, identity: _TIdentity): cb(self, identity)
cb(self, identity)
self.link.set_remote_identified_callback(inner_cb) self.link.set_remote_identified_callback(inner_cb)
def set_link_closed_callback(self, cb: Callable[[LSOutletBase], None]): def set_link_closed_callback(self, cb: Callable[[LSOutletBase], None]):
def inner_cb(link): def inner_cb(link): cb(self)
cb(self)
self.link.set_link_closed_callback(inner_cb) self.link.set_link_closed_callback(inner_cb)
def unset_link_closed_callback(self): def unset_link_closed_callback(self): self.link.set_link_closed_callback(None)
self.link.set_link_closed_callback(None) def teardown(self): self.link.teardown()
def teardown(self):
self.link.teardown()
@property @property
def rtt(self) -> float: def rtt(self) -> float: return self.link.rtt
return self.link.rtt
def __str__(self): def __str__(self): return f"Outlet on {self.link}"
return f"Outlet RNS Link {self.link}"
def __init__(self, link: RNS.Link): def __init__(self, link: RNS.Link):
self.link = link self.link = link
@@ -435,7 +437,5 @@ class RNSOutlet(LSOutletBase):
@staticmethod @staticmethod
def get_outlet(link: RNS.Link): def get_outlet(link: RNS.Link):
if hasattr(link, "lsoutlet"): if hasattr(link, "lsoutlet"): return link.lsoutlet
return link.lsoutlet
return RNSOutlet(link) return RNSOutlet(link)