defmodule Bandit.Headers do @moduledoc false # Conveniences for dealing with headers. @spec is_port_number(integer()) :: Macro.t() defguardp is_port_number(port) when Bitwise.band(port, 0xFFFF) === port @spec get_header(Plug.Conn.headers(), header :: binary()) :: binary() | nil def get_header(headers, header) do case List.keyfind(headers, header, 0) do {_, value} -> value nil -> nil end end # Covers IPv6 addresses, like `[::1]:4000` as defined in RFC3986. @spec parse_hostlike_header(host_header :: binary()) :: {:ok, Plug.Conn.host(), nil | Plug.Conn.port_number()} | {:error, String.t()} def parse_hostlike_header("[" <> _ = host_header) do host_header |> :binary.split("]:") |> case do [host, port] -> case parse_integer(port) do {port, ""} when is_port_number(port) -> {:ok, host <> "]", port} _ -> {:error, "Header contains invalid port"} end [host] -> {:ok, host, nil} end end def parse_hostlike_header(host_header) do host_header |> :binary.split(":") |> case do [host, port] -> case parse_integer(port) do {port, ""} when is_port_number(port) -> {:ok, host, port} _ -> {:error, "Header contains invalid port"} end [host] -> {:ok, host, nil} end end @spec get_content_length(Plug.Conn.headers()) :: {:ok, nil | non_neg_integer()} | {:error, String.t()} def get_content_length(headers) do case get_header(headers, "content-length") do nil -> {:ok, nil} value -> parse_content_length(value) end end @spec get_connection_header_keys(Plug.Conn.headers()) :: {:ok, [String.t()]} | {:error, String.t()} def get_connection_header_keys(headers) do case Bandit.Headers.get_header(headers, "connection") do nil -> {:error, "Expected connection header"} value -> header_keys = value |> String.downcase() |> Plug.Conn.Utils.list() {:ok, header_keys} end end @spec parse_content_length(binary()) :: {:ok, non_neg_integer()} | {:error, String.t()} defp parse_content_length(value) do case parse_integer(value) do {length, ""} -> {:ok, length} {length, _rest} -> if value |> Plug.Conn.Utils.list() |> Enum.all?(&(&1 == to_string(length))), do: {:ok, length}, else: {:error, "invalid content-length header (RFC9112§6.3.5)"} :error -> {:error, "invalid content-length header (RFC9112§6.3.5)"} end end # Parses non-negative integers from strings. Return the valid portion of an # integer and the remaining string as a tuple like `{123, ""}` or `:error`. @spec parse_integer(String.t()) :: {non_neg_integer(), rest :: String.t()} | :error defp parse_integer(<>) when digit >= ?0 and digit <= ?9 do parse_integer(rest, digit - ?0) end defp parse_integer(_), do: :error @spec parse_integer(String.t(), non_neg_integer()) :: {non_neg_integer(), String.t()} defp parse_integer(<>, total) when digit >= ?0 and digit <= ?9 do parse_integer(rest, total * 10 + digit - ?0) end defp parse_integer(rest, total), do: {total, rest} @spec add_content_length(Plug.Conn.headers(), non_neg_integer(), Plug.Conn.int_status()) :: Plug.Conn.headers() def add_content_length(headers, length, status) do headers = Enum.reject(headers, &(elem(&1, 0) == "content-length")) if add_content_length?(status), do: [{"content-length", to_string(length)} | headers], else: headers end # Per RFC9110§8.6 @spec add_content_length?(Plug.Conn.int_status()) :: boolean() defp add_content_length?(status) when status in 100..199, do: false defp add_content_length?(204), do: false defp add_content_length?(304), do: false defp add_content_length?(_), do: true end