Page MenuHomePhorge

No OneTemporary

Size
67 KB
Referenced Files
None
Subscribers
None
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 &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)
{
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

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)

Event Timeline