1use crate::{Error, Result};
4
5use caret::caret_int;
6use std::fmt;
7use std::net::IpAddr;
8
9#[cfg(feature = "arbitrary")]
10use std::net::Ipv6Addr;
11
12use tor_error::bad_api_usage;
13
14#[cfg(feature = "arbitrary")]
15use arbitrary::{Arbitrary, Result as ArbitraryResult, Unstructured};
16
17#[derive(Copy, Clone, Debug, Eq, PartialEq)]
19#[cfg_attr(feature = "arbitrary", derive(Arbitrary))]
20#[non_exhaustive]
21pub enum SocksVersion {
22 V4,
24 V5,
26}
27
28#[derive(Copy, Clone, Debug, Eq, PartialEq, thiserror::Error)]
30#[error("{0} is not a recognized socks version")]
31#[allow(clippy::exhaustive_structs)]
32pub struct InvalidSocksVersion(pub u8);
33
34impl TryFrom<u8> for SocksVersion {
35 type Error = InvalidSocksVersion;
38
39 fn try_from(v: u8) -> std::result::Result<SocksVersion, InvalidSocksVersion> {
40 match v {
41 4 => Ok(SocksVersion::V4),
42 5 => Ok(SocksVersion::V5),
43 _ => Err(InvalidSocksVersion(v)),
47 }
48 }
49}
50
51impl fmt::Display for SocksVersion {
52 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53 match self {
54 SocksVersion::V4 => write!(f, "socks4"),
55 SocksVersion::V5 => write!(f, "socks5"),
56 }
57 }
58}
59
60#[derive(Clone, Debug)]
66#[cfg_attr(test, derive(PartialEq, Eq))]
67pub struct SocksRequest {
68 version: SocksVersion,
70 cmd: SocksCmd,
72 addr: SocksAddr,
74 port: u16,
76 auth: SocksAuth,
81}
82
83#[cfg(feature = "arbitrary")]
84impl<'a> Arbitrary<'a> for SocksRequest {
85 fn arbitrary(u: &mut Unstructured<'a>) -> ArbitraryResult<Self> {
86 let version = SocksVersion::arbitrary(u)?;
87 let cmd = SocksCmd::arbitrary(u)?;
88 let addr = SocksAddr::arbitrary(u)?;
89 let port = u16::arbitrary(u)?;
90 let auth = SocksAuth::arbitrary(u)?;
91
92 SocksRequest::new(version, cmd, addr, port, auth)
93 .map_err(|_| arbitrary::Error::IncorrectFormat)
94 }
95}
96
97#[derive(Clone, Debug, PartialEq, Eq)]
99#[allow(clippy::exhaustive_enums)]
100pub enum SocksAddr {
101 Hostname(SocksHostname),
103 Ip(IpAddr),
107}
108
109#[cfg(feature = "arbitrary")]
110impl<'a> Arbitrary<'a> for SocksAddr {
111 fn arbitrary(u: &mut Unstructured<'a>) -> ArbitraryResult<Self> {
112 use std::net::Ipv4Addr;
113 let b = u8::arbitrary(u)?;
114 Ok(match b % 3 {
115 0 => SocksAddr::Hostname(SocksHostname::arbitrary(u)?),
116 1 => SocksAddr::Ip(IpAddr::V4(Ipv4Addr::arbitrary(u)?)),
117 _ => SocksAddr::Ip(IpAddr::V6(Ipv6Addr::arbitrary(u)?)),
118 })
119 }
120 fn size_hint(_depth: usize) -> (usize, Option<usize>) {
121 (1, Some(256))
122 }
123}
124
125#[derive(Clone, Debug, PartialEq, Eq)]
127pub struct SocksHostname(String);
128
129#[cfg(feature = "arbitrary")]
130impl<'a> Arbitrary<'a> for SocksHostname {
131 fn arbitrary(u: &mut Unstructured<'a>) -> ArbitraryResult<Self> {
132 String::arbitrary(u)?
133 .try_into()
134 .map_err(|_| arbitrary::Error::IncorrectFormat)
135 }
136 fn size_hint(_depth: usize) -> (usize, Option<usize>) {
137 (0, Some(255))
138 }
139}
140
141#[derive(Clone, Debug, PartialEq, Eq, Hash)]
143#[cfg_attr(feature = "arbitrary", derive(Arbitrary))]
144#[non_exhaustive]
145pub enum SocksAuth {
146 NoAuth,
148 Socks4(Vec<u8>),
150 Username(Vec<u8>, Vec<u8>),
152}
153
154caret_int! {
155 #[cfg_attr(feature = "arbitrary", derive(Arbitrary))]
157 pub struct SocksCmd(u8) {
158 CONNECT = 1,
160 BIND = 2,
162 UDP_ASSOCIATE = 3,
164
165 RESOLVE = 0xF0,
167 RESOLVE_PTR = 0xF1,
169 }
170}
171
172caret_int! {
173 #[cfg_attr(feature = "arbitrary", derive(Arbitrary))]
179 pub struct SocksStatus(u8) {
180 SUCCEEDED = 0x00,
182 GENERAL_FAILURE = 0x01,
184 NOT_ALLOWED = 0x02,
189 NETWORK_UNREACHABLE = 0x03,
191 HOST_UNREACHABLE = 0x04,
193 CONNECTION_REFUSED = 0x05,
195 TTL_EXPIRED = 0x06,
199 COMMAND_NOT_SUPPORTED = 0x07,
201 ADDRTYPE_NOT_SUPPORTED = 0x08,
203 HS_DESC_NOT_FOUND = 0xF0,
205 HS_DESC_INVALID = 0xF1,
207 HS_INTRO_FAILED = 0xF2,
209 HS_REND_FAILED = 0xF3,
211 HS_MISSING_CLIENT_AUTH = 0xF4,
213 HS_WRONG_CLIENT_AUTH = 0xF5,
215 HS_BAD_ADDRESS = 0xF6,
219 HS_INTRO_TIMEOUT = 0xF7
223 }
224}
225
226impl SocksCmd {
227 fn recognized(self) -> bool {
229 matches!(
230 self,
231 SocksCmd::CONNECT | SocksCmd::RESOLVE | SocksCmd::RESOLVE_PTR
232 )
233 }
234
235 fn requires_port(self) -> bool {
237 matches!(
238 self,
239 SocksCmd::CONNECT | SocksCmd::BIND | SocksCmd::UDP_ASSOCIATE
240 )
241 }
242}
243
244impl SocksStatus {
245 #[cfg(feature = "proxy-handshake")]
247 pub(crate) fn into_socks4_status(self) -> u8 {
248 match self {
249 SocksStatus::SUCCEEDED => 0x5A,
250 _ => 0x5B,
251 }
252 }
253 #[cfg(feature = "client-handshake")]
255 pub(crate) fn from_socks4_status(status: u8) -> Self {
256 match status {
257 0x5A => SocksStatus::SUCCEEDED,
258 0x5B => SocksStatus::GENERAL_FAILURE,
259 0x5C | 0x5D => SocksStatus::NOT_ALLOWED,
260 _ => SocksStatus::GENERAL_FAILURE,
261 }
262 }
263}
264
265impl TryFrom<String> for SocksHostname {
266 type Error = Error;
267 fn try_from(s: String) -> Result<SocksHostname> {
268 if s.len() > 255 {
269 Err(bad_api_usage!("hostname too long").into())
272 } else if contains_zeros(s.as_bytes()) {
273 Err(Error::Syntax)
276 } else {
277 Ok(SocksHostname(s))
278 }
279 }
280}
281
282impl AsRef<str> for SocksHostname {
283 fn as_ref(&self) -> &str {
284 self.0.as_ref()
285 }
286}
287
288impl SocksAuth {
289 fn validate(&self, version: SocksVersion) -> Result<()> {
294 match self {
295 SocksAuth::NoAuth => {}
296 SocksAuth::Socks4(data) => {
297 if version != SocksVersion::V4 || contains_zeros(data) {
298 return Err(Error::Syntax);
299 }
300 }
301 SocksAuth::Username(user, pass) => {
302 if version != SocksVersion::V5
303 || user.len() > u8::MAX as usize
304 || pass.len() > u8::MAX as usize
305 {
306 return Err(Error::Syntax);
307 }
308 }
309 }
310 Ok(())
311 }
312}
313
314fn contains_zeros(b: &[u8]) -> bool {
318 use subtle::{Choice, ConstantTimeEq};
319 let c: Choice = b
320 .iter()
321 .fold(Choice::from(0), |seen_any, byte| seen_any | byte.ct_eq(&0));
322 c.unwrap_u8() != 0
323}
324
325impl SocksRequest {
326 pub fn new(
330 version: SocksVersion,
331 cmd: SocksCmd,
332 addr: SocksAddr,
333 port: u16,
334 auth: SocksAuth,
335 ) -> Result<Self> {
336 if !cmd.recognized() {
337 return Err(Error::NotImplemented(
338 format!("SOCKS command {}", cmd).into(),
339 ));
340 }
341 if port == 0 && cmd.requires_port() {
342 return Err(Error::Syntax);
343 }
344 auth.validate(version)?;
345
346 Ok(SocksRequest {
347 version,
348 cmd,
349 addr,
350 port,
351 auth,
352 })
353 }
354
355 pub fn version(&self) -> SocksVersion {
357 self.version
358 }
359
360 pub fn command(&self) -> SocksCmd {
362 self.cmd
363 }
364
365 pub fn auth(&self) -> &SocksAuth {
367 &self.auth
368 }
369
370 pub fn port(&self) -> u16 {
372 self.port
373 }
374
375 pub fn addr(&self) -> &SocksAddr {
377 &self.addr
378 }
379}
380
381impl fmt::Display for SocksAddr {
382 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
385 match self {
386 SocksAddr::Ip(a) => write!(f, "{}", a),
387 SocksAddr::Hostname(h) => write!(f, "{}", h.0),
388 }
389 }
390}
391
392#[derive(Debug, Clone)]
394pub struct SocksReply {
395 status: SocksStatus,
397 addr: SocksAddr,
399 port: u16,
401}
402
403impl SocksReply {
404 #[cfg(feature = "client-handshake")]
406 pub(crate) fn new(status: SocksStatus, addr: SocksAddr, port: u16) -> Self {
407 Self { status, addr, port }
408 }
409
410 pub fn status(&self) -> SocksStatus {
412 self.status
413 }
414
415 pub fn addr(&self) -> &SocksAddr {
423 &self.addr
424 }
425
426 pub fn port(&self) -> u16 {
431 self.port
432 }
433}
434
435#[cfg(test)]
436mod test {
437 #![allow(clippy::bool_assert_comparison)]
439 #![allow(clippy::clone_on_copy)]
440 #![allow(clippy::dbg_macro)]
441 #![allow(clippy::mixed_attributes_style)]
442 #![allow(clippy::print_stderr)]
443 #![allow(clippy::print_stdout)]
444 #![allow(clippy::single_char_pattern)]
445 #![allow(clippy::unwrap_used)]
446 #![allow(clippy::unchecked_time_subtraction)]
447 #![allow(clippy::useless_vec)]
448 #![allow(clippy::needless_pass_by_value)]
449 #![allow(clippy::string_slice)] use super::*;
452
453 #[test]
454 fn display_sa() {
455 let a = SocksAddr::Ip(IpAddr::V4("127.0.0.1".parse().unwrap()));
456 assert_eq!(a.to_string(), "127.0.0.1");
457
458 let a = SocksAddr::Ip(IpAddr::V6("f00::9999".parse().unwrap()));
459 assert_eq!(a.to_string(), "f00::9999");
460
461 let a = SocksAddr::Hostname("www.torproject.org".to_string().try_into().unwrap());
462 assert_eq!(a.to_string(), "www.torproject.org");
463 }
464
465 #[test]
466 fn ok_request() {
467 let localhost_v4 = SocksAddr::Ip(IpAddr::V4("127.0.0.1".parse().unwrap()));
468 let r = SocksRequest::new(
469 SocksVersion::V4,
470 SocksCmd::CONNECT,
471 localhost_v4.clone(),
472 1024,
473 SocksAuth::NoAuth,
474 )
475 .unwrap();
476 assert_eq!(r.version(), SocksVersion::V4);
477 assert_eq!(r.command(), SocksCmd::CONNECT);
478 assert_eq!(r.addr(), &localhost_v4);
479 assert_eq!(r.auth(), &SocksAuth::NoAuth);
480 }
481
482 #[test]
483 fn bad_request() {
484 let localhost_v4 = SocksAddr::Ip(IpAddr::V4("127.0.0.1".parse().unwrap()));
485
486 let e = SocksRequest::new(
487 SocksVersion::V4,
488 SocksCmd::BIND,
489 localhost_v4.clone(),
490 1024,
491 SocksAuth::NoAuth,
492 );
493 assert!(matches!(e, Err(Error::NotImplemented(_))));
494
495 let e = SocksRequest::new(
496 SocksVersion::V4,
497 SocksCmd::CONNECT,
498 localhost_v4,
499 0,
500 SocksAuth::NoAuth,
501 );
502 assert!(matches!(e, Err(Error::Syntax)));
503 }
504
505 #[test]
506 fn test_contains_zeros() {
507 assert!(contains_zeros(b"Hello\0world"));
508 assert!(!contains_zeros(b"Hello world"));
509 }
510}