Page MenuHomePhorge

No OneTemporary

Size
153 KB
Referenced Files
None
Subscribers
None
diff --git a/CMakeLists.txt b/CMakeLists.txt
index 3558ab0..0a499e7 100644
--- a/CMakeLists.txt
+++ b/CMakeLists.txt
@@ -1,125 +1,125 @@
# Do not let the option()s override variables here
cmake_minimum_required(VERSION 3.13)
if(NOT DEFINED PROJECT_NAME)
if(NOT DEFINED libkazv_INSTALL_HEADERS)
set(libkazv_INSTALL_HEADERS ON)
endif()
endif()
project(libkazv)
set(libkazvSourceRoot ${CMAKE_CURRENT_SOURCE_DIR})
set(libkazv_VERSION_MAJOR 0)
set(libkazv_VERSION_MINOR 8)
set(libkazv_VERSION_PATCH 0)
set(libkazv_VERSION_STRING ${libkazv_VERSION_MAJOR}.${libkazv_VERSION_MINOR}.${libkazv_VERSION_PATCH})
set(libkazv_SOVERSION 7)
set(CMAKE_MODULE_PATH "${CMAKE_CURRENT_SOURCE_DIR}/cmake" ${CMAKE_MODULE_PATH})
-set(CMAKE_CXX_STANDARD 17)
+set(CMAKE_CXX_STANDARD 20)
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wall -Wextra")
option(libkazv_BUILD_TESTS "Build tests" ON)
option(libkazv_BUILD_EXAMPLES "Build examples" ON)
option(libkazv_BUILD_KAZVJOB "Build libkazvjob the async and networking library" ON)
option(libkazv_OUTPUT_LEVEL "Output level: Debug>=90, Info>=70, Quiet>=20, no output=1" 0)
option(libkazv_ENABLE_COVERAGE "Enable code coverage information" OFF)
if(libkazv_ENABLE_COVERAGE)
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -fprofile-arcs -ftest-coverage -fPIC -O0")
endif()
if((libkazv_BUILD_TESTS OR libkazv_BUILD_EXAMPLES) AND NOT libkazv_BUILD_KAZVJOB)
message(FATAL_ERROR
"You asked kazvjob not to be built, but asked to build tests or examples. Tests and examples both depend on kazvjob. This is not possible.")
endif()
if(libkazv_BUILD_TESTS)
set(BUILD_TESTING ON)
include(CTest)
set(CMAKE_CTEST_ARGUMENTS --output-on-failure --output-junit test-report.xml)
endif()
# Build shared libraries by default
if(NOT DEFINED BUILD_SHARED_LIBS)
set(BUILD_SHARED_LIBS ON)
endif()
if ("${CMAKE_BUILD_TYPE}" STREQUAL "Debug")
set(LIBKAZV_BUILT_WITH_DEBUG 1)
else()
set(LIBKAZV_BUILT_WITH_DEBUG 0)
endif()
find_package(Boost REQUIRED COMPONENTS serialization regex)
if(libkazv_BUILD_KAZVJOB)
find_package(cpr REQUIRED)
set(THREADS_PREFER_PTHREAD_FLAG ON)
find_package(Threads REQUIRED)
endif()
if(libkazv_BUILD_TESTS)
find_package(Catch2 REQUIRED)
endif()
find_package(nlohmann_json REQUIRED)
find_package(Immer REQUIRED)
find_package(Zug REQUIRED)
find_package(Lager REQUIRED)
if(MINGW)
# On MinGW, disregard the Findcryptopp.cmake
find_package(cryptopp REQUIRED CONFIG)
else()
find_package(cryptopp REQUIRED)
endif()
# On MinGW we use the CMake build, and the target name is different
# for the shared and the static libraries.
if(MINGW)
set(CRYPTOPP_TARGET_NAME cryptopp-shared)
else()
set(CRYPTOPP_TARGET_NAME cryptopp)
endif()
find_package(vodozemac REQUIRED)
if(libkazv_BUILD_EXAMPLES)
find_package(PkgConfig REQUIRED)
pkg_check_modules(LIBHTTPSERVER REQUIRED libhttpserver)
endif()
if(libkazv_OUTPUT_LEVEL)
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -DLIBKAZV_OUTPUT_LEVEL=${libkazv_OUTPUT_LEVEL}")
endif()
include(GNUInstallDirs)
set(ConfigPackageLocation ${CMAKE_INSTALL_LIBDIR}/cmake/libkazv)
add_subdirectory(src)
install(EXPORT libkazvTargets
NAMESPACE
libkazv::
DESTINATION
${ConfigPackageLocation}
)
install(
FILES cmake/libkazvConfig.cmake
DESTINATION ${ConfigPackageLocation})
# cpr does not install a good config file.
# We use the find module instead.
install(
FILES cmake/Findcpr.cmake
DESTINATION ${ConfigPackageLocation})
diff --git a/src/base/basejob.cpp b/src/base/basejob.cpp
index fcd7c8c..f7a9517 100644
--- a/src/base/basejob.cpp
+++ b/src/base/basejob.cpp
@@ -1,296 +1,294 @@
/*
* 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 <lager/util.hpp>
#include <vector>
#include <tuple>
#include "basejob.hpp"
namespace Kazv
{
BaseJob::Get BaseJob::GET{};
BaseJob::Post BaseJob::POST{};
BaseJob::Put BaseJob::PUT{};
BaseJob::Delete BaseJob::DELETE{};
struct BaseJob::Private
{
Private(std::string serverUrl,
std::string requestUrl,
Method method,
std::string token,
ReturnType returnType,
Body body,
Query query,
Header header,
std::string jobId,
std::optional<FileDesc> responseFile);
std::string fullRequestUrl;
Method method;
ReturnType returnType;
Body body;
Query query;
Header header;
JsonWrap data;
std::string jobId;
std::optional<std::string> queueId;
JobQueuePolicy queuePolicy;
std::optional<FileDesc> responseFile;
+
+ friend bool operator==(const Private &a, const Private &b) = default;
};
BaseJob::Private::Private(std::string serverUrl,
std::string requestUrl,
Method method,
std::string token,
ReturnType returnType,
Body body,
Query query,
Header header,
std::string jobId,
std::optional<FileDesc> responseFile)
: fullRequestUrl(serverUrl + requestUrl)
, method(std::move(method))
, returnType(returnType)
, body()
, query(std::move(query))
, jobId(std::move(jobId))
, responseFile(std::move(responseFile))
{
auto header_ = header.get();
if (token.size()) {
header_["Authorization"] = "Bearer " + token;
}
// convert to BytesBody, if possible
if (isBodyJson(body)) {
JsonBody j = std::get<JsonBody>(std::move(body));
header_["Content-Type"] = "application/json";
this->body = j.get().dump();
} else if (std::holds_alternative<EmptyBody>(body)) {
this->body = BytesBody();
} else {
this->body = std::move(body);
}
this->header = header_;
}
BaseJob::BaseJob(std::string serverUrl,
std::string requestUrl,
Method method,
std::string jobId,
std::string token,
ReturnType returnType,
Body body,
Query query,
Header header,
std::optional<FileDesc> responseFile)
: m_d(std::make_unique<Private>(serverUrl, requestUrl, method, token,
returnType, body, query, header, jobId, responseFile))
{
}
KAZV_DEFINE_COPYABLE_UNIQUE_PTR(BaseJob, m_d)
BaseJob::~BaseJob() = default;
bool BaseJob::shouldReturnJson() const
{
return m_d->returnType == ReturnType::Json;
};
std::string BaseJob::url() const
{
return m_d->fullRequestUrl;
};
auto BaseJob::requestBody() const -> Body
{
return m_d->body;
}
auto BaseJob::requestHeader() const -> Header
{
return m_d->header;
}
auto BaseJob::returnType() const -> ReturnType
{
return m_d->returnType;
}
auto BaseJob::requestQuery() const -> Query
{
return m_d->query;
}
auto BaseJob::requestMethod() const -> Method
{
return m_d->method;
}
JsonWrap Response::jsonBody() const
{
return std::get<JsonWrap>(body);
}
json Response::dataJson(const std::string &key) const
{
return extraData.get()[key];
}
std::string Response::dataStr(const std::string &key) const
{
return dataJson(key);
}
std::string Response::jobId() const
{
return dataStr("-job-id");
}
bool BaseJob::contentTypeMatches(immer::array<std::string> expected, std::string actual)
{
for (const auto &i : expected) {
if (i == "*/*"s) {
return true;
} else {
std::size_t pos = i.find("/*"s);
if (pos != std::string::npos) {
std::string majorType(i.data(), i.data() + pos + 1); // includes `/'
if (actual.find(majorType) == 0) {
return true;
}
} else if (i == actual) {
return true;
}
}
}
return false;
}
Response BaseJob::genResponse(Response r) const
{
auto j = m_d->data.get();
j["-job-id"] = m_d->jobId;
r.extraData = j;
return r;
}
void BaseJob::attachData(JsonWrap j)
{
m_d->data = j;
}
BaseJob BaseJob::withData(JsonWrap j) &&
{
auto ret = BaseJob(std::move(*this));
ret.attachData(j);
return ret;
}
BaseJob BaseJob::withData(JsonWrap j) const &
{
auto ret = BaseJob(*this);
ret.attachData(j);
return ret;
}
BaseJob BaseJob::withQueue(std::string id, JobQueuePolicy policy) &&
{
auto ret = BaseJob(std::move(*this));
ret.m_d->queueId = id;
ret.m_d->queuePolicy = policy;
return ret;
}
BaseJob BaseJob::withQueue(std::string id, JobQueuePolicy policy) const &
{
auto ret = BaseJob(*this);
ret.m_d->queueId = id;
ret.m_d->queuePolicy = policy;
return ret;
}
json BaseJob::dataJson(const std::string &key) const
{
return m_d->data.get()[key];
}
std::string BaseJob::dataStr(const std::string &key) const
{
return dataJson(key);
}
std::string BaseJob::jobId() const
{
return m_d->jobId;
}
std::optional<std::string> BaseJob::queueId() const
{
return m_d->queueId;
}
JobQueuePolicy BaseJob::queuePolicy() const
{
return m_d->queuePolicy;
}
std::optional<FileDesc> BaseJob::responseFile() const
{
return m_d->responseFile;
}
std::string Response::errorCode() const
{
// https://matrix.org/docs/spec/client_server/latest#api-standards
if (isBodyJson(body)) {
auto jb = jsonBody();
if (jb.get().contains("errcode")) {
auto code = jb.get()["errcode"].get<std::string>();
if (code != "M_UNKNOWN") {
return code;
}
}
}
return std::to_string(statusCode);
}
std::string Response::errorMessage() const
{
if (isBodyJson(body)) {
auto jb = jsonBody();
if (jb.get().contains("error")) {
auto msg = jb.get()["error"].get<std::string>();
return msg;
}
}
return "";
}
- bool operator==(BaseJob a, BaseJob b)
+ bool operator==(const BaseJob &a, const BaseJob &b)
{
- return a.m_d->fullRequestUrl == b.m_d->fullRequestUrl
- && a.m_d->method == b.m_d->method
- && a.m_d->returnType == b.m_d->returnType
- && a.m_d->body == b.m_d->body
- && a.m_d->query == b.m_d->query
- && a.m_d->header == b.m_d->header
- && a.m_d->data == b.m_d->data
- && a.m_d->jobId == b.m_d->jobId;
+ return a.m_d == b.m_d || (
+ a.m_d && b.m_d
+ && *(a.m_d) == *(b.m_d)
+ );
}
- bool operator!=(BaseJob a, BaseJob b)
+ bool operator!=(const BaseJob &a, const BaseJob &b)
{
return !(a == b);
}
}
diff --git a/src/base/basejob.hpp b/src/base/basejob.hpp
index 91281b8..61d9e09 100644
--- a/src/base/basejob.hpp
+++ b/src/base/basejob.hpp
@@ -1,272 +1,268 @@
/*
* 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 <variant>
#include <tuple>
#include <functional>
#include <string>
#include <map>
#include <future>
#include <immer/map.hpp>
#include <immer/array.hpp>
#include <immer/box.hpp>
#include "types.hpp"
#include "copy-helper.hpp"
#include "file-desc.hpp"
namespace Kazv
{
using Header = immer::box<std::map<std::string, std::string>>;
using BytesBody = Bytes;
using JsonBody = JsonWrap;
using FileBody = FileDesc;
- struct EmptyBody {};
+ struct EmptyBody {
+ friend bool operator==(EmptyBody, EmptyBody) = default;
+ };
using Body = std::variant<EmptyBody, JsonBody, BytesBody, FileBody>;
- inline bool operator==(EmptyBody, EmptyBody)
- {
- return true;
- }
inline bool isBodyJson(Body body) {
return std::holds_alternative<JsonBody>(body);
};
enum JobQueuePolicy
{
AlwaysContinue,
CancelFutureIfFailed
};
- struct Response {
+ struct Response
+ {
using StatusCode = int;
StatusCode statusCode;
Body body;
Header header;
JsonWrap extraData;
std::string errorCode() const;
std::string errorMessage() const;
JsonWrap jsonBody() const;
constexpr bool success() const {
return statusCode < 400;
}
json dataJson(const std::string &key) const;
std::string dataStr(const std::string &key) const;
std::string jobId() const;
+ friend bool operator==(const Response &a, const Response &b) = default;
};
- inline bool operator==(Response a, Response b)
- {
- return a.statusCode == b.statusCode
- && a.body == b.body
- && a.header == b.header
- && a.extraData == b.extraData;
- }
-
-
class BaseJob
{
public:
- struct Get {};
- struct Post {};
- struct Put {};
- struct Delete {};
+ struct Get
+ {
+ friend bool operator==(BaseJob::Get, BaseJob::Get) = default;
+ };
+ struct Post
+ {
+ friend bool operator==(BaseJob::Post, BaseJob::Post) = default;
+ };
+ struct Put
+ {
+ friend bool operator==(BaseJob::Put, BaseJob::Put) = default;
+ };
+ struct Delete
+ {
+ friend bool operator==(BaseJob::Delete, BaseJob::Delete) = default;
+ };
using Method = std::variant<Get, Post, Put, Delete>;
static Get GET;
static Post POST;
static Put PUT;
static Delete DELETE;
class Query : public std::vector<std::pair<std::string, std::string>>
{
using BaseT = std::vector<std::pair<std::string, std::string>>;
public:
using BaseT::BaseT;
void add(std::string k, std::string v) {
push_back({k, v});
}
};
using Body = ::Kazv::Body;
using BytesBody = ::Kazv::BytesBody;
using JsonBody = ::Kazv::JsonBody;
using EmptyBody = ::Kazv::EmptyBody;
using Header = ::Kazv::Header;
using Response = ::Kazv::Response;
enum ReturnType {
Json,
File,
};
BaseJob(std::string serverUrl,
std::string requestUrl,
Method method,
std::string jobId,
std::string token = {},
ReturnType returnType = ReturnType::Json,
Body body = EmptyBody{},
Query query = {},
Header header = {},
std::optional<FileDesc> responseFile = std::nullopt);
KAZV_DECLARE_COPYABLE(BaseJob)
~BaseJob();
bool shouldReturnJson() const;
std::string url() const;
Body requestBody() const;
Header requestHeader() const;
ReturnType returnType() const;
/// returns the non-encoded query as an array of pairs
Query requestQuery() const;
Method requestMethod() const;
static bool contentTypeMatches(immer::array<std::string> expected, std::string actual);
Response genResponse(Response r) const;
BaseJob withData(JsonWrap j) &&;
BaseJob withData(JsonWrap j) const &;
BaseJob withQueue(std::string id, JobQueuePolicy policy = AlwaysContinue) &&;
BaseJob withQueue(std::string id, JobQueuePolicy policy = AlwaysContinue) const &;
json dataJson(const std::string &key) const;
std::string dataStr(const std::string &key) const;
std::string jobId() const;
std::optional<std::string> queueId() const;
JobQueuePolicy queuePolicy() const;
std::optional<FileDesc> responseFile() const;
protected:
void attachData(JsonWrap data);
private:
- friend bool operator==(BaseJob a, BaseJob b);
+ friend bool operator==(const BaseJob &a, const BaseJob &b);
+ friend bool operator!=(const BaseJob &a, const BaseJob &b);
struct Private;
std::unique_ptr<Private> m_d;
};
- bool operator==(BaseJob a, BaseJob b);
- bool operator!=(BaseJob a, BaseJob b);
- inline bool operator==(BaseJob::Get, BaseJob::Get) { return true; }
- inline bool operator==(BaseJob::Post, BaseJob::Post) { return true; }
- inline bool operator==(BaseJob::Put, BaseJob::Put) { return true; }
- inline bool operator==(BaseJob::Delete, BaseJob::Delete) { return true; }
-
-
namespace detail
{
template<class T>
struct AddToQueryT
{
template<class U>
static void call(BaseJob::Query &q, std::string name, U &&arg) {
q.add(name, std::to_string(std::forward<U>(arg)));
}
};
template<>
struct AddToQueryT<std::string>
{
template<class U>
static void call(BaseJob::Query &q, std::string name, U &&arg) {
q.add(name, std::forward<U>(arg));
}
};
template<>
struct AddToQueryT<bool>
{
template<class U>
static void call(BaseJob::Query &q, std::string name, U &&arg) {
q.add(name, std::forward<U>(arg) ? "true"s : "false"s);
}
};
template<class T>
struct AddToQueryT<immer::array<T>>
{
template<class U>
static void call(BaseJob::Query &q, std::string name, U &&arg) {
for (auto v : std::forward<U>(arg)) {
q.add(name, v);
}
}
};
template<>
struct AddToQueryT<json>
{
// https://github.com/nlohmann/json/issues/2040
static void call(BaseJob::Query &q, std::string /* name */, const json &arg) {
// assume v is string type
for (auto [k, v] : arg.items()) {
q.add(k, v);
}
}
};
}
template<class T>
inline void addToQuery(BaseJob::Query &q, std::string name, T &&arg)
{
detail::AddToQueryT<std::decay_t<T>>::call(q, name, std::forward<T>(arg));
}
namespace detail
{
template<class T>
struct AddToQueryIfNeededT
{
template<class U>
static void call(BaseJob::Query &q, std::string name, U &&arg) {
using ArgT = std::decay_t<U>;
if constexpr (detail::hasEmptyMethod(boost::hana::type_c<ArgT>)) {
if (! arg.empty()) {
addToQuery(q, name, std::forward<U>(arg));
}
} else {
addToQuery(q, name, std::forward<U>(arg));
}
}
};
template<class T>
struct AddToQueryIfNeededT<std::optional<T>>
{
template<class U>
static void call(BaseJob::Query &q, std::string name, U &&arg) {
if (arg.has_value()) {
addToQuery(q, name, std::forward<U>(arg).value());
}
}
};
}
template<class T>
inline void addToQueryIfNeeded(BaseJob::Query &q, std::string name, T &&arg)
{
detail::AddToQueryIfNeededT<std::decay_t<T>>::call(q, name, std::forward<T>(arg));
}
}
diff --git a/src/base/event.hpp b/src/base/event.hpp
index e8ce5b9..b01f666 100644
--- a/src/base/event.hpp
+++ b/src/base/event.hpp
@@ -1,128 +1,127 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2020-2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include "libkazv-config.hpp"
#include <string>
#include <cstdint>
#include "jsonwrap.hpp"
namespace Kazv
{
using Timestamp = std::int_fast64_t;
class Event
{
public:
static const JsonWrap notYetDecryptedEvent;
enum DecryptionStatus {
NotDecrypted,
Decrypted
};
Event();
Event(JsonWrap j);
static Event fromSync(Event e, std::string roomId);
/// returns the id of this event
std::string id() const;
std::string sender() const;
Timestamp originServerTs() const;
std::string type() const;
std::string stateKey() const;
/**
* @return whether this event is a state event.
* An event is considered a state event if and only if
* it has a maybe empty stateKey.
*/
bool isState() const;
JsonWrap content() const;
/// returns the decrypted json
JsonWrap raw() const;
/// returns the original json we fetched, probably encrypted.
JsonWrap originalJson() const;
JsonWrap decryptedJson() const;
bool encrypted() const;
bool decrypted() const;
/// internal. only to be called from inside the client.
Event setDecryptedJson(JsonWrap decryptedJson, DecryptionStatus decrypted) const;
/// returns whether this event has been redacted.
bool redacted() const;
/**
* Get the event id this event is replying to.
*
* @return The event id this event is replying to, or empty string if this
* event is not a reply.
*/
std::string replyingTo() const;
/**
* Get the relationship that this event contains.
*
* @return A Pair containing the rel_type and event_id in the m.relates_to section
* of this event, or a Pair of empty strings if there are not any.
*/
std::pair<std::string/* relType */, std::string/* eventId */> relationship() const;
/**
* Get the m.relates_to object in the event.
*
* @return The m.relates_to object of this event.
*/
JsonWrap mRelatesTo() const;
template<class Archive>
void serialize(Archive &ar, std::uint32_t const /*version*/ ) {
ar & m_json & m_decryptedJson & m_decrypted & m_encrypted;
}
+ friend bool operator==(const Event &a, const Event &b);
+ inline friend bool operator!=(const Event &a, const Event &b) { return !(a == b); };
+
private:
JsonWrap m_json;
JsonWrap m_decryptedJson;
DecryptionStatus m_decrypted{NotDecrypted};
bool m_encrypted{false};
};
-
- bool operator==(const Event &a, const Event &b);
- inline bool operator!=(const Event &a, const Event &b) { return !(a == b); };
-
}
BOOST_CLASS_VERSION(Kazv::Event, 0)
namespace nlohmann
{
template <>
struct adl_serializer<Kazv::Event> {
static void to_json(json& j, Kazv::Event w) {
j = w.originalJson();
}
static void from_json(const json& j, Kazv::Event &w) {
w = Kazv::Event(Kazv::JsonWrap(j));
}
};
}
diff --git a/src/base/file-desc.cpp b/src/base/file-desc.cpp
index 70838b3..7d53159 100644
--- a/src/base/file-desc.cpp
+++ b/src/base/file-desc.cpp
@@ -1,52 +1,44 @@
/*
* 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 <fstream>
#include "file-desc.hpp"
#include "debug.hpp"
namespace Kazv
{
FileProvider FileDesc::provider(const FileInterface &fh) const
{
if (m_inMemory) {
return FileProvider(DumbFileProvider(m_content));
}
return fh.getProviderFor(*this);
}
DumbFileProvider DumbFileInterface::getProviderFor(FileDesc desc) const
{
using CharT = char;
auto nameOpt = desc.name();
kzo.base.dbg() << "DumbFileInterface::getProviderFor()" << std::endl;
if (! nameOpt) { // should not happen
kzo.base.dbg() << "No name provided, skip" << std::endl;
return DumbFileProvider{FileContent{}};
}
auto stream = std::ifstream(nameOpt.value(), std::ios_base::binary);
if (! stream) {
kzo.base.dbg() << "Cannot open file, skip" << std::endl;
return DumbFileProvider{FileContent{}};
}
auto content = FileContent(std::istreambuf_iterator<CharT>(stream),
std::istreambuf_iterator<CharT>());
return DumbFileProvider{std::move(content)};
}
-
- bool FileDesc::operator==(const FileDesc &that) const
- {
- return m_inMemory == that.m_inMemory
- && m_content == that.m_content
- && m_name == that.m_name
- && m_contentType == that.m_contentType;
- }
}
diff --git a/src/base/file-desc.hpp b/src/base/file-desc.hpp
index 24a4688..7712997 100644
--- a/src/base/file-desc.hpp
+++ b/src/base/file-desc.hpp
@@ -1,331 +1,331 @@
/*
* 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 <optional>
#include <string>
#include <variant>
#include <functional>
#include <immer/box.hpp>
#include <immer/flex_vector.hpp>
namespace Kazv
{
enum struct FileOpRetCode
{
Error,
Success,
Eof
};
enum struct FileOpenMode
{
Read,
Write
};
using FileContent = immer::flex_vector<char>;
class FileStream
{
public:
using DataT = FileContent;
/**
* Constructor.
*
* Construct a FileStream using `o`.
*
* `o` should be of a type that accepts the `read` and
* `write` methods in this class. That is,
* `o.read(int{}, std::function<void(FileOpRetCode, DataT)>{})`
* and `o.write(DataT{}, std::function<void(FileOpRetCode, int)>{})`
* must both be well-formed.
*/
template<class DeriveT>
FileStream(DeriveT &&o)
: m_d(std::make_unique<Model<std::decay_t<DeriveT>>>(
std::forward<DeriveT>(o))) {}
/**
* Read at most maxSize bytes from the stream.
*
* This calls `readCallback(FileOpRetCode::Success, data)` upon success,
* where `data` is a DataT containing the bytes read;
* calls `readCallback(FileOpRetCode::Eof, DataT{})` upon meeting EOF;
* and calls `readCallback(FileOpRetCode::Error, DataT{})` upon other errors.
*/
template<class Callback>
void read(int maxSize, Callback readCallback) { m_d->read(maxSize, readCallback); }
/**
* Write `data` into the stream.
*
* This calls `writeCallback(FileOpRetCode::Success, num)` upon success,
* where `num` is an integer that contains the number of bytes written.
* and calls `readCallback(FileOpRetCode::Error, 0)` upon other errors.
*/
template<class Callback>
void write(DataT data, Callback writeCallback) { m_d->write(data, writeCallback); }
private:
struct Concept
{
virtual ~Concept() = default;
using ReadCallback = std::function<void(FileOpRetCode, DataT)>;
using WriteCallback = std::function<void(FileOpRetCode, int)>;
virtual void read(int maxSize, ReadCallback callback) = 0;
virtual void write(DataT data, WriteCallback callback) = 0;
};
template<class DeriveT>
struct Model : public Concept
{
explicit Model(const DeriveT &o) : obj(o) {}
explicit Model(DeriveT &&o) : obj(std::move(o)) {}
~Model() override = default;
void read(int maxSize, ReadCallback callback) override {
return obj.read(maxSize, callback);
}
void write(DataT data, WriteCallback callback) override {
return obj.write(data, callback);
}
DeriveT obj;
};
std::unique_ptr<Concept> m_d;
};
class FileProvider
{
public:
/**
* Constructor.
*
* Construct a FileProvider using `o`.
*
* `o` should be of a copyable type `T` that has a method
* `getStream()` that takes a FileOpenMode
* and returns something implicitly convertible
* to FileStream.
*
* That is, `[o]() -> FileStream { return o.getStream(FileOpenMode::Read); }`
* must be well-formed.
*
* In addition, if the FileOpenMode passed to getStream is FileOpenMode::Read,
* the returned stream must support read(); if the FileOpenMode passed to getStream
* is FileOpenMode::Write, the returned stream must support write().
*/
template<class DeriveT>
FileProvider(DeriveT &&o)
: m_d(std::make_unique<Model<std::decay_t<DeriveT>>>(
std::forward<DeriveT>(o))) {}
inline FileProvider(const FileProvider &that) : m_d(that.m_d->clone()) {}
inline FileProvider(FileProvider &&that) : m_d(std::move(that.m_d)) {}
inline FileProvider &operator=(const FileProvider &that) {
m_d = that.m_d->clone();
return *this;
}
inline FileProvider &operator=(FileProvider &&that) {
m_d = std::move(that.m_d);
return *this;
}
/**
* Get the file stream provided by this.
*
* @return a FileStream that will contain the content
* of the file provided by this.
*/
inline FileStream getStream(FileOpenMode mode = FileOpenMode::Read) const {
return m_d->getStream(mode);
}
private:
struct Concept
{
virtual ~Concept() = default;
virtual std::unique_ptr<Concept> clone() const = 0;
virtual FileStream getStream(FileOpenMode mode) const = 0;
};
template<class DeriveT>
struct Model : public Concept
{
explicit Model(const DeriveT &o) : obj(o) {}
explicit Model(DeriveT &&o) : obj(std::move(o)) {}
~Model() override = default;
std::unique_ptr<Concept> clone() const override {
return std::make_unique<Model>(obj);
}
FileStream getStream(FileOpenMode mode) const override {
return obj.getStream(mode);
}
DeriveT obj;
};
std::unique_ptr<Concept> m_d;
};
class FileDesc;
class FileInterface;
struct DumbFileStream
{
inline explicit DumbFileStream(FileContent r) : remaining(r) {}
template<class Callback>
void read(int maxSize, Callback callback) {
if (remaining.empty()) {
callback(FileOpRetCode::Eof, remaining);
} else {
auto taken = remaining.take(maxSize);
remaining = remaining.drop(maxSize);
callback(FileOpRetCode::Success, taken);
}
}
template<class Callback>
void write(FileContent data, Callback callback) {
remaining = remaining + data;
callback(FileOpRetCode::Success, data.size());
}
FileContent remaining;
};
struct DumbFileProvider
{
inline explicit DumbFileProvider(FileContent content)
: m_content(content)
{}
inline DumbFileStream getStream(FileOpenMode /* mode */) const {
return DumbFileStream(m_content);
}
FileContent m_content;
};
class FileDesc
{
public:
/**
* Construct an in-memory FileDesc with @c content.
*/
explicit FileDesc(FileContent content)
: m_inMemory{true}
, m_content(content)
{}
/**
* Construct a FileDesc with @c name and @c contentType.
*/
explicit FileDesc(std::string name,
std::optional<std::string> contentType = std::nullopt)
: m_inMemory{false}
, m_name(std::move(name))
, m_contentType(std::move(contentType)) {}
/**
* Get the FileProvider for this FileDesc using FileInterface.
*
* If this is an in-memory file, this will always return a
* DumbFileProvider. Otherwise, it will return what
* `fh.getProviderFor(*this)` returns.
*
* @return a FileProvider for this.
*/
FileProvider provider(const FileInterface &fh) const;
/**
* Get the name for this FileDesc.
*/
inline std::optional<std::string> name() const { return m_name; }
/**
* Get the content type for this FileDesc.
*/
inline std::optional<std::string> contentType() const { return m_contentType; }
- bool operator==(const FileDesc &that) const;
+ friend bool operator==(const FileDesc &a, const FileDesc &b) = default;
private:
bool m_inMemory;
FileContent m_content;
std::optional<std::string> m_name;
std::optional<std::string> m_contentType;
};
class DumbFileInterface
{
public:
DumbFileProvider getProviderFor(FileDesc desc) const;
};
class FileInterface
{
public:
/**
* Constructor.
*
* Construct a FileInterface using `o`.
*
* `o` should be of a type that has the getProviderFor()
* const method that takes a FileDesc and returns an object
* implicitly convertible to FileProvider. That is,
* the following should be well-formed:
* `[&o](FileDesc d) -> FileProvider { return static_cast<std::add_const_t<decltype(o)>>(o).getProviderFor(d); }`
*/
template<class DeriveT>
inline FileInterface(DeriveT &&o)
: m_d(std::make_unique<Model<std::decay_t<DeriveT>>>(
std::forward<DeriveT>(o))) {}
/**
* Get the provider for @c desc from this.
*
* @warning: You should not call this directly. This
* function is called by FileDesc::provider().
*/
inline FileProvider getProviderFor(FileDesc desc) const {
return m_d->getProviderFor(std::move(desc));
}
private:
struct Concept
{
virtual ~Concept() = default;
virtual FileProvider getProviderFor(FileDesc desc) const = 0;
};
template<class DeriveT>
struct Model : public Concept
{
explicit Model(const DeriveT &o) : obj(o) {}
explicit Model(DeriveT &&o) : obj(std::move(o)) {}
~Model() override = default;
FileProvider getProviderFor(FileDesc desc) const override {
return obj.getProviderFor(std::move(desc));
}
DeriveT obj;
};
std::unique_ptr<Concept> m_d;
};
}
diff --git a/src/base/jsonwrap.hpp b/src/base/jsonwrap.hpp
index e5a3453..58dc869 100644
--- a/src/base/jsonwrap.hpp
+++ b/src/base/jsonwrap.hpp
@@ -1,79 +1,74 @@
/*
* 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 <boost/serialization/version.hpp>
#include <boost/serialization/string.hpp>
#include <boost/serialization/split_member.hpp>
#include <nlohmann/json.hpp>
#include <immer/box.hpp>
namespace Kazv
{
using json = nlohmann::json;
class JsonWrap
{
// Cannot directly use box here, because it causes the resulting json
// to be wrapped into an array.
// https://github.com/arximboldi/immer/issues/155
struct Private
{
json j;
+ friend bool operator==(const Private &a, const Private &b) = default;
+ friend bool operator!=(const Private &a, const Private &b) = default;
};
immer::box<Private> m_d;
public:
JsonWrap() : m_d(Private{json()}) {}
JsonWrap(json&& j) : m_d(Private{std::move(j)}) {}
JsonWrap(const json& j) : m_d(Private{j}) {}
const json &get() const { return m_d.get().j; }
operator json() const { return m_d.get().j; }
template <class Archive>
void save(Archive &ar, std::uint32_t const /*version*/) const {
ar << get().dump();
}
template <class Archive>
void load(Archive &ar, std::uint32_t const /*version*/) {
std::string j;
ar >> j;
m_d = immer::box<Private>(Private{json::parse(std::move(j))});
}
BOOST_SERIALIZATION_SPLIT_MEMBER()
+ friend bool operator==(const JsonWrap &a, const JsonWrap &b) = default;
+ friend bool operator!=(const JsonWrap &a, const JsonWrap &b) = default;
};
}
BOOST_CLASS_VERSION(Kazv::JsonWrap, 0)
namespace nlohmann
{
template <>
struct adl_serializer<Kazv::JsonWrap> {
static void to_json(json& j, Kazv::JsonWrap w) {
j = w.get();
}
static void from_json(const json& j, Kazv::JsonWrap &w) {
w = Kazv::JsonWrap(j);
}
};
-
-}
-
-namespace Kazv
-{
- inline bool operator==(JsonWrap a, JsonWrap b)
- {
- return a.get() == b.get();
- }
}
diff --git a/src/client/actions/send.cpp b/src/client/actions/send.cpp
index b953126..9f3a692 100644
--- a/src/client/actions/send.cpp
+++ b/src/client/actions/send.cpp
@@ -1,283 +1,283 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2020 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <debug.hpp>
#include <types.hpp>
#include <immer-utils.hpp>
#include "send.hpp"
#include "status-utils.hpp"
namespace Kazv
{
ClientResult updateClient(ClientModel m, SendMessageAction a)
{
auto event = std::move(a.event);
auto roomId = a.roomId;
auto origJson = event.originalJson().get();
if (!origJson.contains("type") || !origJson.contains("content")) {
m.addTrigger(InvalidMessageFormat{});
return { std::move(m), lager::noop };
}
if (m.roomList.rooms[a.roomId].encrypted && !event.encrypted()) {
if (!event.isState()) {
return { std::move(m), [](auto &&) {
return EffectStatus{
/* succ = */ false,
json{
{"errorCode", "MOE_KAZV_MXC_SENDING_UNENCRYPTED_EVENT_TO_ENCRYPTED_ROOM"},
{"error", "Cannot send unencrypted event to encrypted room"},
}
};
}};
}
}
// We do not use event.type() etc. because we want
// encrypted events stay encrypted.
auto type = origJson["type"];
auto content = origJson["content"];
kzo.client.dbg() << "Sending message of type " << type
<< " with content " << content.dump()
<< " to " << a.roomId
<< " as #" << m.nextTxnId << std::endl;
// We combine the hash of json, the timestamp,
// and a numeric count in the client to avoid collision.
auto txnId = a.txnId.has_value() ? a.txnId.value() : getTxnId(event, m);
m.roomList = RoomListModel::update(std::move(m.roomList),
UpdateRoomAction{
roomId,
AddLocalEchoAction{{txnId, event}},
}
);
auto job = m.job<SendMessageJob>()
.make(a.roomId, type, txnId, content)
.withData(json{
{"roomId", a.roomId},
{"txnId", txnId},
});
m.addJob(std::move(job));
return { std::move(m), lager::noop };
}
ClientResult processResponse(ClientModel m, SendMessageResponse r)
{
auto roomId = r.dataStr("roomId");
if (! r.success()) {
auto txnId = r.dataStr("txnId");
kzo.client.dbg() << "Send message failed" << std::endl;
m.roomList.rooms = std::move(m.roomList.rooms).update(roomId, [txnId](auto room) {
auto maybeLocalEcho = room.getLocalEchoByTxnId(txnId);
if (!maybeLocalEcho.has_value()) {
kzo.client.warn() << "We do not have local echo with txnId " << txnId << " . Have we time-travelled?" << std::endl;
return room;
}
auto localEcho = std::move(maybeLocalEcho).value();
localEcho.status = LocalEchoDesc::Failed;
return RoomModel::update(std::move(room), AddLocalEchoAction{localEcho});
});
m.addTrigger(SendMessageFailed{roomId, r.errorCode(), r.errorMessage()});
return { std::move(m), failWithResponse(r) };
}
m.addTrigger(SendMessageSuccessful{roomId, r.eventId()});
return { std::move(m), lager::noop };
}
ClientResult updateClient(ClientModel m, SendToDeviceMessageAction a)
{
auto origJson = a.event.originalJson().get();
if (!origJson.contains("type") || !origJson.contains("content")) {
return { std::move(m), simpleFail };
}
// We do not use event.type() etc. because we want
// encrypted events stay encrypted.
auto type = origJson["type"];
auto content = origJson["content"];
if (type == "m.room_key" && !a.event.encrypted()) {
kzo.client.err() << "Trying to send room key event unencrypted! Rejecting." << std::endl;
return { std::move(m), failEffect(
"MOE_KAZV_MXC_SENDING_ROOM_KEY_EVENT_UNENCRYPTED",
"Cannot send room key event unencrypted"
)};
}
auto txnId = a.txnId.has_value() ? a.txnId.value() : getTxnId(a.event, m);
kzo.client.info() << "sending to-device message with txnId " << txnId;
auto messages =
immer::map<std::string, immer::map<std::string, JsonWrap>>{};
for (auto [userId, devices] : a.devicesToSend) {
auto deviceIdToContentMap = immer::map<std::string, JsonWrap>{};
for (auto deviceId : devices) {
deviceIdToContentMap = std::move(deviceIdToContentMap).set(deviceId, content);
}
messages = std::move(messages).set(userId, deviceIdToContentMap);
}
auto job = m.job<SendToDeviceJob>()
.make(type, txnId, messages)
.withData(json{{"devicesToSend", a.devicesToSend},
{"txnId", txnId}});
m.addJob(std::move(job));
return { std::move(m), lager::noop };
}
ClientResult updateClient(ClientModel m, SendMultipleToDeviceMessagesAction a)
{
std::optional<std::string> type;
std::optional<std::string> txnId;
immer::map<std::string, immer::map<std::string, JsonWrap>> userToDeviceToContentMap;
for (auto [userId, deviceToEventMap] : a.userToDeviceToEventMap) {
userToDeviceToContentMap = std::move(userToDeviceToContentMap).set(userId, immer::map<std::string, JsonWrap>());
for (auto [deviceId, event] : deviceToEventMap) {
auto origJson = event.originalJson().get();
if (!origJson.contains("type") || !origJson.contains("content")) {
return { std::move(m), failEffect(
"MOE_KAZV_MXC_INVALID_EVENT_FORMAT",
"Invalid event format"
)};
}
if (!type.has_value()) {
type = origJson["type"];
} else {
- if (type.value() != origJson["type"]) {
+ if (origJson["type"] != type.value()) {
return { std::move(m), failEffect(
"MOE_KAZV_MXC_TYPE_NOT_SAME",
"The to-device messages' type must be the same"
)};
}
}
if (origJson["type"] == "m.room_key" && !event.encrypted()) {
kzo.client.err() << "Trying to send room key event unencrypted! Rejecting." << std::endl;
return { std::move(m), failEffect(
"MOE_KAZV_MXC_SENDING_ROOM_KEY_EVENT_UNENCRYPTED",
"Cannot send room key event unencrypted"
)};
}
if (!txnId.has_value()) {
txnId = getTxnId(event, m);
}
auto content = origJson["content"];
userToDeviceToContentMap = setIn(
std::move(userToDeviceToContentMap),
std::move(content),
userId, deviceId
);
}
}
if (!type) {
// nothing to send
return { std::move(m), lager::noop };
}
auto job = m.job<SendToDeviceJob>()
.make(type.value(), txnId.value(), userToDeviceToContentMap)
// XXX not extracting devicesToSend from the map
// those triggers should be deprecated in favour of
// then()-continuation anyway
.withData(json{{"devicesToSend", json::object()},
{"txnId", txnId.value()}});
m.addJob(std::move(job));
return { std::move(m), lager::noop };
}
ClientResult processResponse(ClientModel m, SendToDeviceResponse r)
{
auto devicesToSend = r.dataJson("devicesToSend");
auto txnId = r.dataStr("txnId");
if (! r.success()) {
m.addTrigger(SendToDeviceMessageFailed{devicesToSend, txnId, r.errorCode(), r.errorMessage()});
return { std::move(m), failWithResponse(r) };
}
m.addTrigger(SendToDeviceMessageSuccessful{devicesToSend, txnId});
return { std::move(m), lager::noop };
}
ClientResult updateClient(ClientModel m, SaveLocalEchoAction a)
{
auto txnId = a.txnId.has_value() ? a.txnId.value() : getTxnId(a.event, m);
m.roomList = RoomListModel::update(m.roomList, UpdateRoomAction{
a.roomId,
AddLocalEchoAction{{txnId, a.event}},
});
return { std::move(m), [txnId](auto) {
return EffectStatus(true, json::object({{"txnId", txnId}}));
}};
}
ClientResult updateClient(ClientModel m, UpdateLocalEchoStatusAction a)
{
auto txnId = a.txnId;
auto maybeLocalEcho = m.roomList.rooms[a.roomId].getLocalEchoByTxnId(txnId);
if (!maybeLocalEcho.has_value()) {
return { std::move(m), [](const auto &) {
return EffectStatus(/* succ = */ false, json{
{"errorCode", "MOE_KAZV_MXC_LOCAL_ECHO_NOT_FOUND"},
{"error", "Local echo not found"},
});
} };
}
auto localEcho = maybeLocalEcho.value();
m.roomList = RoomListModel::update(m.roomList, UpdateRoomAction{
a.roomId,
AddLocalEchoAction{{txnId, localEcho.event, a.status}},
});
return { std::move(m), lager::noop };
}
ClientResult updateClient(ClientModel m, RedactEventAction a)
{
auto txnId = getTxnId(Event{json{
{"room_id", a.roomId},
{"event_id", a.eventId},
}}, m);
auto job = m.job<RedactEventJob>()
.make(a.roomId, a.eventId, txnId, a.reason);
m.addJob(std::move(job));
return { std::move(m), lager::noop };
}
ClientResult processResponse(ClientModel m, RedactEventResponse r)
{
if (!r.success()) {
return { std::move(m), failWithResponse(r) };
}
return { std::move(m), lager::noop };
}
}
diff --git a/src/client/clientutil.hpp b/src/client/clientutil.hpp
index 57985c9..b9e636d 100644
--- a/src/client/clientutil.hpp
+++ b/src/client/clientutil.hpp
@@ -1,235 +1,231 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021-2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include <libkazv-config.hpp>
#include <string>
#include <tuple>
#include <immer/map.hpp>
#include <zug/transducer/filter.hpp>
#include <zug/transducer/eager.hpp>
#include <lager/deps.hpp>
#include <boost/container_hash/hash.hpp>
#include <boost/serialization/string.hpp>
#include <cursorutil.hpp>
#include <jobinterface.hpp>
#include <eventinterface.hpp>
#include "thread-safety-helper.hpp"
namespace Kazv
{
struct ClientModel;
template<class K, class V, class List, class Func>
immer::map<K, V> merge(immer::map<K, V> map, List list, Func keyOf)
{
for (auto v : list) {
auto key = keyOf(v);
map = std::move(map).set(key, v);
}
return map;
}
inline std::string keyOfPresence(Event e) {
return e.sender();
}
inline std::string keyOfAccountData(Event e) {
return e.type();
}
inline std::string keyOfTimeline(Event e) {
return e.id();
}
inline std::string keyOfEphemeral(Event e) {
return e.type();
}
struct KeyOfState {
std::string type;
std::string stateKey;
+ friend bool operator==(const KeyOfState &a, const KeyOfState &b) = default;
};
template<class Archive>
void serialize(Archive &ar, KeyOfState &m, std::uint32_t const /* version */)
{
ar & m.type & m.stateKey;
}
- inline bool operator==(KeyOfState a, KeyOfState b)
- {
- return a.type == b.type && a.stateKey == b.stateKey;
- }
-
inline KeyOfState keyOfState(Event e) {
return {e.type(), e.stateKey()};
}
template<class Context>
JobInterface &getJobHandler(Context &&ctx)
{
return lager::get<JobInterface &>(std::forward<Context>(ctx));
}
template<class Context>
EventInterface &getEventEmitter(Context &&ctx)
{
return lager::get<EventInterface &>(std::forward<Context>(ctx));
}
namespace
{
template<class ImmerT>
struct ImmerIterator
{
using value_type = typename ImmerT::value_type;
using reference = typename ImmerT::reference;
using pointer = const value_type *;
using difference_type = long int;
using iterator_category = std::random_access_iterator_tag;
ImmerIterator(const ImmerT &container, std::size_t index)
: m_container(std::ref(container))
, m_index(index)
{}
ImmerIterator &operator+=(difference_type d) {
m_index += d;
return *this;
}
ImmerIterator &operator-=(difference_type d) {
m_index -= d;
return *this;
}
difference_type operator-(ImmerIterator b) const {
return index() - b.index();
}
ImmerIterator &operator++() {
return *this += 1;
}
ImmerIterator operator++(int) {
auto tmp = *this;
*this += 1;
return tmp;
}
ImmerIterator &operator--() {
return *this -= 1;
}
ImmerIterator operator--(int) {
auto tmp = *this;
*this -= 1;
return tmp;
}
reference &operator*() const {
return m_container.get().at(m_index);
}
reference operator[](difference_type d) const;
std::size_t index() const { return m_index; }
private:
std::reference_wrapper<const ImmerT> m_container;
std::size_t m_index;
};
template<class ImmerT>
auto ImmerIterator<ImmerT>::operator[](difference_type d) const -> reference
{
return *(*this + d);
}
template<class ImmerT>
auto operator+(ImmerIterator<ImmerT> a, long int d)
{
return a += d;
};
template<class ImmerT>
auto operator+(long int d, ImmerIterator<ImmerT> a)
{
return a += d;
};
template<class ImmerT>
auto operator-(ImmerIterator<ImmerT> a, long int d)
{
return a -= d;
};
template<class ImmerT>
auto immerBegin(const ImmerT &c)
{
return ImmerIterator<ImmerT>(c, 0);
}
template<class ImmerT>
auto immerEnd(const ImmerT &c)
{
return ImmerIterator<ImmerT>(c, c.size());
}
}
template<class ImmerT1, class RangeT2, class Pred, class Func>
ImmerT1 sortedUniqueMerge(ImmerT1 base, RangeT2 addon, Pred exists, Func keyOf)
{
auto needToAdd = intoImmer(ImmerT1{},
zug::filter([=](auto a) {
return !exists(a);
}),
addon);
auto cmp = [=](auto a, auto b) {
return keyOf(a) < keyOf(b);
};
for (auto item : needToAdd) {
auto it = std::upper_bound(immerBegin(base), immerEnd(base), item, cmp);
auto index = it.index();
base = std::move(base).insert(index, item);
}
return base;
}
std::string increaseTxnId(std::string cur);
std::string getTxnId(Event event, ClientModel &m);
}
namespace std
{
template<> struct hash<Kazv::KeyOfState>
{
std::size_t operator()(const Kazv::KeyOfState & k) const noexcept {
std::size_t seed = 0;
boost::hash_combine(seed, k.type);
boost::hash_combine(seed, k.stateKey);
return seed;
}
};
}
#define KAZV_WRAP_ATTR(_type, _d, _attr) \
inline auto _attr() const { \
KAZV_VERIFY_THREAD_ID(); \
return (_d)[&_type::_attr]; \
}
BOOST_CLASS_VERSION(Kazv::KeyOfState, 0)
diff --git a/src/client/device-list-tracker.cpp b/src/client/device-list-tracker.cpp
index 37d2a85..362be1e 100644
--- a/src/client/device-list-tracker.cpp
+++ b/src/client/device-list-tracker.cpp
@@ -1,196 +1,182 @@
/*
* 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 "device-list-tracker.hpp"
#include <algorithm>
#include <immer/flex_vector_transient.hpp>
#include <zug/transducer/filter.hpp>
#include <zug/sequence.hpp>
#include <zug/transducer/distinct.hpp>
#include <zug/transducer/chain.hpp>
#include <debug.hpp>
namespace Kazv
{
immer::flex_vector<std::string> DeviceListTracker::outdatedUsers() const
{
return intoImmer(
immer::flex_vector<std::string>{},
zug::filter([](auto n) {
auto [userId, outdated] = n;
return outdated;
})
| zug::map([](auto n) {
auto [userId, outdated] = n;
return userId;
}),
usersToTrackDeviceLists);
}
bool DeviceListTracker::addDevice(std::string userId, std::string deviceId, Api::QueryKeysJob::DeviceInformation deviceInfo, Crypto &crypto)
{
using namespace CryptoConstants;
if (userId != deviceInfo.userId
|| deviceId != deviceInfo.deviceId) {
return false;
}
// if the ed25519 key changed, reject
auto curEd25519Key = deviceInfo.keys[ed25519 + ":" + deviceId];
if (deviceLists[userId].find(deviceId)
&& curEd25519Key != deviceLists[userId][deviceId].ed25519Key) {
return false;
}
kzo.client.dbg() << "verifying device info" << std::endl;
if (crypto.verify(deviceInfo, userId, deviceId, curEd25519Key)) {
kzo.client.dbg() << "passed verification" << std::endl;
auto info = DeviceKeyInfo{
deviceId,
deviceInfo.keys[ed25519 + ":" + deviceId],
deviceInfo.keys[curve25519 + ":" + deviceId],
deviceInfo.unsignedData ? deviceInfo.unsignedData.value().deviceDisplayName : std::nullopt
};
deviceLists = std::move(deviceLists)
.update(userId, [=](auto deviceMap) {
return std::move(deviceMap).set(deviceId, info);
});
return true;
}
kzo.client.dbg() << "did not pass verification" << std::endl;
return false;
}
void DeviceListTracker::markUpToDate(std::string userId)
{
usersToTrackDeviceLists = std::move(usersToTrackDeviceLists).set(userId, false);
}
std::optional<DeviceKeyInfo> DeviceListTracker::get(std::string userId, std::string deviceId) const
{
try {
return deviceLists.at(userId).at(deviceId);
} catch (const std::exception &) {
return std::nullopt;
}
}
std::optional<DeviceKeyInfo> DeviceListTracker::findByEd25519Key(
std::string userId, std::string ed25519Key) const
{
auto devices = deviceLists.at(userId);
auto it = std::find_if(devices.begin(), devices.end(),
[=](auto n) {
auto [deviceId, info] = n;
return info.ed25519Key == ed25519Key;
});
if (it != devices.end()) {
return it->second;
} else {
return std::nullopt;
}
}
std::optional<DeviceKeyInfo> DeviceListTracker::findByCurve25519Key(
std::string userId, std::string curve25519Key) const
{
if (!deviceLists.count(userId)) {
return std::nullopt;
}
auto devices = deviceLists.at(userId);
auto it = std::find_if(devices.begin(), devices.end(),
[=](auto n) {
auto [deviceId, info] = n;
return info.curve25519Key == curve25519Key;
});
if (it != devices.end()) {
return it->second;
} else {
return std::nullopt;
}
}
static bool cryptographicallyEqual(DeviceKeyInfo a, DeviceKeyInfo b)
{
a.displayName = std::nullopt;
b.displayName = std::nullopt;
a.trustLevel = Unseen;
b.trustLevel = Unseen;
return std::move(a) == std::move(b);
}
static bool cryptographicallyEqual(const DeviceListTracker::DeviceMapT &a, const DeviceListTracker::DeviceMapT &b)
{
auto changed = false;
auto markChanged = [&changed](const auto &) { changed = true; };
immer::diff(a, b, immer::make_differ(
/* added = */ markChanged,
/* removed = */ markChanged,
/* changed = */ [&changed](const auto &x, const auto &y) {
if (!cryptographicallyEqual(x.second, y.second)) {
changed = true;
}
}
));
return !changed;
}
immer::flex_vector<std::string> DeviceListTracker::diff(DeviceListTracker that) const
{
auto changedUsers = immer::flex_vector_transient<std::string>{};
immer::diff(
that.deviceLists, deviceLists,
immer::make_differ(
/* addedFn = */ [&changedUsers](const auto &pair) {
changedUsers.push_back(pair.first);
},
/* removedFn = */ [&changedUsers](const auto &pair) {
changedUsers.push_back(pair.first);
},
/* changedFn = */ [&changedUsers](const auto &pairA, const auto &pairB) {
// add to changed user list only if device id or trust level changes
if (!cryptographicallyEqual(pairA.second, pairB.second)) {
changedUsers.push_back(pairA.first);
}
}
)
);
return changedUsers.persistent();
}
auto DeviceListTracker::devicesFor(std::string userId) const -> DeviceMapT
{
return deviceLists[userId];
}
-
- bool operator==(DeviceKeyInfo a, DeviceKeyInfo b)
- {
- return a.deviceId == b.deviceId
- && a.ed25519Key == b.ed25519Key
- && a.curve25519Key == b.curve25519Key
- && a.displayName == b.displayName
- && a.trustLevel == b.trustLevel;
- }
-
- bool operator!=(DeviceKeyInfo a, DeviceKeyInfo b)
- {
- return !(a == b);
- }
}
diff --git a/src/client/device-list-tracker.hpp b/src/client/device-list-tracker.hpp
index 8613e81..6edbd55 100644
--- a/src/client/device-list-tracker.hpp
+++ b/src/client/device-list-tracker.hpp
@@ -1,107 +1,106 @@
/*
* 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 <string>
#include <immer/map.hpp>
#include <immer/flex_vector.hpp>
#include <boost/serialization/string.hpp>
#include <serialization/std-optional.hpp>
#include <serialization/immer-map.hpp>
#include <crypto.hpp>
#include <csapi/keys.hpp>
#include "cursorutil.hpp"
namespace Kazv
{
enum DeviceTrustLevel
{
Blocked,
Unseen,
Seen,
Verified,
};
struct DeviceKeyInfo
{
std::string deviceId;
std::string ed25519Key;
std::string curve25519Key;
std::optional<std::string> displayName;
DeviceTrustLevel trustLevel{Unseen};
+ friend bool operator==(const DeviceKeyInfo &a, const DeviceKeyInfo &b) = default;
+ friend bool operator!=(const DeviceKeyInfo &a, const DeviceKeyInfo &b) = default;
};
- bool operator==(DeviceKeyInfo a, DeviceKeyInfo b);
- bool operator!=(DeviceKeyInfo a, DeviceKeyInfo b);
-
template<class Archive>
void serialize(Archive &ar, DeviceKeyInfo &i, std::uint32_t const /*version*/)
{
ar
& i.deviceId
& i.ed25519Key
& i.curve25519Key
& i.displayName
& i.trustLevel
;
}
struct DeviceListTracker
{
using DeviceMapT = immer::map<std::string /* deviceId */, DeviceKeyInfo>;
immer::map<std::string /* userId */, bool /* outdated */> usersToTrackDeviceLists;
immer::map<std::string /* userId */, DeviceMapT> deviceLists;
template<class RangeT>
void track(RangeT &&userIds) {
for (auto userId : std::forward<RangeT>(userIds)) {
usersToTrackDeviceLists = std::move(usersToTrackDeviceLists)
.set(userId, true);
}
}
template<class RangeT>
void untrack(RangeT &&userIds) {
for (auto userId : std::forward<RangeT>(userIds)) {
usersToTrackDeviceLists = std::move(usersToTrackDeviceLists).erase(userId);
}
}
immer::flex_vector<std::string> outdatedUsers() const;
bool addDevice(std::string userId, std::string deviceId, Api::QueryKeysJob::DeviceInformation deviceInfo, Crypto &crypto);
void markUpToDate(std::string userId);
DeviceMapT devicesFor(std::string userId) const;
std::optional<DeviceKeyInfo> get(std::string userId, std::string deviceId) const;
std::optional<DeviceKeyInfo> findByEd25519Key(std::string userId, std::string ed25519Key) const;
std::optional<DeviceKeyInfo> findByCurve25519Key(std::string userId, std::string curve25519Key) const;
/// returns a list of users whose device list has changed
immer::flex_vector<std::string> diff(DeviceListTracker that) const;
};
template<class Archive>
void serialize(Archive &ar, DeviceListTracker &t, std::uint32_t const /*version*/)
{
ar
& t.usersToTrackDeviceLists
& t.deviceLists
;
}
}
BOOST_CLASS_VERSION(Kazv::DeviceKeyInfo, 0)
BOOST_CLASS_VERSION(Kazv::DeviceListTracker, 0)
diff --git a/src/client/encrypted-file.cpp b/src/client/encrypted-file.cpp
index e0d7a82..1101e4f 100644
--- a/src/client/encrypted-file.cpp
+++ b/src/client/encrypted-file.cpp
@@ -1,111 +1,111 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include "encrypted-file.hpp"
namespace Kazv
{
struct EncryptedFileDesc::Private
{
std::string mxcUri;
std::string key;
std::string iv;
std::string sha256Hash;
+ friend bool operator==(const Private &a, const Private &b) = default;
};
EncryptedFileDesc::EncryptedFileDesc()
: m_d(std::make_unique<Private>())
{
}
EncryptedFileDesc::EncryptedFileDesc(std::string mxcUri, std::string key, std::string iv, std::string sha256Hash)
: m_d(std::make_unique<Private>(Private{mxcUri, key, iv, sha256Hash}))
{
}
KAZV_DEFINE_COPYABLE_UNIQUE_PTR(EncryptedFileDesc, m_d)
EncryptedFileDesc::~EncryptedFileDesc() = default;
EncryptedFileDesc EncryptedFileDesc::fromJson(JsonWrap encryptedFile)
{
try {
// First, verify EncryptedFile object
auto jwk = encryptedFile.get().at("key");
auto keyOps = jwk.at("key_ops");
if (! (encryptedFile.get().at("v").template get<std::string>() == "v2"
&& jwk.at("kty").template get<std::string>() == "oct"
&& jwk.at("alg").template get<std::string>() == "A256CTR"
&& jwk.at("ext").template get<bool>() == true
&& std::find(keyOps.cbegin(), keyOps.cend(), "encrypt") != keyOps.cend()
&& std::find(keyOps.cbegin(), keyOps.cend(), "decrypt") != keyOps.cend())) {
// invalid EncryptedFile object
return EncryptedFileDesc();
}
auto mxcUri = encryptedFile.get().at("url").template get<std::string>();
auto key = encryptedFile.get().at("key").at("k").template get<std::string>();
auto iv = encryptedFile.get().at("iv").template get<std::string>();
auto sha256Hash = encryptedFile.get().at("hashes").at("sha256").template get<std::string>();
return EncryptedFileDesc(mxcUri, key, iv, sha256Hash);
} catch (const std::exception &) {
return EncryptedFileDesc();
}
}
JsonWrap EncryptedFileDesc::toJson() const
{
return json::object({
{"url", mxcUri()},
{"key", json::object({
{"kty", "oct"},
{"key_ops", json::array({"encrypt", "decrypt"})},
{"alg", "A256CTR"},
{"k", key()},
{"ext", true},
})},
{"iv", iv()},
{"hashes", json::object({
{"sha256", sha256Hash()}
})},
{"v", "v2"},
});
}
std::string EncryptedFileDesc::mxcUri() const
{
if (! m_d) { return ""; }
return m_d->mxcUri;
}
std::string EncryptedFileDesc::key() const
{
if (! m_d) { return ""; }
return m_d->key;
}
std::string EncryptedFileDesc::iv() const
{
if (! m_d) { return ""; }
return m_d->iv;
}
std::string EncryptedFileDesc::sha256Hash() const
{
if (! m_d) { return ""; }
return m_d->sha256Hash;
}
bool EncryptedFileDesc::operator==(const EncryptedFileDesc &that) const
{
- return (m_d == that.m_d
- || (mxcUri() == that.mxcUri()
- && key() == that.key()
- && iv() == that.iv()
- && sha256Hash() == that.sha256Hash()));
+ return m_d == that.m_d || (
+ m_d && that.m_d
+ && *(m_d) == *(that.m_d)
+ );
}
}
diff --git a/src/client/room/local-echo.hpp b/src/client/room/local-echo.hpp
index db9ffd2..136a516 100644
--- a/src/client/room/local-echo.hpp
+++ b/src/client/room/local-echo.hpp
@@ -1,79 +1,69 @@
/*
* 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 <boost/serialization/split_free.hpp>
#include "event.hpp"
namespace Kazv
{
/**
* Describes a local echo.
*/
struct LocalEchoDesc
{
enum Status
{
Sending,
Failed,
};
std::string txnId;
Event event;
Status status{Sending};
+ friend bool operator==(const LocalEchoDesc &a, const LocalEchoDesc &b) = default;
+ friend bool operator!=(const LocalEchoDesc &a, const LocalEchoDesc &b) = default;
};
-
- inline bool operator==(const LocalEchoDesc &a, const LocalEchoDesc &b)
- {
- return a.txnId == b.txnId
- && a.event == b.event
- && a.status == b.status;
- }
-
- inline bool operator!=(const LocalEchoDesc &a, const LocalEchoDesc &b)
- {
- return !(a == b);
- }
}
namespace boost::serialization
{
template<class Archive>
void save(Archive &ar, const Kazv::LocalEchoDesc &d, std::uint32_t const version)
{
using namespace Kazv;
LocalEchoDesc::Status dummyStatus{LocalEchoDesc::Failed};
ar
& d.txnId
& d.event
& dummyStatus
;
}
template<class Archive>
void load(Archive &ar, Kazv::LocalEchoDesc &d, std::uint32_t const version)
{
using namespace Kazv;
LocalEchoDesc::Status dummyStatus{LocalEchoDesc::Failed};
ar
& d.txnId
& d.event
& dummyStatus
;
d.status = LocalEchoDesc::Failed;
}
template<class Archive>
void serialize(Archive &ar, Kazv::LocalEchoDesc &d, std::uint32_t const version)
{
using boost::serialization::split_free;
split_free(ar, d, version);
}
}
BOOST_CLASS_VERSION(Kazv::LocalEchoDesc, 0)
diff --git a/src/client/room/room-model.cpp b/src/client/room/room-model.cpp
index 64f9551..7c33cbd 100644
--- a/src/client/room/room-model.cpp
+++ b/src/client/room/room-model.cpp
@@ -1,703 +1,683 @@
/*
* 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 <lager/util.hpp>
#include <zug/sequence.hpp>
#include <zug/transducer/map.hpp>
#include <zug/transducer/filter.hpp>
#include "debug.hpp"
#include "room-model.hpp"
#include "cursorutil.hpp"
#include "immer-utils.hpp"
inline const auto receiptTypes = immer::flex_vector<std::string>{"m.read", "m.read.private"};
template<class Func>
static std::string getMaxInTimeline(std::string a, std::string b, Func sortKey)
{
if (a.empty() && b.empty()) {
return std::string();
} else {
// for an unexisting event, event id is empty and timestamp is 0
// for an existing event, event id is not empty and timestamp >= 0
// so this handles all cases even when we do not have the corresponding event
return std::max(a, b, [=](const std::string &x, const std::string &y) {
return sortKey(x) < sortKey(y);
});
}
}
namespace Kazv
{
PendingRoomKeyEvent makePendingRoomKeyEventV0(std::string txnId, Event event, immer::map<std::string, immer::flex_vector<std::string>> devices)
{
immer::map<std::string, immer::map<std::string, Event>> messages;
for (auto [userId, deviceIds] : devices) {
messages = setIn(std::move(messages), immer::map<std::string, Event>(), userId);
for (auto deviceId : deviceIds) {
messages = setIn(
std::move(messages),
event,
userId, deviceId
);
}
}
return PendingRoomKeyEvent{txnId, messages};
}
- bool operator==(const ReadReceipt &a, const ReadReceipt &b)
- {
- return a.eventId == b.eventId && a.timestamp == b.timestamp;
- }
-
- bool operator!=(const ReadReceipt &a, const ReadReceipt &b)
- {
- return !(a == b);
- }
-
- bool operator==(const EventReader &a, const EventReader &b)
- {
- return a.userId == b.userId && a.timestamp == b.timestamp;
- }
-
- bool operator!=(const EventReader &a, const EventReader &b)
- {
- return !(a == b);
- }
-
auto sortKeyForTimelineEvent(Event e) -> std::tuple<Timestamp, std::string>
{
return std::make_tuple(e.originServerTs(), e.id());
}
RoomModel RoomModel::update(RoomModel r, Action a)
{
return lager::match(std::move(a))(
[&](AddStateEventsAction a) {
r.stateEvents = merge(std::move(r.stateEvents), a.stateEvents, keyOfState);
// If m.room.encryption state event appears,
// configure the room to use encryption.
if (r.stateEvents.find(KeyOfState{"m.room.encryption", ""})) {
auto newRoom = update(std::move(r), SetRoomEncryptionAction{});
r = std::move(newRoom);
}
return r;
},
[&](MaybeAddStateEventsAction a) {
for (auto it = a.stateEvents.rbegin();
it != a.stateEvents.rend();
++it) {
const auto &e = *it;
auto k = keyOfState(e);
if (!r.stateEvents.count(k)) {
r.stateEvents = std::move(r.stateEvents).set(k, e);
}
}
return r;
},
[&](AddToTimelineAction a) {
auto eventIds = intoImmer(immer::flex_vector<std::string>(),
zug::map(keyOfTimeline), a.events);
auto oldMessages = r.messages;
r.messages = merge(std::move(r.messages), a.events, keyOfTimeline);
auto exists =
[=](auto eventId) -> bool {
return !! oldMessages.find(eventId);
};
auto key =
[=](auto eventId) {
// sort first by timestamp, then by id
return sortKeyForTimelineEvent(r.messages[eventId]);
};
auto handleRedaction =
[&r](const auto &event) {
if (event.type() == "m.room.redaction") {
auto origJson = event.originalJson().get();
if (origJson.contains("redacts") && origJson.at("redacts").is_string()) {
auto redactedEventId = origJson.at("redacts").template get<std::string>();
if (r.messages.find(redactedEventId)) {
r.messages = std::move(r.messages).update(redactedEventId, [&origJson](const auto &eventToBeRedacted) {
auto newJson = eventToBeRedacted.originalJson().get();
newJson.merge_patch(json{
{"unsigned", {{"redacted_because", std::move(origJson)}}},
});
newJson["content"] = json::object();
return Event(newJson);
});
}
}
}
return event;
};
immer::for_each(a.events, handleRedaction);
r.timeline = sortedUniqueMerge(r.timeline, eventIds, exists, key);
// If this is a pagination request, gapEventId
// should have value. If this is a sync request,
// gapEventId does not have value. The pagination
// request does not have the limited field, and
// whether it has more paginate back token is determined
// by the presence of the prevBatch parameter.
// In sync request, limited may not be specified,
// and thus, if limited does not have value, it means
// it is not limited.
if (((a.limited.has_value() && a.limited.value())
|| a.gapEventId.has_value())
&& a.prevBatch.has_value()) {
// this sync is limited, add a Gap here
if (!eventIds.empty()) {
r.timelineGaps = std::move(r.timelineGaps).set(eventIds[0], a.prevBatch.value());
}
}
// remove the original Gap, as it is resolved
if (a.gapEventId.has_value()) {
r.timelineGaps = std::move(r.timelineGaps).erase(a.gapEventId.value());
}
// remove all Gaps between the gapped event and the first event in this batch
if (!eventIds.empty() && a.gapEventId.has_value()) {
auto cmp = [=](auto a, auto b) {
return key(a) < key(b);
};
auto thisBatchStart = std::equal_range(r.timeline.begin(), r.timeline.end(), eventIds[0], cmp).first;
auto origBatchStart = std::equal_range(thisBatchStart, r.timeline.end(), a.gapEventId.value(), cmp).first;
// Safety assert: we do not want to execute the for_each if the range is empty,
// or it will go out of bounds.
if (thisBatchStart.index() < origBatchStart.index()) {
std::for_each(thisBatchStart + 1, origBatchStart,
[&](auto eventId) {
r.timelineGaps = std::move(r.timelineGaps).erase(eventId);
});
}
}
// remove all local echoes that are received
for (const auto &e : a.events) {
auto jw = e.originalJson();
const auto &json = jw.get();
if (json.contains("unsigned")
&& json["unsigned"].contains("transaction_id")
&& json["unsigned"]["transaction_id"].is_string()) {
r = update(std::move(r), RemoveLocalEchoAction{json["unsigned"]["transaction_id"].template get<std::string>()});
}
}
// calculate event relationships
r.generateRelationships(a.events);
r.addToUndecryptedEvents(a.events);
return r;
},
[&](AddAccountDataAction a) {
r.accountData = merge(std::move(r.accountData), a.events, keyOfAccountData);
return r;
},
[&](ChangeMembershipAction a) {
r.membership = a.membership;
return r;
},
[&](ChangeInviteStateAction a) {
r.inviteState = merge(immer::map<KeyOfState, Event>{}, a.events, keyOfState);
return r;
},
[&](AddEphemeralAction a) {
auto processReceipt = [&](Event e) {
const auto content = e.content().get();
for (auto [eventId, receipts] : content.items()) {
if (!receipts.is_object()) {
continue;
}
for (auto receiptType : receiptTypes) {
if (!(receipts.contains(receiptType)
&& receipts[receiptType].is_object())) {
continue;
}
for (auto [user, receipt]: receipts[receiptType].items()) {
ReadReceipt readReceipt{
eventId,
0,
};
if (receipt.is_object() && receipt.contains("ts")
&& receipt["ts"].is_number()) {
readReceipt.timestamp = receipt["ts"].template get<Timestamp>();
}
// Remove old receipts
if (r.readReceipts.count(user)) {
auto oldReceiptEventId = r.readReceipts[user].eventId;
if (r.eventReadUsers.count(oldReceiptEventId)) {
auto remaining =
intoImmer(
immer::flex_vector<std::string>{},
zug::filter([user=user](auto userId) {
return userId != user;
}),
r.eventReadUsers[oldReceiptEventId]
);
if (remaining.empty()) {
r.eventReadUsers = std::move(r.eventReadUsers).erase(oldReceiptEventId);
} else {
r.eventReadUsers = std::move(r.eventReadUsers).set(oldReceiptEventId, remaining);
}
}
}
// Add new receipt
r.readReceipts = std::move(r.readReceipts).set(user, readReceipt);
auto oldReadUsers = r.eventReadUsers[eventId];
r.eventReadUsers = std::move(r.eventReadUsers).set(eventId, oldReadUsers.push_back(user));
}
}
}
};
for (auto e : a.events) {
if (e.type() == "m.receipt") {
processReceipt(e);
}
}
r.ephemeral = merge(std::move(r.ephemeral), a.events, keyOfEphemeral);
return r;
},
[&](SetLocalDraftAction a) {
r.localDraft = a.localDraft;
return r;
},
[&](SetRoomEncryptionAction) {
r.encrypted = true;
return r;
},
[&](MarkMembersFullyLoadedAction) {
r.membersFullyLoaded = true;
return r;
},
[&](SetHeroIdsAction a) {
r.heroIds = a.heroIds;
return r;
},
[&](AddLocalEchoAction a) {
auto it = std::find_if(r.localEchoes.begin(), r.localEchoes.end(), [a](const auto &desc) {
return desc.txnId == a.localEcho.txnId;
});
if (it == r.localEchoes.end()) {
r.localEchoes = std::move(r.localEchoes).push_back(a.localEcho);
} else {
r.localEchoes = std::move(r.localEchoes).set(it.index(), a.localEcho);
}
return r;
},
[&](RemoveLocalEchoAction a) {
auto it = std::find_if(r.localEchoes.begin(), r.localEchoes.end(), [a](const auto &desc) {
return desc.txnId == a.txnId;
});
if (it != r.localEchoes.end()) {
r.localEchoes = std::move(r.localEchoes).erase(it.index());
}
return r;
},
[&](AddPendingRoomKeyAction a) {
auto it = std::find_if(r.pendingRoomKeyEvents.begin(), r.pendingRoomKeyEvents.end(), [a](const auto &p) {
return p.txnId == a.pendingRoomKeyEvent.txnId;
});
if (it == r.pendingRoomKeyEvents.end()) {
r.pendingRoomKeyEvents = std::move(r.pendingRoomKeyEvents).push_back(a.pendingRoomKeyEvent);
} else {
r.pendingRoomKeyEvents = std::move(r.pendingRoomKeyEvents).set(it.index(), a.pendingRoomKeyEvent);
}
return r;
},
[&](RemovePendingRoomKeyAction a) {
auto it = std::find_if(r.pendingRoomKeyEvents.begin(), r.pendingRoomKeyEvents.end(), [a](const auto &desc) {
return desc.txnId == a.txnId;
});
if (it != r.pendingRoomKeyEvents.end()) {
r.pendingRoomKeyEvents = std::move(r.pendingRoomKeyEvents).erase(it.index());
}
return r;
},
[&](UpdateJoinedMemberCountAction a) {
r.joinedMemberCount = a.joinedMemberCount;
return r;
},
[&](UpdateInvitedMemberCountAction a) {
r.invitedMemberCount = a.invitedMemberCount;
return r;
},
[&](AddLocalNotificationsAction a) {
auto k = [r](const auto &id) {
return sortKeyForTimelineEvent(r.messages[id]);
};
auto readReceiptForCurrentUser = getMaxInTimeline(r.readReceipts[a.myUserId].eventId, r.localReadMarker, k);
auto newEventIds = intoImmer(
immer::flex_vector<std::string>{},
zug::map(&Event::id),
a.newEvents
);
auto needToAddPredicate = [a, r, k, readReceiptForCurrentUser](const auto &eid) {
if (!readReceiptForCurrentUser.empty()
&& k(readReceiptForCurrentUser) >= k(eid)) {
// this means this event is already read
return false;
}
auto e = r.messages[eid];
return e.sender() != a.myUserId
&& a.pushRulesDesc.handle(e, r).shouldNotify;
};
r.unreadNotificationEventIds = sortedUniqueMerge(std::move(r.unreadNotificationEventIds), newEventIds, [needToAddPredicate](const auto &e) { return !needToAddPredicate(e); }, k);
return r;
},
[&](RemoveReadLocalNotificationsAction a) {
auto k = [r](const auto &id) {
return sortKeyForTimelineEvent(r.messages[id]);
};
auto cmp = [k](const auto &a, const auto &b) {
return k(a) < k(b);
};
auto rr = getMaxInTimeline(
r.readReceipts[a.myUserId].eventId,
r.localReadMarker,
k
);
if (rr.empty()) {
return r;
}
auto it = std::upper_bound(
r.unreadNotificationEventIds.begin(),
r.unreadNotificationEventIds.end(),
rr,
cmp
);
// *it > rr, *(it - 1) <= rr (if it - 1 is valid)
if (it == r.unreadNotificationEventIds.end()) {
// If it == end(), it means everything is read
r.unreadNotificationEventIds = {};
} else if (it == r.unreadNotificationEventIds.begin()) {
// If it == begin(), it means everything is unread, so nothing to do
} else {
// it is somewhere in the middle, pointing to the first element that is unread
r.unreadNotificationEventIds = std::move(r.unreadNotificationEventIds).erase(0, it.index());
}
return r;
},
[&](UpdateLocalReadMarkerAction a) {
r.localReadMarker = a.localReadMarker;
auto next = RoomModel::update(std::move(r), RemoveReadLocalNotificationsAction{a.myUserId});
return next;
}
);
}
RoomListModel RoomListModel::update(RoomListModel l, Action a)
{
return lager::match(std::move(a))(
[&](UpdateRoomAction a) {
l.rooms = std::move(l.rooms)
.update(a.roomId,
[=](RoomModel oldRoom) {
oldRoom.roomId = a.roomId; // in case it is a new room
return RoomModel::update(std::move(oldRoom), a.roomAction);
});
return l;
}
);
}
static auto membershipTransducer(const std::string &membership)
{
return zug::filter([](auto val) {
auto [k, v] = val;
auto [type, stateKey] = k;
return type == "m.room.member"s;
})
| zug::map([](auto val) {
auto [k, v] = val;
auto [type, stateKey] = k;
return std::pair<std::string, Kazv::Event>{stateKey, v};
})
| zug::filter([&membership](auto val) {
auto [stateKey, ev] = val;
return ev.content().get()
.at("membership"s) == membership;
});
}
static auto memberIdsByMembership(immer::map<KeyOfState, Event> stateEvents, const std::string &membership)
{
return intoImmer(
immer::flex_vector<std::string>{},
membershipTransducer(membership)
| zug::map([](auto val) {
auto [stateKey, ev] = val;
return stateKey;
}),
stateEvents);
}
auto memberEventsByMembership(immer::map<KeyOfState, Event> stateEvents, const std::string &membership)
{
return intoImmer(
EventList{},
membershipTransducer(membership)
| zug::map([](auto val) {
auto [stateKey, ev] = val;
return ev;
}),
stateEvents);
}
immer::flex_vector<std::string> RoomModel::joinedMemberIds() const
{
return memberIdsByMembership(stateEvents, "join"s);
}
immer::flex_vector<std::string> RoomModel::invitedMemberIds() const
{
return memberIdsByMembership(stateEvents, "invite"s);
}
immer::flex_vector<std::string> RoomModel::knockedMemberIds() const
{
return memberIdsByMembership(stateEvents, "knock"s);
}
immer::flex_vector<std::string> RoomModel::leftMemberIds() const
{
return memberIdsByMembership(stateEvents, "leave"s);
}
immer::flex_vector<std::string> RoomModel::bannedMemberIds() const
{
return memberIdsByMembership(stateEvents, "ban"s);
}
EventList RoomModel::joinedMemberEvents() const
{
return memberEventsByMembership(stateEvents, "join"s);
}
EventList RoomModel::invitedMemberEvents() const
{
return memberEventsByMembership(stateEvents, "invite"s);
}
EventList RoomModel::knockedMemberEvents() const
{
return memberEventsByMembership(stateEvents, "knock"s);
}
EventList RoomModel::leftMemberEvents() const
{
return memberEventsByMembership(stateEvents, "leave"s);
}
EventList RoomModel::bannedMemberEvents() const
{
return memberEventsByMembership(stateEvents, "ban"s);
}
EventList RoomModel::heroMemberEvents() const
{
return intoImmer(
EventList{},
zug::filter([heroIds=heroIds](auto val) {
auto [k, ev] = val;
auto [type, stateKey] = k;
return type == "m.room.member"s &&
std::find(heroIds.begin(), heroIds.end(), stateKey) != heroIds.end();
})
| zug::map([](auto val) {
auto [_, ev] = val;
return ev;
}),
stateEvents);
}
static Timestamp defaultRotateMs = 604800000;
static int defaultRotateMsgs = 100;
MegOlmSessionRotateDesc RoomModel::sessionRotateDesc() const
{
auto k = KeyOfState{"m.room.encryption", ""};
auto content = stateEvents[k].content().get();
auto ms = content.contains("rotation_period_ms")
? content["rotation_period_ms"].get<Timestamp>()
: defaultRotateMs;
auto msgs = content.contains("rotation_period_msgs")
? content["rotation_period_msgs"].get<int>()
: defaultRotateMsgs;
return MegOlmSessionRotateDesc{ ms, msgs };
}
bool RoomModel::hasUser(std::string userId) const
{
try {
auto ev = stateEvents.at(KeyOfState{"m.room.member", userId});
if (ev.content().get().at("membership") == "join") {
return true;
}
} catch (const std::exception &) {
return false;
}
return false;
}
std::optional<LocalEchoDesc> RoomModel::getLocalEchoByTxnId(std::string txnId) const
{
auto it = std::find_if(localEchoes.begin(), localEchoes.end(), [txnId](const auto &desc) {
return txnId == desc.txnId;
});
if (it != localEchoes.end()) {
return *it;
} else {
return std::nullopt;
}
}
std::optional<PendingRoomKeyEvent> RoomModel::getPendingRoomKeyEventByTxnId(std::string txnId) const
{
auto it = std::find_if(pendingRoomKeyEvents.begin(), pendingRoomKeyEvents.end(), [txnId](const auto &desc) {
return txnId == desc.txnId;
});
if (it != pendingRoomKeyEvents.end()) {
return *it;
} else {
return std::nullopt;
}
}
static double getTagOrder(const json &tag)
{
// https://spec.matrix.org/v1.7/client-server-api/#events-12
// If a room has a tag without an order key then it should appear after the rooms with that tag that have an order key.
return tag.contains("order") && tag["order"].is_number()
? tag["order"].template get<double>()
: ROOM_TAG_DEFAULT_ORDER;
}
immer::map<std::string, double> RoomModel::tags() const
{
auto content = accountData["m.tag"].content().get();
if (!content.contains("tags") || !content["tags"].is_object()) {
return {};
}
auto tagsObject = content["tags"];
auto tagsItems = tagsObject.items();
return std::accumulate(tagsItems.begin(), tagsItems.end(), immer::map<std::string, double>(),
[=](auto acc, const auto &cur) {
auto [id, tag] = cur;
return std::move(acc).set(id, getTagOrder(tag));
}
);
}
static auto normalizeTagEventJson(Event e)
{
auto content = e.content().get();
if (!content.contains("tags") || !content["tags"].is_object()) {
content["tags"] = json::object();
}
return json{
{"content", content},
{"type", "m.tag"},
};
}
Event RoomModel::makeAddTagEvent(std::string tagId, std::optional<double> order) const
{
auto eventJson = normalizeTagEventJson(accountData["m.tag"]);
auto tag = json::object();
if (order.has_value()) {
tag["order"] = order.value();
}
eventJson["content"]["tags"][tagId] = tag;
return Event(eventJson);
}
Event RoomModel::makeRemoveTagEvent(std::string tagId) const
{
auto eventJson = normalizeTagEventJson(accountData["m.tag"]);
eventJson["content"]["tags"].erase(tagId);
return Event(eventJson);
}
void RoomModel::generateRelationships(EventList newEvents)
{
for (const auto &event: newEvents) {
auto [relType, eventId] = event.relationship();
if (!relType.empty()) {
reverseEventRelationships = updateIn(std::move(reverseEventRelationships), [event](auto &&evs) {
return evs.push_back(event.id());
}, eventId, relType);
}
}
}
void RoomModel::regenerateRelationships()
{
generateRelationships(intoImmer(EventList{}, zug::map([](const auto &kv) {
return kv.second;
}), messages));
}
void RoomModel::addToUndecryptedEvents(EventList newEvents)
{
if (!encrypted) {
return;
}
for (auto event : newEvents) {
if (event.encrypted() && !event.decrypted()) {
auto original = event.originalJson();
const auto &o = original.get();
if (o.contains("content")
&& o["content"].contains("session_id")
&& o["content"]["session_id"].is_string()
) {
auto sessionId = o["content"]["session_id"].template get<std::string>();
undecryptedEvents = std::move(undecryptedEvents)
.update(sessionId, [id=event.id()](const auto &v) {
return v.push_back(id);
});
}
}
}
}
void RoomModel::recalculateUndecryptedEvents()
{
if (!encrypted) {
return;
}
addToUndecryptedEvents(intoImmer(EventList{}, zug::map([](const auto &kv) {
return kv.second;
}), messages));
}
}
diff --git a/src/client/room/room-model.hpp b/src/client/room/room-model.hpp
index e3f5af1..5eac465 100644
--- a/src/client/room/room-model.hpp
+++ b/src/client/room/room-model.hpp
@@ -1,495 +1,454 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021-2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#pragma once
#include <libkazv-config.hpp>
#include <string>
#include <variant>
#include <immer/flex_vector.hpp>
#include <immer/map.hpp>
#include <serialization/immer-flex-vector.hpp>
#include <serialization/immer-box.hpp>
#include <serialization/immer-map.hpp>
#include <serialization/immer-array.hpp>
#include <csapi/sync.hpp>
#include <event.hpp>
#include <crypto.hpp>
#include "push-rules-desc.hpp"
#include "local-echo.hpp"
#include "clientutil.hpp"
namespace Kazv
{
struct PendingRoomKeyEvent
{
std::string txnId;
immer::map<std::string, immer::map<std::string, Event>> messages;
+
+ friend bool operator==(const PendingRoomKeyEvent &a, const PendingRoomKeyEvent &b) = default;
+ friend bool operator!=(const PendingRoomKeyEvent &a, const PendingRoomKeyEvent &b) = default;
};
PendingRoomKeyEvent makePendingRoomKeyEventV0(std::string txnId, Event event, immer::map<std::string, immer::flex_vector<std::string>> devices);
struct ReadReceipt
{
std::string eventId;
Timestamp timestamp;
+ friend bool operator==(const ReadReceipt &a, const ReadReceipt &b) = default;
+ friend bool operator!=(const ReadReceipt &a, const ReadReceipt &b) = default;
};
template<class Archive>
void serialize(Archive &ar, ReadReceipt &r, std::uint32_t const /* version */)
{
ar & r.eventId & r.timestamp;
}
- bool operator==(const ReadReceipt &a, const ReadReceipt &b);
- bool operator!=(const ReadReceipt &a, const ReadReceipt &b);
struct EventReader
{
std::string userId;
Timestamp timestamp;
+ friend bool operator==(const EventReader &a, const EventReader &b) = default;
+ friend bool operator!=(const EventReader &a, const EventReader &b) = default;
};
- bool operator==(const EventReader &a, const EventReader &b);
- bool operator!=(const EventReader &a, const EventReader &b);
-
struct AddStateEventsAction
{
immer::flex_vector<Event> stateEvents;
};
/// Go from the back of stateEvents to the beginning,
/// adding the event to room state only if the room
/// has no state event with that state key.
struct MaybeAddStateEventsAction
{
immer::flex_vector<Event> stateEvents;
};
struct AddToTimelineAction
{
/// Events from oldest to latest
immer::flex_vector<Event> events;
std::optional<std::string> prevBatch;
std::optional<bool> limited;
std::optional<std::string> gapEventId;
};
struct AddAccountDataAction
{
immer::flex_vector<Event> events;
};
struct ChangeMembershipAction
{
RoomMembership membership;
};
struct ChangeInviteStateAction
{
immer::flex_vector<Event> events;
};
struct AddEphemeralAction
{
EventList events;
};
struct SetLocalDraftAction
{
std::string localDraft;
};
struct SetRoomEncryptionAction
{
};
struct MarkMembersFullyLoadedAction
{
};
struct SetHeroIdsAction
{
immer::flex_vector<std::string> heroIds;
};
struct AddLocalEchoAction
{
LocalEchoDesc localEcho;
};
struct RemoveLocalEchoAction
{
std::string txnId;
};
struct AddPendingRoomKeyAction
{
PendingRoomKeyEvent pendingRoomKeyEvent;
};
struct RemovePendingRoomKeyAction
{
std::string txnId;
};
struct UpdateJoinedMemberCountAction
{
std::size_t joinedMemberCount;
};
struct UpdateInvitedMemberCountAction
{
std::size_t invitedMemberCount;
};
/// Update local notifications to include the new events
///
/// Precondition: newEvents are already in room.messages
struct AddLocalNotificationsAction
{
EventList newEvents;
PushRulesDesc pushRulesDesc;
std::string myUserId;
};
/// Remove local notifications that are already read
struct RemoveReadLocalNotificationsAction
{
std::string myUserId;
};
/// Update the local read marker, removing any read notifications before it.
struct UpdateLocalReadMarkerAction
{
std::string localReadMarker;
std::string myUserId;
};
- inline bool operator==(const PendingRoomKeyEvent &a, const PendingRoomKeyEvent &b)
- {
- return a.txnId == b.txnId && a.messages == b.messages;
- }
-
- inline bool operator!=(const PendingRoomKeyEvent &a, const PendingRoomKeyEvent &b)
- {
- return !(a == b);
- }
-
inline const double ROOM_TAG_DEFAULT_ORDER = 2;
template<class Archive>
void serialize(Archive &ar, PendingRoomKeyEvent &e, std::uint32_t const version)
{
if (version < 1) {
// loading an older version where there is only one event
std::string txnId;
Event event;
immer::map<std::string, immer::flex_vector<std::string>> devices;
ar & txnId & event & devices;
e = makePendingRoomKeyEventV0(
std::move(txnId), std::move(event), std::move(devices));
} else {
ar & e.txnId & e.messages;
}
}
/**
* Get the sort key for a timeline event.
*
* If the key is larger, the event should be placed
* at the more recent end of the timeline.
*
* @param e The event to get the sort key for.
* @return The sort key. You MUST use `auto` to store the
* result.
*/
auto sortKeyForTimelineEvent(Event e) -> std::tuple<Timestamp, std::string>;
struct RoomModel
{
using Membership = RoomMembership;
using ReverseEventRelationshipMap = immer::map<
std::string /* related event id */,
immer::map<std::string /* relation type */, immer::flex_vector<std::string /* relater event id */>>>;
std::string roomId;
immer::map<KeyOfState, Event> stateEvents;
immer::map<KeyOfState, Event> inviteState;
// Smaller indices mean earlier events
// (oldest) 0 --------> n (latest)
immer::flex_vector<std::string> timeline;
immer::map<std::string, Event> messages;
immer::map<std::string, Event> accountData;
Membership membership{};
std::string paginateBackToken;
/// whether this room has earlier events to be fetched
bool canPaginateBack{true};
immer::map<std::string /* eventId */, std::string /* prevBatch */> timelineGaps;
immer::map<std::string, Event> ephemeral;
std::string localDraft;
bool encrypted{false};
/// a marker to indicate whether we need to rotate
/// the session key earlier than it expires
/// (e.g. when a user in the room's device list changed
/// or when someone joins or leaves)
bool shouldRotateSessionKey{true};
bool membersFullyLoaded{false};
immer::flex_vector<std::string> heroIds;
immer::flex_vector<LocalEchoDesc> localEchoes;
immer::flex_vector<PendingRoomKeyEvent> pendingRoomKeyEvents;
ReverseEventRelationshipMap reverseEventRelationships;
std::size_t joinedMemberCount{0};
std::size_t invitedMemberCount{0};
/// The local read marker for this room. Indicates that
/// you have read up to this event.
std::string localReadMarker;
/// The local unread count for this room.
std::size_t localUnreadCount{0};
/// The local unread notification count for this room.
/// XXX this is never used.
std::size_t localNotificationCount{0};
/// Read receipts for all users
immer::map<std::string /* userId */, ReadReceipt> readReceipts;
/// A map from event id to a list of users that has read
/// receipt at that point
immer::map<
std::string /* eventId */,
immer::flex_vector<std::string /* userId */>> eventReadUsers;
/// A map from the session id to a list of event ids of events
/// that cannot (yet) be decrypted.
immer::map<
std::string /* sessionId */,
immer::flex_vector<std::string /* eventId */>> undecryptedEvents;
immer::flex_vector<std::string> unreadNotificationEventIds;
immer::flex_vector<std::string> joinedMemberIds() const;
immer::flex_vector<std::string> invitedMemberIds() const;
immer::flex_vector<std::string> knockedMemberIds() const;
immer::flex_vector<std::string> leftMemberIds() const;
immer::flex_vector<std::string> bannedMemberIds() const;
EventList joinedMemberEvents() const;
EventList invitedMemberEvents() const;
EventList knockedMemberEvents() const;
EventList leftMemberEvents() const;
EventList bannedMemberEvents() const;
EventList heroMemberEvents() const;
MegOlmSessionRotateDesc sessionRotateDesc() const;
bool hasUser(std::string userId) const;
std::optional<LocalEchoDesc> getLocalEchoByTxnId(std::string txnId) const;
std::optional<PendingRoomKeyEvent> getPendingRoomKeyEventByTxnId(std::string txnId) const;
immer::map<std::string, double> tags() const;
Event makeAddTagEvent(std::string tagId, std::optional<double> order) const;
Event makeRemoveTagEvent(std::string tagId) const;
/**
* Fill in reverseEventRelationships by gathering
* the relationships specified in `newEvents`
*
* @param newEvents The events that just came in after last time event relationships
* are gathered.
*/
void generateRelationships(EventList newEvents);
void regenerateRelationships();
/**
* Fill in undecryptedEvents by gathering
* the session ids specified in `newEvents`.
*
* @param newEvents New incoming events.
*/
void addToUndecryptedEvents(EventList newEvents);
void recalculateUndecryptedEvents();
using Action = std::variant<
AddStateEventsAction,
MaybeAddStateEventsAction,
AddToTimelineAction,
AddAccountDataAction,
ChangeMembershipAction,
ChangeInviteStateAction,
AddEphemeralAction,
SetLocalDraftAction,
SetRoomEncryptionAction,
MarkMembersFullyLoadedAction,
SetHeroIdsAction,
AddLocalEchoAction,
RemoveLocalEchoAction,
AddPendingRoomKeyAction,
RemovePendingRoomKeyAction,
UpdateJoinedMemberCountAction,
UpdateInvitedMemberCountAction,
AddLocalNotificationsAction,
RemoveReadLocalNotificationsAction,
UpdateLocalReadMarkerAction
>;
static RoomModel update(RoomModel r, Action a);
+
+ friend bool operator==(const RoomModel &a, const RoomModel &b) = default;
};
using RoomAction = RoomModel::Action;
- inline bool operator==(const RoomModel &a, const RoomModel &b)
- {
- return a.roomId == b.roomId
- && a.stateEvents == b.stateEvents
- && a.inviteState == b.inviteState
- && a.timeline == b.timeline
- && a.messages == b.messages
- && a.accountData == b.accountData
- && a.membership == b.membership
- && a.paginateBackToken == b.paginateBackToken
- && a.canPaginateBack == b.canPaginateBack
- && a.timelineGaps == b.timelineGaps
- && a.ephemeral == b.ephemeral
- && a.localDraft == b.localDraft
- && a.encrypted == b.encrypted
- && a.shouldRotateSessionKey == b.shouldRotateSessionKey
- && a.membersFullyLoaded == b.membersFullyLoaded
- && a.heroIds == b.heroIds
- && a.localEchoes == b.localEchoes
- && a.pendingRoomKeyEvents == b.pendingRoomKeyEvents
- && a.reverseEventRelationships == b.reverseEventRelationships
- && a.joinedMemberCount == b.joinedMemberCount
- && a.localReadMarker == b.localReadMarker
- && a.localUnreadCount == b.localUnreadCount
- && a.localNotificationCount == b.localNotificationCount
- && a.readReceipts == b.readReceipts
- && a.eventReadUsers == b.eventReadUsers
- && a.undecryptedEvents == b.undecryptedEvents
- && a.unreadNotificationEventIds == b.unreadNotificationEventIds
- ;
- }
-
struct UpdateRoomAction
{
std::string roomId;
RoomAction roomAction;
};
struct RoomListModel
{
immer::map<std::string, RoomModel> rooms;
inline auto at(std::string id) const { return rooms.at(id); }
inline auto operator[](std::string id) const { return rooms[id]; }
inline bool has(std::string id) const { return rooms.find(id); }
using Action = std::variant<
UpdateRoomAction
>;
static RoomListModel update(RoomListModel l, Action a);
+
+ friend bool operator==(const RoomListModel &a, const RoomListModel &b) = default;
};
using RoomListAction = RoomListModel::Action;
- inline bool operator==(const RoomListModel &a, const RoomListModel &b)
- {
- return a.rooms == b.rooms;
- }
-
template<class Archive>
void serialize(Archive &ar, RoomModel &r, std::uint32_t const version)
{
ar
& r.roomId
& r.stateEvents
& r.inviteState
& r.timeline
& r.messages
& r.accountData
& r.membership
& r.paginateBackToken
& r.canPaginateBack
& r.timelineGaps
& r.ephemeral
& r.localDraft
& r.encrypted
& r.shouldRotateSessionKey
& r.membersFullyLoaded
;
if (version >= 1) {
ar
& r.heroIds
;
}
if (version >= 2) {
ar & r.localEchoes;
}
if (version >= 3) {
ar & r.pendingRoomKeyEvents;
}
if (version >= 4) {
ar & r.reverseEventRelationships;
} else { // must be reading from an older version
if constexpr (typename Archive::is_loading()) {
r.regenerateRelationships();
}
}
if (version >= 5) {
ar & r.joinedMemberCount & r.invitedMemberCount;
}
if (version >= 6) {
ar
& r.localReadMarker
& r.localUnreadCount
& r.localNotificationCount
& r.readReceipts
& r.eventReadUsers;
}
if (version >= 7) {
ar & r.undecryptedEvents;
} else {
if constexpr (typename Archive::is_loading()) {
r.recalculateUndecryptedEvents();
}
}
if (version >= 8) {
ar & r.unreadNotificationEventIds;
}
}
template<class Archive>
void serialize(Archive &ar, RoomListModel &l, std::uint32_t const /*version*/)
{
ar & l.rooms;
}
}
BOOST_CLASS_VERSION(Kazv::PendingRoomKeyEvent, 1)
BOOST_CLASS_VERSION(Kazv::ReadReceipt, 0)
BOOST_CLASS_VERSION(Kazv::RoomModel, 8)
BOOST_CLASS_VERSION(Kazv::RoomListModel, 0)
diff --git a/src/crypto/crypto-util.hpp b/src/crypto/crypto-util.hpp
index f830b51..046d12e 100644
--- a/src/crypto/crypto-util.hpp
+++ b/src/crypto/crypto-util.hpp
@@ -1,117 +1,107 @@
/*
* 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 <string>
#include <random>
#include <algorithm>
#include <vector>
#include <nlohmann/json.hpp>
#include <boost/container_hash/hash.hpp>
namespace Kazv
{
using ByteArray = std::vector<unsigned char>;
struct KeyOfGroupSession
{
std::string roomId;
std::string sessionId;
+ friend bool operator==(const KeyOfGroupSession &a, const KeyOfGroupSession &b) = default;
};
/**
* The tag to indicate that a constructor should use user-provided random data.
*/
struct RandomTag {};
using RandomData = std::string;
inline void from_json(const nlohmann::json &j, KeyOfGroupSession &k)
{
k.roomId = j.at("roomId");
k.sessionId = j.at("sessionId");
}
inline void to_json(nlohmann::json &j, const KeyOfGroupSession &k)
{
j = nlohmann::json::object({
{"roomId", k.roomId},
{"sessionId", k.sessionId},
});
}
- inline bool operator==(KeyOfGroupSession a, KeyOfGroupSession b)
- {
- return a.roomId == b.roomId
- && a.sessionId == b.sessionId;
- }
-
struct KeyOfOutboundSession
{
std::string userId;
std::string deviceId;
- };
-
- inline bool operator==(KeyOfOutboundSession a, KeyOfOutboundSession b)
- {
- return a.userId == b.userId
- && a.deviceId == b.deviceId;
+ friend bool operator==(const KeyOfOutboundSession &a, const KeyOfOutboundSession &b) = default;
};
[[nodiscard]] inline ByteArray genRandom(int len)
{
auto rd = std::random_device{};
auto ret = ByteArray(len, '\0');
std::generate(ret.begin(), ret.end(), [&] { return rd(); });
return ret;
}
[[nodiscard]] inline RandomData genRandomData(int len)
{
auto rd = std::random_device{};
auto ret = RandomData(len, '\0');
std::generate(ret.begin(), ret.end(), [&] { return rd(); });
return ret;
}
namespace CryptoConstants
{
inline const std::string ed25519{"ed25519"};
inline const std::string curve25519{"curve25519"};
inline const std::string signedCurve25519{"signed_curve25519"};
inline const std::string olmAlgo{"m.olm.v1.curve25519-aes-sha2"};
inline const std::string megOlmAlgo{"m.megolm.v1.aes-sha2"};
}
}
namespace std
{
template<> struct hash<Kazv::KeyOfGroupSession>
{
std::size_t operator()(const Kazv::KeyOfGroupSession & k) const noexcept {
std::size_t seed = 0;
boost::hash_combine(seed, k.roomId);
boost::hash_combine(seed, k.sessionId);
return seed;
}
};
template<> struct hash<Kazv::KeyOfOutboundSession>
{
std::size_t operator()(const Kazv::KeyOfOutboundSession & k) const noexcept {
std::size_t seed = 0;
boost::hash_combine(seed, k.userId);
boost::hash_combine(seed, k.deviceId);
return seed;
}
};
}
diff --git a/src/crypto/inbound-group-session-p.hpp b/src/crypto/inbound-group-session-p.hpp
index 9c926e9..ccc2510 100644
--- a/src/crypto/inbound-group-session-p.hpp
+++ b/src/crypto/inbound-group-session-p.hpp
@@ -1,72 +1,63 @@
/*
* 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 "inbound-group-session.hpp"
#include <vodozemac.h>
#include <immer/map.hpp>
namespace Kazv
{
struct KeyOfDecryptedEvent
{
std::string eventId;
Timestamp originServerTs;
+ friend bool operator==(const KeyOfDecryptedEvent &a, const KeyOfDecryptedEvent &b) = default;
+ friend bool operator!=(const KeyOfDecryptedEvent &a, const KeyOfDecryptedEvent &b) = default;
};
inline void to_json(nlohmann::json &j, const KeyOfDecryptedEvent &k)
{
j = nlohmann::json::object();
j["eventId"] = k.eventId;
j["originServerTs"] = k.originServerTs;
}
inline void from_json(const nlohmann::json &j, KeyOfDecryptedEvent &k)
{
k.eventId = j.at("eventId");
k.originServerTs = j.at("originServerTs");
}
- inline bool operator==(KeyOfDecryptedEvent a, KeyOfDecryptedEvent b)
- {
- return a.eventId == b.eventId
- && a.originServerTs == b.originServerTs;
- }
-
- inline bool operator!=(KeyOfDecryptedEvent a, KeyOfDecryptedEvent b)
- {
- return !(a == b);
- }
-
struct InboundGroupSessionPrivate
{
InboundGroupSessionPrivate();
InboundGroupSessionPrivate(std::string sessionKey, std::string ed25519Key);
InboundGroupSessionPrivate(const InboundGroupSessionPrivate &that);
~InboundGroupSessionPrivate() = default;
std::optional<rust::Box<vodozemac::megolm::InboundGroupSession>> session;
std::string ed25519Key;
bool valid{false};
bool isImported{false};
immer::map<std::uint32_t /* index */, KeyOfDecryptedEvent> decryptedEvents;
std::size_t checkError(std::size_t code) const;
std::string error() const;
std::string pickle() const;
bool unpickle(std::string pickleData);
bool unpickleFromLibolm(std::string pickleData);
};
}
diff --git a/src/tests/client/discovery-test.cpp b/src/tests/client/discovery-test.cpp
index 299d116..13a7d97 100644
--- a/src/tests/client/discovery-test.cpp
+++ b/src/tests/client/discovery-test.cpp
@@ -1,240 +1,240 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2022-2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <catch2/catch_all.hpp>
#include <boost/asio.hpp>
#include <asio-promise-handler.hpp>
#include <cursorutil.hpp>
#include <sdk-model.hpp>
#include <client/client.hpp>
#include "client-test-util.hpp"
#include "factory.hpp"
using namespace Kazv::Factory;
static const json wellKnownResponseJson = R"({
"m.homeserver": {
"base_url": "https://matrix.example.com"
},
"m.identity_server": {
"base_url": "https://identity.example.com"
},
"org.example.custom.property": {
"app_url": "https://custom.app.example.org"
}
})"_json;
static const json versionsResponseJson = R"({
"unstable_features": {
"org.example.my_feature": true
},
"versions": [
"r0.0.1",
"v1.1"
]
})"_json;
TEST_CASE("Auto-discovery tests", "[client][discovery]")
{
using namespace Kazv::CursorOp;
using Catch::Matchers::StartsWith;
WHEN("We do a discovery")
{
ClientModel m;
auto [resModel, dontCareEffect] = ClientModel::update(m, GetWellknownAction{"@foo:example.com"});
THEN("We should send to the server address as in the userId")
{
REQUIRE(resModel.nextJobs.size() == 1);
auto job = resModel.nextJobs[0];
REQUIRE_THAT(job.url(), StartsWith("https://example.com/"));
REQUIRE(job.dataStr("serverUrl") == "https://example.com");
}
}
WHEN("We do a discovery no remote part provided")
{
ClientModel m;
auto [resModel, dontCareEffect] = ClientModel::update(m, GetWellknownAction{"@foo"});
THEN("We should not send any job")
{
REQUIRE(resModel.nextJobs.size() == 0);
}
}
WHEN("We do a discovery 0-length part provided")
{
ClientModel m;
auto [resModel, dontCareEffect] = ClientModel::update(m, GetWellknownAction{"@foo:"});
THEN("We should not send any job")
{
REQUIRE(resModel.nextJobs.size() == 0);
}
}
// Test different server names
// https://spec.matrix.org/v1.1/appendices/#server-name
WHEN("We do a discovery with server names containing ports")
{
ClientModel m;
auto [resModel, dontCareEffect] = ClientModel::update(m, GetWellknownAction{"@foo:example.com:8080"});
THEN("We should send to the server address as in the userId")
{
REQUIRE(resModel.nextJobs.size() == 1);
auto job = resModel.nextJobs[0];
REQUIRE_THAT(job.url(), StartsWith("https://example.com:8080/"));
}
}
WHEN("We do a discovery with server names containing ipv4 address")
{
ClientModel m;
auto [resModel, dontCareEffect] = ClientModel::update(m, GetWellknownAction{"@foo:1.2.3.4:8080"});
THEN("We should send to the server address as in the userId")
{
REQUIRE(resModel.nextJobs.size() == 1);
auto job = resModel.nextJobs[0];
REQUIRE_THAT(job.url(), StartsWith("https://1.2.3.4:8080/"));
}
}
WHEN("We do a discovery with server names containing ipv6 address")
{
ClientModel m;
auto [resModel, dontCareEffect] = ClientModel::update(m, GetWellknownAction{"@foo:[1234:5678::abcd]:5678"});
THEN("We should send to the server address as in the userId")
{
REQUIRE(resModel.nextJobs.size() == 1);
auto job = resModel.nextJobs[0];
REQUIRE_THAT(job.url(), StartsWith("https://[1234:5678::abcd]:5678/"));
}
}
boost::asio::io_context io;
AsioPromiseHandler ph{io.get_executor()};
auto store = createTestClientStore(ph);
WHEN("We got a successful response")
{
auto resp = makeResponse(
"GetWellknown",
withResponseJsonBody(wellKnownResponseJson)
| withResponseDataKV("serverUrl", "https://example.com")
);
THEN("We should return the server url in the response")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
auto data = stat.dataStr("homeserverUrl");
REQUIRE(data == std::string("https://matrix.example.com"));
});
}
}
WHEN("We got 404")
{
auto resp = makeResponse(
"GetWellknown",
withResponseJsonBody(wellKnownResponseJson)
| withResponseDataKV("serverUrl", "https://example.com")
| withResponseStatusCode(404)
);
THEN("We should return the server in the user id")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
auto data = stat.dataStr("homeserverUrl");
REQUIRE(data == std::string("https://example.com"));
});
}
}
WHEN("We got other error codes")
{
auto resp = makeResponse(
"GetWellknown",
withResponseJsonBody(wellKnownResponseJson)
| withResponseDataKV("serverUrl", "https://example.com")
| withResponseStatusCode(500)
);
THEN("We should FAIL_PROMPT")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(!stat.success());
auto data = stat.dataStr("error");
REQUIRE(data == std::string("FAIL_PROMPT"));
});
}
}
io.run();
}
TEST_CASE("GetVersions", "[client][discovery]")
{
using namespace Kazv::CursorOp;
boost::asio::io_context io;
AsioPromiseHandler ph{io.get_executor()};
auto store = createTestClientStore(ph);
WHEN("We got a successful response")
{
auto resp = makeResponse("GetVersions", withResponseJsonBody(versionsResponseJson));
THEN("We should return the versions supported")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
- auto data = stat.dataJson("versions");
+ auto data = stat.dataJson("versions").template get<immer::flex_vector<std::string>>();
REQUIRE(data == immer::flex_vector<std::string>{
"r0.0.1",
"v1.1"
});
});
}
}
WHEN("We got an error")
{
auto resp = makeResponse("GetVersions", withResponseStatusCode(400));
THEN("We should fail")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(!stat.success());
REQUIRE(stat.dataJson("errorCode") == "400");
});
}
}
io.run();
}
diff --git a/src/tests/client/profile-test.cpp b/src/tests/client/profile-test.cpp
index e3f8e77..7a930e4 100644
--- a/src/tests/client/profile-test.cpp
+++ b/src/tests/client/profile-test.cpp
@@ -1,212 +1,212 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2022-2023 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <catch2/catch_all.hpp>
#include <boost/asio.hpp>
#include <asio-promise-handler.hpp>
#include <cursorutil.hpp>
#include <sdk-model.hpp>
#include <client/client.hpp>
#include "client-test-util.hpp"
#include "factory.hpp"
using namespace Kazv::Factory;
static const json getUserProfileResponseJson = R"({
"avatar_url": "mxc://matrix.org/SDGdghriugerRg",
"displayname": "Alice Margatroid"
})"_json;
using Catch::Matchers::ContainsSubstring;
TEST_CASE("GetUserProfile", "[client][profile]")
{
WHEN("We initiate this job")
{
ClientModel loggedInModel = makeClient({});
auto [resModel, dontCareEffect] = ClientModel::update(loggedInModel, GetUserProfileAction{"@alice:example.com"});
THEN("it should be added")
{
REQUIRE(resModel.nextJobs.size() == 1);
REQUIRE(resModel.nextJobs[0].jobId() == "GetUserProfile");
}
THEN("we should not send access token")
{
REQUIRE_FALSE(hasAccessToken(resModel.nextJobs[0]));
}
}
using namespace Kazv::CursorOp;
boost::asio::io_context io;
AsioPromiseHandler ph{io.get_executor()};
auto store = createTestClientStore(ph);
WHEN("We got a successful response")
{
auto resp = makeResponse("GetUserProfile", withResponseJsonBody(getUserProfileResponseJson));
THEN("We should return the avatar url and display name in the response")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
- REQUIRE(stat.dataStr("avatarUrl") == getUserProfileResponseJson["avatar_url"]);
- REQUIRE(stat.dataStr("displayName") == getUserProfileResponseJson["displayname"]);
+ REQUIRE(stat.dataStr("avatarUrl") == getUserProfileResponseJson["avatar_url"].template get<std::string>());
+ REQUIRE(stat.dataStr("displayName") == getUserProfileResponseJson["displayname"].template get<std::string>());
});
}
}
WHEN("We got an empty json")
{
auto resp = makeResponse("GetUserProfile");
THEN("We should return empty string avatar url and display name")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
REQUIRE(stat.dataStr("avatarUrl") == "");
REQUIRE(stat.dataStr("displayName") == "");
});
}
}
WHEN("We got a failed response")
{
auto resp = makeResponse("GetUserProfile");
resp.statusCode = 404;
THEN("We should fail")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(!stat.success());
REQUIRE(stat.dataStr("errorCode") == "404");
});
}
}
io.run();
}
TEST_CASE("SetAvatarUrl", "[client][profile]")
{
auto jobId = std::string("SetAvatarUrl");
WHEN("We initiate this job")
{
ClientModel loggedInModel = makeClient({});
auto [resModel, dontCareEffect] = ClientModel::update(loggedInModel, SetAvatarUrlAction{"mxc://example.com/xxxyyy"});
THEN("it should be added")
{
REQUIRE(resModel.nextJobs.size() == 1);
auto job = resModel.nextJobs[0];
REQUIRE(job.jobId() == jobId);
REQUIRE(hasAccessToken(job));
REQUIRE_THAT(job.url(), ContainsSubstring("/" + loggedInModel.userId + "/"));
}
}
boost::asio::io_context io;
AsioPromiseHandler ph{io.get_executor()};
auto store = createTestClientStore(ph);
WHEN("We got a successful response")
{
auto resp = makeResponse(jobId);
THEN("We should succeed")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
});
}
}
WHEN("We got a failed response")
{
auto resp = makeResponse(jobId);
resp.statusCode = 400;
THEN("We should fail")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(!stat.success());
REQUIRE(stat.dataStr("errorCode") == "400");
});
}
}
io.run();
}
TEST_CASE("SetDisplayName", "[client][profile]")
{
auto jobId = std::string("SetDisplayName");
WHEN("We initiate this job")
{
ClientModel loggedInModel = makeClient({});
auto [resModel, dontCareEffect] = ClientModel::update(loggedInModel, SetDisplayNameAction{"mew mew"});
THEN("it should be added")
{
REQUIRE(resModel.nextJobs.size() == 1);
auto job = resModel.nextJobs[0];
REQUIRE(job.jobId() == jobId);
REQUIRE(hasAccessToken(job));
REQUIRE_THAT(job.url(), ContainsSubstring("/" + loggedInModel.userId + "/"));
}
}
boost::asio::io_context io;
AsioPromiseHandler ph{io.get_executor()};
auto store = createTestClientStore(ph);
WHEN("We got a successful response")
{
auto resp = makeResponse(jobId);
THEN("We should succeed")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(stat.success());
});
}
}
WHEN("We got a failed response")
{
auto resp = makeResponse(jobId);
resp.statusCode = 400;
THEN("We should fail")
{
store.dispatch(ProcessResponseAction{resp})
.then([](auto stat) {
REQUIRE(!stat.success());
REQUIRE(stat.dataStr("errorCode") == "400");
});
}
}
io.run();
}
diff --git a/src/tests/client/room/pinned-events-test.cpp b/src/tests/client/room/pinned-events-test.cpp
index 9f3cad5..e151f22 100644
--- a/src/tests/client/room/pinned-events-test.cpp
+++ b/src/tests/client/room/pinned-events-test.cpp
@@ -1,140 +1,140 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2024 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <catch2/catch_test_macros.hpp>
#include <lager/event_loop/boost_asio.hpp>
#include <cprjobhandler.hpp>
#include <lagerstoreeventemitter.hpp>
#include <asio-promise-handler.hpp>
#include "client-test-util.hpp"
#include "client/action-mock-utils.hpp"
#include "factory.hpp"
using namespace Kazv;
using namespace Kazv::Factory;
TEST_CASE("Room::pinnedEvents()", "[client][room][getter]")
{
WHEN("content is valid") {
auto room = makeRoom(withRoomState({
makeEvent(withEventType("m.room.pinned_events") | withEventContent(json{{"pinned", {"$1", "$2"}}})),
}));
auto client = makeClient(withRoom(room));
auto cursor = lager::make_constant(SdkModel{client});
auto nameCursor = lager::make_constant(room.roomId);
auto r = Room(cursor, nameCursor, dumbContext());
REQUIRE(r.pinnedEvents().get() == immer::flex_vector<std::string>{"$1", "$2"});
}
WHEN("content is invalid") {
auto room = makeRoom(withRoomState({
makeEvent(withEventType("m.room.pinned_events") | withEventContent(json{{"pinned", {"$1", "$2", 3}}})),
}));
auto client = makeClient(withRoom(room));
auto cursor = lager::make_constant(SdkModel{client});
auto nameCursor = lager::make_constant(room.roomId);
auto r = Room(cursor, nameCursor, dumbContext());
REQUIRE(r.pinnedEvents().get() == immer::flex_vector<std::string>{});
}
WHEN("there is no m.room.pinned_events in state") {
auto room = makeRoom();
auto client = makeClient(withRoom(room));
auto cursor = lager::make_constant(SdkModel{client});
auto nameCursor = lager::make_constant(room.roomId);
auto r = Room(cursor, nameCursor, dumbContext());
REQUIRE(r.pinnedEvents().get() == immer::flex_vector<std::string>{});
}
}
TEST_CASE("Room::pinEvents()", "[client][room][getter]")
{
boost::asio::io_context io;
SingleTypePromiseInterface<EffectStatus> sgph{AsioPromiseHandler{io.get_executor()}};
auto room = makeRoom(withRoomState({
makeEvent(withEventType("m.room.pinned_events") | withEventContent(json{{"pinned", {"$1", "$2"}}})),
}));
ClientModel m = makeClient(withRoom(room));
auto jh = Kazv::CprJobHandler{io.get_executor()};
auto ee = Kazv::LagerStoreEventEmitter(lager::with_boost_asio_event_loop{io.get_executor()});
auto sdk = Kazv::makeSdk(
SdkModel{m},
jh,
ee,
Kazv::AsioPromiseHandler{io.get_executor()},
zug::identity
);
auto ctx = sdk.context();
auto dispatcher = getMockDispatcher(
sgph,
ctx,
returnEmpty<SendStateEventAction>()
);
auto mockContext = getMockContext(sgph, dispatcher);
auto client = Client(Client::InEventLoopTag{}, mockContext, sdk.context());
auto r = client.room(room.roomId);
r.pinEvents({"$2", "$3", "$0"})
.then([&](auto stat) {
REQUIRE(stat.success());
REQUIRE(dispatcher.template calledTimes<SendStateEventAction>() == 1);
auto action = dispatcher.template of<SendStateEventAction>()[0];
REQUIRE(action.roomId == room.roomId);
- REQUIRE(action.event.content().get().at("pinned") == immer::flex_vector<std::string>{"$1", "$2", "$3", "$0"});
+ REQUIRE(action.event.content().get().at("pinned") == json{"$1", "$2", "$3", "$0"});
io.stop();
});
io.run();
}
TEST_CASE("Room::unpinEvents()", "[client][room][getter]")
{
boost::asio::io_context io;
SingleTypePromiseInterface<EffectStatus> sgph{AsioPromiseHandler{io.get_executor()}};
auto room = makeRoom(withRoomState({
makeEvent(withEventType("m.room.pinned_events") | withEventContent(json{{"pinned", {"$1", "$2"}}})),
}));
ClientModel m = makeClient(withRoom(room));
auto jh = Kazv::CprJobHandler{io.get_executor()};
auto ee = Kazv::LagerStoreEventEmitter(lager::with_boost_asio_event_loop{io.get_executor()});
auto sdk = Kazv::makeSdk(
SdkModel{m},
jh,
ee,
Kazv::AsioPromiseHandler{io.get_executor()},
zug::identity
);
auto ctx = sdk.context();
auto dispatcher = getMockDispatcher(
sgph,
ctx,
returnEmpty<SendStateEventAction>()
);
auto mockContext = getMockContext(sgph, dispatcher);
auto client = Client(Client::InEventLoopTag{}, mockContext, sdk.context());
auto r = client.room(room.roomId);
r.unpinEvents({"$2", "$3", "$0"})
.then([&](auto stat) {
REQUIRE(stat.success());
REQUIRE(dispatcher.template calledTimes<SendStateEventAction>() == 1);
auto action = dispatcher.template of<SendStateEventAction>()[0];
REQUIRE(action.roomId == room.roomId);
- REQUIRE(action.event.content().get().at("pinned") == immer::flex_vector<std::string>{"$1"});
+ REQUIRE(action.event.content().get().at("pinned") == json{"$1"});
io.stop();
});
io.run();
}
diff --git a/src/tests/crypto/inbound-group-session-test.cpp b/src/tests/crypto/inbound-group-session-test.cpp
index eb9a6c6..f88a0fa 100644
--- a/src/tests/crypto/inbound-group-session-test.cpp
+++ b/src/tests/crypto/inbound-group-session-test.cpp
@@ -1,108 +1,108 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2024 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <catch2/catch_test_macros.hpp>
#include <inbound-group-session.hpp>
#include <outbound-group-session.hpp>
#include <crypto.hpp>
#include "crypto-test-resource.hpp"
using namespace Kazv;
static const auto resource = cryptoDumpResource();
TEST_CASE("InboundGroupSession conversion from libolm to vodozemac")
{
auto sessionJson = resource["a"]["inboundGroupSessions"][0][1];
auto session = sessionJson.template get<InboundGroupSession>();
REQUIRE(session.valid());
REQUIRE(!session.isImported());
- REQUIRE(session.ed25519Key() == sessionJson["ed25519Key"]);
+ REQUIRE(session.ed25519Key() == sessionJson["ed25519Key"].template get<std::string>());
auto encrypted = resource["megolmEncrypted"];
auto plainText = resource["megolmPlainText"];
auto a = Crypto();
a.loadJson(resource["a"]);
auto decrypted = a.decrypt(encrypted);
REQUIRE(decrypted.has_value());
auto decryptedJson = json::parse(decrypted.value());
REQUIRE(decryptedJson == plainText);
}
TEST_CASE("InboundGroupSession::from_json error handling")
{
auto sessionJson = resource["a"]["inboundGroupSessions"][0][1];
sessionJson["session"] = "AAAAAAAAAA";
auto session = sessionJson.template get<InboundGroupSession>();
REQUIRE(!session.valid());
}
TEST_CASE("InboundGroupSession::decrypt error handling")
{
auto sessionJson = resource["a"]["inboundGroupSessions"][0][1];
auto session = sessionJson.template get<InboundGroupSession>();
WHEN("message not decryptable") {
auto res = session.decrypt("AAAAAA", "$1", 1234);
REQUIRE(!res);
}
WHEN("message is not valid base64") {
auto res = session.decrypt("喵喵喵", "$1", 1234);
REQUIRE(!res);
}
WHEN("message is before the index") {
auto ogs = OutboundGroupSession(RandomTag{}, genRandomData(OutboundGroupSession::constructRandomSize()), 0);
auto encrypted1 = ogs.encrypt("text");
auto igs = InboundGroupSession(ogs.sessionKey(), "placeholder");
auto res = igs.decrypt(encrypted1, "$1", 1234);
REQUIRE(!res.has_value());
}
}
TEST_CASE("InboundGroupSession constructor error handling")
{
WHEN("key not valid") {
auto session = InboundGroupSession("AAAAAA", "ed25519Key");
REQUIRE(!session.valid());
}
WHEN("key is not valid base64") {
auto session = InboundGroupSession("喵喵喵", "ed25519Key");
REQUIRE(!session.valid());
}
}
TEST_CASE("invalid InboundGroupSession is copyable")
{
InboundGroupSession session("AAAAAA", "ed25519Key");
REQUIRE(!session.valid());
auto session2 = session;
REQUIRE(!session.valid());
}
TEST_CASE("export and import InboundGroupSession", "[crypto]")
{
auto ogs = OutboundGroupSession(RandomTag{}, genRandomData(OutboundGroupSession::constructRandomSize()), 0);
auto igs = InboundGroupSession(ogs.sessionKey(), "placeholder");
auto exported = igs.toExportFormat();
auto imported = InboundGroupSession(exported, "placeholder");
REQUIRE(imported.valid());
REQUIRE(imported.isImported());
WHEN("try to decrypt") {
auto encrypted1 = ogs.encrypt("text");
auto res = imported.decrypt(encrypted1, "$1", 1234);
REQUIRE(res.has_value());
REQUIRE(res.value() == "text");
}
WHEN("serialization") {
auto j = json(imported);
auto deserialized = j.template get<InboundGroupSession>();
REQUIRE(deserialized.valid());
REQUIRE(deserialized.isImported());
}
}
diff --git a/src/tests/crypto/outbound-group-session-test.cpp b/src/tests/crypto/outbound-group-session-test.cpp
index 9421e3d..6b074c0 100644
--- a/src/tests/crypto/outbound-group-session-test.cpp
+++ b/src/tests/crypto/outbound-group-session-test.cpp
@@ -1,67 +1,67 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2024 tusooa <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <libkazv-config.hpp>
#include <catch2/catch_test_macros.hpp>
#include <outbound-group-session.hpp>
#include <crypto.hpp>
#include "crypto-test-resource.hpp"
using namespace Kazv;
static const auto resource = cryptoDumpResource();
TEST_CASE("OutboundGroupSession conversion from libolm to vodozemac")
{
auto sessionJson = resource["a"]["outboundGroupSessions"]["!foo:example.com"];
auto session = sessionJson.template get<OutboundGroupSession>();
REQUIRE(session.valid());
- REQUIRE(session.initialSessionKey() == sessionJson["initialSessionKey"]);
+ REQUIRE(session.initialSessionKey() == sessionJson["initialSessionKey"].template get<std::string>());
}
TEST_CASE("OutboundGroupSession serialization roundtrip")
{
auto session = OutboundGroupSession(RandomTag{}, genRandomData(OutboundGroupSession::constructRandomSize()), 0);
json j = session;
auto session2 = j.template get<OutboundGroupSession>();
REQUIRE(session2.valid());
REQUIRE(session.initialSessionKey() == session2.initialSessionKey());
REQUIRE(session.sessionId() == session2.sessionId());
}
TEST_CASE("OutboundGroupSession::from_json error handling")
{
auto sessionJson = resource["a"]["outboundGroupSessions"]["!foo:example.com"];
sessionJson["session"] = "AAAAAAAAAA";
auto session = sessionJson.template get<OutboundGroupSession>();
REQUIRE(!session.valid());
}
TEST_CASE("OutboundGroupSession ctor")
{
auto session = OutboundGroupSession();
REQUIRE(!session.valid());
}
TEST_CASE("OutboundGroupSession::encrypt error handling")
{
// invalid utf8 will cause rust::Str to throw
// https://stackoverflow.com/questions/1301402/example-invalid-utf8-string
std::string plainText = "\xc3\x28";
auto session = OutboundGroupSession(RandomTag{}, genRandomData(OutboundGroupSession::constructRandomSize()), 0);
auto originalIndex = session.messageIndex();
auto res = session.encrypt(plainText);
REQUIRE(res.empty());
REQUIRE(originalIndex == session.messageIndex());
}
TEST_CASE("invalid OutboundGroupSession is copyable")
{
OutboundGroupSession session;
REQUIRE(!session.valid());
auto session2 = session;
REQUIRE(!session.valid());
}
diff --git a/src/tests/store-test.cpp b/src/tests/store-test.cpp
index 6341887..4317b6b 100644
--- a/src/tests/store-test.cpp
+++ b/src/tests/store-test.cpp
@@ -1,156 +1,159 @@
/*
* This file is part of libkazv.
* SPDX-FileCopyrightText: 2021 Tusooa Zhu <tusooa@kazv.moe>
* SPDX-License-Identifier: AGPL-3.0-or-later
*/
#include <catch2/catch_all.hpp>
#include <asio-promise-handler.hpp>
#include <store.hpp>
#include <context.hpp>
using namespace Kazv;
struct BackInserter
{
+ BackInserter(std::function<void(int)> f)
+ : insertFunc(f)
+ {}
BackInserter(const BackInserter &) = delete;
BackInserter(BackInserter &&) = default;
void operator()(int a) const { insertFunc(a); }
std::function<void(int)> insertFunc;
};
TEST_CASE("Store should behave properly", "[store]")
{
boost::asio::io_context ioContext;
auto ph = AsioPromiseHandler(ioContext.get_executor());
std::vector<int> results;
- BackInserter bi = { [&results](int a) { results.push_back(a); } };
+ BackInserter bi([&results](int a) { results.push_back(a); });
using Model = int;
using Action = int;
using Deps = lager::deps<BackInserter &>;
using Result = std::pair<Model, Effect<Action, Deps>>;
using Reducer = std::function<Result(Model, Action)>;
Reducer update = [&results](Model m, Action a) -> Result {
if (a > 0) {
return { m + a, lager::noop };
} else {
auto newM = m + a;
return { newM,
[&results, newM](auto &&ctx) {
auto &bi = lager::get<BackInserter &>(ctx);
bi(newM);
return ctx.dispatch(1);
}
};
}
};
auto store = makeStore<Action>(Model{}, update, ph, lager::with_deps(std::ref(bi)));
lager::reader<Model> reader = store;
Context<Action> ctx = store.context();
Context<Action> ctx2 = ctx;
store.dispatch(1)
.then([=](auto res) {
REQUIRE(res);
REQUIRE(reader.get() == 1);
})
.then([=](auto) {
return ctx.dispatch(2);
})
.then([=](auto) {
REQUIRE(reader.get() == 3); // 1 + 2
})
.then([=](auto) {
return ctx2.dispatch(-5);
})
.then([=](auto) {
REQUIRE(reader.get() == -1); // 3 - 5 + 1
})
.then([=, &ioContext](auto) {
return ctx.createWaitingPromise(
[=, &ioContext](auto resolve) {
auto timer = std::make_shared<boost::asio::steady_timer>(ioContext);
timer->expires_after(std::chrono::milliseconds(300));
timer->async_wait(
[ctx, timer, resolve](const boost::system::error_code& error) {
if (! error) {
resolve(ctx.createResolvedPromise(true)
.then([=](auto) { return ctx.dispatch(10); }));
}
});
});
})
.then([=](auto) {
REQUIRE(reader.get() == 9); // -1 + 10
});
ioContext.run();
REQUIRE(results[0] == -2); // 3 - 5
}
TEST_CASE("Store and Context can be moved", "[store]")
{
boost::asio::io_context ioContext;
auto ph = AsioPromiseHandler(ioContext.get_executor());
using Model = int;
using Action = int;
using Result = std::pair<Model, Effect<Action>>;
using Reducer = std::function<Result(Model, Action)>;
Reducer update = [](Model m, Action a) -> Result {
return { m + a, lager::noop };
};
auto store = makeStore<Action>(Model{}, update, ph);
Context<Action> ctx = store.context();
lager::reader<Model> reader = store.reader();
SECTION("Store can be moved without invalidating Contexts") {
auto p1 = ctx.dispatch(1)
.then([=](auto) {
REQUIRE(reader.get() == 1);
});
auto store2 = std::move(store);
p1
.then([=](auto) {
return ctx.dispatch(1);
})
.then([=](auto) {
REQUIRE(reader.get() == 2);
});
ioContext.run();
}
SECTION("Context can be moved without affecting dispatched actions") {
auto p1 = ctx.dispatch(1)
.then([=](auto) {
REQUIRE(reader.get() == 1);
});
auto ctx2 = std::move(ctx);
p1
.then([=](auto) {
return ctx2.dispatch(1);
})
.then([=](auto) {
REQUIRE(reader.get() == 2);
});
ioContext.run();
}
}

File Metadata

Mime Type
text/x-diff
Expires
Fri, Oct 9, 10:17 AM (1 d, 20 h)
Storage Engine
blob
Storage Format
Raw Data
Storage Handle
1784700
Default Alt Text
(153 KB)

Event Timeline