#![forbid(unsafe_code)]
use crate::errors::{SshCliError, SshCliResult};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
const CMD_CONNECT: u8 = 0x01;
const CMD_BIND: u8 = 0x02;
const CMD_UDP_ASSOCIATE: u8 = 0x03;
const ATYP_IPV4: u8 = 0x01;
const ATYP_DOMAIN: u8 = 0x03;
const ATYP_IPV6: u8 = 0x04;
pub const REP_SUCCEEDED: u8 = 0x00;
pub const REP_GENERAL_FAILURE: u8 = 0x01;
pub const REP_HOST_UNREACHABLE: u8 = 0x04;
pub const REP_COMMAND_NOT_SUPPORTED: u8 = 0x07;
pub const REP_ADDRESS_TYPE_NOT_SUPPORTED: u8 = 0x08;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Socks5Target {
pub host: String,
pub port: u16,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Socks5Refusal {
pub reply: u8,
pub reason: String,
}
impl Socks5Refusal {
fn new(reply: u8, reason: impl Into<String>) -> Self {
Self {
reply,
reason: reason.into(),
}
}
}
pub fn parse_greeting_header(header: &[u8; 2]) -> SshCliResult<usize> {
if header[0] != crate::constants::SOCKS5_VERSION {
return Err(SshCliError::InvalidArgument(format!(
"not a SOCKS5 greeting: version byte 0x{:02x}",
header[0]
)));
}
if header[1] == 0 {
return Err(SshCliError::InvalidArgument(
"SOCKS5 greeting advertises zero methods".to_string(),
));
}
Ok(usize::from(header[1]))
}
#[must_use]
pub fn offers_no_auth(methods: &[u8]) -> bool {
methods.contains(&crate::constants::SOCKS5_METHOD_NO_AUTH)
}
#[must_use]
pub fn fixed_address_len(atyp: u8) -> Option<usize> {
match atyp {
ATYP_IPV4 => Some(4),
ATYP_IPV6 => Some(16),
_ => None,
}
}
pub fn parse_request_header(header: &[u8; 4]) -> SshCliResult<Result<u8, Socks5Refusal>> {
if header[0] != crate::constants::SOCKS5_VERSION {
return Err(SshCliError::InvalidArgument(format!(
"not a SOCKS5 request: version byte 0x{:02x}",
header[0]
)));
}
match header[1] {
CMD_CONNECT => {}
CMD_BIND => {
return Ok(Err(Socks5Refusal::new(
REP_COMMAND_NOT_SUPPORTED,
"BIND is not supported: this proxy only opens outbound SSH channels",
)))
}
CMD_UDP_ASSOCIATE => {
return Ok(Err(Socks5Refusal::new(
REP_COMMAND_NOT_SUPPORTED,
"UDP ASSOCIATE is not supported: SSH forwarding is stream-only",
)))
}
other => {
return Ok(Err(Socks5Refusal::new(
REP_COMMAND_NOT_SUPPORTED,
format!("unknown SOCKS5 command 0x{other:02x}"),
)))
}
}
let atyp = header[3];
if !matches!(atyp, ATYP_IPV4 | ATYP_IPV6 | ATYP_DOMAIN) {
return Ok(Err(Socks5Refusal::new(
REP_ADDRESS_TYPE_NOT_SUPPORTED,
format!("unknown SOCKS5 address type 0x{atyp:02x}"),
)));
}
Ok(Ok(atyp))
}
pub fn decode_address(atyp: u8, bytes: &[u8]) -> SshCliResult<String> {
match atyp {
ATYP_IPV4 => {
let octets: [u8; 4] = bytes.try_into().map_err(|_| {
SshCliError::InvalidArgument(format!(
"SOCKS5 IPv4 address needs 4 bytes, got {}",
bytes.len()
))
})?;
Ok(std::net::Ipv4Addr::from(octets).to_string())
}
ATYP_IPV6 => {
let octets: [u8; 16] = bytes.try_into().map_err(|_| {
SshCliError::InvalidArgument(format!(
"SOCKS5 IPv6 address needs 16 bytes, got {}",
bytes.len()
))
})?;
Ok(std::net::Ipv6Addr::from(octets).to_string())
}
ATYP_DOMAIN => {
if bytes.is_empty() {
return Err(SshCliError::InvalidArgument(
"SOCKS5 domain name is empty".to_string(),
));
}
std::str::from_utf8(bytes).map(str::to_owned).map_err(|_| {
SshCliError::InvalidArgument("SOCKS5 domain name is not valid UTF-8".to_string())
})
}
other => Err(SshCliError::InvalidArgument(format!(
"unsupported SOCKS5 address type 0x{other:02x}"
))),
}
}
#[must_use]
pub fn encode_reply(reply_code: u8) -> [u8; 10] {
let mut frame = [0_u8; 10];
frame[0] = crate::constants::SOCKS5_VERSION;
frame[1] = reply_code;
frame[2] = 0x00; frame[3] = ATYP_IPV4;
frame
}
pub async fn handshake<S>(stream: &mut S) -> SshCliResult<Result<Socks5Target, Socks5Refusal>>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let mut budget = crate::constants::SOCKS5_HANDSHAKE_MAX_BYTES;
let mut greeting = [0_u8; 2];
read_exact_capped(stream, &mut greeting, &mut budget).await?;
let n_methods = parse_greeting_header(&greeting)?;
let mut methods = vec![0_u8; n_methods];
read_exact_capped(stream, &mut methods, &mut budget).await?;
if !offers_no_auth(&methods) {
write_all(
stream,
&[
crate::constants::SOCKS5_VERSION,
crate::constants::SOCKS5_METHOD_NONE_ACCEPTABLE,
],
)
.await?;
return Ok(Err(Socks5Refusal::new(
REP_GENERAL_FAILURE,
"client offered no acceptable authentication method",
)));
}
write_all(
stream,
&[
crate::constants::SOCKS5_VERSION,
crate::constants::SOCKS5_METHOD_NO_AUTH,
],
)
.await?;
let mut header = [0_u8; 4];
read_exact_capped(stream, &mut header, &mut budget).await?;
let atyp = match parse_request_header(&header)? {
Ok(atyp) => atyp,
Err(refusal) => {
write_all(stream, &encode_reply(refusal.reply)).await?;
return Ok(Err(refusal));
}
};
let addr_len = match fixed_address_len(atyp) {
Some(len) => len,
None => {
let mut len_byte = [0_u8; 1];
read_exact_capped(stream, &mut len_byte, &mut budget).await?;
usize::from(len_byte[0])
}
};
let mut addr = vec![0_u8; addr_len];
read_exact_capped(stream, &mut addr, &mut budget).await?;
let host = decode_address(atyp, &addr)?;
let mut port_bytes = [0_u8; 2];
read_exact_capped(stream, &mut port_bytes, &mut budget).await?;
let port = u16::from_be_bytes(port_bytes);
Ok(Ok(Socks5Target { host, port }))
}
pub async fn write_reply<S>(stream: &mut S, reply_code: u8) -> SshCliResult<()>
where
S: AsyncWrite + Unpin,
{
write_all(stream, &encode_reply(reply_code)).await
}
async fn write_all<S: AsyncWrite + Unpin>(stream: &mut S, bytes: &[u8]) -> SshCliResult<()> {
stream.write_all(bytes).await.map_err(SshCliError::Io)?;
stream.flush().await.map_err(SshCliError::Io)
}
async fn read_exact_capped<S: AsyncRead + Unpin>(
stream: &mut S,
buf: &mut [u8],
budget: &mut usize,
) -> SshCliResult<()> {
let want = buf.len();
if want > *budget {
return Err(SshCliError::InvalidArgument(format!(
"SOCKS5 handshake exceeds {} bytes",
crate::constants::SOCKS5_HANDSHAKE_MAX_BYTES
)));
}
stream.read_exact(buf).await.map_err(SshCliError::Io)?;
*budget -= want;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn greeting_rejects_socks4() {
let err = parse_greeting_header(&[0x04, 0x01]).expect_err("SOCKS4 must be rejected");
assert!(matches!(err, SshCliError::InvalidArgument(_)));
}
#[test]
fn greeting_rejects_zero_methods() {
assert!(parse_greeting_header(&[0x05, 0x00]).is_err());
}
#[test]
fn greeting_reports_method_count() {
assert_eq!(parse_greeting_header(&[0x05, 0x03]).unwrap(), 3);
}
#[test]
fn no_auth_is_detected_anywhere_in_the_list() {
assert!(offers_no_auth(&[0x02, 0x00]));
assert!(!offers_no_auth(&[0x01, 0x02]));
}
#[test]
fn bind_is_refused_with_the_rfc_code() {
let refusal = parse_request_header(&[0x05, CMD_BIND, 0x00, ATYP_IPV4])
.unwrap()
.expect_err("BIND must be refused");
assert_eq!(refusal.reply, REP_COMMAND_NOT_SUPPORTED);
}
#[test]
fn udp_associate_is_refused_with_the_rfc_code() {
let refusal = parse_request_header(&[0x05, CMD_UDP_ASSOCIATE, 0x00, ATYP_IPV4])
.unwrap()
.expect_err("UDP ASSOCIATE must be refused");
assert_eq!(refusal.reply, REP_COMMAND_NOT_SUPPORTED);
}
#[test]
fn unknown_address_type_is_refused_not_errored() {
let refusal = parse_request_header(&[0x05, CMD_CONNECT, 0x00, 0x09])
.unwrap()
.expect_err("unknown ATYP must be refused");
assert_eq!(refusal.reply, REP_ADDRESS_TYPE_NOT_SUPPORTED);
}
#[test]
fn connect_with_ipv4_is_accepted() {
assert_eq!(
parse_request_header(&[0x05, CMD_CONNECT, 0x00, ATYP_IPV4])
.unwrap()
.unwrap(),
ATYP_IPV4
);
}
#[test]
fn address_lengths_follow_the_rfc() {
assert_eq!(fixed_address_len(ATYP_IPV4), Some(4));
assert_eq!(fixed_address_len(ATYP_IPV6), Some(16));
assert_eq!(fixed_address_len(ATYP_DOMAIN), None);
}
#[test]
fn ipv4_decodes_to_dotted_quad() {
assert_eq!(
decode_address(ATYP_IPV4, &[127, 0, 0, 1]).unwrap(),
"127.0.0.1"
);
}
#[test]
fn ipv6_decodes_to_canonical_form() {
let mut bytes = [0_u8; 16];
bytes[15] = 1;
assert_eq!(decode_address(ATYP_IPV6, &bytes).unwrap(), "::1");
}
#[test]
fn domain_is_passed_through_unresolved() {
assert_eq!(
decode_address(ATYP_DOMAIN, b"internal.db.lan").unwrap(),
"internal.db.lan"
);
}
#[test]
fn empty_domain_is_rejected() {
assert!(decode_address(ATYP_DOMAIN, b"").is_err());
}
#[test]
fn non_utf8_domain_is_rejected() {
assert!(decode_address(ATYP_DOMAIN, &[0xFF, 0xFE]).is_err());
}
#[test]
fn short_ipv4_is_rejected_rather_than_padded() {
assert!(decode_address(ATYP_IPV4, &[127, 0, 0]).is_err());
}
#[test]
fn reply_frame_is_well_formed() {
let frame = encode_reply(REP_SUCCEEDED);
assert_eq!(frame[0], crate::constants::SOCKS5_VERSION);
assert_eq!(frame[1], REP_SUCCEEDED);
assert_eq!(frame[3], ATYP_IPV4);
assert_eq!(&frame[4..], &[0, 0, 0, 0, 0, 0]);
}
async fn drive(payload: &[u8]) -> (SshCliResult<Result<Socks5Target, Socks5Refusal>>, Vec<u8>) {
let (mut client, mut server) = tokio::io::duplex(4096);
client.write_all(payload).await.expect("feed handshake");
let outcome = super::handshake(&mut server).await;
drop(server);
let mut written = Vec::new();
client
.read_to_end(&mut written)
.await
.expect("drain replies");
(outcome, written)
}
#[tokio::test]
async fn handshake_parses_a_domain_connect() {
let mut payload = vec![0x05, 0x01, 0x00]; payload.extend_from_slice(&[0x05, CMD_CONNECT, 0x00, ATYP_DOMAIN]);
payload.push(9);
payload.extend_from_slice(b"localhost");
payload.extend_from_slice(&443_u16.to_be_bytes());
let (outcome, written) = drive(&payload).await;
let target = outcome
.expect("handshake must parse")
.expect("CONNECT must be accepted");
assert_eq!(target.host, "localhost");
assert_eq!(target.port, 443);
assert_eq!(written, vec![0x05, 0x00]);
}
#[tokio::test]
async fn handshake_accepts_ipv4_connect() {
let mut payload = vec![0x05, 0x02, 0x00, 0x02];
payload.extend_from_slice(&[0x05, CMD_CONNECT, 0x00, ATYP_IPV4]);
payload.extend_from_slice(&[10, 0, 0, 7]);
payload.extend_from_slice(&5432_u16.to_be_bytes());
let (outcome, _) = drive(&payload).await;
let target = outcome.unwrap().unwrap();
assert_eq!(target.host, "10.0.0.7");
assert_eq!(target.port, 5432);
}
#[tokio::test]
async fn handshake_refuses_a_client_without_no_auth() {
let (outcome, written) = drive(&[0x05, 0x01, 0x02]).await; let refusal = outcome
.expect("handshake must not error")
.expect_err("client without no-auth must be refused");
assert_eq!(refusal.reply, REP_GENERAL_FAILURE);
assert_eq!(written, vec![0x05, 0xFF]);
}
#[tokio::test]
async fn handshake_answers_bind_with_command_not_supported() {
let mut payload = vec![0x05, 0x01, 0x00];
payload.extend_from_slice(&[0x05, CMD_BIND, 0x00, ATYP_IPV4]);
let (outcome, written) = drive(&payload).await;
let refusal = outcome.unwrap().expect_err("BIND must be refused");
assert_eq!(refusal.reply, REP_COMMAND_NOT_SUPPORTED);
assert_eq!(written.len(), 2 + 10);
assert_eq!(written[2], crate::constants::SOCKS5_VERSION);
assert_eq!(written[3], REP_COMMAND_NOT_SUPPORTED);
}
#[tokio::test]
async fn handshake_rejects_a_socks4_client() {
let (outcome, _) = drive(&[0x04, 0x01, 0x00]).await;
let err = outcome.expect_err("SOCKS4 must not be parsed as SOCKS5");
assert!(matches!(err, SshCliError::InvalidArgument(_)));
}
#[test]
fn handshake_budget_covers_the_largest_legal_request() {
let largest = 2 + 255 + 4 + 1 + 255 + 2;
assert!(crate::constants::SOCKS5_HANDSHAKE_MAX_BYTES >= largest);
}
}