From 7239301cef258b7051333dee8f806b81ada0ed0c Mon Sep 17 00:00:00 2001 From: Paolo Insogna Date: Sun, 4 Oct 2026 15:04:01 +0200 Subject: [PATCH] fix: Upgrade only on requests or 101 responses. Signed-off-by: Paolo Insogna --- macros/src/actions.rs | 3 +- parser/tests/issue.rs | 67 ++++++++++++++++++++++++++++++++- parser/wasm/test/issues.test.js | 65 ++++++++++++++++++++++++++++++++ 3 files changed, 133 insertions(+), 2 deletions(-) diff --git a/macros/src/actions.rs b/macros/src/actions.rs index 2b11ddd..1c8f972 100644 --- a/macros/src/actions.rs +++ b/macros/src/actions.rs @@ -88,7 +88,8 @@ pub fn event_with_metadata(input: TokenStream) -> TokenStream { let status_or_method = if self.is_request { self.method as u16 } else { self.status as u16 }; let body_kind = if self.has_content_length { 0u8 } else if self.has_chunked_transfer_encoding { 1u8 } else { 2u8 }; let should_keep_alive = (!self.has_connection_close) as u8; - let should_upgrade = (self.has_upgrade && self.has_connection_upgrade) as u8; + // Responses switch protocols only with status 101; requests can propose an upgrade. + let should_upgrade = (self.has_upgrade && self.has_connection_upgrade && (self.is_request || self.status == 101)) as u8; let has_trailers = self.has_trailers as u8; let content_length = if self.has_content_length { self.content_length } else { 0 }; diff --git a/parser/tests/issue.rs b/parser/tests/issue.rs index 8d3d27d..c9ff6dd 100644 --- a/parser/tests/issue.rs +++ b/parser/tests/issue.rs @@ -1,6 +1,8 @@ mod helpers; -use milo_parser::{ERROR_NONE, Parser, STATE_ERROR}; +use milo_parser::{ + ERROR_NONE, EVENT_ACTIVE_ON_HEADERS, EVENT_HEADERS, METHOD_CONNECT, METHOD_POST, Parser, STATE_ERROR, STATE_TUNNEL, +}; use crate::helpers::{create_parser, parse}; @@ -82,3 +84,66 @@ fn issue_22__bare_lf_rejected() { assert_eq!(parser.state, STATE_ERROR); } + +#[test] +#[allow(non_snake_case)] +fn issue_25__headers_upgrade_metadata() { + let mut cases: Vec<_> = [100, 101, 103, 200, 204, 301, 304, 400, 426, 500] + .into_iter() + .map(|status| (format!("HTTP/1.1 {status} Test"), status, false, false)) + .collect(); + cases.extend([ + ("POST / HTTP/1.1".into(), METHOD_POST as u16, true, false), + ( + "CONNECT example.com:443 HTTP/1.1".into(), + METHOD_CONNECT as u16, + true, + true, + ), + ("HTTP/1.1 200 Connection Established".into(), 200, false, true), + ]); + + for (start, method_or_status, request, connect) in cases { + for upgrade in [false, true] { + let expected = upgrade && (request || method_or_status == 101); + let headers = if upgrade { + "Connection: upgrade\r\nUpgrade: h2c\r\n" + } else { + "" + }; + let message = format!("{start}\r\n{headers}\r\n"); + let mut parser = Parser::new(); + parser.autodetect = false; + parser.is_request = request; + parser.suspend_after_headers = true; + parser.active_events = EVENT_ACTIVE_ON_HEADERS; + + assert_eq!( + parser.parse(message.as_ptr(), message.len()), + message.len(), + "{message}" + ); + assert_eq!(parser.error_code, ERROR_NONE, "{message}"); + // SAFETY: The live parser owns the event buffer, and the complete headers emit + // a 19-byte event. + let events = unsafe { std::slice::from_raw_parts(parser.events, 19) }; + assert_eq!(events[0], EVENT_HEADERS, "{message}"); + assert_eq!( + u16::from_le_bytes([events[5], events[6]]), + method_or_status, + "{message}" + ); + assert_eq!(events[8], expected as u8, "{message}"); + + // Supply CONNECT response context after parsing headers, before deciding the + // body framing. + if !request && connect { + parser.is_connect = true; + } + parser.suspend_after_headers = false; + parser.parse(message.as_ptr(), 0); + assert_eq!(parser.error_code, ERROR_NONE, "{message}"); + assert_eq!(parser.state == STATE_TUNNEL, connect || expected, "{message}"); + } + } +} diff --git a/parser/wasm/test/issues.test.js b/parser/wasm/test/issues.test.js index a5e21b5..dcfd028 100644 --- a/parser/wasm/test/issues.test.js +++ b/parser/wasm/test/issues.test.js @@ -87,3 +87,68 @@ it('issue-24 - memory_deallocation', async () => { assert.equal(milo.memory.buffer.byteLength, before, `memory grew in batch ${batch + 1}`) } }) + +it('issue-25 - headers_upgrade_metadata', () => { + const cases = [ + ...[100, 101, 103, 200, 204, 301, 304, 400, 426, 500].map(status => ({ + start: `HTTP/1.1 ${status} Test`, + status, + request: false, + connect: false + })), + { start: 'POST / HTTP/1.1', request: true, connect: false }, + { start: 'CONNECT example.com:443 HTTP/1.1', request: true, connect: true }, + { start: 'HTTP/1.1 200 Connection Established', status: 200, request: false, connect: true } + ] + + for (const { start, status, request, connect } of cases) { + for (const upgrade of [false, true]) { + for (const callbacks of [false, true]) { + const label = `${start}, upgrade=${upgrade}, callbacks=${callbacks}` + const expected = upgrade && (request || status === 101) + const received = [] + const milo = setup({ + on_headers (parser, at, methodOrStatus, keepAlive, shouldUpgrade) { + received.push({ methodOrStatus, shouldUpgrade: Boolean(shouldUpgrade) }) + } + }) + const parser = milo.create() + const message = Buffer.from(`${start}\r\n${upgrade ? 'Connection: upgrade\r\nUpgrade: h2c\r\n' : ''}\r\n`) + const ptr = milo.alloc(message.length) + try { + milo.setShouldAutodetect(parser, false) + milo.setIsRequest(parser, request) + milo.setShouldSuspendAfterHeaders(parser, true) + milo.setActiveEvents(parser, callbacks ? 0n : milo.EVENT_ACTIVE_ON_HEADERS) + milo.setActiveCallbacks(parser, callbacks ? milo.CALLBACK_ACTIVE_ON_HEADERS : 0n) + new Uint8Array(milo.memory.buffer, ptr, message.length).set(message) + assert.equal(milo.parse(parser, ptr, message.length), message.length, label) + assert.equal(milo.getErrorCode(parser), milo.ERROR_NONE, label) + + const methodOrStatus = request ? (connect ? milo.METHOD_CONNECT : milo.METHOD_POST) : status + if (callbacks) { + assert.deepEqual(received, [{ methodOrStatus, shouldUpgrade: expected }], label) + } else { + const fields = new DataView(milo.memory.buffer) + const events = fields.getUint32(parser + milo.ParserFields.EVENTS, true) + assert.equal(fields.getUint8(events), milo.EVENT_HEADERS, label) + assert.equal(fields.getUint16(events + 5, true), methodOrStatus, label) + assert.equal(fields.getUint8(events + 8), Number(expected), label) + } + + // Supply CONNECT response context after parsing headers, before deciding the body framing. + if (!request && connect) { + milo.setIsConnect(parser, true) + } + milo.setShouldSuspendAfterHeaders(parser, false) + milo.parse(parser, ptr, 0) + assert.equal(milo.getErrorCode(parser), milo.ERROR_NONE, label) + assert.equal(milo.getState(parser) === milo.STATE_TUNNEL, connect || expected, label) + } finally { + milo.destroy(parser) + milo.dealloc(ptr, message.length) + } + } + } + } +})