Page MenuHomePhorge

D365.1791464414.diff
No OneTemporary

Size
23 KB
Referenced Files
None
Subscribers
None

D365.1791464414.diff

diff --git a/src/base/immer-utils.hpp b/src/base/immer-utils.hpp
--- a/src/base/immer-utils.hpp
+++ b/src/base/immer-utils.hpp
@@ -8,6 +8,7 @@
#include "libkazv-config.hpp"
#include <tuple>
+#include <optional>
namespace Kazv
{
@@ -58,4 +59,61 @@
auto newVal = std::move(func)(std::move(oldVal));
return setIn(std::move(container), std::move(newVal), std::forward<Keys>(keys)...);
}
+
+ /**
+ * Do something with the box, updating its value, then return the return value
+ * of the function.
+ *
+ * @param box The immer::box<T>.
+ * @param func A function taking T& as an argument. It is expected to modify its argument in-place.
+ * @return void if func returns void. func(T&) otherwise.
+ * If box holds a unique value, no copy of T is called.
+ */
+ template<class ImmerT, class Func, class ValueT = typename ImmerT::value_type>
+ auto withBox(ImmerT &box, Func &&func) -> std::decay_t<std::invoke_result_t<Func &&, ValueT &>>
+ {
+ using ResT = std::decay_t<std::invoke_result_t<Func &&, ValueT &>>;
+ if constexpr (std::is_same_v<ResT, void>) {
+ box = std::move(box)
+ .update([f=std::forward<Func>(func)](ValueT v) mutable {
+ std::forward<Func>(f)(v);
+ return v;
+ });
+ } else {
+ std::optional<ResT> res;
+ box = std::move(box)
+ .update([f=std::forward<Func>(func), &res](ValueT v) mutable {
+ res = std::forward<Func>(f)(v);
+ return v;
+ });
+ return std::move(res).value();
+ }
+ }
+
+ template<
+ class ImmerT,
+ class Key,
+ class Func,
+ class BoxT = typename ImmerT::mapped_type,
+ class ValueT = typename BoxT::value_type
+ >
+ auto withMapBox(ImmerT &map, Key &&k, Func &&func) -> std::decay_t<std::invoke_result_t<Func &&, ValueT &>>
+ {
+ using ResT = std::decay_t<std::invoke_result_t<Func &&, ValueT &>>;
+ if constexpr (std::is_same_v<ResT, void>) {
+ map = std::move(map)
+ .update(std::forward<Key>(k), [f=std::forward<Func>(func)](BoxT v) mutable {
+ withBox(v, std::forward<Func>(f));
+ return v;
+ });
+ } else {
+ std::optional<ResT> res;
+ map = std::move(map)
+ .update(std::forward<Key>(k), [f=std::forward<Func>(func), &res](BoxT v) mutable {
+ res = withBox(v, std::forward<Func>(f));
+ return v;
+ });
+ return std::move(res).value();
+ }
+ }
}
diff --git a/src/base/types.hpp b/src/base/types.hpp
--- a/src/base/types.hpp
+++ b/src/base/types.hpp
@@ -159,11 +159,11 @@
immer::map<T, V> ret;
if constexpr (std::is_same_v<T, std::string>) {
for (const auto &[k, v] : j.items()) {
- ret = std::move(ret).set(k, v);
+ ret = std::move(ret).set(k, v.template get<V>());
}
} else {
for (std::size_t i = 0; i < j.size(); i += 2) {
- ret = std::move(ret).set(j[i], j[i+1]);
+ ret = std::move(ret).set(j[i].template get<T>(), j[i+1].template get<V>());
}
}
m = ret;
@@ -190,6 +190,17 @@
}
};
+ template <class T>
+ struct adl_serializer<immer::box<T>> {
+ static void to_json(json& j, immer::box<T> box) {
+ j = box.get();
+ }
+
+ static void from_json(const json& j, immer::box<T> &box) {
+ box = j.template get<T>();
+ }
+ };
+
template <class T>
struct adl_serializer<immer::flex_vector<T>> {
static void to_json(json& j, immer::flex_vector<T> arr) {
diff --git a/src/crypto/crypto-p.hpp b/src/crypto/crypto-p.hpp
--- a/src/crypto/crypto-p.hpp
+++ b/src/crypto/crypto-p.hpp
@@ -9,14 +9,14 @@
#include <vodozemac.h>
-#include <unordered_map>
-
#include "crypto.hpp"
#include "crypto-util.hpp"
#include "crypto-util-p.hpp"
#include "session.hpp"
#include "inbound-group-session.hpp"
#include "outbound-group-session.hpp"
+#include <immer/map.hpp>
+#include <immer/box.hpp>
namespace Kazv
{
@@ -32,10 +32,10 @@
std::optional<rust::Box<vodozemac::olm::Account>> account;
immer::map<std::string /* algorithm */, int> uploadedOneTimeKeysCount;
int numUnpublishedKeys{0};
- std::unordered_map<std::string /* theirCurve25519IdentityKey */, Session> knownSessions;
- std::unordered_map<KeyOfGroupSession, InboundGroupSession> inboundGroupSessions;
+ immer::map<std::string /* theirCurve25519IdentityKey */, immer::box<Session>> knownSessions;
+ immer::map<KeyOfGroupSession, immer::box<InboundGroupSession>> inboundGroupSessions;
- std::unordered_map<std::string /* roomId */, OutboundGroupSession> outboundGroupSessions;
+ immer::map<std::string /* roomId */, immer::box<OutboundGroupSession>> outboundGroupSessions;
bool valid{true};
diff --git a/src/crypto/crypto.cpp b/src/crypto/crypto.cpp
--- a/src/crypto/crypto.cpp
+++ b/src/crypto/crypto.cpp
@@ -18,12 +18,14 @@
#include <event.hpp>
#include <cursorutil.hpp>
#include <types.hpp>
+#include <immer-utils.hpp>
#include <validator.hpp>
#include "crypto-p.hpp"
#include "session-p.hpp"
#include "crypto-util-p.hpp"
#include "crypto-util.hpp"
#include "time-util.hpp"
+#include <immer/map_transient.hpp>
namespace Kazv
{
@@ -101,14 +103,14 @@
auto type = content.at("ciphertext").at(ourCurve25519IdentityKey).at("type").get<int>();
auto body = content.at("ciphertext").at(ourCurve25519IdentityKey).at("body").get<std::string>();
- auto hasKnownSession = knownSessions.find(theirCurve25519IdentityKey) != knownSessions.end();
+ auto hasKnownSession = knownSessions.find(theirCurve25519IdentityKey);
if (type == 0) { // pre-key message
bool shouldCreateNewSession =
// there is no possible session
(! hasKnownSession)
// the possible session does not match this message
- || (! knownSessions.at(theirCurve25519IdentityKey).matches(body));
+ || (! knownSessions.at(theirCurve25519IdentityKey)->matches(body));
if (shouldCreateNewSession) {
auto created = createInboundSession(theirCurve25519IdentityKey, body);
@@ -117,20 +119,25 @@
return NotBut("Cannot create session");
}
- auto &session = knownSessions.at(theirCurve25519IdentityKey);
- return session.m_d->takeFirstDecrypted();
+ auto res = withMapBox(knownSessions, theirCurve25519IdentityKey, [](Session &session) {
+ return session.m_d->takeFirstDecrypted();
+ });
+ return res;
}
- auto &session = knownSessions.at(theirCurve25519IdentityKey);
-
- return session.decrypt(type, body);
+ auto res = withMapBox(knownSessions, theirCurve25519IdentityKey, [&](Session &session) {
+ return session.decrypt(type, body);
+ });
+ return res;
} else {
if (! hasKnownSession) {
return NotBut("No available session");
}
- auto &session = knownSessions.at(theirCurve25519IdentityKey);
- return session.decrypt(type, body);
+ auto res = withMapBox(knownSessions, theirCurve25519IdentityKey, [&](Session &session) {
+ return session.decrypt(type, body);
+ });
+ return res;
}
}
@@ -143,15 +150,16 @@
auto k = KeyOfGroupSession{roomId, sessionId};
- if (inboundGroupSessions.find(k) == inboundGroupSessions.end()) {
+ if (!inboundGroupSessions.find(k)) {
return NotBut("We do not have the keys for this");
} else {
auto msg = content.at("ciphertext").get<std::string>();
auto eventId = eventJson.at("event_id").get<std::string>();
auto originServerTs = eventJson.at("origin_server_ts").get<Timestamp>();
- auto &session = inboundGroupSessions.at(k);
-
- return session.decrypt(msg, eventId, originServerTs);
+ auto res = withMapBox(inboundGroupSessions, k, [&](InboundGroupSession &session) {
+ return session.decrypt(msg, eventId, originServerTs);
+ });
+ return res;
}
}
@@ -162,7 +170,7 @@
theirCurve25519IdentityKey, message);
if (s.valid()) {
- knownSessions.insert_or_assign(theirCurve25519IdentityKey, std::move(s));
+ knownSessions = std::move(knownSessions).set(theirCurve25519IdentityKey, std::move(s));
return true;
}
@@ -178,23 +186,23 @@
valid = false;
} else {
auto it = outboundGroupSessions.find(roomId);
- if (it == outboundGroupSessions.end()) {
+ if (!it) {
valid = false;
} else {
- auto &session = it->second;
- if (timeMs - session.creationTimeMs() >= desc.value().ms) {
+ auto sessionBox = *it;
+ if (timeMs - sessionBox->creationTimeMs() >= desc.value().ms) {
valid = false;
- } else if (session.messageIndex() >= desc.value().messages) {
+ } else if (sessionBox->messageIndex() >= desc.value().messages) {
valid = false;
}
}
}
if (! valid) {
- outboundGroupSessions.insert_or_assign(roomId, OutboundGroupSession(RandomTag{}, random, timeMs));
- auto &session = outboundGroupSessions.at(roomId);
- auto sessionId = session.sessionId();
- auto sessionKey = session.sessionKey();
+ outboundGroupSessions = std::move(outboundGroupSessions).set(roomId, OutboundGroupSession(RandomTag{}, random, timeMs));
+ auto sessionBox = outboundGroupSessions.at(roomId);
+ auto sessionId = sessionBox->sessionId();
+ auto sessionKey = sessionBox->sessionKey();
auto k = KeyOfGroupSession{roomId, sessionId};
@@ -412,7 +420,7 @@
bool Crypto::hasInboundGroupSession(KeyOfGroupSession k) const
{
- return m_d->inboundGroupSessions.find(k) != m_d->inboundGroupSessions.end();
+ return m_d->inboundGroupSessions.find(k);
}
bool CryptoPrivate::createInboundGroupSession(KeyOfGroupSession k, std::string sessionKey, std::string ed25519Key)
@@ -424,19 +432,21 @@
}
auto currentSessionIt = inboundGroupSessions.find(k);
- if (currentSessionIt == inboundGroupSessions.end()) {
+ if (!currentSessionIt) {
// the session is new, insert it
- inboundGroupSessions.insert({k, std::move(session)});
+ inboundGroupSessions = std::move(inboundGroupSessions).set(k, std::move(session));
return true;
}
// the session already exists, do some check
- auto &currentSession = currentSessionIt->second;
- if (currentSession.ed25519Key() != ed25519Key) {
+ if ((*currentSessionIt)->ed25519Key() != ed25519Key) {
return false;
}
- return currentSession.merge(session);
+ auto res = withMapBox(inboundGroupSessions, k, [&session](InboundGroupSession &currentSession) {
+ return currentSession.merge(session);
+ });
+ return res;
}
bool Crypto::verify(nlohmann::json object, std::string userId, std::string deviceId, std::string ed25519Key)
@@ -475,11 +485,11 @@
auto k = KeyOfGroupSession{roomId, sessionId};
- if (m_d->inboundGroupSessions.find(k) == m_d->inboundGroupSessions.end()) {
+ if (!m_d->inboundGroupSessions.find(k)) {
return NotBut("We do not have the keys for this");
} else {
- auto &session = m_d->inboundGroupSessions.at(k);
- return session.ed25519Key();
+ auto sessionBox = m_d->inboundGroupSessions.at(k);
+ return sessionBox->ed25519Key();
}
}
@@ -515,8 +525,9 @@
{
assert(random.size() >= encryptOlmRandomSize(theirCurve25519IdentityKey));
try {
- auto &session = m_d->knownSessions.at(theirCurve25519IdentityKey);
- auto [type, body] = session.encryptWithRandom(random, eventJson.dump());
+ auto [type, body] = withMapBox(m_d->knownSessions, theirCurve25519IdentityKey, [&](Session &session) {
+ return session.encryptWithRandom(random, eventJson.dump());
+ });
return nlohmann::json{
{
theirCurve25519IdentityKey, {
@@ -543,9 +554,9 @@
auto textToEncrypt = std::move(jsonToEncrypt).dump();
- auto &session = m_d->outboundGroupSessions.at(roomId);
-
- auto ciphertext = session.encrypt(std::move(textToEncrypt));
+ auto [sessionId, ciphertext] = withMapBox(m_d->outboundGroupSessions, roomId, [&textToEncrypt](OutboundGroupSession &session) {
+ return std::make_pair(session.sessionId(), session.encrypt(std::move(textToEncrypt)));
+ });
return
json{
@@ -554,7 +565,7 @@
// as per the Matrix spec
{"sender_key", curve25519IdentityKey()},
{"ciphertext", ciphertext},
- {"session_id", session.sessionId()},
+ {"session_id", sessionId},
};
}
@@ -581,14 +592,14 @@
std::string Crypto::outboundGroupSessionInitialKey(std::string roomId)
{
- auto &session = m_d->outboundGroupSessions.at(roomId);
- return session.initialSessionKey();
+ auto sessionBox = m_d->outboundGroupSessions.at(roomId);
+ return sessionBox->initialSessionKey();
}
std::string Crypto::outboundGroupSessionCurrentKey(std::string roomId)
{
- auto &session = m_d->outboundGroupSessions.at(roomId);
- return session.sessionKey();
+ auto sessionBox = m_d->outboundGroupSessions.at(roomId);
+ return sessionBox->sessionKey();
}
auto Crypto::devicesMissingOutboundSessionKey(
@@ -601,8 +612,7 @@
intoImmer(immer::flex_vector<std::string>{},
zug::filter([this](auto kv) {
auto [deviceId, theirCurve25519IdentityKey] = kv;
- return m_d->knownSessions.find(theirCurve25519IdentityKey)
- == m_d->knownSessions.end();
+ return !m_d->knownSessions.find(theirCurve25519IdentityKey);
})
| zug::map([](auto kv) {
auto [deviceId, key] = kv;
@@ -635,8 +645,8 @@
theirOneTimeKey);
if (session.valid()) {
- m_d->knownSessions.insert_or_assign(theirIdentityKey,
- std::move(session));
+ m_d->knownSessions = std::move(m_d->knownSessions)
+ .set(theirIdentityKey, std::move(session));
}
}
@@ -645,7 +655,7 @@
std::string pickledData = m_d->valid ? m_d->pickle() : std::string();
auto j = nlohmann::json::object({
{"valid", m_d->valid},
- {"version", 1},
+ {"version", 2},
{"account", std::move(pickledData)},
{"uploadedOneTimeKeysCount", m_d->uploadedOneTimeKeysCount},
{"numUnpublishedKeys", m_d->numUnpublishedKeys},
@@ -657,12 +667,22 @@
return j;
}
+ template<class Sess, class K>
+ immer::map<K, immer::box<Sess>> sessionStdMapToImmer(std::unordered_map<K, Sess> &&stdMap)
+ {
+ auto res = immer::map_transient<K, immer::box<Sess>>{};
+ for (auto &&[k, v] : stdMap) {
+ res.set(k, std::move(v));
+ }
+ return res.persistent();
+ }
+
void Crypto::loadJson(const nlohmann::json &j)
{
m_d->valid = j.contains("valid") ? j["valid"].template get<bool>() : true;
const auto &pickledData = j.at("account").template get<std::string>();
if (m_d->valid) {
- if (j.contains("version") && j["version"] == 1) {
+ if (j.contains("version") && j["version"] >= 1) {
m_d->valid = m_d->unpickle(pickledData);
} else {
m_d->valid = m_d->unpickleFromLibolm(pickledData);
@@ -671,8 +691,14 @@
m_d->uploadedOneTimeKeysCount = j.at("uploadedOneTimeKeysCount");
m_d->numUnpublishedKeys = j.at("numUnpublishedKeys");
- m_d->knownSessions = j.at("knownSessions").template get<decltype(m_d->knownSessions)>();
- m_d->inboundGroupSessions = j.at("inboundGroupSessions").template get<decltype(m_d->inboundGroupSessions)>();
- m_d->outboundGroupSessions = j.at("outboundGroupSessions").template get<decltype(m_d->outboundGroupSessions)>();
+ if (j.contains("version") && j["version"] >= 2) {
+ m_d->knownSessions = j.at("knownSessions").template get<decltype(m_d->knownSessions)>();
+ m_d->inboundGroupSessions = j.at("inboundGroupSessions").template get<decltype(m_d->inboundGroupSessions)>();
+ m_d->outboundGroupSessions = j.at("outboundGroupSessions").template get<decltype(m_d->outboundGroupSessions)>();
+ } else {
+ m_d->knownSessions = sessionStdMapToImmer(j.at("knownSessions").template get<std::unordered_map<std::string, Session>>());
+ m_d->inboundGroupSessions = sessionStdMapToImmer(j.at("inboundGroupSessions").template get<std::unordered_map<KeyOfGroupSession, InboundGroupSession>>());
+ m_d->outboundGroupSessions = sessionStdMapToImmer(j.at("outboundGroupSessions").template get<std::unordered_map<std::string, OutboundGroupSession>>());
+ }
}
}
diff --git a/src/crypto/outbound-group-session-p.hpp b/src/crypto/outbound-group-session-p.hpp
--- a/src/crypto/outbound-group-session-p.hpp
+++ b/src/crypto/outbound-group-session-p.hpp
@@ -37,7 +37,7 @@
bool unpickle(std::string pickleData);
bool unpickleFromLibolm(std::string pickleData);
- std::string sessionKey();
+ std::string sessionKey() const;
};
}
diff --git a/src/crypto/outbound-group-session.hpp b/src/crypto/outbound-group-session.hpp
--- a/src/crypto/outbound-group-session.hpp
+++ b/src/crypto/outbound-group-session.hpp
@@ -47,11 +47,11 @@
bool valid() const;
- std::string sessionKey();
+ std::string sessionKey() const;
std::string initialSessionKey() const;
- std::string sessionId();
+ std::string sessionId() const;
- int messageIndex();
+ int messageIndex() const;
Timestamp creationTimeMs() const;
private:
diff --git a/src/crypto/outbound-group-session.cpp b/src/crypto/outbound-group-session.cpp
--- a/src/crypto/outbound-group-session.cpp
+++ b/src/crypto/outbound-group-session.cpp
@@ -126,14 +126,14 @@
return static_cast<std::string>(res.value()->to_base64());
}
- std::string OutboundGroupSessionPrivate::sessionKey()
+ std::string OutboundGroupSessionPrivate::sessionKey() const
{
auto key = session.value()->session_key()->to_base64();
return static_cast<std::string>(key);
}
- std::string OutboundGroupSession::sessionKey()
+ std::string OutboundGroupSession::sessionKey() const
{
return m_d->sessionKey();
}
@@ -143,13 +143,13 @@
return m_d->initialSessionKey;
}
- std::string OutboundGroupSession::sessionId()
+ std::string OutboundGroupSession::sessionId() const
{
auto id = m_d->session.value()->session_id();
return static_cast<std::string>(id);
}
- int OutboundGroupSession::messageIndex()
+ int OutboundGroupSession::messageIndex() const
{
return m_d->session.value()->message_index();
}
diff --git a/src/crypto/session.hpp b/src/crypto/session.hpp
--- a/src/crypto/session.hpp
+++ b/src/crypto/session.hpp
@@ -58,7 +58,7 @@
Session &operator=(Session &&that);
~Session();
- bool matches(std::string message);
+ bool matches(std::string message) const;
bool valid() const;
diff --git a/src/crypto/session.cpp b/src/crypto/session.cpp
--- a/src/crypto/session.cpp
+++ b/src/crypto/session.cpp
@@ -95,6 +95,7 @@
{
if (that.valid) {
valid = unpickle(that.pickle());
+ firstDecrypted = that.firstDecrypted;
}
}
@@ -198,7 +199,7 @@
}
- bool Session::matches(std::string message)
+ bool Session::matches(std::string message) const
{
auto res = checkVodozemacError([&]() {
auto msg = vodozemac::olm::olm_message_from_parts(vodozemac::olm::OlmMessageParts{
diff --git a/src/tests/base/immer-utils-test.cpp b/src/tests/base/immer-utils-test.cpp
--- a/src/tests/base/immer-utils-test.cpp
+++ b/src/tests/base/immer-utils-test.cpp
@@ -14,6 +14,7 @@
#include <immer/map.hpp>
#include <immer/flex_vector.hpp>
+#include <immer/box.hpp>
#include <immer-utils.hpp>
@@ -77,3 +78,82 @@
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
});
}
+
+struct TestTypeA
+{
+ int a{0};
+ TestTypeA() {}
+ TestTypeA(int c) : a(c) {}
+ TestTypeA(const TestTypeA &that) : a(that.a + 1) {}
+ TestTypeA(TestTypeA &&that) = default;
+ TestTypeA &operator=(const TestTypeA &that)
+ {
+ a = that.a + 1;
+ return *this;
+ }
+ TestTypeA &operator=(TestTypeA &&that) = default;
+};
+
+TEST_CASE("withBox()", "[base][immer-utils]")
+{
+ immer::box<TestTypeA> box = TestTypeA(2);
+ auto res = withBox(box, [](TestTypeA &t) {
+ REQUIRE(t.a == 2);
+ t.a = 4;
+ return t.a;
+ });
+ REQUIRE(box->a == 4);
+ REQUIRE(res == 4);
+
+ withBox(box, [](TestTypeA &t) {
+ REQUIRE(t.a == 4);
+ t.a = 6;
+ });
+ REQUIRE(box->a == 6);
+}
+
+TEST_CASE("withBox() non-unique", "[base][immer-utils]")
+{
+ immer::box<TestTypeA> box = TestTypeA(2);
+ auto box2 = box;
+ auto res = withBox(box, [](TestTypeA &t) {
+ REQUIRE(t.a == 3);
+ t.a = 5;
+ return t.a;
+ });
+ REQUIRE(box->a == 5);
+ REQUIRE(res == 5);
+}
+
+struct TestTypeB
+{
+ int num{0};
+};
+
+TEST_CASE("withMapBox()", "[base][immer-utils]")
+{
+ immer::map<std::string, immer::box<TestTypeB>> map = {
+ {"a", TestTypeB{1}},
+ {"b", TestTypeB{2}},
+ };
+
+ withMapBox(map, "a", [](TestTypeB &t) {
+ REQUIRE(t.num == 1);
+ t.num = 3;
+ });
+ REQUIRE(map["a"]->num == 3);
+
+ auto res = withMapBox(map, "b", [](TestTypeB &t) {
+ REQUIRE(t.num == 2);
+ t.num = 6;
+ return 8;
+ });
+ REQUIRE(map["b"]->num == 6);
+ REQUIRE(res == 8);
+
+ withMapBox(map, "c", [](TestTypeB &t) {
+ REQUIRE(t.num == 0);
+ t.num = 9;
+ });
+ REQUIRE(map.at("c")->num == 9);
+}

File Metadata

Mime Type
text/plain
Expires
Thu, Oct 8, 6:00 AM (11 h, 35 m)
Storage Engine
blob
Storage Format
Raw Data
Storage Handle
1783453
Default Alt Text
D365.1791464414.diff (23 KB)

Event Timeline