aboutsummaryrefslogtreecommitdiffhomepage
path: root/tde2e/td/e2e/Call.cpp
diff options
context:
space:
mode:
Diffstat (limited to 'tde2e/td/e2e/Call.cpp')
-rw-r--r--tde2e/td/e2e/Call.cpp73
1 files changed, 53 insertions, 20 deletions
diff --git a/tde2e/td/e2e/Call.cpp b/tde2e/td/e2e/Call.cpp
index f249c1f29..dfbb4659e 100644
--- a/tde2e/td/e2e/Call.cpp
+++ b/tde2e/td/e2e/Call.cpp
@@ -38,6 +38,15 @@ CallVerificationChain::State CallVerificationChain::get_state() const {
return state_;
}
+template <typename... Args>
+std::string concat(Args... args) {
+ const size_t total_size = (args.size() + ...);
+ std::string result;
+ result.reserve(total_size);
+ (result.append(args.data(), args.size()), ...);
+ return result;
+}
+
template <class F>
struct LambdaStorer {
const F &store_;
@@ -138,7 +147,7 @@ std::string CallVerificationChain::to_short_string(e2e::object_ptr<e2e::e2e_chai
return sb.as_cslice().str();
}
-td::Status CallVerificationChain::process_broadcast(std::string message,
+td::Status CallVerificationChain::process_broadcast(td::Slice message,
e2e::object_ptr<e2e::e2e_chain_GroupBroadcast> broadcast) {
td::Status status;
td::UInt256 broadcast_chain_hash{};
@@ -274,8 +283,19 @@ void CallEncryption::forget_shared_key(td::int32 epoch, td::UInt256 epoch_hash)
epochs_to_forget_.emplace(td::Timestamp::in(FORGET_EPOCH_DELAY), epoch);
}
-td::Result<std::string> CallEncryption::decrypt(td::int64 user_id, td::int32 channel_id, td::Slice encrypted_data) {
+td::Result<std::string> CallEncryption::decrypt(td::int64 user_id, td::int32 channel_id, td::Slice packet) {
sync();
+ if (packet.size() < 4) {
+ return td::Status::Error("Packet too small");
+ }
+ td::uint32 unencrypted_prefix_size = td::as<td::uint32>(packet.data() + packet.size() - 4);
+ packet.remove_suffix(4);
+ if (unencrypted_prefix_size > packet.size() || unencrypted_prefix_size >= (1 << 16)) {
+ return td::Status::Error("Unencrypted prefix size is too large");
+ }
+ auto unencrypted_prefix = packet.substr(0, unencrypted_prefix_size);
+ auto encrypted_data = packet.substr(unencrypted_prefix_size);
+
if (user_id == user_id_) {
return td::Status::Error("Packet is encrypted by us");
}
@@ -321,17 +341,25 @@ td::Result<std::string> CallEncryption::decrypt(td::int64 user_id, td::int32 cha
auto &epoch_info = it2->second;
TRY_RESULT(one_time_secret,
MessageEncryption::decrypt_header(encrypted_header, encrypted_packet, epoch_info.secret_));
- return decrypt_packet_with_secret(user_id, channel_id, unencrypted_header, encrypted_packet, one_time_secret,
- epoch_info.group_state_);
+ return decrypt_packet_with_secret(user_id, channel_id, unencrypted_header, unencrypted_prefix, encrypted_packet,
+ one_time_secret, epoch_info.group_state_);
}
}
}
return Error(E::Decrypt_UnknownEpoch);
}
-td::Result<std::string> CallEncryption::encrypt(td::int32 channel_id, td::Slice decrypted_data) {
+td::Result<std::string> CallEncryption::encrypt(td::int32 channel_id, td::Slice data,
+ size_t unencrypted_header_length) {
sync();
+ if (unencrypted_header_length > data.size() || unencrypted_header_length >= (1 << 16)) {
+ return td::Status::Error("Unencrypted header length is too large");
+ }
+ TRY_RESULT(unencrypted_header_length_u32, td::narrow_cast_safe<td::uint32>(unencrypted_header_length));
+ auto unencrypted_prefix = data.substr(0, unencrypted_header_length);
+ auto decrypted_data = data.substr(unencrypted_header_length);
+
// use all active epochs
if (epochs_.empty()) {
return Error(E::Encrypt_UnknownEpoch);
@@ -348,7 +376,8 @@ td::Result<std::string> CallEncryption::encrypt(td::int32 channel_id, td::Slice
td::SecureString one_time_secret(32, 0);
td::Random::secure_bytes(one_time_secret.as_mutable_slice());
- TRY_RESULT(encrypted_packet, encrypt_packet_with_secret(channel_id, header_a, decrypted_data, one_time_secret));
+ TRY_RESULT(encrypted_packet, encrypt_packet_with_secret(channel_id, header_a + unencrypted_prefix.str(),
+ decrypted_data, one_time_secret));
std::vector<td::SecureString> encrypted_headers;
for (auto &[epoch_i, epoch] : epochs_) {
@@ -363,14 +392,16 @@ td::Result<std::string> CallEncryption::encrypt(td::int32 channel_id, td::Slice
}
});
- //LOG(ERROR) << decrypted_data.size() << " -> " << header_a.size() << " + " << header_b.size() << " + " << encrypted_packet.size();
- return header_a + header_b + encrypted_packet;
+ std::string trailer(4, '\0');
+ td::as<td::uint32>(trailer.data()) = unencrypted_header_length_u32;
+
+ //LOG(ERROR) << decrypted_data.size() << " -> " << unencrypted_prefix.size() << " + " << header_a.size() << " + " << header_b.size() << " + " << encrypted_packet.size() << " + " << trailer.size();
+ return concat(unencrypted_prefix, header_a, header_b, encrypted_packet, trailer);
}
-std::string add_magic(td::int32 magic, td::Slice header) {
- std::string res(4 + header.size(), '\0');
+std::string make_magic(td::int32 magic) {
+ std::string res(4, '\0');
td::as<td::int32>(res.data()) = magic;
- td::MutableSlice(res).substr(4).copy_from(header);
return res;
}
@@ -393,16 +424,17 @@ td::Result<std::string> CallEncryption::encrypt_packet_with_secret(td::int32 cha
// TODO: there is too much copies happening here. Almost all of them could be avoided
td::UInt256 large_msg_id{};
auto encrypted_payload = MessageEncryption::encrypt_data(
- payload, one_time_secret, add_magic(td::e2e_api::e2e_callPacket::ID, unencrypted_part), &large_msg_id);
- auto to_sign = add_magic(td::e2e_api::e2e_callPacketLargeMsgId::ID, large_msg_id.as_slice());
+ payload, one_time_secret, concat(make_magic(td::e2e_api::e2e_callPacket::ID), unencrypted_part), &large_msg_id);
+ auto to_sign = concat(make_magic(td::e2e_api::e2e_callPacketLargeMsgId::ID), large_msg_id.as_slice());
TRY_RESULT(signature, private_key_.sign(to_sign));
return encrypted_payload.as_slice().str() + signature.to_slice().str();
}
td::Result<std::string> CallEncryption::decrypt_packet_with_secret(
- td::int64 expected_user_id, td::int32 expected_channel_id, td::Slice unencrypted_header, td::Slice encrypted_packet,
- td::Slice one_time_secret, const GroupStateRef &group_state) {
+ td::int64 expected_user_id, td::int32 expected_channel_id, td::Slice unencrypted_header,
+ td::Slice unencrypted_prefix, td::Slice encrypted_packet, td::Slice one_time_secret,
+ const GroupStateRef &group_state) {
TRY_RESULT(participant, group_state->get_participant(expected_user_id));
if (encrypted_packet.size() < 64) {
return td::Status::Error("Not enough encryption data");
@@ -411,13 +443,14 @@ td::Result<std::string> CallEncryption::decrypt_packet_with_secret(
encrypted_packet.remove_suffix(64);
td::UInt256 large_msg_id{};
- TRY_RESULT(payload_str, MessageEncryption::decrypt_data(
- encrypted_packet, one_time_secret,
- add_magic(td::e2e_api::e2e_callPacket::ID, unencrypted_header), &large_msg_id));
+ TRY_RESULT(payload_str, MessageEncryption::decrypt_data(encrypted_packet, one_time_secret,
+ concat(make_magic(td::e2e_api::e2e_callPacket::ID),
+ unencrypted_header, unencrypted_prefix),
+ &large_msg_id));
// we know that this is packet created by some participant
auto payload = td::Slice(payload_str);
- auto to_verify = add_magic(td::e2e_api::e2e_callPacketLargeMsgId::ID, large_msg_id.as_slice());
+ auto to_verify = concat(make_magic(td::e2e_api::e2e_callPacketLargeMsgId::ID), large_msg_id.as_slice());
TRY_STATUS(participant.public_key.verify(to_verify, signature));
td::TlParser parser(payload);
@@ -436,7 +469,7 @@ td::Result<std::string> CallEncryption::decrypt_packet_with_secret(
}
TRY_STATUS(check_not_seen(participant.public_key, channel_id, seqno));
mark_as_seen(participant.public_key, channel_id, seqno);
- return result;
+ return concat(unencrypted_prefix, result);
}
td::Status CallEncryption::check_not_seen(const PublicKey &public_key, td::int32 channel_id, td::uint32 seqno) {