diff --git a/CHANGELOG.md b/CHANGELOG.md index 1e36d39..24188e4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,16 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [0.2.1] - 2026-08-29 + +### Added + +- `Server.port` reports the port the transport is listening on, so + `TcpTransport(port=0)` can leave the choice to the OS. To make the assigned + port knowable, the listener now binds in `start()` rather than in the accept + task it spawns — which also means a failed bind raises out of `start()` + instead of only reaching stderr. + ## [0.2.0] - 2026-08-03 ### Added diff --git a/Cargo.lock b/Cargo.lock index 08dd1de..2560c39 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -215,7 +215,7 @@ checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" [[package]] name = "echo" -version = "0.2.0" +version = "0.2.1" dependencies = [ "arc-swap", "criterion", diff --git a/Cargo.toml b/Cargo.toml index 197014a..26ad370 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "echo" -version = "0.2.0" +version = "0.2.1" edition = "2021" description = "A fast distributed replay buffer for reinforcement learning." license = "Apache-2.0" diff --git a/docs/src/guides/transports.md b/docs/src/guides/transports.md index 2b26e30..d8f7c1f 100644 --- a/docs/src/guides/transports.md +++ b/docs/src/guides/transports.md @@ -26,7 +26,7 @@ message. Throughput is governed by the network and the drainer rate. `TcpTransport` takes: -- `port`: bind port on the server side. +- `port`: bind port on the server side. Pass `0` to let the OS assign one. - `num_threads` (default 8): worker threads used by the server to accept connections and push into SPSC queues. Bumping this only helps if you have many concurrent connections. @@ -34,6 +34,20 @@ message. Throughput is governed by the network and the drainer rate. `TcpClient` takes `max_inflight_msgs` (default 32), which caps the number of un-acked sends in flight via a `BoundedSemaphore`. +## Finding the port + +`Server.start()` binds before it returns, so `Server.port` reports where +the transport actually landed: + +```python +server = Server(example, batch_size=32, transport=TcpTransport(port=0)) +server.start() +print(server.port) # e.g. 42201 — publish this to your clients +``` + +Passing `0` and reading the port back lets you run several servers on a +host without hand-allocating port numbers. + ## Backpressure model Client sends, server receives, server pushes into the connection's SPSC diff --git a/pyproject.toml b/pyproject.toml index 81bb457..c6cbf34 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "maturin" [project] name = "id-echo" -version = "0.2.0" +version = "0.2.1" description = "A fast distributed replay buffer for reinforcement learning." readme = "README.md" license = { file = "LICENSE" } diff --git a/python/echo/echo.pyi b/python/echo/echo.pyi index 0d6c69a..f401026 100644 --- a/python/echo/echo.pyi +++ b/python/echo/echo.pyi @@ -48,6 +48,8 @@ class _Server: producer_queue_size: int = 8, pin_host_memory: bool = False, ) -> None: ... + @property + def port(self) -> int: ... def start(self) -> None: ... def sample(self) -> tuple[list[np.ndarray[Any, np.dtype[np.uint8]]], SampleInfo] | None: ... def submit(self, data: Sequence[bytes]) -> None: ... diff --git a/python/echo/server.py b/python/echo/server.py index f9cea91..72e9f70 100644 --- a/python/echo/server.py +++ b/python/echo/server.py @@ -78,6 +78,11 @@ def start(self) -> None: """Start the transport (bind port, accept connections).""" self._server.start() + @property + def port(self) -> int: + """The port the transport is listening on, valid after ``start()``.""" + return self._server.port + def sample(self) -> Sample | None: """ Block until a batch is ready. Returns None on shutdown. diff --git a/python/tests/test_tcp.py b/python/tests/test_tcp.py index d38340d..16f6cab 100644 --- a/python/tests/test_tcp.py +++ b/python/tests/test_tcp.py @@ -148,3 +148,21 @@ def test_byte_by_byte_payload_is_reassembled(self, make_server): sample = server.sample() assert sample is not None np.testing.assert_array_equal(sample.batch["obs"], [[1, 2, 3, 4]]) + + +class TestPort: + def test_ephemeral_port_is_reported_and_listening(self): + example = {"obs": np.zeros((4,), dtype=np.float32)} + server = Server(example, batch_size=10, transport=TcpTransport(port=0)) + server.start() + try: + assert server.port != 0 + wait_for_listen(server.port) + finally: + server.close() + + def test_port_before_start_raises(self): + example = {"obs": np.zeros((4,), dtype=np.float32)} + server = Server(example, batch_size=10, transport=TcpTransport(port=0)) + with pytest.raises(RuntimeError, match="not started"): + server.port diff --git a/src/py_bindings.rs b/src/py_bindings.rs index 1470638..47082c9 100644 --- a/src/py_bindings.rs +++ b/src/py_bindings.rs @@ -195,6 +195,16 @@ impl PyServer { .map_err(|e| PyRuntimeError::new_err(format!("failed to start transport: {e}"))) } + /// The port the transport is listening on, once started. + #[getter] + fn port(&self) -> PyResult { + let t = self.transport.as_ref().ok_or_else(|| { + PyRuntimeError::new_err("no transport configured; pass TcpTransport to __init__") + })?; + t.port() + .ok_or_else(|| PyRuntimeError::new_err("transport not started; call start() first")) + } + /// Block until a batch is ready. /// /// Returns `(arrays, info)` where `arrays` is a list of uint8 numpy views diff --git a/src/transport/mod.rs b/src/transport/mod.rs index 04c2eea..d4736d4 100644 --- a/src/transport/mod.rs +++ b/src/transport/mod.rs @@ -11,4 +11,8 @@ pub use tcp::TcpTransport; pub trait Transport: Send + Sync { fn start(&self) -> Result<(), Box>; fn shutdown(&self); + + fn port(&self) -> Option { + None + } } diff --git a/src/transport/tcp.rs b/src/transport/tcp.rs index 0e4deff..664aedd 100644 --- a/src/transport/tcp.rs +++ b/src/transport/tcp.rs @@ -40,6 +40,7 @@ pub fn encode_handshake(specs: &[ArraySpec]) -> Vec { struct State { _runtime: tokio::runtime::Runtime, shutdown_tx: Option>, + port: u16, } pub struct TcpTransport { @@ -74,6 +75,13 @@ impl super::Transport for TcpTransport { return Err("TCP server already started".into()); } + // Bind here rather than inside the spawned task: a `start` that + // returns then means the port is bound. + let addr: std::net::SocketAddr = ([0, 0, 0, 0], self.port).into(); + let listener = std::net::TcpListener::bind(addr)?; + listener.set_nonblocking(true)?; + let port = listener.local_addr()?.port(); + // Bound transport threads so the drainer pool has CPU left. let rt = tokio::runtime::Builder::new_multi_thread() .worker_threads(self.num_threads) @@ -84,10 +92,9 @@ impl super::Transport for TcpTransport { let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); let drainer_pool = self.drainer_pool.clone(); let specs = self.specs.clone(); - let addr: std::net::SocketAddr = ([0, 0, 0, 0], self.port).into(); rt.spawn(async move { - if let Err(e) = run_server(drainer_pool, specs, addr, shutdown_rx).await { + if let Err(e) = run_server(drainer_pool, specs, listener, shutdown_rx).await { eprintln!("TCP server error: {e}"); } }); @@ -95,6 +102,7 @@ impl super::Transport for TcpTransport { *state = Some(State { _runtime: rt, shutdown_tx: Some(shutdown_tx), + port, }); Ok(()) } @@ -108,17 +116,22 @@ impl super::Transport for TcpTransport { } *state = None; } + + fn port(&self) -> Option { + self.state.lock().as_ref().map(|s| s.port) + } } async fn run_server( drainer_pool: Arc, specs: Vec, - addr: std::net::SocketAddr, + listener: std::net::TcpListener, shutdown_rx: tokio::sync::oneshot::Receiver<()>, ) -> Result<(), Box> { let handshake = encode_handshake(&specs); let payload_size: usize = specs.iter().map(|s| s.num_bytes()).sum(); - let listener = TcpListener::bind(addr).await?; + // Registers with this runtime's reactor, so it has to happen in-task. + let listener = TcpListener::from_std(listener)?; tokio::select! { _ = async { diff --git a/tests/tcp.rs b/tests/tcp.rs index 214f8ca..47ab38e 100644 --- a/tests/tcp.rs +++ b/tests/tcp.rs @@ -254,3 +254,27 @@ async fn test_server_clean_disconnect_between_frames() { pool.shutdown(); cleanup(transport, pool).await; } + +#[tokio::test(flavor = "multi_thread")] +async fn test_port_reports_the_os_assigned_bind() { + let specs = vec![ArraySpec::new(vec![4], 1)]; + let (_store, pool) = make_pool(&specs, 4, 2); + let transport = TcpTransport::new(1, pool.clone(), specs.clone(), 0); + assert_eq!(transport.port(), None, "nothing is bound before start"); + + transport.start().unwrap(); + let port = transport.port().expect("a started transport has a port"); + assert_ne!( + port, 0, + "port 0 asks for an assignment, it is never the answer" + ); + + // The number is only worth anything if that's where connections land. + let mut stream = connect_with_retry(port).await; + let expected = encode_handshake(&specs); + assert_eq!(read_exact(&mut stream, expected.len()).await, expected); + + drop(stream); + pool.shutdown(); + cleanup(transport, pool).await; +} diff --git a/uv.lock b/uv.lock index c510412..82aec38 100644 --- a/uv.lock +++ b/uv.lock @@ -423,7 +423,7 @@ wheels = [ [[package]] name = "id-echo" -version = "0.2.0" +version = "0.2.1" source = { editable = "." } dependencies = [ { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },