Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 55 additions & 34 deletions h11/_headers.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,15 @@
import re
from typing import AnyStr, cast, List, overload, Sequence, Tuple, TYPE_CHECKING, Union
from typing import (
AnyStr,
cast,
List,
Optional,
overload,
Sequence,
Tuple,
TYPE_CHECKING,
Union,
)

from ._abnf import field_name, field_value
from ._util import bytesify, LocalProtocolError, validate
Expand Down Expand Up @@ -132,6 +142,41 @@ def raw_items(self) -> List[Tuple[bytes, bytes]]:
]


def _normalize_special_header(
name: bytes,
value: bytes,
seen_content_length: Optional[bytes],
saw_transfer_encoding: bool,
) -> Tuple[bytes, Optional[bytes], bool, bool]:
if name == b"content-length":
lengths = {length.strip() for length in value.split(b",")}
if len(lengths) != 1:
raise LocalProtocolError("conflicting Content-Length headers")
value = lengths.pop()
validate(_content_length_re, value, "bad Content-Length")
if len(value) > CONTENT_LENGTH_MAX_DIGITS:
raise LocalProtocolError("bad Content-Length")
if seen_content_length is None:
seen_content_length = value
elif seen_content_length != value:
raise LocalProtocolError("conflicting Content-Length headers")
else:
return value, seen_content_length, saw_transfer_encoding, False
elif name == b"transfer-encoding":
if saw_transfer_encoding:
raise LocalProtocolError(
"multiple Transfer-Encoding headers", error_status_hint=501
)
value = value.lower()
if value != b"chunked":
raise LocalProtocolError(
"Only Transfer-Encoding: chunked is supported",
error_status_hint=501,
)
saw_transfer_encoding = True
return value, seen_content_length, saw_transfer_encoding, True


@overload
def normalize_and_validate(headers: Headers, _parsed: Literal[True]) -> Headers:
...
Expand Down Expand Up @@ -169,39 +214,15 @@ def normalize_and_validate(

raw_name = name
name = name.lower()
if name == b"content-length":
lengths = {length.strip() for length in value.split(b",")}
if len(lengths) != 1:
raise LocalProtocolError("conflicting Content-Length headers")
value = lengths.pop()
validate(_content_length_re, value, "bad Content-Length")
if len(value) > CONTENT_LENGTH_MAX_DIGITS:
raise LocalProtocolError("bad Content-Length")
if seen_content_length is None:
seen_content_length = value
new_headers.append((raw_name, name, value))
elif seen_content_length != value:
raise LocalProtocolError("conflicting Content-Length headers")
elif name == b"transfer-encoding":
# "A server that receives a request message with a transfer coding
# it does not understand SHOULD respond with 501 (Not
# Implemented)."
# https://tools.ietf.org/html/rfc7230#section-3.3.1
if saw_transfer_encoding:
raise LocalProtocolError(
"multiple Transfer-Encoding headers", error_status_hint=501
)
# "All transfer-coding names are case-insensitive"
# -- https://tools.ietf.org/html/rfc7230#section-4
value = value.lower()
if value != b"chunked":
raise LocalProtocolError(
"Only Transfer-Encoding: chunked is supported",
error_status_hint=501,
)
saw_transfer_encoding = True
new_headers.append((raw_name, name, value))
else:
(
value,
seen_content_length,
saw_transfer_encoding,
keep,
) = _normalize_special_header(
name, value, seen_content_length, saw_transfer_encoding
)
if keep:
new_headers.append((raw_name, name, value))
return Headers(new_headers)

Expand Down
40 changes: 30 additions & 10 deletions h11/_readers.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@

from ._abnf import chunk_header, header_field, request_line, status_line
from ._events import Data, EndOfMessage, InformationalResponse, Request, Response
from ._headers import _normalize_special_header, Headers
from ._receivebuffer import ReceiveBuffer
from ._state import (
CLIENT,
Expand Down Expand Up @@ -61,12 +62,33 @@ def _obsolete_line_fold(lines: Iterable[bytes]) -> Iterable[bytes]:
yield last


def _decode_header_lines(
lines: Iterable[bytes],
) -> Iterable[Tuple[bytes, bytes]]:
def _decode_header_lines(lines: Iterable[bytes]) -> Headers:
full_items = []
for line in _obsolete_line_fold(lines):
matches = validate(header_field_re, line, "illegal header line: {!r}", line)
yield (matches["field_name"], matches["field_value"])
match = header_field_re.fullmatch(line)
if match is None:
raise LocalProtocolError(f"illegal header line: {line!r}")
raw_name = match["field_name"]
full_items.append((raw_name, raw_name.lower(), match["field_value"]))

seen_content_length = None
saw_transfer_encoding = False
write_index = 0
for raw_name, name, value in full_items:
(
value,
seen_content_length,
saw_transfer_encoding,
keep,
) = _normalize_special_header(
name, value, seen_content_length, saw_transfer_encoding
)
if keep:
full_items[write_index] = (raw_name, name, value)
write_index += 1

del full_items[write_index:]
return Headers(full_items)


request_line_re = re.compile(request_line.encode("ascii"))
Expand All @@ -83,9 +105,7 @@ def maybe_read_from_IDLE_client(buf: ReceiveBuffer) -> Optional[Request]:
matches = validate(
request_line_re, lines[0], "illegal request line: {!r}", lines[0]
)
return Request(
headers=list(_decode_header_lines(lines[1:])), _parsed=True, **matches
)
return Request(headers=_decode_header_lines(lines[1:]), _parsed=True, **matches)


status_line_re = re.compile(status_line.encode("ascii"))
Expand All @@ -111,7 +131,7 @@ def maybe_read_from_SEND_RESPONSE_server(
InformationalResponse if status_code < 200 else Response
)
return class_(
headers=list(_decode_header_lines(lines[1:])),
headers=_decode_header_lines(lines[1:]),
_parsed=True,
status_code=status_code,
reason=reason,
Expand Down Expand Up @@ -158,7 +178,7 @@ def __call__(self, buf: ReceiveBuffer) -> Union[Data, EndOfMessage, None]:
lines = buf.maybe_extract_lines()
if lines is None:
return None
return EndOfMessage(headers=list(_decode_header_lines(lines)))
return EndOfMessage(headers=_decode_header_lines(lines))
if self._bytes_to_discard:
data = buf.maybe_extract_at_most(len(self._bytes_to_discard))
if data is None:
Expand Down
10 changes: 10 additions & 0 deletions h11/tests/test_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -547,6 +547,16 @@ def test_reject_garbage_in_header_line() -> None:
)


def test_header_syntax_error_precedes_semantic_error() -> None:
buf = makebuf(
b"GET / HTTP/1.1\r\n" b"Host: example.com \n" b"Content-Length:\n" b"0\r\n\r\n"
)
reader = READERS[CLIENT, IDLE]
assert callable(reader)
with pytest.raises(LocalProtocolError, match="illegal header line"):
reader(buf)


def test_reject_non_vchar_in_path() -> None:
for bad_char in b"\x00\x20\x7f\xee":
message = bytearray(b"HEAD /")
Expand Down
1 change: 1 addition & 0 deletions newsfragments/207.misc.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Reduce temporary allocations and redundant header normalization when parsing headers from the wire.