diff --git a/src/client.rs b/src/client.rs index 47cf59416..53611dc00 100644 --- a/src/client.rs +++ b/src/client.rs @@ -68,8 +68,7 @@ impl Future for ClientFuture { loop { waiting = true; if let Some(ref mut client) = self.client { - if let Some(p) = client.endpoint.queued() { - let buf = client.endpoint.encode_packet(p)?; + if let Some(buf) = client.endpoint.queued() { let len = try_ready!(client.socket.poll_send(&buf)); debug_assert_eq!(len, buf.len()); waiting = false; diff --git a/src/endpoint.rs b/src/endpoint.rs index 739cf08c2..77b403643 100644 --- a/src/endpoint.rs +++ b/src/endpoint.rs @@ -23,7 +23,7 @@ pub struct Endpoint { secret: Secret, prev_secret: Option, streams: Streams, - queue: VecDeque, + queue: VecDeque>, tls: T, } @@ -65,7 +65,7 @@ where } } - pub fn queued(&self) -> Option<&Packet> { + pub fn queued(&self) -> Option<&Vec> { self.queue.front() } @@ -106,7 +106,7 @@ where self.prev_secret = Some(old); } - pub fn build_initial_packet(&mut self, mut payload: Vec) { + pub fn build_initial_packet(&mut self, mut payload: Vec) -> QuicResult<()> { let number = self.src_pn; self.src_pn += 1; @@ -116,61 +116,66 @@ where payload_len = 1200; } - debug_assert_eq!(self.local.cid.len, GENERATED_CID_LENGTH); - self.queue.push_back(Packet { + let (dst_cid, src_cid) = (self.remote.cid, self.local.cid); + debug_assert_eq!(src_cid.len, GENERATED_CID_LENGTH); + self.queue_packet(Packet { header: Header::Long { ptype: LongType::Initial, version: QUIC_VERSION, - dst_cid: self.remote.cid, - src_cid: self.local.cid, + dst_cid, + src_cid, len: payload_len as u64, number, }, payload, - }); + }) } - pub fn build_handshake_packet(&mut self, payload: Vec) { + pub fn build_handshake_packet(&mut self, payload: Vec) -> QuicResult<()> { let number = self.src_pn; self.src_pn += 1; - debug_assert_eq!(self.local.cid.len, GENERATED_CID_LENGTH); - self.queue.push_back(Packet { + let len = (payload.buf_len() + self.secret.tag_len()) as u64; + let (dst_cid, src_cid) = (self.remote.cid, self.local.cid); + debug_assert_eq!(src_cid.len, GENERATED_CID_LENGTH); + self.queue_packet(Packet { header: Header::Long { ptype: LongType::Handshake, version: QUIC_VERSION, - dst_cid: self.remote.cid, - src_cid: self.local.cid, - len: (payload.buf_len() + self.secret.tag_len()) as u64, + dst_cid, + src_cid, + len, number, }, payload, - }); + }) } - fn build_short_packet(&mut self, payload: Vec) { + fn build_short_packet(&mut self, payload: Vec) -> QuicResult<()> { let number = self.src_pn; self.src_pn += 1; + let dst_cid = self.remote.cid; debug_assert_eq!(self.state, State::Connected); debug_assert_eq!(self.local.cid.len, GENERATED_CID_LENGTH); - self.queue.push_back(Packet { + self.queue_packet(Packet { header: Header::Short { key_phase: false, ptype: ShortType::Four, - dst_cid: self.remote.cid, - number: number, + dst_cid, + number, }, payload, - }); + }) } - pub fn encode_packet(&self, packet: &Packet) -> QuicResult> { + pub fn queue_packet(&mut self, packet: Packet) -> QuicResult<()> { let key = self.encode_key(&packet.header); let len = packet.buf_len() + key.algorithm().tag_len(); let mut buf = vec![0u8; len]; packet.encode(&key, &mut buf)?; - Ok(buf) + self.queue.push_back(buf); + Ok(()) } pub(crate) fn handle(&mut self, buf: &mut [u8]) -> QuicResult<()> { @@ -296,7 +301,6 @@ where } else { self.build_handshake_packet(payload) } - Ok(()) } } @@ -321,8 +325,7 @@ impl Endpoint { len: Some(handshake.len() as u64), data: handshake, }), - ]); - Ok(()) + ]) } } diff --git a/src/server.rs b/src/server.rs index 26dab734a..5c9c21169 100644 --- a/src/server.rs +++ b/src/server.rs @@ -71,8 +71,7 @@ impl Future for Server { endpoint.handle_partial(partial)?; let mut sent = false; - if let Some(rsp) = endpoint.queued() { - let buf = endpoint.encode_packet(rsp)?; + if let Some(buf) = endpoint.queued() { try_ready!(self.socket.poll_send_to(&buf, &addr)); sent = true; } diff --git a/src/tests.rs b/src/tests.rs index 97532c27b..cee42b3f5 100644 --- a/src/tests.rs +++ b/src/tests.rs @@ -35,21 +35,21 @@ fn test_client_connect_resolves() { fn test_encoded_handshake() { let mut c = client_endpoint(); c.initial().unwrap(); - let mut c_initial = c.encode_packet(c.queued().unwrap()).unwrap(); + let mut c_initial = c.queued().unwrap().clone(); c.pop_queue(); let mut s = server_endpoint(Packet::start_decode(&mut c_initial).dst_cid()); s.handle(&mut c_initial).unwrap(); - let mut s_sh = s.encode_packet(s.queued().unwrap()).unwrap(); + let mut s_sh = s.queued().unwrap().clone(); s.pop_queue(); c.handle(&mut s_sh).unwrap(); - let mut c_fin = c.encode_packet(c.queued().unwrap()).unwrap(); + let mut c_fin = c.queued().unwrap().clone(); c.pop_queue(); s.handle(&mut c_fin).unwrap(); - let mut s_short = s.encode_packet(s.queued().unwrap()).unwrap(); + let mut s_short = s.queued().unwrap().clone(); s.pop_queue(); let c_short = { let partial = Packet::start_decode(&mut s_short); @@ -63,12 +63,12 @@ fn test_encoded_handshake() { fn test_handshake() { let mut c = client_endpoint(); c.initial().unwrap(); - let mut initial = c.encode_packet(&c.queued().unwrap()).unwrap(); + let mut initial = c.queued().unwrap().clone(); c.pop_queue(); let mut s = server_endpoint(Packet::start_decode(&mut initial).dst_cid()); s.handle(&mut initial).unwrap(); - let mut server_hello = s.encode_packet(s.queued().unwrap()).unwrap(); + let mut server_hello = s.queued().unwrap().clone(); c.handle(&mut server_hello).unwrap(); assert!(c.queued().is_some());