Skip to repository content192 lines · 6.3 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T04:02:49.295Z 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
futures_io.rs
1//! A driver that runs a [`Handshake`] over [`futures`] I/O traits.
2
3use std::io;
4use std::pin::Pin;
5use std::task::{Context, Poll};
6
7use futures::io::{AsyncRead, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _};
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 [u8],
66 ) -> Poll<io::Result<usize>> {
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.len());
71 buf[..length].copy_from_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(length));
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_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
97 Pin::new(&mut self.get_mut().stream).poll_close(cx)
98 }
99}
100
101#[cfg(test)]
102mod tests {
103 use futures::executor::block_on;
104 use futures::join;
105
106 use super::*;
107
108 #[test]
109 fn establishes_http_tunnel_and_preserves_leftover() {
110 let (client, mut server) = duplex();
111 let spec = ProxySpec::parse(&"http://proxy:8080".parse().unwrap()).unwrap();
112 let target = Target::Domain("cloud.example.com".to_string(), 443);
113 block_on(async {
114 join!(
115 async {
116 let mut stream = establish(client, &spec, &target).await.unwrap();
117 let mut buffer = [0u8; 12];
118 stream.read_exact(&mut buffer).await.unwrap();
119 assert_eq!(&buffer, b"early-tunnel");
120 stream.write_all(b"hello").await.unwrap();
121 stream.flush().await.unwrap();
122 },
123 async {
124 let mut head = Vec::new();
125 let mut byte = [0u8; 1];
126 while !head.ends_with(b"\r\n\r\n") {
127 server.read_exact(&mut byte).await.unwrap();
128 head.push(byte[0]);
129 }
130 // Send the response and some tunnel bytes in one write so
131 // the driver has to hand them back through `leftover`.
132 server
133 .write_all(b"HTTP/1.1 200 OK\r\n\r\nearly-tunnel")
134 .await
135 .unwrap();
136 let mut buffer = [0u8; 5];
137 server.read_exact(&mut buffer).await.unwrap();
138 assert_eq!(&buffer, b"hello");
139 },
140 );
141 });
142 }
143
144 fn duplex() -> (PipeStream, PipeStream) {
145 let (client_reader, server_writer) = piper::pipe(1024);
146 let (server_reader, client_writer) = piper::pipe(1024);
147 (
148 PipeStream {
149 reader: client_reader,
150 writer: client_writer,
151 },
152 PipeStream {
153 reader: server_reader,
154 writer: server_writer,
155 },
156 )
157 }
158
159 struct PipeStream {
160 reader: piper::Reader,
161 writer: piper::Writer,
162 }
163
164 impl AsyncRead for PipeStream {
165 fn poll_read(
166 self: Pin<&mut Self>,
167 cx: &mut Context<'_>,
168 buf: &mut [u8],
169 ) -> Poll<io::Result<usize>> {
170 Pin::new(&mut self.get_mut().reader).poll_read(cx, buf)
171 }
172 }
173
174 impl AsyncWrite for PipeStream {
175 fn poll_write(
176 self: Pin<&mut Self>,
177 cx: &mut Context<'_>,
178 buf: &[u8],
179 ) -> Poll<io::Result<usize>> {
180 Pin::new(&mut self.get_mut().writer).poll_write(cx, buf)
181 }
182
183 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
184 Pin::new(&mut self.get_mut().writer).poll_flush(cx)
185 }
186
187 fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
188 Pin::new(&mut self.get_mut().writer).poll_close(cx)
189 }
190 }
191}
192