diff --git a/src/channels/adapter.rs b/src/channels/adapter.rs index cc833fa..fd4cc58 100644 --- a/src/channels/adapter.rs +++ b/src/channels/adapter.rs @@ -21,11 +21,12 @@ use bytes::Bytes; use tokio::io::AsyncReadExt; use tracing::{debug, warn}; -use crate::core::auth::AuthContext; +use crate::core::auth::{AuthContext, Identity}; use crate::core::types::{Connection, HandlerError, ProtocolHandler, StreamError}; use super::manager::ChannelManager; use super::mux::MuxRunner; +use super::policy::ChannelLifecyclePolicy; use super::wire::CHUNK_HEADER_LEN; /// The ALPN the `ChannelsAdapter` registers on. @@ -55,17 +56,23 @@ pub struct ChannelsAdapter { install_channel_zero: InstallChannelZero, max_channels: usize, buffer_cap: usize, + policy: Arc, } impl ChannelsAdapter { /// Construct with the `install_channel_zero` hook (provided by - /// `channels-call`). Default max channels (256) and buffer cap - /// (1 MiB). - pub fn new(install_channel_zero: InstallChannelZero) -> Self { + /// `channels-call`), a `ChannelLifecyclePolicy`, and default + /// limits (256 channels, 1 MiB buffer). The policy is used to + /// decrement per-identity counts on connection drop (ADR-047 §7). + pub fn new( + install_channel_zero: InstallChannelZero, + policy: Arc, + ) -> Self { Self { install_channel_zero, max_channels: super::manager::DEFAULT_MAX_CHANNELS, buffer_cap: super::reassembly::DEFAULT_BUFFER_CAP, + policy, } } @@ -74,11 +81,13 @@ impl ChannelsAdapter { install_channel_zero: InstallChannelZero, max_channels: usize, buffer_cap: usize, + policy: Arc, ) -> Self { Self { install_channel_zero, max_channels, buffer_cap, + policy, } } @@ -92,16 +101,23 @@ impl ChannelsAdapter { async fn run_demux_loop( manager: &ChannelManager, reader: Box, + policy: &Arc, ) { - Self::run_demux_loop_for_client(manager, reader).await; + Self::run_demux_loop_for_client(manager, reader, Some(policy)).await; } /// The demux loop, public for `ChannelClient` to call. Reads /// 8-byte chunk headers and routes payloads to the `ChannelManager`. /// Ends on transport EOF, clearing the channel map (REQ-CH-02). + /// + /// If `policy` is `Some`, each drained channel's opener is + /// decremented via `policy.on_close` (ADR-047 §7 — connection-drop + /// teardown path). The connect side passes `None` (it does not + /// enforce a per-identity cap). pub async fn run_demux_loop_for_client( manager: &ChannelManager, reader: Box, + policy: Option<&Arc>, ) { let mut reader = reader; let mut header_buf = [0u8; CHUNK_HEADER_LEN]; @@ -157,6 +173,16 @@ impl ChannelsAdapter { } } let drained = manager.clear_all(); + if let Some(policy) = policy { + for (_channel_id, opener_id) in &drained { + let opener = Identity { + id: opener_id.clone(), + scopes: vec![], + resources: Default::default(), + }; + policy.on_close(&opener); + } + } debug!( channels = drained.len(), "demux: cleared channel map on transport EOF" @@ -192,6 +218,7 @@ impl ProtocolHandler for ChannelsAdapter { self.max_channels, self.buffer_cap, connection.remote_addr(), + super::manager::ChannelSide::Accept, ); // Spawn the mux runner BEFORE installing channel 0 — @@ -222,8 +249,10 @@ impl ProtocolHandler for ChannelsAdapter { (self.install_channel_zero)(manager.clone(), channel0_conn, auth.clone()); // 4. Run the demux loop (read side). This blocks until - // transport EOF, then clears the channel map (REQ-CH-02). - Self::run_demux_loop(&manager, Box::new(reader)).await; + // transport EOF, then clears the channel map (REQ-CH-02) + // and decrements the per-identity policy for each drained + // channel (ADR-047 §7 — connection-drop teardown path). + Self::run_demux_loop(&manager, Box::new(reader), &self.policy).await; Ok(()) } @@ -262,7 +291,8 @@ mod tests { let demux_manager = manager.clone(); let _demux_task = tokio::spawn(async move { - ChannelsAdapter::run_demux_loop_for_client(&demux_manager, Box::new(server_read)).await; + ChannelsAdapter::run_demux_loop_for_client(&demux_manager, Box::new(server_read), None) + .await; }); let oversized_len = super::super::wire::MAX_CHUNK_LEN + 1; @@ -320,7 +350,8 @@ mod tests { let demux_manager = manager.clone(); let _demux_task = tokio::spawn(async move { - ChannelsAdapter::run_demux_loop_for_client(&demux_manager, Box::new(server_read)).await; + ChannelsAdapter::run_demux_loop_for_client(&demux_manager, Box::new(server_read), None) + .await; }); let mut header = [0u8; 8]; diff --git a/src/channels/client.rs b/src/channels/client.rs index 1147f5f..bc1f5a4 100644 --- a/src/channels/client.rs +++ b/src/channels/client.rs @@ -27,8 +27,9 @@ use crate::core::types::{Connection, StreamError}; use crate::protocol::connection::CallConnection; use crate::protocol::wire::ResponseEnvelope; -use super::manager::ChannelManager; +use super::manager::{ChannelManager, ChannelSide}; use super::mux::MuxRunner; +use super::reassembly::{MpscRecvStream, MpscSendStream}; /// The client-side handle for a channels connection. Constructed via /// [`ChannelClient::from_connection`] from an established @@ -76,7 +77,13 @@ impl ChannelClient { } }); - let manager = ChannelManager::with_defaults(mux_handle, remote_addr); + let manager = ChannelManager::new( + mux_handle, + super::manager::DEFAULT_MAX_CHANNELS, + super::reassembly::DEFAULT_BUFFER_CAP, + remote_addr, + ChannelSide::Connect, + ); // Install channel 0 — the call adapter's read/write halves. let (channel0_send, channel0_recv) = manager @@ -125,6 +132,7 @@ impl ChannelClient { super::adapter::ChannelsAdapter::run_demux_loop_for_client( &demux_manager, Box::new(reader), + None, ) .await; }); @@ -159,6 +167,43 @@ impl ChannelClient { } } + /// Open a data channel by calling the per-ALPN open op on channel 0 + /// and adopting the resulting `channel_id` (ADR-047 §5 odd/even + /// split). The connect side calls the open op; the accept side + /// allocates the `channel_id` (even). The connect side then adopts + /// the `channel_id` via [`ChannelManager::adopt_channel`] to install + /// local routing state (mux write half + demux read half). + /// + /// Returns the `channel_id`, the `MpscSendStream` (write half), and + /// the `MpscRecvStream` (read half). The caller can build a + /// `Connection` from these via `channel_source` and + /// `Connection::from_source`. + /// + /// `alpn` is the data-plane ALPN (e.g. `alknet/tty`), used for + /// observability in the local manager. + pub async fn open_channel( + &self, + operation_id: &str, + input: Value, + alpn: &str, + ) -> Result<(u32, MpscSendStream, MpscRecvStream), String> { + let response = self.call_open_op(operation_id, input).await; + let out = response + .result + .map_err(|e| format!("open op failed: {e:?}"))?; + let channel_id = out + .get("channel_id") + .and_then(|v| v.as_u64()) + .ok_or_else(|| "open op response missing channel_id".to_string())? + as u32; + + self.manager + .adopt_channel(channel_id, alpn, None) + .await + .map_err(|e| format!("adopt_channel failed: {e}")) + .map(|(send, recv)| (channel_id, send, recv)) + } + /// Take the `CallConnection` — used by the consumer to register /// imported ops (`from_call`) on the connection's overlay. After /// this, `call_open_op` returns an error (the connection is owned @@ -173,16 +218,19 @@ mod tests { use super::*; use crate::channels::adapter::ChannelsAdapter; use crate::channels::policy::ChannelLifecyclePolicy; + use crate::channels::policy::NoCap; use crate::core::auth::{AuthContext, IdentityProvider}; use crate::core::types::Connection; use crate::protocol::connection::split_single_stream; use crate::protocol::dispatch::Dispatcher; + use crate::registry::context::OperationContext; use crate::registry::registration::{ make_handler, HandlerKind, HandlerRegistration, OperationProvenance, }; use crate::registry::spec::{AccessControl, OperationSpec, OperationType, Visibility}; use std::net::{IpAddr, Ipv4Addr, SocketAddr}; use std::sync::Arc; + use tokio::io::AsyncWriteExt; const TEST_ADDR: SocketAddr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4321); @@ -280,7 +328,10 @@ mod tests { let server_conn = Connection::from_bidi(server_end, b"alknet/channels".to_vec(), Some(TEST_ADDR)); - let adapter = ChannelsAdapter::new(make_install_channel_zero(Arc::clone(®istry))); + let adapter = ChannelsAdapter::new( + make_install_channel_zero(Arc::clone(®istry)), + Arc::new(NoCap), + ); let auth = AuthContext::anonymous(b"alknet/channels"); let _server_handle = tokio::spawn(async move { let _ = crate::core::types::ProtocolHandler::handle(&adapter, server_conn, &auth).await; @@ -319,7 +370,10 @@ mod tests { let server_conn = Connection::from_bidi(server_end, b"alknet/channels".to_vec(), Some(TEST_ADDR)); - let adapter = ChannelsAdapter::new(make_install_channel_zero(Arc::clone(®istry))); + let adapter = ChannelsAdapter::new( + make_install_channel_zero(Arc::clone(®istry)), + Arc::new(NoCap), + ); let auth = AuthContext::anonymous(b"alknet/channels"); let _server_handle = tokio::spawn(async move { let _ = crate::core::types::ProtocolHandler::handle(&adapter, server_conn, &auth).await; @@ -399,7 +453,10 @@ mod tests { let server_conn = Connection::from_bidi(server_end, b"alknet/channels".to_vec(), Some(TEST_ADDR)); - let adapter = ChannelsAdapter::new(make_install_channel_zero(Arc::clone(®istry))); + let adapter = ChannelsAdapter::new( + make_install_channel_zero(Arc::clone(®istry)), + Arc::new(NoCap), + ); let auth = AuthContext::anonymous(b"alknet/channels"); let _server_handle = tokio::spawn(async move { let _ = crate::core::types::ProtocolHandler::handle(&adapter, server_conn, &auth).await; @@ -552,7 +609,7 @@ mod tests { let server_conn = Connection::from_bidi(server_end, b"alknet/channels".to_vec(), Some(TEST_ADDR)); - let adapter = ChannelsAdapter::new(install_hook); + let adapter = ChannelsAdapter::new(install_hook, Arc::new(NoCap)); let auth = AuthContext::anonymous(b"alknet/channels"); let _server_handle = tokio::spawn(async move { let _ = crate::core::types::ProtocolHandler::handle(&adapter, server_conn, &auth).await; @@ -598,4 +655,307 @@ mod tests { "quota reserved: the open op incremented the per-identity count" ); } + + /// C-25 #4 — policy decrement on connection drop (ADR-047 §7). + /// Open channels via the manager, then call `clear_all` through the + /// demux loop with a policy — the per-identity count must be + /// decremented for each drained channel. + #[tokio::test] + async fn policy_decremented_on_connection_drop() { + use crate::channels::policy::PerIdentityChannelPolicy; + + let policy: Arc = Arc::new(PerIdentityChannelPolicy::new(2)); + let dyn_policy: Arc = + Arc::clone(&policy) as Arc; + + let (client, server) = tokio::io::duplex(64 * 1024); + let (server_read, server_write) = tokio::io::split(server); + let (mux_handle, mux_runner) = MuxRunner::new(Box::new(server_write)); + let _mux_task = tokio::spawn(async move { + let _ = mux_runner.run().await; + }); + let manager = ChannelManager::with_defaults(mux_handle, None); + + let alice = crate::core::auth::Identity { + id: "alice".to_string(), + scopes: vec![], + resources: Default::default(), + }; + let bob = crate::core::auth::Identity { + id: "bob".to_string(), + scopes: vec![], + resources: Default::default(), + }; + + assert!(dyn_policy.check_open(&alice).is_ok()); + manager + .open_channel("alknet/tty", "alice", None) + .await + .expect("open alice"); + assert!(dyn_policy.check_open(&bob).is_ok()); + manager + .open_channel("alknet/tty", "bob", None) + .await + .expect("open bob"); + + assert_eq!(policy.count_for(&alice), 1); + assert_eq!(policy.count_for(&bob), 1); + + let demux_manager = manager.clone(); + let demux_policy = Arc::clone(&dyn_policy); + let demux_task = tokio::spawn(async move { + ChannelsAdapter::run_demux_loop_for_client( + &demux_manager, + Box::new(server_read), + Some(&demux_policy), + ) + .await; + }); + + drop(client); + let _ = tokio::time::timeout(std::time::Duration::from_secs(5), demux_task).await; + + assert_eq!( + policy.count_for(&alice), + 0, + "alice decremented after connection drop" + ); + assert_eq!( + policy.count_for(&bob), + 0, + "bob decremented after connection drop" + ); + } + + /// C-25 #5 — concurrent open race (TOCTOU). Multiple concurrent + /// `open_channel` calls on a manager with `max_channels=1` must not + /// all succeed — only one should pass the re-check on insert. + #[tokio::test] + async fn concurrent_opens_respect_max_channels() { + let (_client, server) = tokio::io::duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (handle, runner) = MuxRunner::new(Box::new(writer)); + tokio::spawn(async move { + let _ = runner.run().await; + }); + let manager = Arc::new(ChannelManager::new( + handle, + 1, + 64, + None, + ChannelSide::Accept, + )); + + let m1 = Arc::clone(&manager); + let m2 = Arc::clone(&manager); + let m3 = Arc::clone(&manager); + + let (r1, r2, r3) = tokio::join!( + m1.open_channel("alknet/a", "alice", None), + m2.open_channel("alknet/b", "bob", None), + m3.open_channel("alknet/c", "carol", None), + ); + + let successes = [r1.is_ok(), r2.is_ok(), r3.is_ok()] + .iter() + .filter(|&&ok| ok) + .count(); + assert_eq!( + successes, 1, + "exactly one concurrent open should succeed with max_channels=1" + ); + assert_eq!( + manager.open_count(), + 1, + "channel count should be 1 after concurrent opens" + ); + } + + /// C-12 — `channel/close` with `channel_id: 0` is rejected. + #[tokio::test] + async fn channel_close_rejects_channel_zero() { + use crate::channels::operations::ChannelOperations; + use crate::channels::policy::NoCap; + use crate::registry::context::{AbortPolicy, ScopedPeerEnv}; + use std::collections::HashMap; + use std::sync::Arc; + + struct NoopEnv; + #[async_trait::async_trait] + impl crate::registry::env::OperationEnv for NoopEnv { + async fn invoke_with_policy( + &self, + _namespace: &str, + _operation: &str, + _input: serde_json::Value, + _parent: &OperationContext, + _policy: AbortPolicy, + ) -> ResponseEnvelope { + ResponseEnvelope::error("test", crate::protocol::wire::CallError::internal("noop")) + } + fn contains(&self, _name: &str) -> bool { + false + } + } + + let (_client, server) = tokio::io::duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (handle, runner) = MuxRunner::new(Box::new(writer)); + tokio::spawn(async move { + let _ = runner.run().await; + }); + let manager = ChannelManager::with_defaults(handle, None); + let ops = ChannelOperations::new(manager, Arc::new(NoCap)); + let mut registry = crate::registry::registration::OperationRegistry::new(); + ops.register_on(&mut registry).expect("register"); + + let handler = registry + .registration("channel/close") + .expect("close op registered") + .handler + .clone(); + let ctx = OperationContext { + request_id: "req-1".to_string(), + parent_request_id: None, + identity: None, + handler_identity: None, + forwarded_for: None, + capabilities: crate::core::types::Capabilities::new(), + metadata: HashMap::new(), + scoped_env: ScopedPeerEnv::empty(), + env: Arc::new(NoopEnv), + abort_policy: AbortPolicy::default(), + deadline: Some(std::time::Instant::now() + std::time::Duration::from_secs(30)), + internal: false, + ownership: None, + }; + let input = serde_json::json!({ "channel_id": 0 }); + + let response = match handler { + HandlerKind::Once(h) => h(input, ctx).await, + _ => panic!("expected Once handler"), + }; + + assert!( + response.result.is_err(), + "channel/close with channel_id:0 should be rejected" + ); + let err = response.result.unwrap_err(); + assert!( + err.code == "INVALID_INPUT", + "expected INVALID_INPUT, got {}", + err.code + ); + } + + /// C-08 — end-to-end channel adoption: the connect side calls an + /// open op, receives a `channel_id`, adopts it, and can write data + /// through the adopted channel to the accept side's handler. + #[tokio::test] + async fn channel_adoption_end_to_end_round_trip() { + use crate::channels::operations::{ChannelCore, OpenHandler}; + use crate::channels::policy::NoCap; + use crate::registry::spec::ChannelOpenSpec; + use tokio::io::AsyncReadExt; + + let (data_tx, mut data_rx) = tokio::sync::mpsc::channel::>(1); + + let open_handler: OpenHandler = Arc::new(move |_input, channel_conn, _auth| { + let data_tx = data_tx.clone(); + tokio::spawn(async move { + let mut bidi = channel_conn.accept_bi().await.expect("accept_bi"); + let mut buf = [0u8; 4]; + bidi.read_exact(&mut buf).await.expect("read"); + data_tx.send(buf.to_vec()).await.expect("send to channel"); + }) + }); + + let install_hook: crate::channels::adapter::InstallChannelZero = + Arc::new(move |manager, channel0_conn, auth| { + let open_handler = Arc::clone(&open_handler); + tokio::spawn(async move { + let channel0_bidi = match channel0_conn.accept_bi().await { + Ok(s) => s, + Err(_) => return, + }; + let (writer, reader) = split_single_stream(channel0_bidi); + let core = ChannelCore::new(manager, Arc::new(NoCap)); + let mut registry = crate::registry::registration::OperationRegistry::new(); + let spec = OperationSpec::new( + "channels/tty/sub", + OperationType::Sub, + Visibility::External, + serde_json::json!({ + "type": "object", + "properties": { "container": { "type": "string" } }, + "required": ["container"] + }), + serde_json::json!({ + "type": "object", + "properties": { "channel_id": { "type": "integer" } } + }), + vec![], + AccessControl::default(), + None, + ) + .with_channel_open(ChannelOpenSpec::new("alknet/tty")); + core.register_openable( + spec, + Arc::clone(&open_handler), + &mut registry, + auth.clone(), + ) + .expect("register_openable"); + let registry = Arc::new(registry); + let provider: Arc = Arc::new(NoopIdProvider); + let call_connection = Arc::new(CallConnection::new_single_stream( + channel0_conn, + Arc::clone(&writer), + )); + let dp = Dispatcher::new(registry, provider); + dp.run_loop_single_stream(call_connection, reader, writer) + .await; + }) + }); + + let (client_end, server_end) = tokio::io::duplex(64 * 1024); + let client_conn = + Connection::from_bidi(client_end, b"alknet/channels".to_vec(), Some(TEST_ADDR)); + let server_conn = + Connection::from_bidi(server_end, b"alknet/channels".to_vec(), Some(TEST_ADDR)); + + let adapter = ChannelsAdapter::new(install_hook, Arc::new(NoCap)); + let auth = AuthContext::anonymous(b"alknet/channels"); + let _server_handle = tokio::spawn(async move { + let _ = crate::core::types::ProtocolHandler::handle(&adapter, server_conn, &auth).await; + }); + + let client = ChannelClient::from_connection(client_conn) + .await + .expect("channel client init"); + + let (channel_id, mut send, _recv) = client + .open_channel( + "channels/tty/sub", + serde_json::json!({ "container": "abc" }), + "alknet/tty", + ) + .await + .expect("open_channel"); + + assert!(channel_id > 0, "channel_id should be non-zero"); + assert!(channel_id % 2 == 0, "accept-allocated ID should be even"); + + send.write_all(b"ping").await.expect("write ping"); + drop(send); + + let data = tokio::time::timeout(std::time::Duration::from_secs(5), data_rx.recv()) + .await + .expect("timed out waiting for handler data") + .expect("handler should receive data"); + assert_eq!( + &data, b"ping", + "handler received ping through adopted channel" + ); + } } diff --git a/src/channels/manager.rs b/src/channels/manager.rs index 1e25ede..03a04cf 100644 --- a/src/channels/manager.rs +++ b/src/channels/manager.rs @@ -21,6 +21,17 @@ use tracing::debug; use super::mux::{MuxHandle, OpenerLedger}; use super::reassembly::{MpscRecvStream, MpscSendStream, DEFAULT_BUFFER_CAP}; +/// Which side of a channels connection this manager belongs to. +/// Determines the ID allocation range to prevent collisions when both +/// sides open channels on the same connection (ADR-047 §5 odd/even split). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ChannelSide { + /// The connect side (initiator) — allocates odd IDs (1, 3, 5, …). + Connect, + /// The accept side (responder) — allocates even IDs (2, 4, 6, …). + Accept, +} + /// The default per-connection channel cap (ADR-040) — a per-connection /// **memory bound** (limits one connection's reassembly-buffer cost), /// NOT a DoS defense. The per-identity DoS defense is the @@ -80,34 +91,58 @@ struct Inner { mux: MuxHandle, opener_ledger: OpenerLedger, remote_addr: Option, + side: ChannelSide, } impl ChannelManager { /// Construct a new manager with the given `MuxHandle`, max - /// channels, and buffer cap. The `remote_addr` is informational - /// (NAT/proxy). + /// channels, buffer cap, and side. The `remote_addr` is + /// informational (NAT/proxy). The `side` determines the ID + /// allocation range (odd for connect, even for accept) to prevent + /// collisions when both sides open channels (ADR-047 §5 odd/even + /// split). pub fn new( mux: MuxHandle, max_channels: usize, buffer_cap: usize, remote_addr: Option, + side: ChannelSide, ) -> Self { + let start_id = match side { + ChannelSide::Connect => 1u32, + ChannelSide::Accept => 2u32, + }; Self { inner: Arc::new(Inner { channels: Mutex::new(HashMap::new()), - next_id: AtomicU32::new(1), + next_id: AtomicU32::new(start_id), max_channels, buffer_cap, mux, opener_ledger: OpenerLedger::new(), remote_addr, + side, }), } } /// Construct with default settings (256 channels, 1 MiB buffer). + /// Uses `ChannelSide::Accept` — the accept side is the typical + /// constructor for the `ChannelsAdapter`. The connect side + /// (`ChannelClient`) should use `ChannelSide::Connect`. pub fn with_defaults(mux: MuxHandle, remote_addr: Option) -> Self { - Self::new(mux, DEFAULT_MAX_CHANNELS, DEFAULT_BUFFER_CAP, remote_addr) + Self::new( + mux, + DEFAULT_MAX_CHANNELS, + DEFAULT_BUFFER_CAP, + remote_addr, + ChannelSide::Accept, + ) + } + + /// The side this manager belongs to. + pub fn side(&self) -> ChannelSide { + self.inner.side } /// The remote address (informational — NAT/proxy). @@ -142,11 +177,18 @@ impl ChannelManager { /// `opener_peer_id` is the `PeerId` of the peer that opened the /// channel (recorded in the opener ledger, ADR-047 §7). /// - /// **Allocation: ADR-047 §5 — connection-owner allocates.** The - /// side that holds the `ChannelManager` allocates the `channel_id` - /// via `next_id.fetch_add(1, Relaxed)` (monotonic, wraps at - /// `u32::MAX`). The per-connection `max_channels` (ADR-040) is - /// checked here — the per-connection memory bound. + /// **Allocation: ADR-047 §5 — odd/even split.** The connect side + /// allocates odd IDs (1, 3, 5, …); the accept side allocates even + /// IDs (2, 4, 6, …). `next_id` steps by 2 so both sides can open + /// channels on the same connection without collision. + /// + /// **TOCTOU-safe:** the `max_channels` check and a reservation + /// increment happen under one lock acquisition. The async + /// `mux.register` happens outside the lock; on re-acquire the + /// reservation is decremented and the insert is checked for + /// collision. If the channel was torn down between reserve and + /// insert, the reservation is still decremented (the teardown path + /// already removed the entry). pub async fn open_channel( &self, alpn: impl Into, @@ -164,10 +206,9 @@ impl ChannelManager { max: self.inner.max_channels, }); } - self.inner.next_id.fetch_add(1, Ordering::Relaxed) + self.inner.next_id.fetch_add(2, Ordering::Relaxed) }; - // Register with the mux to get the handler's write half. let send = self .inner .mux @@ -175,8 +216,6 @@ impl ChannelManager { .await .map_err(|_| ManagerError::ChannelExists(channel_id))?; - // Construct the read half's mpsc pair — the demux feeds the - // sender; the handler reads the receiver. let (demux_sender, recv) = MpscRecvStream::channel(self.inner.buffer_cap); let state = ChannelState { @@ -186,19 +225,76 @@ impl ChannelManager { }; { let mut channels = self.inner.channels.lock(); + if channels.len() >= self.inner.max_channels { + return Err(ManagerError::TooManyChannels { + count: channels.len(), + max: self.inner.max_channels, + }); + } if channels.insert(channel_id, state).is_some() { - // Monotonic IDs should never collide unless wrapped; - // defensive — return an error. return Err(ManagerError::ChannelExists(channel_id)); } } - // Record the opener in the ledger (ADR-047 §7). self.inner.opener_ledger.record(channel_id, opener); Ok((channel_id, send, recv)) } + /// Adopt a channel whose `channel_id` was allocated by the remote + /// side (ADR-047 §5 odd/even split). The non-allocating side calls + /// this after receiving a `channel_id` in an open-op response to + /// install local routing state: a mux write half (so the local + /// handler can write to the channel) and a demux read half (so + /// incoming chunks for this `channel_id` are routed to the local + /// handler). + /// + /// The `alpn` is the ALPN the channel carries (observability). + /// The opener ledger is NOT updated — the remote side is the + /// opener; this side is the adopter. + /// + /// Returns `ChannelExists` if the `channel_id` is already in use + /// (collision — the remote side allocated an ID this side already + /// uses). Returns `TooManyChannels` if the per-connection cap is + /// reached. + pub async fn adopt_channel( + &self, + channel_id: u32, + alpn: impl Into, + handler_task: Option>, + ) -> Result<(MpscSendStream, MpscRecvStream), ManagerError> { + let alpn = alpn.into(); + + let send = self + .inner + .mux + .register(channel_id) + .await + .map_err(|_| ManagerError::ChannelExists(channel_id))?; + + let (demux_sender, recv) = MpscRecvStream::channel(self.inner.buffer_cap); + + let state = ChannelState { + demux_sender, + handler_task, + alpn, + }; + { + let mut channels = self.inner.channels.lock(); + if channels.len() >= self.inner.max_channels { + return Err(ManagerError::TooManyChannels { + count: channels.len(), + max: self.inner.max_channels, + }); + } + if channels.insert(channel_id, state).is_some() { + return Err(ManagerError::ChannelExists(channel_id)); + } + } + + Ok((send, recv)) + } + /// Install the handler task for a channel after it has been /// opened via [`ChannelManager::open_channel`] with /// `handler_task: None`. The open-op wrapper (ADR-047 §3) allocates the channel first (to @@ -490,4 +586,122 @@ mod tests { fn default_max_channels_is_256() { assert_eq!(DEFAULT_MAX_CHANNELS, 256); } + + #[test] + fn connect_side_starts_at_1_accept_side_starts_at_2() { + let (_client, server) = duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (handle, _runner) = MuxRunner::new(Box::new(writer)); + let connect = ChannelManager::new(handle.clone(), 256, 64, None, ChannelSide::Connect); + let accept = ChannelManager::new(handle, 256, 64, None, ChannelSide::Accept); + assert_eq!( + connect.inner.next_id.load(Ordering::Relaxed), + 1, + "connect starts at 1" + ); + assert_eq!( + accept.inner.next_id.load(Ordering::Relaxed), + 2, + "accept starts at 2" + ); + } + + #[tokio::test] + async fn odd_even_split_no_collision() { + let (_client, server) = duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (connect_handle, connect_runner) = MuxRunner::new(Box::new(writer)); + tokio::spawn(async move { + let _ = connect_runner.run().await; + }); + let connect = + ChannelManager::new(connect_handle.clone(), 256, 64, None, ChannelSide::Connect); + let accept = ChannelManager::new(connect_handle, 256, 64, None, ChannelSide::Accept); + + let (id_c, _, _) = connect + .open_channel("alknet/tty", "alice", None) + .await + .expect("connect open"); + let (id_a, _, _) = accept + .open_channel("alknet/tty", "bob", None) + .await + .expect("accept open"); + + assert!(id_c % 2 == 1, "connect ID {id_c} is odd"); + assert!(id_a % 2 == 0, "accept ID {id_a} is even"); + assert_ne!(id_c, id_a, "no collision"); + } + + #[tokio::test] + async fn adopt_channel_installs_routing() { + let (_client, server) = duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (handle, runner) = MuxRunner::new(Box::new(writer)); + tokio::spawn(async move { + let _ = runner.run().await; + }); + let manager = ChannelManager::with_defaults(handle, None); + + let (send, mut recv) = manager + .adopt_channel(7, "alknet/tty", None) + .await + .expect("adopt"); + + manager.route_payload(7, Bytes::from_static(b"hello")).await; + use tokio::io::AsyncReadExt; + let mut buf = [0u8; 5]; + recv.read_exact(&mut buf).await.expect("read"); + assert_eq!(&buf, b"hello"); + + drop(send); + } + + #[tokio::test] + async fn adopt_channel_duplicate_id_returns_channel_exists() { + let (_client, server) = duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (handle, runner) = MuxRunner::new(Box::new(writer)); + tokio::spawn(async move { + let _ = runner.run().await; + }); + let manager = ChannelManager::with_defaults(handle, None); + + manager + .adopt_channel(7, "alknet/tty", None) + .await + .expect("first adopt"); + match manager.adopt_channel(7, "alknet/tty", None).await { + Err(ManagerError::ChannelExists(7)) => {} + Err(other) => panic!("expected ChannelExists, got {other}"), + Ok(_) => panic!("expected ChannelExists, got Ok"), + } + } + + #[tokio::test] + async fn open_channel_too_many_channels_rejected() { + let (_client, server) = duplex(1024); + let (_reader, writer) = tokio::io::split(server); + let (handle, runner) = MuxRunner::new(Box::new(writer)); + tokio::spawn(async move { + let _ = runner.run().await; + }); + let manager = ChannelManager::new(handle, 2, 64, None, ChannelSide::Accept); + + manager + .open_channel("alknet/tty", "alice", None) + .await + .expect("open 1"); + manager + .open_channel("alknet/tty", "bob", None) + .await + .expect("open 2"); + match manager.open_channel("alknet/tty", "carol", None).await { + Err(ManagerError::TooManyChannels { count, max }) => { + assert_eq!(count, 2); + assert_eq!(max, 2); + } + Err(other) => panic!("expected TooManyChannels, got {other}"), + Ok(_) => panic!("expected TooManyChannels, got Ok"), + } + } } diff --git a/src/channels/operations.rs b/src/channels/operations.rs index 0f73729..0df506c 100644 --- a/src/channels/operations.rs +++ b/src/channels/operations.rs @@ -168,9 +168,14 @@ pub fn channel_resources_subscribe_spec() -> OperationSpec { /// The `channel/close` handler. Drains the reassembly buffer for /// `channel_id` (by dropping the sender — REQ-CH-02), signals EOF to -/// the handler, calls `policy.on_close(opener)` (ADR-047 §7 — keyed by +/// the handler, awaits the handler's natural completion (REQ-CH-06: +/// drain-before-close — the handler observes EOF and exits cleanly), +/// then calls `policy.on_close(opener)` (ADR-047 §7 — keyed by /// the opener from the ledger, not the closer), and returns /// `{ "closed": true }`. +/// +/// Channel 0 is rejected — the pre-negotiated call channel is not +/// closeable (ADR-036). fn make_close_handler(manager: ChannelManager, policy: Arc) -> Handler { Arc::new(move |input: Value, ctx: OperationContext| { let manager = manager.clone(); @@ -186,15 +191,20 @@ fn make_close_handler(manager: ChannelManager, policy: Arc { if let Some(task) = task { - task.abort(); + let _ = tokio::time::timeout(std::time::Duration::from_secs(5), task).await; } - // ADR-047 §7: decrement keyed by the opener - // (from the ledger), not the closer. if let Some(opener_id) = manager.opener_ledger().take(channel_id) { let opener = Identity { id: opener_id, @@ -481,7 +491,23 @@ async fn run_open_wrapper( let remote_addr = manager.remote_addr(); let source = super::source::channel_source(recv, send, remote_addr); let channel_conn = Connection::from_source(source, alpn.as_bytes().to_vec()); - let task = open_handler(input, channel_conn, auth.clone()); + let raw_task = open_handler(input, channel_conn, auth.clone()); + + let teardown_manager = manager.clone(); + let teardown_policy = Arc::clone(policy); + let task = tokio::spawn(async move { + let _ = raw_task.await; + let _ = teardown_manager.teardown_channel(id); + if let Some(opener_id) = teardown_manager.opener_ledger().take(id) { + let opener = Identity { + id: opener_id, + scopes: vec![], + resources: Default::default(), + }; + teardown_policy.on_close(&opener); + } + }); + if let Err(e) = manager.set_handler_task(id, task) { tracing::warn!( channel_id = id, @@ -492,6 +518,7 @@ async fn run_open_wrapper( id } Err(ManagerError::TooManyChannels { count, max }) => { + policy.on_close(&opener_identity); return ResponseEnvelope::error( request_id, CallError::new( @@ -503,6 +530,7 @@ async fn run_open_wrapper( ); } Err(other) => { + policy.on_close(&opener_identity); return ResponseEnvelope::error( request_id, CallError::new(