8#include <botan/internal/tls_handshake_io.h>
10#include <botan/assert.h>
11#include <botan/exceptn.h>
12#include <botan/tls_exceptn.h>
13#include <botan/tls_handshake_msg.h>
14#include <botan/internal/fmt.h>
15#include <botan/internal/loadstor.h>
16#include <botan/internal/tls_record.h>
17#include <botan/internal/tls_seq_numbers.h>
23constexpr size_t DTLS_HANDSHAKE_HEADER_SIZE = 12;
27constexpr size_t DEFAULT_PEER_REPLAY_BUDGET = 12;
29inline size_t load_be24(
const uint8_t q[3]) {
39 throw TLS_Exception(Alert::UnexpectedMessage,
"Invalid handshake message type");
45void store_be24(uint8_t out[3],
size_t val) {
54 return Protocol_Version::TLS_V12;
62 m_queue.insert(m_queue.end(), record, record + record_len);
64 if(record_len != 1 || record[0] != 1) {
70 m_queue.insert(m_queue.end(), ccs_hs, ccs_hs +
sizeof(ccs_hs));
72 throw Decoding_Error(
"Unknown message type " + std::to_string(
static_cast<size_t>(record_type)) +
73 " in handshake processing");
78 size_t max_message_size) {
79 if(m_queue.size() >= 4) {
82 const size_t rec_length =
make_uint32(0, m_queue[1], m_queue[2], m_queue[3]);
91 throw TLS_Exception(Alert::UnexpectedMessage,
"Expected ChangeCipherSpec but got a handshake message");
94 verify_is_expected_wire_handshake_type(type);
96 if(max_message_size > 0 && rec_length > max_message_size) {
98 Alert::HandshakeFailure,
99 Botan::fmt(
"Handshake message is {} bytes, policy maximum is {}", rec_length, max_message_size));
103 const size_t length = 4 + rec_length;
105 if(m_queue.size() >= length) {
106 const std::vector<uint8_t> contents(m_queue.begin() + 4, m_queue.begin() + length);
108 m_queue.erase(m_queue.begin(), m_queue.begin() + length);
110 return std::make_pair(type, contents);
118 std::vector<uint8_t> send_buf(4 + msg.size());
120 const size_t buf_size = msg.size();
122 send_buf[0] =
static_cast<uint8_t
>(type);
124 store_be24(&send_buf[1], buf_size);
127 copy_mem(&send_buf[4], msg.data(), msg.size());
134 throw Invalid_State(
"Not possible to send under arbitrary epoch with stream based TLS");
138 const std::vector<uint8_t> msg_bits = msg.
serialize();
142 return std::vector<uint8_t>();
150#if defined(BOTAN_HAS_TLS_DOWNGRADE_SUPPORT)
152std::vector<uint8_t> Stream_Handshake_IO::start_with_client_hello_from_downgrade(
155 "Expected ClientHello message for TLS downgrade");
166size_t max_pending_reassembly(
size_t policy_hs_max) {
180 constexpr size_t overall_cap = 16 * 1024 * 1024;
182 if(policy_hs_max == 0 || policy_hs_max >= overall_cap / 4) {
187 return policy_hs_max * 4;
197 uint64_t initial_timeout_ms,
198 uint64_t max_timeout_ms,
199 std::optional<size_t> max_retransmissions,
200 size_t max_handshake_msg_size,
201 uint16_t initial_epoch) :
205 m_initial_timeout(initial_timeout_ms),
206 m_max_timeout(max_timeout_ms),
207 m_max_retransmissions(max_retransmissions),
208 m_initial_epoch(initial_epoch),
209 m_last_delivered_epoch(initial_epoch),
210 m_send_hs(std::move(writer)),
211 m_steady_clock_ms(std::move(steady_clock_ms)),
213 m_max_handshake_msg_size(max_handshake_msg_size),
214 m_max_pending_reassembly(max_pending_reassembly(m_max_handshake_msg_size)) {}
217 return Protocol_Version::DTLS_V12;
220std::optional<size_t> Datagram_Handshake_IO::last_completed_flight_index()
const {
223 const size_t flight_idx = (m_flights.size() == 1) ? 0 : (m_flights.size() - 2);
228 if(m_flights[flight_idx].empty()) {
235void Datagram_Handshake_IO::retransmit_last_flight() {
236 if(
const auto flight_idx = last_completed_flight_index()) {
237 retransmit_flight(*flight_idx);
238 m_last_write = m_steady_clock_ms();
258void Datagram_Handshake_IO::replay_last_flight_for_peer() {
262 const size_t bound = m_max_retransmissions.value_or(DEFAULT_PEER_REPLAY_BUDGET);
264 if(m_peer_replay_count >= bound) {
268 if(
const auto flight_idx = last_completed_flight_index()) {
269 m_peer_replay_count += 1;
270 retransmit_flight(*flight_idx);
274void Datagram_Handshake_IO::retransmit_flight(
size_t flight_idx) {
275 const auto& flight = m_flights.at(flight_idx);
276 const auto& ccs_records = m_flight_ccs.at(flight_idx);
277 const std::vector<uint8_t> ccs = {1};
279 BOTAN_ASSERT(!flight.empty(),
"Nonempty flight to retransmit");
282 for(
size_t msg_idx = 0; msg_idx != flight.size(); ++msg_idx) {
283 while(ccs_idx != ccs_records.size() && ccs_records[ccs_idx].first == msg_idx) {
288 const auto msg_seq = flight[msg_idx];
289 const auto& msg = m_flight_data.at(msg_seq);
290 send_message(msg_seq, msg.epoch, msg.msg_type, msg.msg_bits);
293 while(ccs_idx != ccs_records.size() && ccs_records[ccs_idx].first == flight.size()) {
302 const auto next = m_messages.find(m_in_message_seq);
303 return next != m_messages.end() && next->second.complete();
310 if(!m_flights.rbegin()->empty()) {
311 m_flights.emplace_back();
312 m_flight_ccs.emplace_back();
320 m_retransmit_terminal_flight = retransmit_terminal_flight;
325 if(!timeout || timeout->count() > 0) {
334 m_retransmit_count += 1;
335 if(m_max_retransmissions.has_value() && m_retransmit_count > m_max_retransmissions.value()) {
336 throw TLS_Exception(Alert::None,
"DTLS handshake timed out: maximum retransmissions exceeded");
342 retransmit_last_flight();
344 m_next_timeout = std::min(2 * m_next_timeout, m_max_timeout);
355 if(!m_last_write.has_value() || (m_flights.size() > 1 && !m_flights.rbegin()->empty())) {
359 const uint64_t ms_since_write = m_steady_clock_ms() - m_last_write.value();
360 if(ms_since_write >= m_next_timeout) {
361 return std::chrono::milliseconds(0);
368 const uint64_t fudge_ms = std::min<uint64_t>(15, m_initial_timeout / 2);
369 const uint64_t remaining_ms = m_next_timeout - ms_since_write;
371 return std::chrono::milliseconds(remaining_ms <= fudge_ms ? 0 : remaining_ms);
377 uint64_t record_sequence) {
378 add_record(record, record_len, record_type, record_sequence,
false);
384 uint64_t record_sequence) {
385 add_record(record, record_len, record_type, record_sequence,
true);
388bool Datagram_Handshake_IO::reassemble_retransmitted_fragment(
const uint8_t fragment[],
389 size_t fragment_length,
390 size_t fragment_offset,
394 uint16_t message_seq) {
395 auto [i, inserted] = m_retransmitted_messages.try_emplace(msg_type, message_seq, Handshake_Reassembly{});
397 if(!inserted && i->second.first != message_seq) {
398 release_reassembly_bytes(i->second.second);
399 i->second = std::make_pair(message_seq, Handshake_Reassembly());
402 auto& reassembly = i->second.second;
406 if(!charged_add_fragment(reassembly,
407 m_max_pending_reassembly,
417 if(!reassembly.complete()) {
421 release_reassembly_bytes(reassembly);
422 m_retransmitted_messages.erase(i);
426bool Datagram_Handshake_IO::charged_add_fragment(Handshake_Reassembly& reassembly,
428 const uint8_t fragment[],
429 size_t fragment_length,
430 size_t fragment_offset,
438 if(!reassembly.initialized()) {
439 if(m_pending_reassembly_bytes + msg_length > ceiling) {
442 m_pending_reassembly_bytes += msg_length;
445 reassembly.add_fragment(fragment, fragment_length, fragment_offset, epoch, msg_type, msg_length);
449void Datagram_Handshake_IO::release_reassembly_bytes(
const Handshake_Reassembly& reassembly) {
451 m_pending_reassembly_bytes -= reassembly.msg_length();
454bool Datagram_Handshake_IO::process_previous_handshake_fragment(
const uint8_t fragment[],
455 size_t fragment_length,
456 size_t fragment_offset,
460 uint16_t message_seq,
461 bool retransmitted_flight) {
465 if(fragment_length == 0 && msg_length != 0) {
473 if(!retransmitted_flight && m_awaiting_cookie_client_hello) {
474 if(!m_retransmitted_client_hello.has_value() || m_retransmitted_client_hello->first != message_seq) {
475 if(m_retransmitted_client_hello.has_value()) {
476 release_reassembly_bytes(m_retransmitted_client_hello->second);
478 m_retransmitted_client_hello = std::make_pair(message_seq, Handshake_Reassembly());
481 charged_add_fragment(m_retransmitted_client_hello->second,
482 m_max_pending_reassembly,
500 return reassemble_retransmitted_fragment(
501 fragment, fragment_length, fragment_offset, epoch, msg_type, msg_length, message_seq);
516 if(!retransmitted_flight) {
525 reassemble_retransmitted_fragment(
526 fragment, fragment_length, fragment_offset, epoch, msg_type, msg_length, message_seq)) {
529 if(m_retransmitted_ccs_epoch == epoch) {
530 m_retransmitted_ccs_epoch.reset();
537 m_retransmitted_ccs_epoch.reset();
538 m_retransmitted_finished_epoch = epoch;
547 uint64_t record_sequence,
548 bool retransmitted_flight) {
549 const uint16_t epoch =
static_cast<uint16_t
>(record_sequence >> 48);
562 if(epoch > 0 && m_finished) {
563 m_peer_replay_count = 0;
567 if(record_len != 1 || record[0] != 1) {
568 throw Decoding_Error(
"Invalid ChangeCipherSpec");
572 m_ccs_epochs.insert(epoch);
573 if(retransmitted_flight) {
577 const uint16_t finished_epoch =
static_cast<uint16_t
>(epoch + 1);
578 if(m_retransmitted_finished_epoch == finished_epoch) {
579 m_retransmitted_finished_epoch.reset();
580 if(m_retransmit_terminal_flight) {
581 replay_last_flight_for_peer();
584 m_retransmitted_finished_epoch.reset();
585 m_retransmitted_ccs_epoch = finished_epoch;
591 bool retransmit_response =
false;
593 while(record_len > 0) {
594 if(record_len < DTLS_HANDSHAKE_HEADER_SIZE) {
600 verify_is_expected_wire_handshake_type(msg_type);
602 const size_t msg_len = load_be24(&record[1]);
604 if(m_max_handshake_msg_size > 0 && msg_len > m_max_handshake_msg_size) {
606 Alert::HandshakeFailure,
607 Botan::fmt(
"Handshake message is {} bytes, policy maximum is {}", msg_len, m_max_handshake_msg_size));
611 const size_t fragment_offset = load_be24(&record[6]);
612 const size_t fragment_length = load_be24(&record[9]);
614 const size_t total_size = DTLS_HANDSHAKE_HEADER_SIZE + fragment_length;
616 if(record_len < total_size) {
617 throw Decoding_Error(
"Bad lengths in DTLS header");
621 constexpr uint16_t reassembly_window = 16;
623 if(message_seq >= m_in_message_seq && (message_seq - m_in_message_seq) < reassembly_window) {
626 if(m_in_message_seq_wrapped) {
627 record += total_size;
628 record_len -= total_size;
637 if(epoch < m_last_delivered_epoch) {
638 record += total_size;
639 record_len -= total_size;
657 record += total_size;
658 record_len -= total_size;
662 if(retransmitted_flight) {
663 if(fragment_length == 0) {
664 record += total_size;
665 record_len -= total_size;
669 throw TLS_Exception(Alert::UnexpectedMessage,
"Unexpected new DTLS handshake message");
674 if(fragment_length == 0 && msg_len > 0) {
675 record += total_size;
676 record_len -= total_size;
684 const size_t ceiling =
685 (message_seq == m_in_message_seq) ? m_max_pending_reassembly : m_max_pending_reassembly / 2;
687 auto [it, inserted] = m_messages.try_emplace(message_seq);
689 const bool accepted = charged_add_fragment(it->second,
691 &record[DTLS_HANDSHAKE_HEADER_SIZE],
697 if(!accepted && inserted) {
698 m_messages.erase(it);
700 }
else if(message_seq < m_in_message_seq) {
701 retransmit_response |= process_previous_handshake_fragment(&record[DTLS_HANDSHAKE_HEADER_SIZE],
708 retransmitted_flight);
714 record += total_size;
715 record_len -= total_size;
718 if(retransmit_response && (!m_finished || m_retransmit_terminal_flight)) {
719 replay_last_flight_for_peer();
723void Datagram_Handshake_IO::discard_stale_epoch_messages() {
724 for(
auto i = m_messages.lower_bound(m_in_message_seq); i != m_messages.end();) {
725 if(i->second.epoch() < m_last_delivered_epoch) {
726 release_reassembly_bytes(i->second);
727 i = m_messages.erase(i);
737 if(!m_flights.rbegin()->empty()) {
738 m_flights.emplace_back();
739 m_flight_ccs.emplace_back();
745 if(m_first_delivered_epoch.has_value() && m_ccs_epochs.contains(*m_first_delivered_epoch)) {
751 if(m_retransmitted_client_hello.has_value() && m_retransmitted_client_hello->second.complete()) {
752 auto result = m_retransmitted_client_hello->second.message();
753 release_reassembly_bytes(m_retransmitted_client_hello->second);
754 m_retransmitted_client_hello.reset();
755 m_recreating_hello_verify_request =
true;
759 auto i = m_messages.find(m_in_message_seq);
761 if(i == m_messages.end() || !i->second.complete()) {
765 m_in_message_seq += 1;
766 if(m_in_message_seq == 0) {
767 m_in_message_seq_wrapped =
true;
770 if(!m_first_delivered_epoch.has_value()) {
771 m_first_delivered_epoch = i->second.epoch();
774 auto result = i->second.message();
777 m_awaiting_cookie_client_hello =
false;
780 const uint16_t delivered_epoch = i->second.epoch();
782 release_reassembly_bytes(i->second);
785 if(delivered_epoch > m_last_delivered_epoch) {
786 m_last_delivered_epoch = delivered_epoch;
787 discard_stale_epoch_messages();
793void Datagram_Handshake_IO::Handshake_Reassembly::add_fragment(
const uint8_t fragment[],
794 size_t fragment_length,
795 size_t fragment_offset,
802 m_msg_type = msg_type;
803 m_msg_length = msg_length;
804 m_message.resize(msg_length);
805 m_received_mask.assign(msg_length, 0);
814 if(msg_type != m_msg_type || msg_length != m_msg_length || epoch != m_epoch) {
815 throw Decoding_Error(
"Inconsistent values in fragmented DTLS handshake header");
819 if(fragment_offset > m_msg_length) {
820 throw Decoding_Error(
"Fragment offset past end of message");
823 if(fragment_offset + fragment_length > m_msg_length) {
824 throw Decoding_Error(
"Fragment overlaps past end of message");
829 for(
size_t i = 0; i != fragment_length; ++i) {
830 const size_t off = fragment_offset + i;
831 if(m_received_mask[off] != 0) {
834 if(m_message[off] != fragment[i]) {
835 throw Decoding_Error(
"Inconsistent overlapping DTLS handshake fragment");
838 m_message[off] = fragment[i];
839 m_received_mask[off] = 1;
845bool Datagram_Handshake_IO::Handshake_Reassembly::complete()
const {
849std::pair<Handshake_Type, std::vector<uint8_t>> Datagram_Handshake_IO::Handshake_Reassembly::message()
const {
851 throw Internal_Error(
"Datagram_Handshake_IO - message not complete");
854 return std::make_pair(m_msg_type, m_message);
857std::vector<uint8_t> Datagram_Handshake_IO::format_fragment(
const uint8_t fragment[],
859 uint32_t frag_offset,
862 uint16_t msg_sequence)
const {
863 std::vector<uint8_t> send_buf(12 + frag_len);
865 send_buf[0] =
static_cast<uint8_t
>(type);
867 store_be24(&send_buf[1], msg_len);
869 store_be(msg_sequence, &send_buf[4]);
871 store_be24(&send_buf[6], frag_offset);
872 store_be24(&send_buf[9], frag_len);
875 copy_mem(&send_buf[12], fragment, frag_len);
881std::vector<uint8_t> Datagram_Handshake_IO::format_w_seq(
const std::vector<uint8_t>& msg,
883 uint16_t msg_sequence)
const {
884 return format_fragment(msg.data(), msg.size(), 0,
static_cast<uint32_t
>(msg.size()), type, msg_sequence);
893 return format_w_seq(msg, type,
static_cast<uint16_t
>(m_in_message_seq - 1));
901 const std::vector<uint8_t> msg_bits = msg.
serialize();
905 m_flight_ccs.rbegin()->emplace_back(m_flights.rbegin()->size(), epoch);
912 const uint16_t msg_seq = m_recreating_hello_verify_request ? m_out_message_seq - 1 : m_out_message_seq++;
913 m_awaiting_cookie_client_hello =
true;
914 m_recreating_hello_verify_request =
false;
915 send_message(msg_seq, epoch, msg_type, msg_bits);
919 m_flights.rbegin()->push_back(m_out_message_seq);
920 m_flight_data.insert_or_assign(m_out_message_seq, Message_Info(epoch, msg_type, msg_bits));
922 m_out_message_seq += 1;
923 m_last_write = m_steady_clock_ms();
924 m_next_timeout = m_initial_timeout;
927 m_retransmit_count = 0;
928 m_peer_replay_count = 0;
930 return send_message(m_out_message_seq - 1, epoch, msg_type, msg_bits);
933#if defined(BOTAN_HAS_TLS_DOWNGRADE_SUPPORT)
935std::vector<uint8_t> Datagram_Handshake_IO::start_with_client_hello_from_downgrade(
938 "Expected ClientHello message for DTLS downgrade");
941 return format_w_seq(client_hello.
serialize(), client_hello.
wire_type(), m_out_message_seq++);
946std::vector<uint8_t> Datagram_Handshake_IO::send_message(uint16_t msg_seq,
949 const std::vector<uint8_t>& msg_bits) {
950 auto no_fragment = format_w_seq(msg_bits, msg_type, msg_seq);
957 const size_t ciphersuite_overhead = (epoch > 0) ? 48 : 0;
963 size_t frag_offset = 0;
965 constexpr size_t DTLS_HANDSHAKE_OVERHEAD =
DTLS_HEADER_SIZE + DTLS_HANDSHAKE_HEADER_SIZE;
967 if(m_mtu <= (DTLS_HANDSHAKE_OVERHEAD + ciphersuite_overhead)) {
968 throw Invalid_Argument(
"DTLS MTU is too small to send headers");
971 const size_t max_rec_size = m_mtu - (DTLS_HANDSHAKE_OVERHEAD + ciphersuite_overhead);
973 while(frag_offset != msg_bits.size()) {
974 const size_t frag_len = std::min<size_t>(msg_bits.size() - frag_offset, max_rec_size);
976 const std::vector<uint8_t> frag = format_fragment(&msg_bits[frag_offset],
978 static_cast<uint32_t
>(frag_offset),
979 static_cast<uint32_t
>(msg_bits.size()),
985 frag_offset += frag_len;
#define BOTAN_ASSERT_NOMSG(expr)
#define BOTAN_STATE_CHECK(expr)
#define BOTAN_ARG_CHECK(expr, msg)
#define BOTAN_ASSERT(expr, assertion_made)
virtual uint16_t current_write_epoch() const =0
std::function< uint64_t()> steady_clock_fn
std::function< void(uint16_t, Record_Type, const std::vector< uint8_t > &)> writer_fn
std::vector< uint8_t > send_under_epoch(const Handshake_Message &msg, uint16_t epoch) override
std::optional< std::chrono::milliseconds > next_retransmission_timeout() const override
std::pair< Handshake_Type, std::vector< uint8_t > > get_next_record(bool expecting_ccs, size_t max_message_size) override
bool timeout_check() override
void add_record(const uint8_t record[], size_t record_len, Record_Type type, uint64_t sequence_number) override
std::vector< uint8_t > format(const std::vector< uint8_t > &handshake_msg, Handshake_Type handshake_type) const override
Protocol_Version initial_record_version() const override
void finalize_handshake(bool retransmit_terminal_flight)
void add_retransmitted_record(const uint8_t record[], size_t record_len, Record_Type type, uint64_t sequence_number)
std::vector< uint8_t > send(const Handshake_Message &msg) override
bool have_more_data() const override
Datagram_Handshake_IO(writer_fn writer, steady_clock_fn clock_ms, class Connection_Sequence_Numbers &seq, uint16_t mtu, uint64_t initial_timeout_ms, uint64_t max_timeout_ms, std::optional< size_t > max_retransmissions, size_t max_handshake_msg_size, uint16_t initial_epoch=0)
virtual Handshake_Type type() const =0
virtual std::vector< uint8_t > serialize() const =0
virtual Handshake_Type wire_type() const
std::vector< uint8_t > send_under_epoch(const Handshake_Message &msg, uint16_t epoch) override
std::vector< uint8_t > format(const std::vector< uint8_t > &handshake_msg, Handshake_Type handshake_type) const override
Protocol_Version initial_record_version() const override
std::vector< uint8_t > send(const Handshake_Message &msg) override
std::pair< Handshake_Type, std::vector< uint8_t > > get_next_record(bool expecting_ccs, size_t max_message_size) override
void add_record(const uint8_t record[], size_t record_len, Record_Type type, uint64_t sequence_number) override
constexpr uint8_t get_byte(T input)
std::string fmt(std::string_view format, const T &... args)
constexpr void copy_mem(T *out, const T *in, size_t n)
constexpr uint32_t make_uint32(uint8_t i0, uint8_t i1, uint8_t i2, uint8_t i3)
constexpr auto store_be(ParamTs &&... params)
constexpr auto load_be(ParamTs &&... params)