Route Windows sandbox proxy traffic by restricting SID (#34613)

## Why

Elevated Windows sandboxes need stable managed-proxy ports while preserving the network policy and environment attribution of each sandboxed process.

## What changed

- Keep shared HTTP and SOCKS5 loopback ingress listeners alive across managed-proxy instances.
- Add a per-route restricting SID to elevated sandbox tokens and dispatch incoming connections to the matching proxy policy after attributing the client process.
- Reject connections without exactly one registered route, remove routes when their proxy handle is dropped, and keep unsandboxed Windows launches off the managed ingress.
- Provision the elevated sandbox with the configured proxy ports and local-binding setting, honoring the selected profile and CLI overrides.

## Testing

- Add Windows unit tests for TCP ownership attribution, route selection, restricting-token propagation, and setup settings.
- Add an end-to-end Windows test covering stable ports, isolated environment policies, HTTP and SOCKS5 routing, and route teardown.

GitOrigin-RevId: 783fac6e0f904dc9bb1955b75d4a5895e8bb9690
This commit is contained in:
iceweasel-oai
2026-07-21 21:01:25 +00:00
committed by copyberry
parent dfd2d8133c
commit 999a715089
40 changed files with 2792 additions and 191 deletions

View File

@@ -51,3 +51,14 @@ security-framework = "3"
[target.'cfg(windows)'.dependencies]
schannel = "0.1"
windows-sys = { version = "0.52", features = [
"Win32_Foundation",
"Win32_NetworkManagement_IpHelper",
"Win32_Networking_WinSock",
"Win32_Security",
"Win32_Security_Authorization",
"Win32_System_Threading",
] }
[target.'cfg(windows)'.dev-dependencies]
codex-windows-sandbox = { path = "../windows-sandbox-rs" }

View File

@@ -430,6 +430,26 @@ pub fn resolve_runtime(cfg: &NetworkProxyConfig) -> Result<RuntimeConfig> {
})
}
/// Returns the sorted loopback ports used by the configured managed proxy listeners.
pub fn managed_proxy_ports(cfg: &NetworkProxyConfig) -> Result<Vec<u16>> {
let runtime = resolve_runtime(cfg)?;
if runtime.http_addr.port() == 0 {
bail!("network.proxy_url must use a fixed non-zero port for managed proxy provisioning");
}
let mut ports = vec![runtime.http_addr.port()];
if cfg.enable_socks5 {
if runtime.socks_addr.port() == 0 {
bail!(
"network.socks_url must use a fixed non-zero port for managed proxy provisioning"
);
}
ports.push(runtime.socks_addr.port());
}
ports.sort_unstable();
ports.dedup();
Ok(ports)
}
fn resolve_addr(url: &str, default_port: u16) -> Result<SocketAddr> {
let addr_parts = parse_host_port(url, default_port)?;
let host = if addr_parts.host.eq_ignore_ascii_case("localhost") {
@@ -605,6 +625,32 @@ mod tests {
);
}
#[test]
fn managed_proxy_ports_reject_ephemeral_ports() {
let mut config = NetworkProxyConfig {
proxy_url: "http://127.0.0.1:0".to_string(),
..Default::default()
};
assert_eq!(
managed_proxy_ports(&config).unwrap_err().to_string(),
"network.proxy_url must use a fixed non-zero port for managed proxy provisioning"
);
config.proxy_url = "http://127.0.0.1:3128".to_string();
config.socks_url = "socks5h://127.0.0.1:48081".to_string();
assert_eq!(managed_proxy_ports(&config).unwrap(), vec![3128, 48081]);
config.socks_url = "socks5h://127.0.0.1:0".to_string();
assert_eq!(
managed_proxy_ports(&config).unwrap_err().to_string(),
"network.socks_url must use a fixed non-zero port for managed proxy provisioning"
);
config.enable_socks5 = false;
assert_eq!(managed_proxy_ports(&config).unwrap(), vec![3128]);
}
#[test]
fn network_proxy_config_uses_struct_defaults_for_missing_fields() {
let config: NetworkProxyConfig = serde_json::from_str(r#"{ "enabled": true }"#).unwrap();

View File

@@ -40,6 +40,7 @@ use rama_core::error::ErrorExt as _;
use rama_core::error::OpaqueError;
use rama_core::extensions::ExtensionsMut;
use rama_core::extensions::ExtensionsRef;
use rama_core::service::BoxService;
use rama_core::service::service_fn;
use rama_core::stream::Stream;
use rama_http::Body;
@@ -66,6 +67,7 @@ use rama_net::proxy::ProxyRequest;
use rama_net::proxy::ProxyTarget;
use rama_net::proxy::StreamForwardService;
use rama_net::stream::SocketInfo;
use rama_tcp::TcpStream;
use rama_tcp::client::Request as TcpRequest;
use rama_tcp::server::TcpListener;
use rama_tls_rustls::client::TlsConnectorDataBuilder;
@@ -124,12 +126,25 @@ async fn run_http_proxy_with_listener(
policy_decider: Option<Arc<dyn NetworkPolicyDecider>>,
environment_id: Option<String>,
) -> Result<()> {
ensure_rustls_crypto_provider();
let addr = listener
.local_addr()
.context("read HTTP proxy listener local addr")?;
info!("HTTP proxy listening on {addr}");
listener
.serve(http_proxy_service(state, policy_decider, environment_id))
.await;
Ok(())
}
pub(crate) fn http_proxy_service(
state: Arc<NetworkProxyState>,
policy_decider: Option<Arc<dyn NetworkPolicyDecider>>,
environment_id: Option<String>,
) -> BoxService<TcpStream, (), rama_core::error::BoxError> {
ensure_rustls_crypto_provider();
// This proxy listener only needs HTTP/1 proxy semantics. Using Rama's auto builder
// forces every accepted socket through the HTTP version sniffing pre-read path before proxy
// request parsing, which can stall some local clients on macOS before CONNECT/absolute-form
@@ -156,16 +171,7 @@ async fn run_http_proxy_with_listener(
})),
);
info!("HTTP proxy listening on {addr}");
listener
.serve(BindConnectionAttribution::new(
http_service,
state,
environment_id,
))
.await;
Ok(())
BindConnectionAttribution::new(http_service, state, environment_id).boxed()
}
async fn http_connect_accept(

View File

@@ -19,6 +19,10 @@ mod runtime;
mod socks5;
mod state;
mod upstream;
#[cfg(target_os = "windows")]
mod windows_proxy_ingress;
#[cfg(target_os = "windows")]
mod windows_tcp_attribution;
pub use attribution::PROXY_ATTRIBUTION_TOKEN_ENV_KEY;
pub use attribution::write_attribution_frame;
@@ -32,6 +36,7 @@ pub use config::NetworkProxyConfig;
pub use config::NetworkUnixSocketPermission;
pub use config::NetworkUnixSocketPermissions;
pub use config::host_and_port_from_network_addr;
pub use config::managed_proxy_ports;
pub use credential_broker::CREDENTIAL_BROKER_ACTIVE_ENV_KEY;
pub use credential_broker::brokered_credential_dummy_env_keys;
pub use credential_broker::brokered_credential_env_keys;
@@ -67,7 +72,9 @@ pub use proxy::PROXY_GIT_SSH_COMMAND_ENV_KEY;
pub use proxy::PROXY_URL_ENV_KEYS;
pub use proxy::PreparedManagedNetwork;
pub use proxy::has_proxy_url_env_vars;
pub use proxy::is_managed_proxy_env_var;
pub use proxy::proxy_url_env_value;
pub use proxy::strip_managed_proxy_env;
pub use remote_config::RemoteNetworkProxyConfig;
pub use remote_config::RemoteNetworkProxyLaunchConfig;
pub use runtime::BlockedRequest;

File diff suppressed because it is too large Load Diff

View File

@@ -29,6 +29,7 @@ use rama_core::error::BoxError;
use rama_core::extensions::Extensions;
use rama_core::extensions::ExtensionsMut;
use rama_core::extensions::ExtensionsRef;
use rama_core::service::BoxService;
use rama_core::service::service_fn;
use rama_net::address::HostWithPort;
use rama_net::client::EstablishedClientConnection;
@@ -129,6 +130,23 @@ async fn run_socks5_with_listener(
}
}
listener
.serve(socks5_proxy_service(
state,
policy_decider,
environment_id,
enable_socks5_udp,
))
.await;
Ok(())
}
pub(crate) fn socks5_proxy_service(
state: Arc<NetworkProxyState>,
policy_decider: Option<Arc<dyn NetworkPolicyDecider>>,
environment_id: Option<String>,
enable_socks5_udp: bool,
) -> BoxService<TcpStream, (), BoxError> {
let tcp_connector = TargetCheckedTcpConnector::new(state.clone());
let policy_tcp_connector = service_fn({
let policy_decider = policy_decider.clone();
@@ -163,19 +181,10 @@ async fn run_socks5_with_listener(
}
}));
let socks_acceptor = base.with_udp_associator(udp_relay);
listener
.serve(BindConnectionAttribution::new(
socks_acceptor,
state,
environment_id,
))
.await;
BindConnectionAttribution::new(socks_acceptor, state, environment_id).boxed()
} else {
listener
.serve(BindConnectionAttribution::new(base, state, environment_id))
.await;
BindConnectionAttribution::new(base, state, environment_id).boxed()
}
Ok(())
}
async fn handle_socks5_tcp(

View File

@@ -0,0 +1,368 @@
use crate::proxy::reserve_windows_managed_listeners;
use crate::proxy::reserve_windows_managed_socks_listener;
use crate::proxy::windows_managed_loopback_addr;
use crate::windows_tcp_attribution::restricting_sids_for_tcp_connection;
use anyhow::Context;
use anyhow::Result;
use rama_core::Service;
use rama_core::error::BoxError;
use rama_core::service::BoxService;
use rama_net::stream::Socket;
use rama_tcp::TcpStream;
use rama_tcp::server::TcpListener;
use std::collections::HashMap;
use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::LazyLock;
use std::sync::Mutex;
#[cfg(test)]
use std::sync::Weak;
use tokio::runtime::Handle;
use tokio::task::JoinHandle;
use tracing::info;
pub(crate) type WindowsRouteService = BoxService<TcpStream, (), BoxError>;
// Production keeps the listeners alive for the process lifetime so their ports remain stable even
// when no routes are registered. Crate tests use independent Tokio runtimes and requested ports, so
// they retain only a weak reference and can tear each ingress down between tests.
#[cfg(not(test))]
static SHARED_INGRESS: LazyLock<Mutex<Option<Arc<WindowsProxyIngress>>>> =
LazyLock::new(|| Mutex::new(None));
#[cfg(test)]
static SHARED_INGRESS: LazyLock<Mutex<Weak<WindowsProxyIngress>>> =
LazyLock::new(|| Mutex::new(Weak::new()));
#[derive(Clone)]
struct RouteServices {
http: WindowsRouteService,
socks: Option<WindowsRouteService>,
}
type RouteRegistry = Arc<Mutex<HashMap<String, Arc<RouteServices>>>>;
#[derive(Clone, Copy)]
enum ProxyProtocol {
Http,
Socks,
}
#[derive(Clone)]
struct IngressDispatcher {
routes: RouteRegistry,
protocol: ProxyProtocol,
}
impl Service<TcpStream> for IngressDispatcher {
type Output = ();
type Error = BoxError;
async fn serve(&self, stream: TcpStream) -> Result<(), BoxError> {
let local_addr = stream.local_addr()?;
let peer_addr = stream.peer_addr()?;
let restricting_sids = restricting_sids_for_tcp_connection(local_addr, peer_addr)?;
let route = {
let routes = self
.routes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
registered_route_for_sids(&routes, &restricting_sids)?
};
let service = match self.protocol {
ProxyProtocol::Http => route.http.clone(),
ProxyProtocol::Socks => route.socks.clone().ok_or_else(|| {
io::Error::new(
io::ErrorKind::PermissionDenied,
"network proxy route does not enable SOCKS5",
)
})?,
};
service.serve(stream).await
}
}
pub(crate) struct WindowsProxyIngress {
http_addr: SocketAddr,
routes: RouteRegistry,
runtime: Handle,
http_task: JoinHandle<()>,
socks: Mutex<SocksListenerState>,
}
struct SocksListenerState {
addr: SocketAddr,
task: Option<JoinHandle<()>>,
}
impl WindowsProxyIngress {
pub(crate) fn shared(
requested_http_addr: SocketAddr,
requested_socks_addr: SocketAddr,
reserve_socks_listener: bool,
) -> Result<Arc<Self>> {
let requested_http_addr = windows_managed_loopback_addr(requested_http_addr);
let requested_socks_addr = windows_managed_loopback_addr(requested_socks_addr);
let mut shared = SHARED_INGRESS
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
#[cfg(not(test))]
if let Some(ingress) = shared.as_ref()
&& ingress.is_running()
{
if reserve_socks_listener {
ingress.ensure_socks_listener(requested_socks_addr)?;
}
return Ok(Arc::clone(ingress));
}
#[cfg(not(test))]
shared.take();
#[cfg(test)]
if let Some(ingress) = shared.upgrade()
&& ingress.is_running()
{
if reserve_socks_listener {
ingress.ensure_socks_listener(requested_socks_addr)?;
}
return Ok(ingress);
}
let listeners = reserve_windows_managed_listeners(
requested_http_addr,
requested_socks_addr,
reserve_socks_listener,
)
.context("reserve shared managed Windows proxy ingress")?;
let http_addr = listeners.http_addr()?;
let socks_addr = listeners.socks_addr(requested_socks_addr)?;
let (http_listener, socks_listener) = listeners.into_listeners();
let http_listener =
TcpListener::try_from(http_listener).context("convert shared HTTP ingress listener")?;
let socks_listener = socks_listener
.map(TcpListener::try_from)
.transpose()
.context("convert shared SOCKS5 ingress listener")?;
let runtime =
Handle::try_current().context("start shared managed Windows proxy ingress")?;
let routes = Arc::new(Mutex::new(HashMap::new()));
let http_task = runtime.spawn(run_listener(
http_listener,
IngressDispatcher {
routes: Arc::clone(&routes),
protocol: ProxyProtocol::Http,
},
"HTTP",
http_addr,
));
let socks_task = socks_listener
.map(|listener| spawn_socks_listener(&runtime, &routes, listener, socks_addr));
let ingress = Arc::new(Self {
http_addr,
routes,
runtime,
http_task,
socks: Mutex::new(SocksListenerState {
addr: socks_addr,
task: socks_task,
}),
});
#[cfg(not(test))]
{
*shared = Some(Arc::clone(&ingress));
}
#[cfg(test)]
{
*shared = Arc::downgrade(&ingress);
}
Ok(ingress)
}
pub(crate) fn http_addr(&self) -> SocketAddr {
self.http_addr
}
pub(crate) fn socks_addr(&self) -> SocketAddr {
self.socks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.addr
}
pub(crate) fn active_socks_addr(&self) -> Option<SocketAddr> {
let socks = self
.socks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
socks
.task
.as_ref()
.filter(|task| !task.is_finished())
.map(|_| socks.addr)
}
pub(crate) fn register_route(
self: &Arc<Self>,
http: WindowsRouteService,
socks: Option<WindowsRouteService>,
) -> WindowsProxyRoute {
let services = Arc::new(RouteServices { http, socks });
let mut routes = self
.routes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let sid = loop {
let sid = random_restricting_sid();
if !routes.contains_key(&sid) {
break sid;
}
};
routes.insert(sid.clone(), Arc::clone(&services));
WindowsProxyRoute {
sid,
services,
ingress: Arc::clone(self),
}
}
fn is_running(&self) -> bool {
!self.http_task.is_finished()
&& self
.socks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.task
.as_ref()
.is_none_or(|task| !task.is_finished())
}
fn ensure_socks_listener(&self, requested_addr: SocketAddr) -> Result<()> {
let mut socks = self
.socks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(task) = socks.task.as_ref() {
anyhow::ensure!(
!task.is_finished(),
"shared managed Windows SOCKS5 ingress stopped"
);
return Ok(());
}
let listener = reserve_windows_managed_socks_listener(requested_addr)
.context("reserve shared managed Windows SOCKS5 ingress")?;
let addr = listener
.local_addr()
.context("read shared managed Windows SOCKS5 ingress address")?;
let listener = {
let _runtime = self.runtime.enter();
TcpListener::try_from(listener)
}
.context("convert shared SOCKS5 ingress listener")?;
let task = spawn_socks_listener(&self.runtime, &self.routes, listener, addr);
socks.addr = addr;
socks.task = Some(task);
Ok(())
}
}
impl Drop for WindowsProxyIngress {
fn drop(&mut self) {
self.http_task.abort();
let socks = self
.socks
.get_mut()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(socks_task) = socks.task.as_ref() {
socks_task.abort();
}
}
}
fn spawn_socks_listener(
runtime: &Handle,
routes: &RouteRegistry,
listener: TcpListener,
addr: SocketAddr,
) -> JoinHandle<()> {
runtime.spawn(run_listener(
listener,
IngressDispatcher {
routes: Arc::clone(routes),
protocol: ProxyProtocol::Socks,
},
"SOCKS5",
addr,
))
}
pub(crate) struct WindowsProxyRoute {
sid: String,
services: Arc<RouteServices>,
ingress: Arc<WindowsProxyIngress>,
}
impl WindowsProxyRoute {
pub(crate) fn sid(&self) -> &str {
&self.sid
}
}
impl Drop for WindowsProxyRoute {
fn drop(&mut self) {
let mut routes = self
.ingress
.routes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if routes
.get(&self.sid)
.is_some_and(|services| Arc::ptr_eq(services, &self.services))
{
routes.remove(&self.sid);
}
}
}
async fn run_listener(
listener: TcpListener,
dispatcher: IngressDispatcher,
protocol: &'static str,
addr: SocketAddr,
) {
info!("shared managed Windows {protocol} proxy ingress listening on {addr}");
listener.serve(dispatcher).await;
}
fn registered_route_for_sids(
routes: &HashMap<String, Arc<RouteServices>>,
restricting_sids: &[String],
) -> io::Result<Arc<RouteServices>> {
let mut matching_routes = restricting_sids
.iter()
.filter_map(|sid| routes.get(sid).cloned());
let route = matching_routes.next().ok_or_else(|| {
io::Error::new(
io::ErrorKind::PermissionDenied,
"proxy client token has no registered network proxy route SID",
)
})?;
if matching_routes.next().is_some() {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"proxy client token has multiple registered network proxy route SIDs",
));
}
Ok(route)
}
fn random_restricting_sid() -> String {
let a = rand::random::<u32>();
let b = rand::random::<u32>();
let c = rand::random::<u32>();
let d = rand::random::<u32>();
format!("S-1-5-21-{a}-{b}-{c}-{d}")
}
#[cfg(test)]
#[path = "windows_proxy_ingress_tests.rs"]
mod tests;

View File

@@ -0,0 +1,43 @@
use super::*;
use rama_core::service::service_fn;
#[test]
fn selects_exactly_one_registered_route() {
let route = route_services();
let routes = HashMap::from([("registered".to_string(), Arc::clone(&route))]);
let selected = registered_route_for_sids(
&routes,
&["unrelated".to_string(), "registered".to_string()],
)
.expect("one registered route should be selected");
assert!(Arc::ptr_eq(&selected, &route));
}
#[test]
fn rejects_missing_or_ambiguous_registered_routes() {
let first = route_services();
let second = route_services();
let routes = HashMap::from([("first".to_string(), first), ("second".to_string(), second)]);
let Err(missing) = registered_route_for_sids(&routes, &["missing".to_string()]) else {
panic!("an unknown SID should fail closed");
};
let Err(ambiguous) =
registered_route_for_sids(&routes, &["first".to_string(), "second".to_string()])
else {
panic!("multiple registered SIDs should fail closed");
};
assert_eq!(missing.kind(), io::ErrorKind::PermissionDenied);
assert_eq!(ambiguous.kind(), io::ErrorKind::PermissionDenied);
}
fn route_services() -> Arc<RouteServices> {
let service = service_fn(|_stream: TcpStream| async { Ok::<(), BoxError>(()) }).boxed();
Arc::new(RouteServices {
http: service,
socks: None,
})
}

View File

@@ -0,0 +1,315 @@
use std::ffi::c_void;
use std::io;
use std::mem::offset_of;
use std::mem::size_of;
use std::net::Ipv4Addr;
use std::net::SocketAddr;
use std::net::SocketAddrV4;
use std::os::windows::io::AsRawHandle;
use std::os::windows::io::FromRawHandle;
use std::os::windows::io::OwnedHandle;
use std::os::windows::io::RawHandle;
use windows_sys::Win32::Foundation::ERROR_INSUFFICIENT_BUFFER;
use windows_sys::Win32::Foundation::GetLastError;
use windows_sys::Win32::Foundation::HANDLE;
use windows_sys::Win32::Foundation::HLOCAL;
use windows_sys::Win32::Foundation::LocalFree;
use windows_sys::Win32::Foundation::NO_ERROR;
use windows_sys::Win32::Foundation::PSID;
use windows_sys::Win32::NetworkManagement::IpHelper::GetExtendedTcpTable;
use windows_sys::Win32::NetworkManagement::IpHelper::MIB_TCPROW_OWNER_PID;
use windows_sys::Win32::NetworkManagement::IpHelper::MIB_TCPTABLE_OWNER_PID;
use windows_sys::Win32::NetworkManagement::IpHelper::TCP_TABLE_OWNER_PID_CONNECTIONS;
use windows_sys::Win32::Networking::WinSock::AF_INET;
use windows_sys::Win32::Security::Authorization::ConvertSidToStringSidW;
use windows_sys::Win32::Security::GetTokenInformation;
use windows_sys::Win32::Security::SID_AND_ATTRIBUTES;
use windows_sys::Win32::Security::TOKEN_GROUPS;
use windows_sys::Win32::Security::TOKEN_QUERY;
use windows_sys::Win32::Security::TokenRestrictedSids;
use windows_sys::Win32::System::Threading::OpenProcess;
use windows_sys::Win32::System::Threading::OpenProcessToken;
use windows_sys::Win32::System::Threading::PROCESS_QUERY_LIMITED_INFORMATION;
/// Returns the restricting SIDs on the process that opened an accepted loopback connection.
///
/// `accepted_local_addr` and `accepted_peer_addr` must come from the accepted server socket. The
/// owning-PID table describes the client side in the opposite direction, so the lookup matches the
/// exact reversed four-tuple.
pub(crate) fn restricting_sids_for_tcp_connection(
accepted_local_addr: SocketAddr,
accepted_peer_addr: SocketAddr,
) -> io::Result<Vec<String>> {
let (SocketAddr::V4(accepted_local_addr), SocketAddr::V4(accepted_peer_addr)) =
(accepted_local_addr, accepted_peer_addr)
else {
return Err(io::Error::new(
io::ErrorKind::Unsupported,
"Windows proxy connection attribution currently supports IPv4 only",
));
};
let process_id = owning_process_id(accepted_local_addr, accepted_peer_addr)?;
restricting_sids_for_process(process_id)
}
fn owning_process_id(
accepted_local_addr: SocketAddrV4,
accepted_peer_addr: SocketAddrV4,
) -> io::Result<u32> {
let mut byte_len = 0_u32;
let result = unsafe {
GetExtendedTcpTable(
std::ptr::null_mut(),
&mut byte_len,
0,
AF_INET as u32,
TCP_TABLE_OWNER_PID_CONNECTIONS,
0,
)
};
if result != ERROR_INSUFFICIENT_BUFFER {
return Err(win32_error("query IPv4 TCP owner table size", result));
}
let buffer = loop {
let mut buffer = aligned_buffer(byte_len as usize)?;
let result = unsafe {
GetExtendedTcpTable(
buffer.as_mut_ptr().cast::<c_void>(),
&mut byte_len,
0,
AF_INET as u32,
TCP_TABLE_OWNER_PID_CONNECTIONS,
0,
)
};
match result {
NO_ERROR => break buffer,
ERROR_INSUFFICIENT_BUFFER => continue,
_ => return Err(win32_error("read IPv4 TCP owner table", result)),
}
};
let rows = parse_tcp_owner_rows(&buffer, byte_len as usize)?;
unique_client_process_id(rows, accepted_local_addr, accepted_peer_addr)
}
fn parse_tcp_owner_rows(buffer: &[usize], byte_len: usize) -> io::Result<&[MIB_TCPROW_OWNER_PID]> {
let rows_offset = offset_of!(MIB_TCPTABLE_OWNER_PID, table);
if byte_len > size_of_val(buffer) || byte_len < rows_offset {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid IPv4 TCP owner table length",
));
}
let row_count = unsafe { std::ptr::read_unaligned(buffer.as_ptr().cast::<u32>()) } as usize;
let rows_byte_len = row_count
.checked_mul(size_of::<MIB_TCPROW_OWNER_PID>())
.and_then(|len| rows_offset.checked_add(len))
.ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
"IPv4 TCP owner table length overflow",
)
})?;
if rows_byte_len > byte_len {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"truncated IPv4 TCP owner table",
));
}
let rows = unsafe {
let rows_ptr = buffer
.as_ptr()
.cast::<u8>()
.add(rows_offset)
.cast::<MIB_TCPROW_OWNER_PID>();
std::slice::from_raw_parts(rows_ptr, row_count)
};
Ok(rows)
}
fn unique_client_process_id(
rows: &[MIB_TCPROW_OWNER_PID],
accepted_local_addr: SocketAddrV4,
accepted_peer_addr: SocketAddrV4,
) -> io::Result<u32> {
let mut matching_process_ids = rows
.iter()
.filter(|row| client_row_matches(row, accepted_local_addr, accepted_peer_addr))
.map(|row| row.dwOwningPid);
let process_id = matching_process_ids.next().ok_or_else(|| {
io::Error::new(
io::ErrorKind::NotFound,
"accepted connection is absent from the IPv4 TCP owner table",
)
})?;
if matching_process_ids.next().is_some() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"accepted connection has multiple IPv4 TCP owner rows",
));
}
Ok(process_id)
}
fn client_row_matches(
row: &MIB_TCPROW_OWNER_PID,
accepted_local_addr: SocketAddrV4,
accepted_peer_addr: SocketAddrV4,
) -> bool {
ipv4_addr_matches(row.dwLocalAddr, *accepted_peer_addr.ip())
&& tcp_port(row.dwLocalPort) == accepted_peer_addr.port()
&& ipv4_addr_matches(row.dwRemoteAddr, *accepted_local_addr.ip())
&& tcp_port(row.dwRemotePort) == accepted_local_addr.port()
}
fn ipv4_addr_matches(table_addr: u32, socket_addr: Ipv4Addr) -> bool {
table_addr.to_ne_bytes() == socket_addr.octets()
}
fn tcp_port(table_port: u32) -> u16 {
u16::from_be(table_port as u16)
}
fn restricting_sids_for_process(process_id: u32) -> io::Result<Vec<String>> {
let process_handle = unsafe { OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION, 0, process_id) };
let process = owned_handle(process_handle, "open proxy client process")?;
let mut token_handle: HANDLE = 0;
let opened = unsafe {
OpenProcessToken(
process.as_raw_handle() as HANDLE,
TOKEN_QUERY,
&mut token_handle,
)
};
if opened == 0 {
return Err(last_error("open proxy client process token"));
}
let token = owned_handle(token_handle, "open proxy client process token")?;
let mut byte_len = 0_u32;
let queried = unsafe {
GetTokenInformation(
token.as_raw_handle() as HANDLE,
TokenRestrictedSids,
std::ptr::null_mut(),
0,
&mut byte_len,
)
};
if queried != 0 || unsafe { GetLastError() } != ERROR_INSUFFICIENT_BUFFER {
return Err(last_error("query proxy client restricting SID buffer size"));
}
let mut buffer = aligned_buffer(byte_len as usize)?;
let queried = unsafe {
GetTokenInformation(
token.as_raw_handle() as HANDLE,
TokenRestrictedSids,
buffer.as_mut_ptr().cast::<c_void>(),
byte_len,
&mut byte_len,
)
};
if queried == 0 {
return Err(last_error("read proxy client restricting SIDs"));
}
parse_token_groups(&buffer, byte_len as usize)?
.iter()
.map(|entry| sid_to_string(entry.Sid))
.collect()
}
fn parse_token_groups(buffer: &[usize], byte_len: usize) -> io::Result<&[SID_AND_ATTRIBUTES]> {
let groups_offset = offset_of!(TOKEN_GROUPS, Groups);
if byte_len > size_of_val(buffer) || byte_len < groups_offset {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid restricting SID buffer length",
));
}
let group_count = unsafe { std::ptr::read_unaligned(buffer.as_ptr().cast::<u32>()) } as usize;
let groups_byte_len = group_count
.checked_mul(size_of::<SID_AND_ATTRIBUTES>())
.and_then(|len| groups_offset.checked_add(len))
.ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
"restricting SID buffer length overflow",
)
})?;
if groups_byte_len > byte_len {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"truncated restricting SID buffer",
));
}
let groups = unsafe {
let groups_ptr = buffer
.as_ptr()
.cast::<u8>()
.add(groups_offset)
.cast::<SID_AND_ATTRIBUTES>();
std::slice::from_raw_parts(groups_ptr, group_count)
};
Ok(groups)
}
fn sid_to_string(sid: PSID) -> io::Result<String> {
let mut string_sid = std::ptr::null_mut();
if unsafe { ConvertSidToStringSidW(sid, &mut string_sid) } == 0 {
return Err(last_error("convert proxy client restricting SID to string"));
}
let value = unsafe {
let mut len = 0;
while *string_sid.add(len) != 0 {
len += 1;
}
String::from_utf16_lossy(std::slice::from_raw_parts(string_sid, len))
};
unsafe {
LocalFree(string_sid as HLOCAL);
}
Ok(value)
}
fn aligned_buffer(byte_len: usize) -> io::Result<Vec<usize>> {
if byte_len == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Windows API returned an empty buffer length",
));
}
Ok(vec![0; byte_len.div_ceil(size_of::<usize>())])
}
fn owned_handle(handle: HANDLE, operation: &str) -> io::Result<OwnedHandle> {
if handle == 0 {
return Err(last_error(operation));
}
Ok(unsafe { OwnedHandle::from_raw_handle(handle as RawHandle) })
}
fn win32_error(operation: &str, error_code: u32) -> io::Error {
let error = io::Error::from_raw_os_error(error_code as i32);
io::Error::new(error.kind(), format!("{operation}: {error}"))
}
fn last_error(operation: &str) -> io::Error {
let error = io::Error::last_os_error();
io::Error::new(error.kind(), format!("{operation}: {error}"))
}
#[cfg(test)]
#[path = "windows_tcp_attribution_tests.rs"]
mod tests;

View File

@@ -0,0 +1,112 @@
use super::*;
use pretty_assertions::assert_eq;
use std::net::TcpListener;
use std::net::TcpStream;
#[test]
fn parses_owner_table_and_matches_reversed_client_tuple() -> io::Result<()> {
let proxy_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 3128);
let client_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152);
let rows = [
tcp_row(proxy_addr, client_addr, 100),
tcp_row(client_addr, proxy_addr, 200),
];
let (buffer, byte_len) = owner_table_buffer(&rows);
let parsed = parse_tcp_owner_rows(&buffer, byte_len)?;
assert_eq!(
unique_client_process_id(parsed, proxy_addr, client_addr)?,
200
);
Ok(())
}
#[test]
fn rejects_multiple_matching_owner_rows() {
let proxy_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 3128);
let client_addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152);
let rows = [
tcp_row(client_addr, proxy_addr, 200),
tcp_row(client_addr, proxy_addr, 201),
];
let error = unique_client_process_id(&rows, proxy_addr, client_addr)
.expect_err("duplicate connection rows should fail closed");
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn rejects_truncated_owner_table() {
let byte_len = offset_of!(MIB_TCPTABLE_OWNER_PID, table);
let mut buffer = aligned_buffer(byte_len).expect("aligned table buffer");
unsafe {
std::ptr::write_unaligned(buffer.as_mut_ptr().cast::<u32>(), 1);
}
let Err(error) = parse_tcp_owner_rows(&buffer, byte_len) else {
panic!("truncated connection row should fail closed");
};
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn resolves_loopback_connection_to_current_process() -> io::Result<()> {
let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))?;
let client = TcpStream::connect(listener.local_addr()?)?;
let (accepted, _) = listener.accept()?;
let local_addr = accepted.local_addr()?;
let peer_addr = accepted.peer_addr()?;
let process_id = owning_process_id(socket_addr_v4(local_addr)?, socket_addr_v4(peer_addr)?)?;
let restricting_sids = restricting_sids_for_tcp_connection(local_addr, peer_addr)?;
assert_eq!(process_id, std::process::id());
assert!(restricting_sids.iter().all(|sid| sid.starts_with("S-")));
drop(client);
Ok(())
}
fn socket_addr_v4(addr: SocketAddr) -> io::Result<SocketAddrV4> {
match addr {
SocketAddr::V4(addr) => Ok(addr),
SocketAddr::V6(_) => Err(io::Error::new(
io::ErrorKind::InvalidData,
"test listener unexpectedly used IPv6",
)),
}
}
fn tcp_row(
local_addr: SocketAddrV4,
remote_addr: SocketAddrV4,
process_id: u32,
) -> MIB_TCPROW_OWNER_PID {
MIB_TCPROW_OWNER_PID {
dwState: 0,
dwLocalAddr: u32::from_ne_bytes(local_addr.ip().octets()),
dwLocalPort: local_addr.port().to_be() as u32,
dwRemoteAddr: u32::from_ne_bytes(remote_addr.ip().octets()),
dwRemotePort: remote_addr.port().to_be() as u32,
dwOwningPid: process_id,
}
}
fn owner_table_buffer(rows: &[MIB_TCPROW_OWNER_PID]) -> (Vec<usize>, usize) {
let rows_offset = offset_of!(MIB_TCPTABLE_OWNER_PID, table);
let rows_byte_len = size_of_val(rows);
let byte_len = rows_offset + rows_byte_len;
let mut buffer = aligned_buffer(byte_len).expect("aligned table buffer");
unsafe {
let buffer_ptr = buffer.as_mut_ptr().cast::<u8>();
std::ptr::write_unaligned(buffer_ptr.cast::<u32>(), rows.len() as u32);
std::ptr::copy_nonoverlapping(
rows.as_ptr().cast::<u8>(),
buffer_ptr.add(rows_offset),
rows_byte_len,
);
}
(buffer, byte_len)
}

View File

@@ -0,0 +1,590 @@
#![cfg(target_os = "windows")]
use codex_network_proxy::ConfigReloader;
use codex_network_proxy::ConfigReloaderFuture;
use codex_network_proxy::ConfigState;
use codex_network_proxy::NetworkDecision;
use codex_network_proxy::NetworkMode;
use codex_network_proxy::NetworkPolicyDecider;
use codex_network_proxy::NetworkPolicyRequest;
use codex_network_proxy::NetworkProtocol;
use codex_network_proxy::NetworkProxy;
use codex_network_proxy::NetworkProxyConfig;
use codex_network_proxy::NetworkProxyState;
use codex_network_proxy::build_config_state;
use codex_windows_sandbox::ConsoleMode;
use codex_windows_sandbox::LocalSid;
use codex_windows_sandbox::create_process_as_user;
use codex_windows_sandbox::create_readonly_token_with_caps_and_user_from;
use codex_windows_sandbox::get_current_token_for_restriction;
use pretty_assertions::assert_eq;
use std::collections::HashMap;
use std::io::BufRead;
use std::io::BufReader;
use std::io::Read;
use std::io::Write;
use std::net::Ipv4Addr;
use std::net::SocketAddr;
use std::net::TcpListener;
use std::net::TcpStream;
use std::os::windows::io::AsRawHandle;
use std::os::windows::io::FromRawHandle;
use std::os::windows::io::OwnedHandle;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Duration;
use tokio::io::AsyncReadExt;
use tokio::io::AsyncWriteExt;
use windows_sys::Win32::System::Threading::GetExitCodeProcess;
use windows_sys::Win32::System::Threading::TerminateProcess;
use windows_sys::Win32::System::Threading::WaitForSingleObject;
const CHILD_MODE_ENV: &str = "CODEX_WINDOWS_PROXY_TEST_CHILD";
const HTTP_ADDR_ENV: &str = "CODEX_WINDOWS_PROXY_TEST_HTTP_ADDR";
const SOCKS_ADDR_ENV: &str = "CODEX_WINDOWS_PROXY_TEST_SOCKS_ADDR";
const ORIGIN_PORT_ENV: &str = "CODEX_WINDOWS_PROXY_TEST_ORIGIN_PORT";
const ALLOWED_HOST_ENV: &str = "CODEX_WINDOWS_PROXY_TEST_ALLOWED_HOST";
const DENIED_HOST_ENV: &str = "CODEX_WINDOWS_PROXY_TEST_DENIED_HOST";
const FIRST_ENVIRONMENT_ID: &str = "first-environment";
const SECOND_ENVIRONMENT_ID: &str = "second-environment";
const DECIDER_DENIED_HOST: &str = "not-allowed.invalid";
const CHILD_TIMEOUT_MS: u32 = 30_000;
const WAIT_OBJECT_0: u32 = 0;
#[derive(Clone)]
struct StaticReloader(ConfigState);
impl ConfigReloader for StaticReloader {
fn source_label(&self) -> String {
"test config".to_string()
}
fn maybe_reload(&self) -> ConfigReloaderFuture<'_, Option<ConfigState>> {
Box::pin(async { Ok(None) })
}
fn reload_now(&self) -> ConfigReloaderFuture<'_, ConfigState> {
let state = self.0.clone();
Box::pin(async move { Ok(state) })
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn restricted_tokens_select_stable_routes_and_cleanup() -> anyhow::Result<()> {
let (origin_port, origin_task) = start_http_origin().await?;
let (first_decider, first_requests) = recording_decider();
let (second_decider, second_requests) = recording_decider();
let first_requested = requested_addrs()?;
let first = build_proxy(
first_requested,
"localhost",
/*enable_socks5*/ false,
Some(first_decider),
)
.await?;
let initial_addrs = (first.http_addr(), first.socks_addr());
let first_handle = first.run().await?;
let first_sid = first
.network_proxy_restricting_sid(None)
.expect("running proxy should have a route SID");
first.prepare_for_optional_environment(HashMap::new(), Some(FIRST_ENVIRONMENT_ID))?;
let first_environment_sid = first
.network_proxy_restricting_sid(Some(FIRST_ENVIRONMENT_ID))
.expect("first environment should have a route SID");
let second_requested = requested_addrs()?;
assert_ne!(second_requested.0, initial_addrs.0);
let second = build_proxy(
second_requested,
"127.0.0.1",
/*enable_socks5*/ true,
Some(second_decider),
)
.await?;
let stable_addrs = (second.http_addr(), second.socks_addr());
assert_eq!(stable_addrs.0, initial_addrs.0);
assert_eq!(stable_addrs.1, second_requested.1);
assert_eq!((first.http_addr(), first.socks_addr()), stable_addrs);
let second_handle = second.run().await?;
let second_sid = second
.network_proxy_restricting_sid(None)
.expect("running proxy should have a route SID");
assert_ne!(second_sid, first_sid);
second.prepare_for_optional_environment(HashMap::new(), Some(SECOND_ENVIRONMENT_ID))?;
let second_environment_sid = second
.network_proxy_restricting_sid(Some(SECOND_ENVIRONMENT_ID))
.expect("second environment should have a route SID");
run_restricted_child(
&first_environment_sid,
stable_addrs,
origin_port,
Some(("localhost", DECIDER_DENIED_HOST)),
/*expect_socks*/ false,
)
.await?;
assert_recorded_requests(
&first_requests,
FIRST_ENVIRONMENT_ID,
DECIDER_DENIED_HOST,
origin_port,
&[NetworkProtocol::Http],
);
assert!(
second_requests
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.is_empty()
);
run_restricted_child(
&second_environment_sid,
stable_addrs,
origin_port,
Some(("127.0.0.1", DECIDER_DENIED_HOST)),
/*expect_socks*/ true,
)
.await?;
assert_recorded_requests(
&first_requests,
FIRST_ENVIRONMENT_ID,
DECIDER_DENIED_HOST,
origin_port,
&[NetworkProtocol::Http],
);
assert_recorded_requests(
&second_requests,
SECOND_ENVIRONMENT_ID,
DECIDER_DENIED_HOST,
origin_port,
&[NetworkProtocol::Http, NetworkProtocol::Socks5Tcp],
);
run_restricted_child(
&first_sid,
stable_addrs,
origin_port,
Some(("localhost", "127.0.0.1")),
/*expect_socks*/ false,
)
.await?;
run_restricted_child(
&second_sid,
stable_addrs,
origin_port,
Some(("127.0.0.1", "localhost")),
/*expect_socks*/ true,
)
.await?;
first_handle.shutdown().await?;
run_restricted_child(
&first_sid,
stable_addrs,
origin_port,
None,
/*expect_socks*/ false,
)
.await?;
run_restricted_child(
&first_environment_sid,
stable_addrs,
origin_port,
None,
/*expect_socks*/ false,
)
.await?;
run_restricted_child(
&second_sid,
stable_addrs,
origin_port,
Some(("127.0.0.1", "localhost")),
/*expect_socks*/ true,
)
.await?;
second_handle.shutdown().await?;
drop((first, second));
let third_requested = requested_addrs()?;
assert_ne!(third_requested, stable_addrs);
let third = build_proxy(
third_requested,
"localhost",
/*enable_socks5*/ false,
None,
)
.await?;
assert_eq!((third.http_addr(), third.socks_addr()), stable_addrs);
let third_handle = third.run().await?;
let third_sid = third
.network_proxy_restricting_sid(None)
.expect("running proxy should have a route SID");
assert_ne!(third_sid, first_sid);
assert_ne!(third_sid, second_sid);
run_restricted_child(
&third_sid,
stable_addrs,
origin_port,
Some(("localhost", "127.0.0.1")),
/*expect_socks*/ false,
)
.await?;
drop(third_handle);
assert_eq!(third.network_proxy_restricting_sid(None), None);
run_restricted_child(
&third_sid,
stable_addrs,
origin_port,
None,
/*expect_socks*/ false,
)
.await?;
origin_task.abort();
Ok(())
}
#[test]
fn restricted_child_exercises_http_and_socks() -> anyhow::Result<()> {
let Ok(mode) = std::env::var(CHILD_MODE_ENV) else {
return Ok(());
};
let http_addr = required_env(HTTP_ADDR_ENV)?.parse::<SocketAddr>()?;
let socks_addr = required_env(SOCKS_ADDR_ENV)?.parse::<SocketAddr>()?;
let origin_port = required_env(ORIGIN_PORT_ENV)?.parse::<u16>()?;
if mode == "missing-route" {
let authority = format!("localhost:{origin_port}");
assert!(http_status(http_addr, &authority).is_err());
assert!(socks_status(socks_addr, "localhost", origin_port).is_err());
return Ok(());
}
let allowed_host = required_env(ALLOWED_HOST_ENV)?;
let denied_host = required_env(DENIED_HOST_ENV)?;
let allowed_authority = format!("{allowed_host}:{origin_port}");
let denied_authority = format!("{denied_host}:{origin_port}");
assert_eq!(http_status(http_addr, &allowed_authority)?, 200);
assert_eq!(http_status(http_addr, &denied_authority)?, 403);
if mode == "http-only" {
assert!(socks_status(socks_addr, &allowed_host, origin_port).is_err());
return Ok(());
}
assert_eq!(
socks_status(socks_addr, &allowed_host, origin_port)?,
SocksOutcome::Connected
);
assert!(matches!(
socks_status(socks_addr, &denied_host, origin_port)?,
SocksOutcome::Denied(_)
));
Ok(())
}
async fn build_proxy(
requested_addrs: (SocketAddr, SocketAddr),
allowed_domain: &str,
enable_socks5: bool,
policy_decider: Option<Arc<dyn NetworkPolicyDecider>>,
) -> anyhow::Result<NetworkProxy> {
let (http_addr, socks_addr) = requested_addrs;
let mut config = NetworkProxyConfig {
enabled: true,
proxy_url: format!("http://{http_addr}"),
socks_url: format!("socks5://{socks_addr}"),
enable_socks5,
enable_socks5_udp: false,
allow_local_binding: true,
mode: NetworkMode::Full,
..NetworkProxyConfig::default()
};
config.set_allowed_domains(vec![allowed_domain.to_string()]);
let config_state = build_config_state(config, Default::default())?;
let reloader = Arc::new(StaticReloader(config_state.clone()));
let state = Arc::new(NetworkProxyState::with_reloader(config_state, reloader));
let mut builder = NetworkProxy::builder().state(state);
if let Some(policy_decider) = policy_decider {
builder = builder.policy_decider_arc(policy_decider);
}
builder.build().await
}
fn recording_decider() -> (
Arc<dyn NetworkPolicyDecider>,
Arc<Mutex<Vec<NetworkPolicyRequest>>>,
) {
let requests = Arc::new(Mutex::new(Vec::new()));
let recorded_requests = Arc::clone(&requests);
let decider: Arc<dyn NetworkPolicyDecider> = Arc::new(move |request: NetworkPolicyRequest| {
recorded_requests
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.push(request);
async { NetworkDecision::deny("integration test denial") }
});
(decider, requests)
}
fn assert_recorded_requests(
requests: &Arc<Mutex<Vec<NetworkPolicyRequest>>>,
environment_id: &str,
host: &str,
port: u16,
expected_protocols: &[NetworkProtocol],
) {
let requests = requests
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
assert_eq!(requests.len(), expected_protocols.len());
let actual_protocols = requests
.iter()
.map(|request| request.protocol)
.collect::<Vec<_>>();
assert_eq!(actual_protocols, expected_protocols);
assert!(requests.iter().all(|request| {
request.environment_id.as_deref() == Some(environment_id)
&& request.host == host
&& request.port == port
}));
}
fn requested_addrs() -> std::io::Result<(SocketAddr, SocketAddr)> {
let http = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))?;
let socks = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))?;
Ok((http.local_addr()?, socks.local_addr()?))
}
async fn start_http_origin() -> std::io::Result<(u16, tokio::task::JoinHandle<()>)> {
let listener = tokio::net::TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).await?;
let port = listener.local_addr()?.port();
let task = tokio::spawn(async move {
while let Ok((mut stream, _)) = listener.accept().await {
tokio::spawn(async move {
let mut request = [0_u8; 1024];
let _ = stream.read(&mut request).await;
let _ = stream
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nOK",
)
.await;
});
}
});
Ok((port, task))
}
async fn run_restricted_child(
route_sid: &str,
proxy_addrs: (SocketAddr, SocketAddr),
origin_port: u16,
policy: Option<(&str, &str)>,
expect_socks: bool,
) -> anyhow::Result<()> {
let route_sid = route_sid.to_string();
let policy = policy.map(|(allowed, denied)| (allowed.to_string(), denied.to_string()));
tokio::task::spawn_blocking(move || {
run_restricted_child_blocking(&route_sid, proxy_addrs, origin_port, policy, expect_socks)
})
.await??;
Ok(())
}
fn run_restricted_child_blocking(
route_sid: &str,
(http_addr, socks_addr): (SocketAddr, SocketAddr),
origin_port: u16,
policy: Option<(String, String)>,
expect_socks: bool,
) -> anyhow::Result<()> {
let route_sid = LocalSid::from_string(route_sid)?;
let capability_sid = LocalSid::from_string("S-1-5-21-10-20-30-40")?;
let base_token = unsafe {
OwnedHandle::from_raw_handle(get_current_token_for_restriction()? as *mut std::ffi::c_void)
};
let restricted_token = unsafe {
create_readonly_token_with_caps_and_user_from(
base_token.as_raw_handle() as isize,
&[capability_sid.as_ptr()],
&[route_sid.as_ptr()],
)?
};
let restricted_token =
unsafe { OwnedHandle::from_raw_handle(restricted_token as *mut std::ffi::c_void) };
let mut env = std::env::vars().collect::<HashMap<_, _>>();
env.insert(HTTP_ADDR_ENV.to_string(), http_addr.to_string());
env.insert(SOCKS_ADDR_ENV.to_string(), socks_addr.to_string());
env.insert(ORIGIN_PORT_ENV.to_string(), origin_port.to_string());
match policy {
Some((allowed, denied)) => {
let mode = if expect_socks { "policy" } else { "http-only" };
env.insert(CHILD_MODE_ENV.to_string(), mode.to_string());
env.insert(ALLOWED_HOST_ENV.to_string(), allowed);
env.insert(DENIED_HOST_ENV.to_string(), denied);
}
None => {
env.insert(CHILD_MODE_ENV.to_string(), "missing-route".to_string());
}
}
let test_exe = std::env::current_exe()?;
let command = vec![
test_exe.to_string_lossy().into_owned(),
"--exact".to_string(),
"restricted_child_exercises_http_and_socks".to_string(),
"--nocapture".to_string(),
"--test-threads=1".to_string(),
];
let cwd = std::env::current_dir()?;
let spawned = unsafe {
create_process_as_user(
restricted_token.as_raw_handle() as isize,
&command,
&cwd,
&env,
/*logs_base_dir*/ None,
/*stdio*/ None,
/*console_mode*/ ConsoleMode::Inherit,
/*use_private_desktop*/ false,
)?
};
let process = unsafe {
OwnedHandle::from_raw_handle(spawned.process_info.hProcess as *mut std::ffi::c_void)
};
let _thread = unsafe {
OwnedHandle::from_raw_handle(spawned.process_info.hThread as *mut std::ffi::c_void)
};
let wait = unsafe {
WaitForSingleObject(
process.as_raw_handle() as isize,
/*dwMilliseconds*/ CHILD_TIMEOUT_MS,
)
};
if wait != WAIT_OBJECT_0 {
unsafe {
TerminateProcess(process.as_raw_handle() as isize, 1);
}
}
let mut exit_code = 1_u32;
unsafe {
GetExitCodeProcess(process.as_raw_handle() as isize, &mut exit_code);
}
anyhow::ensure!(
wait == WAIT_OBJECT_0 && exit_code == 0,
"restricted proxy child failed (wait={wait}, exit={exit_code})"
);
Ok(())
}
fn required_env(key: &str) -> anyhow::Result<String> {
std::env::var(key).map_err(Into::into)
}
fn http_status(proxy_addr: SocketAddr, authority: &str) -> std::io::Result<u16> {
let mut stream = TcpStream::connect(proxy_addr)?;
configure_stream(&stream)?;
write!(
stream,
"GET http://{authority}/ HTTP/1.1\r\nHost: {authority}\r\nConnection: close\r\n\r\n"
)?;
read_http_status(&mut stream)
}
#[derive(Debug, Eq, PartialEq)]
enum SocksOutcome {
Connected,
Denied(u8),
}
fn socks_status(
proxy_addr: SocketAddr,
host: &str,
origin_port: u16,
) -> std::io::Result<SocksOutcome> {
let mut stream = TcpStream::connect(proxy_addr)?;
configure_stream(&stream)?;
stream.write_all(&[5, 1, 0])?;
let mut greeting = [0_u8; 2];
stream.read_exact(&mut greeting)?;
if greeting != [5, 0] {
return Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"SOCKS5 proxy rejected no-authentication method",
));
}
let mut request = vec![5, 1, 0];
if let Ok(ip) = host.parse::<Ipv4Addr>() {
request.push(1);
request.extend_from_slice(&ip.octets());
} else {
let host_len = u8::try_from(host.len()).map_err(|_| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, "SOCKS5 hostname too long")
})?;
request.extend_from_slice(&[3, host_len]);
request.extend_from_slice(host.as_bytes());
}
request.extend_from_slice(&origin_port.to_be_bytes());
stream.write_all(&request)?;
let mut reply = [0_u8; 4];
stream.read_exact(&mut reply)?;
if reply[0] != 5 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"invalid SOCKS5 response version",
));
}
if reply[1] != 0 {
return Ok(SocksOutcome::Denied(reply[1]));
}
consume_socks_bound_address(&mut stream, reply[3])?;
Ok(SocksOutcome::Connected)
}
fn consume_socks_bound_address(stream: &mut TcpStream, address_type: u8) -> std::io::Result<()> {
let address_len = match address_type {
1 => 4,
3 => {
let mut len = [0_u8; 1];
stream.read_exact(&mut len)?;
usize::from(len[0])
}
4 => 16,
_ => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"invalid SOCKS5 bound address type",
));
}
};
let mut address_and_port = vec![0_u8; address_len + 2];
stream.read_exact(&mut address_and_port)
}
fn configure_stream(stream: &TcpStream) -> std::io::Result<()> {
let timeout = Some(Duration::from_secs(5));
stream.set_read_timeout(timeout)?;
stream.set_write_timeout(timeout)
}
fn read_http_status(stream: &mut TcpStream) -> std::io::Result<u16> {
let mut status_line = String::new();
if BufReader::new(stream).read_line(&mut status_line)? == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"proxy closed before an HTTP status line",
));
}
status_line
.split_ascii_whitespace()
.nth(1)
.ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidData, "missing HTTP status code")
})?
.parse::<u16>()
.map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidData, err))
}