From 275ccff2c6c8580fa763796424573e6caad77e3f Mon Sep 17 00:00:00 2001 From: Andrew Tridgell Date: Thu, 10 Sep 2026 14:13:33 +1000 Subject: [PATCH 1/2] Preserve TLS write buffers across WebSocket retries Appending video while SSL_write was pending could move its input buffer or change its length, disconnecting TLS viewers with a bad write retry error. Retain one stable TLS record until the write succeeds and rebuild all WebSocket users when its layout changes. Add real TLS 1.2 and TLS 1.3 backpressure regressions that queue additional video before the pending write completes. --- Makefile | 4 + tests/test_websocket_tls_retry.py | 164 ++++++++++++++++++++++++++++++ websocket.cpp | 9 +- websocket.h | 3 + 4 files changed, 179 insertions(+), 1 deletion(-) create mode 100644 tests/test_websocket_tls_retry.py diff --git a/Makefile b/Makefile index 1429f55..bb2579d 100644 --- a/Makefile +++ b/Makefile @@ -71,6 +71,10 @@ mavlink.o: mavlink.cpp mavlink.h $(MAVLINK_DIR)/protocol.h @echo "Compiling $<..." $(CXX) $(CXXFLAGS) -Wno-stringop-truncation -c $< -o $@ +# WebSocket is allocated in both MAVLink and video code. Rebuild all users +# when its layout or inline accessors change, including transitive includes. +supportproxy.o mavlink.o binlog.o video.o videoview.o websocket.o: websocket.h + # Dependencies. mavlink.h includes keydb.h, so any object that pulls in # mavlink.h transitively depends on keydb.h too. supportproxy.o: supportproxy.cpp mavlink.h util.h keydb.h conntdb.h tlog.h binlog.h session.h cleanup.h websocket.h video.h videots.h diff --git a/tests/test_websocket_tls_retry.py b/tests/test_websocket_tls_retry.py new file mode 100644 index 0000000..c6d8d23 --- /dev/null +++ b/tests/test_websocket_tls_retry.py @@ -0,0 +1,164 @@ +"""Exercise real TLS writes while the WebSocket output queue grows under backpressure.""" +import base64 +from pathlib import Path +import select +import socket +import ssl +import subprocess + +import pytest + + +HARNESS = r''' +#include "websocket.h" +#include +#include +#include +#include +#include +#include +#include + +int main() +{ + setbuf(stdout, nullptr); + alarm(20); + int listener = socket(AF_INET, SOCK_STREAM, 0); + assert(listener >= 0); + sockaddr_in address {}; + address.sin_family = AF_INET; + address.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + assert(bind(listener, reinterpret_cast(&address), sizeof(address)) == 0); + socklen_t length = sizeof(address); + assert(getsockname(listener, reinterpret_cast(&address), &length) == 0); + assert(listen(listener, 1) == 0); + printf("PORT %u\n", ntohs(address.sin_port)); + int fd = accept(listener, nullptr, nullptr); + assert(fd >= 0); + close(listener); + int size = 4096; + assert(setsockopt(fd, SOL_SOCKET, SO_SNDBUF, &size, sizeof(size)) == 0); + assert(fcntl(fd, F_SETFL, O_NONBLOCK) == 0); + pollfd event {fd, POLLIN, 0}; + assert(poll(&event, 1, 5000) == 1); + WebSocket ws(fd); + assert(ws.is_SSL()); + char input[32]; + ssize_t count; + do { + count = ws.recv(input, sizeof(input)); + assert(count >= 0); + usleep(1000); + } while (count == 0); + assert(count == 5 && memcmp(input, "start", 5) == 0); + std::vector first(65536, 0x41); + std::vector second(196608, 0x42); + assert(ws.send(first.data(), first.size()) == ssize_t(first.size())); + assert(ws.has_pending_output()); // force a nonblocking TLS write retry + // Appending a larger message changes tx's length and forces a reallocation + // while SSL_write is pending. The pre-fix implementation disconnects here. + assert(ws.send(second.data(), second.size()) == ssize_t(second.size())); + printf("QUEUED\n"); + for (;;) { + assert(ws.flush()); + count = ws.recv(input, sizeof(input)); + assert(count >= 0); + if (count == 4 && memcmp(input, "done", 4) == 0) break; + usleep(1000); + } + assert(!ws.has_pending_output()); + close(fd); + printf("PASS\n"); +} +''' + + +@pytest.fixture(scope='module') +def tls_writer(tmp_path_factory): + directory = tmp_path_factory.mktemp('tls-write-retry') + source = directory / 'writer.cpp' + source.write_text(HARNESS) + repo = Path(__file__).resolve().parents[1] + binary = directory / 'writer' + subprocess.run(['g++', '-std=c++11', '-O2', '-Wall', '-Wextra', '-Werror', + '-I' + str(repo), str(source), str(repo / 'websocket.cpp'), + '-lssl', '-lcrypto', '-o', str(binary)], check=True) + subprocess.run(['openssl', 'req', '-x509', '-newkey', 'rsa:2048', '-nodes', + '-keyout', str(directory / 'privkey.pem'), + '-out', str(directory / 'fullchain.pem'), + '-days', '1', '-subj', '/CN=localhost'], + check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) + return binary + + +def wait_line(process, marker): + while True: + assert select.select([process.stdout], [], [], 10)[0], marker + line = process.stdout.readline() + assert line, f'writer exited before {marker}' + if line.startswith(marker): + return line + + +def read_exact(connection, length): + result = b'' + while len(result) < length: + data = connection.recv(length - len(result)) + assert data, 'TLS writer disconnected' + result += data + return result + + +def frame(connection): + header = read_exact(connection, 2) + assert header[0] == 0x82 and header[1] & 0x80 == 0 + length = header[1] & 127 + if length == 126: + length = int.from_bytes(read_exact(connection, 2), 'big') + elif length == 127: + length = int.from_bytes(read_exact(connection, 8), 'big') + return read_exact(connection, length) + + +def send_frame(connection, payload): + connection.sendall(bytes([0x82, 0x80 | len(payload)]) + b'\0' * 4 + payload) + + +@pytest.mark.parametrize('version', [ssl.TLSVersion.TLSv1_2, ssl.TLSVersion.TLSv1_3]) +def test_tls_retry_preserves_queued_video(tls_writer, version): + process = subprocess.Popen([str(tls_writer)], cwd=tls_writer.parent, + stdout=subprocess.PIPE, stderr=subprocess.STDOUT, + bufsize=0) + try: + port = int(wait_line(process, b'PORT ').split()[1]) + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + context.check_hostname = False + context.verify_mode = ssl.CERT_NONE + context.minimum_version = context.maximum_version = version + with socket.socket() as raw: + # A small receive window makes the first SSL_write wait while the + # next WebSocket message is queued, even on a fast loopback link. + raw.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, 4096) + raw.settimeout(10) + raw.connect(('127.0.0.1', port)) + with context.wrap_socket(raw, server_hostname='localhost') as connection: + key = base64.b64encode(b'x' * 16).decode() + connection.sendall(('GET /v1 HTTP/1.1\r\nHost: localhost\r\n' + 'Upgrade: websocket\r\nConnection: Upgrade\r\n' + f'Sec-WebSocket-Key: {key}\r\n' + 'Sec-WebSocket-Version: 13\r\n\r\n').encode()) + response = b'' + while not response.endswith(b'\r\n\r\n'): + response += read_exact(connection, 1) + assert b'101 Switching Protocols' in response + send_frame(connection, b'start') + wait_line(process, b'QUEUED') + assert frame(connection) == b'A' * 65536 + assert frame(connection) == b'B' * 196608 + send_frame(connection, b'done') + wait_line(process, b'PASS') + assert process.wait(timeout=5) == 0 + finally: + if process.poll() is None: + process.kill() + process.wait(timeout=5) diff --git a/websocket.cpp b/websocket.cpp index 6e9fc48..93e6c2b 100644 --- a/websocket.cpp +++ b/websocket.cpp @@ -514,7 +514,13 @@ bool WebSocket::flush(void) const size_t remain = tx.size() - tx_sent; ssize_t wret; if (_is_SSL && ssl) { - wret = SSL_write(ssl, &tx[tx_sent], int(remain)); + if (tls_tx.empty()) { + // Keep one TLS record stable until SSL_write succeeds. The + // main output queue can grow while a nonblocking write waits. + const size_t chunk = remain < 16384U ? remain : 16384U; + tls_tx.assign(tx.begin() + tx_sent, tx.begin() + tx_sent + chunk); + } + wret = SSL_write(ssl, tls_tx.data(), int(tls_tx.size())); if (wret <= 0) { int err = SSL_get_error(ssl, wret); if (err == SSL_ERROR_WANT_WRITE || err == SSL_ERROR_WANT_READ) { @@ -534,6 +540,7 @@ bool WebSocket::flush(void) return false; } } + tls_tx.clear(); tx_sent += size_t(wret); } tx.clear(); diff --git a/websocket.h b/websocket.h index c8dd70f..c00a63b 100644 --- a/websocket.h +++ b/websocket.h @@ -71,6 +71,9 @@ class WebSocket { std::vector tx; // framed output size_t tx_sent = 0; // how much of tx has reached the socket + // SSL_write retries must use the same bytes, address and length even if + // queue_frame() appends more output and reallocates tx in the meantime. + std::vector tls_tx; std::string req_target; From 8053ccaf97af5e8977f033006c71b844cf33d8af Mon Sep 17 00:00:00 2001 From: Andrew Tridgell Date: Thu, 10 Sep 2026 14:57:03 +1000 Subject: [PATCH 2/2] Fix concurrent TDB test isolation and run TLS retry tests in CI TDB refuses duplicate handles in one process, so the threaded proxy simulation intermittently fails with EBUSY. Use a spawned process to match production, check its exit status and exact counter updates, and include both TLS retry regressions in the CI test list. --- scripts/run_tests.py | 1 + tests/webadmin/test_concurrent.py | 42 ++++++++++++++++++++----------- 2 files changed, 29 insertions(+), 14 deletions(-) diff --git a/scripts/run_tests.py b/scripts/run_tests.py index dae8de3..90d96ec 100755 --- a/scripts/run_tests.py +++ b/scripts/run_tests.py @@ -37,6 +37,7 @@ 'tests/test_websocket_decode.py', 'tests/test_ws_handshake_ordering.py', 'tests/test_websocket_framing.py', + 'tests/test_websocket_tls_retry.py', 'tests/test_bidi_video_preauth.py', 'tests/test_video_schema.py', 'tests/test_video_ports.py', diff --git a/tests/webadmin/test_concurrent.py b/tests/webadmin/test_concurrent.py index 8738938..b5320f5 100644 --- a/tests/webadmin/test_concurrent.py +++ b/tests/webadmin/test_concurrent.py @@ -2,7 +2,7 @@ when the web UI also writes to the same entry. Simulate the proxy's connection-close path (read entry, bump counters, -save) running in a separate thread while the web UI POSTs a name change +save) running in a separate process while the web UI POSTs a name change to /admin/. Both writes use TDB transactions, so they must serialize, and both updates must land in the final record. @@ -10,7 +10,7 @@ under TDB's per-DB lock; if either path bypassed the transaction we'd expect to see the counters revert to their initial values. """ -import threading +import multiprocessing import time import keydb_lib @@ -20,7 +20,7 @@ fetch_entry, login_as) -def _bump_counters(keydb_path, port2, iterations, stop_event): +def _bump_counters(keydb_path, port2, iterations, stop_event, ready, completed): """Imitate supportproxy.cpp's connection-close counter update.""" for _ in range(iterations): if stop_event.is_set(): @@ -41,6 +41,8 @@ def _bump_counters(keydb_path, port2, iterations, stop_event): raise finally: db.close() + completed.value += 1 + ready.set() time.sleep(0.001) @@ -48,29 +50,41 @@ class TestConcurrent: def test_ui_rename_does_not_lose_counter_updates(self, client, keydb_path): login_as(client, BOB_PORT1, BOB_PASS) - stop = threading.Event() - # bump alice's counters from a background thread while we POST to /admin - t = threading.Thread(target=_bump_counters, - args=(keydb_path, ALICE_PORT2, 200, stop)) - t.start() + # TDB rejects two handles for the same database in one process with + # EBUSY. Match the separate proxy/web workers and avoid inheriting any + # TDB state from the test runner. + ctx = multiprocessing.get_context('spawn') + stop, ready = ctx.Event(), ctx.Event() + completed = ctx.Value('I', 0) + worker = ctx.Process(target=_bump_counters, + args=(keydb_path, ALICE_PORT2, 200, stop, + ready, completed)) + worker.start() try: + assert ready.wait(10), "counter worker did not start" for i in range(20): - client.post('/admin/' + str(ALICE_PORT2), data={ + response = client.post('/admin/' + str(ALICE_PORT2), data={ 'name': 'alice rename %d' % i, 'port1': ALICE_PORT1, 'submit': 'Save', }) + assert response.status_code == 302 finally: stop.set() - t.join(timeout=10) + worker.join(timeout=10) + if worker.is_alive(): + worker.terminate() + worker.join(timeout=10) + assert worker.exitcode == 0, "counter worker failed or timed out" # final record must reflect both axes of mutation: # counters were bumped by the worker (>= 1), and the final name # we POSTed is the latest of our 20 attempts. ke = fetch_entry(keydb_path, ALICE_PORT2) assert ke is not None - assert ke.count1 >= 1, "background counter increments lost" - assert ke.count2 >= 2 - assert ke.connections >= 1 - assert ke.name.startswith('alice rename'), \ + assert completed.value >= 1, "counter worker made no updates" + assert ke.count1 == completed.value, "background counter increments lost" + assert ke.count2 == 2 * completed.value + assert ke.connections == completed.value + assert ke.name == 'alice rename 19', \ "UI rename did not land (got %r)" % ke.name