use std::sync::Arc;
use snafu::ResultExt;
use tracing::instrument;
use super::error::{
ConnectLocalStreamSnafu, ConnectRemoteStreamSnafu, ControlIoTimeoutSnafu,
DecodePbConnStreamRespSnafu, EncodePbConnStreamReqSnafu, PbConnStreamRespNotMatchSnafu, Result,
WritePbConnStreamReqSnafu,
};
use crate::server::error::CreateHeaderToolSnafu;
use pb_mapper_core::checksum::Credential;
use pb_mapper_core::config::{ResolvedAddrs, control_io_timeout};
use pb_mapper_core::snafu_error_handle;
use pb_mapper_protocol::command::{MessageSerializer, PbConnRequest, PbConnResponse};
use pb_mapper_protocol::forward::StreamForward;
use pb_mapper_protocol::secure::ClientHeaderSession;
use uni_stream::stream::{StreamProvider, StreamSplit, set_tcp_keep_alive, set_tcp_nodelay};
#[derive(Clone, Debug)]
pub struct StreamConnect {
pub local_addr: ResolvedAddrs,
pub remote_addr: ResolvedAddrs,
pub keep_alive: bool,
pub namespace: Option<u64>,
pub credential: Credential,
}
#[instrument(skip(connect, setup_permit), fields(key = %key))]
pub async fn handle_stream<LocalStream: StreamProvider>(
key: Arc<str>,
client_id: u32,
server_generation: u64,
connect: StreamConnect,
setup_permit: tokio::sync::OwnedSemaphorePermit,
) -> Result<()>
where
LocalStream::Item: StreamForward,
{
let StreamConnect {
local_addr,
remote_addr,
keep_alive,
namespace,
credential,
} = connect;
let request = match namespace {
Some(namespace) => PbConnRequest::StreamScoped {
key: key.to_string(),
namespace,
dst_id: client_id,
server_generation,
},
None => PbConnRequest::Stream {
key: key.to_string(),
dst_id: client_id,
server_generation,
},
};
let msg = request.encode().context(EncodePbConnStreamReqSnafu)?;
let timeout = control_io_timeout().min(std::time::Duration::from_secs(5));
let deadline = tokio::time::Instant::now() + timeout;
let mut remote_stream =
match tokio::time::timeout_at(deadline, crate::addr::connect_tcp(&remote_addr)).await {
Ok(result) => result.context(ConnectRemoteStreamSnafu)?,
Err(_) => ControlIoTimeoutSnafu {
action: "connect remote stream",
timeout,
}
.fail()?,
};
if keep_alive {
snafu_error_handle!(
set_tcp_keep_alive(&remote_stream),
"remote stream set keepalive"
);
}
snafu_error_handle!(set_tcp_nodelay(&remote_stream), "remote stream set nodelay");
let codec_key = {
let session = ClientHeaderSession::new_v2(&credential)
.context(CreateHeaderToolSnafu { action: "session" })?;
let response = session
.exchange(
&mut remote_stream,
&msg,
deadline.saturating_duration_since(tokio::time::Instant::now()),
)
.await
.context(WritePbConnStreamReqSnafu)?;
let resp = PbConnResponse::decode(&response).context(DecodePbConnStreamRespSnafu)?;
match resp {
PbConnResponse::Stream { codec_key } => codec_key,
PbConnResponse::Error(error) => PbConnStreamRespNotMatchSnafu {
resp: format!("{}: {}", error.code, error.message),
}
.fail()?,
_ => PbConnStreamRespNotMatchSnafu {
resp: format!("{resp:?}"),
}
.fail()?,
}
};
let mut local_stream = match tokio::time::timeout_at(
deadline,
LocalStream::from_addr(local_addr.as_slice()),
)
.await
{
Ok(result) => result.context(ConnectLocalStreamSnafu)?,
Err(_) => {
return ControlIoTimeoutSnafu {
action: "connect local service",
timeout,
}
.fail();
}
};
drop(setup_permit);
let (client_reader, client_writer) = remote_stream.split();
let (server_reader, server_writer) = local_stream.split();
snafu_error_handle!(
<LocalStream::Item as StreamForward>::forward_local_to_remote(
codec_key,
*credential.key(),
server_reader,
server_writer,
client_reader,
client_writer,
)
.await
);
Ok(())
}