Skip to content

Commit 9d1d535

Browse files
committed
refactor(server): reuse gateway listener spec
Signed-off-by: Evan Lezar <elezar@nvidia.com>
1 parent 345ca15 commit 9d1d535

2 files changed

Lines changed: 50 additions & 58 deletions

File tree

crates/openshell-server/src/gateway_listener.rs

Lines changed: 38 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -20,29 +20,51 @@ pub struct CoveredGatewayAddress {
2020
}
2121

2222
#[derive(Clone, Debug, Eq, PartialEq)]
23-
struct GatewayListenerSpec {
24-
address: SocketAddr,
25-
scope: GatewayListenerScope,
23+
pub struct GatewayListenerSpec {
24+
pub address: SocketAddr,
25+
pub scope: GatewayListenerScope,
2626
covered_addresses: Vec<CoveredGatewayAddress>,
2727
}
2828

2929
/// A gateway listener together with the context needed to serve it.
3030
pub struct BoundGatewayListener {
3131
pub listener: TcpListener,
32-
pub address: SocketAddr,
33-
pub scope: GatewayListenerScope,
34-
pub covered_addresses: Vec<CoveredGatewayAddress>,
32+
pub spec: GatewayListenerSpec,
33+
}
34+
35+
impl GatewayListenerSpec {
36+
pub fn new(address: SocketAddr, scope: GatewayListenerScope) -> Self {
37+
Self {
38+
address,
39+
scope,
40+
covered_addresses: Vec::new(),
41+
}
42+
}
43+
44+
pub fn scope_for_local_addr(&self, local_addr: SocketAddr) -> GatewayListenerScope {
45+
self.covered_addresses
46+
.iter()
47+
.find(|covered| covered.address == local_addr)
48+
.map_or(self.scope, |covered| covered.scope)
49+
}
50+
51+
fn bind_to(mut self, local_addr: SocketAddr) -> Self {
52+
let requested_addr = self.address;
53+
self.address = local_addr;
54+
self.covered_addresses =
55+
resolve_bound_covered_addresses(&self.covered_addresses, requested_addr, local_addr);
56+
self
57+
}
3558
}
3659

3760
fn gateway_listener_specs(
3861
bind_address: SocketAddr,
3962
extra_addresses: &[SocketAddr],
4063
) -> Vec<GatewayListenerSpec> {
41-
let mut specs = vec![GatewayListenerSpec {
42-
address: bind_address,
43-
scope: GatewayListenerScope::Primary,
44-
covered_addresses: Vec::new(),
45-
}];
64+
let mut specs = vec![GatewayListenerSpec::new(
65+
bind_address,
66+
GatewayListenerScope::Primary,
67+
)];
4668
for address in extra_addresses {
4769
let scope = GatewayListenerScope::ComputeDriverCallback;
4870
if let Some(existing) = specs
@@ -62,11 +84,7 @@ fn gateway_listener_specs(
6284
});
6385
}
6486
} else {
65-
specs.push(GatewayListenerSpec {
66-
address: *address,
67-
scope,
68-
covered_addresses: Vec::new(),
69-
});
87+
specs.push(GatewayListenerSpec::new(*address, scope));
7088
}
7189
}
7290
specs
@@ -86,29 +104,12 @@ pub async fn bind_gateway_listeners(
86104
info!(address = %local_addr, "Server listening");
87105
listeners.push(BoundGatewayListener {
88106
listener,
89-
address: local_addr,
90-
scope: spec.scope,
91-
covered_addresses: resolve_bound_covered_addresses(
92-
&spec.covered_addresses,
93-
spec.address,
94-
local_addr,
95-
),
107+
spec: spec.bind_to(local_addr),
96108
});
97109
}
98110
Ok(listeners)
99111
}
100112

101-
pub fn gateway_listener_scope_for_local_addr(
102-
default_scope: GatewayListenerScope,
103-
covered_addresses: &[CoveredGatewayAddress],
104-
local_addr: SocketAddr,
105-
) -> GatewayListenerScope {
106-
covered_addresses
107-
.iter()
108-
.find(|covered| covered.address == local_addr)
109-
.map_or(default_scope, |covered| covered.scope)
110-
}
111-
112113
fn resolve_bound_covered_addresses(
113114
covered_addresses: &[CoveredGatewayAddress],
114115
requested_listener_addr: SocketAddr,
@@ -158,7 +159,7 @@ fn listener_covers(existing: SocketAddr, requested: SocketAddr) -> bool {
158159
mod tests {
159160
use super::{
160161
CoveredGatewayAddress, GatewayListenerScope, GatewayListenerSpec, bind_gateway_listeners,
161-
gateway_listener_scope_for_local_addr, gateway_listener_specs,
162+
gateway_listener_specs,
162163
};
163164
use std::net::SocketAddr;
164165
use std::sync::atomic::{AtomicBool, Ordering};
@@ -192,11 +193,11 @@ mod tests {
192193
.unwrap();
193194

194195
assert_eq!(
195-
gateway_listener_scope_for_local_addr(spec.scope, &spec.covered_addresses, docker),
196+
spec.scope_for_local_addr(docker),
196197
GatewayListenerScope::ComputeDriverCallback,
197198
);
198199
assert_eq!(
199-
gateway_listener_scope_for_local_addr(spec.scope, &spec.covered_addresses, loopback),
200+
spec.scope_for_local_addr(loopback),
200201
GatewayListenerScope::Primary,
201202
);
202203
}

crates/openshell-server/src/lib.rs

Lines changed: 12 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -72,10 +72,9 @@ use tracing::{debug, error, info, warn};
7272
pub(crate) static TEST_ENV_LOCK: LazyLock<Mutex<()>> = LazyLock::new(|| Mutex::new(()));
7373

7474
use compute::ComputeRuntime;
75-
use gateway_listener::{
76-
BoundGatewayListener, GatewayListenerScope, bind_gateway_listeners,
77-
gateway_listener_scope_for_local_addr,
78-
};
75+
#[cfg(test)]
76+
use gateway_listener::GatewayListenerSpec;
77+
use gateway_listener::{BoundGatewayListener, GatewayListenerScope, bind_gateway_listeners};
7978
pub use grpc::OpenShellService;
8079
pub use http::{health_router, http_router, metrics_router, service_http_router};
8180
pub use multiplex::{MultiplexService, MultiplexedService};
@@ -550,12 +549,8 @@ async fn serve_gateway_listener(
550549
enable_loopback_service_http: bool,
551550
mut shutdown: watch::Receiver<bool>,
552551
) {
553-
let BoundGatewayListener {
554-
listener,
555-
address: listen_addr,
556-
scope,
557-
covered_addresses,
558-
} = bound_listener;
552+
let BoundGatewayListener { listener, spec } = bound_listener;
553+
let listen_addr = spec.address;
559554

560555
loop {
561556
let accepted = tokio::select! {
@@ -576,12 +571,10 @@ async fn serve_gateway_listener(
576571
}
577572
};
578573
let listener_scope = match stream.local_addr() {
579-
Ok(local_addr) => {
580-
gateway_listener_scope_for_local_addr(scope, &covered_addresses, local_addr)
581-
}
574+
Ok(local_addr) => spec.scope_for_local_addr(local_addr),
582575
Err(e) => {
583576
debug!(error = %e, client = %addr, listen = %listen_addr, "Failed to inspect accepted local address");
584-
scope
577+
spec.scope
585578
}
586579
};
587580

@@ -1003,10 +996,10 @@ pub(crate) async fn ensure_default_workspace(store: &Store) -> Result<()> {
1003996
mod tests {
1004997
use super::{
1005998
BoundGatewayListener, ConfiguredComputeDriver, ConnectionProtocol, GatewayListenerScope,
1006-
MultiplexService, ServerState, TlsAcceptor, allow_plaintext_service_http,
1007-
bind_gateway_listeners, classify_initial_bytes, configured_compute_driver,
1008-
is_benign_tls_handshake_failure, kubernetes_sandbox_jwt_expiry_disabled,
1009-
serve_gateway_listener,
999+
GatewayListenerSpec, MultiplexService, ServerState, TlsAcceptor,
1000+
allow_plaintext_service_http, bind_gateway_listeners, classify_initial_bytes,
1001+
configured_compute_driver, is_benign_tls_handshake_failure,
1002+
kubernetes_sandbox_jwt_expiry_disabled, serve_gateway_listener,
10101003
};
10111004
use openshell_core::{
10121005
ComputeDriverKind, Config,
@@ -1102,9 +1095,7 @@ mod tests {
11021095
let handle = tokio::spawn(serve_gateway_listener(
11031096
BoundGatewayListener {
11041097
listener,
1105-
address: listen_addr,
1106-
scope: GatewayListenerScope::Primary,
1107-
covered_addresses: Vec::new(),
1098+
spec: GatewayListenerSpec::new(listen_addr, GatewayListenerScope::Primary),
11081099
},
11091100
service,
11101101
Some(tls_acceptor),

0 commit comments

Comments
 (0)