Skip to content

Commit 1882f98

Browse files
authored
Allow server accept calls to proceed concurrently (#88)
* Allow server accept calls to proceed concurrently The sprockets handshake includes both TLS handshake and bidirectional remote attestation. On real Oxide hardware this results in talking over IPCC to the SP and then over SPI to the RoT. This introduces delays in the protocol. And while each message that must reach the RoT may have to go over IPCC serially (pipelined?), there are other operations that occur interleaved with these operations, such as the messages sent between nodes over the bootstrap network. We therefore want to allow the accept operations as a whole to proceed concurrently, which will speed up cold boot and also match the behavior we allow on the client side. As in most async networking code, we wait for each individual connection by polling the TCP socket acceptor. Once that happens we return a future that can be polled concurrently or in parallel as desired by the caller. This matches the tokio-rustls pattern. A test was added to ensure that the pattern works as we expect with `spawn`, which we plan to use in Omicron. Some minor test deduplicaton was also done.
1 parent b3b15d9 commit 1882f98

3 files changed

Lines changed: 365 additions & 256 deletions

File tree

tls/examples/server.rs

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -88,12 +88,18 @@ async fn main() {
8888
resolve,
8989
};
9090

91-
let mut server = Server::new(server_config, listen_addr, log.clone())
91+
let server = Server::new(server_config, listen_addr, log.clone())
9292
.await
9393
.unwrap();
9494

9595
loop {
96-
let (stream, _) = server.accept(args.corpus.as_slice()).await.unwrap();
96+
let (stream, _) = server
97+
.accept(args.corpus.clone())
98+
.await
99+
.unwrap()
100+
.handshake()
101+
.await
102+
.unwrap();
97103
let platform_id = stream.peer_platform_id().as_str().unwrap();
98104
info!(log, "connected to attested peer: {platform_id}");
99105
let (mut reader, mut writer) = split(stream);

tls/src/lib.rs

Lines changed: 154 additions & 84 deletions
Original file line numberDiff line numberDiff line change
@@ -320,6 +320,8 @@ mod tests {
320320
use slog::Drain;
321321
use std::net::SocketAddrV6;
322322
use std::str::FromStr;
323+
use std::sync::atomic::AtomicUsize;
324+
use std::sync::Arc;
323325
use std::time::Duration;
324326
use tokio::io::{AsyncReadExt, AsyncWriteExt};
325327
use tokio::time::sleep;
@@ -331,6 +333,38 @@ mod tests {
331333
slog::Logger::root(drain, slog::o!("component" => "sprockets"))
332334
}
333335

336+
pub fn pki_keydir() -> Utf8PathBuf {
337+
let mut pki_keydir = Utf8PathBuf::from(env!("CARGO_MANIFEST_DIR"));
338+
pki_keydir.push("test-keys");
339+
pki_keydir
340+
}
341+
342+
fn local_config(n: usize) -> keys::SprocketsConfig {
343+
let pki_keydir = pki_keydir();
344+
345+
let attest_priv_key =
346+
pki_keydir.join(format!("test-alias-{n}.key.pem"));
347+
let attest_cert_chain =
348+
pki_keydir.join(format!("test-alias-{n}.certlist.pem"));
349+
let resolve_priv_key =
350+
pki_keydir.join(format!("test-sprockets-auth-{n}.key.pem"));
351+
let resolve_cert_chain =
352+
pki_keydir.join(format!("test-sprockets-auth-{n}.certlist.pem"));
353+
354+
keys::SprocketsConfig {
355+
attest: keys::AttestConfig::Local {
356+
priv_key: attest_priv_key,
357+
cert_chain: attest_cert_chain,
358+
log: pki_keydir.join("log.bin"),
359+
},
360+
roots: vec![pki_keydir.join("test-root-a.cert.pem")],
361+
resolve: keys::ResolveSetting::Local {
362+
priv_key: resolve_priv_key,
363+
cert_chain: resolve_cert_chain,
364+
},
365+
}
366+
}
367+
334368
#[tokio::test]
335369
async fn toml_config() {
336370
let ipcc = r#"
@@ -353,26 +387,9 @@ mod tests {
353387

354388
#[tokio::test]
355389
async fn no_corpus() {
356-
let mut pki_keydir = Utf8PathBuf::from(env!("CARGO_MANIFEST_DIR"));
357-
pki_keydir.push("test-keys");
358390
let log = logger();
359-
360391
let addr: SocketAddrV6 = SocketAddrV6::from_str("[::1]:46457").unwrap();
361392

362-
let server_config = keys::SprocketsConfig {
363-
attest: keys::AttestConfig::Local {
364-
priv_key: pki_keydir.join("test-alias-1.key.pem"),
365-
cert_chain: pki_keydir.join("test-alias-1.certlist.pem"),
366-
log: pki_keydir.join("log.bin"),
367-
},
368-
roots: vec![pki_keydir.join("test-root-a.cert.pem")],
369-
resolve: keys::ResolveSetting::Local {
370-
priv_key: pki_keydir.join("test-sprockets-auth-1.key.pem"),
371-
cert_chain: pki_keydir
372-
.join("test-sprockets-auth-1.certlist.pem"),
373-
},
374-
};
375-
376393
// Message to send over TLS
377394
const MSG: &str = "Hello Joe";
378395

@@ -383,12 +400,18 @@ mod tests {
383400
];
384401

385402
tokio::spawn(async move {
386-
let mut server = Server::new(server_config, addr, log2.clone())
403+
let server_config = local_config(1);
404+
let server = Server::new(server_config, addr, log2.clone())
387405
.await
388406
.unwrap();
389407

390-
let (mut stream, _) =
391-
server.accept(corpus.as_slice()).await.unwrap();
408+
let (mut stream, _) = server
409+
.accept(corpus.clone())
410+
.await
411+
.unwrap()
412+
.handshake()
413+
.await
414+
.unwrap();
392415
let mut buf = String::new();
393416
stream.read_to_string(&mut buf).await.unwrap();
394417

@@ -400,19 +423,7 @@ mod tests {
400423

401424
// Loop until we succesfully connect
402425
let mut stream = loop {
403-
let client_config = keys::SprocketsConfig {
404-
attest: keys::AttestConfig::Local {
405-
priv_key: pki_keydir.join("test-alias-2.key.pem"),
406-
cert_chain: pki_keydir.join("test-alias-2.certlist.pem"),
407-
log: pki_keydir.join("log.bin"),
408-
},
409-
roots: vec![pki_keydir.join("test-root-a.cert.pem")],
410-
resolve: keys::ResolveSetting::Local {
411-
priv_key: pki_keydir.join("test-sprockets-auth-2.key.pem"),
412-
cert_chain: pki_keydir
413-
.join("test-sprockets-auth-2.certlist.pem"),
414-
},
415-
};
426+
let client_config = local_config(2);
416427

417428
let corpus = vec![
418429
// We don't use a corpus
@@ -439,25 +450,10 @@ mod tests {
439450

440451
#[tokio::test]
441452
async fn basic() {
442-
let mut pki_keydir = Utf8PathBuf::from(env!("CARGO_MANIFEST_DIR"));
443-
pki_keydir.push("test-keys");
444453
let log = logger();
445-
454+
let pki_keydir = pki_keydir();
446455
let addr: SocketAddrV6 = SocketAddrV6::from_str("[::1]:46456").unwrap();
447-
448-
let server_config = keys::SprocketsConfig {
449-
attest: keys::AttestConfig::Local {
450-
priv_key: pki_keydir.join("test-alias-1.key.pem"),
451-
cert_chain: pki_keydir.join("test-alias-1.certlist.pem"),
452-
log: pki_keydir.join("log.bin"),
453-
},
454-
roots: vec![pki_keydir.join("test-root-a.cert.pem")],
455-
resolve: keys::ResolveSetting::Local {
456-
priv_key: pki_keydir.join("test-sprockets-auth-1.key.pem"),
457-
cert_chain: pki_keydir
458-
.join("test-sprockets-auth-1.certlist.pem"),
459-
},
460-
};
456+
let server_config = local_config(1);
461457

462458
// Message to send over TLS
463459
const MSG: &str = "Hello Joe";
@@ -470,12 +466,17 @@ mod tests {
470466
];
471467

472468
tokio::spawn(async move {
473-
let mut server = Server::new(server_config, addr, log2.clone())
469+
let server = Server::new(server_config, addr, log2.clone())
474470
.await
475471
.unwrap();
476472

477-
let (mut stream, _) =
478-
server.accept(corpus.as_slice()).await.unwrap();
473+
let (mut stream, _) = server
474+
.accept(corpus.clone())
475+
.await
476+
.unwrap()
477+
.handshake()
478+
.await
479+
.unwrap();
479480
let mut buf = String::new();
480481
stream.read_to_string(&mut buf).await.unwrap();
481482

@@ -487,19 +488,7 @@ mod tests {
487488

488489
// Loop until we succesfully connect
489490
let mut stream = loop {
490-
let client_config = keys::SprocketsConfig {
491-
attest: keys::AttestConfig::Local {
492-
priv_key: pki_keydir.join("test-alias-2.key.pem"),
493-
cert_chain: pki_keydir.join("test-alias-2.certlist.pem"),
494-
log: pki_keydir.join("log.bin"),
495-
},
496-
roots: vec![pki_keydir.join("test-root-a.cert.pem")],
497-
resolve: keys::ResolveSetting::Local {
498-
priv_key: pki_keydir.join("test-sprockets-auth-2.key.pem"),
499-
cert_chain: pki_keydir
500-
.join("test-sprockets-auth-2.certlist.pem"),
501-
},
502-
};
491+
let client_config = local_config(2);
503492

504493
let corpus = vec![
505494
pki_keydir.join("corim-rot.cbor"),
@@ -527,25 +516,11 @@ mod tests {
527516

528517
#[tokio::test]
529518
async fn unattested_client() {
530-
let mut pki_keydir = Utf8PathBuf::from(env!("CARGO_MANIFEST_DIR"));
531-
pki_keydir.push("test-keys");
532519
let log = logger();
533-
520+
let pki_keydir = pki_keydir();
534521
let addr: SocketAddrV6 = SocketAddrV6::from_str("[::1]:46459").unwrap();
535522

536-
let server_config = keys::SprocketsConfig {
537-
attest: keys::AttestConfig::Local {
538-
priv_key: pki_keydir.join("test-alias-1.key.pem"),
539-
cert_chain: pki_keydir.join("test-alias-1.certlist.pem"),
540-
log: pki_keydir.join("log.bin"),
541-
},
542-
roots: vec![pki_keydir.join("test-root-a.cert.pem")],
543-
resolve: keys::ResolveSetting::Local {
544-
priv_key: pki_keydir.join("test-sprockets-auth-1.key.pem"),
545-
cert_chain: pki_keydir
546-
.join("test-sprockets-auth-1.certlist.pem"),
547-
},
548-
};
523+
let server_config = local_config(1);
549524

550525
// Message to send over TLS
551526
const MSG: &str = "Hello Joe";
@@ -558,12 +533,18 @@ mod tests {
558533
];
559534

560535
let handle = tokio::spawn(async move {
561-
let mut server = Server::new(server_config, addr, log2.clone())
536+
let server = Server::new(server_config, addr, log2.clone())
562537
.await
563538
.unwrap();
564539

565540
// We never expect this to succeed
566-
let _ = match server.accept(corpus.as_slice()).await {
541+
let _ = match server
542+
.accept(corpus.clone())
543+
.await
544+
.unwrap()
545+
.handshake()
546+
.await
547+
{
567548
Ok(_) => panic!("This should not succed"),
568549
Err(_) => done_tx.send(()),
569550
};
@@ -604,4 +585,93 @@ mod tests {
604585

605586
handle.await.unwrap();
606587
}
588+
589+
#[tokio::test]
590+
async fn spawn_accept() {
591+
let log = logger();
592+
let pki_keydir = pki_keydir();
593+
594+
let addr: SocketAddrV6 = SocketAddrV6::from_str("[::1]:46466").unwrap();
595+
596+
let server_config = local_config(1);
597+
598+
// Message to send over TLS
599+
const MSG: &str = "Hello Joe";
600+
601+
let log2 = log.clone();
602+
let corpus = vec![
603+
pki_keydir.join("corim-rot.cbor"),
604+
pki_keydir.join("corim-sp.cbor"),
605+
];
606+
607+
// Accept connections from `max_connections` clients in different tasks
608+
//
609+
// For this test, the clients all share a set of keys, because
610+
// we only generate 2 sets of keys from the KDL.
611+
let max_connections = 3;
612+
let done_count = Arc::new(AtomicUsize::new(0));
613+
let dc2 = done_count.clone();
614+
tokio::spawn(async move {
615+
let server = Server::new(server_config, addr, log2.clone())
616+
.await
617+
.unwrap();
618+
619+
for _ in 0..max_connections {
620+
let acceptor = server.accept(corpus.clone()).await.unwrap();
621+
let done_count = dc2.clone();
622+
tokio::spawn(async move {
623+
let (mut stream, _) = acceptor.handshake().await.unwrap();
624+
let mut buf = String::new();
625+
stream.read_to_string(&mut buf).await.unwrap();
626+
627+
assert_eq!(buf.as_str(), MSG);
628+
done_count
629+
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
630+
});
631+
}
632+
});
633+
634+
// Spawn `max_connections` tasks to concurrently connect
635+
for _ in 0..max_connections {
636+
let pki_keydir = pki_keydir.clone();
637+
let log = log.clone();
638+
tokio::spawn(async move {
639+
// Loop until we succesfully connect
640+
let mut stream = loop {
641+
let client_config = local_config(2);
642+
643+
let corpus = vec![
644+
pki_keydir.join("corim-rot.cbor"),
645+
pki_keydir.join("corim-sp.cbor"),
646+
];
647+
648+
if let Ok(stream) = Client::connect(
649+
client_config,
650+
addr,
651+
corpus,
652+
log.clone(),
653+
)
654+
.await
655+
{
656+
break stream;
657+
}
658+
sleep(Duration::from_millis(1)).await;
659+
};
660+
661+
stream.write_all(MSG.as_bytes()).await.unwrap();
662+
663+
// Trigger an EOF so that `read_to_string` in the acceptor task
664+
// completes.
665+
stream.shutdown().await.unwrap();
666+
});
667+
}
668+
669+
// Wait each spawned server task to receive and assert the message from
670+
// a single client.
671+
while done_count.load(std::sync::atomic::Ordering::Relaxed)
672+
!= max_connections
673+
{
674+
sleep(Duration::from_millis(1)).await;
675+
}
676+
}
607677
}

0 commit comments

Comments
 (0)