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
+240
View File
@@ -0,0 +1,240 @@
use base64::Engine as _;
use crate::crypto::SecureRandom;
/// Browser security policy for the transient Telegram Desktop bridge page.
pub(crate) const PERMISSIONS_POLICY: &str = "accelerometer=(), autoplay=(), camera=(), clipboard-read=(), clipboard-write=(), display-capture=(), encrypted-media=(), fullscreen=(), geolocation=(), gyroscope=(), hid=(), idle-detection=(), magnetometer=(), microphone=(), midi=(), payment=(), picture-in-picture=(), publickey-credentials-create=(), publickey-credentials-get=(), screen-wake-lock=(), serial=(), usb=(), web-share=(), xr-spatial-tracking=()";
/// Fully rendered bridge response and its per-response script policy.
pub(crate) struct BridgePage {
/// Complete transient HTML document.
pub(crate) body: String,
/// Nonce-bound policy that authorizes only the embedded bridge script.
pub(crate) content_security_policy: String,
}
/// Renders the HTTPS-only WEB carrier bridge with a fresh CSP nonce.
pub(crate) fn render(
host: &str,
bootstrap: &str,
batch_limit: usize,
queue_limit: usize,
queue_items: usize,
rng: &SecureRandom,
) -> BridgePage {
let mut nonce = [0u8; 18];
rng.fill(&mut nonce);
let nonce = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(nonce);
let body = DOCUMENT
.replace("__NONCE__", &nonce)
.replace("__HOST__", host)
.replace("__BOOTSTRAP__", bootstrap)
.replace("__BATCH_LIMIT__", &batch_limit.to_string())
.replace("__QUEUE_LIMIT__", &queue_limit.to_string())
.replace("__QUEUE_ITEMS__", &queue_items.to_string());
BridgePage {
body,
content_security_policy: format!(
"default-src 'none'; base-uri 'none'; child-src 'none'; connect-src 'self' wss://{host}; font-src 'none'; form-action 'none'; frame-ancestors http://127.0.0.1:*; frame-src 'none'; img-src 'none'; manifest-src 'none'; media-src 'none'; object-src 'none'; script-src 'nonce-{nonce}'; style-src 'none'; worker-src 'none'; sandbox allow-same-origin allow-scripts"
),
}
}
const DOCUMENT: &str = r##"<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>Connection</title>
</head>
<body>
<script nonce="__NONCE__">
(()=>{
'use strict';
const relayOrigin='https://__HOST__',bootstrap='__BOOTSTRAP__';
const batchLimit=__BATCH_LIMIT__,queueLimit=__QUEUE_LIMIT__,queueItemLimit=__QUEUE_ITEMS__;
const fragment=location.hash,androidNonce=/^#android=([A-Za-z0-9_-]{43})$/.exec(fragment)?.[1]||'';
history.replaceState(null,'',location.pathname);
let initialized=false,closed=false,port=null,sessionToken='',createStarted=false;
let queuedBytes=0,queuedItems=0,upSequence=1,downCursor='0',upRunning=false,pollController=null;
const pending=[],upPending=[];
const status=state=>{if(port&&!closed)port.postMessage({t:'status',state})};
const pause=milliseconds=>new Promise(resolve=>setTimeout(resolve,milliseconds));
const options=(method,token,body,headers,signal,keepalive)=>({
method,body,signal,keepalive:!!keepalive,mode:'same-origin',credentials:'omit',cache:'no-store',redirect:'error',referrerPolicy:'no-referrer',
headers:Object.assign(token?{Authorization:'Bearer '+token}:{},body?{'Content-Type':'application/octet-stream'}:{},headers||{})
});
function reserve(data){
if(!data.byteLength||data.byteLength>queueLimit-queuedBytes||queuedItems>=queueItemLimit)return false;
queuedBytes+=data.byteLength;queuedItems++;return true;
}
function release(bytes,items){queuedBytes-=bytes;queuedItems-=items}
function frameBound(value,maxFrames,maxBytes){
const view=new DataView(value);let offset=0,frames=0;
while(offset<value.byteLength){
if(value.byteLength-offset<8)throw new Error('invalid frame batch');
const size=view.getUint32(offset+4),end=offset+8+size;
if(size>1048576||end>value.byteLength)throw new Error('invalid frame');
if(frames>0&&(frames>=maxFrames||end>maxBytes))break;
frames++;offset=end;
}
if(!frames)throw new Error('empty frame batch');
return {frames,bytes:offset};
}
function splitFrames(value){
const view=new DataView(value),result=[];let offset=0;
while(offset<value.byteLength){
if(value.byteLength-offset<8||result.length>=4096)throw new Error('invalid frame batch');
const size=view.getUint32(offset+4),end=offset+8+size;
if(size>1048576||end>value.byteLength)throw new Error('invalid frame');
result.push(offset===0&&end===value.byteLength?value:value.slice(offset,end));offset=end;
}
if(!result.length)throw new Error('empty frame batch');return result;
}
function joinPending(values){
let total=0,count=0,frames=0;
while(count<values.length){
const bound=frameBound(values[count],4096,batchLimit),whole=bound.bytes===values[count].byteLength;
if(count===0&&!whole){
const head=new Uint8Array(values[0],0,bound.bytes).slice();
values[0]=values[0].slice(bound.bytes);queuedItems++;
return {body:head.buffer,total:bound.bytes,count:1};
}
if(count&&(total+values[count].byteLength>batchLimit||frames+bound.frames>4096))break;
total+=values[count].byteLength;frames+=bound.frames;count++;
}
const joined=new Uint8Array(total);let offset=0;
for(const data of values.splice(0,count)){joined.set(new Uint8Array(data),offset);offset+=data.byteLength}
return {body:joined.buffer,total,count};
}
function retryAfterMs(response){
const value=Number(response.headers.get('Retry-After'));
return Number.isFinite(value)&&value>=0?Math.min(value*1000,30000):0;
}
async function request(path,makeOptions){
let delay=250,attempt=0;const deadline=Date.now()+90000;
while(true){
const requestOptions=makeOptions(),controller=new AbortController(),external=requestOptions.signal;
const abort=()=>controller.abort();if(external)external.addEventListener('abort',abort,{once:true});
requestOptions.signal=controller.signal;const timer=setTimeout(abort,90000);
let serviceUnavailable=false,wait=0;
try{
const response=await fetch(relayOrigin+path,requestOptions);
if(response.status!==503)return response;
serviceUnavailable=true;wait=retryAfterMs(response);await response.arrayBuffer();
}catch(error){
if(closed||(external&&external.aborted))throw error;
if(++attempt===9)throw new Error('carrier retry limit reached');
}finally{clearTimeout(timer);if(external)external.removeEventListener('abort',abort)}
if(serviceUnavailable&&Date.now()>=deadline)throw new Error('carrier retry limit reached');
status('reconnecting');await pause(wait||(delay+Math.floor(Math.random()*Math.max(1,delay/4))));
if(!serviceUnavailable)delay=Math.min(delay*2,5000);
}
}
function fail(){if(closed)return;status('failed');if(port)port.postMessage({t:'close'});close(true)}
async function createSession(first){
try{
status('connecting');
const response=await request('/api/v1/session',()=>options('POST',bootstrap,first));
if(response.status!==200||response.headers.get('X-Carrier-Mode')!=='https')throw new Error('session rejected');
sessionToken=response.headers.get('X-Session-Token')||'';downCursor=response.headers.get('X-Down-Cursor')||'0';
if(!/^[A-Za-z0-9_-]{43}$/.test(sessionToken)||downCursor!=='0')throw new Error('invalid session metadata');
if(closed){deleteSession();return}
const welcome=await response.arrayBuffer();
const welcomeBytes=new Uint8Array(welcome);
if(welcomeBytes.length!==8||welcomeBytes[0]!==17||welcomeBytes.slice(1).some(value=>value!==0))throw new Error('invalid welcome');
port.postMessage(welcome,[welcome]);status('connected');
for(const data of pending.splice(0)){release(data.byteLength,1);queueUp(data)}
poll();
}catch(error){fail()}
}
function queueUp(data){if(!reserve(data)){fail();return}upPending.push(data);runUp()}
async function runUp(){
if(upRunning)return;upRunning=true;
try{
while(!closed&&sessionToken&&upPending.length){
const batch=joinPending(upPending),sequence=String(upSequence);
const response=await request('/api/v1/up',()=>options('POST',sessionToken,batch.body,{'X-Up-Seq':sequence}));
if(response.status!==204||response.headers.get('X-Up-Ack')!==sequence)throw new Error('uplink rejected');
release(batch.total,batch.count);port.postMessage({t:'traffic',up:batch.total,down:0});upSequence++;
}
}catch(error){fail()}
finally{upRunning=false;if(!closed&&sessionToken&&upPending.length)runUp()}
}
async function poll(){
while(!closed&&sessionToken){
try{
pollController=new AbortController();
const response=await request('/api/v1/down',()=>options('POST',sessionToken,null,{'X-Down-Cursor':downCursor},pollController.signal));
if(response.status===204){status('connected');continue}
if(response.status!==200)throw new Error('downlink rejected');
const next=response.headers.get('X-Down-Cursor')||'',data=await response.arrayBuffer();
if(!next||!data.byteLength)throw new Error('invalid downlink response');
if(closed)return;
port.postMessage({t:'traffic',up:0,down:data.byteLength});port.postMessage(data,[data]);downCursor=next;status('connected');
}catch(error){if(!closed)fail();return}
}
}
function deleteSession(){
if(sessionToken)fetch(relayOrigin+'/api/v1/session',options('DELETE',sessionToken,null,null,undefined,true)).catch(()=>{});
}
function close(notifyServer){
if(closed)return;closed=true;if(pollController)pollController.abort();if(notifyServer)deleteSession();
pending.length=0;upPending.length=0;queuedBytes=0;queuedItems=0;if(port)port.close();
}
function activatePort(nextPort){
initialized=true;port=nextPort;
port.onmessage=message=>{
if(message.data instanceof ArrayBuffer){
if(!createStarted){createStarted=true;createSession(message.data)}
else if(!sessionToken){if(!reserve(message.data)){fail();return}pending.push(message.data)}
else queueUp(message.data);
}else if(message.data&&message.data.t==='close')close(true);
};
port.start();status('connecting');
}
addEventListener('message',event=>{
if(initialized||event.source!==parent||event.data===null||typeof event.data!=='object')return;
const keys=Object.keys(event.data).sort();
if(keys.length!==2||keys[0]!=='t'||keys[1]!=='v'||event.data.t!=='tproxy-init'||event.data.v!==1||event.ports.length!==1)return;
let source;try{source=new URL(event.origin)}catch(error){return}
if(source.protocol!=='http:'||source.hostname!=='127.0.0.1'||!source.port||source.origin!==event.origin)return;
activatePort(event.ports[0]);
});
const androidBridge=globalThis.TelegramWebProxy;
if(!initialized&&androidNonce&&androidBridge&&typeof androidBridge.postMessage==='function'){
const androidPort={onmessage:null,start(){},close(){androidBridge.onmessage=null},postMessage(value){
if(value instanceof ArrayBuffer){for(const item of splitFrames(value))androidBridge.postMessage(item)}else androidBridge.postMessage(JSON.stringify(value));
}};
androidBridge.onmessage=event=>{let data=event.data;if(typeof data==='string'){try{data=JSON.parse(data)}catch(error){return}}if(androidPort.onmessage)androidPort.onmessage({data})};
activatePort(androidPort);androidBridge.postMessage(JSON.stringify({t:'tproxy-android-init',v:1,nonce:androidNonce}));
}
addEventListener('pagehide',()=>close(true),{once:true});
})();
</script>
</body>
</html>
"##;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rendered_page_contains_no_template_markers_or_capability() {
let page = render(
"proxy.example.com",
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA",
2 * 1024 * 1024,
32 * 1024 * 1024,
16 * 1024,
&SecureRandom::new(),
);
assert!(!page.body.contains("__"));
assert!(!page.body.contains("bridge="));
assert!(page.body.contains("X-Up-Seq"));
assert!(page
.content_security_policy
.contains("frame-ancestors http://127.0.0.1:*"));
}
}
+272
View File
@@ -0,0 +1,272 @@
use bytes::{BufMut, Bytes, BytesMut};
use crate::config::WebLimitsConfig;
/// Fixed WEB frame header size.
pub(crate) const HEADER_BYTES: usize = 8;
/// Initial bidirectional stream credit.
pub(crate) const INITIAL_STREAM_WINDOW: u32 = 4 * 1024 * 1024;
/// Maximum data chunk emitted by the server.
pub(crate) const DATA_CHUNK_BYTES: usize = 64 * 1024;
/// WEB frame type codes shared with Telegram Desktop.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[repr(u8)]
pub(crate) enum FrameType {
/// Opens a logical MTProxy stream.
Open = 0x01,
/// Carries logical-stream payload bytes.
Data = 0x02,
/// Closes a logical stream.
Close = 0x03,
/// Returns consumed flow-control credit.
Window = 0x04,
/// Requests an application-level liveness response.
Ping = 0x05,
/// Answers application-level liveness traffic.
Pong = 0x06,
/// Starts one WEB carrier session.
Hello = 0x10,
/// Confirms WEB carrier session creation.
Welcome = 0x11,
/// Terminates a WEB carrier session.
Bye = 0x1f,
}
impl FrameType {
fn parse(value: u8) -> Option<Self> {
match value {
0x01 => Some(Self::Open),
0x02 => Some(Self::Data),
0x03 => Some(Self::Close),
0x04 => Some(Self::Window),
0x05 => Some(Self::Ping),
0x06 => Some(Self::Pong),
0x10 => Some(Self::Hello),
0x11 => Some(Self::Welcome),
0x1f => Some(Self::Bye),
_ => None,
}
}
}
/// One parsed frame borrowing its payload from the HTTP request body.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct Frame<'a> {
/// Parsed frame type.
pub(crate) frame_type: FrameType,
/// Logical 24-bit stream identifier.
pub(crate) stream_id: u32,
/// Borrowed frame payload.
pub(crate) payload: &'a [u8],
}
/// Protocol parse or shape failure.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum FrameError {
/// A carrier body contained no frame.
EmptyBatch,
/// A carrier body exceeded the configured frame count.
TooManyFrames,
/// A frame header or payload was truncated.
Incomplete,
/// A frame payload exceeded its configured ceiling.
PayloadLimit,
/// The frame type code is not defined.
UnknownType,
/// A known frame violated direction-specific grammar.
InvalidShape,
}
/// Parses and validates all frame boundaries without copying payloads.
pub(crate) fn parse_all<'a>(
input: &'a [u8],
limits: &WebLimitsConfig,
) -> std::result::Result<Vec<Frame<'a>>, FrameError> {
if input.is_empty() {
return Err(FrameError::EmptyBatch);
}
let mut remaining = input;
let mut frames = Vec::with_capacity(remaining.len().div_ceil(HEADER_BYTES).min(16));
while !remaining.is_empty() {
if frames.len() >= limits.max_frames_per_body {
return Err(FrameError::TooManyFrames);
}
if remaining.len() < HEADER_BYTES {
return Err(FrameError::Incomplete);
}
let frame_type = FrameType::parse(remaining[0]).ok_or(FrameError::UnknownType)?;
let stream_id = u32::from(remaining[1]) << 16
| u32::from(remaining[2]) << 8
| u32::from(remaining[3]);
let payload_len = u32::from_be_bytes([
remaining[4],
remaining[5],
remaining[6],
remaining[7],
]) as usize;
if payload_len > limits.max_frame_payload_bytes {
return Err(FrameError::PayloadLimit);
}
let frame_len = HEADER_BYTES
.checked_add(payload_len)
.ok_or(FrameError::PayloadLimit)?;
if frame_len > remaining.len() {
return Err(FrameError::Incomplete);
}
frames.push(Frame {
frame_type,
stream_id,
payload: &remaining[HEADER_BYTES..frame_len],
});
remaining = &remaining[frame_len..];
}
Ok(frames)
}
/// Enforces the client-to-server frame grammar.
pub(crate) fn validate_client_shape(frame: Frame<'_>) -> std::result::Result<(), FrameError> {
if frame.stream_id == 0 {
return if frame.frame_type == FrameType::Pong && frame.payload.len() <= 64 {
Ok(())
} else {
Err(FrameError::InvalidShape)
};
}
match frame.frame_type {
FrameType::Open | FrameType::Close if frame.payload.is_empty() => Ok(()),
FrameType::Data if !frame.payload.is_empty() => Ok(()),
FrameType::Window => window_amount(frame.payload).map(|_| ()),
_ => Err(FrameError::InvalidShape),
}
}
/// Validates the exact first-session HELLO body.
pub(crate) fn validate_hello(input: &[u8], limits: &WebLimitsConfig) -> bool {
let Ok(frames) = parse_all(input, limits) else {
return false;
};
frames.len() == 1
&& frames[0].frame_type == FrameType::Hello
&& frames[0].stream_id == 0
&& frames[0].payload == [1]
}
/// Encodes one complete WEB frame.
pub(crate) fn encode(frame_type: FrameType, stream_id: u32, payload: &[u8]) -> Bytes {
let mut output = BytesMut::with_capacity(HEADER_BYTES + payload.len());
output.put_u8(frame_type as u8);
output.put_u8((stream_id >> 16) as u8);
output.put_u8((stream_id >> 8) as u8);
output.put_u8(stream_id as u8);
output.put_u32(payload.len() as u32);
output.extend_from_slice(payload);
output.freeze()
}
/// Decodes a non-zero WINDOW delta.
pub(crate) fn window_amount(payload: &[u8]) -> std::result::Result<u32, FrameError> {
let bytes: [u8; 4] = payload.try_into().map_err(|_| FrameError::InvalidShape)?;
let amount = u32::from_be_bytes(bytes);
(amount != 0)
.then_some(amount)
.ok_or(FrameError::InvalidShape)
}
/// Encodes a WINDOW delta payload.
pub(crate) fn window_payload(amount: u32) -> [u8; 4] {
amount.to_be_bytes()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hello_and_welcome_match_reference_bytes() {
let limits = WebLimitsConfig::default();
let hello = encode(FrameType::Hello, 0, &[1]);
assert_eq!(hello.as_ref(), &hex::decode("100000000000000101").unwrap());
assert!(validate_hello(&hello, &limits));
assert_eq!(
encode(FrameType::Welcome, 0, &[]).as_ref(),
&hex::decode("1100000000000000").unwrap()
);
}
#[test]
fn parser_rejects_excessive_payload_before_slicing() {
let limits = WebLimitsConfig {
max_frame_payload_bytes: 4,
..WebLimitsConfig::default()
};
let frame = encode(FrameType::Data, 1, &[0; 5]);
assert_eq!(parse_all(&frame, &limits), Err(FrameError::PayloadLimit));
}
#[test]
fn client_shape_rejects_control_types_on_stream_zero() {
let frame = Frame {
frame_type: FrameType::Ping,
stream_id: 0,
payload: &[],
};
assert_eq!(validate_client_shape(frame), Err(FrameError::InvalidShape));
}
#[test]
fn stream_frames_match_client_reference_vectors() {
assert_eq!(
encode(FrameType::Open, 17, &[]).as_ref(),
&hex::decode("0100001100000000").unwrap()
);
assert_eq!(
encode(FrameType::Data, 17, b"round trip").as_ref(),
&hex::decode("020000110000000a726f756e642074726970").unwrap()
);
assert_eq!(
encode(FrameType::Window, 17, &10u32.to_be_bytes()).as_ref(),
&hex::decode("04000011000000040000000a").unwrap()
);
assert_eq!(
encode(FrameType::Open, 0x00ff_ffff, &[]).as_ref(),
&hex::decode("01ffffff00000000").unwrap()
);
}
#[test]
fn parser_rejects_empty_truncated_and_excessive_batches() {
let mut limits = WebLimitsConfig::default();
assert_eq!(parse_all(&[], &limits), Err(FrameError::EmptyBatch));
assert_eq!(
parse_all(&hex::decode("0200000100000001").unwrap(), &limits),
Err(FrameError::Incomplete)
);
limits.max_frames_per_body = 1;
let mut body = encode(FrameType::Pong, 0, &[]).to_vec();
body.extend_from_slice(&encode(FrameType::Pong, 0, &[]));
assert_eq!(parse_all(&body, &limits), Err(FrameError::TooManyFrames));
}
#[test]
fn client_shape_rejects_empty_data_and_zero_window() {
let empty_data = Frame {
frame_type: FrameType::Data,
stream_id: 1,
payload: &[],
};
let zero_window = Frame {
frame_type: FrameType::Window,
stream_id: 1,
payload: &[0; 4],
};
assert_eq!(
validate_client_shape(empty_data),
Err(FrameError::InvalidShape)
);
assert_eq!(
validate_client_shape(zero_window),
Err(FrameError::InvalidShape)
);
}
}
+511
View File
@@ -0,0 +1,511 @@
use std::convert::Infallible;
use std::error::Error;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, Full};
use hyper::body::Incoming;
use hyper::header::{self, HeaderName, HeaderValue};
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Method, Request, Response, StatusCode};
use hyper_util::rt::{TokioIo, TokioTimer};
use ipnetwork::IpNetwork;
use parking_lot::Mutex;
use tokio::net::TcpStream;
use tokio_util::sync::CancellationToken;
use crate::config::{WebClientIpSource, WebRuntimeVhost};
use crate::web::bridge;
use crate::web::frame::{self, FrameType};
use crate::web::manager::{ManagerError, WebProcessRuntime};
// Response-body activity keeps connection idle accounting lifecycle-correct.
mod activity;
// Body collection retains allocation permits through request processing.
mod body;
// Decoy routing and upstream proxying are isolated from carrier authentication.
mod decoy;
// Canonical request parsing rejects ambiguous credentials before routing.
mod request;
#[cfg(test)]
mod tests;
use decoy::serve_decoy;
use activity::{ActivityBody, RequestActivity};
use body::{CollectBodyError, CollectedBody, collect_body};
use request::{
bearer_token_hash, binary_content_type, bridge_candidate, canonical_request_host,
canonical_u64_header, client_ip, match_profile,
};
type BoxError = Box<dyn Error + Send + Sync>;
type HttpBody = UnsyncBoxBody<Bytes, BoxError>;
type HttpResponse = Response<HttpBody>;
const CREATE_BODY_LIMIT: usize = 64;
const TRANSPORT_PATHS: [&str; 3] = ["/api/v1/session", "/api/v1/up", "/api/v1/down"];
/// Serves one bounded HTTP/1.1 connection accepted from an external TLS terminator.
pub(crate) async fn serve_connection(
stream: TcpStream,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: Arc<[IpNetwork]>,
runtime: Arc<WebProcessRuntime>,
cancellation: CancellationToken,
connection_permit: tokio::sync::OwnedSemaphorePermit,
) {
let config = runtime.active_generation().config();
let max_header_bytes = config.web.limits.max_header_bytes;
let header_timeout = Duration::from_secs(config.web.timeouts.header_secs);
let idle_timeout = Duration::from_secs(config.web.timeouts.http_idle_secs);
let last_activity = Arc::new(Mutex::new(Instant::now()));
let service_last_activity = Arc::clone(&last_activity);
let service = service_fn(move |request| {
let runtime = Arc::clone(&runtime);
let trusted_proxy_cidrs = Arc::clone(&trusted_proxy_cidrs);
let last_activity = Arc::clone(&service_last_activity);
let client_ip_source = client_ip_source;
async move {
let activity = RequestActivity::begin(last_activity);
let response = if let Some(_handler_permit) = runtime.try_http_handler() {
handle_request(
request,
peer,
client_ip_source,
&trusted_proxy_cidrs,
runtime,
)
.await
} else {
service_unavailable()
};
let response = response.map(|body| {
ActivityBody::new(body, activity)
.boxed_unsync()
});
Ok::<_, Infallible>(response)
}
});
let connection = http1::Builder::new()
.timer(TokioTimer::new())
.header_read_timeout(header_timeout)
.max_buf_size(max_header_bytes)
.keep_alive(true)
.serve_connection(TokioIo::new(stream), service);
tokio::pin!(connection);
let mut idle_check = tokio::time::interval((idle_timeout / 2).max(Duration::from_secs(1)));
idle_check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
biased;
_ = cancellation.cancelled() => break,
_ = &mut connection => break,
_ = idle_check.tick() => {
if Instant::now().saturating_duration_since(*last_activity.lock())
>= idle_timeout
{
break;
}
}
}
}
drop(connection_permit);
}
async fn handle_request(
request: Request<Incoming>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
) -> HttpResponse {
let generation = runtime.active_generation();
let config = generation.config();
let Some(web_runtime) = config.web.runtime.as_ref() else {
return generic_not_found();
};
let Some(host) = canonical_request_host(&request) else {
return generic_not_found();
};
let Some(vhost) = web_runtime.vhosts.get(host).cloned() else {
return generic_not_found();
};
let path = request.uri().path();
if TRANSPORT_PATHS.contains(&path) {
return handle_api(
request,
peer,
client_ip_source,
trusted_proxy_cidrs,
runtime,
vhost,
)
.await;
}
if path == "/" && matches!(*request.method(), Method::GET | Method::HEAD) {
return handle_root(
request,
peer,
client_ip_source,
trusted_proxy_cidrs,
runtime,
vhost,
)
.await;
}
serve_decoy(request, vhost, false, &runtime).await
}
async fn handle_root(
mut request: Request<Incoming>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
) -> HttpResponse {
let (candidate, canonical) = bridge_candidate(request.uri().query());
let profile = match_profile(&vhost, &candidate);
let Some(profile) = profile.filter(|_| canonical && request.method() == Method::GET) else {
return serve_decoy(request, vhost, false, &runtime).await;
};
let Some(client_ip) = client_ip(
&request,
peer,
client_ip_source,
trusted_proxy_cidrs,
) else {
strip_query(&mut request);
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(bootstrap) = runtime.issue_bootstrap(profile, client_ip) else {
strip_query(&mut request);
return serve_decoy(request, vhost, true, &runtime).await;
};
let generation = runtime.active_generation();
let config = generation.config();
let page = bridge::render(
&vhost.host,
&bootstrap,
config.web.limits.carrier_batch_bytes,
config.web.limits.pending_bytes_per_session,
config.web.limits.pending_items_per_session,
&generation.rng,
);
let mut response = full_response(StatusCode::OK, Bytes::from(page.body));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("text/html; charset=utf-8"),
);
insert_header(
&mut response,
header::CONTENT_SECURITY_POLICY,
&page.content_security_policy,
);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
response.headers_mut().insert(
header::REFERRER_POLICY,
HeaderValue::from_static("no-referrer"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
response.headers_mut().insert(
HeaderName::from_static("x-dns-prefetch-control"),
HeaderValue::from_static("off"),
);
insert_header(
&mut response,
HeaderName::from_static("permissions-policy"),
bridge::PERMISSIONS_POLICY,
);
response
}
async fn handle_api(
request: Request<Incoming>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
) -> HttpResponse {
if request.uri().query().is_some()
|| request.headers().contains_key(header::COOKIE)
|| request.headers().contains_key("x-lane-id")
{
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(client_ip) = client_ip(
&request,
peer,
client_ip_source,
trusted_proxy_cidrs,
) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Some(token_hash) = bearer_token_hash(&request) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
match request.uri().path() {
"/api/v1/session" => {
handle_session(request, runtime, vhost, token_hash, client_ip).await
}
"/api/v1/up" => handle_up(request, runtime, vhost, token_hash).await,
"/api/v1/down" => handle_down(request, runtime, vhost, token_hash).await,
_ => serve_decoy(request, vhost, true, &runtime).await,
}
}
async fn handle_session(
request: Request<Incoming>,
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
token_hash: crate::web::manager::TokenHash,
client_ip: IpAddr,
) -> HttpResponse {
if request.method() == Method::DELETE {
if request.headers().contains_key(header::CONTENT_TYPE) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, 1, true).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
if !body.is_empty() || runtime.close_token(token_hash, &vhost.host).is_err() {
return serve_decoy(request, vhost, true, &runtime).await;
}
return carrier_empty(StatusCode::NO_CONTENT);
}
if request.method() != Method::POST || !binary_content_type(&request) {
return serve_decoy(request, vhost, true, &runtime).await;
}
if !runtime.has_bootstrap(token_hash, &vhost.host) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, CREATE_BODY_LIMIT, false).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
match runtime.create_session(token_hash, &vhost.host, client_ip, &body) {
Ok(result) => {
let welcome = frame::encode(FrameType::Welcome, 0, &[]);
let mut response = full_response(StatusCode::OK, welcome);
carrier_headers(&mut response);
insert_header(
&mut response,
HeaderName::from_static("x-session-token"),
&result.token,
);
response.headers_mut().insert(
HeaderName::from_static("x-carrier-mode"),
HeaderValue::from_static("https"),
);
response.headers_mut().insert(
HeaderName::from_static("x-down-cursor"),
HeaderValue::from_static("0"),
);
response
}
Err(ManagerError::Limit | ManagerError::Backpressure | ManagerError::Concurrent) => {
service_unavailable()
}
Err(_) => serve_decoy(request, vhost, true, &runtime).await,
}
}
async fn handle_up(
request: Request<Incoming>,
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
token_hash: crate::web::manager::TokenHash,
) -> HttpResponse {
if request.method() != Method::POST || !binary_content_type(&request) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(sequence) = canonical_u64_header(&request, "x-up-seq").filter(|value| *value != 0)
else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let limit = runtime.active_generation().config().web.limits.max_body_bytes;
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, limit, false).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
match session.process_up(sequence, &body) {
Ok(ack) => {
let mut response = carrier_empty(StatusCode::NO_CONTENT);
insert_header(
&mut response,
HeaderName::from_static("x-up-ack"),
&ack.to_string(),
);
response
}
Err(ManagerError::Backpressure | ManagerError::Concurrent | ManagerError::Limit) => {
service_unavailable()
}
Err(_) => serve_decoy(request, vhost, true, &runtime).await,
}
}
async fn handle_down(
request: Request<Incoming>,
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
token_hash: crate::web::manager::TokenHash,
) -> HttpResponse {
if request.method() != Method::POST || request.headers().contains_key(header::CONTENT_TYPE) {
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(cursor) = canonical_u64_header(&request, "x-down-cursor") else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let CollectedBody {
request,
body,
_body_budget,
} = match collect_body(request, &runtime, 1, true).await {
Ok(result) => result,
Err(CollectBodyError::Limit) => return service_unavailable(),
Err(CollectBodyError::Invalid(request)) => {
return serve_decoy(request, vhost, true, &runtime).await;
}
};
if !body.is_empty() {
return serve_decoy(request, vhost, true, &runtime).await;
}
match session.poll_down(cursor).await {
Ok(result) if result.body.is_empty() => {
let mut response = carrier_empty(StatusCode::NO_CONTENT);
insert_header(
&mut response,
HeaderName::from_static("x-down-cursor"),
&result.next_cursor.to_string(),
);
response
}
Ok(result) => {
let mut response = full_response(StatusCode::OK, result.body);
carrier_headers(&mut response);
insert_header(
&mut response,
HeaderName::from_static("x-down-cursor"),
&result.next_cursor.to_string(),
);
response
}
Err(ManagerError::Concurrent | ManagerError::Backpressure | ManagerError::Limit) => {
service_unavailable()
}
Err(_) => serve_decoy(request, vhost, true, &runtime).await,
}
}
fn carrier_headers(response: &mut HttpResponse) {
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/octet-stream"),
);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
}
fn carrier_empty(status: StatusCode) -> HttpResponse {
let mut response = empty_response(status);
response
.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
response
}
fn service_unavailable() -> HttpResponse {
let mut response = carrier_empty(StatusCode::SERVICE_UNAVAILABLE);
response
.headers_mut()
.insert(header::RETRY_AFTER, HeaderValue::from_static("1"));
response
}
fn bad_gateway() -> HttpResponse {
full_response(
StatusCode::BAD_GATEWAY,
Bytes::from_static(b"site unavailable\n"),
)
}
fn generic_not_found() -> HttpResponse {
full_response(
StatusCode::NOT_FOUND,
Bytes::from_static(b"not found\n"),
)
}
fn full_response(status: StatusCode, body: Bytes) -> HttpResponse {
let length = body.len();
let body = Full::new(body)
.map_err(|never| -> BoxError { match never {} })
.boxed_unsync();
let mut response = Response::new(body);
*response.status_mut() = status;
insert_header(
&mut response,
header::CONTENT_LENGTH,
&length.to_string(),
);
response
}
fn empty_response(status: StatusCode) -> HttpResponse {
full_response(status, Bytes::new())
}
fn insert_header(response: &mut HttpResponse, name: HeaderName, value: &str) {
if let Ok(value) = HeaderValue::from_str(value) {
response.headers_mut().insert(name, value);
}
}
fn strip_query<B>(request: &mut Request<B>) {
if request.uri().query().is_some()
&& let Ok(uri) = request.uri().path().parse()
{
*request.uri_mut() = uri;
}
}
+69
View File
@@ -0,0 +1,69 @@
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Instant;
use bytes::Bytes;
use hyper::body::{Body, Frame, SizeHint};
use parking_lot::Mutex;
use super::{BoxError, HttpBody};
/// Request lifecycle guard that refreshes HTTP connection activity on completion.
pub(super) struct RequestActivity {
last_activity: Arc<Mutex<Instant>>,
}
impl RequestActivity {
/// Starts activity accounting for one HTTP request.
pub(super) fn begin(last_activity: Arc<Mutex<Instant>>) -> Self {
*last_activity.lock() = Instant::now();
Self { last_activity }
}
}
impl Drop for RequestActivity {
fn drop(&mut self) {
*self.last_activity.lock() = Instant::now();
}
}
/// Response body wrapper that refreshes activity while downstream data progresses.
pub(super) struct ActivityBody {
inner: HttpBody,
activity: RequestActivity,
}
impl ActivityBody {
/// Binds one response body to its request activity guard.
pub(super) fn new(inner: HttpBody, activity: RequestActivity) -> Self {
Self {
inner,
activity,
}
}
}
impl Body for ActivityBody {
type Data = Bytes;
type Error = BoxError;
fn poll_frame(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
let result = Pin::new(&mut self.inner).poll_frame(context);
if result.is_ready() {
*self.activity.last_activity.lock() = Instant::now();
}
result
}
fn is_end_stream(&self) -> bool {
self.inner.is_end_stream()
}
fn size_hint(&self) -> SizeHint {
self.inner.size_hint()
}
}
+80
View File
@@ -0,0 +1,80 @@
use std::time::Duration;
use bytes::Bytes;
use http_body_util::{BodyExt, Empty, Limited};
use hyper::body::{Body as _, Incoming};
use hyper::Request;
use crate::web::manager::WebProcessRuntime;
/// Collected carrier request retaining its process-wide body reservation.
pub(super) struct CollectedBody {
/// Request head reconstructed without the consumed network body.
pub(super) request: Request<Empty<Bytes>>,
/// Fully collected bounded carrier payload.
pub(super) body: Bytes,
/// Byte-budget reservation held through request processing.
pub(super) _body_budget: tokio::sync::OwnedSemaphorePermit,
}
// Keep rejected requests inline to avoid attacker-controlled allocations on invalid bodies.
#[allow(clippy::large_enum_variant)]
/// Body collection failure with sanitized request context when decoy routing is safe.
pub(super) enum CollectBodyError {
/// The body shape, size, or deadline failed after retaining the request head.
Invalid(Request<Empty<Bytes>>),
/// Process-wide body reader or byte capacity is temporarily exhausted.
Limit,
}
/// Collects one bounded carrier body under reader, byte, and deadline ownership.
pub(super) async fn collect_body(
request: Request<Incoming>,
runtime: &WebProcessRuntime,
limit: usize,
allow_empty: bool,
) -> Result<CollectedBody, CollectBodyError> {
let exceeds_limit = request.body().size_hint().lower() > limit as u64
|| request
.body()
.size_hint()
.upper()
.is_some_and(|upper| upper > limit as u64);
let (parts, body) = request.into_parts();
if exceeds_limit {
return Err(CollectBodyError::Invalid(Request::from_parts(
parts,
Empty::new(),
)));
}
let Some((reader_budget, body_budget)) = runtime.try_body_budget(limit) else {
return Err(CollectBodyError::Limit);
};
let body_timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.body_secs,
);
let body = match tokio::time::timeout(body_timeout, Limited::new(body, limit).collect()).await {
Ok(Ok(body)) => body.to_bytes(),
_ => {
return Err(CollectBodyError::Invalid(Request::from_parts(
parts,
Empty::new(),
)));
}
};
drop(reader_budget);
let request = Request::from_parts(parts, Empty::new());
if !allow_empty && body.is_empty() {
return Err(CollectBodyError::Invalid(request));
}
Ok(CollectedBody {
request,
body,
_body_budget: body_budget,
})
}
+284
View File
@@ -0,0 +1,284 @@
use std::error::Error;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use http_body_util::{BodyExt, Empty};
use hyper::header::{self, HeaderName, HeaderValue};
use hyper::{Method, Request, StatusCode, Uri};
use hyper_util::rt::TokioIo;
use tokio::net::TcpStream;
use super::{
BoxError, HttpBody, HttpResponse, bad_gateway, full_response, generic_not_found,
insert_header,
};
use crate::config::{WebRuntimeDecoy, WebRuntimeVhost};
use crate::web::manager::WebProcessRuntime;
/// Serves the configured ordinary site after optionally removing carrier material.
pub(super) async fn serve_decoy<B>(
mut request: Request<B>,
vhost: Arc<WebRuntimeVhost>,
sanitize_transport: bool,
runtime: &WebProcessRuntime,
) -> HttpResponse
where
B: hyper::body::Body<Data = Bytes> + Send + 'static,
B::Error: Error + Send + Sync + 'static,
{
if sanitize_transport {
sanitize_transport_request(&mut request);
}
let (parts, body) = request.into_parts();
let body = if sanitize_transport {
Empty::<Bytes>::new()
.map_err(|never| -> BoxError { match never {} })
.boxed_unsync()
} else {
body.map_err(|error| -> BoxError { Box::new(error) })
.boxed_unsync()
};
let request = Request::from_parts(parts, body);
match &vhost.decoy {
WebRuntimeDecoy::StaticDirectory(site) => serve_static(request, site),
WebRuntimeDecoy::HttpUpstream { addr, authority } => {
proxy_to_upstream(
request,
*addr,
authority,
Duration::from_secs(vhost.decoy_header_secs),
runtime,
)
.await
}
}
}
fn serve_static<B>(request: Request<B>, site: &crate::config::WebStaticSite) -> HttpResponse {
if !matches!(*request.method(), Method::GET | Method::HEAD) {
return static_entry(request, site, None, StatusCode::NOT_FOUND);
}
let path = request.uri().path();
let resolved = resolve_static_path(path, site);
let status = if resolved.is_some() {
StatusCode::OK
} else {
StatusCode::NOT_FOUND
};
static_entry(request, site, resolved, status)
}
fn static_entry<B>(
request: Request<B>,
site: &crate::config::WebStaticSite,
route: Option<&str>,
status: StatusCode,
) -> HttpResponse {
let fallback = format!("/{}", site.index);
let not_found = site.assets.contains_key("/404.html").then_some("/404.html");
let route = route.or(not_found).unwrap_or(&fallback);
let Some(asset) = site.assets.get(route) else {
return generic_not_found();
};
let not_modified = status == StatusCode::OK
&& request
.headers()
.get(header::IF_NONE_MATCH)
.and_then(|value| value.to_str().ok())
== Some(asset.etag.as_str());
let response_body = if request.method() == Method::HEAD || not_modified {
Bytes::new()
} else {
asset.body.clone()
};
let mut response = full_response(
if not_modified {
StatusCode::NOT_MODIFIED
} else {
status
},
response_body,
);
insert_header(&mut response, header::CONTENT_TYPE, asset.content_type);
insert_header(&mut response, header::ETAG, &asset.etag);
insert_header(
&mut response,
header::CONTENT_LENGTH,
&asset.body.len().to_string(),
);
response.headers_mut().insert(
header::CACHE_CONTROL,
if status.is_client_error() || request.uri().query().is_some() {
HeaderValue::from_static("no-store")
} else {
HeaderValue::from_static("public, max-age=300")
},
);
response.headers_mut().insert(
header::CONTENT_SECURITY_POLICY,
HeaderValue::from_static("default-src 'self'; style-src 'self'; img-src 'self'; worker-src 'none'; frame-ancestors 'none'; base-uri 'none'; form-action 'none'"),
);
response.headers_mut().insert(
header::REFERRER_POLICY,
HeaderValue::from_static("strict-origin-when-cross-origin"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
response.headers_mut().insert(header::X_FRAME_OPTIONS, HeaderValue::from_static("DENY"));
response
}
fn resolve_static_path<'a>(
path: &str,
site: &'a crate::config::WebStaticSite,
) -> Option<&'a str> {
if !path.starts_with('/')
|| path.contains('\\')
|| path.contains("//")
|| path.split('/').any(|part| matches!(part, "." | ".."))
{
return None;
}
let root;
let route = if path == "/" {
root = format!("/{}", site.index);
root.as_str()
} else {
path
};
if site.assets.contains_key(route) {
return site.assets.get_key_value(route).map(|(key, _)| key.as_str());
}
if route == "/favicon.ico" && site.assets.contains_key("/favicon.svg") {
return Some("/favicon.svg");
}
if !route.rsplit('/').next().unwrap_or_default().contains('.') {
let html = format!("{route}.html");
return site
.assets
.get_key_value(&html)
.map(|(key, _)| key.as_str());
}
None
}
async fn proxy_to_upstream(
mut request: Request<HttpBody>,
addr: SocketAddr,
authority: &str,
header_timeout: Duration,
runtime: &WebProcessRuntime,
) -> HttpResponse {
remove_hop_by_hop(request.headers_mut());
if let Ok(host) = HeaderValue::from_str(authority) {
request.headers_mut().insert(header::HOST, host);
}
let path_and_query = request
.uri()
.path_and_query()
.map(|value| value.as_str())
.unwrap_or("/");
let Ok(uri) = path_and_query.parse::<Uri>() else {
return bad_gateway();
};
*request.uri_mut() = uri;
let stream = match tokio::time::timeout(header_timeout, TcpStream::connect(addr)).await {
Ok(Ok(stream)) => stream,
_ => return bad_gateway(),
};
let max_header_bytes = runtime
.active_generation()
.config()
.web
.limits
.max_header_bytes;
let mut builder = hyper::client::conn::http1::Builder::new();
builder.max_buf_size(max_header_bytes);
let (mut sender, connection) = match tokio::time::timeout(
header_timeout,
builder.handshake(TokioIo::new(stream)),
)
.await
{
Ok(Ok(parts)) => parts,
_ => return bad_gateway(),
};
runtime.spawn_auxiliary(async move {
let _ = connection.await;
});
let mut response = match tokio::time::timeout(header_timeout, sender.send_request(request)).await
{
Ok(Ok(response)) => response,
_ => return bad_gateway(),
};
remove_hop_by_hop(response.headers_mut());
response.map(|body| {
body.map_err(|error| -> BoxError { Box::new(error) })
.boxed_unsync()
})
}
fn sanitize_transport_request<B>(request: &mut Request<B>) {
for name in [
header::AUTHORIZATION,
header::CONTENT_LENGTH,
header::CONTENT_TYPE,
header::UPGRADE,
HeaderName::from_static("sec-websocket-key"),
HeaderName::from_static("sec-websocket-protocol"),
HeaderName::from_static("sec-websocket-version"),
HeaderName::from_static("x-down-cursor"),
HeaderName::from_static("x-lane-id"),
HeaderName::from_static("x-up-seq"),
] {
request.headers_mut().remove(name);
}
request
.headers_mut()
.insert(header::CONNECTION, HeaderValue::from_static("close"));
}
fn remove_hop_by_hop(headers: &mut hyper::HeaderMap) {
let nominated = headers
.get_all(header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.filter_map(|value| HeaderName::from_bytes(value.trim().as_bytes()).ok())
.collect::<Vec<_>>();
for name in nominated {
headers.remove(name);
}
for name in [
header::CONNECTION,
header::PROXY_AUTHENTICATE,
header::PROXY_AUTHORIZATION,
header::TE,
header::TRAILER,
header::TRANSFER_ENCODING,
header::UPGRADE,
HeaderName::from_static("keep-alive"),
HeaderName::from_static("proxy-connection"),
] {
headers.remove(name);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn static_resolver_rejects_rewritten_paths() {
let site = crate::config::WebStaticSite {
assets: std::collections::BTreeMap::new(),
index: "index.html".to_string(),
};
assert!(resolve_static_path("/../index.html", &site).is_none());
assert!(resolve_static_path("//index.html", &site).is_none());
}
}
+240
View File
@@ -0,0 +1,240 @@
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use base64::Engine as _;
use hyper::header;
use hyper::Request;
use ipnetwork::IpNetwork;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use crate::config::{
WebClientIpSource, WebRuntimeProfile, WebRuntimeVhost,
};
use crate::web::manager::TokenHash;
/// Parses one lowercase canonical Host value restricted to the public HTTPS port.
pub(super) fn canonical_request_host<B>(request: &Request<B>) -> Option<&str> {
let values = request.headers().get_all(header::HOST);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some() {
return None;
}
let authority = value.parse::<hyper::http::uri::Authority>().ok()?;
if authority.port_u16().is_some_and(|port| port != 443) {
return None;
}
let host = value.strip_suffix(":443").unwrap_or(value);
if authority.host() != host || host.bytes().any(|byte| byte.is_ascii_uppercase())
{
return None;
}
Some(host)
}
/// Accepts one canonical forwarded client address from an explicitly trusted peer.
pub(super) fn client_ip<B>(
request: &Request<B>,
peer: SocketAddr,
source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
) -> Option<IpAddr> {
if !trusted_proxy_cidrs
.iter()
.any(|network| network.contains(peer.ip()))
{
return None;
}
let header_name = match source {
WebClientIpSource::XForwardedFor => "x-forwarded-for",
};
let values = request.headers().get_all(header_name);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some()
|| value.is_empty()
|| value.trim() != value
|| value.contains(',')
{
return None;
}
let ip = value.parse::<IpAddr>().ok()?;
(ip.to_string() == value).then_some(ip)
}
/// Decodes an exact canonical bridge query without allocating credential strings.
pub(super) fn bridge_candidate(query: Option<&str>) -> ([u8; 32], bool) {
let mut candidate = [0u8; 32];
let Some(value) = query.and_then(|query| query.strip_prefix("bridge=")) else {
return (candidate, false);
};
if value.len() != 43 {
return (candidate, false);
}
let mut decoded = [0u8; 32];
let Ok(decoded_len) = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode_slice(value, &mut decoded)
else {
return (candidate, false);
};
let mut canonical = [0u8; 43];
let Ok(encoded_len) = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode_slice(decoded, &mut canonical)
else {
return (candidate, false);
};
if decoded_len != decoded.len()
|| encoded_len != canonical.len()
|| !bool::from(canonical.ct_eq(value.as_bytes()))
{
return (candidate, false);
}
candidate = decoded;
(candidate, true)
}
/// Matches a capability in constant time across every profile of one virtual host.
pub(super) fn match_profile(
vhost: &WebRuntimeVhost,
candidate: &[u8; 32],
) -> Option<Arc<WebRuntimeProfile>> {
let mut matched = None;
for profile in &vhost.profiles {
if bool::from(profile.capability.ct_eq(candidate)) {
matched = Some(Arc::clone(profile));
}
}
matched
}
/// Validates and hashes one canonical bearer credential for map lookup.
pub(super) fn bearer_token_hash<B>(request: &Request<B>) -> Option<TokenHash> {
let values = request.headers().get_all(header::AUTHORIZATION);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some() || !value.starts_with("Bearer ") || value.matches(' ').count() != 1
{
return None;
}
let token = value.strip_prefix("Bearer ")?;
if token.len() != 43 {
return None;
}
let mut decoded = [0u8; 32];
let decoded_len = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode_slice(token, &mut decoded)
.ok()?;
let mut canonical = [0u8; 43];
let encoded_len = base64::engine::general_purpose::URL_SAFE_NO_PAD
.encode_slice(decoded, &mut canonical)
.ok()?;
(decoded_len == decoded.len()
&& encoded_len == canonical.len()
&& bool::from(canonical.ct_eq(token.as_bytes())))
.then(|| Sha256::digest(decoded).into())
}
/// Checks the exact carrier media type without accepting duplicate headers.
pub(super) fn binary_content_type<B>(request: &Request<B>) -> bool {
let values = request.headers().get_all(header::CONTENT_TYPE);
let mut values = values.iter();
let value = values.next().and_then(|value| value.to_str().ok());
values.next().is_none()
&& value.is_some_and(|value| value.eq_ignore_ascii_case("application/octet-stream"))
}
/// Parses one canonical unsigned decimal carrier sequence header.
pub(super) fn canonical_u64_header<B>(
request: &Request<B>,
name: &'static str,
) -> Option<u64> {
let values = request.headers().get_all(name);
let mut values = values.iter();
let value = values.next()?.to_str().ok()?;
if values.next().is_some()
|| value.is_empty()
|| value.starts_with('+')
|| (value.len() > 1 && value.starts_with('0'))
{
return None;
}
let parsed = value.parse::<u64>().ok()?;
(parsed.to_string() == value).then_some(parsed)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn canonical_bridge_query_rejects_aliases() {
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode([7u8; 32]);
assert!(bridge_candidate(Some(&format!("bridge={token}"))).1);
assert!(!bridge_candidate(Some(&format!("x=1&bridge={token}"))).1);
assert!(!bridge_candidate(Some(&format!("bridge={token}="))).1);
}
#[test]
fn host_and_forwarded_identity_require_canonical_single_values() {
let request = Request::builder()
.header(header::HOST, "proxy.example.com:443")
.header("x-forwarded-for", "192.0.2.10")
.body(())
.unwrap();
assert_eq!(
canonical_request_host(&request),
Some("proxy.example.com")
);
let trusted: [IpNetwork; 1] = ["127.0.0.1/32".parse().unwrap()];
assert_eq!(
client_ip(
&request,
"127.0.0.1:40000".parse().unwrap(),
WebClientIpSource::XForwardedFor,
&trusted,
),
Some("192.0.2.10".parse().unwrap())
);
let uppercase = Request::builder()
.header(header::HOST, "Proxy.Example.com")
.body(())
.unwrap();
assert!(canonical_request_host(&uppercase).is_none());
let appended = Request::builder()
.header("x-forwarded-for", "192.0.2.10, 198.51.100.4")
.body(())
.unwrap();
assert!(
client_ip(
&appended,
"127.0.0.1:40000".parse().unwrap(),
WebClientIpSource::XForwardedFor,
&trusted,
)
.is_none()
);
}
#[test]
fn bearer_and_sequence_headers_reject_noncanonical_aliases() {
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode([1u8; 32]);
let request = Request::builder()
.header(header::AUTHORIZATION, format!("Bearer {token}"))
.header("x-up-seq", "17")
.body(())
.unwrap();
assert_eq!(
bearer_token_hash(&request),
Some(Sha256::digest([1u8; 32]).into())
);
assert_eq!(canonical_u64_header(&request, "x-up-seq"), Some(17));
let leading_zero = Request::builder()
.header("x-up-seq", "017")
.body(())
.unwrap();
assert!(canonical_u64_header(&leading_zero, "x-up-seq").is_none());
}
}
+206
View File
@@ -0,0 +1,206 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use arc_swap::ArcSwap;
use base64::Engine as _;
use bytes::Bytes;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio_util::sync::CancellationToken;
use super::serve_connection;
use crate::config::{
ProxyConfig, WebClientIpSource, WebRuntimeConfig, WebRuntimeDecoy,
WebRuntimeProfile, WebRuntimeVhost, WebSecretMode, WebStaticAsset, WebStaticSite,
};
use crate::maestro::generation::test_runtime_generation;
use crate::web::frame::{self, FrameType};
use crate::web::manager::WebProcessRuntime;
fn runtime_config(capability: [u8; 32]) -> ProxyConfig {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: "203.0.113.10:443".parse().unwrap(),
user: "alice".to_string(),
secret_mode: WebSecretMode::Plain,
capability,
max_sessions: 4,
max_streams: 16,
max_streams_per_session: 4,
});
let mut assets = BTreeMap::new();
assets.insert(
"/index.html".to_string(),
WebStaticAsset {
body: Bytes::from_static(b"<!doctype html><title>decoy</title>"),
content_type: "text/html; charset=utf-8",
etag: "\"test\"".to_string(),
},
);
let site = Arc::new(WebStaticSite {
assets,
index: "index.html".to_string(),
});
let vhost = Arc::new(WebRuntimeVhost {
host: "proxy.example.com".to_string(),
decoy: WebRuntimeDecoy::StaticDirectory(Arc::clone(&site)),
decoy_header_secs: 1,
profiles: vec![Arc::clone(&profile)],
});
let mut vhosts = BTreeMap::new();
vhosts.insert("proxy.example.com".to_string(), vhost);
vhosts.insert(
"other.example.com".to_string(),
Arc::new(WebRuntimeVhost {
host: "other.example.com".to_string(),
decoy: WebRuntimeDecoy::StaticDirectory(site),
decoy_header_secs: 1,
profiles: Vec::new(),
}),
);
let mut config = ProxyConfig::default();
config.web.enabled = true;
config.web.limits.max_bootstraps_per_ip = 1;
config.web.timeouts.shutdown_secs = 1;
config.web.runtime = Some(Arc::new(WebRuntimeConfig {
vhosts,
profiles: vec![profile],
}));
config
}
async fn request(
listener: &TcpListener,
runtime: &Arc<WebProcessRuntime>,
request: Vec<u8>,
) -> Vec<u8> {
let addr = listener.local_addr().unwrap();
let (accepted, client) = tokio::join!(listener.accept(), TcpStream::connect(addr));
let (server, peer) = accepted.unwrap();
let mut client = client.unwrap();
let permit = runtime.try_http_connection().unwrap();
let task = tokio::spawn(serve_connection(
server,
peer,
WebClientIpSource::XForwardedFor,
Arc::from(["127.0.0.1/32".parse().unwrap()]),
Arc::clone(runtime),
CancellationToken::new(),
permit,
));
client.write_all(&request).await.unwrap();
let mut response = Vec::new();
client.read_to_end(&mut response).await.unwrap();
task.await.unwrap();
response
}
fn split_response(response: &[u8]) -> (&[u8], &[u8]) {
let separator = response
.windows(4)
.position(|window| window == b"\r\n\r\n")
.unwrap();
(&response[..separator], &response[separator + 4..])
}
fn response_header<'a>(headers: &'a [u8], name: &str) -> &'a str {
std::str::from_utf8(headers)
.unwrap()
.lines()
.filter_map(|line| line.split_once(':'))
.find_map(|(header, value)| header.eq_ignore_ascii_case(name).then_some(value.trim()))
.unwrap()
}
#[tokio::test]
async fn https_carrier_bootstraps_and_closes_one_session() {
let capability = [7u8; 32];
let generation = test_runtime_generation(1, runtime_config(capability));
let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation)));
let runtime = WebProcessRuntime::start(Arc::clone(&active_runtime));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(capability);
let wrong_family = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 2001:db8::10\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let wrong_family_response = request(&listener, &runtime, wrong_family).await;
let (_, wrong_family_body) = split_response(&wrong_family_response);
assert!(!wrong_family_body
.windows(11)
.any(|value| value == b"bootstrap='"));
let root = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let root_response = request(&listener, &runtime, root).await;
let (root_headers, root_body) = split_response(&root_response);
assert!(root_headers.starts_with(b"HTTP/1.1 200"));
let root_body = std::str::from_utf8(root_body).unwrap();
let bootstrap = root_body
.split_once("bootstrap='")
.and_then(|(_, suffix)| suffix.split_once('\''))
.map(|(token, _)| token)
.unwrap();
assert_eq!(bootstrap.len(), 43);
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let mut wrong_host = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: other.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
wrong_host.extend_from_slice(&hello);
let wrong_host_response = request(&listener, &runtime, wrong_host).await;
assert!(wrong_host_response.starts_with(b"HTTP/1.1 404"));
let mut create = format!(
"POST /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
hello.len()
)
.into_bytes();
let create_retry = create.clone();
create.extend_from_slice(&hello);
let mut create_retry = create_retry;
create_retry.extend_from_slice(&hello);
let create_response = request(&listener, &runtime, create).await;
let (create_headers, create_body) = split_response(&create_response);
assert!(create_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(create_headers, "x-carrier-mode"), "https");
assert_eq!(create_body, frame::encode(FrameType::Welcome, 0, &[]));
let session = response_header(create_headers, "x-session-token");
assert_eq!(session.len(), 43);
let replacement = test_runtime_generation(2, runtime_config(capability));
active_runtime.store(Arc::clone(&replacement));
tokio::time::sleep(std::time::Duration::from_millis(1100)).await;
let retry_response = request(&listener, &runtime, create_retry).await;
let (retry_headers, retry_body) = split_response(&retry_response);
assert!(retry_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(retry_headers, "x-session-token"), session);
assert_eq!(retry_body, frame::encode(FrameType::Welcome, 0, &[]));
let next_root = format!(
"GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let next_root_response = request(&listener, &runtime, next_root).await;
let (_, next_root_body) = split_response(&next_root_response);
assert!(next_root_body.windows(11).any(|value| value == b"bootstrap='"));
let close = format!(
"DELETE /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
)
.into_bytes();
let close_retry = close.clone();
let close_response = request(&listener, &runtime, close).await;
assert!(close_response.starts_with(b"HTTP/1.1 204"));
let close_retry_response = request(&listener, &runtime, close_retry).await;
assert!(close_retry_response.starts_with(b"HTTP/1.1 204"));
runtime.shutdown().await;
generation.stop_sessions().await;
generation.stop_background_tasks().await;
replacement.stop_sessions().await;
replacement.stop_background_tasks().await;
}
+532
View File
@@ -0,0 +1,532 @@
use std::future::Future;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use arc_swap::ArcSwap;
use parking_lot::Mutex;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore};
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use zeroize::Zeroizing;
use crate::config::{WebLimitsConfig, WebRuntimeProfile};
use crate::maestro::generation::RuntimeGeneration;
use crate::web::frame;
use crate::web::session::WebSession;
// Credential maps, quotas, and token-bucket helpers remain private to the manager.
mod state;
// Stream admission and synthetic tuple ownership are process-scoped.
mod admission;
// Shutdown and expiry work remain outside request-path coordination.
mod lifecycle;
use state::{
Bootstrap, ManagerState, allow_rate, control_item_reserve, decrement_map,
evict_oldest_unused_bootstrap, matching_profile, new_unique_token, profile_key,
remove_expired_locked,
};
const TOKEN_BYTES: usize = 32;
const CLEANUP_INTERVAL: Duration = Duration::from_secs(1);
/// Stable hash key used for bootstrap and session credentials.
pub(crate) type TokenHash = [u8; TOKEN_BYTES];
/// Stable non-allocating key used for per-profile quotas.
pub(crate) type ProfileKey = [u8; TOKEN_BYTES];
/// WEB manager operation failure category.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum ManagerError {
/// Credential, hostname, or ownership validation failed.
Authentication,
/// Bounded queue capacity is temporarily unavailable.
Backpressure,
/// A configured admission or rate ceiling was reached.
Limit,
/// Carrier framing or sequencing violated the protocol.
Protocol,
/// The operation conflicts with another in-flight operation.
Concurrent,
/// The process or session has stopped accepting work.
Closed,
}
/// Successful idempotent session creation result.
pub(crate) struct CreateResult {
/// Opaque bearer token for the created or replayed session.
pub(crate) token: String,
}
/// Process-owned bounded WEB credential, session, and memory coordinator.
pub(crate) struct WebProcessRuntime {
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
limits: WebLimitsConfig,
state: Mutex<ManagerState>,
http_connections: Arc<Semaphore>,
http_handlers: Arc<Semaphore>,
body_readers: Arc<Semaphore>,
body_bytes: Arc<Semaphore>,
stream_handshakes: Arc<Semaphore>,
budget_notify: Arc<Notify>,
budget_saturated: AtomicBool,
shutdown: CancellationToken,
tasks: TaskTracker,
sessions_created: AtomicU64,
sessions_closed: AtomicU64,
streams_opened: AtomicU64,
streams_rejected: AtomicU64,
bytes_up: AtomicU64,
bytes_down: AtomicU64,
limit_hits: AtomicU64,
}
impl WebProcessRuntime {
/// Starts one process-scoped manager using immutable allocation ceilings.
pub(crate) fn start(
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) -> Arc<Self> {
let limits = active_runtime.load().config().web.limits.clone();
let runtime = Arc::new(Self {
active_runtime,
http_connections: Arc::new(Semaphore::new(limits.max_http_connections)),
http_handlers: Arc::new(Semaphore::new(limits.max_http_handlers)),
body_readers: Arc::new(Semaphore::new(limits.max_body_readers)),
body_bytes: Arc::new(Semaphore::new(limits.max_body_bytes_global)),
stream_handshakes: Arc::new(Semaphore::new(limits.max_stream_handshakes)),
limits,
state: Mutex::new(ManagerState::default()),
budget_notify: Arc::new(Notify::new()),
budget_saturated: AtomicBool::new(false),
shutdown: CancellationToken::new(),
tasks: TaskTracker::new(),
sessions_created: AtomicU64::new(0),
sessions_closed: AtomicU64::new(0),
streams_opened: AtomicU64::new(0),
streams_rejected: AtomicU64::new(0),
bytes_up: AtomicU64::new(0),
bytes_down: AtomicU64::new(0),
limit_hits: AtomicU64::new(0),
});
let weak = Arc::downgrade(&runtime);
let shutdown = runtime.shutdown.clone();
runtime.tasks.spawn(async move {
let mut interval = tokio::time::interval(CLEANUP_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
_ = shutdown.cancelled() => break,
_ = interval.tick() => {
let Some(runtime) = weak.upgrade() else {
break;
};
runtime.cleanup();
}
}
}
});
runtime
}
/// Loads the currently active generation without retaining older generations.
pub(crate) fn active_generation(&self) -> Arc<RuntimeGeneration> {
self.active_runtime.load_full()
}
/// Reserves one accepted HTTP connection.
pub(crate) fn try_http_connection(&self) -> Option<OwnedSemaphorePermit> {
let permit = Arc::clone(&self.http_connections).try_acquire_owned().ok();
if permit.is_none() {
self.record_limit_hit();
}
permit
}
/// Reserves one concurrently executing HTTP request handler.
pub(crate) fn try_http_handler(&self) -> Option<OwnedSemaphorePermit> {
let permit = Arc::clone(&self.http_handlers).try_acquire_owned().ok();
if permit.is_none() {
self.record_limit_hit();
}
permit
}
/// Reserves one logical stream in the inner MTProxy handshake phase.
pub(crate) fn try_stream_handshake(&self) -> Option<OwnedSemaphorePermit> {
let permit = Arc::clone(&self.stream_handshakes)
.try_acquire_owned()
.ok();
if permit.is_none() {
self.record_stream_rejected();
}
permit
}
/// Spawns one process-owned auxiliary task with shutdown cancellation.
pub(crate) fn spawn_auxiliary<F>(&self, future: F)
where
F: Future<Output = ()> + Send + 'static,
{
let shutdown = self.shutdown.clone();
self.tasks.spawn(async move {
tokio::select! {
_ = shutdown.cancelled() => {}
_ = future => {}
}
});
}
/// Reserves one body reader and its declared bounded body allocation.
pub(crate) fn try_body_budget(
&self,
bytes: usize,
) -> Option<(OwnedSemaphorePermit, OwnedSemaphorePermit)> {
let Some(bytes) = u32::try_from(bytes).ok() else {
self.record_limit_hit();
return None;
};
let Some(reader) = Arc::clone(&self.body_readers).try_acquire_owned().ok() else {
self.record_limit_hit();
return None;
};
let Some(body) = Arc::clone(&self.body_bytes)
.try_acquire_many_owned(bytes)
.ok()
else {
self.record_limit_hit();
return None;
};
Some((reader, body))
}
/// Issues a one-use bootstrap credential for the active generation.
pub(crate) fn issue_bootstrap(
&self,
profile: Arc<WebRuntimeProfile>,
client_ip: IpAddr,
) -> std::result::Result<String, ManagerError> {
let generation = self.active_generation();
let config = generation.config();
let profile = config
.web
.runtime
.as_ref()
.and_then(|runtime| matching_profile(runtime, &profile))
.ok_or(ManagerError::Authentication)?;
if !config.web.enabled
|| profile.public_addr.is_ipv4() != client_ip.is_ipv4()
|| !generation.proxy_shared.is_user_enabled(&profile.user)
{
return Err(ManagerError::Closed);
}
let now = Instant::now();
let mut state = self.state.lock();
remove_expired_locked(&mut state, now);
if state.closed
|| state.bootstraps_per_ip.get(&client_ip).copied().unwrap_or(0)
>= self.limits.max_bootstraps_per_ip
|| !allow_rate(
&mut state.bootstrap_rate,
now,
self.limits.new_bootstraps_per_minute,
self.limits.new_bootstraps_burst,
)
{
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
}
if state.bootstraps.len() >= self.limits.max_bootstraps_global
&& !evict_oldest_unused_bootstrap(&mut state)
{
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
}
let Some((token, hash)) = new_unique_token(&generation, &state) else {
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
};
state.bootstraps.insert(
hash,
Bootstrap {
generation_id: generation.id,
expires_at: now + Duration::from_secs(config.web.timeouts.bootstrap_lifetime_secs),
issued_at: now,
issuance_ip: client_ip,
profile,
body_digest: [0; TOKEN_BYTES],
session_token: Zeroizing::new(String::new()),
session: None,
used: false,
},
);
*state.bootstraps_per_ip.entry(client_ip).or_insert(0) += 1;
Ok(token)
}
/// Checks whether a bootstrap token is live before reading a request body.
pub(crate) fn has_bootstrap(&self, hash: TokenHash, host: &str) -> bool {
let generation_id = self.active_runtime.load().id;
let now = Instant::now();
let state = self.state.lock();
state.bootstraps.get(&hash).is_some_and(|entry| {
entry.profile.host == host
&& now <= entry.expires_at
&& (entry.generation_id == generation_id
|| entry.used && entry.session.is_some())
})
}
/// Creates a session exactly once or replays the original successful result.
pub(crate) fn create_session(
self: &Arc<Self>,
bootstrap_hash: TokenHash,
host: &str,
client_ip: IpAddr,
body: &[u8],
) -> std::result::Result<CreateResult, ManagerError> {
if !frame::validate_hello(body, &self.limits) {
return Err(ManagerError::Protocol);
}
let body_digest: TokenHash = Sha256::digest(body).into();
let generation = self.active_generation();
let config = generation.config();
let now = Instant::now();
let mut state = self.state.lock();
remove_expired_locked(&mut state, now);
let Some(entry) = state.bootstraps.get(&bootstrap_hash) else {
return Err(ManagerError::Authentication);
};
if entry.profile.host != host || now > entry.expires_at {
return Err(ManagerError::Authentication);
}
if entry.used {
let digest_matches = bool::from(entry.body_digest.ct_eq(&body_digest));
if !digest_matches {
return Err(ManagerError::Authentication);
}
if entry.session.is_none() {
return Err(ManagerError::Authentication);
}
return Ok(CreateResult {
token: entry.session_token.as_str().to_owned(),
});
}
if entry.generation_id != generation.id {
return Err(ManagerError::Authentication);
}
if state.closed || !config.web.enabled {
return Err(ManagerError::Closed);
}
let profile = config
.web
.runtime
.as_ref()
.and_then(|runtime| matching_profile(runtime, &entry.profile))
.filter(|profile| {
profile.public_addr.is_ipv4() == client_ip.is_ipv4()
&& generation.proxy_shared.is_user_enabled(&profile.user)
})
.ok_or(ManagerError::Authentication)?;
let profile_key = profile_key(&profile);
if state.sessions.len() >= self.limits.max_sessions_global
|| state.sessions_per_ip.get(&client_ip).copied().unwrap_or(0)
>= self.limits.max_sessions_per_ip
|| state
.sessions_per_profile
.get(&profile_key)
.copied()
.unwrap_or(0)
>= profile.max_sessions
|| !allow_rate(
&mut state.session_rate,
now,
self.limits.new_sessions_per_minute,
self.limits.new_sessions_burst,
)
{
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
}
let Some((session_token, session_hash)) = new_unique_token(&generation, &state) else {
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return Err(ManagerError::Limit);
};
let session = WebSession::new(
Arc::downgrade(self),
session_hash,
client_ip,
profile,
profile_key,
self.limits.clone(),
config.web.timeouts.clone(),
);
state.sessions.insert(session_hash, Arc::clone(&session));
*state.sessions_per_ip.entry(client_ip).or_insert(0) += 1;
*state.sessions_per_profile.entry(profile_key).or_insert(0) += 1;
let entry = state
.bootstraps
.get_mut(&bootstrap_hash)
.ok_or(ManagerError::Authentication)?;
entry.used = true;
entry.body_digest = body_digest;
entry.session_token = Zeroizing::new(session_token.clone());
entry.session = Some(Arc::clone(&session));
let issuance_ip = entry.issuance_ip;
decrement_map(&mut state.bootstraps_per_ip, &issuance_ip);
self.sessions_created.fetch_add(1, Ordering::Relaxed);
Ok(CreateResult {
token: session_token,
})
}
/// Resolves an authenticated session token.
pub(crate) fn get_session(
&self,
hash: TokenHash,
host: &str,
) -> std::result::Result<Arc<WebSession>, ManagerError> {
self.state
.lock()
.sessions
.get(&hash)
.cloned()
.filter(|session| session.matches_host(host))
.ok_or(ManagerError::Authentication)
}
/// Closes a live token and accepts bounded tombstone retries.
pub(crate) fn close_token(
&self,
hash: TokenHash,
host: &str,
) -> std::result::Result<(), ManagerError> {
let state = self.state.lock();
let session = state
.sessions
.get(&hash)
.filter(|session| session.matches_host(host))
.cloned();
let closed = state
.closed_tokens
.get(&hash)
.is_some_and(|closed| closed.host == host);
drop(state);
if let Some(session) = session {
session.close();
return Ok(());
}
closed.then_some(()).ok_or(ManagerError::Authentication)
}
/// Reserves bounded process-wide queue capacity for data or control traffic.
pub(crate) fn try_reserve_pending(
&self,
bytes: usize,
items: usize,
control: bool,
downlink: bool,
) -> bool {
let mut state = self.state.lock();
let data_byte_limit = self
.limits
.pending_bytes_global
.saturating_sub(self.limits.control_bytes_global);
let control_item_reserve = control_item_reserve(&self.limits);
let data_item_limit = self
.limits
.pending_items_global
.saturating_sub(control_item_reserve);
if state.closed {
return false;
}
let fits = if control {
bytes <= self.limits.control_bytes_global
&& items <= control_item_reserve
&& state.pending_bytes
<= self.limits.pending_bytes_global.saturating_sub(bytes)
&& state.pending_items
<= self.limits.pending_items_global.saturating_sub(items)
&& state.pending_control_bytes
<= self.limits.control_bytes_global.saturating_sub(bytes)
&& state.pending_control_items
<= control_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 downlink {
let uplink_bytes = self
.limits
.max_body_bytes
.saturating_add(
self.limits
.max_frames_per_body
.saturating_mul(crate::web::session::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 {
self.budget_saturated.store(true, Ordering::Release);
self.record_limit_hit();
return false;
}
state.pending_bytes += bytes;
state.pending_items += items;
if control {
state.pending_control_bytes += bytes;
state.pending_control_items += items;
}
true
}
/// Releases process-wide queue capacity and wakes blocked relay writers.
pub(crate) fn release_pending(&self, bytes: usize, items: usize, control: bool) {
let mut state = self.state.lock();
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);
}
drop(state);
if self.budget_saturated.swap(false, Ordering::AcqRel) {
self.budget_notify.notify_waiters();
}
}
/// Returns the shared notification source for global queue capacity changes.
pub(crate) fn budget_notify(&self) -> Arc<Notify> {
Arc::clone(&self.budget_notify)
}
/// Accounts one successfully committed carrier uplink body.
pub(crate) fn record_up(&self, bytes: usize) {
self.bytes_up.fetch_add(bytes as u64, Ordering::Relaxed);
}
/// Accounts one emitted carrier downlink body.
pub(crate) fn record_down(&self, bytes: usize) {
self.bytes_down.fetch_add(bytes as u64, Ordering::Relaxed);
}
fn record_limit_hit(&self) {
self.limit_hits.fetch_add(1, Ordering::Relaxed);
}
}
+130
View File
@@ -0,0 +1,130 @@
use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::Ordering;
use std::time::Instant;
use super::state::{
allocate_stream_port, allow_rate, decrement_map, release_stream_port,
};
use super::{ProfileKey, WebProcessRuntime};
impl WebProcessRuntime {
/// Reserves one process-wide and per-profile live logical-stream slot.
pub(crate) fn try_acquire_stream(
&self,
profile_key: ProfileKey,
max_streams: usize,
client_ip: IpAddr,
public_addr: SocketAddr,
) -> Option<u16> {
let now = Instant::now();
let mut state = self.state.lock();
if state.closed
|| state.streams_live >= self.limits.max_streams_global
|| state
.streams_per_profile
.get(&profile_key)
.copied()
.unwrap_or(0)
>= max_streams
|| !allow_rate(
&mut state.stream_rate,
now,
self.limits.new_streams_per_minute,
self.limits.new_streams_burst,
)
{
self.streams_rejected.fetch_add(1, Ordering::Relaxed);
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return None;
}
let Some(peer_port) = allocate_stream_port(&mut state, client_ip, public_addr) else {
self.streams_rejected.fetch_add(1, Ordering::Relaxed);
self.limit_hits.fetch_add(1, Ordering::Relaxed);
return None;
};
state.streams_live += 1;
*state
.streams_per_profile
.entry(profile_key)
.or_insert(0) += 1;
self.streams_opened.fetch_add(1, Ordering::Relaxed);
Some(peer_port)
}
/// Releases one live logical-stream slot after its relay task exits.
pub(crate) fn release_stream(
&self,
profile_key: ProfileKey,
client_ip: IpAddr,
public_addr: SocketAddr,
peer_port: u16,
) {
let mut state = self.state.lock();
if !release_stream_port(&mut state, client_ip, public_addr, peer_port) {
return;
}
state.streams_live = state.streams_live.saturating_sub(1);
decrement_map(&mut state.streams_per_profile, &profile_key);
}
/// Records a logical stream rejected outside manager quota acquisition.
pub(crate) fn record_stream_rejected(&self) {
self.streams_rejected.fetch_add(1, Ordering::Relaxed);
self.record_limit_hit();
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use arc_swap::ArcSwap;
use super::*;
use crate::config::ProxyConfig;
use crate::maestro::generation::test_runtime_generation;
use crate::web::session::QUEUE_ITEM_COST;
#[tokio::test]
async fn global_downlink_budget_preserves_one_maximum_uplink_batch() {
let generation = test_runtime_generation(1, ProxyConfig::default());
let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(generation)));
let control_items = super::super::state::control_item_reserve(&runtime.limits);
let data_bytes = runtime
.limits
.pending_bytes_global
.saturating_sub(runtime.limits.control_bytes_global);
let data_items = runtime
.limits
.pending_items_global
.saturating_sub(control_items);
let uplink_bytes = runtime
.limits
.max_body_bytes
.saturating_add(runtime.limits.max_frames_per_body * QUEUE_ITEM_COST);
let downlink_bytes = data_bytes - uplink_bytes;
let downlink_items = data_items - runtime.limits.max_frames_per_body;
assert!(runtime.try_reserve_pending(
downlink_bytes,
downlink_items,
false,
true,
));
assert!(runtime.try_reserve_pending(
uplink_bytes,
runtime.limits.max_frames_per_body,
false,
false,
));
assert!(!runtime.try_reserve_pending(1, 1, false, true));
runtime.release_pending(downlink_bytes, downlink_items, false);
runtime.release_pending(
uplink_bytes,
runtime.limits.max_frames_per_body,
false,
);
runtime.shutdown().await;
}
}
+149
View File
@@ -0,0 +1,149 @@
use std::net::IpAddr;
use std::sync::atomic::Ordering;
use std::time::{Duration, Instant};
use tracing::info;
use super::{ProfileKey, TokenHash, WebProcessRuntime};
use super::state::{
ClosedToken, decrement_map, remove_bootstrap_locked, remove_expired_locked,
};
impl WebProcessRuntime {
/// Removes one closed session and retains a bounded host-bound replay marker.
pub(crate) fn session_finished(
&self,
hash: TokenHash,
client_ip: IpAddr,
profile_key: ProfileKey,
profile_host: &str,
) {
let mut state = self.state.lock();
if state.sessions.remove(&hash).is_none() {
return;
}
decrement_map(&mut state.sessions_per_ip, &client_ip);
decrement_map(&mut state.sessions_per_profile, &profile_key);
let expiry = Instant::now()
+ Duration::from_secs(
self.active_runtime
.load()
.config()
.web
.timeouts
.bootstrap_lifetime_secs,
);
state.closed_tokens.insert(
hash,
ClosedToken {
expires_at: expiry,
host: profile_host.to_string(),
},
);
while state.closed_tokens.len() > self.limits.max_sessions_global.saturating_mul(16) {
let Some(oldest) = state
.closed_tokens
.iter()
.min_by_key(|(_, closed)| closed.expires_at)
.map(|(hash, _)| *hash)
else {
break;
};
state.closed_tokens.remove(&oldest);
}
let bootstrap_hashes = state
.bootstraps
.iter()
.filter_map(|(bootstrap_hash, bootstrap)| {
bootstrap
.session
.as_ref()
.is_some_and(|session| session.token_hash() == hash)
.then_some(*bootstrap_hash)
})
.collect::<Vec<_>>();
for bootstrap_hash in bootstrap_hashes {
remove_bootstrap_locked(&mut state, bootstrap_hash);
}
self.sessions_closed.fetch_add(1, Ordering::Relaxed);
}
/// Stops issuance, closes all sessions, and joins bounded child work.
pub(crate) async fn shutdown(&self) {
self.shutdown.cancel();
let sessions = {
let mut state = self.state.lock();
state.closed = true;
state.bootstraps.clear();
state.bootstraps_per_ip.clear();
state.sessions.values().cloned().collect::<Vec<_>>()
};
for session in &sessions {
session.close();
}
let timeout_secs = self
.active_runtime
.load()
.config()
.web
.timeouts
.shutdown_secs;
let waits = async {
for session in sessions {
session.wait().await;
}
};
let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), waits).await;
self.tasks.close();
let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), self.tasks.wait()).await;
let (sessions_live, streams_live, pending_bytes, pending_items) = {
let state = self.state.lock();
(
state.sessions.len(),
state.streams_live,
state.pending_bytes,
state.pending_items,
)
};
info!(
target: "telemt::web",
sessions_created = self.sessions_created.load(Ordering::Relaxed),
sessions_closed = self.sessions_closed.load(Ordering::Relaxed),
sessions_live,
streams_opened = self.streams_opened.load(Ordering::Relaxed),
streams_rejected = self.streams_rejected.load(Ordering::Relaxed),
streams_live,
pending_bytes,
pending_items,
bytes_up = self.bytes_up.load(Ordering::Relaxed),
bytes_down = self.bytes_down.load(Ordering::Relaxed),
limit_hits = self.limit_hits.load(Ordering::Relaxed),
"WEB runtime stopped"
);
}
/// Expires credentials and closes idle sessions without holding locks across callbacks.
pub(super) fn cleanup(&self) {
let generation_id = self.active_runtime.load().id;
let now = Instant::now();
let sessions = {
let mut state = self.state.lock();
remove_expired_locked(&mut state, now);
let stale_bootstraps = state
.bootstraps
.iter()
.filter_map(|(hash, bootstrap)| {
(bootstrap.generation_id != generation_id && !bootstrap.used)
.then_some(*hash)
})
.collect::<Vec<_>>();
for hash in stale_bootstraps {
remove_bootstrap_locked(&mut state, hash);
}
state.sessions.values().cloned().collect::<Vec<_>>()
};
for session in sessions.into_iter().filter(|session| session.is_idle(now)) {
session.close();
}
}
}
+293
View File
@@ -0,0 +1,293 @@
use std::collections::{HashMap, HashSet};
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::Instant;
use base64::Engine as _;
use sha2::{Digest, Sha256};
use zeroize::Zeroizing;
use super::{ProfileKey, TOKEN_BYTES, TokenHash};
use crate::config::{WebLimitsConfig, WebRuntimeConfig, WebRuntimeProfile};
use crate::maestro::generation::RuntimeGeneration;
use crate::web::session::WebSession;
/// One issued bootstrap and optional idempotent session-creation replay state.
pub(super) struct Bootstrap {
/// Generation that issued the bootstrap.
pub(super) generation_id: u64,
/// Credential and replay-state expiry deadline.
pub(super) expires_at: Instant,
/// Stable ordering point used for bounded eviction.
pub(super) issued_at: Instant,
/// Forwarded client address that owns this credential.
pub(super) issuance_ip: IpAddr,
/// Immutable profile selected during capability validation.
pub(super) profile: Arc<WebRuntimeProfile>,
/// Digest of the accepted HELLO body for idempotent retry matching.
pub(super) body_digest: TokenHash,
/// Zeroizing copy returned only for an exact session-creation retry.
pub(super) session_token: Zeroizing<String>,
/// Created session retained while retry replay remains valid.
pub(super) session: Option<Arc<WebSession>>,
/// Distinguishes unused issuance quota from completed creation replay state.
pub(super) used: bool,
}
/// Bounded replay marker for one explicitly or naturally closed session token.
pub(super) struct ClosedToken {
/// Deadline after which the token hash may be forgotten.
pub(super) expires_at: Instant,
/// Canonical host that owned the session.
pub(super) host: String,
}
/// Token-bucket state for one process-wide creation class.
#[derive(Default)]
pub(super) struct RateState {
tokens: f64,
last: Option<Instant>,
}
struct StreamPortState {
active: HashSet<u16>,
next: u16,
}
/// Process-wide WEB registries and quota accounting protected by one short lock.
#[derive(Default)]
pub(super) struct ManagerState {
/// Bootstrap credentials indexed by their SHA-256 token hash.
pub(super) bootstraps: HashMap<TokenHash, Bootstrap>,
/// Unused bootstrap ownership counts by forwarded client address.
pub(super) bootstraps_per_ip: HashMap<IpAddr, usize>,
/// Live sessions indexed by bearer-token hash.
pub(super) sessions: HashMap<TokenHash, Arc<WebSession>>,
/// Recently closed token hashes retained for idempotent DELETE semantics.
pub(super) closed_tokens: HashMap<TokenHash, ClosedToken>,
/// Live session counts by forwarded client address.
pub(super) sessions_per_ip: HashMap<IpAddr, usize>,
/// Live session counts by stable profile key.
pub(super) sessions_per_profile: HashMap<ProfileKey, usize>,
/// Live relay-task counts by stable profile key.
pub(super) streams_per_profile: HashMap<ProfileKey, usize>,
/// Process-wide live relay-task count.
pub(super) streams_live: usize,
stream_ports: HashMap<(IpAddr, SocketAddr), StreamPortState>,
/// Total process-wide queued byte reservation.
pub(super) pending_bytes: usize,
/// Total process-wide queued item reservation.
pub(super) pending_items: usize,
/// Portion of queued bytes charged to the control reserve.
pub(super) pending_control_bytes: usize,
/// Portion of queued items charged to the control reserve.
pub(super) pending_control_items: usize,
/// Bootstrap issuance rate limiter.
pub(super) bootstrap_rate: RateState,
/// Session creation rate limiter.
pub(super) session_rate: RateState,
/// Logical-stream creation rate limiter.
pub(super) stream_rate: RateState,
/// Process shutdown admission latch.
pub(super) closed: bool,
}
/// Generates one collision-checked credential and its stable hash key.
pub(super) fn new_unique_token(
generation: &RuntimeGeneration,
state: &ManagerState,
) -> Option<(String, TokenHash)> {
for _ in 0..8 {
let mut raw = [0u8; TOKEN_BYTES];
generation.rng.fill(&mut raw);
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw);
let hash = Sha256::digest(raw).into();
if !state.bootstraps.contains_key(&hash)
&& !state.sessions.contains_key(&hash)
&& !state.closed_tokens.contains_key(&hash)
{
return Some((token, hash));
}
}
None
}
/// Returns the precomputed capability as the stable process profile key.
pub(super) fn profile_key(profile: &WebRuntimeProfile) -> ProfileKey {
profile.capability
}
/// Re-resolves an issued profile against the active generation without weakening identity.
pub(super) fn matching_profile(
runtime: &WebRuntimeConfig,
expected: &WebRuntimeProfile,
) -> Option<Arc<WebRuntimeProfile>> {
runtime
.profiles
.iter()
.find(|profile| {
profile.host == expected.host
&& profile.public_addr == expected.public_addr
&& profile.user == expected.user
&& profile.secret_mode == expected.secret_mode
&& profile.capability == expected.capability
})
.cloned()
}
/// Applies one token-bucket admission decision at a caller-supplied monotonic time.
pub(super) fn allow_rate(
state: &mut RateState,
now: Instant,
per_minute: u32,
burst: u32,
) -> bool {
let burst = f64::from(burst);
if let Some(last) = state.last {
let elapsed = now.saturating_duration_since(last).as_secs_f64();
state.tokens =
(state.tokens + elapsed * f64::from(per_minute) / 60.0).min(burst);
} else {
state.tokens = burst;
}
state.last = Some(now);
if state.tokens < 1.0 {
return false;
}
state.tokens -= 1.0;
true
}
/// Evicts the oldest unused bootstrap while preserving used retry state.
pub(super) fn evict_oldest_unused_bootstrap(state: &mut ManagerState) -> bool {
let Some(hash) = state
.bootstraps
.iter()
.filter(|(_, bootstrap)| !bootstrap.used)
.min_by_key(|(_, bootstrap)| bootstrap.issued_at)
.map(|(hash, _)| *hash)
else {
return false;
};
remove_bootstrap_locked(state, hash);
true
}
/// Removes expired bootstrap and closed-token entries while the manager lock is held.
pub(super) fn remove_expired_locked(state: &mut ManagerState, now: Instant) {
let expired = state
.bootstraps
.iter()
.filter_map(|(hash, bootstrap)| (now > bootstrap.expires_at).then_some(*hash))
.collect::<Vec<_>>();
for hash in expired {
remove_bootstrap_locked(state, hash);
}
state
.closed_tokens
.retain(|_, closed| now <= closed.expires_at);
}
/// Removes one bootstrap and releases its per-address issuance quota when unused.
pub(super) fn remove_bootstrap_locked(state: &mut ManagerState, hash: TokenHash) {
let Some(bootstrap) = state.bootstraps.remove(&hash) else {
return;
};
if !bootstrap.used {
decrement_map(&mut state.bootstraps_per_ip, &bootstrap.issuance_ip);
}
}
/// Decrements one counted owner and removes its map entry at zero.
pub(super) fn decrement_map<K, Q>(values: &mut HashMap<K, usize>, key: &Q)
where
K: std::borrow::Borrow<Q> + std::hash::Hash + Eq,
Q: std::hash::Hash + Eq + ?Sized,
{
let remove = if let Some(value) = values.get_mut(key) {
*value = value.saturating_sub(1);
*value == 0
} else {
false
};
if remove {
values.remove(key);
}
}
/// Computes the process-wide item reserve required for session control progress.
pub(super) fn control_item_reserve(limits: &WebLimitsConfig) -> usize {
limits.max_sessions_global.saturating_mul(
16usize.saturating_add(limits.max_streams_per_session.saturating_mul(3)),
)
}
/// Allocates a non-zero source port unique among live streams for one KDF route.
pub(super) fn allocate_stream_port(
state: &mut ManagerState,
client_ip: IpAddr,
public_addr: SocketAddr,
) -> Option<u16> {
let ports = state
.stream_ports
.entry((client_ip, public_addr))
.or_insert_with(|| StreamPortState {
active: HashSet::new(),
next: 1,
});
for _ in 0..u16::MAX {
let candidate = ports.next;
ports.next = ports.next.checked_add(1).unwrap_or(1);
if ports.active.insert(candidate) {
return Some(candidate);
}
}
None
}
/// Releases one source port and reclaims empty per-route allocator state.
pub(super) fn release_stream_port(
state: &mut ManagerState,
client_ip: IpAddr,
public_addr: SocketAddr,
peer_port: u16,
) -> bool {
let key = (client_ip, public_addr);
let Some(ports) = state.stream_ports.get_mut(&key) else {
return false;
};
let removed = ports.active.remove(&peer_port);
if ports.active.is_empty() {
state.stream_ports.remove(&key);
}
removed
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn synthetic_ports_are_unique_per_live_route_and_state_is_reclaimed() {
let mut state = ManagerState::default();
let client_ip = "192.0.2.10".parse().unwrap();
let public_addr = "203.0.113.10:443".parse().unwrap();
let first = allocate_stream_port(&mut state, client_ip, public_addr).unwrap();
let second = allocate_stream_port(&mut state, client_ip, public_addr).unwrap();
assert_ne!(first, second);
assert!(release_stream_port(
&mut state,
client_ip,
public_addr,
first,
));
assert!(release_stream_port(
&mut state,
client_ip,
public_addr,
second,
));
assert!(state.stream_ports.is_empty());
}
}
+14
View File
@@ -0,0 +1,14 @@
//! Bounded WEB carrier ingress behind a trusted external TLS terminator.
/// Browser bridge generation for the serialized HTTPS carrier.
pub(crate) mod bridge;
/// Shared binary frame codec and protocol constants.
pub(crate) mod frame;
/// Plain HTTP ingress and decoy routing behind external TLS termination.
pub(crate) mod http;
/// Process-wide credentials, quotas, memory budgets, and shutdown ownership.
pub(crate) mod manager;
/// Resumable carrier sessions and logical-stream state machines.
pub(crate) mod session;
/// AsyncRead and AsyncWrite adapter for one logical MTProxy stream.
pub(crate) mod stream;
+355
View File
@@ -0,0 +1,355 @@
use std::collections::{HashMap, HashSet, VecDeque};
use std::io;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
use bytes::{Bytes, BytesMut};
use parking_lot::Mutex;
use tokio::io::ReadBuf;
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken;
use crate::config::{WebLimitsConfig, WebRuntimeProfile, WebTimeoutsConfig};
use crate::web::frame::{self, FrameType};
use crate::web::manager::{ProfileKey, TokenHash, WebProcessRuntime};
// Backend tasks own generation admission and authenticated MTProxy relay lifetimes.
mod backend;
// Downlink queues own cursor replay, flow control, and memory reservations.
mod downlink;
// Uplink batches own exactly-once sequencing and client-frame validation.
mod uplink;
/// Conservative allocator and container overhead charged to every queued item.
pub(crate) const QUEUE_ITEM_COST: usize = 256;
#[derive(Clone, Copy, PartialEq, Eq)]
enum PendingClass {
Uplink,
Downlink,
Control,
}
struct InboundChunk {
bytes: Bytes,
offset: usize,
}
struct StreamState {
inbound: VecDeque<InboundChunk>,
receive_window: u32,
send_credit: u64,
read_waker: Option<Waker>,
write_waker: Option<Waker>,
}
struct QueuedFrame {
encoded: BytesMut,
frame_type: FrameType,
stream_id: u32,
control: bool,
cost: usize,
}
struct DownBatch {
body: Bytes,
base_cursor: u64,
next_cursor: u64,
data_bytes: usize,
data_items: usize,
control_bytes: usize,
control_items: usize,
}
struct SessionState {
streams: HashMap<u32, StreamState>,
active_peer_ports: HashSet<u16>,
closed_streams: HashSet<u32>,
closed_order: VecDeque<u32>,
pending_frames: VecDeque<QueuedFrame>,
pending_windows: HashMap<u32, usize>,
unacked: Option<DownBatch>,
down_cursor: u64,
down_epoch: u64,
last_up_sequence: u64,
last_up_digest: TokenHash,
pending_bytes: usize,
pending_items: usize,
pending_control_bytes: usize,
pending_control_items: usize,
last_activity: Instant,
closed: bool,
}
/// One bounded WEB carrier session containing logical MTProxy streams.
pub(crate) struct WebSession {
manager: std::sync::Weak<WebProcessRuntime>,
token_hash: TokenHash,
client_ip: IpAddr,
profile: Arc<WebRuntimeProfile>,
profile_key: ProfileKey,
limits: WebLimitsConfig,
timeouts: WebTimeoutsConfig,
state: Mutex<SessionState>,
down_notify: Arc<Notify>,
cancel: CancellationToken,
tasks_live: AtomicUsize,
tasks_done: Arc<Notify>,
finished: AtomicBool,
up_active: AtomicBool,
}
/// One successful downlink poll result.
pub(crate) struct PollResult {
/// Encoded downlink frame batch, or an empty long-poll result.
pub(crate) body: Bytes,
/// Cursor the client must present on its next downlink request.
pub(crate) next_cursor: u64,
}
impl WebSession {
#[allow(clippy::too_many_arguments)]
/// Creates one carrier session with immutable ownership and allocation policy.
pub(crate) fn new(
manager: std::sync::Weak<WebProcessRuntime>,
token_hash: TokenHash,
client_ip: IpAddr,
profile: Arc<WebRuntimeProfile>,
profile_key: ProfileKey,
limits: WebLimitsConfig,
timeouts: WebTimeoutsConfig,
) -> Arc<Self> {
Arc::new(Self {
manager,
token_hash,
client_ip,
profile,
profile_key,
limits,
timeouts,
state: Mutex::new(SessionState {
streams: HashMap::new(),
active_peer_ports: HashSet::new(),
closed_streams: HashSet::new(),
closed_order: VecDeque::new(),
pending_frames: VecDeque::new(),
pending_windows: HashMap::new(),
unacked: None,
down_cursor: 0,
down_epoch: 0,
last_up_sequence: 0,
last_up_digest: [0; 32],
pending_bytes: 0,
pending_items: 0,
pending_control_bytes: 0,
pending_control_items: 0,
last_activity: Instant::now(),
closed: false,
}),
down_notify: Arc::new(Notify::new()),
cancel: CancellationToken::new(),
tasks_live: AtomicUsize::new(0),
tasks_done: Arc::new(Notify::new()),
finished: AtomicBool::new(false),
up_active: AtomicBool::new(false),
})
}
/// Returns the stable hashed token identity without exposing the credential.
pub(crate) fn token_hash(&self) -> TokenHash {
self.token_hash
}
/// Checks the canonical virtual host that owns this bearer session.
pub(crate) fn matches_host(&self, host: &str) -> bool {
self.profile.host == host
}
/// Closes carrier state while relay tasks retain their admission until exit.
pub(crate) fn close(&self) {
let (data_bytes, data_items, control_bytes, control_items) = {
let mut state = self.state.lock();
if state.closed {
return;
}
state.closed = true;
for stream in state.streams.values_mut() {
if let Some(waker) = stream.read_waker.take() {
waker.wake();
}
if let Some(waker) = stream.write_waker.take() {
waker.wake();
}
}
state.streams.clear();
state.pending_frames.clear();
state.pending_windows.clear();
state.unacked = None;
let control_bytes = state.pending_control_bytes;
let control_items = state.pending_control_items;
let data_bytes = state.pending_bytes.saturating_sub(control_bytes);
let data_items = state.pending_items.saturating_sub(control_items);
state.pending_bytes = 0;
state.pending_items = 0;
state.pending_control_bytes = 0;
state.pending_control_items = 0;
(data_bytes, data_items, control_bytes, control_items)
};
self.cancel.cancel();
self.down_notify.notify_waiters();
if let Some(manager) = self.manager.upgrade() {
manager.release_pending(data_bytes, data_items, false);
manager.release_pending(control_bytes, control_items, true);
if !self.finished.swap(true, Ordering::AcqRel) {
manager.session_finished(
self.token_hash,
self.client_ip,
self.profile_key,
&self.profile.host,
);
}
}
}
/// Waits for all logical-stream tasks after admission has closed.
pub(crate) async fn wait(&self) {
loop {
let notified = self.tasks_done.notified();
if self.tasks_live.load(Ordering::Acquire) == 0 {
return;
}
notified.await;
}
}
/// Returns whether reconnect grace elapsed without activity.
pub(crate) fn is_idle(&self, now: Instant) -> bool {
let state = self.state.lock();
!state.closed
&& now.saturating_duration_since(state.last_activity)
>= Duration::from_secs(self.timeouts.reconnect_grace_secs)
}
/// Polls client-to-server bytes and returns consumed flow-control credit.
pub(super) fn poll_read(
&self,
stream_id: u32,
cx: &mut Context<'_>,
output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let mut state = self.state.lock();
let (count, finished) = {
let Some(stream) = state.streams.get_mut(&stream_id) else {
return Poll::Ready(Ok(()));
};
let Some(chunk) = stream.inbound.front_mut() else {
stream.read_waker = Some(cx.waker().clone());
return Poll::Pending;
};
let available = &chunk.bytes[chunk.offset..];
let count = available.len().min(output.remaining());
output.put_slice(&available[..count]);
chunk.offset += count;
let finished = chunk.offset == chunk.bytes.len();
if finished {
stream.inbound.pop_front();
}
stream.receive_window = stream.receive_window.saturating_add(count as u32);
(count, finished)
};
let overhead = if finished { QUEUE_ITEM_COST } else { 0 };
self.release_locked(&mut state, count + overhead, usize::from(finished), false);
if !self.queue_window_locked(&mut state, stream_id, count as u32) {
drop(state);
self.close();
return Poll::Ready(Err(io::Error::other("WEB session control budget exhausted")));
}
Poll::Ready(Ok(()))
}
/// Polls server-to-client writes against stream credit and bounded queues.
pub(super) fn poll_write(
&self,
stream_id: u32,
cx: &mut Context<'_>,
input: &[u8],
) -> Poll<io::Result<usize>> {
if input.is_empty() {
return Poll::Ready(Ok(0));
}
let mut state = self.state.lock();
let Some(stream) = state.streams.get_mut(&stream_id) else {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"WEB logical stream is closed",
)));
};
let count = input
.len()
.min(frame::DATA_CHUNK_BYTES)
.min(self.limits.max_frame_payload_bytes)
.min(stream.send_credit as usize);
if count == 0 {
stream.write_waker = Some(cx.waker().clone());
return Poll::Pending;
}
if !self.queue_data_locked(&mut state, stream_id, &input[..count]) {
if let Some(stream) = state.streams.get_mut(&stream_id) {
stream.write_waker = Some(cx.waker().clone());
}
return Poll::Pending;
}
let Some(stream) = state.streams.get_mut(&stream_id) else {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"WEB logical stream is closed",
)));
};
stream.send_credit -= count as u64;
state.last_activity = Instant::now();
drop(state);
self.down_notify.notify_waiters();
Poll::Ready(Ok(count))
}
/// Returns the process queue-capacity notification source while the manager lives.
pub(super) fn budget_notify(&self) -> Option<Arc<Notify>> {
self.manager.upgrade().map(|manager| manager.budget_notify())
}
fn release_stream_reservation(&self, peer_port: u16) {
let removed = self.state.lock().active_peer_ports.remove(&peer_port);
if removed
&& let Some(manager) = self.manager.upgrade()
{
manager.release_stream(
self.profile_key,
self.client_ip,
self.profile.public_addr,
peer_port,
);
}
}
}
fn inbound_queue_cost(queue: &VecDeque<InboundChunk>) -> (usize, usize) {
let bytes = queue.iter().fold(0usize, |total, chunk| {
total.saturating_add(chunk.bytes.len().saturating_sub(chunk.offset) + QUEUE_ITEM_COST)
});
(bytes, queue.len())
}
fn remember_closed(state: &mut SessionState, stream_id: u32, limit: usize) {
if !state.closed_streams.insert(stream_id) {
return;
}
state.closed_order.push_back(stream_id);
while state.closed_order.len() > limit {
if let Some(oldest) = state.closed_order.pop_front() {
state.closed_streams.remove(&oldest);
}
}
}
+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);
}
}
+83
View File
@@ -0,0 +1,83 @@
use std::future::Future;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::sync::futures::OwnedNotified;
use crate::web::session::WebSession;
/// Async byte stream that maps one WEB stream identifier onto carrier frames.
pub(crate) struct WebLogicalStream {
session: Arc<WebSession>,
stream_id: u32,
budget_wait: Option<Pin<Box<OwnedNotified>>>,
}
impl WebLogicalStream {
/// Binds a virtual byte stream to one live carrier stream identifier.
pub(crate) fn new(session: Arc<WebSession>, stream_id: u32) -> Self {
Self {
session,
stream_id,
budget_wait: None,
}
}
}
impl AsyncRead for WebLogicalStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
self.session.poll_read(self.stream_id, cx, output)
}
}
impl AsyncWrite for WebLogicalStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
input: &[u8],
) -> Poll<io::Result<usize>> {
let result = self.session.poll_write(self.stream_id, cx, input);
if !result.is_pending() {
self.budget_wait = None;
return result;
}
// Register before retrying so a concurrent global-capacity release cannot be lost.
loop {
if self.budget_wait.is_none()
&& let Some(notify) = self.session.budget_notify()
{
self.budget_wait = Some(Box::pin(notify.notified_owned()));
}
let Some(wait) = self.budget_wait.as_mut() else {
break;
};
if wait.as_mut().poll(cx).is_pending() {
break;
}
self.budget_wait = None;
}
match self.session.poll_write(self.stream_id, cx, input) {
Poll::Ready(result) => {
self.budget_wait = None;
Poll::Ready(result)
}
Poll::Pending => Poll::Pending,
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}