mirror of
https://github.com/telemt/telemt.git
synced 2026-09-16 07:24:10 +03:00
WEB
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com> Co-Authored-By: John Preston <17900494+john-preston@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,168 @@
|
||||
use std::io;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::Ordering;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::web::frame::FrameType;
|
||||
use crate::web::stream::WebLogicalStream;
|
||||
use crate::proxy::shared_state::ConntrackClosePolicy;
|
||||
|
||||
use super::{WebSession, inbound_queue_cost, remember_closed};
|
||||
|
||||
impl WebSession {
|
||||
/// Starts one owned inner handshake and relay task for an admitted stream.
|
||||
pub(super) fn spawn_stream(self: &Arc<Self>, stream_id: u32, peer_port: u16) {
|
||||
let Some(manager) = self.manager.upgrade() else {
|
||||
self.stream_finished(stream_id, peer_port);
|
||||
return;
|
||||
};
|
||||
let generation = manager.active_generation();
|
||||
let Ok(connection_permit) = generation.max_connections.clone().try_acquire_owned() else {
|
||||
manager.record_stream_rejected();
|
||||
self.stream_finished(stream_id, peer_port);
|
||||
return;
|
||||
};
|
||||
let Some(handshake_permit) = manager.try_stream_handshake() else {
|
||||
self.stream_finished(stream_id, peer_port);
|
||||
return;
|
||||
};
|
||||
let deps = generation.client_runtime_deps();
|
||||
let replay_checker = Arc::clone(&generation.replay_checker);
|
||||
let session = Arc::clone(self);
|
||||
let cancel = self.cancel.clone();
|
||||
self.tasks_live.fetch_add(1, Ordering::AcqRel);
|
||||
let spawned = generation.spawn_session(async move {
|
||||
let _connection_permit = connection_permit;
|
||||
let _completion = StreamCompletion {
|
||||
session: Arc::clone(&session),
|
||||
stream_id,
|
||||
peer_port,
|
||||
};
|
||||
let stream = WebLogicalStream::new(Arc::clone(&session), stream_id);
|
||||
tokio::select! {
|
||||
_ = cancel.cancelled() => {}
|
||||
_ = run_stream(
|
||||
Arc::clone(&session),
|
||||
stream,
|
||||
deps,
|
||||
replay_checker,
|
||||
handshake_permit,
|
||||
peer_port,
|
||||
) => {}
|
||||
}
|
||||
});
|
||||
if !spawned {
|
||||
self.tasks_live.fetch_sub(1, Ordering::AcqRel);
|
||||
self.stream_finished(stream_id, peer_port);
|
||||
self.tasks_done.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
fn stream_finished(&self, stream_id: u32, peer_port: u16) {
|
||||
let (queued, reserved) = {
|
||||
let mut state = self.state.lock();
|
||||
let reserved = state.active_peer_ports.remove(&peer_port);
|
||||
let queued = state.streams.remove(&stream_id).map(|stream| {
|
||||
let (bytes, items) = inbound_queue_cost(&stream.inbound);
|
||||
self.release_locked(&mut state, bytes, items, false);
|
||||
remember_closed(
|
||||
&mut state,
|
||||
stream_id,
|
||||
self.limits.max_tombstones_per_session,
|
||||
);
|
||||
self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[])
|
||||
});
|
||||
(queued, reserved)
|
||||
};
|
||||
if reserved
|
||||
&& let Some(manager) = self.manager.upgrade()
|
||||
{
|
||||
manager.release_stream(
|
||||
self.profile_key,
|
||||
self.client_ip,
|
||||
self.profile.public_addr,
|
||||
peer_port,
|
||||
);
|
||||
}
|
||||
if let Some(queued) = queued {
|
||||
if !queued {
|
||||
self.close();
|
||||
}
|
||||
self.down_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct StreamCompletion {
|
||||
session: Arc<WebSession>,
|
||||
stream_id: u32,
|
||||
peer_port: u16,
|
||||
}
|
||||
|
||||
impl Drop for StreamCompletion {
|
||||
fn drop(&mut self) {
|
||||
self.session.stream_finished(self.stream_id, self.peer_port);
|
||||
if self.session.tasks_live.fetch_sub(1, Ordering::AcqRel) == 1 {
|
||||
self.session.tasks_done.notify_waiters();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_stream(
|
||||
session: Arc<WebSession>,
|
||||
stream: WebLogicalStream,
|
||||
deps: crate::proxy::authenticated::ClientRuntimeDeps,
|
||||
replay_checker: Arc<crate::stats::ReplayChecker>,
|
||||
handshake_permit: tokio::sync::OwnedSemaphorePermit,
|
||||
peer_port: u16,
|
||||
) {
|
||||
use tokio::io::AsyncReadExt;
|
||||
|
||||
use crate::protocol::constants::HANDSHAKE_LEN;
|
||||
use crate::proxy::authenticated::run_authenticated;
|
||||
use crate::proxy::handshake::handle_mtproto_handshake_for_web_user;
|
||||
|
||||
let (mut reader, writer) = tokio::io::split(stream);
|
||||
let mut handshake = [0u8; HANDSHAKE_LEN];
|
||||
let peer = std::net::SocketAddr::new(session.client_ip, peer_port);
|
||||
deps.stats.increment_connects_all();
|
||||
let handshake_result = tokio::time::timeout(
|
||||
Duration::from_secs(session.timeouts.stream_handshake_secs),
|
||||
async {
|
||||
reader.read_exact(&mut handshake).await?;
|
||||
Ok::<_, io::Error>(
|
||||
handle_mtproto_handshake_for_web_user(
|
||||
&handshake,
|
||||
reader,
|
||||
writer,
|
||||
peer,
|
||||
&deps.config,
|
||||
&replay_checker,
|
||||
&session.profile.user,
|
||||
session.profile.secret_mode,
|
||||
&deps.shared,
|
||||
)
|
||||
.await,
|
||||
)
|
||||
},
|
||||
)
|
||||
.await;
|
||||
drop(handshake_permit);
|
||||
let Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) =
|
||||
handshake_result
|
||||
else {
|
||||
deps.stats
|
||||
.increment_connects_bad_with_class("web_mtproto_bad_client");
|
||||
return;
|
||||
};
|
||||
let _ = run_authenticated(
|
||||
reader,
|
||||
writer,
|
||||
success,
|
||||
deps,
|
||||
session.profile.public_addr,
|
||||
peer,
|
||||
ConntrackClosePolicy::Suppress,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -0,0 +1,494 @@
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use bytes::{BufMut, Bytes, BytesMut};
|
||||
|
||||
use super::{
|
||||
DownBatch, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState,
|
||||
WebSession,
|
||||
};
|
||||
use crate::web::frame::{self, FrameType};
|
||||
use crate::web::manager::ManagerError;
|
||||
|
||||
impl WebSession {
|
||||
/// Polls pending downlink frames with cursor replay and newest-poll-wins semantics.
|
||||
pub(crate) async fn poll_down(&self, cursor: u64) -> Result<PollResult, ManagerError> {
|
||||
let epoch = {
|
||||
let mut state = self.state.lock();
|
||||
if state.closed {
|
||||
return Err(ManagerError::Closed);
|
||||
}
|
||||
state.last_activity = Instant::now();
|
||||
if let Some(unacked) = &state.unacked {
|
||||
if cursor == unacked.base_cursor {
|
||||
return Ok(PollResult {
|
||||
body: unacked.body.clone(),
|
||||
next_cursor: unacked.next_cursor,
|
||||
});
|
||||
}
|
||||
if cursor != unacked.next_cursor {
|
||||
drop(state);
|
||||
self.close();
|
||||
return Err(ManagerError::Protocol);
|
||||
}
|
||||
self.release_unacked_locked(&mut state);
|
||||
} else if cursor != state.down_cursor {
|
||||
drop(state);
|
||||
self.close();
|
||||
return Err(ManagerError::Protocol);
|
||||
}
|
||||
state.down_epoch = state.down_epoch.wrapping_add(1).max(1);
|
||||
state.down_epoch
|
||||
};
|
||||
self.down_notify.notify_waiters();
|
||||
|
||||
let deadline = Duration::from_secs(self.timeouts.long_poll_secs);
|
||||
let poll = async {
|
||||
loop {
|
||||
let notified = self.down_notify.notified();
|
||||
{
|
||||
let mut state = self.state.lock();
|
||||
if state.down_epoch != epoch {
|
||||
return Ok(PollResult {
|
||||
body: Bytes::new(),
|
||||
next_cursor: cursor,
|
||||
});
|
||||
}
|
||||
if !state.pending_frames.is_empty() {
|
||||
let batch = match self.take_down_batch_locked(&mut state, cursor) {
|
||||
Ok(batch) => batch,
|
||||
Err(error) => {
|
||||
drop(state);
|
||||
self.close();
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let result = PollResult {
|
||||
body: batch.body.clone(),
|
||||
next_cursor: batch.next_cursor,
|
||||
};
|
||||
if let Some(manager) = self.manager.upgrade() {
|
||||
manager.record_down(result.body.len());
|
||||
}
|
||||
state.unacked = Some(batch);
|
||||
return Ok(result);
|
||||
}
|
||||
if state.closed {
|
||||
return Err(ManagerError::Closed);
|
||||
}
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
};
|
||||
match tokio::time::timeout(deadline, poll).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
let mut state = self.state.lock();
|
||||
if state.down_epoch == epoch {
|
||||
state.last_activity = Instant::now();
|
||||
}
|
||||
Ok(PollResult {
|
||||
body: Bytes::new(),
|
||||
next_cursor: cursor,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Reserves session and process queue capacity while the session lock is held.
|
||||
pub(super) fn reserve_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
bytes: usize,
|
||||
items: usize,
|
||||
class: PendingClass,
|
||||
) -> bool {
|
||||
if bytes == 0 && items == 0 {
|
||||
return true;
|
||||
}
|
||||
let data_byte_limit = self
|
||||
.limits
|
||||
.pending_bytes_per_session
|
||||
.saturating_sub(self.limits.control_bytes_per_session);
|
||||
let item_reserve = 16usize.saturating_add(
|
||||
self.limits.max_streams_per_session.saturating_mul(3),
|
||||
);
|
||||
let data_item_limit = self
|
||||
.limits
|
||||
.pending_items_per_session
|
||||
.saturating_sub(item_reserve);
|
||||
if state.closed {
|
||||
return false;
|
||||
}
|
||||
let control = class == PendingClass::Control;
|
||||
let fits = if control {
|
||||
bytes <= self.limits.control_bytes_per_session
|
||||
&& items <= item_reserve
|
||||
&& state.pending_bytes
|
||||
<= self.limits.pending_bytes_per_session.saturating_sub(bytes)
|
||||
&& state.pending_items
|
||||
<= self.limits.pending_items_per_session.saturating_sub(items)
|
||||
&& state.pending_control_bytes
|
||||
<= self.limits.control_bytes_per_session.saturating_sub(bytes)
|
||||
&& state.pending_control_items <= item_reserve.saturating_sub(items)
|
||||
} else {
|
||||
let data_bytes = state
|
||||
.pending_bytes
|
||||
.saturating_sub(state.pending_control_bytes);
|
||||
let data_items = state
|
||||
.pending_items
|
||||
.saturating_sub(state.pending_control_items);
|
||||
let (byte_limit, item_limit) = if class == PendingClass::Downlink {
|
||||
let uplink_bytes = self
|
||||
.limits
|
||||
.max_body_bytes
|
||||
.saturating_add(
|
||||
self.limits
|
||||
.max_frames_per_body
|
||||
.saturating_mul(QUEUE_ITEM_COST),
|
||||
);
|
||||
(
|
||||
data_byte_limit.saturating_sub(uplink_bytes),
|
||||
data_item_limit.saturating_sub(self.limits.max_frames_per_body),
|
||||
)
|
||||
} else {
|
||||
(data_byte_limit, data_item_limit)
|
||||
};
|
||||
bytes <= byte_limit
|
||||
&& items <= item_limit
|
||||
&& data_bytes <= byte_limit - bytes
|
||||
&& data_items <= item_limit - items
|
||||
};
|
||||
if !fits {
|
||||
return false;
|
||||
}
|
||||
let Some(manager) = self.manager.upgrade() else {
|
||||
return false;
|
||||
};
|
||||
if !manager.try_reserve_pending(
|
||||
bytes,
|
||||
items,
|
||||
control,
|
||||
class == PendingClass::Downlink,
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
state.pending_bytes += bytes;
|
||||
state.pending_items += items;
|
||||
if control {
|
||||
state.pending_control_bytes += bytes;
|
||||
state.pending_control_items += items;
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
/// Releases session and process queue capacity while the session lock is held.
|
||||
pub(super) fn release_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
bytes: usize,
|
||||
items: usize,
|
||||
control: bool,
|
||||
) {
|
||||
state.pending_bytes = state.pending_bytes.saturating_sub(bytes);
|
||||
state.pending_items = state.pending_items.saturating_sub(items);
|
||||
if control {
|
||||
state.pending_control_bytes = state.pending_control_bytes.saturating_sub(bytes);
|
||||
state.pending_control_items = state.pending_control_items.saturating_sub(items);
|
||||
}
|
||||
if let Some(manager) = self.manager.upgrade() {
|
||||
manager.release_pending(bytes, items, control);
|
||||
}
|
||||
}
|
||||
|
||||
/// Coalesces one flow-control update into the bounded control queue.
|
||||
pub(super) fn queue_window_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
stream_id: u32,
|
||||
amount: u32,
|
||||
) -> bool {
|
||||
if amount == 0 {
|
||||
return true;
|
||||
}
|
||||
if let Some(index) = state.pending_windows.get(&stream_id).copied()
|
||||
&& let Some(queued) = state.pending_frames.get_mut(index)
|
||||
{
|
||||
let previous = u32::from_be_bytes(
|
||||
queued.encoded[frame::HEADER_BYTES..frame::HEADER_BYTES + 4]
|
||||
.try_into()
|
||||
.unwrap_or([0; 4]),
|
||||
);
|
||||
if let Some(total) = previous.checked_add(amount) {
|
||||
queued.encoded[frame::HEADER_BYTES..frame::HEADER_BYTES + 4]
|
||||
.copy_from_slice(&total.to_be_bytes());
|
||||
self.down_notify.notify_waiters();
|
||||
return true;
|
||||
}
|
||||
}
|
||||
self.queue_control_locked(
|
||||
state,
|
||||
FrameType::Window,
|
||||
stream_id,
|
||||
&frame::window_payload(amount),
|
||||
)
|
||||
}
|
||||
|
||||
/// Appends one control frame under both reserved queue budgets.
|
||||
pub(super) fn queue_control_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
frame_type: FrameType,
|
||||
stream_id: u32,
|
||||
payload: &[u8],
|
||||
) -> bool {
|
||||
self.queue_frame_locked(state, frame_type, stream_id, payload, true)
|
||||
}
|
||||
|
||||
/// Appends one server-to-client DATA frame under downlink data budgets.
|
||||
pub(super) fn queue_data_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
stream_id: u32,
|
||||
payload: &[u8],
|
||||
) -> bool {
|
||||
let can_coalesce = state.pending_frames.back().is_some_and(|last| {
|
||||
last.frame_type == FrameType::Data
|
||||
&& last.stream_id == stream_id
|
||||
&& last.encoded.len() - frame::HEADER_BYTES + payload.len()
|
||||
<= self.limits.max_frame_payload_bytes
|
||||
});
|
||||
if can_coalesce {
|
||||
if !self.reserve_locked(state, payload.len(), 0, PendingClass::Downlink) {
|
||||
return false;
|
||||
}
|
||||
let Some(last) = state.pending_frames.back_mut() else {
|
||||
return false;
|
||||
};
|
||||
last.encoded.extend_from_slice(payload);
|
||||
last.cost += payload.len();
|
||||
let payload_len = (last.encoded.len() - frame::HEADER_BYTES) as u32;
|
||||
last.encoded[4..8].copy_from_slice(&payload_len.to_be_bytes());
|
||||
return true;
|
||||
}
|
||||
self.queue_frame_locked(state, FrameType::Data, stream_id, payload, false)
|
||||
}
|
||||
|
||||
fn queue_frame_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
frame_type: FrameType,
|
||||
stream_id: u32,
|
||||
payload: &[u8],
|
||||
control: bool,
|
||||
) -> bool {
|
||||
let cost = frame::HEADER_BYTES + payload.len() + QUEUE_ITEM_COST;
|
||||
let class = if control {
|
||||
PendingClass::Control
|
||||
} else {
|
||||
PendingClass::Downlink
|
||||
};
|
||||
if !self.reserve_locked(state, cost, 1, class) {
|
||||
return false;
|
||||
}
|
||||
let mut encoded = BytesMut::with_capacity(frame::HEADER_BYTES + payload.len());
|
||||
encoded.put_u8(frame_type as u8);
|
||||
encoded.put_u8((stream_id >> 16) as u8);
|
||||
encoded.put_u8((stream_id >> 8) as u8);
|
||||
encoded.put_u8(stream_id as u8);
|
||||
encoded.put_u32(payload.len() as u32);
|
||||
encoded.extend_from_slice(payload);
|
||||
let index = state.pending_frames.len();
|
||||
state.pending_frames.push_back(QueuedFrame {
|
||||
encoded,
|
||||
frame_type,
|
||||
stream_id,
|
||||
control,
|
||||
cost,
|
||||
});
|
||||
if frame_type == FrameType::Window {
|
||||
state.pending_windows.insert(stream_id, index);
|
||||
}
|
||||
self.down_notify.notify_waiters();
|
||||
true
|
||||
}
|
||||
|
||||
fn take_down_batch_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
cursor: u64,
|
||||
) -> Result<DownBatch, ManagerError> {
|
||||
let next_cursor = state
|
||||
.down_cursor
|
||||
.checked_add(1)
|
||||
.ok_or(ManagerError::Protocol)?;
|
||||
let mut count = 0usize;
|
||||
let mut body_len = 0usize;
|
||||
for queued in &state.pending_frames {
|
||||
if count >= self.limits.max_frames_per_body
|
||||
|| (count != 0
|
||||
&& body_len.saturating_add(queued.encoded.len())
|
||||
> self.limits.carrier_batch_bytes)
|
||||
{
|
||||
break;
|
||||
}
|
||||
body_len += queued.encoded.len();
|
||||
count += 1;
|
||||
}
|
||||
let mut body = BytesMut::with_capacity(body_len);
|
||||
let mut data_bytes = 0usize;
|
||||
let mut data_items = 0usize;
|
||||
let mut control_bytes = 0usize;
|
||||
let mut control_items = 0usize;
|
||||
for index in 0..count {
|
||||
let Some(queued) = state.pending_frames.get(index) else {
|
||||
break;
|
||||
};
|
||||
if queued.frame_type == FrameType::Window
|
||||
&& state.pending_windows.get(&queued.stream_id) == Some(&index)
|
||||
{
|
||||
state.pending_windows.remove(&queued.stream_id);
|
||||
}
|
||||
}
|
||||
for _ in 0..count {
|
||||
let Some(queued) = state.pending_frames.pop_front() else {
|
||||
break;
|
||||
};
|
||||
body.extend_from_slice(&queued.encoded);
|
||||
if queued.control {
|
||||
control_bytes += queued.cost;
|
||||
control_items += 1;
|
||||
} else {
|
||||
data_bytes += queued.cost;
|
||||
data_items += 1;
|
||||
}
|
||||
}
|
||||
for index in state.pending_windows.values_mut() {
|
||||
*index = index.saturating_sub(count);
|
||||
}
|
||||
state.down_cursor = next_cursor;
|
||||
Ok(DownBatch {
|
||||
body: body.freeze(),
|
||||
base_cursor: cursor,
|
||||
next_cursor,
|
||||
data_bytes,
|
||||
data_items,
|
||||
control_bytes,
|
||||
control_items,
|
||||
})
|
||||
}
|
||||
|
||||
fn release_unacked_locked(&self, state: &mut SessionState) {
|
||||
let Some(batch) = state.unacked.take() else {
|
||||
return;
|
||||
};
|
||||
self.release_locked(state, batch.data_bytes, batch.data_items, false);
|
||||
self.release_locked(state, batch.control_bytes, batch.control_items, true);
|
||||
for stream in state.streams.values_mut() {
|
||||
if let Some(waker) = stream.write_waker.take() {
|
||||
waker.wake();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::config::{
|
||||
WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
|
||||
};
|
||||
use crate::web::manager::WebProcessRuntime;
|
||||
|
||||
fn session() -> Arc<WebSession> {
|
||||
let profile = Arc::new(WebRuntimeProfile {
|
||||
host: "proxy.example.com".to_string(),
|
||||
public_addr: SocketAddr::from(([203, 0, 113, 10], 443)),
|
||||
user: "alice".to_string(),
|
||||
secret_mode: WebSecretMode::Plain,
|
||||
capability: [0; 32],
|
||||
max_sessions: 1,
|
||||
max_streams: 1,
|
||||
max_streams_per_session: 1,
|
||||
});
|
||||
WebSession::new(
|
||||
std::sync::Weak::<WebProcessRuntime>::new(),
|
||||
[1; 32],
|
||||
"192.0.2.10".parse().unwrap(),
|
||||
profile,
|
||||
[2; 32],
|
||||
WebLimitsConfig::default(),
|
||||
WebTimeoutsConfig::default(),
|
||||
)
|
||||
}
|
||||
|
||||
fn queue_close(session: &WebSession) {
|
||||
let encoded = frame::encode(FrameType::Close, 1, &[]);
|
||||
session.state.lock().pending_frames.push_back(QueuedFrame {
|
||||
encoded: BytesMut::from(encoded.as_ref()),
|
||||
frame_type: FrameType::Close,
|
||||
stream_id: 1,
|
||||
control: true,
|
||||
cost: frame::HEADER_BYTES + QUEUE_ITEM_COST,
|
||||
});
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn downlink_replays_unacknowledged_batch_byte_for_byte() {
|
||||
let session = session();
|
||||
queue_close(&session);
|
||||
let first = session.poll_down(0).await.unwrap();
|
||||
let replay = session.poll_down(0).await.unwrap();
|
||||
assert_eq!(first.next_cursor, 1);
|
||||
assert_eq!(replay.next_cursor, 1);
|
||||
assert_eq!(first.body, replay.body);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_or_overflowing_cursor_closes_session() {
|
||||
let invalid = session();
|
||||
assert!(matches!(
|
||||
invalid.poll_down(1).await,
|
||||
Err(ManagerError::Protocol)
|
||||
));
|
||||
assert!(invalid.state.lock().closed);
|
||||
|
||||
let overflow = session();
|
||||
{
|
||||
let mut state = overflow.state.lock();
|
||||
state.down_cursor = u64::MAX;
|
||||
}
|
||||
queue_close(&overflow);
|
||||
assert!(matches!(
|
||||
overflow.poll_down(u64::MAX).await,
|
||||
Err(ManagerError::Protocol)
|
||||
));
|
||||
assert!(overflow.state.lock().closed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn newer_poll_supersedes_older_poll_without_closing_session() {
|
||||
let session = session();
|
||||
let first_session = Arc::clone(&session);
|
||||
let first = tokio::spawn(async move { first_session.poll_down(0).await });
|
||||
while session.state.lock().down_epoch < 1 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
let second_session = Arc::clone(&session);
|
||||
let second = tokio::spawn(async move { second_session.poll_down(0).await });
|
||||
while session.state.lock().down_epoch < 2 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
let superseded = tokio::time::timeout(Duration::from_secs(1), first)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert!(superseded.body.is_empty());
|
||||
assert_eq!(superseded.next_cursor, 0);
|
||||
assert!(!session.state.lock().closed);
|
||||
second.abort();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,431 @@
|
||||
use std::collections::{HashMap, HashSet, VecDeque};
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::time::Instant;
|
||||
|
||||
use bytes::Bytes;
|
||||
use sha2::{Digest, Sha256};
|
||||
use subtle::ConstantTimeEq;
|
||||
|
||||
use super::{
|
||||
InboundChunk, PendingClass, QUEUE_ITEM_COST, SessionState, StreamState, WebSession,
|
||||
inbound_queue_cost, remember_closed,
|
||||
};
|
||||
use crate::web::frame::{self, Frame, FrameType};
|
||||
use crate::web::manager::{ManagerError, TokenHash};
|
||||
|
||||
impl WebSession {
|
||||
/// Applies one exactly-once uplink batch.
|
||||
pub(crate) fn process_up(
|
||||
self: &Arc<Self>,
|
||||
sequence: u64,
|
||||
body: &[u8],
|
||||
) -> Result<u64, ManagerError> {
|
||||
if self
|
||||
.up_active
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.is_err()
|
||||
{
|
||||
return Err(ManagerError::Concurrent);
|
||||
}
|
||||
let _uplink = UplinkGuard(&self.up_active);
|
||||
let frames = match frame::parse_all(body, &self.limits) {
|
||||
Ok(frames) => frames,
|
||||
Err(_) => {
|
||||
self.close();
|
||||
return Err(ManagerError::Protocol);
|
||||
}
|
||||
};
|
||||
if frames
|
||||
.iter()
|
||||
.copied()
|
||||
.any(|value| frame::validate_client_shape(value).is_err())
|
||||
{
|
||||
self.close();
|
||||
return Err(ManagerError::Protocol);
|
||||
}
|
||||
let digest: TokenHash = Sha256::digest(body).into();
|
||||
let mut opened = Vec::new();
|
||||
let result = {
|
||||
let mut state = self.state.lock();
|
||||
if state.closed {
|
||||
return Err(ManagerError::Closed);
|
||||
}
|
||||
state.last_activity = Instant::now();
|
||||
if sequence == state.last_up_sequence && sequence != 0 {
|
||||
return if bool::from(state.last_up_digest.ct_eq(&digest)) {
|
||||
Ok(sequence)
|
||||
} else {
|
||||
drop(state);
|
||||
self.close();
|
||||
Err(ManagerError::Protocol)
|
||||
};
|
||||
}
|
||||
if sequence == 0 || sequence != state.last_up_sequence.saturating_add(1) {
|
||||
drop(state);
|
||||
self.close();
|
||||
return Err(ManagerError::Protocol);
|
||||
}
|
||||
if !validate_batch(&state, &frames) {
|
||||
drop(state);
|
||||
self.close();
|
||||
return Err(ManagerError::Protocol);
|
||||
}
|
||||
let (reserve_bytes, reserve_items) = inbound_reservation(&state, &frames);
|
||||
if !self.reserve_locked(
|
||||
&mut state,
|
||||
reserve_bytes,
|
||||
reserve_items,
|
||||
PendingClass::Uplink,
|
||||
) {
|
||||
return Err(ManagerError::Backpressure);
|
||||
}
|
||||
let mut unused_bytes = reserve_bytes;
|
||||
let mut unused_items = reserve_items;
|
||||
let applied = self.apply_batch_locked(
|
||||
&mut state,
|
||||
&frames,
|
||||
&mut opened,
|
||||
&mut unused_bytes,
|
||||
&mut unused_items,
|
||||
);
|
||||
self.release_locked(&mut state, unused_bytes, unused_items, false);
|
||||
if !applied {
|
||||
Err(ManagerError::Closed)
|
||||
} else {
|
||||
state.last_up_sequence = sequence;
|
||||
state.last_up_digest = digest;
|
||||
Ok(sequence)
|
||||
}
|
||||
};
|
||||
if matches!(result, Err(ManagerError::Backpressure)) {
|
||||
return result;
|
||||
}
|
||||
if result.is_err() {
|
||||
self.close();
|
||||
for (_, peer_port) in opened {
|
||||
self.release_stream_reservation(peer_port);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
for (stream_id, peer_port) in opened {
|
||||
self.spawn_stream(stream_id, peer_port);
|
||||
}
|
||||
if let Some(manager) = self.manager.upgrade() {
|
||||
manager.record_up(body.len());
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn apply_batch_locked(
|
||||
&self,
|
||||
state: &mut SessionState,
|
||||
frames: &[Frame<'_>],
|
||||
opened: &mut Vec<(u32, u16)>,
|
||||
unused_bytes: &mut usize,
|
||||
unused_items: &mut usize,
|
||||
) -> bool {
|
||||
for value in frames {
|
||||
if value.stream_id == 0 {
|
||||
continue;
|
||||
}
|
||||
let was_closed = state.closed_streams.contains(&value.stream_id);
|
||||
match value.frame_type {
|
||||
FrameType::Open => {
|
||||
let Some(peer_port) = self.reserve_stream_locked(state) else {
|
||||
remember_closed(
|
||||
state,
|
||||
value.stream_id,
|
||||
self.limits.max_tombstones_per_session,
|
||||
);
|
||||
if !self.queue_control_locked(
|
||||
state,
|
||||
FrameType::Close,
|
||||
value.stream_id,
|
||||
&[],
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
continue;
|
||||
};
|
||||
state.streams.insert(
|
||||
value.stream_id,
|
||||
StreamState {
|
||||
inbound: VecDeque::new(),
|
||||
receive_window: frame::INITIAL_STREAM_WINDOW,
|
||||
send_credit: u64::from(frame::INITIAL_STREAM_WINDOW),
|
||||
read_waker: None,
|
||||
write_waker: None,
|
||||
},
|
||||
);
|
||||
opened.push((value.stream_id, peer_port));
|
||||
}
|
||||
FrameType::Data if !was_closed => {
|
||||
let Some(stream) = state.streams.get_mut(&value.stream_id) else {
|
||||
return false;
|
||||
};
|
||||
stream.receive_window -= value.payload.len() as u32;
|
||||
stream.inbound.push_back(InboundChunk {
|
||||
bytes: Bytes::copy_from_slice(value.payload),
|
||||
offset: 0,
|
||||
});
|
||||
*unused_bytes = unused_bytes
|
||||
.saturating_sub(value.payload.len() + QUEUE_ITEM_COST);
|
||||
*unused_items = unused_items.saturating_sub(1);
|
||||
if let Some(waker) = stream.read_waker.take() {
|
||||
waker.wake();
|
||||
}
|
||||
}
|
||||
FrameType::Window if !was_closed => {
|
||||
let Some(stream) = state.streams.get_mut(&value.stream_id) else {
|
||||
return false;
|
||||
};
|
||||
let amount = frame::window_amount(value.payload).unwrap_or(0);
|
||||
stream.send_credit = stream
|
||||
.send_credit
|
||||
.saturating_add(u64::from(amount))
|
||||
.min(u64::from(u32::MAX));
|
||||
if let Some(waker) = stream.write_waker.take() {
|
||||
waker.wake();
|
||||
}
|
||||
}
|
||||
FrameType::Close if !was_closed => {
|
||||
let Some(stream) = state.streams.remove(&value.stream_id) else {
|
||||
return false;
|
||||
};
|
||||
let (bytes, items) = inbound_queue_cost(&stream.inbound);
|
||||
self.release_locked(state, bytes, items, false);
|
||||
remember_closed(
|
||||
state,
|
||||
value.stream_id,
|
||||
self.limits.max_tombstones_per_session,
|
||||
);
|
||||
if let Some(waker) = stream.read_waker {
|
||||
waker.wake();
|
||||
}
|
||||
if let Some(waker) = stream.write_waker {
|
||||
waker.wake();
|
||||
}
|
||||
}
|
||||
FrameType::Data | FrameType::Window | FrameType::Close => {}
|
||||
_ => return false,
|
||||
}
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
fn reserve_stream_locked(&self, state: &mut SessionState) -> Option<u16> {
|
||||
if state.active_peer_ports.len() >= self.profile.max_streams_per_session {
|
||||
return None;
|
||||
}
|
||||
let manager = self.manager.upgrade()?;
|
||||
let peer_port = manager.try_acquire_stream(
|
||||
self.profile_key,
|
||||
self.profile.max_streams,
|
||||
self.client_ip,
|
||||
self.profile.public_addr,
|
||||
)?;
|
||||
if state.active_peer_ports.insert(peer_port) {
|
||||
return Some(peer_port);
|
||||
}
|
||||
manager.release_stream(
|
||||
self.profile_key,
|
||||
self.client_ip,
|
||||
self.profile.public_addr,
|
||||
peer_port,
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
struct UplinkGuard<'a>(&'a AtomicBool);
|
||||
|
||||
impl Drop for UplinkGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(false, Ordering::Release);
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_batch(state: &SessionState, frames: &[Frame<'_>]) -> bool {
|
||||
let mut live = state
|
||||
.streams
|
||||
.iter()
|
||||
.map(|(id, stream)| (*id, (stream.receive_window, stream.send_credit)))
|
||||
.collect::<HashMap<_, _>>();
|
||||
let mut closed = HashSet::new();
|
||||
for value in frames {
|
||||
if value.stream_id == 0 {
|
||||
if value.frame_type != FrameType::Pong {
|
||||
return false;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let was_closed = state.closed_streams.contains(&value.stream_id)
|
||||
|| closed.contains(&value.stream_id);
|
||||
match value.frame_type {
|
||||
FrameType::Open => {
|
||||
if live.contains_key(&value.stream_id) || was_closed {
|
||||
return false;
|
||||
}
|
||||
live.insert(
|
||||
value.stream_id,
|
||||
(
|
||||
frame::INITIAL_STREAM_WINDOW,
|
||||
u64::from(frame::INITIAL_STREAM_WINDOW),
|
||||
),
|
||||
);
|
||||
}
|
||||
FrameType::Data if !was_closed => {
|
||||
let Some((receive_window, send_credit)) = live.get_mut(&value.stream_id) else {
|
||||
return false;
|
||||
};
|
||||
let Ok(payload_len) = u32::try_from(value.payload.len()) else {
|
||||
return false;
|
||||
};
|
||||
if payload_len > *receive_window {
|
||||
return false;
|
||||
}
|
||||
*receive_window -= payload_len;
|
||||
let _ = send_credit;
|
||||
}
|
||||
FrameType::Window if !was_closed => {
|
||||
let Some((_, send_credit)) = live.get_mut(&value.stream_id) else {
|
||||
return false;
|
||||
};
|
||||
let Ok(amount) = frame::window_amount(value.payload) else {
|
||||
return false;
|
||||
};
|
||||
*send_credit = send_credit
|
||||
.saturating_add(u64::from(amount))
|
||||
.min(u64::from(u32::MAX));
|
||||
}
|
||||
FrameType::Close if !was_closed => {
|
||||
if live.remove(&value.stream_id).is_none() {
|
||||
return false;
|
||||
}
|
||||
closed.insert(value.stream_id);
|
||||
}
|
||||
FrameType::Data | FrameType::Window | FrameType::Close => {}
|
||||
_ => return false,
|
||||
}
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
fn inbound_reservation(state: &SessionState, frames: &[Frame<'_>]) -> (usize, usize) {
|
||||
let mut live = state.streams.keys().copied().collect::<HashSet<_>>();
|
||||
let mut bytes = 0usize;
|
||||
let mut items = 0usize;
|
||||
for value in frames {
|
||||
match value.frame_type {
|
||||
FrameType::Open => {
|
||||
live.insert(value.stream_id);
|
||||
}
|
||||
FrameType::Data if live.contains(&value.stream_id) => {
|
||||
bytes = bytes.saturating_add(value.payload.len() + QUEUE_ITEM_COST);
|
||||
items = items.saturating_add(1);
|
||||
}
|
||||
FrameType::Close => {
|
||||
live.remove(&value.stream_id);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
(bytes, items)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use crate::config::{
|
||||
WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
|
||||
};
|
||||
use crate::web::manager::WebProcessRuntime;
|
||||
|
||||
fn session() -> Arc<WebSession> {
|
||||
let profile = Arc::new(WebRuntimeProfile {
|
||||
host: "proxy.example.com".to_string(),
|
||||
public_addr: SocketAddr::from(([203, 0, 113, 10], 443)),
|
||||
user: "alice".to_string(),
|
||||
secret_mode: WebSecretMode::Plain,
|
||||
capability: [0; 32],
|
||||
max_sessions: 1,
|
||||
max_streams: 1,
|
||||
max_streams_per_session: 1,
|
||||
});
|
||||
WebSession::new(
|
||||
std::sync::Weak::<WebProcessRuntime>::new(),
|
||||
[1; 32],
|
||||
"192.0.2.10".parse().unwrap(),
|
||||
profile,
|
||||
[2; 32],
|
||||
WebLimitsConfig::default(),
|
||||
WebTimeoutsConfig::default(),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uplink_retry_commits_only_one_exact_body() {
|
||||
let session = session();
|
||||
let first = frame::encode(FrameType::Pong, 0, &[1, 2, 3]);
|
||||
assert_eq!(session.process_up(1, &first), Ok(1));
|
||||
assert_eq!(session.process_up(1, &first), Ok(1));
|
||||
|
||||
let changed = frame::encode(FrameType::Pong, 0, &[1, 2, 4]);
|
||||
assert_eq!(session.process_up(1, &changed), Err(ManagerError::Protocol));
|
||||
assert!(session.state.lock().closed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_uplink_does_not_commit_sequence() {
|
||||
let session = session();
|
||||
let body = frame::encode(FrameType::Pong, 0, &[]);
|
||||
session.up_active.store(true, Ordering::Release);
|
||||
assert_eq!(
|
||||
session.process_up(1, &body),
|
||||
Err(ManagerError::Concurrent)
|
||||
);
|
||||
assert_eq!(session.state.lock().last_up_sequence, 0);
|
||||
session.up_active.store(false, Ordering::Release);
|
||||
assert_eq!(session.process_up(1, &body), Ok(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backpressured_uplink_does_not_commit_or_close() {
|
||||
let session = session();
|
||||
{
|
||||
let mut state = session.state.lock();
|
||||
state.streams.insert(
|
||||
1,
|
||||
StreamState {
|
||||
inbound: VecDeque::new(),
|
||||
receive_window: frame::INITIAL_STREAM_WINDOW,
|
||||
send_credit: u64::from(frame::INITIAL_STREAM_WINDOW),
|
||||
read_waker: None,
|
||||
write_waker: None,
|
||||
},
|
||||
);
|
||||
state.pending_bytes = session.limits.pending_bytes_per_session;
|
||||
}
|
||||
let body = frame::encode(FrameType::Data, 1, &[1]);
|
||||
|
||||
assert_eq!(
|
||||
session.process_up(1, &body),
|
||||
Err(ManagerError::Backpressure)
|
||||
);
|
||||
let state = session.state.lock();
|
||||
assert!(!state.closed);
|
||||
assert_eq!(state.last_up_sequence, 0);
|
||||
assert!(state.streams.get(&1).unwrap().inbound.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uplink_gap_is_fatal() {
|
||||
let session = session();
|
||||
let body = frame::encode(FrameType::Pong, 0, &[]);
|
||||
assert_eq!(session.process_up(2, &body), Err(ManagerError::Protocol));
|
||||
assert!(session.state.lock().closed);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user