use std::{
fs,
net::SocketAddr,
path::{Path, PathBuf},
time::Duration,
};
use anyhow::{Context, Result, anyhow, bail};
use serde::Deserialize;
use tokio_tungstenite::tungstenite::http::{HeaderName, HeaderValue};
use crate::upstream::UpstreamProxy;
pub const DEFAULT_BUFFER_SIZE: usize = 16 * 1024;
pub const DEFAULT_LISTEN: &str = "127.0.0.1:3128";
pub const DEFAULT_RULE_REFRESH_INTERVAL_SECS: u64 = 60;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum ProxyMode {
Auto,
Global,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum AuthMode {
#[default]
Token,
Basic,
}
#[derive(Debug, Clone)]
pub struct Settings {
pub listen: SocketAddr,
pub socks_listen: Option<SocketAddr>,
pub gateway: String,
pub basic_auth: Option<String>,
pub buffer_size: usize,
pub log_level: Option<String>,
pub custom_domain_rules: Option<PathBuf>,
pub rule_refresh_interval: Duration,
pub proxy_mode: ProxyMode,
pub insecure: bool,
pub auth_mode: AuthMode,
pub upstream_proxy: Option<UpstreamProxy>,
pub headers: Vec<(HeaderName, HeaderValue)>,
}
impl Settings {
pub fn add_header(&mut self, name: &str, value: &str) -> Result<()> {
let name = HeaderName::from_bytes(name.trim().as_bytes())
.with_context(|| format!("invalid header name {name:?}"))?;
let value = HeaderValue::from_str(value)
.with_context(|| format!("invalid value for header {name}"))?;
let lowered = name.as_str();
if matches!(lowered, "host" | "connection" | "upgrade")
|| lowered.starts_with("sec-websocket-")
{
bail!("header {name} is reserved for the websocket handshake");
}
self.headers.retain(|(existing, _)| *existing != name);
self.headers.push((name, value));
Ok(())
}
}
#[derive(Debug, Default, Deserialize)]
struct FileSettings {
listen: Option<SocketAddr>,
socks_listen: Option<SocketAddr>,
gateway: Option<String>,
basic_auth: Option<String>,
buffer_size: Option<usize>,
log_level: Option<String>,
custom_domain_rules: Option<PathBuf>,
rule_refresh_interval_secs: Option<u64>,
proxy_mode: Option<ProxyMode>,
insecure: Option<bool>,
auth_mode: Option<AuthMode>,
upstream_proxy: Option<String>,
}
#[derive(Debug, Default)]
pub struct SettingsOverrides {
pub config: Option<PathBuf>,
pub listen: Option<SocketAddr>,
pub socks_listen: Option<SocketAddr>,
pub gateway: Option<String>,
pub basic_auth: Option<String>,
pub buffer_size: Option<usize>,
pub log_level: Option<String>,
pub custom_domain_rules: Option<PathBuf>,
pub rule_refresh_interval_secs: Option<u64>,
pub proxy_mode: Option<ProxyMode>,
pub insecure: bool,
pub auth_mode: Option<AuthMode>,
pub upstream_proxy: Option<String>,
}
impl Settings {
pub fn resolve(overrides: SettingsOverrides) -> Result<Self> {
let (file_settings, config_dir) = match &overrides.config {
Some(path) => (
read_file_settings(path)?,
path.parent().map(Path::to_path_buf),
),
None => (FileSettings::default(), None),
};
let listen = overrides
.listen
.or(file_settings.listen)
.or_else(|| DEFAULT_LISTEN.parse().ok())
.ok_or_else(|| anyhow!("invalid default listen address {DEFAULT_LISTEN}"))?;
let gateway = overrides
.gateway
.or(file_settings.gateway)
.ok_or_else(|| anyhow!("--gateway is required unless provided by --config"))?;
let buffer_size = overrides
.buffer_size
.or(file_settings.buffer_size)
.unwrap_or(DEFAULT_BUFFER_SIZE);
if buffer_size == 0 {
bail!("--buffer-size must be greater than 0");
}
let rule_refresh_interval_secs = overrides
.rule_refresh_interval_secs
.or(file_settings.rule_refresh_interval_secs)
.unwrap_or(DEFAULT_RULE_REFRESH_INTERVAL_SECS);
if rule_refresh_interval_secs == 0 {
bail!("--rule-refresh-interval-secs must be greater than 0");
}
Ok(Self {
listen,
socks_listen: overrides.socks_listen.or(file_settings.socks_listen),
gateway,
basic_auth: overrides.basic_auth.or(file_settings.basic_auth),
buffer_size,
log_level: overrides.log_level.or(file_settings.log_level),
custom_domain_rules: overrides.custom_domain_rules.or_else(|| {
file_settings
.custom_domain_rules
.map(|path| resolve_config_relative_path(path, config_dir.as_deref()))
}),
rule_refresh_interval: Duration::from_secs(rule_refresh_interval_secs),
proxy_mode: overrides
.proxy_mode
.or(file_settings.proxy_mode)
.unwrap_or(ProxyMode::Auto),
insecure: if overrides.insecure {
true
} else {
file_settings.insecure.unwrap_or(false)
},
auth_mode: overrides
.auth_mode
.or(file_settings.auth_mode)
.unwrap_or_default(),
upstream_proxy: UpstreamProxy::parse_optional(
overrides
.upstream_proxy
.or(file_settings.upstream_proxy)
.as_deref(),
)?,
headers: Vec::new(),
})
}
}
fn resolve_config_relative_path(path: PathBuf, config_dir: Option<&Path>) -> PathBuf {
if path.is_absolute() {
return path;
}
match config_dir {
Some(dir) => dir.join(path),
None => path,
}
}
fn read_file_settings(path: &Path) -> Result<FileSettings> {
let contents = fs::read_to_string(path)
.with_context(|| format!("failed to read config file {}", path.display()))?;
toml::from_str(&contents)
.with_context(|| format!("failed to parse config file {}", path.display()))
}
#[cfg(test)]
mod tests {
use super::*;
fn overrides_with_config(config: Option<std::path::PathBuf>) -> SettingsOverrides {
SettingsOverrides {
config,
listen: None,
socks_listen: None,
gateway: None,
basic_auth: None,
buffer_size: None,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval_secs: None,
proxy_mode: None,
insecure: false,
auth_mode: None,
upstream_proxy: None,
}
}
fn empty_settings() -> Settings {
Settings {
listen: DEFAULT_LISTEN.parse().unwrap(),
socks_listen: None,
gateway: "ws://example.com".to_owned(),
basic_auth: None,
buffer_size: DEFAULT_BUFFER_SIZE,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval: Duration::from_secs(DEFAULT_RULE_REFRESH_INTERVAL_SECS),
proxy_mode: ProxyMode::Global,
insecure: false,
auth_mode: AuthMode::Basic,
upstream_proxy: None,
headers: Vec::new(),
}
}
#[test]
fn add_header_replaces_same_name() {
let mut settings = empty_settings();
settings.add_header("User-Agent", "first").unwrap();
settings.add_header("user-agent", "second").unwrap();
settings.add_header("X-Extra", "1").unwrap();
assert_eq!(settings.headers.len(), 2);
assert_eq!(settings.headers[0].0, "user-agent");
assert_eq!(settings.headers[0].1, "second");
assert_eq!(settings.headers[1].0, "x-extra");
}
#[test]
fn add_header_rejects_invalid_and_reserved() {
let mut settings = empty_settings();
assert!(settings.add_header("bad name", "x").is_err());
assert!(settings.add_header("X-Bad", "line\nbreak").is_err());
assert!(settings.add_header("Host", "x").is_err());
assert!(settings.add_header("Upgrade", "x").is_err());
assert!(settings.add_header("Sec-WebSocket-Key", "x").is_err());
assert!(settings.headers.is_empty());
}
#[test]
fn rejects_missing_gateway() {
assert!(Settings::resolve(overrides_with_config(None)).is_err());
}
#[test]
fn resolves_cli_only_settings() {
let settings = Settings::resolve(SettingsOverrides {
config: None,
listen: Some("127.0.0.1:9000".parse().unwrap()),
socks_listen: None,
gateway: Some("wss://example.com/ws".to_owned()),
basic_auth: Some("user:pass".to_owned()),
buffer_size: Some(4096),
log_level: Some("debug".to_owned()),
custom_domain_rules: Some("cli-domains.txt".into()),
rule_refresh_interval_secs: Some(30),
proxy_mode: Some(ProxyMode::Global),
insecure: true,
auth_mode: None,
upstream_proxy: None,
})
.unwrap();
assert_eq!(settings.listen, "127.0.0.1:9000".parse().unwrap());
assert_eq!(settings.gateway, "wss://example.com/ws");
assert_eq!(settings.basic_auth.as_deref(), Some("user:pass"));
assert_eq!(settings.buffer_size, 4096);
assert_eq!(settings.log_level.as_deref(), Some("debug"));
assert_eq!(
settings.custom_domain_rules.as_deref(),
Some(Path::new("cli-domains.txt"))
);
assert_eq!(settings.rule_refresh_interval, Duration::from_secs(30));
assert_eq!(settings.proxy_mode, ProxyMode::Global);
assert!(settings.insecure);
}
#[test]
fn cli_overrides_file_settings() {
let config_path = std::env::temp_dir().join(format!(
"ws2tcp-local-test-{}-{}.toml",
std::process::id(),
"cli-overrides"
));
fs::write(
&config_path,
r#"
listen = "127.0.0.1:8000"
gateway = "wss://file.example/ws"
basic_auth = "file:secret"
buffer_size = 1024
log_level = "info"
custom_domain_rules = "file-domains.txt"
rule_refresh_interval_secs = 45
proxy_mode = "auto"
insecure = true
"#,
)
.unwrap();
let settings = Settings::resolve(SettingsOverrides {
config: Some(config_path.clone()),
listen: Some("127.0.0.1:9000".parse().unwrap()),
socks_listen: None,
gateway: Some("wss://cli.example/ws".to_owned()),
basic_auth: Some("cli:secret".to_owned()),
buffer_size: Some(2048),
log_level: Some("debug".to_owned()),
custom_domain_rules: Some("cli-domains.txt".into()),
rule_refresh_interval_secs: Some(30),
proxy_mode: Some(ProxyMode::Global),
insecure: false,
auth_mode: None,
upstream_proxy: None,
})
.unwrap();
let _ = fs::remove_file(&config_path);
assert_eq!(settings.listen, "127.0.0.1:9000".parse().unwrap());
assert_eq!(settings.gateway, "wss://cli.example/ws");
assert_eq!(settings.basic_auth.as_deref(), Some("cli:secret"));
assert_eq!(settings.buffer_size, 2048);
assert_eq!(settings.log_level.as_deref(), Some("debug"));
assert_eq!(
settings.custom_domain_rules.as_deref(),
Some(Path::new("cli-domains.txt"))
);
assert_eq!(settings.rule_refresh_interval, Duration::from_secs(30));
assert_eq!(settings.proxy_mode, ProxyMode::Global);
assert!(settings.insecure);
}
#[test]
fn resolves_file_only_settings() {
let config_path = std::env::temp_dir().join(format!(
"ws2tcp-local-test-{}-{}.toml",
std::process::id(),
"file-only"
));
fs::write(
&config_path,
r#"
listen = "127.0.0.1:7000"
gateway = "wss://file.example/ws"
buffer_size = 8192
log_level = "info"
custom_domain_rules = "custom-domains.txt"
rule_refresh_interval_secs = 45
proxy_mode = "global"
insecure = true
"#,
)
.unwrap();
let settings = Settings::resolve(overrides_with_config(Some(config_path.clone()))).unwrap();
let _ = fs::remove_file(&config_path);
assert_eq!(settings.listen, "127.0.0.1:7000".parse().unwrap());
assert_eq!(settings.gateway, "wss://file.example/ws");
assert_eq!(settings.basic_auth, None);
assert_eq!(settings.buffer_size, 8192);
assert_eq!(settings.log_level.as_deref(), Some("info"));
assert_eq!(
settings.custom_domain_rules.as_deref(),
Some(config_path.with_file_name("custom-domains.txt").as_path())
);
assert_eq!(settings.rule_refresh_interval, Duration::from_secs(45));
assert_eq!(settings.proxy_mode, ProxyMode::Global);
assert!(settings.insecure);
}
#[test]
fn verifies_server_certificate_by_default() {
let settings = Settings::resolve(SettingsOverrides {
config: None,
listen: None,
socks_listen: None,
gateway: Some("wss://example.com/ws".to_owned()),
basic_auth: None,
buffer_size: None,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval_secs: None,
proxy_mode: None,
insecure: false,
auth_mode: None,
upstream_proxy: None,
})
.unwrap();
assert!(!settings.insecure);
}
#[test]
fn uses_auto_proxy_mode_by_default() {
let settings = Settings::resolve(SettingsOverrides {
config: None,
listen: None,
socks_listen: None,
gateway: Some("wss://example.com/ws".to_owned()),
basic_auth: None,
buffer_size: None,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval_secs: None,
proxy_mode: None,
insecure: false,
auth_mode: None,
upstream_proxy: None,
})
.unwrap();
assert_eq!(settings.proxy_mode, ProxyMode::Auto);
}
#[test]
fn uses_default_listen_address() {
let settings = Settings::resolve(SettingsOverrides {
config: None,
listen: None,
socks_listen: None,
gateway: Some("wss://example.com/ws".to_owned()),
basic_auth: None,
buffer_size: None,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval_secs: None,
proxy_mode: None,
insecure: false,
auth_mode: None,
upstream_proxy: None,
})
.unwrap();
assert_eq!(settings.listen, "127.0.0.1:3128".parse().unwrap());
}
#[test]
fn uses_default_rule_refresh_interval() {
let settings = Settings::resolve(SettingsOverrides {
config: None,
listen: None,
socks_listen: None,
gateway: Some("wss://example.com/ws".to_owned()),
basic_auth: None,
buffer_size: None,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval_secs: None,
proxy_mode: None,
insecure: false,
auth_mode: None,
upstream_proxy: None,
})
.unwrap();
assert_eq!(
settings.rule_refresh_interval,
Duration::from_secs(DEFAULT_RULE_REFRESH_INTERVAL_SECS)
);
}
#[test]
fn rejects_zero_rule_refresh_interval() {
assert!(
Settings::resolve(SettingsOverrides {
config: None,
listen: None,
socks_listen: None,
gateway: Some("wss://example.com/ws".to_owned()),
basic_auth: None,
buffer_size: None,
log_level: None,
custom_domain_rules: None,
rule_refresh_interval_secs: Some(0),
proxy_mode: None,
insecure: false,
auth_mode: None,
upstream_proxy: None,
})
.is_err()
);
}
#[test]
fn auth_mode_defaults_to_token_and_can_come_from_the_file_or_the_flag() {
let resolve = |config: Option<PathBuf>, flag: Option<AuthMode>| {
Settings::resolve(SettingsOverrides {
config,
gateway: Some("wss://example.com/ws".to_owned()),
auth_mode: flag,
..SettingsOverrides::default()
})
.unwrap()
.auth_mode
};
assert_eq!(resolve(None, None), AuthMode::Token);
assert_eq!(resolve(None, Some(AuthMode::Basic)), AuthMode::Basic);
let config_path = std::env::temp_dir().join(format!(
"ws2tcp-local-test-{}-auth-mode.toml",
std::process::id()
));
fs::write(&config_path, "auth_mode = \"basic\"\n").unwrap();
assert_eq!(resolve(Some(config_path.clone()), None), AuthMode::Basic);
assert_eq!(
resolve(Some(config_path.clone()), Some(AuthMode::Token)),
AuthMode::Token
);
fs::write(&config_path, "auth_mode = \"both\"\n").unwrap();
assert!(
Settings::resolve(SettingsOverrides {
config: Some(config_path.clone()),
gateway: Some("wss://example.com/ws".to_owned()),
..SettingsOverrides::default()
})
.is_err()
);
let _ = fs::remove_file(&config_path);
}
#[test]
fn upstream_proxy_comes_from_the_file_and_the_command_line_wins() {
let resolve = |config: Option<PathBuf>, flag: Option<&str>| {
Settings::resolve(SettingsOverrides {
config,
gateway: Some("wss://example.com/ws".to_owned()),
upstream_proxy: flag.map(str::to_owned),
..SettingsOverrides::default()
})
.map(|settings| settings.upstream_proxy.map(|proxy| proxy.to_string()))
};
assert_eq!(resolve(None, None).unwrap(), None);
assert_eq!(
resolve(None, Some("socks5h://u:p@127.0.0.1:1080"))
.unwrap()
.as_deref(),
Some("socks5h://127.0.0.1:1080")
);
assert!(resolve(None, Some("ftp://127.0.0.1:21")).is_err());
let config_path = std::env::temp_dir().join(format!(
"ws2tcp-local-test-{}-upstream-proxy.toml",
std::process::id()
));
fs::write(
&config_path,
"upstream_proxy = \"http://file.example:3128\"\n",
)
.unwrap();
assert_eq!(
resolve(Some(config_path.clone()), None).unwrap().as_deref(),
Some("http://file.example:3128")
);
assert_eq!(
resolve(
Some(config_path.clone()),
Some("socks5h://cli.example:1080")
)
.unwrap()
.as_deref(),
Some("socks5h://cli.example:1080")
);
assert_eq!(resolve(Some(config_path.clone()), Some("")).unwrap(), None);
let _ = fs::remove_file(&config_path);
}
}