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