mirror of
https://github.com/openai/codex.git
synced 2026-08-23 13:09:46 +00:00
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:
@@ -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" }
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
368
codex-rs/network-proxy/src/windows_proxy_ingress.rs
Normal file
368
codex-rs/network-proxy/src/windows_proxy_ingress.rs
Normal 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;
|
||||
43
codex-rs/network-proxy/src/windows_proxy_ingress_tests.rs
Normal file
43
codex-rs/network-proxy/src/windows_proxy_ingress_tests.rs
Normal 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,
|
||||
})
|
||||
}
|
||||
315
codex-rs/network-proxy/src/windows_tcp_attribution.rs
Normal file
315
codex-rs/network-proxy/src/windows_tcp_attribution.rs
Normal 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;
|
||||
112
codex-rs/network-proxy/src/windows_tcp_attribution_tests.rs
Normal file
112
codex-rs/network-proxy/src/windows_tcp_attribution_tests.rs
Normal 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)
|
||||
}
|
||||
590
codex-rs/network-proxy/tests/windows_stable_ingress.rs
Normal file
590
codex-rs/network-proxy/tests/windows_stable_ingress.rs
Normal 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))
|
||||
}
|
||||
Reference in New Issue
Block a user