Page Menu
Home
Phorge
Search
Configure Global Search
Log In
Files
F85803378
D365.1791464414.diff
No One
Temporary
Actions
View File
Edit File
Delete File
View Transforms
Subscribe
Award Token
Flag For Later
Size
23 KB
Referenced Files
None
Subscribers
None
D365.1791464414.diff
View Options
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 ¤tSession = currentSessionIt->second;
- if (currentSession.ed25519Key() != ed25519Key) {
+ if ((*currentSessionIt)->ed25519Key() != ed25519Key) {
return false;
}
- return currentSession.merge(session);
+ auto res = withMapBox(inboundGroupSessions, k, [&session](InboundGroupSession ¤tSession) {
+ 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
Details
Attached
Mime Type
text/plain
Expires
Thu, Oct 8, 6:00 AM (9 h, 19 m)
Storage Engine
blob
Storage Format
Raw Data
Storage Handle
1783453
Default Alt Text
D365.1791464414.diff (23 KB)
Attached To
Mode
D365: Use immer for kazvcrypto sessions
Attached
Detach File
Event Timeline
Log In to Comment