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:
Alexey
2026-08-23 03:12:11 +03:00
parent 8dbd24b11b
commit 1029703c2c
58 changed files with 7460 additions and 2646 deletions
+168
View File
@@ -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;
}
+494
View File
@@ -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();
}
}
+431
View File
@@ -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);
}
}