Store encoded packets in Endpoint's send queue

This commit is contained in:
Dirkjan Ochtman
2018-05-16 20:32:43 +02:00
parent f142ca761c
commit 44c8a4cfa6
4 changed files with 36 additions and 35 deletions
+1 -2
View File
@@ -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;
+28 -25
View File
@@ -23,7 +23,7 @@ pub struct Endpoint<T> {
secret: Secret,
prev_secret: Option<Secret>,
streams: Streams,
queue: VecDeque<Packet>,
queue: VecDeque<Vec<u8>>,
tls: T,
}
@@ -65,7 +65,7 @@ where
}
}
pub fn queued(&self) -> Option<&Packet> {
pub fn queued(&self) -> Option<&Vec<u8>> {
self.queue.front()
}
@@ -106,7 +106,7 @@ where
self.prev_secret = Some(old);
}
pub fn build_initial_packet(&mut self, mut payload: Vec<Frame>) {
pub fn build_initial_packet(&mut self, mut payload: Vec<Frame>) -> 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<Frame>) {
pub fn build_handshake_packet(&mut self, payload: Vec<Frame>) -> 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<Frame>) {
fn build_short_packet(&mut self, payload: Vec<Frame>) -> 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<Vec<u8>> {
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<tls::ClientSession> {
len: Some(handshake.len() as u64),
data: handshake,
}),
]);
Ok(())
])
}
}
+1 -2
View File
@@ -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;
}
+6 -6
View File
@@ -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());