From a6398738bb0ddee3e3764e6235f5f7bc8fd04897 Mon Sep 17 00:00:00 2001 From: emeric Date: Fri, 30 Apr 2021 17:41:03 +0200 Subject: [PATCH] listenbrainz: added a listens synchronizer. fixes #142 --- conf/lms.conf | 8 +- src/libs/database/impl/Track.cpp | 17 +- src/libs/database/impl/TrackList.cpp | 12 + src/libs/database/impl/User.cpp | 9 + src/libs/database/include/database/Track.hpp | 30 +- .../database/include/database/TrackList.hpp | 3 + src/libs/database/include/database/User.hpp | 1 + src/libs/scanner/impl/AcousticBrainzUtils.cpp | 4 +- src/libs/scrobbling/CMakeLists.txt | 8 +- src/libs/scrobbling/impl/IScrobbler.hpp | 2 +- src/libs/scrobbling/impl/Scrobbling.cpp | 18 +- src/libs/scrobbling/impl/Scrobbling.hpp | 4 +- .../impl/internal/InternalScrobbler.cpp | 6 +- .../impl/internal/InternalScrobbler.hpp | 2 +- .../listenbrainz/ListenBrainzScrobbler.cpp | 282 ++-------- .../listenbrainz/ListenBrainzScrobbler.hpp | 64 +-- .../impl/listenbrainz/ListensSynchronizer.cpp | 514 ++++++++++++++++++ .../impl/listenbrainz/ListensSynchronizer.hpp | 98 ++++ .../impl/listenbrainz/SendQueue.cpp | 228 ++++++++ .../impl/listenbrainz/SendQueue.hpp | 118 ++++ .../scrobbling/impl/listenbrainz/Utils.cpp | 63 +++ .../scrobbling/impl/listenbrainz/Utils.hpp | 38 ++ .../include/scrobbling/IScrobbling.hpp | 6 +- .../scrobbling/include/scrobbling/Listen.hpp | 7 + src/libs/subsonic/impl/SubsonicResource.cpp | 2 +- src/libs/utils/CMakeLists.txt | 1 + src/libs/utils/impl/ChildProcessManager.cpp | 36 +- src/libs/utils/impl/ChildProcessManager.hpp | 12 +- src/libs/utils/impl/IOContextRunner.cpp | 49 ++ .../include/utils/IChildProcessManager.hpp | 7 +- .../utils/include/utils/IOContextRunner.hpp | 43 ++ src/lms/main.cpp | 11 +- src/test/database/DatabaseTest.cpp | 19 + 33 files changed, 1359 insertions(+), 363 deletions(-) create mode 100644 src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.cpp create mode 100644 src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.hpp create mode 100644 src/libs/scrobbling/impl/listenbrainz/SendQueue.cpp create mode 100644 src/libs/scrobbling/impl/listenbrainz/SendQueue.hpp create mode 100644 src/libs/scrobbling/impl/listenbrainz/Utils.cpp create mode 100644 src/libs/scrobbling/impl/listenbrainz/Utils.hpp create mode 100644 src/libs/utils/impl/IOContextRunner.cpp create mode 100644 src/libs/utils/include/utils/IOContextRunner.hpp diff --git a/conf/lms.conf b/conf/lms.conf index 5ec9768d..f62d8d83 100644 --- a/conf/lms.conf +++ b/conf/lms.conf @@ -35,10 +35,14 @@ deploy-path = "/"; http-server-thread-count = 0; # ListenBrainz root API -listenbrainz-api-url = "https://api.listenbrainz.org/1/"; +listenbrainz-api-base-url = "https://api.listenbrainz.org"; +# How many listens to retrieve when syncing (0 disables sync) +listenbrainz-max-sync-listen-count = 1000; +# How often to resync listens (0 disables sync) +listenbrainz-sync-listens-period-hours = 1; # Acousticbrainz root API -acousticbrainz-api-url = "https://acousticbrainz.org/api/v1/"; +acousticbrainz-api-base-url = "https://acousticbrainz.org/api"; # Authentication # Available backends: "internal", "PAM", "http-headers" diff --git a/src/libs/database/impl/Track.cpp b/src/libs/database/impl/Track.cpp index 604e7062..860e4ccb 100644 --- a/src/libs/database/impl/Track.cpp +++ b/src/libs/database/impl/Track.cpp @@ -30,6 +30,7 @@ #include "utils/Logger.hpp" #include "SqlQuery.hpp" +#include "StringViewTraits.hpp" namespace Database { @@ -151,12 +152,12 @@ Track::getById(Session& session, IdType id) } Track::pointer -Track::getByMBID(Session& session, const UUID& mbid) +Track::getByRecordingMBID(Session& session, const UUID& mbid) { session.checkSharedLocked(); return session.getDboSession().find() - .where("mbid = ?").bind(std::string {mbid.getAsString()}); + .where("recording_mbid = ?").bind(std::string {mbid.getAsString()}); } Track::pointer @@ -354,6 +355,18 @@ Track::getByFilter(Session& session, return res; } +std::vector +Track::getByNameAndReleaseName(Session& session, std::string_view trackName, std::string_view releaseName) +{ + session.checkSharedLocked(); + Wt::Dbo::collection collection = session.getDboSession().query("SELECT t from track t") + .join("release r ON t.release_id = r.id") + .where("t.name = ?").bind(trackName) + .where("r.name = ?").bind(releaseName); + + return std::vector(collection.begin(), collection.end()); +} + std::vector Track::getSimilarTracks(Session& session, const std::unordered_set& tracks, diff --git a/src/libs/database/impl/TrackList.cpp b/src/libs/database/impl/TrackList.cpp index 54b685e3..73c1067a 100644 --- a/src/libs/database/impl/TrackList.cpp +++ b/src/libs/database/impl/TrackList.cpp @@ -148,6 +148,18 @@ TrackList::getEntries(std::optional offset, std::optional>(entries.begin(), entries.end()); } +Wt::Dbo::ptr +TrackList::getEntryByTrackAndDateTime(Wt::Dbo::ptr track, const Wt::WDateTime& dateTime) const +{ + assert(session()); + assert(IdIsValid(self()->id())); + + return session()->find() + .where("tracklist_id = ?").bind(self().id()) + .where("track_id = ?").bind(track.id()) + .where("date_time = ?").bind(Wt::WDateTime::fromTime_t(dateTime.toTime_t())); +} + static Wt::Dbo::Query createArtistsQuery(Wt::Dbo::Session& session, const std::string& queryStr, IdType tracklistId, const std::set& clusterIds, std::optional linkType) diff --git a/src/libs/database/impl/User.cpp b/src/libs/database/impl/User.cpp index 0e3b09bd..b0dc3f01 100644 --- a/src/libs/database/impl/User.cpp +++ b/src/libs/database/impl/User.cpp @@ -84,6 +84,15 @@ User::getAll(Session& session) return std::vector(res.begin(), res.end()); } +std::vector +User::getAllIds(Session& session) +{ + session.checkSharedLocked(); + + Wt::Dbo::collection res = session.getDboSession().query("SELECT id FROM user"); + return std::vector(res.begin(), res.end()); +} + User::pointer User::getDemo(Session& session) { diff --git a/src/libs/database/include/database/Track.hpp b/src/libs/database/include/database/Track.hpp index 119b7d75..fc8055ae 100644 --- a/src/libs/database/include/database/Track.hpp +++ b/src/libs/database/include/database/Track.hpp @@ -23,6 +23,7 @@ #include #include #include +#include #include #include @@ -61,7 +62,7 @@ class Track : public Wt::Dbo::Dbo static std::size_t getCount(Session& session); static pointer getByPath(Session& session, const std::filesystem::path& p); static pointer getById(Session& session, IdType id); - static pointer getByMBID(Session& session, const UUID& MBID); + static pointer getByRecordingMBID(Session& session, const UUID& MBID); static std::vector getSimilarTracks(Session& session, const std::unordered_set& trackIds, std::optional offset = {}, @@ -73,6 +74,7 @@ class Track : public Wt::Dbo::Dbo const std::vector& keywords, // if non empty, name must match all of these keywords std::optional range, bool& moreExpected); + static std::vector getByNameAndReleaseName(Session& session, std::string_view trackName, std::string_view releaseName); static std::vector getAll(Session& session, std::optional limit = std::nullopt); static std::vector getAllRandom(Session& session, const std::set& clusters, std::optional limit = std::nullopt); @@ -188,28 +190,28 @@ class Track : public Wt::Dbo::Dbo static const std::size_t _maxCopyrightLength = 128; static const std::size_t _maxCopyrightURLLength = 128; - int _scanVersion {}; - int _trackNumber {}; - int _discNumber {}; - std::string _discSubtitle; - int _totalTrack {}; - int _totalDisc {}; + int _scanVersion {}; + int _trackNumber {}; + int _discNumber {}; + std::string _discSubtitle; + int _totalTrack {}; + int _totalDisc {}; std::string _name; std::string _artistName; std::string _releaseName; - std::chrono::duration _duration; - int _year {}; - int _originalYear {}; + std::chrono::duration _duration {}; + int _year {}; + int _originalYear {}; std::string _filePath; - Wt::WDateTime _fileLastWrite; - Wt::WDateTime _fileAdded; + Wt::WDateTime _fileLastWrite; + Wt::WDateTime _fileAdded; bool _hasCover {}; std::string _trackMBID; std::string _recordingMBID; std::string _copyright; std::string _copyrightURL; - std::optional _trackReplayGain; - std::optional _releaseReplayGain; + std::optional _trackReplayGain; + std::optional _releaseReplayGain; Wt::Dbo::ptr _release; Wt::Dbo::collection> _trackArtistLinks; diff --git a/src/libs/database/include/database/TrackList.hpp b/src/libs/database/include/database/TrackList.hpp index 7e7099da..d8505d23 100644 --- a/src/libs/database/include/database/TrackList.hpp +++ b/src/libs/database/include/database/TrackList.hpp @@ -84,6 +84,9 @@ class TrackList : public Wt::Dbo::Dbo std::size_t getCount() const; Wt::Dbo::ptr getEntry(std::size_t pos) const; std::vector> getEntries(std::optional offset = {}, std::optional size = {}) const; + Wt::Dbo::ptr getEntryByTrackAndDateTime(Wt::Dbo::ptr track, const Wt::WDateTime& dateTime) const; + + // Get track bya std::vector> getArtistsReverse(const std::set& clusterIds, std::optional linkType, std::optional range, bool& moreResults) const; std::vector> getReleasesReverse(const std::set& clusterIds, std::optional range, bool& moreResults) const; diff --git a/src/libs/database/include/database/User.hpp b/src/libs/database/include/database/User.hpp index 0f5b5758..82980188 100644 --- a/src/libs/database/include/database/User.hpp +++ b/src/libs/database/include/database/User.hpp @@ -139,6 +139,7 @@ class User : public Wt::Dbo::Dbo static pointer getById(Session& session, IdType id); static pointer getByLoginName(Session& session, std::string_view loginName); static std::vector getAll(Session& session); + static std::vector getAllIds(Session& session); static pointer getDemo(Session& session); static std::size_t getCount(Session& session); diff --git a/src/libs/scanner/impl/AcousticBrainzUtils.cpp b/src/libs/scanner/impl/AcousticBrainzUtils.cpp index 37f06166..64980664 100644 --- a/src/libs/scanner/impl/AcousticBrainzUtils.cpp +++ b/src/libs/scanner/impl/AcousticBrainzUtils.cpp @@ -38,9 +38,9 @@ static std::string getJsonData(const UUID& mbid) { - static constexpr std::string_view defaultAPIURL {"https://acousticbrainz.org/api/v1/"}; + static constexpr std::string_view defaultAPIURL {"https://acousticbrainz.org/api"}; - const std::string url {std::string {Service::get()->getString("acousticbrainz-api-url", defaultAPIURL)} + std::string {mbid.getAsString()} + "/low-level"}; + const std::string url {std::string {Service::get()->getString("acousticbrainz-api-base-url", defaultAPIURL)} + std::string {mbid.getAsString()} + "/low-level"}; boost::asio::io_service ioService; diff --git a/src/libs/scrobbling/CMakeLists.txt b/src/libs/scrobbling/CMakeLists.txt index 162f1c9c..bd8a0960 100644 --- a/src/libs/scrobbling/CMakeLists.txt +++ b/src/libs/scrobbling/CMakeLists.txt @@ -2,6 +2,9 @@ add_library(lmsscrobbling SHARED impl/internal/InternalScrobbler.cpp impl/listenbrainz/ListenBrainzScrobbler.cpp + impl/listenbrainz/ListensSynchronizer.cpp + impl/listenbrainz/SendQueue.cpp + impl/listenbrainz/Utils.cpp impl/Scrobbling.cpp ) @@ -14,9 +17,12 @@ target_include_directories(lmsscrobbling PRIVATE impl ) +target_link_libraries(lmsscrobbling PRIVATE + lmsutils + ) + target_link_libraries(lmsscrobbling PUBLIC lmsdatabase - lmsutils ) install(TARGETS lmsscrobbling DESTINATION lib) diff --git a/src/libs/scrobbling/impl/IScrobbler.hpp b/src/libs/scrobbling/impl/IScrobbler.hpp index 840d24ef..01fae716 100644 --- a/src/libs/scrobbling/impl/IScrobbler.hpp +++ b/src/libs/scrobbling/impl/IScrobbler.hpp @@ -46,7 +46,7 @@ namespace Scrobbling virtual void listenStarted(const Listen& listen) = 0; virtual void listenFinished(const Listen& listen, std::optional duration) = 0; - virtual void addListen(const Listen& listen, const Wt::WDateTime& timePoint) = 0; + virtual void addTimedListen(const TimedListen& listen) = 0; virtual Wt::Dbo::ptr getListensTrackList(Database::Session& session, Wt::Dbo::ptr user) = 0; }; diff --git a/src/libs/scrobbling/impl/Scrobbling.cpp b/src/libs/scrobbling/impl/Scrobbling.cpp index 14f84c7e..65df227d 100644 --- a/src/libs/scrobbling/impl/Scrobbling.cpp +++ b/src/libs/scrobbling/impl/Scrobbling.cpp @@ -30,37 +30,37 @@ namespace Scrobbling { std::unique_ptr - createScrobbling(Database::Db& db) + createScrobbling(boost::asio::io_context& ioContext, Database::Db& db) { - return std::make_unique(db); + return std::make_unique(ioContext, db); } - Scrobbling::Scrobbling(Database::Db& db) + Scrobbling::Scrobbling(boost::asio::io_context& ioContext, Database::Db& db) : _db {db} { _scrobblers.emplace(Database::Scrobbler::Internal, std::make_unique(_db)); - _scrobblers.emplace(Database::Scrobbler::ListenBrainz, std::make_unique(_db)); + _scrobblers.emplace(Database::Scrobbler::ListenBrainz, std::make_unique(ioContext, _db)); } void Scrobbling::listenStarted(const Listen& listen) { - if (auto scrobbler {getUserScrobbler(listen.userId)}) + if (std::optional scrobbler {getUserScrobbler(listen.userId)}) _scrobblers[*scrobbler]->listenStarted(listen); } void Scrobbling::listenFinished(const Listen& listen, std::optional duration) { - if (auto scrobbler {getUserScrobbler(listen.userId)}) + if (std::optional scrobbler {getUserScrobbler(listen.userId)}) _scrobblers[*scrobbler]->listenFinished(listen, duration); } void - Scrobbling::addListen(const Listen& listen, Wt::WDateTime timePoint) + Scrobbling::addTimedListen(const TimedListen& listen) { - if (auto scrobbler {getUserScrobbler(listen.userId)}) - _scrobblers[*scrobbler]->addListen(listen, timePoint); + if (std::optional scrobbler {getUserScrobbler(listen.userId)}) + _scrobblers[*scrobbler]->addTimedListen(listen); } std::optional diff --git a/src/libs/scrobbling/impl/Scrobbling.hpp b/src/libs/scrobbling/impl/Scrobbling.hpp index 26f3cbb1..613335c8 100644 --- a/src/libs/scrobbling/impl/Scrobbling.hpp +++ b/src/libs/scrobbling/impl/Scrobbling.hpp @@ -31,12 +31,12 @@ namespace Scrobbling class Scrobbling : public IScrobbling { public: - Scrobbling(Database::Db& db); + Scrobbling(boost::asio::io_context& ioContext, Database::Db& db); private: void listenStarted(const Listen& listen) override; void listenFinished(const Listen& listen, std::optional duration) override; - void addListen(const Listen& listen, Wt::WDateTime timePoint) override; + void addTimedListen(const TimedListen& listen) override; std::vector> getRecentArtists(Database::Session& session, Wt::Dbo::ptr user, diff --git a/src/libs/scrobbling/impl/internal/InternalScrobbler.cpp b/src/libs/scrobbling/impl/internal/InternalScrobbler.cpp index 14a86adf..019759f2 100644 --- a/src/libs/scrobbling/impl/internal/InternalScrobbler.cpp +++ b/src/libs/scrobbling/impl/internal/InternalScrobbler.cpp @@ -47,11 +47,11 @@ namespace Scrobbling if (duration && *duration < std::chrono::seconds {5}) return; - addListen(listen, Wt::WDateTime::currentDateTime()); + addTimedListen({listen, Wt::WDateTime::currentDateTime()}); } void - InternalScrobbler::addListen(const Listen& listen, const Wt::WDateTime& timePoint) + InternalScrobbler::addTimedListen(const TimedListen& listen) { Database::Session& session {_db.getTLSSession()}; @@ -69,7 +69,7 @@ namespace Scrobbling if (!track) return; - Database::TrackListEntry::create(session, track, getListensTrackList(session, user), timePoint); + Database::TrackListEntry::create(session, track, getListensTrackList(session, user), listen.listenedAt); } Wt::Dbo::ptr diff --git a/src/libs/scrobbling/impl/internal/InternalScrobbler.hpp b/src/libs/scrobbling/impl/internal/InternalScrobbler.hpp index e11ff43e..4d3323ea 100644 --- a/src/libs/scrobbling/impl/internal/InternalScrobbler.hpp +++ b/src/libs/scrobbling/impl/internal/InternalScrobbler.hpp @@ -32,7 +32,7 @@ namespace Scrobbling void listenStarted(const Listen& listen) override; void listenFinished(const Listen& listen, std::optional duration) override; - void addListen(const Listen& listen, const Wt::WDateTime& timePoint) override; + void addTimedListen(const TimedListen& listen) override; Wt::Dbo::ptr getListensTrackList(Database::Session& session, Wt::Dbo::ptr user) override; diff --git a/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.cpp b/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.cpp index 39a05c4a..0cdb8443 100644 --- a/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.cpp +++ b/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.cpp @@ -34,41 +34,12 @@ #include "utils/IConfig.hpp" #include "utils/Logger.hpp" #include "utils/Service.hpp" +#include "Utils.hpp" #define LOG(sev) LMS_LOG(SCROBBLING, sev) << "[listenbrainz] - " -namespace StringUtils -{ - template<> - std::optional - readAs(const std::string& str) - { - std::optional res; - - if (const std::optional value {StringUtils::readAs(str)}) - res = std::chrono::seconds {*value}; - - return res; - } -} - namespace { - std::optional - getListenBrainzToken(Database::Session& session, Database::IdType userId) - { - auto transaction {session.createSharedTransaction()}; - - const Database::User::pointer user {Database::User::getById(session, userId)}; - if (!user) - return std::nullopt; - - if (user->getScrobbler() != Database::Scrobbler::ListenBrainz) - return std::nullopt; - - return user->getListenBrainzToken(); - } - bool canBeScrobbled(Database::Session& session, Database::IdType trackId, std::chrono::seconds duration) { @@ -164,252 +135,105 @@ namespace res = Wt::Json::serialize(root); return res; } - - template - std::optional - headerReadAs(const Wt::Http::Message& msg, std::string_view headerName) - { - std::optional res; - - if (const std::string* headerValue {msg.getHeader(std::string {headerName})}) - res = StringUtils::readAs(*headerValue); - - return res; - } } -namespace Scrobbling +namespace Scrobbling::ListenBrainz { - static const std::string historyTracklistName {"__scrobbler_listenbrainz_history__"}; - - ListenBrainzScrobbler::ListenBrainzScrobbler(Database::Db& db) - : _apiEndpoint {Service::get()->getString("listenbrainz-api-url", "https://api.listenbrainz.org/1/")} + Scrobbler::Scrobbler(boost::asio::io_context& ioContext, Database::Db& db) + : _ioContext {ioContext} , _db {db} + , _sendQueue {_ioContext, Service::get()->getString("listenbrainz-api-base-url", "https://api.listenbrainz.org")} + , _listensSynchronizer {_ioContext, db, _sendQueue} { - LOG(INFO) << "Starting ListenBrainz scrobbler... API endpoint = '" << _apiEndpoint << "'"; - - _client.done().connect([this](Wt::AsioWrapper::error_code ec, const Wt::Http::Message& msg) - { - onClientDone(ec, msg); - }); - - _ioService.setThreadCount(1); - _ioService.start(); + LOG(INFO) << "Starting ListenBrainz scrobbler... API endpoint = '" << _sendQueue.getAPIBaseURL(); } - ListenBrainzScrobbler::~ListenBrainzScrobbler() + Scrobbler::~Scrobbler() { - _ioService.stop(); - - LOG(INFO) << "Stopped ListenBrainz scrobbler"; + LOG(INFO) << "Stopped ListenBrainz scrobbler!"; } void - ListenBrainzScrobbler::listenStarted(const Listen& listen) + Scrobbler::listenStarted(const Listen& listen) { - _ioService.post([=] - { - enqueListen(listen, Wt::WDateTime {}); - }); + enqueListen(listen, Wt::WDateTime {}); } void - ListenBrainzScrobbler::listenFinished(const Listen& listen, std::optional duration) + Scrobbler::listenFinished(const Listen& listen, std::optional duration) { if (duration && !canBeScrobbled(_db.getTLSSession(), listen.trackId, *duration)) return; - Listen timedListen {listen}; + const Listen timedListen {listen}; const Wt::WDateTime now {Wt::WDateTime::currentDateTime()}; - _ioService.post([=] - { - enqueListen(timedListen, now); - }); + enqueListen(timedListen, now); } void - ListenBrainzScrobbler::addListen(const Listen& listen, const Wt::WDateTime& timePoint) + Scrobbler::addTimedListen(const TimedListen& listen) { - assert(timePoint.isValid()); - - _ioService.post([=] - { - enqueListen(listen, timePoint); - }); + assert(listen.listenedAt.isValid()); + enqueListen(listen, listen.listenedAt); } - Wt::Dbo::ptr - ListenBrainzScrobbler::getListensTrackList(Database::Session& session, Wt::Dbo::ptr user) + Database::TrackList::pointer + Scrobbler::getListensTrackList(Database::Session& session, Database::User::pointer user) { - return Database::TrackList::get(session, historyTracklistName, Database::TrackList::Type::Internal, user); + return Utils::getListensTrackList(session, user); } void - ListenBrainzScrobbler::enqueListen(const Listen& listen, const Wt::WDateTime& timePoint) + Scrobbler::enqueListen(const Listen& listen, const Wt::WDateTime& timePoint) { - if (!timePoint.isValid()) - { - // If we are currently throttled, just replace the entry if it has no timePoint - // in order to only report the newest track listened to - // If not throttled, just search past the next current first message as it is being sent - - const std::size_t offset {_state == State::Throttled ? std::size_t {0} : std::size_t {1}}; - if (_sendQueue.size() > offset) - { - _sendQueue.erase(std::remove_if(std::next(std::begin(_sendQueue), offset), std::end(_sendQueue), - [&](const QueuedListen& queuedListen) { return queuedListen.listen.userId == listen.userId && !queuedListen.timePoint.isValid(); }), std::end(_sendQueue)); - } - } - - _sendQueue.emplace_back(QueuedListen {listen, timePoint}); - - LOG(DEBUG) << "listen queue size = " << _sendQueue.size(); - - if (_state == State::Idle) - sendNextQueuedListen(); - } - - void - ListenBrainzScrobbler::sendNextQueuedListen() - { - assert(_state == State::Idle); - - while (!_sendQueue.empty()) - { - if (sendListen(_sendQueue.front().listen, _sendQueue.front().timePoint)) - { - _state = State::Sending; - break; - } - - _sendQueue.pop_front(); - } - } - - bool - ListenBrainzScrobbler::sendListen(const Listen& listen, const Wt::WDateTime& timePoint) - { - Database::Session& session {_db.getTLSSession()}; - - const std::optional listenBrainzToken {getListenBrainzToken(session, listen.userId)}; - if (!listenBrainzToken) - return false; - - std::string payload {listenToJsonString(session, listen, timePoint, timePoint.isValid() ? "single" : "playing_now")}; - if (payload.empty()) - { - LOG(DEBUG) << "Cannot convert listen to json: skipping"; - return false; - } - - // now send this - Wt::Http::Message message; - message.addHeader("Authorization", "Token " + std::string {listenBrainzToken->getAsString()}); - message.addHeader("Content-Type", "application/json"); - message.addBodyText(payload); - - const std::string endPoint {_apiEndpoint + "submit-listens"}; - if (!_client.post(endPoint, message)) - { - LOG(ERROR) << "Cannot post to '" << endPoint << "': invalid scheme or URL?"; - return false; - } - - LOG(DEBUG) << "Listen POST done to '" << endPoint << "'"; - return true; - } - - void - ListenBrainzScrobbler::onClientDone(Wt::AsioWrapper::error_code ec, const Wt::Http::Message& msg) - { - assert(!_sendQueue.empty()); - QueuedListen& queuedListen {_sendQueue.front()}; - - _state = State::Idle; - - LOG(DEBUG) << "POST done. status = " << msg.status() << ", msg = '" << msg.body() << "'"; - if (ec) - { - LOG(ERROR) << "Retry " << queuedListen.retryCount << ", client error: '" << ec.message() << "'"; - // may be a network error, try again later - if (++queuedListen.retryCount > _maxRetryCount) - _sendQueue.pop_front(); - - throttle(_defaultRetryWaitDuration); + std::optional requestData {createSubmitListenRequestData(listen, timePoint)}; + if (!requestData) return; - } - bool mustThrottle{}; - - switch (msg.status()) + SendQueue::Request submitListen {std::move(*requestData)}; + if (timePoint.isValid()) { - case 429: - mustThrottle = true; - break; - - case 200: - if (queuedListen.timePoint.isValid()) - cacheListen(queuedListen.listen, queuedListen.timePoint); - _sendQueue.pop_front(); - break; - - default: - LOG(ERROR) << "Submit error: '" << msg.body() << "'"; - _sendQueue.pop_front(); - break; - } - - const auto remainingCount {headerReadAs(msg, "X-RateLimit-Remaining")}; - LOG(DEBUG) << "Remaining messages = " << (remainingCount ? *remainingCount : 0); - if (mustThrottle || (remainingCount && *remainingCount == 0)) - { - const auto waitDuration {headerReadAs(msg, "X-RateLimit-Reset-In")}; - throttle(waitDuration.value_or(_defaultRetryWaitDuration)); + submitListen.setPriority(SendQueue::Request::Priority::Normal); + submitListen.setOnSuccessFunc([=](std::string_view) + { + _listensSynchronizer.saveListen(TimedListen {listen, timePoint}); + }); } else { - sendNextQueuedListen(); + // We want "listen now" to appear as soon as possible + submitListen.setPriority(SendQueue::Request::Priority::High); } + + _sendQueue.enqueueRequest(std::move(submitListen)); } - void - ListenBrainzScrobbler::throttle(std::chrono::seconds requestedDuration) - { - assert(_state == State::Idle); - - const std::chrono::seconds duration {clamp(requestedDuration, _minRetryWaitDuration, _maxRetryWaitDuration)}; - LOG(DEBUG) << "Throttling for " << duration.count() << " seconds"; - - _ioService.schedule(duration, [this] - { - _state = State::Idle; - sendNextQueuedListen(); - }); - _state = State::Throttled; - } - - void - ListenBrainzScrobbler::cacheListen(const Listen& listen, const Wt::WDateTime& timePoint) + std::optional + Scrobbler::createSubmitListenRequestData(const Listen& listen, const Wt::WDateTime& timePoint) { Database::Session& session {_db.getTLSSession()}; - auto transaction {session.createUniqueTransaction()}; + const std::optional listenBrainzToken {Utils::getListenBrainzToken(session, listen.userId)}; + if (!listenBrainzToken) + return std::nullopt; - const Database::User::pointer user {Database::User::getById(session, listen.userId)}; - if (!user) - return; + SendQueue::RequestData requestData; + requestData.endpoint = "/1/submit-listens"; + requestData.type = SendQueue::RequestData::Type::POST; - const Database::Track::pointer track {Database::Track::getById(session, listen.trackId)}; - if (!track) - return; + std::string bodyText {listenToJsonString(session, listen, timePoint, timePoint.isValid() ? "single" : "playing_now")}; + if (bodyText.empty()) + { + LOG(DEBUG) << "Cannot convert listen to json: skipping"; + return std::nullopt; + } - Database::TrackList::pointer tracklist {getListensTrackList(session, user)}; - if (!tracklist) - tracklist = Database::TrackList::create(session, historyTracklistName, Database::TrackList::Type::Internal, false, user); + requestData.message.addBodyText(bodyText); + requestData.message.addHeader("Authorization", "Token " + std::string {listenBrainzToken->getAsString()}); + requestData.message.addHeader("Content-Type", "application/json"); - Database::TrackListEntry::create(session, track, getListensTrackList(session, user), timePoint); + return requestData; } - -} // Scrobbling +} // namespace Scrobbling::ListenBrainz diff --git a/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.hpp b/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.hpp index 690a7afc..ea6e5d1d 100644 --- a/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.hpp +++ b/src/libs/scrobbling/impl/listenbrainz/ListenBrainzScrobbler.hpp @@ -19,12 +19,12 @@ #pragma once -#include - -#include -#include +#include +#include #include "IScrobbler.hpp" +#include "ListensSynchronizer.hpp" +#include "SendQueue.hpp" namespace Database { @@ -33,59 +33,33 @@ namespace Database class TrackList; } -namespace Scrobbling +namespace Scrobbling::ListenBrainz { - class ListenBrainzScrobbler final : public IScrobbler + class Scrobbler final : public IScrobbler { public: - ListenBrainzScrobbler(Database::Db& db); - ~ListenBrainzScrobbler(); + Scrobbler(boost::asio::io_context& ioContext, Database::Db& db); + ~Scrobbler(); - ListenBrainzScrobbler(const ListenBrainzScrobbler&) = delete; - ListenBrainzScrobbler(const ListenBrainzScrobbler&&) = delete; - ListenBrainzScrobbler& operator=(const ListenBrainzScrobbler&) = delete; - ListenBrainzScrobbler& operator=(const ListenBrainzScrobbler&&) = delete; + Scrobbler(const Scrobbler&) = delete; + Scrobbler(const Scrobbler&&) = delete; + Scrobbler& operator=(const Scrobbler&) = delete; + Scrobbler& operator=(const Scrobbler&&) = delete; private: void listenStarted(const Listen& listen) override; void listenFinished(const Listen& listen, std::optional duration) override; - void addListen(const Listen& listen, const Wt::WDateTime& timePoint) override; - + void addTimedListen(const TimedListen& listen) override; Wt::Dbo::ptr getListensTrackList(Database::Session& session, Wt::Dbo::ptr user) override; + // Submit listens void enqueListen(const Listen& listen, const Wt::WDateTime& timePoint); - void sendNextQueuedListen(); - bool sendListen(const Listen& listen, const Wt::WDateTime& timePoint); - void onClientDone(Wt::AsioWrapper::error_code ec, const Wt::Http::Message& msg); - void throttle(std::chrono::seconds duration); - - void cacheListen(const Listen& listen, const Wt::WDateTime& timePoint); - - enum class State - { - Idle, - Throttled, - Sending, - }; - State _state {State::Idle}; - - const std::string _apiEndpoint; - const std::size_t _maxRetryCount {2}; - const std::chrono::seconds _defaultRetryWaitDuration {30}; - const std::chrono::seconds _minRetryWaitDuration {1}; - const std::chrono::seconds _maxRetryWaitDuration {300}; + std::optional createSubmitListenRequestData(const Listen& listen, const Wt::WDateTime& timePoint); + boost::asio::io_context& _ioContext; Database::Db& _db; - Wt::WIOService _ioService; - Wt::Http::Client _client {_ioService}; - - struct QueuedListen - { - Listen listen; - Wt::WDateTime timePoint; - std::size_t retryCount {}; - }; - std::deque _sendQueue; + SendQueue _sendQueue; + ListensSynchronizer _listensSynchronizer; }; -} // Scrobbling +} // Scrobbling::ListenBrainz diff --git a/src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.cpp b/src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.cpp new file mode 100644 index 00000000..03395ae0 --- /dev/null +++ b/src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.cpp @@ -0,0 +1,514 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#include "ListenBrainzScrobbler.hpp" + +#include +#include +#include +#include +#include + +#include "database/Artist.hpp" +#include "database/Db.hpp" +#include "database/Release.hpp" +#include "database/Session.hpp" +#include "database/Track.hpp" +#include "database/TrackList.hpp" +#include "database/User.hpp" +#include "utils/IConfig.hpp" +#include "utils/Logger.hpp" +#include "utils/Service.hpp" + +#include "Utils.hpp" + +#define LOG(sev) LMS_LOG(SCROBBLING, sev) << "[listenbrainz Synchronizer] - " + +namespace +{ + using namespace Scrobbling::ListenBrainz; + + SendQueue::RequestData + createValidateTokenRequestData(std::string_view authToken) + { + SendQueue::RequestData requestData; + requestData.type = SendQueue::RequestData::Type::GET; + requestData.endpoint = "/1/validate-token"; + requestData.headers = { {"Authorization", "Token " + std::string {authToken}} }; + + return requestData; + } + + std::string + parseValidateToken(std::string_view msgBody) + { + std::string listenBrainzUserName; + + Wt::Json::ParseError error; + Wt::Json::Object root; + if (!Wt::Json::parse(std::string {msgBody}, root, error)) + { + LOG(ERROR) << "Cannot parse 'validate-token' result: " << error.what(); + return listenBrainzUserName; + } + + if (!root.get("valid").orIfNull(false)) + { + LOG(INFO) << "Invalid listenbrainz user"; + return listenBrainzUserName; + } + + listenBrainzUserName = root.get("user_name").orIfNull(""); + return listenBrainzUserName; + } + + SendQueue::RequestData + createListenCountRequestData(std::string_view listenBrainzUserName) + { + LOG(DEBUG) << "Getting listen count for listenbrainz user '" << listenBrainzUserName << "'"; + + SendQueue::RequestData requestData; + requestData.type = SendQueue::RequestData::Type::GET; + requestData.endpoint = "/1/user/" + std::string {listenBrainzUserName} + "/listen-count"; + + return requestData; + } + + std::optional + parseListenCount(std::string_view msgBody) + { + try + { + Wt::Json::Object root; + Wt::Json::parse(std::string {msgBody}, root); + + const Wt::Json::Object& payload {static_cast(root.get("payload"))}; + return static_cast(payload.get("count")); + } + catch (const Wt::WException& e) + { + LOG(ERROR) << "Cannot parse listen count response: " << e.what(); + return std::nullopt; + } + } + + SendQueue::RequestData + createGetListensRequestData(std::string_view listenBrainzUserName, const Wt::WDateTime& maxDateTime) + { + LOG(DEBUG) << "Getting listens for listenbrainz user '" << listenBrainzUserName << "' with max_ts = " << maxDateTime.toString(); + + SendQueue::RequestData requestData; + requestData.type = SendQueue::RequestData::Type::GET; + requestData.endpoint = "/1/user/" + std::string {listenBrainzUserName} + "/listens?max_ts=" + std::to_string(maxDateTime.toTime_t()); + + return requestData; + } + + Database::Track::pointer + tryMatchListen(Database::Session& session, const Wt::Json::Object& metadata) + { + Database::Track::pointer track; + + // first try to get the associated track using MBIDs, and then fallback on names + if (metadata.type("additional_info") == Wt::Json::Type::Object) + { + const Wt::Json::Object& additionalInfo = metadata.get("additional_info"); + if (std::optional recordingMBID {UUID::fromString(additionalInfo.get("recording_mbid").orIfNull(""))}) + track = Database::Track::getByRecordingMBID(session, *recordingMBID); + } + + if (track) + return track; + + // these fields are mandatory + const std::string trackName {static_cast(metadata.get("track_name"))}; + const std::string releaseName {static_cast(metadata.get("release_name"))}; + + auto tracks {Database::Track::getByNameAndReleaseName(session, trackName, releaseName)}; + if (tracks.size() > 1) + { + tracks.erase(std::remove_if(std::begin(tracks), std::end(tracks), + [&](const Database::Track::pointer track) + { + if (std::string artistName {metadata.get("artist_name").orIfNull("")}; !artistName.empty()) + { + const auto& artists {track->getArtists({Database::TrackArtistLinkType::Artist})}; + if (std::none_of(std::begin(artists), std::end(artists), [&](const Database::Artist::pointer& artist) { return artist->getName() == artistName; })) + return true; + } + if (metadata.type("additional_info") == Wt::Json::Type::Object) + { + const Wt::Json::Object& additionalInfo = metadata.get("additional_info"); + if (track->getTrackNumber()) + { + int otherTrackNumber {additionalInfo.get("tracknumber").orIfNull(-1)}; + if (otherTrackNumber > 0 && static_cast(otherTrackNumber) != *track->getTrackNumber()) + return true; + } + + if (auto releaseMBID {track->getRelease()->getMBID()}) + { + if (std::optional otherReleaseMBID {UUID::fromString(additionalInfo.get("release_mbid").orIfNull(""))}) + { + if (otherReleaseMBID->getAsString() != releaseMBID->getAsString()) + return true; + } + } + } + + return false; + }), std::end(tracks)); + } + + if (tracks.size() == 1) + track = tracks.front(); + + return track; + } + + struct ParseGetListensResult + { + Wt::WDateTime oldestEntry; + std::size_t listenCount{}; + std::vector matchedListens; + }; + ParseGetListensResult + parseGetListens(Database::Session& session, std::string_view msgBody, Database::IdType userId) + { + ParseGetListensResult result; + + try + { + Wt::Json::Object root; + Wt::Json::parse(std::string {msgBody}, root); + + const Wt::Json::Object& payload = root.get("payload"); + const Wt::Json::Array& listens = payload.get("listens"); + + LOG(DEBUG) << "Got " << listens.size() << " listens"; + + if (listens.empty()) + return result; + + auto transaction {session.createSharedTransaction()}; + + for (const Wt::Json::Value& value : listens) + { + const Wt::Json::Object& listen = value; + const Wt::WDateTime listenedAt {Wt::WDateTime::fromTime_t(static_cast(listen.get("listened_at")))}; + const Wt::Json::Object& metadata = listen.get("track_metadata"); + + if (!listenedAt.isValid()) + { + LOG(ERROR) << "bad listened_at field!"; + continue; + } + + result.listenCount++; + if (!result.oldestEntry.isValid()) + result.oldestEntry = listenedAt; + else if (listenedAt < result.oldestEntry) + result.oldestEntry = listenedAt; + + if (const Database::Track::pointer track {tryMatchListen(session, metadata)}) + result.matchedListens.emplace_back(Scrobbling::TimedListen {userId, track.id(), listenedAt}); + } + } + catch (const Wt::WException& error) + { + LOG(ERROR) << "Cannot parse 'get-listens' result: " << error.what(); + } + + return result; + } +} + +namespace Scrobbling::ListenBrainz +{ + ListensSynchronizer::ListensSynchronizer(boost::asio::io_context& ioContext, Database::Db& db, SendQueue& sendQueue) + : _ioContext {ioContext} + , _db {db} + , _sendQueue {sendQueue} + , _maxSyncListenCount {Service::get()->getULong("listenbrainz-max-sync-listen-count", 1000)} + , _syncListensPeriod {Service::get()->getULong("listenbrainz-sync-listens-period-hours", 1)} + { + LOG(INFO) << "Starting Listens synchronizer, maxSyncListenCount = " << _maxSyncListenCount << ", _syncListensPeriod = " << _syncListensPeriod.count() << " hours"; + + scheduleGetListens(std::chrono::seconds {30}); + } + + void + ListensSynchronizer::saveListen(const TimedListen& listen) + { + _strand.dispatch([=] + { + Database::Session& session {_db.getTLSSession()}; + + auto transaction {session.createUniqueTransaction()}; + + const Database::User::pointer user {Database::User::getById(session, listen.userId)}; + if (!user) + return; + + const Database::Track::pointer track {Database::Track::getById(session, listen.trackId)}; + if (!track) + return; + + Database::TrackListEntry::create(session, track, Utils::getOrCreateListensTrackList(session, user), listen.listenedAt); + + UserContext& context {getUserContext(listen.userId)}; + if (context.listenCount) + (*context.listenCount)++; + }); + } + + ListensSynchronizer::UserContext& + ListensSynchronizer::getUserContext(Database::IdType userId) + { + auto itContext {_userContexts.find(userId)}; + if (itContext == std::cend(_userContexts)) + { + auto [itNewContext, inserted] {_userContexts.emplace(userId, userId)}; + itContext = itNewContext; + } + + return itContext->second; + } + + bool + ListensSynchronizer::isFetching() const + { + return std::any_of(std::cbegin(_userContexts), std::cend(_userContexts), [](const auto& contextEntry) + { + const auto& [userId, context] {contextEntry}; + return context.fetching; + }); + } + + void + ListensSynchronizer::scheduleGetListens(std::chrono::seconds fromNow) + { + if (_syncListensPeriod.count() == 0 || _maxSyncListenCount == 0) + return; + + LOG(DEBUG) << "Scheduled sync in " << fromNow.count() << " seconds..."; + _getListensTimer.expires_after(fromNow); + _getListensTimer.async_wait(boost::asio::bind_executor(_strand, [this] (const boost::system::error_code& ec) + { + if (ec == boost::asio::error::operation_aborted) + { + LOG(DEBUG) << "getListens aborted"; + return; + } + + startGetListens(); + })); + } + + void + ListensSynchronizer::startGetListens() + { + LOG(DEBUG) << "GetListens started!!!"; + + assert(!isFetching()); + + std::vector userIds; + { + Database::Session& session {_db.getTLSSession()}; + auto transaction {session.createSharedTransaction()}; + userIds = Database::User::getAllIds(_db.getTLSSession()); + } + + for (const Database::IdType userId : userIds) + { + if (Utils::getListenBrainzToken(_db.getTLSSession(), userId)) + startGetListens(getUserContext(userId)); + } + + if (!isFetching()) + scheduleGetListens(_syncListensPeriod); + } + + void + ListensSynchronizer::startGetListens(UserContext& context) + { + context.fetching = true; + context.listenBrainzUserName = ""; + context.maxDateTime = {}; + context.fetchedListenCount = 0; + context.matchedListenCount = 0; + context.importedListenCount = 0; + + enqueValidateToken(context); + } + + void + ListensSynchronizer::onGetListensEnded(UserContext& context) + { + _strand.dispatch([this, &context] + { + LOG(DEBUG) << "Fetch done for user " << context.userId << ", fetched: " << context.fetchedListenCount << ", matched: " << context.matchedListenCount << ", imported: " << context.importedListenCount; + context.fetching = false; + + if (!isFetching()) + scheduleGetListens(_syncListensPeriod); + }); + } + + void + ListensSynchronizer::enqueValidateToken(UserContext& context) + { + assert(context.listenBrainzUserName.empty()); + + std::optional requestData {createValidateTokenRequestData(context.userId)}; + if (!requestData) + { + onGetListensEnded(context); + return; + } + + SendQueue::Request validateTokenRequest {std::move(*requestData)}; + validateTokenRequest.setOnSuccessFunc([this, &context] (std::string_view msgBody) + { + context.listenBrainzUserName = parseValidateToken(msgBody); + if (context.listenBrainzUserName.empty()) + { + onGetListensEnded(context); + return; + } + enqueGetListenCount(context); + }); + validateTokenRequest.setOnFailureFunc([this, &context] + { + onGetListensEnded(context); + }); + + validateTokenRequest.setPriority(SendQueue::Request::Priority::Low); + _sendQueue.enqueueRequest(std::move(validateTokenRequest)); + } + + void + ListensSynchronizer::enqueGetListenCount(UserContext& context) + { + assert(!context.listenBrainzUserName.empty()); + + SendQueue::Request getListenCountRequest {createListenCountRequestData(context.listenBrainzUserName)}; + getListenCountRequest.setOnSuccessFunc([=, &context] (std::string_view msgBody) + { + const auto listenCount = parseListenCount(msgBody); + if (listenCount) + LOG(DEBUG) << "Listen count for listenbrainz user '" << context.listenBrainzUserName << "' = " << *listenCount; + + bool needSync {listenCount && (!context.listenCount || *context.listenCount != *listenCount)}; + context.listenCount = listenCount; + + if (!needSync) + { + onGetListensEnded(context); + return; + } + + context.maxDateTime = Wt::WDateTime::currentDateTime(); + enqueGetListens(context); + }); + getListenCountRequest.setOnFailureFunc([this, &context] + { + onGetListensEnded(context); + }); + + getListenCountRequest.setPriority(SendQueue::Request::Priority::Low); + _sendQueue.enqueueRequest(std::move(getListenCountRequest)); + } + + void + ListensSynchronizer::enqueGetListens(UserContext& context) + { + assert(!context.listenBrainzUserName.empty()); + + SendQueue::Request getListensRequest {::createGetListensRequestData(context.listenBrainzUserName, context.maxDateTime)}; + getListensRequest.setOnSuccessFunc([=, &context] (std::string_view msgBody) + { + processGetListensResponse(msgBody, context); + if (context.fetchedListenCount >= _maxSyncListenCount || !context.maxDateTime.isValid()) + { + onGetListensEnded(context); + return; + } + + enqueGetListens(context); + }); + getListensRequest.setOnFailureFunc([=, &context] + { + onGetListensEnded(context); + }); + + getListensRequest.setPriority(SendQueue::Request::Priority::Low); + _sendQueue.enqueueRequest(std::move(getListensRequest)); + } + + std::optional + ListensSynchronizer::createValidateTokenRequestData(Database::IdType userId) + { + Database::Session& session {_db.getTLSSession()}; + + const std::optional listenBrainzToken {Utils::getListenBrainzToken(session, userId)}; + if (!listenBrainzToken) + return std::nullopt; + + return ::createValidateTokenRequestData(listenBrainzToken->getAsString()); + } + + void + ListensSynchronizer::processGetListensResponse(std::string_view msgBody, UserContext& context) + { + Database::Session& session {_db.getTLSSession()}; + + const ParseGetListensResult parseResult {parseGetListens(session, msgBody, context.userId)}; + + context.fetchedListenCount += parseResult.listenCount; + context.matchedListenCount += parseResult.matchedListens.size(); + context.maxDateTime = parseResult.oldestEntry; + + if (parseResult.matchedListens.empty()) + return; + + auto transaction {session.createUniqueTransaction()}; + + Database::User::pointer user {Database::User::getById(session, context.userId)}; + if (!user) + return; + + Database::TrackList::pointer tracklist {Utils::getOrCreateListensTrackList(session, user)}; + + for (const TimedListen& listen : parseResult.matchedListens) + { + const Database::Track::pointer track {Database::Track::getById(session, listen.trackId)}; + if (!track) + continue; + + if (!tracklist->getEntryByTrackAndDateTime(track, listen.listenedAt)) + { + context.importedListenCount++; + Database::TrackListEntry::create(session, track, tracklist, listen.listenedAt); + } + } + } + +} // namespace Scrobbling::ListenBrainz + diff --git a/src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.hpp b/src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.hpp new file mode 100644 index 00000000..1a49ed98 --- /dev/null +++ b/src/libs/scrobbling/impl/listenbrainz/ListensSynchronizer.hpp @@ -0,0 +1,98 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#pragma once + +#include +#include +#include +#include +#include +#include + +#include "database/Types.hpp" +#include "scrobbling/Listen.hpp" +#include "SendQueue.hpp" + +namespace Database +{ + class Db; + class Session; + class TrackList; + class User; +} + +namespace Scrobbling::ListenBrainz +{ + class ListensSynchronizer + { + public: + ListensSynchronizer(boost::asio::io_context& ioContext, Database::Db& db, SendQueue& sendQueue); + + void saveListen(const TimedListen& listen); + + private: + struct UserContext + { + UserContext(Database::IdType id) : userId {id} {} + + UserContext(const UserContext&) = delete; + UserContext(UserContext&&) = delete; + UserContext& operator=(const UserContext&) = delete; + UserContext& operator=(UserContext&&) = delete; + + const Database::IdType userId; + bool fetching {}; + std::optional listenCount {}; + + // resetted at each fetch + std::string listenBrainzUserName; // need to be resolved first + Wt::WDateTime maxDateTime; + std::size_t fetchedListenCount{}; + std::size_t matchedListenCount{}; + std::size_t importedListenCount{}; + + }; + + UserContext& getUserContext(Database::IdType userId); + bool isFetching() const; + void scheduleGetListens(std::chrono::seconds fromNow); + void startGetListens(); + void startGetListens(UserContext& context); + void onGetListensEnded(UserContext& context); + void enqueValidateToken(UserContext& context); + void enqueGetListenCount(UserContext& context); + void enqueGetListens(UserContext& context); + std::optional createValidateTokenRequestData(Database::IdType userId); + std::optional createGetListensRequestData(std::string_view listenBrainzUserName, const Wt::WDateTime& maxDateTime); + void processGetListensResponse(std::string_view body, UserContext& context); + + boost::asio::io_context& _ioContext; + boost::asio::io_context::strand _strand {_ioContext}; + Database::Db& _db; + SendQueue& _sendQueue; + boost::asio::steady_timer _getListensTimer {_ioContext}; + + std::unordered_map _userContexts; + + const std::size_t _maxSyncListenCount; + const std::chrono::hours _syncListensPeriod; + }; +} // Scrobbling::ListenBrainz + diff --git a/src/libs/scrobbling/impl/listenbrainz/SendQueue.cpp b/src/libs/scrobbling/impl/listenbrainz/SendQueue.cpp new file mode 100644 index 00000000..dbede5e0 --- /dev/null +++ b/src/libs/scrobbling/impl/listenbrainz/SendQueue.cpp @@ -0,0 +1,228 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#include "SendQueue.hpp" + +#include + +#include "utils/Logger.hpp" +#include "utils/String.hpp" + +#define LOG(sev) LMS_LOG(SCROBBLING, sev) << "[listenbrainz SendQueue] - " + +namespace StringUtils +{ + template<> + std::optional + readAs(const std::string& str) + { + std::optional res; + + if (const std::optional value {StringUtils::readAs(str)}) + res = std::chrono::seconds {*value}; + + return res; + } +} + +namespace +{ + template + std::optional + headerReadAs(const Wt::Http::Message& msg, std::string_view headerName) + { + std::optional res; + + if (const std::string* headerValue {msg.getHeader(std::string {headerName})}) + res = StringUtils::readAs(*headerValue); + + return res; + } +} + +namespace Scrobbling::ListenBrainz +{ + SendQueue::SendQueue(boost::asio::io_context& ioContext, std::string_view apiBaseURL) + : _ioContext {ioContext} + , _apiBaseURL {apiBaseURL} + { + _client.done().connect([this](Wt::AsioWrapper::error_code ec, const Wt::Http::Message& msg) + { + _strand.dispatch([=, msg = std::move(msg)] + { + onClientDone(ec, msg); + }); + }); + } + + SendQueue::~SendQueue() + { + _client.abort(); + } + + void + SendQueue::enqueueRequest(Request request) + { + _strand.dispatch([this, request = std::move(request)]() + { + _sendQueue[request._priority].emplace_back(std::move(request)); + + if (_state == State::Idle) + sendNextQueuedRequest(); + }); + } + + void + SendQueue::sendNextQueuedRequest() + { + assert(_state == State::Idle); + + for (auto& [prio, requests] : _sendQueue) + { + LOG(DEBUG) << "Processing prio " << static_cast(prio) << ", request count = " << requests.size(); + while (!requests.empty()) + { + Request request {std::move(requests.front())}; + requests.pop_front(); + + if (!sendRequest(request._requestData)) + continue; + + _state = State::Sending; + _currentRequest = std::move(request); + return; + } + } + } + + bool + SendQueue::sendRequest(const RequestData& requestData) + { + const std::string url {_apiBaseURL + requestData.endpoint}; + + LOG(DEBUG) << "Sending request type " << (requestData.type == RequestData::Type::GET ? "GET" : "POST") << " to url '" << url << "'"; + + bool res{}; + switch (requestData.type) + { + case RequestData::Type::GET: + res = _client.get(url, requestData.headers); + break; + case RequestData::Type::POST: + res = _client.post(url, requestData.message); + break; + } + + if (!res) + LOG(ERROR) << "Send failed, bad url or unsupported scheme?"; + + return res; + } + + void + SendQueue::onClientDone(Wt::AsioWrapper::error_code ec, const Wt::Http::Message& msg) + { + if (ec == boost::asio::error::operation_aborted) + { + LOG(DEBUG) << "SendQueue: client aborted"; + return; + } + + assert(_currentRequest); + Request request {std::move(*_currentRequest)}; + _state = State::Idle; + + LOG(DEBUG) << "Client done. status = " << msg.status(); + if (ec) + { + LOG(ERROR) << "Retry " << request._retryCount << ", client error: '" << ec.message() << "'"; + + // may be a network error, try again later + throttle(_defaultRetryWaitDuration); + + if (request._retryCount++ < _maxRetryCount) + { + _sendQueue[request._priority].emplace_front(std::move(request)); + } + else + { + LOG(ERROR) << "Too many retries, giving up operation and throttle"; + if (request._onFailureFunc) + request._onFailureFunc(); + } + return; + } + + bool mustThrottle{}; + if (msg.status() == 429) + _sendQueue[request._priority].emplace_front(std::move(request)); + + const auto remainingCount {headerReadAs(msg, "X-RateLimit-Remaining")}; + LOG(DEBUG) << "Remaining messages = " << (remainingCount ? *remainingCount : 0); + if (mustThrottle || (remainingCount && *remainingCount == 0)) + { + const auto waitDuration {headerReadAs(msg, "X-RateLimit-Reset-In")}; + throttle(waitDuration.value_or(_defaultRetryWaitDuration)); + } + + if (!mustThrottle) + { + if (msg.status() == 200) + { + if (request._onSuccessFunc) + request._onSuccessFunc(msg.body()); + } + else + { + LOG(ERROR) << "Send error: '" << msg.body() << "'"; + if (request._onFailureFunc) + request._onFailureFunc(); + } + } + + if (_state == State::Idle) + sendNextQueuedRequest(); + } + + void + SendQueue::throttle(std::chrono::seconds requestedDuration) + { + assert(_state == State::Idle); + + const std::chrono::seconds duration {clamp(requestedDuration, _minRetryWaitDuration, _maxRetryWaitDuration)}; + LOG(DEBUG) << "Throttling for " << duration.count() << " seconds"; + + _throttleTimer.expires_after(duration); + _throttleTimer.async_wait([this](const boost::system::error_code& ec) + { + if (ec == boost::asio::error::operation_aborted) + { + LOG(DEBUG) << "SendQueue: throttle aborted"; + return; + } + + if (ec) + LOG(ERROR) << "async_wait failed:" << ec.message(); + + _state = State::Idle; + sendNextQueuedRequest(); + }); + _state = State::Throttled; + } +} // namespace Scrobbling::ListenBrainz diff --git a/src/libs/scrobbling/impl/listenbrainz/SendQueue.hpp b/src/libs/scrobbling/impl/listenbrainz/SendQueue.hpp new file mode 100644 index 00000000..7345212e --- /dev/null +++ b/src/libs/scrobbling/impl/listenbrainz/SendQueue.hpp @@ -0,0 +1,118 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#pragma once + +#include + +#include +#include +#include + +#include + +namespace Scrobbling::ListenBrainz +{ + class SendQueue + { + public: + SendQueue(boost::asio::io_context& ioContext, std::string_view apiBaseURL); + ~SendQueue(); + + SendQueue(const SendQueue&) = delete; + SendQueue(const SendQueue&&) = delete; + SendQueue& operator=(const SendQueue&) = delete; + SendQueue& operator=(const SendQueue&&) = delete; + + // generic queue operations + struct RequestData + { + enum class Type + { + GET, + POST, + }; + + Type type; + std::string endpoint; // relative URL to the base API + std::vector headers; // used by GET + Wt::Http::Message message; // used by POST + }; + + class Request + { + public: + + enum class Priority + { + High, + Normal, + Low, + }; + + Request(RequestData requestData) : _requestData {std::move(requestData)} {} + + using OnSuccessFunc = std::function; + using OnFailureFunc = std::function; + + void setOnSuccessFunc(OnSuccessFunc onSuccessFunc) { _onSuccessFunc = onSuccessFunc; } + void setOnFailureFunc(OnFailureFunc onFailureFunc) { _onFailureFunc = onFailureFunc; } + void setPriority(Priority priority) { _priority = priority; } + + private: + friend class SendQueue; + RequestData _requestData; + Priority _priority {Priority::Normal}; + std::size_t _retryCount {}; + OnSuccessFunc _onSuccessFunc; + OnFailureFunc _onFailureFunc; + }; + + std::string_view getAPIBaseURL() const { return _apiBaseURL; } + void enqueueRequest(Request request); + + private: + void sendNextQueuedRequest(); + bool sendRequest(const RequestData& request); + void onClientDone(Wt::AsioWrapper::error_code ec, const Wt::Http::Message& msg); + void throttle(std::chrono::seconds duration); + + const std::size_t _maxRetryCount {2}; + const std::chrono::seconds _defaultRetryWaitDuration {30}; + const std::chrono::seconds _minRetryWaitDuration {1}; + const std::chrono::seconds _maxRetryWaitDuration {300}; + + enum class State + { + Idle, + Throttled, + Sending, + }; + boost::asio::io_context& _ioContext; + boost::asio::io_context::strand _strand {_ioContext}; + boost::asio::steady_timer _throttleTimer {_ioContext}; + std::string _apiBaseURL; + State _state {State::Idle}; + Wt::Http::Client _client {_ioContext}; + std::map> _sendQueue; + std::optional _currentRequest; + }; + +} // namespace Scrobbling::ListenBrainz + diff --git a/src/libs/scrobbling/impl/listenbrainz/Utils.cpp b/src/libs/scrobbling/impl/listenbrainz/Utils.cpp new file mode 100644 index 00000000..47cf901c --- /dev/null +++ b/src/libs/scrobbling/impl/listenbrainz/Utils.cpp @@ -0,0 +1,63 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#include "Utils.hpp" + +#include + +#include "database/Session.hpp" +#include "database/TrackList.hpp" +#include "database/User.hpp" + +static constexpr std::string_view historyTracklistName {"__scrobbler_listenbrainz_history__"}; + +namespace Scrobbling::ListenBrainz::Utils +{ + std::optional + getListenBrainzToken(Database::Session& session, Database::IdType userId) + { + auto transaction {session.createSharedTransaction()}; + + const Database::User::pointer user {Database::User::getById(session, userId)}; + if (!user) + return std::nullopt; + + if (user->getScrobbler() != Database::Scrobbler::ListenBrainz) + return std::nullopt; + + return user->getListenBrainzToken(); + } + + Database::TrackList::pointer + getListensTrackList(Database::Session& session, Database::User::pointer user) + { + return Database::TrackList::get(session, historyTracklistName, Database::TrackList::Type::Internal, user); + } + + Database::TrackList::pointer + getOrCreateListensTrackList(Database::Session& session, Database::User::pointer user) + { + Database::TrackList::pointer tracklist {getListensTrackList(session, user)}; + if (!tracklist) + tracklist = Database::TrackList::create(session, historyTracklistName, Database::TrackList::Type::Internal, false, user); + + return tracklist; + } + +} diff --git a/src/libs/scrobbling/impl/listenbrainz/Utils.hpp b/src/libs/scrobbling/impl/listenbrainz/Utils.hpp new file mode 100644 index 00000000..84074086 --- /dev/null +++ b/src/libs/scrobbling/impl/listenbrainz/Utils.hpp @@ -0,0 +1,38 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#pragma once + +#include +#include "utils/UUID.hpp" +#include "database/Types.hpp" + +namespace Database +{ + class Session; + class TrackList; + class User; +} + +namespace Scrobbling::ListenBrainz::Utils +{ + std::optional getListenBrainzToken(Database::Session& session, Database::IdType userId); + Wt::Dbo::ptr getOrCreateListensTrackList(Database::Session& session, Wt::Dbo::ptr user); + Wt::Dbo::ptr getListensTrackList(Database::Session& session, Wt::Dbo::ptr user); +} diff --git a/src/libs/scrobbling/include/scrobbling/IScrobbling.hpp b/src/libs/scrobbling/include/scrobbling/IScrobbling.hpp index 376ad085..e1fc5ce9 100644 --- a/src/libs/scrobbling/include/scrobbling/IScrobbling.hpp +++ b/src/libs/scrobbling/include/scrobbling/IScrobbling.hpp @@ -19,6 +19,8 @@ #pragma once +#include + #include #include #include @@ -51,7 +53,7 @@ namespace Scrobbling virtual void listenStarted(const Listen& listen) = 0; virtual void listenFinished(const Listen& listen, std::optional playedDuration = std::nullopt) = 0; - virtual void addListen(const Listen& listen, Wt::WDateTime timePoint) = 0; + virtual void addTimedListen(const TimedListen& listen) = 0; // Stats // From most recent to oldest @@ -95,7 +97,7 @@ namespace Scrobbling bool& moreResults) = 0; }; - std::unique_ptr createScrobbling(Database::Db& db); + std::unique_ptr createScrobbling(boost::asio::io_service& ioService, Database::Db& db); } // ns Scrobbling diff --git a/src/libs/scrobbling/include/scrobbling/Listen.hpp b/src/libs/scrobbling/include/scrobbling/Listen.hpp index 48c10456..b67a2ae9 100644 --- a/src/libs/scrobbling/include/scrobbling/Listen.hpp +++ b/src/libs/scrobbling/include/scrobbling/Listen.hpp @@ -19,6 +19,8 @@ #pragma once +#include + #include "database/Types.hpp" namespace Scrobbling @@ -28,5 +30,10 @@ namespace Scrobbling Database::IdType userId {}; Database::IdType trackId {}; }; + + struct TimedListen : public Listen + { + Wt::WDateTime listenedAt; + }; } // ns Scrobbling diff --git a/src/libs/subsonic/impl/SubsonicResource.cpp b/src/libs/subsonic/impl/SubsonicResource.cpp index a9caad89..cccd9a2d 100644 --- a/src/libs/subsonic/impl/SubsonicResource.cpp +++ b/src/libs/subsonic/impl/SubsonicResource.cpp @@ -1679,7 +1679,7 @@ handleScrobble(RequestContext& context) { const Database::IdType trackId {ids[i].value}; const unsigned long time {times[i]}; - Service::get()->addListen({context.userId, trackId}, Wt::WDateTime::fromTime_t(static_cast(time / 1000))); + Service::get()->addTimedListen({context.userId, trackId, Wt::WDateTime::fromTime_t(static_cast(time / 1000))}); } } } diff --git a/src/libs/utils/CMakeLists.txt b/src/libs/utils/CMakeLists.txt index 39d6ee8d..f6ef50b7 100644 --- a/src/libs/utils/CMakeLists.txt +++ b/src/libs/utils/CMakeLists.txt @@ -4,6 +4,7 @@ add_library(lmsutils SHARED impl/ChildProcessManager.cpp impl/Config.cpp impl/FileResourceHandler.cpp + impl/IOContextRunner.cpp impl/Logger.cpp impl/NetAddress.cpp impl/Path.cpp diff --git a/src/libs/utils/impl/ChildProcessManager.cpp b/src/libs/utils/impl/ChildProcessManager.cpp index 128fcd9f..ab563329 100644 --- a/src/libs/utils/impl/ChildProcessManager.cpp +++ b/src/libs/utils/impl/ChildProcessManager.cpp @@ -25,42 +25,14 @@ std::unique_ptr -createChildProcessManager() +createChildProcessManager(boost::asio::io_context& ioContext) { - return std::make_unique(); + return std::make_unique(ioContext); } -ChildProcessManager::ChildProcessManager() -: _work {boost::asio::make_work_guard(_ioContext)} +ChildProcessManager::ChildProcessManager(boost::asio::io_context& ioContext) +: _ioContext {ioContext} { - start(); -} - -ChildProcessManager::~ChildProcessManager() -{ - stop(); -} - -void -ChildProcessManager::start() -{ - LMS_LOG(CHILDPROCESS, INFO) << "Starting child process manager..."; - - _thread = std::make_unique([&]() - { - _ioContext.run(); - }); - - LMS_LOG(CHILDPROCESS, INFO) << "Child process manager started!"; -} - -void -ChildProcessManager::stop() -{ - LMS_LOG(CHILDPROCESS, INFO) << "Stopping child process manager"; - _work.reset(); - _thread->join(); - LMS_LOG(CHILDPROCESS, INFO) << "Stopped child process manager"; } std::unique_ptr diff --git a/src/libs/utils/impl/ChildProcessManager.hpp b/src/libs/utils/impl/ChildProcessManager.hpp index eb8f0353..79cf3943 100644 --- a/src/libs/utils/impl/ChildProcessManager.hpp +++ b/src/libs/utils/impl/ChildProcessManager.hpp @@ -23,15 +23,14 @@ #include #include -#include #include "utils/IChildProcessManager.hpp" class ChildProcessManager : public IChildProcessManager { public: - ChildProcessManager(); - ~ChildProcessManager(); + ChildProcessManager(boost::asio::io_context& ioContext); + ~ChildProcessManager() = default; ChildProcessManager(const ChildProcessManager&) = delete; ChildProcessManager(ChildProcessManager&&) = delete; @@ -41,12 +40,7 @@ class ChildProcessManager : public IChildProcessManager private: std::unique_ptr spawnChildProcess(const std::filesystem::path& path, const IChildProcess::Args& args) override; - void start(); - void stop(); - - boost::asio::io_context _ioContext; - std::unique_ptr _thread; - boost::asio::executor_work_guard _work; + boost::asio::io_context& _ioContext; }; diff --git a/src/libs/utils/impl/IOContextRunner.cpp b/src/libs/utils/impl/IOContextRunner.cpp new file mode 100644 index 00000000..038f7ff5 --- /dev/null +++ b/src/libs/utils/impl/IOContextRunner.cpp @@ -0,0 +1,49 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#include "utils/Logger.hpp" +#include "utils/IOContextRunner.hpp" + +IOContextRunner::IOContextRunner(boost::asio::io_service& ioService, std::size_t threadCount) +: _ioService {ioService} +, _work {ioService} +{ + LMS_LOG(UTILS, INFO) << "Starting IO Context with " << threadCount << " threads..."; + for (std::size_t i {}; i < threadCount; ++i) + _threads.emplace_back([&] { _ioService.run(); }); +} + +void +IOContextRunner::stop() +{ + LMS_LOG(UTILS, INFO) << "Stopping IO Context"; + _work.reset(); + _ioService.stop(); + LMS_LOG(UTILS, INFO) << "Stopped IO Context"; +} + +IOContextRunner::~IOContextRunner() +{ + + stop(); + + for (std::thread& t : _threads) + t.join(); + +} diff --git a/src/libs/utils/include/utils/IChildProcessManager.hpp b/src/libs/utils/include/utils/IChildProcessManager.hpp index fd9f8e08..ab4b34d3 100644 --- a/src/libs/utils/include/utils/IChildProcessManager.hpp +++ b/src/libs/utils/include/utils/IChildProcessManager.hpp @@ -20,10 +20,7 @@ #include #include -#pragma once - -#include -#include +#include #include "IChildProcess.hpp" @@ -35,6 +32,6 @@ class IChildProcessManager virtual std::unique_ptr spawnChildProcess(const std::filesystem::path& path, const IChildProcess::Args& args) = 0; }; -std::unique_ptr createChildProcessManager(); +std::unique_ptr createChildProcessManager(boost::asio::io_service& ioService); diff --git a/src/libs/utils/include/utils/IOContextRunner.hpp b/src/libs/utils/include/utils/IOContextRunner.hpp new file mode 100644 index 00000000..804a297c --- /dev/null +++ b/src/libs/utils/include/utils/IOContextRunner.hpp @@ -0,0 +1,43 @@ +/* + * Copyright (C) 2021 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#pragma once + +#include +#include +#include + +class IOContextRunner +{ + public: + IOContextRunner(boost::asio::io_service& ioService, std::size_t threadCount); + ~IOContextRunner(); + + IOContextRunner(const IOContextRunner&) = delete; + IOContextRunner(IOContextRunner&&) = delete; + IOContextRunner& operator=(const IOContextRunner&) = delete; + IOContextRunner& operator=(IOContextRunner&&) = delete; + + void stop(); + + private: + boost::asio::io_service& _ioService; + std::optional _work; + std::vector _threads; +}; diff --git a/src/lms/main.cpp b/src/lms/main.cpp index 15f04cdc..65d626d7 100644 --- a/src/lms/main.cpp +++ b/src/lms/main.cpp @@ -19,6 +19,7 @@ #include +#include #include #include @@ -38,6 +39,7 @@ #include "ui/LmsApplicationManager.hpp" #include "utils/IChildProcessManager.hpp" #include "utils/IConfig.hpp" +#include "utils/IOContextRunner.hpp" #include "utils/Service.hpp" #include "utils/String.hpp" #include "utils/WtLogger.hpp" @@ -212,9 +214,12 @@ int main(int argc, char* argv[]) wtArgv[i] = wtServerArgs[i].c_str(); } + boost::asio::io_context ioContext; // ioContext used to dispatch all the services that are out of the Wt event loop Wt::WServer server {argv[0]}; server.setServerConfiguration(wtServerArgs.size(), const_cast(&wtArgv[0])); + IOContextRunner ioContextRunner {ioContext, std::max(2, std::thread::hardware_concurrency())}; + // Initializing a connection pool to the database that will be shared along services Database::Db database {config->getPath("working-dir") / "lms.db"}; { @@ -226,7 +231,7 @@ int main(int argc, char* argv[]) UserInterface::LmsApplicationManager appManager; // Service initialization order is important (reverse-order for deinit) - Service childProcessManagerService {createChildProcessManager()}; + Service childProcessManagerService {createChildProcessManager(ioContext)}; Service authTokenService; Service authPasswordService; @@ -251,7 +256,7 @@ int main(int argc, char* argv[]) config->getULong("cover-max-file-size", 10) * 1000 * 1000, config->getULong("cover-jpeg-quality", 75))}; Service recommendationEngineService {Recommendation::createEngine(database)}; - Service scannerService {Scanner::createScanner(database, *recommendationEngineService)}; + Service scannerService {Scanner::createScanner(/*ioContext,*/ database, *recommendationEngineService)}; scannerService->getEvents().scanComplete.connect([&] { @@ -260,7 +265,7 @@ int main(int argc, char* argv[]) coverArtService->flushCache(); }); - Service scrobblingService {Scrobbling::createScrobbling(database)}; + Service scrobblingService {Scrobbling::createScrobbling(ioContext, database)}; API::Subsonic::SubsonicResource subsonicResource {database}; diff --git a/src/test/database/DatabaseTest.cpp b/src/test/database/DatabaseTest.cpp index bd612339..edcfb4f5 100644 --- a/src/test/database/DatabaseTest.cpp +++ b/src/test/database/DatabaseTest.cpp @@ -481,6 +481,8 @@ testSingleTrackSingleRelease(Session& session) auto transaction {session.createUniqueTransaction()}; track.get().modify()->setRelease(release.get()); + track.get().modify()->setName("MyTrackName"); + release.get().modify()->setName("MyReleaseName"); } { @@ -498,6 +500,23 @@ testSingleTrackSingleRelease(Session& session) CHECK(track->getRelease()); CHECK(track->getRelease().id() == release.getId()); } + + { + auto transaction {session.createUniqueTransaction()}; + auto tracks {Track::getByNameAndReleaseName(session, "MyTrackName", "MyReleaseName")}; + CHECK(tracks.size() == 1); + CHECK(tracks.front().id() == track.getId()); + } + { + auto transaction {session.createUniqueTransaction()}; + auto tracks {Track::getByNameAndReleaseName(session, "MyTrackName", "MyReleaseFoo")}; + CHECK(tracks.size() == 0); + } + { + auto transaction {session.createUniqueTransaction()}; + auto tracks {Track::getByNameAndReleaseName(session, "MyTrackFoo", "MyReleaseName")}; + CHECK(tracks.size() == 0); + } } {