Drive SDK maintenance while receiving
ci / test (push) Canceled after 0s
ci / fuzz-smoke (push) Canceled after 0s
ci / windows-client (push) Canceled after 0s
ci / package-release (linux-x86_64, ubuntu-latest) (push) Canceled after 0s
ci / package-release (macos-aarch64, macos-14) (push) Canceled after 0s
ci / package-release (macos-x86_64, macos-13) (push) Canceled after 0s
ci / package-release (windows-x86_64, windows-latest) (push) Canceled after 0s
ci / remote-bench (push) Canceled after 0s
ci / publish-gitea-release (push) Canceled after 0s

This commit is contained in:
DuProcess
2026-07-13 00:58:54 -04:00
parent 8b33a5d20b
commit 0125966d2a
3 changed files with 250 additions and 74 deletions
+124 -26
View File
@@ -10,8 +10,8 @@ use crate::protocol::{
NativeServerHelloBody, NativeUserAuthBody, PacketKind, SERVER_TO_CLIENT,
};
use crate::transport::{
DoshTransport, SessionEvent, SessionRole, SessionTransportConfig, TransportConfig,
service_name_from_target,
ADAPTIVE_RETRANSMIT_MIN, DoshTransport, SessionEvent, SessionRole, SessionTransportConfig,
TransportConfig, service_name_from_target,
};
use crate::udp::{is_transient_udp_error, is_transient_udp_send_error};
use anyhow::{Context, Result, anyhow, bail};
@@ -169,35 +169,54 @@ impl DoshServer {
let mut buf = vec![0u8; 65535];
loop {
self.expire_pending();
let (n, peer) = match self.socket.recv_from(&mut buf).await {
Ok(value) => value,
Err(err) if is_transient_udp_error(&err) => continue,
Err(err) => return Err(err.into()),
let (n, peer) = match tokio::time::timeout(
ADAPTIVE_RETRANSMIT_MIN,
self.socket.recv_from(&mut buf),
)
.await
{
Ok(Ok(value)) => value,
Ok(Err(err)) if is_transient_udp_error(&err) => continue,
Ok(Err(err)) => return Err(err.into()),
Err(_) => {
self.maintenance().await?;
continue;
}
};
let packet = match protocol::decode(&buf[..n]) {
Ok(packet) => packet,
Err(_) => continue,
};
if packet.header.kind == PacketKind::NativeClientHello {
self.handle_client_hello(peer, packet.body).await?;
return Ok(DoshServerEvent::Ignored);
if let Some(event) = self.handle_datagram(peer, &buf[..n]).await? {
return Ok(event);
}
if packet.header.kind == PacketKind::NativeUserAuth {
return self.handle_user_auth(peer, &packet).await;
}
if let Some(transport) = self.transports.get_mut(&packet.header.conn_id) {
let event = transport.handle_datagram(&buf[..n], peer).await?;
return Ok(DoshServerEvent::Session {
conn_id: packet.header.conn_id,
event,
});
}
self.send_reject(peer, packet.header.conn_id, "unknown Dosh connection")
.await?;
return Ok(DoshServerEvent::Ignored);
}
}
async fn handle_datagram(
&mut self,
peer: SocketAddr,
datagram: &[u8],
) -> Result<Option<DoshServerEvent>> {
let packet = match protocol::decode(datagram) {
Ok(packet) => packet,
Err(_) => return Ok(None),
};
if packet.header.kind == PacketKind::NativeClientHello {
self.handle_client_hello(peer, packet.body).await?;
return Ok(Some(DoshServerEvent::Ignored));
}
if packet.header.kind == PacketKind::NativeUserAuth {
return self.handle_user_auth(peer, &packet).await.map(Some);
}
if let Some(transport) = self.transports.get_mut(&packet.header.conn_id) {
let event = transport.handle_datagram(datagram, peer).await?;
return Ok(Some(DoshServerEvent::Session {
conn_id: packet.header.conn_id,
event,
}));
}
self.send_reject(peer, packet.header.conn_id, "unknown Dosh connection")
.await?;
Ok(Some(DoshServerEvent::Ignored))
}
pub async fn accept_stream(&mut self, conn_id: [u8; 16], stream_id: u64) -> Result<()> {
self.transport_mut(&conn_id)
.ok_or_else(|| anyhow!("unknown Dosh connection"))?
@@ -620,6 +639,85 @@ mod tests {
}
}
#[tokio::test]
async fn sdk_server_recv_drives_transport_maintenance_while_idle() {
let dir = tempfile::tempdir().unwrap();
let host_key = dir.path().join("host_key");
let authorized_keys = dir.path().join("authorized_keys");
let signing = SigningKey::from_bytes(&[94u8; 32]);
let keypair = ssh_key::private::Ed25519Keypair::from(&signing);
let private =
ssh_key::PrivateKey::new(ssh_key::private::KeypairData::from(keypair), "").unwrap();
std::fs::write(
&authorized_keys,
format!("{}\n", private.public_key().to_openssh().unwrap()),
)
.unwrap();
let server_config = ServerConfig {
host_key: host_key.to_string_lossy().to_string(),
authorized_keys: vec![authorized_keys.to_string_lossy().to_string()],
..ServerConfig::default()
};
let server_config = DoshServerConfig::new(server_config)
.bind_addr("127.0.0.1:0".parse().unwrap())
.service("echo")
.unwrap()
.require_current_user(false)
.transport(TransportConfig {
retransmit_after: Duration::from_millis(20),
keepalive_after: Duration::from_secs(60),
..TransportConfig::default()
});
let mut server = DoshServer::bind(server_config).await.unwrap();
let client_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client_socket.local_addr().unwrap();
let conn_id = [95u8; 16];
let session_key = [96u8; 32];
let transport = DoshTransport::new(
Arc::clone(&server.socket),
SessionTransportConfig {
role: SessionRole::Server,
conn_id,
session_key,
peer_addr: client_addr,
initial_send_seq: 1,
initial_ack: 0,
stream: server.config.transport.clone(),
},
);
server.transports.insert(conn_id, transport);
server
.transport_mut(&conn_id)
.unwrap()
.open_service("echo")
.await
.unwrap();
let mut buf = [0u8; 65535];
tokio::time::timeout(
Duration::from_millis(200),
client_socket.recv_from(&mut buf),
)
.await
.unwrap()
.unwrap();
let recv_task = tokio::spawn(async move { server.recv().await });
let (n, _) =
tokio::time::timeout(Duration::from_secs(1), client_socket.recv_from(&mut buf))
.await
.unwrap()
.unwrap();
recv_task.abort();
let packet = protocol::decode(&buf[..n]).unwrap();
assert_eq!(packet.header.kind, PacketKind::StreamOpen);
let plain = protocol::decrypt_body(&packet, &session_key, SERVER_TO_CLIENT).unwrap();
let open: protocol::StreamOpen = protocol::from_body(&plain).unwrap();
assert_eq!(open.target_host, "@dosh-echo");
}
#[tokio::test]
async fn bad_native_auth_does_not_consume_pending_challenge() {
let dir = tempfile::tempdir().unwrap();