Page Menu
Home
Phorge
Search
Configure Global Search
Log In
Files
F85803688
No One
Temporary
Actions
View File
Edit File
Delete File
View Transforms
Subscribe
Award Token
Flag For Later
Size
153 KB
Referenced Files
None
Subscribers
None
View Options
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
Details
Attached
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)
Attached To
Mode
rL libkazv
Attached
Detach File
Event Timeline
Log In to Comment