Skip to main content

microsandbox_network/engine/netstack/
shared.rs

1//! Shared state between the NetWorker thread, smoltcp poll thread, and tokio
2//! proxy tasks.
3//!
4//! All inter-thread communication flows through [`SharedState`], which holds
5//! lock-free frame queues and cross-platform [`WakePipe`] notifications.
6
7use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
8use std::sync::{
9    Arc, Mutex, OnceLock,
10    atomic::{AtomicU64, Ordering},
11};
12use std::time::{Duration, Instant};
13
14use crossbeam_queue::ArrayQueue;
15use ipnetwork::Ipv6Network;
16use microsandbox_types::HttpConfig;
17use microsandbox_utils::ttl_reverse_index::TtlReverseIndex;
18pub use microsandbox_utils::wake_pipe::WakePipe;
19use parking_lot::RwLock;
20
21use crate::addr::normalize_ip_addr;
22use crate::engine::http_deny::{self, DEFAULT_HTTP_DENY_MESSAGE};
23
24//--------------------------------------------------------------------------------------------------
25// Constants
26//--------------------------------------------------------------------------------------------------
27
28/// Default frame queue capacity. Matches libkrun's virtio queue size.
29pub const DEFAULT_QUEUE_CAPACITY: usize = 1024;
30
31//--------------------------------------------------------------------------------------------------
32// Types
33//--------------------------------------------------------------------------------------------------
34
35/// All shared state between the three threads:
36///
37/// - **NetWorker** (libkrun) — pushes guest frames to `tx_ring`, pops
38///   response frames from `rx_ring`.
39/// - **smoltcp poll thread** — pops from `tx_ring`, processes through smoltcp,
40///   pushes responses to `rx_ring`.
41/// - **tokio proxy tasks** — relay data between smoltcp sockets and real
42///   network connections.
43///
44/// Queue naming follows the **guest's perspective** (matching libkrun's
45/// convention): `tx_ring` = "transmit from guest", `rx_ring` = "receive at
46/// guest".
47pub struct SharedState {
48    /// Frames from guest → smoltcp (NetWorker writes, smoltcp reads).
49    pub tx_ring: ArrayQueue<Vec<u8>>,
50
51    /// Frames from smoltcp → guest (smoltcp writes, NetWorker reads).
52    pub rx_ring: ArrayQueue<Vec<u8>>,
53
54    /// Wakes NetWorker: "rx_ring has frames for the guest."
55    /// Written by `SmoltcpDevice::transmit()`. Read end polled by NetWorker's
56    /// epoll loop.
57    pub rx_wake: WakePipe,
58
59    /// Wakes smoltcp poll thread: "tx_ring has frames from the guest."
60    /// Written by `SmoltcpBackend::write_frame()`. Read end polled by the
61    /// poll loop.
62    pub tx_wake: WakePipe,
63
64    /// Wakes smoltcp poll thread: "proxy task has data to write to a smoltcp
65    /// socket." Written by proxy tasks via channels. Read end polled by the
66    /// poll loop.
67    pub proxy_wake: WakePipe,
68
69    /// Optional host-side termination hook used for fatal policy violations.
70    termination_hook: Mutex<Option<Arc<dyn Fn() + Send + Sync>>>,
71
72    /// Resolved hostname index used to map destination IPs back to queried hostnames.
73    resolved_hostnames: RwLock<TtlReverseIndex<ResolvedHostnameKey, IpAddr>>,
74
75    /// Per-sandbox gateway IPv4. Set once at boot; used by
76    /// `DestinationGroup::Host` rule matching and `host.microsandbox.internal`
77    /// DNS synthesis. `None` in isolated unit tests.
78    gateway_ipv4: OnceLock<Ipv4Addr>,
79
80    /// Per-sandbox gateway IPv6. Set once at boot. See `gateway_ipv4`.
81    gateway_ipv6: OnceLock<Ipv6Addr>,
82
83    /// NAT64 `/96` prefixes used by policy classification.
84    nat64_prefixes: OnceLock<Vec<Ipv6Network>>,
85
86    /// Aggregate network byte counters at the guest/runtime boundary.
87    metrics: NetworkMetrics,
88
89    /// HTTP denial settings installed before the network starts.
90    http: OnceLock<HttpConfig>,
91}
92
93/// Aggregate network byte counters shared with the runtime metrics sampler.
94pub struct NetworkMetrics {
95    tx_bytes: AtomicU64,
96    rx_bytes: AtomicU64,
97}
98
99/// Address family for resolved hostname entries.
100#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
101pub enum ResolvedHostnameFamily {
102    Ipv4,
103    Ipv6,
104}
105
106/// Composite cache key for a single DNS resolution.
107///
108/// `family` partitions entries so that `A` and `AAAA` responses for the
109/// same hostname refresh independently instead of overwriting each other.
110#[derive(Clone, Debug, PartialEq, Eq, Hash)]
111struct ResolvedHostnameKey {
112    hostname: String,
113    family: ResolvedHostnameFamily,
114}
115
116//--------------------------------------------------------------------------------------------------
117// Methods
118//--------------------------------------------------------------------------------------------------
119
120impl SharedState {
121    /// Create shared state with the given queue capacity.
122    pub fn new(queue_capacity: usize) -> Self {
123        Self {
124            tx_ring: ArrayQueue::new(queue_capacity),
125            rx_ring: ArrayQueue::new(queue_capacity),
126            rx_wake: WakePipe::new(),
127            tx_wake: WakePipe::new(),
128            proxy_wake: WakePipe::new(),
129            termination_hook: Mutex::new(None),
130            resolved_hostnames: RwLock::new(TtlReverseIndex::default()),
131            gateway_ipv4: OnceLock::new(),
132            gateway_ipv6: OnceLock::new(),
133            nat64_prefixes: OnceLock::new(),
134            metrics: NetworkMetrics::default(),
135            http: OnceLock::new(),
136        }
137    }
138
139    /// Set the per-sandbox gateway IPs. Called once at boot. Each family is
140    /// only published when active for this sandbox.
141    pub fn set_gateway_ips(&self, ipv4: Option<Ipv4Addr>, ipv6: Option<Ipv6Addr>) {
142        if let Some(ipv4) = ipv4 {
143            let _ = self.gateway_ipv4.set(ipv4);
144        }
145        if let Some(ipv6) = ipv6 {
146            let _ = self.gateway_ipv6.set(ipv6);
147        }
148    }
149
150    /// Gateway IPv4 address, if set.
151    pub fn gateway_ipv4(&self) -> Option<Ipv4Addr> {
152        self.gateway_ipv4.get().copied()
153    }
154
155    /// Gateway IPv6 address, if set.
156    pub fn gateway_ipv6(&self) -> Option<Ipv6Addr> {
157        self.gateway_ipv6.get().copied()
158    }
159
160    /// Install HTTP denial settings. The first call wins.
161    pub fn set_http_config(&self, config: HttpConfig) {
162        let _ = self.http.set(config);
163    }
164
165    /// Whether readable denial responses were explicitly enabled.
166    pub fn http_deny_response_enabled(&self) -> bool {
167        self.http.get().is_some_and(|http| http.deny_response)
168    }
169
170    /// Render the HTTP/HTTPS deny body for `host`.
171    pub fn http_deny_body(&self, host: &str) -> String {
172        let template = self
173            .http
174            .get()
175            .and_then(|http| http.deny_message.as_deref())
176            .unwrap_or(DEFAULT_HTTP_DENY_MESSAGE);
177        http_deny::render_http_deny_message(template, host)
178    }
179
180    /// Set NAT64 prefixes. Called once before policy evaluation starts.
181    pub fn set_nat64_prefixes(&self, prefixes: Vec<Ipv6Network>) {
182        let _ = self.nat64_prefixes.set(prefixes);
183    }
184
185    /// NAT64 prefixes used by policy classification.
186    pub fn nat64_prefixes(&self) -> &[Ipv6Network] {
187        self.nat64_prefixes.get().map(Vec::as_slice).unwrap_or(&[])
188    }
189
190    /// Install a host-side termination hook.
191    pub fn set_termination_hook(&self, hook: Arc<dyn Fn() + Send + Sync>) {
192        *self.termination_hook.lock().unwrap() = Some(hook);
193    }
194
195    /// Trigger host-side termination if a hook is installed.
196    pub fn trigger_termination(&self) {
197        let hook = self.termination_hook.lock().unwrap().clone();
198        if let Some(hook) = hook {
199            hook();
200        }
201    }
202
203    /// Replace the resolved addresses for a hostname within the given address family.
204    pub fn cache_resolved_hostname(
205        &self,
206        domain: &str,
207        family: ResolvedHostnameFamily,
208        addrs: impl IntoIterator<Item = IpAddr>,
209        ttl: Duration,
210    ) {
211        let hostname = normalize_hostname(domain);
212        let key = ResolvedHostnameKey { hostname, family };
213        let addrs = addrs.into_iter().map(normalize_ip_addr);
214        self.resolved_hostnames
215            .write()
216            .insert(key, addrs, ttl, Instant::now());
217    }
218
219    /// Clear the resolved addresses for a hostname within the given address family.
220    pub fn clear_resolved_hostname(&self, domain: &str, family: ResolvedHostnameFamily) {
221        let hostname = normalize_hostname(domain);
222        let key = ResolvedHostnameKey { hostname, family };
223        self.resolved_hostnames.write().remove(&key, Instant::now());
224    }
225
226    /// Returns `true` when any resolved hostname for `addr` satisfies `predicate`.
227    pub fn any_resolved_hostname(
228        &self,
229        addr: IpAddr,
230        mut predicate: impl FnMut(&str) -> bool,
231    ) -> bool {
232        let addr = normalize_ip_addr(addr);
233
234        self.resolved_hostnames
235            .read()
236            .member_matches(&addr, Instant::now(), |key| predicate(&key.hostname))
237    }
238
239    /// Best-effort expiry maintenance for resolved hostnames.
240    ///
241    /// This runs outside the hot egress read path. If the index is currently
242    /// busy, cleanup is skipped and retried on the next maintenance pass.
243    pub fn cleanup_resolved_hostnames(&self) {
244        if let Some(mut idx) = self.resolved_hostnames.try_write() {
245            idx.evict_expired(Instant::now());
246        }
247    }
248
249    /// Increment the guest -> runtime byte counter.
250    pub fn add_tx_bytes(&self, bytes: usize) {
251        self.metrics
252            .tx_bytes
253            .fetch_add(bytes as u64, Ordering::Relaxed);
254    }
255
256    /// Increment the runtime -> guest byte counter.
257    pub fn add_rx_bytes(&self, bytes: usize) {
258        self.metrics
259            .rx_bytes
260            .fetch_add(bytes as u64, Ordering::Relaxed);
261    }
262
263    /// Push a runtime -> guest ethernet frame and update RX metrics on success.
264    pub(crate) fn push_rx_frame(&self, frame: Vec<u8>) -> bool {
265        let frame_len = frame.len();
266        if self.rx_ring.push(frame).is_err() {
267            return false;
268        }
269
270        self.add_rx_bytes(frame_len);
271        true
272    }
273
274    /// Push a runtime -> guest ethernet frame, update RX metrics, and wake libkrun.
275    pub(crate) fn push_rx_frame_and_wake(&self, frame: Vec<u8>) -> bool {
276        if !self.push_rx_frame(frame) {
277            return false;
278        }
279
280        self.rx_wake.wake();
281        true
282    }
283
284    /// Total bytes transmitted by the guest into the runtime.
285    pub fn tx_bytes(&self) -> u64 {
286        self.metrics.tx_bytes.load(Ordering::Relaxed)
287    }
288
289    /// Total bytes delivered by the runtime to the guest.
290    pub fn rx_bytes(&self) -> u64 {
291        self.metrics.rx_bytes.load(Ordering::Relaxed)
292    }
293}
294
295impl Default for NetworkMetrics {
296    fn default() -> Self {
297        Self {
298            tx_bytes: AtomicU64::new(0),
299            rx_bytes: AtomicU64::new(0),
300        }
301    }
302}
303
304pub(crate) fn normalize_hostname(domain: &str) -> String {
305    domain.trim_end_matches('.').to_ascii_lowercase()
306}
307
308//--------------------------------------------------------------------------------------------------
309// Tests
310//--------------------------------------------------------------------------------------------------
311
312#[cfg(test)]
313mod tests {
314    use super::*;
315
316    #[test]
317    fn shared_state_queue_push_pop() {
318        let state = SharedState::new(4);
319
320        // Push frames to tx_ring.
321        state.tx_ring.push(vec![1, 2, 3]).unwrap();
322        state.tx_ring.push(vec![4, 5, 6]).unwrap();
323
324        // Pop in FIFO order.
325        assert_eq!(state.tx_ring.pop(), Some(vec![1, 2, 3]));
326        assert_eq!(state.tx_ring.pop(), Some(vec![4, 5, 6]));
327        assert_eq!(state.tx_ring.pop(), None);
328    }
329
330    #[test]
331    fn shared_state_queue_full() {
332        let state = SharedState::new(2);
333
334        state.rx_ring.push(vec![1]).unwrap();
335        state.rx_ring.push(vec![2]).unwrap();
336        // Queue is full — push returns the frame back.
337        assert!(state.rx_ring.push(vec![3]).is_err());
338    }
339
340    #[test]
341    fn push_rx_frame_counts_only_successful_pushes() {
342        let state = SharedState::new(1);
343
344        assert!(state.push_rx_frame(vec![1, 2, 3]));
345        assert_eq!(state.rx_bytes(), 3);
346
347        assert!(!state.push_rx_frame(vec![4, 5]));
348        assert_eq!(state.rx_bytes(), 3);
349    }
350
351    #[test]
352    fn resolved_hostnames_are_isolated_per_family() {
353        let state = SharedState::new(4);
354        let v4: IpAddr = "1.1.1.1".parse().unwrap();
355        let v6: IpAddr = "2606:4700:4700::1111".parse().unwrap();
356
357        state.cache_resolved_hostname(
358            "Example.com.",
359            ResolvedHostnameFamily::Ipv4,
360            [v4],
361            Duration::from_secs(30),
362        );
363        state.cache_resolved_hostname(
364            "example.com",
365            ResolvedHostnameFamily::Ipv6,
366            [v6],
367            Duration::from_secs(30),
368        );
369
370        assert!(state.any_resolved_hostname(v4, |h| h == "example.com"));
371        assert!(state.any_resolved_hostname(v6, |h| h == "example.com"));
372        assert!(!state.any_resolved_hostname(v4, |h| h == "other.example"));
373    }
374
375    #[test]
376    fn resolved_hostnames_normalize_ipv4_mapped_ipv6() {
377        let state = SharedState::new(4);
378        let mapped: IpAddr = "::ffff:169.254.169.254".parse().unwrap();
379        let embedded: IpAddr = "169.254.169.254".parse().unwrap();
380
381        state.cache_resolved_hostname(
382            "metadata.example",
383            ResolvedHostnameFamily::Ipv6,
384            [mapped],
385            Duration::from_secs(30),
386        );
387
388        assert!(state.any_resolved_hostname(embedded, |h| h == "metadata.example"));
389        assert!(state.any_resolved_hostname(mapped, |h| h == "metadata.example"));
390    }
391}