Page Menu
Home
Phorge
Search
Configure Global Search
Log In
Files
F85804091
No One
Temporary
Actions
View File
Edit File
Delete File
View Transforms
Subscribe
Award Token
Flag For Later
Size
67 KB
Referenced Files
None
Subscribers
None
View Options
diff --git a/src/base/immer-utils.hpp b/src/base/immer-utils.hpp
index 2022479..a1cba12 100644
--- a/src/base/immer-utils.hpp
+++ b/src/base/immer-utils.hpp
@@ -1,61 +1,119 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include "libkazv-config.hpp"
#include <tuple>
+#include <optional>
namespace Kazv
{
template<class T>
[[nodiscard]] auto getIn(T &&item)
{
return std::forward<T>(item);
}
template<class ImmerT, class K, class ...Keys>
[[nodiscard]] auto getIn(ImmerT &&container, K &&key, Keys &&...keys)
{
return getIn(std::forward<ImmerT>(container)[std::forward<K>(key)], std::forward<Keys>(keys)...);
}
template<class T>
[[nodiscard]] auto setIn(T, T newVal)
{
return newVal;
}
template<class ImmerT, class K>
[[nodiscard]] auto setIn(ImmerT container, std::decay_t<decltype(getIn(std::declval<ImmerT>(), std::declval<K>()))> newVal, K &&key) -> std::decay_t<ImmerT>
{
return std::move(container).set(std::forward<K>(key), std::move(newVal));
}
template<class ImmerT, class K, class ...Keys>
[[nodiscard]] auto setIn(ImmerT container, std::decay_t<decltype(getIn(std::declval<ImmerT>(), std::declval<K>(), std::declval<Keys>()...))> newVal, K &&key, Keys &&...keys) -> std::decay_t<ImmerT>
{
auto oldItem = getIn(container, key);
return std::move(container).set(
std::forward<K>(key),
setIn(oldItem, std::move(newVal), std::forward<Keys>(keys)...)
);
}
template<class T, class Func>
[[nodiscard]] auto updateIn(T oldVal, Func func) -> T
{
return std::move(func)(oldVal);
}
template<class ImmerT, class Func, class ...Keys>
[[nodiscard]] auto updateIn(ImmerT container, Func func, Keys &&...keys) -> ImmerT
{
auto oldVal = getIn(container, keys...);
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
index c23eaf2..b2c43d2 100644
--- a/src/base/types.hpp
+++ b/src/base/types.hpp
@@ -1,233 +1,244 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2020-2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include "libkazv-config.hpp"
#include <optional>
#include <string>
#include <variant>
#include <nlohmann/json.hpp>
#include <immer/array.hpp>
#include <immer/flex_vector.hpp>
#include <immer/map.hpp>
#include <boost/hana/type.hpp>
#include <lager/util.hpp>
#include "jsonwrap.hpp"
#include "event.hpp"
namespace Kazv
{
using Bytes = std::string;
enum Status : bool
{
FAIL,
SUCC,
};
namespace detail
{
constexpr auto hasEmptyMethod = boost::hana::is_valid(
[](auto t) -> decltype((void)std::declval<typename decltype(t)::type>().empty()) {});
template<class U>
struct AddToJsonIfNeededT
{
template<class T>
static void call(json &j, std::string name, T &&arg) {
using Type = std::decay_t<T>;
if constexpr (detail::hasEmptyMethod(boost::hana::type_c<Type>)) {
if (! arg.empty()) {
j[name] = std::forward<T>(arg);
}
} else {
j[name] = std::forward<T>(arg);
}
}
};
template<class U>
struct AddToJsonIfNeededT<std::optional<U>>
{
template<class T>
static void call(json &j, std::string name, T &&arg) {
if (arg.has_value()) {
j[name] = std::forward<T>(arg).value();
}
}
};
template<>
struct AddToJsonIfNeededT<JsonWrap>
{
template<class T>
static void call(json &j, std::string name, T &&arg) {
if (!arg.get().is_null()) {
j[name] = std::forward<T>(arg).get();
}
}
};
}
template<class T>
inline void addToJsonIfNeeded(json &j, std::string name, T &&arg)
{
detail::AddToJsonIfNeededT<std::decay_t<T>>::call(j, name, std::forward<T>(arg));
};
// Provide a non-destructive way to add the map
// to json.
template<class MapT,
// disallow json object here
std::enable_if_t<!std::is_same_v<std::decay_t<MapT>, json>
&& !std::is_same_v<std::decay_t<MapT>, JsonWrap>, int> = 0>
inline void addPropertyMapToJson(json &j, MapT &&arg)
{
for (auto kv : std::forward<MapT>(arg)) {
auto [k, v] = kv;
j[k] = v;
}
};
inline void addPropertyMapToJson(json &j, const json &arg)
{
for (auto kv : arg.items()) {
auto [k, v] = kv;
j[k] = v;
}
};
using EventList = immer::flex_vector<Event>;
using namespace std::string_literals;
struct Null {};
using Variant = std::variant<std::string, JsonWrap, Null>;
namespace detail
{
struct DefaultValT
{
template<class T>
constexpr operator T() const {
return T();
}
};
}
constexpr detail::DefaultValT DEFVAL;
enum RoomMembership
{
Invite, Join, Leave
};
namespace detail
{
// emulates declval() but returns lvalue reference
template<class T>
typename std::add_lvalue_reference<T>::type declref() noexcept;
}
}
namespace nlohmann {
template <class T, class V>
struct adl_serializer<immer::map<T, V>> {
static void to_json(json& j, immer::map<T, V> map) {
if constexpr (std::is_same_v<T, std::string>) {
j = json::object();
for (auto [k, v] : map) {
j[k] = v;
}
} else {
j = json::array();
for (auto [k, v] : map) {
j.push_back(k);
j.push_back(v);
}
}
}
static void from_json(const json& j, immer::map<T, V> &m) {
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;
}
};
template <class T>
struct adl_serializer<immer::array<T>> {
static void to_json(json& j, immer::array<T> arr) {
j = json::array();
for (auto i : arr) {
j.push_back(json(i));
}
}
static void from_json(const json& j, immer::array<T> &a) {
immer::array<T> ret;
if (j.is_array()) {
for (const auto &i : j) {
ret = std::move(ret).push_back(i);
}
}
a = ret;
}
};
+ 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) {
j = json::array();
for (auto i : arr) {
j.push_back(json(i));
}
}
static void from_json(const json& j, immer::flex_vector<T> &a) {
immer::flex_vector<T> ret;
if (j.is_array()) {
for (const auto &i : j) {
ret = std::move(ret).push_back(i.get<T>());
}
}
a = ret;
}
};
template <>
struct adl_serializer<Kazv::Variant> {
static void to_json(json& j, const Kazv::Variant &var) {
std::visit(lager::visitor{
[&j](std::string i) { j = i; },
[&j](Kazv::JsonWrap i) { j = i; },
[&j](Kazv::Null) { j = nullptr; }
}, var);
}
static void from_json(const json& j, Kazv::Variant &var) {
if (j.is_string()) {
var = j.get<std::string>();
} else if (j.is_null()) {
var = Kazv::Null{};
} else { // is object
var = Kazv::Variant(Kazv::JsonWrap(j));
}
}
};
}
diff --git a/src/crypto/crypto-p.hpp b/src/crypto/crypto-p.hpp
index ad6e147..da77f7b 100644
--- a/src/crypto/crypto-p.hpp
+++ b/src/crypto/crypto-p.hpp
@@ -1,63 +1,63 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021-2024 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include <libkazv-config.hpp>
#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
{
using SessionList = std::vector<Session>;
struct CryptoPrivate
{
CryptoPrivate();
CryptoPrivate(RandomTag, RandomData data);
CryptoPrivate(const CryptoPrivate &that);
~CryptoPrivate();
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};
std::string pickle() const;
bool unpickle(std::string data);
bool unpickleFromLibolm(std::string data);
std::string ed25519IdentityKey() const;
std::string curve25519IdentityKey() const;
MaybeString decryptOlm(nlohmann::json content);
// Here we need the full event for eventId and originServerTs
MaybeString decryptMegOlm(nlohmann::json eventJson);
/// returns whether the session is successfully established
bool createInboundSession(std::string theirCurve25519IdentityKey,
std::string message);
bool createInboundGroupSession(KeyOfGroupSession k, std::string sessionKey, std::string ed25519Key);
bool reuseOrCreateOutboundGroupSession(RandomData random, Timestamp timeMs,
std::string roomId, std::optional<MegOlmSessionRotateDesc> desc);
};
}
diff --git a/src/crypto/crypto.cpp b/src/crypto/crypto.cpp
index 85237a8..73206a6 100644
--- a/src/crypto/crypto.cpp
+++ b/src/crypto/crypto.cpp
@@ -1,678 +1,704 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2020-2024 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <vector>
#include <zug/transducer/filter.hpp>
#include <vodozemac.h>
#include <nlohmann/json.hpp>
#include <debug.hpp>
#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
{
using namespace CryptoConstants;
CryptoPrivate::CryptoPrivate()
: account(std::nullopt)
, valid(false)
{
}
CryptoPrivate::CryptoPrivate(RandomTag, [[maybe_unused]] RandomData data)
: account(std::nullopt)
, valid(true)
{
account = vodozemac::olm::new_account();
}
CryptoPrivate::~CryptoPrivate()
{
}
CryptoPrivate::CryptoPrivate(const CryptoPrivate &that)
: account(std::nullopt)
, uploadedOneTimeKeysCount(that.uploadedOneTimeKeysCount)
, numUnpublishedKeys(that.numUnpublishedKeys)
, knownSessions(that.knownSessions)
, inboundGroupSessions(that.inboundGroupSessions)
, outboundGroupSessions(that.outboundGroupSessions)
{
if (that.valid) {
valid = unpickle(that.pickle());
}
}
std::string CryptoPrivate::pickle() const
{
auto pickleData = account.value()->pickle(VODOZEMAC_PICKLE_KEY);
return static_cast<std::string>(pickleData);
}
bool CryptoPrivate::unpickle(std::string pickleData)
{
account = checkVodozemacError([&]() {
return vodozemac::olm::account_from_pickle(
rust::Str(pickleData),
VODOZEMAC_PICKLE_KEY
);
});
return account.has_value();
}
bool CryptoPrivate::unpickleFromLibolm(std::string pickleData)
{
account = checkVodozemacError([&]() {
return vodozemac::olm::account_from_libolm_pickle(
rust::Str(pickleData),
rust::Slice<const unsigned char>(OLM_PICKLE_KEY.data(), OLM_PICKLE_KEY.size())
);
});
return account.has_value();
}
MaybeString CryptoPrivate::decryptOlm(nlohmann::json content)
{
auto theirCurve25519IdentityKey = content.at("sender_key").get<std::string>();
auto ourCurve25519IdentityKey = curve25519IdentityKey();
if (! content.at("ciphertext").contains(ourCurve25519IdentityKey)) {
return NotBut("Message not intended for us");
}
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);
if (! created) { // cannot create session, thus cannot decrypt
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;
}
}
MaybeString CryptoPrivate::decryptMegOlm(nlohmann::json eventJson)
{
auto content = eventJson.at("content");
auto sessionId = content.at("session_id").get<std::string>();
auto roomId = eventJson.at("room_id").get<std::string>();
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;
}
}
bool CryptoPrivate::createInboundSession(std::string theirCurve25519IdentityKey,
std::string message)
{
auto s = Session(InboundSessionTag{}, *this,
theirCurve25519IdentityKey, message);
if (s.valid()) {
- knownSessions.insert_or_assign(theirCurve25519IdentityKey, std::move(s));
+ knownSessions = std::move(knownSessions).set(theirCurve25519IdentityKey, std::move(s));
return true;
}
return false;
}
bool CryptoPrivate::reuseOrCreateOutboundGroupSession(
RandomData random, Timestamp timeMs,
std::string roomId, std::optional<MegOlmSessionRotateDesc> desc)
{
bool valid = true;
if (! desc.has_value()) { // force rotate
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};
if (! createInboundGroupSession(k, sessionKey, ed25519IdentityKey())) {
kzo.client.warn() << "Create inbound group session from outbound group session failed. We may not be able to read our own messages." << std::endl;
}
}
return valid;
}
std::size_t Crypto::constructRandomSize()
{
return 0;
}
Crypto::Crypto()
: m_d(new CryptoPrivate{})
{
}
Crypto::Crypto(RandomTag, RandomData data)
: m_d(new CryptoPrivate(RandomTag{}, std::move(data)))
{
}
Crypto::~Crypto() = default;
Crypto::Crypto(const Crypto &that)
: m_d(new CryptoPrivate(*that.m_d))
{
}
Crypto::Crypto(Crypto &&that)
: m_d(std::move(that.m_d))
{
}
Crypto &Crypto::operator=(const Crypto &that)
{
m_d.reset(new CryptoPrivate(*that.m_d));
return *this;
}
Crypto &Crypto::operator=(Crypto &&that)
{
m_d = std::move(that.m_d);
return *this;
}
bool Crypto::operator==(const Crypto &that) const
{
return this->m_d == that.m_d;
}
bool Crypto::valid() const
{
return m_d->valid;
}
std::string CryptoPrivate::ed25519IdentityKey() const
{
auto key = account.value()->ed25519_key()->to_base64();
return static_cast<std::string>(key);
}
std::string CryptoPrivate::curve25519IdentityKey() const
{
auto key = account.value()->curve25519_key()->to_base64();
return static_cast<std::string>(key);
}
std::string Crypto::ed25519IdentityKey() const
{
return m_d->ed25519IdentityKey();
}
std::string Crypto::curve25519IdentityKey() const
{
return m_d->curve25519IdentityKey();
}
std::string Crypto::sign(nlohmann::json j)
{
j.erase("signatures");
j.erase("unsigned");
auto str = j.dump();
auto signature = checkVodozemacError([&]() {
return m_d->account.value()->sign(rust::Slice<const std::uint8_t>(reinterpret_cast<const std::uint8_t *>(str.data()), str.size()));
});
if (!signature.has_value()) {
return "";
}
return static_cast<std::string>(signature.value()->to_base64());
}
void Crypto::setUploadedOneTimeKeysCount(immer::map<std::string /* algorithm */, int> uploadedOneTimeKeysCount)
{
m_d->uploadedOneTimeKeysCount = uploadedOneTimeKeysCount;
}
std::size_t Crypto::maxNumberOfOneTimeKeys() const
{
return m_d->account.value()->max_number_of_one_time_keys();
}
std::size_t Crypto::genOneTimeKeysRandomSize([[maybe_unused]] int num)
{
return 0;
}
void Crypto::genOneTimeKeysWithRandom([[maybe_unused]] RandomData random, int num)
{
assert(random.size() >= genOneTimeKeysRandomSize(num));
m_d->account.value()->generate_one_time_keys(num);
m_d->numUnpublishedKeys += num;
}
nlohmann::json Crypto::unpublishedOneTimeKeys() const
{
auto keys = m_d->account.value()->one_time_keys();
auto ret = nlohmann::json{
{curve25519, nlohmann::json::object()},
};
for (const auto &k : keys) {
auto keyId = static_cast<std::string>(k.key_id);
auto key = static_cast<std::string>(k.key->to_base64());
ret[curve25519][keyId] = key;
}
return ret;
}
void Crypto::markOneTimeKeysAsPublished()
{
m_d->account.value()->mark_keys_as_published();
m_d->numUnpublishedKeys = 0;
}
int Crypto::numUnpublishedOneTimeKeys() const
{
return m_d->numUnpublishedKeys;
}
int Crypto::uploadedOneTimeKeysCount(std::string algorithm) const
{
return m_d->uploadedOneTimeKeysCount[algorithm];
}
MaybeString Crypto::decrypt(nlohmann::json eventJson)
{
try {
auto content = eventJson.at("content");
auto algo = content.contains("algorithm") ? content.at("algorithm").template get<std::string>() : std::string();
if (algo == olmAlgo) {
return m_d->decryptOlm(std::move(content));
} else if (algo == megOlmAlgo) {
return m_d->decryptMegOlm(eventJson);
}
return NotBut("Algorithm " + algo + " not supported");
} catch (const std::exception &e) {
return NotBut("Malformed event");
}
}
bool Crypto::createInboundGroupSession(KeyOfGroupSession k, std::string sessionKey, std::string ed25519Key)
{
return m_d->createInboundGroupSession(std::move(k), std::move(sessionKey), std::move(ed25519Key));
}
std::size_t Crypto::importInboundGroupSessions(const nlohmann::json &keys)
{
if (!keys.is_array()) {
return 0;
}
auto validateStr = identValidate(&nlohmann::json::is_string);
std::size_t count = 0;
for (const auto &data : keys) {
if (!data.is_object()) {
continue;
}
auto key = nlohmann::json::object();
if (!(cast(key, data, "algorithm", identValidate([](const auto &j) {
return j == megOlmAlgo;
})) && cast(key, data, "room_id", validateStr)
&& cast(key, data, "session_key", validateStr)
&& cast(key, data, "session_id", validateStr)
&& cast(key, data, "/sender_claimed_keys/ed25519"_json_pointer, validateStr)
)) {
continue;
}
auto keyOfGroupSession = KeyOfGroupSession{
key["room_id"].template get<std::string>(),
key["session_id"].template get<std::string>(),
};
if (createInboundGroupSession(
keyOfGroupSession,
key["session_key"].template get<std::string>(),
key["sender_claimed_keys"]["ed25519"].template get<std::string>()
)) {
++count;
}
}
return count;
}
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)
{
auto session = InboundGroupSession(sessionKey, ed25519Key);
if (!session.valid()) {
kzo.crypto.warn() << "Invalid session key for: " << k.roomId << ", " << k.sessionId << std::endl;
return false;
}
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)
{
if (! object.contains("signatures")) {
return false;
}
std::string signature;
try {
signature = object.at("signatures").at(userId).at(ed25519 + ":" + deviceId);
} catch(const std::exception &) {
return false;
}
object.erase("signatures");
object.erase("unsigned");
auto message = object.dump();
auto res = checkVodozemacError([&]() {
auto key = vodozemac::types::ed25519_key_from_base64(ed25519Key);
auto sig = vodozemac::types::ed25519_signature_from_base64(signature);
key->verify(rust::Slice<const std::uint8_t>(reinterpret_cast<const std::uint8_t *>(message.data()), message.size()), *sig);
// It throws if the signature cannot be verified
return true;
});
return res.has_value() && res.value();
}
MaybeString Crypto::getInboundGroupSessionEd25519KeyFromEvent(const nlohmann::json &eventJson) const
{
auto content = eventJson.at("content");
auto sessionId = content.at("session_id").get<std::string>();
auto roomId = eventJson.at("room_id").get<std::string>();
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();
}
}
std::size_t Crypto::encryptOlmRandomSize(std::string /* theirCurve25519IdentityKey */) const
{
// HACK: To prevent a possible race condition where we call
// encryptedOlmRandomSize() -> randomGenerator.generateRange() ~> [encryptOlm()]
//
// Here, encryptOlm() must be called in the reducer,
// as it changes the status of
// Crypto (so we should not risk any infomation lose in Crypto).
// Also, for the reducer to be pure, we may not call generateRange() in the reducer,
// as *that* is not a pure function. This means it can only be called in an effect,
// or .then() continuation. This means, the sequence from encryptedOlmRandomSize()
// to encryptOlm() can never be atomic. That is, there may be other encryptOlm()
// calls within, and that may increase the random data needed, and as a result,
// we will not have enough random data.
//
// According to the olm headers:
// https://gitlab.matrix.org/matrix-org/olm/-/blob/master/include/olm/ratchet.hh
// The maximum random size needed to encrypt is 32. We use this to ensure we
// will always have enough random data fot the encryption.
return encryptOlmMaxRandomSize();
}
std::size_t Crypto::encryptOlmMaxRandomSize()
{
return 0;
}
nlohmann::json Crypto::encryptOlmWithRandom(
RandomData random, nlohmann::json eventJson, std::string theirCurve25519IdentityKey)
{
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, {
{"type", type},
{"body", body}
}
}
};
} catch (const std::exception &) {
return nlohmann::json::object();
}
}
nlohmann::json Crypto::encryptMegOlm(nlohmann::json eventJson)
{
auto roomId = eventJson.at("room_id").get<std::string>();
auto content = eventJson.at("content");
auto type = eventJson.at("type").get<std::string>();
auto jsonToEncrypt = nlohmann::json::object();
jsonToEncrypt["room_id"] = roomId;
jsonToEncrypt["content"] = std::move(content);
jsonToEncrypt["type"] = type;
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{
{"algorithm", CryptoConstants::megOlmAlgo},
// NOTE: we might stop sending sender_key in the future
// as per the Matrix spec
{"sender_key", curve25519IdentityKey()},
{"ciphertext", ciphertext},
- {"session_id", session.sessionId()},
+ {"session_id", sessionId},
};
}
std::size_t Crypto::rotateMegOlmSessionRandomSize()
{
return OutboundGroupSession::constructRandomSize();
}
std::string Crypto::rotateMegOlmSessionWithRandom(RandomData random, Timestamp timeMs, std::string roomId)
{
m_d->reuseOrCreateOutboundGroupSession(
random, timeMs,
roomId, std::nullopt);
return outboundGroupSessionCurrentKey(roomId);
}
std::optional<std::string> Crypto::rotateMegOlmSessionWithRandomIfNeeded(
RandomData random, Timestamp timeMs,
std::string roomId, MegOlmSessionRotateDesc desc)
{
auto oldSessionValid = m_d->reuseOrCreateOutboundGroupSession(random, timeMs, roomId, std::move(desc));
return oldSessionValid ? std::nullopt : std::optional(outboundGroupSessionCurrentKey(roomId));
}
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(
immer::map<std::string, immer::map<std::string /* deviceId */,
std::string /* curve25519IdentityKey */>> keyMap) const -> UserIdToDeviceIdMap
{
auto ret = UserIdToDeviceIdMap{};
for (auto [userId, devices] : keyMap) {
auto unknownDevices =
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;
return deviceId;
}),
devices);
if (! unknownDevices.empty()) {
ret = std::move(ret).set(userId, std::move(unknownDevices));
}
}
return ret;
}
std::size_t Crypto::createOutboundSessionRandomSize()
{
return Session::constructOutboundRandomSize();
}
void Crypto::createOutboundSessionWithRandom(
RandomData random,
std::string theirIdentityKey,
std::string theirOneTimeKey)
{
assert(random.size() >= createOutboundSessionRandomSize());
auto session = Session(OutboundSessionTag{},
RandomTag{},
random,
*m_d,
theirIdentityKey,
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));
}
}
nlohmann::json Crypto::toJson() const
{
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},
{"knownSessions", nlohmann::json(m_d->knownSessions)},
{"inboundGroupSessions", nlohmann::json(m_d->inboundGroupSessions)},
{"outboundGroupSessions", nlohmann::json(m_d->outboundGroupSessions)},
});
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);
}
}
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
index 98e0637..9e515c0 100644
--- a/src/crypto/outbound-group-session-p.hpp
+++ b/src/crypto/outbound-group-session-p.hpp
@@ -1,43 +1,43 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include <libkazv-config.hpp>
#include "outbound-group-session.hpp"
#include <vodozemac.h>
#include <immer/map.hpp>
namespace Kazv
{
struct OutboundGroupSessionPrivate
{
/// to be deprecated
OutboundGroupSessionPrivate();
OutboundGroupSessionPrivate(RandomTag,
RandomData random,
Timestamp creationTime);
OutboundGroupSessionPrivate(const OutboundGroupSessionPrivate &that);
~OutboundGroupSessionPrivate() = default;
std::optional<rust::Box<vodozemac::megolm::GroupSession>> session;
bool valid{false};
Timestamp creationTime;
std::string initialSessionKey;
std::string pickle() const;
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.cpp b/src/crypto/outbound-group-session.cpp
index 8d6b046..6a6ba08 100644
--- a/src/crypto/outbound-group-session.cpp
+++ b/src/crypto/outbound-group-session.cpp
@@ -1,188 +1,188 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021-2024 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include "outbound-group-session-p.hpp"
#include "crypto-util-p.hpp"
#include <debug.hpp>
#include "time-util.hpp"
namespace Kazv
{
OutboundGroupSessionPrivate::OutboundGroupSessionPrivate()
: session(std::nullopt)
{
}
OutboundGroupSessionPrivate::OutboundGroupSessionPrivate(
RandomTag,
[[maybe_unused]] RandomData random,
Timestamp creationTime)
: session(std::nullopt)
, creationTime(creationTime)
{
session = checkVodozemacError([&]() {
return vodozemac::megolm::new_group_session();
});
if (session.has_value()) {
valid = true;
initialSessionKey = sessionKey();
}
}
OutboundGroupSessionPrivate::OutboundGroupSessionPrivate(const OutboundGroupSessionPrivate &that)
: session(std::nullopt)
, creationTime(that.creationTime)
, initialSessionKey(that.initialSessionKey)
{
if (that.valid) {
valid = unpickle(that.pickle());
}
}
std::string OutboundGroupSessionPrivate::pickle() const
{
auto pickleData = session.value()->pickle(VODOZEMAC_PICKLE_KEY);
return static_cast<std::string>(pickleData);
}
bool OutboundGroupSessionPrivate::unpickle(std::string pickleData)
{
session = checkVodozemacError([&]() {
return vodozemac::megolm::group_session_from_pickle(pickleData, VODOZEMAC_PICKLE_KEY);
});
return session.has_value();
}
bool OutboundGroupSessionPrivate::unpickleFromLibolm(std::string pickleData)
{
session = checkVodozemacError([&]() {
return vodozemac::megolm::group_session_from_libolm_pickle(pickleData, rust::Slice<const unsigned char>(OLM_PICKLE_KEY.data(), OLM_PICKLE_KEY.size()));
});
return session.has_value();
}
std::size_t OutboundGroupSession::constructRandomSize()
{
return 0;
}
OutboundGroupSession::OutboundGroupSession()
: m_d(new OutboundGroupSessionPrivate)
{
}
OutboundGroupSession::OutboundGroupSession(RandomTag, RandomData random, Timestamp creationTime)
: m_d(new OutboundGroupSessionPrivate(RandomTag{}, std::move(random), std::move(creationTime)))
{
}
OutboundGroupSession::~OutboundGroupSession() = default;
OutboundGroupSession::OutboundGroupSession(const OutboundGroupSession &that)
: m_d(new OutboundGroupSessionPrivate(*that.m_d))
{
}
OutboundGroupSession::OutboundGroupSession(OutboundGroupSession &&that)
: m_d(std::move(that.m_d))
{
}
OutboundGroupSession &OutboundGroupSession::operator=(const OutboundGroupSession &that)
{
m_d.reset(new OutboundGroupSessionPrivate(*that.m_d));
return *this;
}
OutboundGroupSession &OutboundGroupSession::operator=(OutboundGroupSession &&that)
{
m_d = std::move(that.m_d);
return *this;
}
bool OutboundGroupSession::valid() const
{
return m_d && m_d->valid;
}
std::string OutboundGroupSession::encrypt(std::string plainText)
{
auto res = checkVodozemacError([&]() {
return m_d->session.value()->encrypt(rust::Str(plainText));
});
if (!res.has_value()) {
return std::string();
}
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();
}
std::string OutboundGroupSession::initialSessionKey() const
{
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();
}
Timestamp OutboundGroupSession::creationTimeMs() const
{
return m_d->creationTime;
}
void to_json(nlohmann::json &j, const OutboundGroupSession &s)
{
j = nlohmann::json::object();
j["version"] = 1;
j["valid"] = s.m_d->valid;
j["creationTime"] = s.m_d->creationTime;
j["initialSessionKey"] = s.m_d->initialSessionKey;
if (s.m_d->valid) {
j["session"] = s.m_d->pickle();
}
}
void from_json(const nlohmann::json &j, OutboundGroupSession &s)
{
s.m_d->valid = j.at("valid");
s.m_d->creationTime = j.at("creationTime");
s.m_d->initialSessionKey = j.at("initialSessionKey");
if (s.m_d->valid) {
if (j.contains("version") && j["version"] == 1) {
s.m_d->valid = s.m_d->unpickle(j.at("session"));
} else {
s.m_d->valid = s.m_d->unpickleFromLibolm(j.at("session"));
}
}
}
}
diff --git a/src/crypto/outbound-group-session.hpp b/src/crypto/outbound-group-session.hpp
index 6d14019..c023812 100644
--- a/src/crypto/outbound-group-session.hpp
+++ b/src/crypto/outbound-group-session.hpp
@@ -1,62 +1,62 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include <libkazv-config.hpp>
#include <memory>
#include <maybe.hpp>
#include <event.hpp>
#include "crypto-util.hpp"
namespace Kazv
{
struct OutboundGroupSessionPrivate;
class OutboundGroupSession
{
public:
/**
* @return The size of random data needed to construct an OutboundGroupSession.
*/
static std::size_t constructRandomSize();
explicit OutboundGroupSession();
/**
* Constructs an OutboundGroupSession from custom random data.
*
* @param random The random data to use. Must be of at least size
* `constructRandomSize()`.
* @param creationTime The creation time of this OutboundGroupSession.
*/
OutboundGroupSession(RandomTag, RandomData random, Timestamp creationTime);
OutboundGroupSession(const OutboundGroupSession &that);
OutboundGroupSession(OutboundGroupSession &&that);
OutboundGroupSession &operator=(const OutboundGroupSession &that);
OutboundGroupSession &operator=(OutboundGroupSession &&that);
~OutboundGroupSession();
std::string encrypt(std::string plainText);
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:
friend void to_json(nlohmann::json &j, const OutboundGroupSession &s);
friend void from_json(const nlohmann::json &j, OutboundGroupSession &s);
std::unique_ptr<OutboundGroupSessionPrivate> m_d;
};
}
diff --git a/src/crypto/session.cpp b/src/crypto/session.cpp
index 1736b8b..b48d582 100644
--- a/src/crypto/session.cpp
+++ b/src/crypto/session.cpp
@@ -1,280 +1,281 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <debug.hpp>
#include "crypto-p.hpp"
#include "session-p.hpp"
#include "base64.hpp"
namespace Kazv
{
static rust::Vec<std::uint8_t> decodeBase64ToVec(std::string encoded)
{
auto decodedMsg = decodeBase64(encoded);
auto decodedMsgVec = rust::Vec<std::uint8_t>();
decodedMsgVec.reserve(decodedMsg.size());
std::transform(
decodedMsg.begin(), decodedMsg.end(),
std::back_inserter(decodedMsgVec),
[](char v) {
return static_cast<std::uint8_t>(v);
}
);
return decodedMsgVec;
}
static const vodozemac::olm::SessionConfig &sessionConfigV1()
{
static rust::Box<vodozemac::olm::SessionConfig> v1 = vodozemac::olm::new_session_config_version_1();
return *v1;
}
SessionPrivate::SessionPrivate()
: session(std::nullopt)
{
}
SessionPrivate::SessionPrivate(OutboundSessionTag,
RandomTag,
[[maybe_unused]] RandomData random,
CryptoPrivate &cryptoD,
std::string theirIdentityKey,
std::string theirOneTimeKey)
: SessionPrivate()
{
assert(random.size() >= Session::constructOutboundRandomSize());
session = checkVodozemacError([&]() {
auto identityKey = vodozemac::types::curve_key_from_base64(rust::Str(theirIdentityKey));
auto oneTimeKey = vodozemac::types::curve_key_from_base64(rust::Str(theirOneTimeKey));
return cryptoD.account.value()->create_outbound_session(
sessionConfigV1(),
*identityKey,
*oneTimeKey
);
});
valid = session.has_value();
}
SessionPrivate::SessionPrivate(InboundSessionTag,
CryptoPrivate &cryptoD,
std::string theirIdentityKey,
std::string message)
: SessionPrivate()
{
auto res = checkVodozemacError([&]() {
auto identityKey = vodozemac::types::curve_key_from_base64(rust::Str(theirIdentityKey));
auto msg = vodozemac::olm::olm_message_from_parts(vodozemac::olm::OlmMessageParts{
0, // pre-key message
decodeBase64ToVec(message),
});
return cryptoD.account.value()->create_inbound_session(
sessionConfigV1(),
*identityKey,
*msg
);
});
valid = res.has_value();
if (valid) {
session = std::move(res->session);
firstDecrypted = std::string(res->plaintext.begin(), res->plaintext.end());
}
}
SessionPrivate::SessionPrivate(const SessionPrivate &that)
: SessionPrivate()
{
if (that.valid) {
valid = unpickle(that.pickle());
+ firstDecrypted = that.firstDecrypted;
}
}
std::string SessionPrivate::pickle() const
{
auto pickleData = session.value()->pickle(VODOZEMAC_PICKLE_KEY);
return static_cast<std::string>(pickleData);
}
bool SessionPrivate::unpickle(std::string pickleData)
{
session = checkVodozemacError([&]() {
return vodozemac::olm::session_from_pickle(pickleData, VODOZEMAC_PICKLE_KEY);
});
return session.has_value();
}
bool SessionPrivate::unpickleFromLibolm(std::string pickleData)
{
session = checkVodozemacError([&]() {
return vodozemac::olm::session_from_libolm_pickle(
pickleData,
rust::Slice<const unsigned char>(OLM_PICKLE_KEY.data(), OLM_PICKLE_KEY.size()));
});
return session.has_value();
}
MaybeString SessionPrivate::takeFirstDecrypted()
{
auto res = std::move(firstDecrypted);
firstDecrypted = std::nullopt;
if (res.has_value()) {
return res.value();
} else {
return NotBut("No first decrypted available");
}
}
std::size_t Session::constructOutboundRandomSize()
{
return 0;
}
Session::Session()
: m_d(new SessionPrivate)
{
}
Session::Session(OutboundSessionTag,
RandomTag,
RandomData data,
CryptoPrivate &cryptoD,
std::string theirIdentityKey,
std::string theirOneTimeKey)
: m_d(new SessionPrivate{
OutboundSessionTag{},
RandomTag{},
std::move(data),
cryptoD,
theirIdentityKey,
theirOneTimeKey})
{
}
Session::Session(InboundSessionTag,
CryptoPrivate &cryptoD,
std::string theirIdentityKey,
std::string theirOneTimeKey)
: m_d(new SessionPrivate{
InboundSessionTag{},
cryptoD,
theirIdentityKey,
theirOneTimeKey})
{
}
Session::~Session() = default;
Session::Session(const Session &that)
: m_d(new SessionPrivate(*that.m_d))
{
}
Session::Session(Session &&that)
: m_d(std::move(that.m_d))
{
}
Session &Session::operator=(const Session &that)
{
m_d.reset(new SessionPrivate(*that.m_d));
return *this;
}
Session &Session::operator=(Session &&that)
{
m_d = std::move(that.m_d);
return *this;
}
- 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{
0, // pre-key
decodeBase64ToVec(message),
});
return m_d->session.value()->session_matches(*msg);
});
return res.has_value() && res.value();
}
bool Session::valid() const
{
// maybe a moved-from state, so check m_d first
return m_d && m_d->valid;
}
MaybeString Session::decrypt(int type, std::string message)
{
auto res = checkVodozemacError([&]() {
auto msg = vodozemac::olm::olm_message_from_parts(vodozemac::olm::OlmMessageParts{
static_cast<std::size_t>(type),
decodeBase64ToVec(message),
});
return m_d->session.value()->decrypt(*msg);
});
if (!res.has_value()) {
return NotBut(res.reason());
}
return std::string(res.value().begin(), res.value().end());
}
std::size_t Session::encryptRandomSize() const
{
return 0;
}
std::pair<int, std::string> Session::encryptWithRandom([[maybe_unused]] RandomData random, std::string plainText)
{
assert(random.size() >= encryptRandomSize());
auto res = checkVodozemacError([&]() {
return m_d->session.value()->encrypt(plainText);
});
if (!res.has_value()) {
return { -1, "" };
}
auto [type, msg] = res.value()->to_parts();
auto base64Msg = encodeBase64(std::string(msg.begin(), msg.end()));
return { type, std::move(base64Msg) };
}
void to_json(nlohmann::json &j, const Session &s)
{
j = nlohmann::json::object({
{"valid", s.m_d->valid},
{"version", 1},
{"data", s.m_d->valid ? s.m_d->pickle() : std::string()}
});
}
void from_json(const nlohmann::json &j, Session &s)
{
if (j.at("valid").template get<bool>()) {
if (j.contains("version") && j["version"] == 1) {
s.m_d->valid = s.m_d->unpickle(j.at("data"));
} else {
s.m_d->valid = s.m_d->unpickleFromLibolm(j.at("data"));
}
}
}
}
diff --git a/src/crypto/session.hpp b/src/crypto/session.hpp
index 73ab937..5cf1c11 100644
--- a/src/crypto/session.hpp
+++ b/src/crypto/session.hpp
@@ -1,94 +1,94 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include <libkazv-config.hpp>
#include <memory>
#include <tuple>
#include <maybe.hpp>
#include "crypto-util.hpp"
namespace Kazv
{
class Crypto;
struct InboundSessionTag {};
struct OutboundSessionTag {};
struct SessionPrivate;
struct CryptoPrivate;
class Session
{
/**
* @return The size of random data needed to construct an outbound session.
*/
static std::size_t constructOutboundRandomSize();
/**
* Construct an outbound session with custom random data.
*
* @param data The custom random data. Must be of at least
* size `constructOutboundRandomSize()`.
*/
explicit Session(OutboundSessionTag,
RandomTag,
RandomData data,
CryptoPrivate &cryptoD,
std::string theirIdentityKey,
std::string theirOneTimeKey);
// Creates an inbound session
explicit Session(InboundSessionTag,
CryptoPrivate &cryptoD,
std::string theirIdentityKey,
std::string message);
public:
explicit Session();
Session(const Session &that);
Session(Session &&that);
Session &operator=(const Session &that);
Session &operator=(Session &&that);
~Session();
- bool matches(std::string message);
+ bool matches(std::string message) const;
bool valid() const;
MaybeString decrypt(int type, std::string message);
/**
* @return The size of random data needed for the next encryption.
*/
std::size_t encryptRandomSize() const;
/**
* Encrypt plainText.
*
* @param random The random data needed for the encryption. Must be
* of at least size `encryptRandomSize()`.
* @param plainText The plain text to encrypt.
*
* @return A pair containing the type and encrypted message string if
* the encryption is successful. Otherwise, return {-1, ""}.
*/
std::pair<int /* type */, std::string /* message */> encryptWithRandom(
RandomData random, std::string plainText);
private:
friend class Crypto;
friend struct CryptoPrivate;
friend struct SessionPrivate;
friend void to_json(nlohmann::json &j, const Session &s);
friend void from_json(const nlohmann::json &j, Session &s);
std::unique_ptr<SessionPrivate> m_d;
};
}
diff --git a/src/tests/base/immer-utils-test.cpp b/src/tests/base/immer-utils-test.cpp
index 0c23dbf..a5b0df0 100644
--- a/src/tests/base/immer-utils-test.cpp
+++ b/src/tests/base/immer-utils-test.cpp
@@ -1,79 +1,159 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <catch2/catch_test_macros.hpp>
#include <catch2/matchers/catch_matchers_predicate.hpp>
#include <catch2/matchers/catch_matchers_contains.hpp>
#include <string>
#include <immer/map.hpp>
#include <immer/flex_vector.hpp>
+#include <immer/box.hpp>
#include <immer-utils.hpp>
using namespace Kazv;
using ComplexMapT = immer::map<std::string, immer::map<std::string, immer::flex_vector<std::string>>>;
TEST_CASE("getIn()", "[base][immer-utils]")
{
ComplexMapT complexMap{
{"a", {{"a/b", {"a/b/0", "a/b/1"}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
};
REQUIRE(getIn(complexMap) == complexMap);
REQUIRE(getIn(complexMap, "a") == complexMap["a"]);
REQUIRE(getIn(complexMap, "a", "a/b") == complexMap["a"]["a/b"]);
REQUIRE(getIn(complexMap, "a", "a/b", 0) == complexMap["a"]["a/b"][0]);
}
TEST_CASE("setIn()", "[base][immer-utils]")
{
ComplexMapT complexMap{
{"a", {{"a/b", {"a/b/0", "a/b/1"}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
};
REQUIRE(setIn(complexMap, {}) == ComplexMapT{});
REQUIRE(setIn(complexMap, {}, "a") == ComplexMapT{
{"a", {}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
});
REQUIRE(setIn(complexMap, {}, "a", "a/b") == ComplexMapT{
{"a", {{"a/b", {}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
});
REQUIRE(setIn(complexMap, "test", "a", "a/b", 0) == ComplexMapT{
{"a", {{"a/b", {"test", "a/b/1"}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
});
}
TEST_CASE("updateIn()", "[base][immer-utils]")
{
ComplexMapT complexMap{
{"a", {{"a/b", {"a/b/0", "a/b/1"}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
};
REQUIRE(updateIn(complexMap, [](auto) { return ComplexMapT{}; }) == ComplexMapT{});
REQUIRE(updateIn(complexMap, [](auto a) { return a.set("a/b2", {"a/b2/0", "a/b2/1"}); }, "a") == ComplexMapT{
{"a", {{"a/b", {"a/b/0", "a/b/1"}}, {"a/b2", {"a/b2/0", "a/b2/1"}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
});
REQUIRE(updateIn(complexMap, [](auto a) { return a.set(0, "new"); }, "a", "a/b") == ComplexMapT{
{"a", {{"a/b", {"new", "a/b/1"}}}},
{"c", {{"c/d", {"c/d/0", "c/d/1"}}}},
});
REQUIRE(updateIn(complexMap, [](auto a) { return a + "test"; }, "a", "a/b", 0) == ComplexMapT{
{"a", {{"a/b", {"a/b/0test", "a/b/1"}}}},
{"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/x-diff
Expires
Fri, Oct 9, 9:29 PM (1 d, 19 h)
Storage Engine
blob
Storage Format
Raw Data
Storage Handle
1784884
Default Alt Text
(67 KB)
Attached To
Mode
rL libkazv
Attached
Detach File
Event Timeline
Log In to Comment