Skip to repository content140 lines · 4.9 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T01:52:26.047Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
tokio.rs
1//! A driver that runs a [`Handshake`] over [`tokio`](::tokio) I/O traits.
2
3use std::io;
4use std::pin::Pin;
5use std::task::{Context, Poll};
6
7use ::tokio::io::{AsyncRead, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _, ReadBuf};
8
9use crate::{EstablishError, Handshake, ProxySpec, Step, Target};
10
11/// Runs the proxy handshake over `stream` and returns the tunneled stream.
12///
13/// The stream must already be connected to the proxy (and wrapped in TLS if
14/// [`ProxySpec::tls`] asks for it). See [`Handshake::new`] for how `target`
15/// interacts with DNS resolution.
16pub async fn establish<S>(
17 mut stream: S,
18 spec: &ProxySpec,
19 target: &Target,
20) -> Result<Tunneled<S>, EstablishError>
21where
22 S: AsyncRead + AsyncWrite + Unpin,
23{
24 let mut handshake = Handshake::new(spec, target)?;
25 let mut buffer = [0u8; 4096];
26 let mut received_length = 0;
27 loop {
28 let step = handshake.advance(&buffer[..received_length])?;
29 received_length = 0;
30 match step {
31 Step::Send(bytes) => {
32 stream.write_all(&bytes).await?;
33 stream.flush().await?;
34 }
35 Step::NeedMoreInput => {
36 received_length = stream.read(&mut buffer).await?;
37 if received_length == 0 {
38 return Err(EstablishError::Io(io::ErrorKind::UnexpectedEof.into()));
39 }
40 }
41 Step::Done { leftover } => {
42 return Ok(Tunneled {
43 leftover,
44 offset: 0,
45 stream,
46 });
47 }
48 }
49 }
50}
51
52/// A tunneled stream returned by [`establish`]: replays the handshake's
53/// leftover bytes before reading from the underlying transport. Writes pass
54/// straight through.
55pub struct Tunneled<S> {
56 leftover: Vec<u8>,
57 offset: usize,
58 stream: S,
59}
60
61impl<S: AsyncRead + Unpin> AsyncRead for Tunneled<S> {
62 fn poll_read(
63 self: Pin<&mut Self>,
64 cx: &mut Context<'_>,
65 buf: &mut ReadBuf<'_>,
66 ) -> Poll<io::Result<()>> {
67 let this = self.get_mut();
68 if this.offset < this.leftover.len() {
69 let remaining = &this.leftover[this.offset..];
70 let length = remaining.len().min(buf.remaining());
71 buf.put_slice(&remaining[..length]);
72 this.offset += length;
73 if this.offset == this.leftover.len() {
74 this.leftover = Vec::new();
75 this.offset = 0;
76 }
77 return Poll::Ready(Ok(()));
78 }
79 Pin::new(&mut this.stream).poll_read(cx, buf)
80 }
81}
82
83impl<S: AsyncWrite + Unpin> AsyncWrite for Tunneled<S> {
84 fn poll_write(
85 self: Pin<&mut Self>,
86 cx: &mut Context<'_>,
87 buf: &[u8],
88 ) -> Poll<io::Result<usize>> {
89 Pin::new(&mut self.get_mut().stream).poll_write(cx, buf)
90 }
91
92 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
93 Pin::new(&mut self.get_mut().stream).poll_flush(cx)
94 }
95
96 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
97 Pin::new(&mut self.get_mut().stream).poll_shutdown(cx)
98 }
99}
100
101#[cfg(test)]
102mod tests {
103 use super::*;
104
105 #[tokio::test]
106 async fn establishes_socks5_tunnel_and_preserves_leftover() {
107 let (client, mut server) = ::tokio::io::duplex(1024);
108 let spec = ProxySpec::parse(&"socks5h://proxy:1080".parse().unwrap()).unwrap();
109 let target = Target::Domain("cloud.example.com".to_string(), 443);
110 ::tokio::join!(
111 async {
112 let mut stream = establish(client, &spec, &target).await.unwrap();
113 let mut buffer = [0u8; 12];
114 stream.read_exact(&mut buffer).await.unwrap();
115 assert_eq!(&buffer, b"early-tunnel");
116 stream.write_all(b"hello").await.unwrap();
117 stream.flush().await.unwrap();
118 },
119 async {
120 let mut greeting = [0u8; 3];
121 server.read_exact(&mut greeting).await.unwrap();
122 assert_eq!(greeting, [0x05, 0x01, 0x00]);
123 server.write_all(&[0x05, 0x00]).await.unwrap();
124
125 let mut connect_request = vec![0u8; 4 + 1 + 17 + 2];
126 server.read_exact(&mut connect_request).await.unwrap();
127 // Send the reply and some tunnel bytes in one write so the
128 // driver has to hand them back through `leftover`.
129 let mut reply = vec![0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0];
130 reply.extend_from_slice(b"early-tunnel");
131 server.write_all(&reply).await.unwrap();
132
133 let mut buffer = [0u8; 5];
134 server.read_exact(&mut buffer).await.unwrap();
135 assert_eq!(&buffer, b"hello");
136 },
137 );
138 }
139}
140