Skip to content

Commit b44e2ea

Browse files
committed
Add strictly-typed class mtproto::MessageId.
1 parent e47cea5 commit b44e2ea

18 files changed

Lines changed: 356 additions & 266 deletions

‎CMakeLists.txt‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -518,6 +518,7 @@ set(TDLIB_SOURCE
518518
td/mtproto/HttpTransport.h
519519
td/mtproto/IStreamTransport.h
520520
td/mtproto/KDF.h
521+
td/mtproto/MessageId.h
521522
td/mtproto/MtprotoQuery.h
522523
td/mtproto/NoCryptoStorer.h
523524
td/mtproto/PacketInfo.h

‎benchmark/bench_misc.cpp‎

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -416,7 +416,7 @@ class IdDuplicateCheckerOld {
416416
static td::string get_description() {
417417
return "Old";
418418
}
419-
td::Status check(td::int64 message_id) {
419+
td::Status check(td::uint64 message_id) {
420420
if (saved_message_ids_.size() == MAX_SAVED_MESSAGE_IDS) {
421421
auto oldest_message_id = *saved_message_ids_.begin();
422422
if (message_id < oldest_message_id) {
@@ -437,7 +437,7 @@ class IdDuplicateCheckerOld {
437437

438438
private:
439439
static constexpr size_t MAX_SAVED_MESSAGE_IDS = 1000;
440-
std::set<td::int64> saved_message_ids_;
440+
std::set<td::uint64> saved_message_ids_;
441441
};
442442

443443
template <size_t MAX_SAVED_MESSAGE_IDS>
@@ -446,7 +446,7 @@ class IdDuplicateCheckerNew {
446446
static td::string get_description() {
447447
return PSTRING() << "New" << MAX_SAVED_MESSAGE_IDS;
448448
}
449-
td::Status check(td::int64 message_id) {
449+
td::Status check(td::uint64 message_id) {
450450
auto insert_result = saved_message_ids_.insert(message_id);
451451
if (!insert_result.second) {
452452
return td::Status::Error(1, PSLICE() << "Ignore already processed message " << message_id);
@@ -464,15 +464,15 @@ class IdDuplicateCheckerNew {
464464
}
465465

466466
private:
467-
std::set<td::int64> saved_message_ids_;
467+
std::set<td::uint64> saved_message_ids_;
468468
};
469469

470470
class IdDuplicateCheckerNewOther {
471471
public:
472472
static td::string get_description() {
473473
return "NewOther";
474474
}
475-
td::Status check(td::int64 message_id) {
475+
td::Status check(td::uint64 message_id) {
476476
if (!saved_message_ids_.insert(message_id).second) {
477477
return td::Status::Error(1, PSLICE() << "Ignore already processed message " << message_id);
478478
}
@@ -490,15 +490,15 @@ class IdDuplicateCheckerNewOther {
490490

491491
private:
492492
static constexpr size_t MAX_SAVED_MESSAGE_IDS = 1000;
493-
std::set<td::int64> saved_message_ids_;
493+
std::set<td::uint64> saved_message_ids_;
494494
};
495495

496496
class IdDuplicateCheckerNewSimple {
497497
public:
498498
static td::string get_description() {
499499
return "NewSimple";
500500
}
501-
td::Status check(td::int64 message_id) {
501+
td::Status check(td::uint64 message_id) {
502502
auto insert_result = saved_message_ids_.insert(message_id);
503503
if (!insert_result.second) {
504504
return td::Status::Error(1, "Ignore already processed message");
@@ -516,7 +516,7 @@ class IdDuplicateCheckerNewSimple {
516516

517517
private:
518518
static constexpr size_t MAX_SAVED_MESSAGE_IDS = 1000;
519-
std::set<td::int64> saved_message_ids_;
519+
std::set<td::uint64> saved_message_ids_;
520520
};
521521

522522
template <size_t max_size>
@@ -525,7 +525,7 @@ class IdDuplicateCheckerArray {
525525
static td::string get_description() {
526526
return PSTRING() << "Array" << max_size;
527527
}
528-
td::Status check(td::int64 message_id) {
528+
td::Status check(td::uint64 message_id) {
529529
if (end_pos_ == 2 * max_size) {
530530
std::copy_n(&saved_message_ids_[max_size], max_size, &saved_message_ids_[0]);
531531
end_pos_ = max_size;
@@ -550,7 +550,7 @@ class IdDuplicateCheckerArray {
550550
}
551551

552552
private:
553-
std::array<td::int64, 2 * max_size> saved_message_ids_;
553+
std::array<td::uint64, 2 * max_size> saved_message_ids_;
554554
std::size_t end_pos_ = 0;
555555
};
556556

‎td/mtproto/AuthData.cpp‎

Lines changed: 20 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@
66
//
77
#include "td/mtproto/AuthData.h"
88

9-
#include "td/utils/format.h"
109
#include "td/utils/logging.h"
1110
#include "td/utils/Random.h"
1211
#include "td/utils/SliceBuilder.h"
@@ -17,7 +16,8 @@
1716
namespace td {
1817
namespace mtproto {
1918

20-
Status check_message_id_duplicates(uint64 *saved_message_ids, size_t max_size, size_t &end_pos, uint64 message_id) {
19+
Status check_message_id_duplicates(MessageId *saved_message_ids, size_t max_size, size_t &end_pos,
20+
MessageId message_id) {
2121
// In addition, the identifiers (msg_id) of the last N messages received from the other side must be stored, and if
2222
// a message comes in with msg_id lower than all or equal to any of the stored values, that message is to be
2323
// ignored. Otherwise, the new message msg_id is added to the set, and, if the number of stored msg_id values is
@@ -32,13 +32,12 @@ Status check_message_id_duplicates(uint64 *saved_message_ids, size_t max_size, s
3232
return Status::OK();
3333
}
3434
if (end_pos >= max_size && message_id < saved_message_ids[0]) {
35-
return Status::Error(2, PSLICE() << "Ignore very old message " << format::as_hex(message_id)
36-
<< " older than the oldest known message "
37-
<< format::as_hex(saved_message_ids[0]));
35+
return Status::Error(
36+
2, PSLICE() << "Ignore very old " << message_id << " older than the oldest known " << saved_message_ids[0]);
3837
}
3938
auto it = std::lower_bound(&saved_message_ids[0], &saved_message_ids[end_pos], message_id);
4039
if (*it == message_id) {
41-
return Status::Error(1, PSLICE() << "Ignore already processed message " << format::as_hex(message_id));
40+
return Status::Error(1, PSLICE() << "Ignore already processed " << message_id);
4241
}
4342
std::copy_backward(it, &saved_message_ids[end_pos], &saved_message_ids[end_pos + 1]);
4443
*it = message_id;
@@ -105,39 +104,39 @@ std::vector<ServerSalt> AuthData::get_future_salts() const {
105104
return res;
106105
}
107106

108-
uint64 AuthData::next_message_id(double now) {
107+
MessageId AuthData::next_message_id(double now) {
109108
double server_time = get_server_time(now);
110109
auto t = static_cast<uint64>(server_time * (static_cast<uint64>(1) << 32));
111110

112111
// randomize lower bits for clocks with low precision
113112
// TODO(perf) do not do this for systems with good precision?..
114113
auto rx = Random::secure_int32();
115114
auto to_xor = rx & ((1 << 22) - 1);
116-
auto to_mul = ((rx >> 22) & 1023) + 1;
117115

118116
t ^= to_xor;
119-
auto result = t & static_cast<uint64>(-4);
117+
auto result = MessageId(t & static_cast<uint64>(-4));
120118
if (last_message_id_ >= result) {
121-
result = last_message_id_ + 8 * to_mul;
119+
auto to_mul = ((rx >> 22) & 1023) + 1;
120+
result = MessageId(last_message_id_.get() + 8 * to_mul);
122121
}
123-
LOG(DEBUG) << "Create message identifier " << format::as_hex(result) << " at " << now;
122+
LOG(DEBUG) << "Create identifier for " << result << " at " << now;
124123
last_message_id_ = result;
125124
return result;
126125
}
127126

128-
bool AuthData::is_valid_outbound_msg_id(uint64 message_id, double now) const {
127+
bool AuthData::is_valid_outbound_msg_id(MessageId message_id, double now) const {
129128
double server_time = get_server_time(now);
130-
auto id_time = static_cast<double>(message_id) / static_cast<double>(static_cast<uint64>(1) << 32);
129+
auto id_time = static_cast<double>(message_id.get()) / static_cast<double>(static_cast<uint64>(1) << 32);
131130
return server_time - 150 < id_time && id_time < server_time + 30;
132131
}
133132

134-
bool AuthData::is_valid_inbound_msg_id(uint64 message_id, double now) const {
133+
bool AuthData::is_valid_inbound_msg_id(MessageId message_id, double now) const {
135134
double server_time = get_server_time(now);
136-
auto id_time = static_cast<double>(message_id) / static_cast<double>(static_cast<uint64>(1) << 32);
135+
auto id_time = static_cast<double>(message_id.get()) / static_cast<double>(static_cast<uint64>(1) << 32);
137136
return server_time - 300 < id_time && id_time < server_time + 30;
138137
}
139138

140-
Status AuthData::check_packet(uint64 session_id, uint64 message_id, double now, bool &time_difference_was_updated) {
139+
Status AuthData::check_packet(uint64 session_id, MessageId message_id, double now, bool &time_difference_was_updated) {
141140
// Client is to check that the session_id field in the decrypted message indeed equals to that of an active session
142141
// created by the client.
143142
if (get_session_id() != session_id) {
@@ -147,22 +146,21 @@ Status AuthData::check_packet(uint64 session_id, uint64 message_id, double now,
147146

148147
// Client must check that msg_id has even parity for messages from client to server, and odd parity for messages
149148
// from server to client.
150-
if ((message_id & 1) == 0) {
151-
return Status::Error(PSLICE() << "Receive invalid message identifier " << format::as_hex(message_id));
149+
if ((message_id.get() & 1) == 0) {
150+
return Status::Error(PSLICE() << "Receive invalid " << message_id);
152151
}
153152

154153
TRY_STATUS(duplicate_checker_.check(message_id));
155154

156-
LOG(DEBUG) << "Receive packet " << format::as_hex(message_id) << " from session " << format::as_hex(session_id)
157-
<< " at " << now;
158-
time_difference_was_updated = update_server_time_difference(static_cast<uint32>(message_id >> 32) - now);
155+
LOG(DEBUG) << "Receive packet in " << message_id << " from session " << session_id << " at " << now;
156+
time_difference_was_updated = update_server_time_difference(static_cast<uint32>(message_id.get() >> 32) - now);
159157

160158
// In addition, msg_id values that belong over 30 seconds in the future or over 300 seconds in the past are to be
161159
// ignored (recall that msg_id approximately equals unixtime * 2^32). This is especially important for the server.
162160
// The client would also find this useful (to protect from a replay attack), but only if it is certain of its time
163161
// (for example, if its time has been synchronized with that of the server).
164162
if (server_time_difference_was_updated_ && !is_valid_inbound_msg_id(message_id, now)) {
165-
return Status::Error(PSLICE() << "Ignore too old or too new message " << format::as_hex(message_id));
163+
return Status::Error(PSLICE() << "Ignore too old or too new " << message_id);
166164
}
167165

168166
return Status::OK();

‎td/mtproto/AuthData.h‎

Lines changed: 12 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
#pragma once
88

99
#include "td/mtproto/AuthKey.h"
10+
#include "td/mtproto/MessageId.h"
1011

1112
#include "td/utils/common.h"
1213
#include "td/utils/Slice.h"
@@ -37,17 +38,18 @@ void parse(ServerSalt &salt, ParserT &parser) {
3738
salt.valid_until = parser.fetch_double();
3839
}
3940

40-
Status check_message_id_duplicates(uint64 *saved_message_ids, size_t max_size, size_t &end_pos, uint64 message_id);
41+
Status check_message_id_duplicates(MessageId *saved_message_ids, size_t max_size, size_t &end_pos,
42+
MessageId message_id);
4143

4244
template <size_t max_size>
4345
class MessageIdDuplicateChecker {
4446
public:
45-
Status check(uint64 message_id) {
47+
Status check(MessageId message_id) {
4648
return check_message_id_duplicates(&saved_message_ids_[0], max_size, end_pos_, message_id);
4749
}
4850

4951
private:
50-
std::array<uint64, 2 * max_size> saved_message_ids_;
52+
std::array<MessageId, 2 * max_size> saved_message_ids_;
5153
size_t end_pos_ = 0;
5254
};
5355

@@ -232,19 +234,19 @@ class AuthData {
232234

233235
std::vector<ServerSalt> get_future_salts() const;
234236

235-
uint64 next_message_id(double now);
237+
MessageId next_message_id(double now);
236238

237-
bool is_valid_outbound_msg_id(uint64 message_id, double now) const;
239+
bool is_valid_outbound_msg_id(MessageId message_id, double now) const;
238240

239-
bool is_valid_inbound_msg_id(uint64 message_id, double now) const;
241+
bool is_valid_inbound_msg_id(MessageId message_id, double now) const;
240242

241-
Status check_packet(uint64 session_id, uint64 message_id, double now, bool &time_difference_was_updated);
243+
Status check_packet(uint64 session_id, MessageId message_id, double now, bool &time_difference_was_updated);
242244

243-
Status check_update(uint64 message_id) {
245+
Status check_update(MessageId message_id) {
244246
return updates_duplicate_checker_.check(message_id);
245247
}
246248

247-
Status recheck_update(uint64 message_id) {
249+
Status recheck_update(MessageId message_id) {
248250
return updates_duplicate_rechecker_.check(message_id);
249251
}
250252

@@ -275,7 +277,7 @@ class AuthData {
275277
bool server_time_difference_was_updated_ = false;
276278
double server_time_difference_ = 0;
277279
ServerSalt server_salt_;
278-
uint64 last_message_id_ = 0;
280+
MessageId last_message_id_;
279281
int32 seq_no_ = 0;
280282
string header_;
281283
uint64 session_id_ = 0;

‎td/mtproto/CryptoStorer.h‎

Lines changed: 11 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
#pragma once
88

99
#include "td/mtproto/AuthData.h"
10+
#include "td/mtproto/MessageId.h"
1011
#include "td/mtproto/MtprotoQuery.h"
1112
#include "td/mtproto/PacketStorer.h"
1213
#include "td/mtproto/utils.h"
@@ -57,15 +58,15 @@ class ObjectImpl {
5758
bool empty() const {
5859
return !not_empty_;
5960
}
60-
uint64 get_message_id() const {
61+
MessageId get_message_id() const {
6162
return message_id_;
6263
}
6364

6465
private:
6566
bool not_empty_;
6667
Object object_;
6768
ObjectStorer object_storer_;
68-
uint64 message_id_;
69+
MessageId message_id_;
6970
int32 seq_no_;
7071
};
7172

@@ -96,7 +97,7 @@ class CancelVectorImpl {
9697
bool not_empty() const {
9798
return !storers_.empty();
9899
}
99-
uint64 get_message_id() const {
100+
MessageId get_message_id() const {
100101
CHECK(storers_.size() == 1);
101102
return storers_[0].get_message_id();
102103
}
@@ -107,7 +108,7 @@ class CancelVectorImpl {
107108

108109
class InvokeAfter {
109110
public:
110-
explicit InvokeAfter(Span<uint64> message_ids) : message_ids_(message_ids) {
111+
explicit InvokeAfter(Span<MessageId> message_ids) : message_ids_(message_ids) {
111112
}
112113
template <class StorerT>
113114
void store(StorerT &storer) const {
@@ -116,20 +117,20 @@ class InvokeAfter {
116117
}
117118
if (message_ids_.size() == 1) {
118119
storer.store_int(static_cast<int32>(0xcb9f372d));
119-
storer.store_binary(message_ids_[0]);
120+
storer.store_binary(message_ids_[0].get());
120121
return;
121122
}
122123
// invokeAfterMsgs#3dc4b4f0 {X:Type} msg_ids:Vector<long> query:!X = X;
123124
storer.store_int(static_cast<int32>(0x3dc4b4f0));
124125
storer.store_int(static_cast<int32>(0x1cb5c415));
125126
storer.store_int(narrow_cast<int32>(message_ids_.size()));
126127
for (auto message_id : message_ids_) {
127-
storer.store_binary(message_id);
128+
storer.store_binary(message_id.get());
128129
}
129130
}
130131

131132
private:
132-
Span<uint64> message_ids_;
133+
Span<MessageId> message_ids_;
133134
};
134135

135136
class QueryImpl {
@@ -206,8 +207,8 @@ class CryptoImpl {
206207
CryptoImpl(const vector<MtprotoQuery> &to_send, Slice header, vector<int64> &&to_ack, int64 ping_id, int ping_timeout,
207208
int max_delay, int max_after, int max_wait, int future_salt_n, vector<int64> get_info,
208209
vector<int64> resend, const vector<int64> &cancel, bool destroy_key, AuthData *auth_data,
209-
uint64 *container_message_id, uint64 *get_info_message_id, uint64 *resend_message_id,
210-
uint64 *ping_message_id, uint64 *parent_message_id)
210+
MessageId *container_message_id, MessageId *get_info_message_id, MessageId *resend_message_id,
211+
MessageId *ping_message_id, MessageId *parent_message_id)
211212
: query_storer_(to_send, header)
212213
, ack_empty_(to_ack.empty())
213214
, ack_storer_(!ack_empty_, mtproto_api::msgs_ack(std::move(to_ack)), auth_data)
@@ -362,7 +363,7 @@ class CryptoImpl {
362363
Mixed
363364
};
364365
Type type_;
365-
uint64 message_id_;
366+
MessageId message_id_;
366367
int32 seq_no_;
367368
};
368369

‎td/mtproto/HandshakeConnection.h‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
#include "td/mtproto/AuthKey.h"
1010
#include "td/mtproto/Handshake.h"
11+
#include "td/mtproto/MessageId.h"
1112
#include "td/mtproto/NoCryptoStorer.h"
1213
#include "td/mtproto/PacketInfo.h"
1314
#include "td/mtproto/PacketStorer.h"
@@ -61,7 +62,7 @@ class HandshakeConnection final
6162
unique_ptr<AuthKeyHandshakeContext> context_;
6263

6364
void send_no_crypto(const Storer &storer) final {
64-
raw_connection_->send_no_crypto(PacketStorer<NoCryptoImpl>(0, storer));
65+
raw_connection_->send_no_crypto(PacketStorer<NoCryptoImpl>(MessageId(), storer));
6566
}
6667

6768
Status on_raw_packet(const PacketInfo &packet_info, BufferSlice packet) final {

0 commit comments

Comments
 (0)