Skip to main content

eggress_testkit/
pproxy_oracle.rs

1use std::net::SocketAddr;
2use std::process::{Child, Command, Stdio};
3use std::sync::{Arc, Mutex};
4use std::time::{Duration, Instant};
5
6use thiserror::Error;
7use tokio::net::TcpStream;
8
9#[derive(Debug, Error)]
10pub enum PproxyOracleError {
11    #[error("io error: {0}")]
12    Io(#[from] std::io::Error),
13
14    #[error("process failed: {0}")]
15    ProcessFailed(String),
16
17    #[error("startup timeout after {0:?}")]
18    StartupTimeout(Duration),
19
20    #[error("not ready: {0}")]
21    NotReady(String),
22
23    #[error("version mismatch: expected {expected}, got {actual}")]
24    VersionMismatch { expected: String, actual: String },
25
26    #[error("version detection failed: {0}")]
27    VersionDetectionFailed(String),
28}
29
30#[derive(Debug, Clone)]
31pub struct OracleConfig {
32    pub python_binary: String,
33    pub pproxy_version: String,
34    pub startup_timeout: Duration,
35    pub shutdown_timeout: Duration,
36    pub io_timeout: Duration,
37}
38
39impl Default for OracleConfig {
40    fn default() -> Self {
41        Self {
42            python_binary: "python3".to_string(),
43            pproxy_version: "2.7.9".to_string(),
44            startup_timeout: Duration::from_secs(15),
45            shutdown_timeout: Duration::from_secs(5),
46            io_timeout: Duration::from_secs(3),
47        }
48    }
49}
50
51pub struct PproxyProcess {
52    child: Option<Child>,
53    bound_addr: SocketAddr,
54    stdout_buf: Arc<Mutex<Vec<u8>>>,
55    stderr_buf: Arc<Mutex<Vec<u8>>>,
56    #[allow(dead_code)]
57    work_dir: Option<tempfile::TempDir>,
58}
59
60impl PproxyProcess {
61    pub async fn start(config: &OracleConfig, args: &[String]) -> Result<Self, PproxyOracleError> {
62        let work_dir = tempfile::TempDir::new().map_err(PproxyOracleError::Io)?;
63
64        let stdout_buf = Arc::new(Mutex::new(Vec::new()));
65        let stderr_buf = Arc::new(Mutex::new(Vec::new()));
66
67        let stdout_clone = Arc::clone(&stdout_buf);
68        let stderr_clone = Arc::clone(&stderr_buf);
69
70        let mut child = Command::new(&config.python_binary)
71            .arg("-m")
72            .arg("pproxy")
73            .args(args)
74            .current_dir(work_dir.path())
75            .stdout(Stdio::piped())
76            .stderr(Stdio::piped())
77            .spawn()
78            .map_err(PproxyOracleError::Io)?;
79
80        let child_stdout = child.stdout.take().expect("stdout piped");
81        let child_stderr = child.stderr.take().expect("stderr piped");
82
83        std::thread::spawn(move || {
84            let mut reader = child_stderr;
85            let mut tmp = [0u8; 4096];
86            loop {
87                match std::io::Read::read(&mut reader, &mut tmp) {
88                    Ok(0) => break,
89                    Ok(n) => {
90                        if let Ok(mut guard) = stderr_clone.lock() {
91                            guard.extend_from_slice(&tmp[..n]);
92                        }
93                    }
94                    Err(_) => break,
95                }
96            }
97        });
98
99        std::thread::spawn(move || {
100            let mut reader = child_stdout;
101            let mut tmp = [0u8; 4096];
102            loop {
103                match std::io::Read::read(&mut reader, &mut tmp) {
104                    Ok(0) => break,
105                    Ok(n) => {
106                        if let Ok(mut guard) = stdout_clone.lock() {
107                            guard.extend_from_slice(&tmp[..n]);
108                        }
109                    }
110                    Err(_) => break,
111                }
112            }
113        });
114
115        let bound_addr = if let Some(addr) = parse_listen_addr_from_args(args) {
116            addr
117        } else {
118            wait_for_output_ready(&stderr_buf, &stdout_buf, config.startup_timeout).await?
119        };
120
121        let proc = Self {
122            child: Some(child),
123            bound_addr,
124            stdout_buf,
125            stderr_buf,
126            work_dir: Some(work_dir),
127        };
128
129        proc.wait_ready(config).await?;
130
131        Ok(proc)
132    }
133
134    pub async fn wait_ready(&self, config: &OracleConfig) -> Result<(), PproxyOracleError> {
135        let start = Instant::now();
136        let interval = Duration::from_millis(100);
137
138        loop {
139            match TcpStream::connect(self.bound_addr).await {
140                Ok(_) => return Ok(()),
141                Err(_) if start.elapsed() < config.startup_timeout => {
142                    tokio::time::sleep(interval).await;
143                }
144                Err(e) => {
145                    return Err(PproxyOracleError::NotReady(format!(
146                        "tcp connect to {} failed: {}",
147                        self.bound_addr, e
148                    )));
149                }
150            }
151        }
152    }
153
154    pub fn shutdown(&mut self) {
155        if let Some(ref mut child) = self.child {
156            let _ = child.kill();
157            let _ = child.wait();
158        }
159        self.child = None;
160    }
161
162    pub fn stdout(&self) -> Vec<u8> {
163        self.stdout_buf
164            .lock()
165            .map(|g| g.clone())
166            .unwrap_or_default()
167    }
168
169    pub fn stderr(&self) -> Vec<u8> {
170        self.stderr_buf
171            .lock()
172            .map(|g| g.clone())
173            .unwrap_or_default()
174    }
175
176    pub fn bound_addr(&self) -> SocketAddr {
177        self.bound_addr
178    }
179
180    pub fn redacted_stderr(&self) -> String {
181        let raw = self.stderr();
182        redact_credentials(&raw)
183    }
184}
185
186impl Drop for PproxyProcess {
187    fn drop(&mut self) {
188        self.shutdown();
189    }
190}
191
192pub fn redact_credentials(data: &[u8]) -> String {
193    let text = String::from_utf8_lossy(data);
194    redact_uri_credentials(&text)
195}
196
197fn redact_uri_credentials(text: &str) -> String {
198    let mut result = text.to_string();
199
200    let patterns = [
201        "socks4://",
202        "socks4a://",
203        "socks5://",
204        "http://",
205        "https://",
206        "ss://",
207        "trojan://",
208    ];
209
210    for scheme in &patterns {
211        let mut offset = 0;
212        while let Some(scheme_pos) = result[offset..].find(scheme) {
213            let abs_pos = offset + scheme_pos;
214            let after_scheme = abs_pos + scheme.len();
215            if let Some(at_rel) = result[after_scheme..].find('@') {
216                let cred_start = after_scheme;
217                let cred_end = after_scheme + at_rel;
218                let colon_pos = result[cred_start..cred_end].find(':');
219                if let Some(colon_pos) = colon_pos {
220                    let user = result[cred_start..cred_start + colon_pos].to_string();
221                    let rest = result[cred_end + 1..].to_string();
222                    let redacted = format!("{}:***@{}", user, rest);
223                    result = format!("{}{}", &result[..cred_start], redacted);
224                    offset = cred_start + user.len() + 5;
225                } else {
226                    offset = cred_end;
227                }
228            } else {
229                break;
230            }
231        }
232    }
233
234    result
235}
236
237async fn wait_for_output_ready(
238    stderr_buf: &Arc<Mutex<Vec<u8>>>,
239    stdout_buf: &Arc<Mutex<Vec<u8>>>,
240    timeout: Duration,
241) -> Result<SocketAddr, PproxyOracleError> {
242    let start = Instant::now();
243    let interval = Duration::from_millis(100);
244
245    loop {
246        {
247            let guard = stderr_buf.lock().map_err(|e| {
248                PproxyOracleError::ProcessFailed(format!("stderr lock poisoned: {}", e))
249            })?;
250            let text = String::from_utf8_lossy(&guard);
251            if let Some(addr) = parse_bound_addr(&text) {
252                return Ok(addr);
253            }
254        }
255        {
256            let guard = stdout_buf.lock().map_err(|e| {
257                PproxyOracleError::ProcessFailed(format!("stdout lock poisoned: {}", e))
258            })?;
259            let text = String::from_utf8_lossy(&guard);
260            if let Some(addr) = parse_bound_addr(&text) {
261                return Ok(addr);
262            }
263        }
264
265        if start.elapsed() >= timeout {
266            let stderr_text = stderr_buf
267                .lock()
268                .map(|g| String::from_utf8_lossy(&g).trim().to_string())
269                .unwrap_or_default();
270            let stdout_text = stdout_buf
271                .lock()
272                .map(|g| String::from_utf8_lossy(&g).trim().to_string())
273                .unwrap_or_default();
274            return Err(PproxyOracleError::ProcessFailed(format!(
275                "startup timeout after {:?}, stderr: [{}], stdout: [{}]",
276                timeout, stderr_text, stdout_text
277            )));
278        }
279
280        tokio::time::sleep(interval).await;
281    }
282}
283
284fn parse_listen_addr_from_args(args: &[String]) -> Option<SocketAddr> {
285    let mut i = 0;
286    while i < args.len() {
287        if (args[i] == "-l" || args[i] == "--listen") && i + 1 < args.len() {
288            let uri = &args[i + 1];
289            let host_port = if let Some(at_pos) = uri.rfind('@') {
290                &uri[at_pos + 1..]
291            } else {
292                uri.as_str()
293            };
294            let stripped = host_port
295                .strip_prefix("socks5://")
296                .or_else(|| host_port.strip_prefix("http://"))
297                .or_else(|| host_port.strip_prefix("socks4://"))
298                .or_else(|| host_port.strip_prefix("socks4a://"))
299                .or_else(|| host_port.strip_prefix("ss://"))
300                .or_else(|| host_port.strip_prefix("trojan://"))
301                .or_else(|| host_port.strip_prefix("direct://"))
302                .unwrap_or(host_port);
303            if let Some(colon_pos) = stripped.rfind(':') {
304                let port_str = &stripped[colon_pos + 1..];
305                if let Ok(port) = port_str.parse::<u16>() {
306                    if port > 0 {
307                        let host_str = &stripped[..colon_pos];
308                        let host = if host_str.is_empty() || host_str == "127.0.0.1" {
309                            std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
310                        } else if host_str == "0.0.0.0" {
311                            std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)
312                        } else {
313                            host_str.parse().ok()?
314                        };
315                        return Some(SocketAddr::new(host, port));
316                    }
317                }
318            }
319            return None;
320        }
321        i += 1;
322    }
323    None
324}
325
326fn parse_bound_addr(text: &str) -> Option<SocketAddr> {
327    for line in text.lines() {
328        let line = line.trim();
329        if line.is_empty() {
330            continue;
331        }
332
333        if let Some(idx) = line.find("127.0.0.1:") {
334            let port_str = &line[idx + "127.0.0.1:".len()..];
335            let port_str = port_str
336                .chars()
337                .take_while(|c| c.is_ascii_digit())
338                .collect::<String>();
339            if let Ok(port) = port_str.parse::<u16>() {
340                if port > 0 {
341                    return Some(SocketAddr::new(
342                        std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST),
343                        port,
344                    ));
345                }
346            }
347        }
348
349        if let Some(idx) = line.find("0.0.0.0:") {
350            let port_str = &line[idx + "0.0.0.0:".len()..];
351            let port_str = port_str
352                .chars()
353                .take_while(|c| c.is_ascii_digit())
354                .collect::<String>();
355            if let Ok(port) = port_str.parse::<u16>() {
356                if port > 0 {
357                    return Some(SocketAddr::new(
358                        std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED),
359                        port,
360                    ));
361                }
362            }
363        }
364    }
365
366    None
367}
368
369pub async fn verify_pproxy_version(config: &OracleConfig) -> Result<String, PproxyOracleError> {
370    let output = Command::new(&config.python_binary)
371        .args([
372            "-c",
373            "import pproxy; print(getattr(pproxy, '__version__', 'unknown'))",
374        ])
375        .stdout(Stdio::piped())
376        .stderr(Stdio::piped())
377        .output()
378        .map_err(PproxyOracleError::Io)?;
379
380    if !output.status.success() {
381        let stderr = String::from_utf8_lossy(&output.stderr);
382        return Err(PproxyOracleError::VersionDetectionFailed(
383            stderr.trim().to_string(),
384        ));
385    }
386
387    let version = String::from_utf8_lossy(&output.stdout).trim().to_string();
388
389    Ok(version)
390}
391
392pub async fn assert_pproxy_version(config: &OracleConfig) -> Result<(), PproxyOracleError> {
393    let actual = verify_pproxy_version(config).await?;
394    if actual != config.pproxy_version && actual != "unknown" {
395        return Err(PproxyOracleError::VersionMismatch {
396            expected: config.pproxy_version.clone(),
397            actual,
398        });
399    }
400    Ok(())
401}
402
403#[cfg(test)]
404mod tests {
405    use super::*;
406
407    fn require_pproxy() -> bool {
408        std::process::Command::new("python3")
409            .args(["-c", "import pproxy"])
410            .stdout(Stdio::null())
411            .stderr(Stdio::null())
412            .status()
413            .map(|s| s.success())
414            .unwrap_or(false)
415    }
416
417    #[tokio::test]
418    #[ignore]
419    async fn test_process_guard_drop_kills_child() {
420        if !require_pproxy() {
421            eprintln!("pproxy not available, skipping");
422            return;
423        }
424
425        let config = OracleConfig::default();
426        let port = crate::get_free_port().await;
427        let listen = format!("socks5://127.0.0.1:{}", port);
428        let args = vec![
429            "-l".to_string(),
430            listen,
431            "-r".to_string(),
432            "direct".to_string(),
433        ];
434
435        let proc = PproxyProcess::start(&config, &args).await.unwrap();
436        let addr = proc.bound_addr();
437        assert_eq!(addr.port(), port, "should parse port from args");
438
439        drop(proc);
440
441        tokio::time::sleep(Duration::from_millis(200)).await;
442
443        let result = TcpStream::connect(addr).await;
444        assert!(result.is_err(), "process should be dead after drop");
445    }
446
447    #[tokio::test]
448    #[ignore]
449    async fn test_readiness_probe() {
450        if !require_pproxy() {
451            eprintln!("pproxy not available, skipping");
452            return;
453        }
454
455        let config = OracleConfig::default();
456        let port = crate::get_free_port().await;
457        let listen = format!("socks5://127.0.0.1:{}", port);
458        let args = vec![
459            "-l".to_string(),
460            listen,
461            "-r".to_string(),
462            "direct".to_string(),
463        ];
464
465        let proc = PproxyProcess::start(&config, &args).await.unwrap();
466
467        let result = TcpStream::connect(proc.bound_addr()).await;
468        assert!(result.is_ok(), "process should be ready after start");
469
470        drop(proc);
471    }
472
473    #[test]
474    fn test_log_redaction() {
475        let input = b"socks5://user:secret123@127.0.0.1:1080\nhttp://admin:pw0rd@0.0.0.0:8080\n";
476        let redacted = redact_credentials(input);
477
478        assert!(!redacted.contains("secret123"));
479        assert!(!redacted.contains("pw0rd"));
480        assert!(redacted.contains("user:***@127.0.0.1:1080"));
481        assert!(redacted.contains("admin:***@0.0.0.0:8080"));
482    }
483
484    #[test]
485    fn test_log_redaction_no_credentials() {
486        let input = b"socks5://127.0.0.1:1080\nlistening on port 8080\n";
487        let redacted = redact_credentials(input);
488
489        assert_eq!(redacted, String::from_utf8_lossy(input));
490    }
491
492    #[test]
493    fn test_parse_bound_addr() {
494        assert_eq!(
495            parse_bound_addr("Listen: socks5://127.0.0.1:9090"),
496            Some(SocketAddr::new(
497                std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST),
498                9090
499            ))
500        );
501
502        assert_eq!(
503            parse_bound_addr("Listen: http://0.0.0.0:8080"),
504            Some(SocketAddr::new(
505                std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED),
506                8080
507            ))
508        );
509
510        assert_eq!(parse_bound_addr("no address here"), None);
511        assert_eq!(parse_bound_addr(""), None);
512    }
513
514    #[test]
515    fn test_redact_uri_credentials() {
516        assert_eq!(
517            redact_uri_credentials("socks5://user:pass@host:1080"),
518            "socks5://user:***@host:1080"
519        );
520
521        assert_eq!(
522            redact_uri_credentials("http://admin:secret@proxy:8080"),
523            "http://admin:***@proxy:8080"
524        );
525
526        assert_eq!(
527            redact_uri_credentials("socks5://127.0.0.1:1080"),
528            "socks5://127.0.0.1:1080"
529        );
530    }
531
532    #[tokio::test]
533    async fn test_version_detection() {
534        if !require_pproxy() {
535            eprintln!("pproxy not available, skipping");
536            return;
537        }
538
539        let config = OracleConfig::default();
540        let result = verify_pproxy_version(&config).await;
541        assert!(
542            result.is_ok(),
543            "version detection should succeed: {:?}",
544            result.err()
545        );
546    }
547}