Made database ID manipulations safer

This commit is contained in:
emeric
2021-09-20 23:53:38 +02:00
parent 598f01069e
commit 441aed622c
138 changed files with 2164 additions and 2054 deletions
+3
View File
@@ -6,6 +6,9 @@ set(CMAKE_MODULE_PATH ${CMAKE_MODULE_PATH} ${CMAKE_SOURCE_DIR}/cmake/modules/)
set(CMAKE_CXX_STANDARD 17) set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED True) set(CMAKE_CXX_STANDARD_REQUIRED True)
if (UNIX)
set(CMAKE_SHARED_LINKER_FLAGS "-Wl,--no-undefined")
endif ()
include(CTest) include(CTest)
find_package(PkgConfig REQUIRED) find_package(PkgConfig REQUIRED)
+3 -3
View File
@@ -25,7 +25,7 @@
namespace Auth namespace Auth
{ {
Database::IdType Database::UserId
AuthServiceBase::getOrCreateUser(Database::Session& session, std::string_view loginName) AuthServiceBase::getOrCreateUser(Database::Session& session, std::string_view loginName)
{ {
auto transaction {session.createUniqueTransaction()}; auto transaction {session.createUniqueTransaction()};
@@ -41,11 +41,11 @@ namespace Auth
user.modify()->setType(type); user.modify()->setType(type);
} }
return user.id(); return user->getId();
} }
void void
AuthServiceBase::onUserAuthenticated(Database::Session& session, Database::IdType userId) AuthServiceBase::onUserAuthenticated(Database::Session& session, Database::UserId userId)
{ {
auto transaction {session.createUniqueTransaction()}; auto transaction {session.createUniqueTransaction()};
Database::User::pointer user {Database::User::getById(session, userId)}; Database::User::pointer user {Database::User::getById(session, userId)};
+2 -2
View File
@@ -32,7 +32,7 @@ namespace Auth
class AuthServiceBase class AuthServiceBase
{ {
protected: protected:
Database::IdType getOrCreateUser(Database::Session& session, std::string_view loginName); Database::UserId getOrCreateUser(Database::Session& session, std::string_view loginName);
void onUserAuthenticated(Database::Session& session, Database::IdType userId); void onUserAuthenticated(Database::Session& session, Database::UserId userId);
}; };
} }
+3 -3
View File
@@ -45,7 +45,7 @@ namespace Auth
} }
std::string std::string
AuthTokenService::createAuthToken(Database::Session& session, Database::IdType userId, const Wt::WDateTime& expiry) AuthTokenService::createAuthToken(Database::Session& session, Database::UserId userId, const Wt::WDateTime& expiry)
{ {
const std::string secret {Wt::WRandom::generateId(32)}; const std::string secret {Wt::WRandom::generateId(32)};
const std::string secretHash {sha1Function.compute(secret, {})}; const std::string secretHash {sha1Function.compute(secret, {})};
@@ -86,7 +86,7 @@ namespace Auth
LMS_LOG(UI, DEBUG) << "Found auth token for user '" << authToken->getUser()->getLoginName() << "'!"; LMS_LOG(UI, DEBUG) << "Found auth token for user '" << authToken->getUser()->getLoginName() << "'!";
AuthTokenService::AuthTokenProcessResult::AuthTokenInfo res {authToken->getUser().id(), authToken->getExpiry()}; AuthTokenService::AuthTokenProcessResult::AuthTokenInfo res {authToken->getUser()->getId(), authToken->getExpiry()};
authToken.remove(); authToken.remove();
return res; return res;
@@ -123,7 +123,7 @@ namespace Auth
} }
void void
AuthTokenService::clearAuthTokens(Database::Session& session, Database::IdType userId) AuthTokenService::clearAuthTokens(Database::Session& session, Database::UserId userId)
{ {
auto transaction {session.createUniqueTransaction()}; auto transaction {session.createUniqueTransaction()};
+2 -2
View File
@@ -44,8 +44,8 @@ namespace Auth
private: private:
AuthTokenProcessResult processAuthToken(Database::Session& session, const boost::asio::ip::address& clientAddress, std::string_view tokenValue) override; AuthTokenProcessResult processAuthToken(Database::Session& session, const boost::asio::ip::address& clientAddress, std::string_view tokenValue) override;
std::string createAuthToken(Database::Session& session, Database::IdType userId, const Wt::WDateTime& expiry) override; std::string createAuthToken(Database::Session& session, Database::UserId userId, const Wt::WDateTime& expiry) override;
void clearAuthTokens(Database::Session& session, Database::IdType userId) override; void clearAuthTokens(Database::Session& session, Database::UserId userId) override;
std::shared_mutex _mutex; std::shared_mutex _mutex;
LoginThrottler _loginThrottler; LoginThrottler _loginThrottler;
+1 -1
View File
@@ -84,7 +84,7 @@ namespace Auth
{ {
_loginThrottler.onGoodClientAttempt(clientAddress); _loginThrottler.onGoodClientAttempt(clientAddress);
const Database::IdType userId {getOrCreateUser(session, loginName)}; const Database::UserId userId {getOrCreateUser(session, loginName)};
onUserAuthenticated(session, userId); onUserAuthenticated(session, userId);
return {CheckResult::State::Granted, userId}; return {CheckResult::State::Granted, userId};
} }
@@ -43,7 +43,7 @@ namespace Auth
LMS_LOG(AUTH, DEBUG) << "Extracted login name = '" << loginName << "' from HTTP header"; LMS_LOG(AUTH, DEBUG) << "Extracted login name = '" << loginName << "' from HTTP header";
const Database::IdType userId {getOrCreateUser(session, loginName)}; const Database::UserId userId {getOrCreateUser(session, loginName)};
onUserAuthenticated(session, userId); onUserAuthenticated(session, userId);
return {CheckResult::State::Granted, userId}; return {CheckResult::State::Granted, userId};
} }
@@ -57,7 +57,7 @@ namespace Auth
LMS_LOG(AUTH, DEBUG) << "Extracted login name = '" << loginName << "' from HTTP header"; LMS_LOG(AUTH, DEBUG) << "Extracted login name = '" << loginName << "' from HTTP header";
const Database::IdType userId {getOrCreateUser(session, loginName)}; const Database::UserId userId {getOrCreateUser(session, loginName)};
onUserAuthenticated(session, userId); onUserAuthenticated(session, userId);
return {CheckResult::State::Granted, userId}; return {CheckResult::State::Granted, userId};
} }
@@ -96,13 +96,13 @@ namespace Auth
} }
void void
InternalPasswordService::setPassword(Database::Session& session, Database::IdType userId, std::string_view newPassword) InternalPasswordService::setPassword(Database::Session& session, Database::UserId userId, std::string_view newPassword)
{ {
const Database::User::PasswordHash passwordHash {hashPassword(newPassword)}; const Database::User::PasswordHash passwordHash {hashPassword(newPassword)};
auto transaction {session.createUniqueTransaction()}; auto transaction {session.createUniqueTransaction()};
const Database::User::pointer user {Database::User::getById(session, userId)}; Database::User::pointer user {Database::User::getById(session, userId)};
if (!user) if (!user)
throw Exception {"User not found!"}; throw Exception {"User not found!"};
@@ -42,7 +42,7 @@ namespace Auth
bool canSetPasswords() const override; bool canSetPasswords() const override;
PasswordAcceptabilityResult checkPasswordAcceptability(std::string_view loginName, const PasswordValidationContext& context) const override; PasswordAcceptabilityResult checkPasswordAcceptability(std::string_view loginName, const PasswordValidationContext& context) const override;
void setPassword(Database::Session& session, Database::IdType userId, std::string_view newPassword) override; void setPassword(Database::Session& session, Database::UserId userId, std::string_view newPassword) override;
Database::User::PasswordHash hashPassword(std::string_view password) const; Database::User::PasswordHash hashPassword(std::string_view password) const;
void hashRandomPassword() const; void hashRandomPassword() const;
@@ -193,7 +193,7 @@ namespace Auth
} }
void void
PAMPasswordService::setPassword(Database::Session&, Database::IdType, std::string_view) PAMPasswordService::setPassword(Database::Session&, Database::UserId, std::string_view)
{ {
throw NotImplementedException {}; throw NotImplementedException {};
} }
@@ -38,7 +38,7 @@ namespace Auth
bool canSetPasswords() const override; bool canSetPasswords() const override;
PasswordAcceptabilityResult checkPasswordAcceptability(std::string_view loginName, const PasswordValidationContext& context) const override; PasswordAcceptabilityResult checkPasswordAcceptability(std::string_view loginName, const PasswordValidationContext& context) const override;
void setPassword(Database::Session& session, void setPassword(Database::Session& session,
Database::IdType userId, Database::UserId userId,
std::string_view newPassword) override; std::string_view newPassword) override;
}; };
} }
@@ -54,7 +54,7 @@ namespace Auth
struct AuthTokenInfo struct AuthTokenInfo
{ {
Database::IdType userId; Database::UserId userId;
Wt::WDateTime expiry; Wt::WDateTime expiry;
}; };
@@ -66,8 +66,8 @@ namespace Auth
virtual AuthTokenProcessResult processAuthToken(Database::Session& session, const boost::asio::ip::address& clientAddress, std::string_view tokenValue) = 0; virtual AuthTokenProcessResult processAuthToken(Database::Session& session, const boost::asio::ip::address& clientAddress, std::string_view tokenValue) = 0;
// Returns a one time token // Returns a one time token
virtual std::string createAuthToken(Database::Session& session, Database::IdType userid, const Wt::WDateTime& expiry) = 0; virtual std::string createAuthToken(Database::Session& session, Database::UserId userid, const Wt::WDateTime& expiry) = 0;
virtual void clearAuthTokens(Database::Session& session, Database::IdType userid) = 0; virtual void clearAuthTokens(Database::Session& session, Database::UserId userid) = 0;
}; };
std::unique_ptr<IAuthTokenService> createAuthTokenService(std::size_t maxThrottlerEntryCount); std::unique_ptr<IAuthTokenService> createAuthTokenService(std::size_t maxThrottlerEntryCount);
+1 -1
View File
@@ -57,7 +57,7 @@ namespace Auth
}; };
State state {State::Denied}; State state {State::Denied};
std::optional<Database::IdType> userId {}; std::optional<Database::UserId> userId {};
}; };
virtual CheckResult processEnv(Database::Session& session, const Wt::WEnvironment& env) = 0; virtual CheckResult processEnv(Database::Session& session, const Wt::WEnvironment& env) = 0;
@@ -53,7 +53,7 @@ namespace Auth
Throttled, Throttled,
}; };
State state {State::Denied}; State state {State::Denied};
std::optional<Database::IdType> userId {}; std::optional<Database::UserId> userId {};
std::optional<Wt::WDateTime> expiry {}; std::optional<Wt::WDateTime> expiry {};
}; };
virtual CheckResult checkUserPassword(Database::Session& session, virtual CheckResult checkUserPassword(Database::Session& session,
@@ -70,7 +70,7 @@ namespace Auth
MustMatchLoginName, MustMatchLoginName,
}; };
virtual PasswordAcceptabilityResult checkPasswordAcceptability(std::string_view password, const PasswordValidationContext& context) const = 0; virtual PasswordAcceptabilityResult checkPasswordAcceptability(std::string_view password, const PasswordValidationContext& context) const = 0;
virtual void setPassword(Database::Session& session, Database::IdType userId, std::string_view newPassword) = 0; virtual void setPassword(Database::Session& session, Database::UserId userId, std::string_view newPassword) = 0;
}; };
std::unique_ptr<IPasswordService> createPasswordService(std::string_view authPasswordBackend, std::size_t maxThrottlerEntryCount, IAuthTokenService& authTokenService); std::unique_ptr<IPasswordService> createPasswordService(std::string_view authPasswordBackend, std::size_t maxThrottlerEntryCount, IAuthTokenService& authTokenService);
+11 -15
View File
@@ -45,11 +45,11 @@ namespace
bool hasCover {}; bool hasCover {};
bool isMultiDisc {}; bool isMultiDisc {};
std::filesystem::path trackPath; std::filesystem::path trackPath;
std::optional<Database::IdType> releaseId; std::optional<Database::ReleaseId> releaseId;
}; };
std::optional<TrackInfo> std::optional<TrackInfo>
getTrackInfo(Database::Session& dbSession, Database::IdType trackId) getTrackInfo(Database::Session& dbSession, Database::TrackId trackId)
{ {
std::optional<TrackInfo> res; std::optional<TrackInfo> res;
@@ -66,7 +66,7 @@ namespace
if (const Database::Release::pointer& release {track->getRelease()}) if (const Database::Release::pointer& release {track->getRelease()})
{ {
res->releaseId = release.id(); res->releaseId = release->getId();
if (release->getTotalDisc() > 1) if (release->getTotalDisc() > 1)
res->isMultiDisc = true; res->isMultiDisc = true;
} }
@@ -75,7 +75,6 @@ namespace
} }
} }
namespace CoverArt { namespace CoverArt {
static static
@@ -101,7 +100,7 @@ Grabber::Grabber(const std::filesystem::path& execPath,
: _defaultCoverPath {defaultCoverPath} : _defaultCoverPath {defaultCoverPath}
, _maxCacheSize {maxCacheSize} , _maxCacheSize {maxCacheSize}
, _maxFileSize {maxFileSize} , _maxFileSize {maxFileSize}
, _jpegQuality {clamp<unsigned>(jpegQuality, 1, 100)} , _jpegQuality {Utils::clamp<unsigned>(jpegQuality, 1, 100)}
{ {
LMS_LOG(COVER, INFO) << "Default cover path = '" << _defaultCoverPath.string() << "'"; LMS_LOG(COVER, INFO) << "Default cover path = '" << _defaultCoverPath.string() << "'";
LMS_LOG(COVER, INFO) << "Max cache size = " << _maxCacheSize; LMS_LOG(COVER, INFO) << "Max cache size = " << _maxCacheSize;
@@ -314,20 +313,17 @@ Grabber::getFromTrack(const std::filesystem::path& p, ImageSize width) const
} }
std::shared_ptr<IEncodedImage> std::shared_ptr<IEncodedImage>
Grabber::getFromTrack(Database::Session& dbSession, Database::IdType trackId, ImageSize width) Grabber::getFromTrack(Database::Session& dbSession, Database::TrackId trackId, ImageSize width)
{ {
return getFromTrack(dbSession, trackId, width, true /* allow release fallback*/); return getFromTrack(dbSession, trackId, width, true /* allow release fallback*/);
} }
std::shared_ptr<IEncodedImage> std::shared_ptr<IEncodedImage>
Grabber::getFromTrack(Database::Session& dbSession, Database::IdType trackId, ImageSize width, bool allowReleaseFallback) Grabber::getFromTrack(Database::Session& dbSession, Database::TrackId trackId, ImageSize width, bool allowReleaseFallback)
{ {
using namespace Database; using namespace Database;
const CacheEntryDesc cacheEntryDesc {CacheEntryDesc::Type::Track, trackId, width}; const CacheEntryDesc cacheEntryDesc {trackId, width};
std::shared_ptr<IEncodedImage> cover {loadFromCache(cacheEntryDesc)}; std::shared_ptr<IEncodedImage> cover {loadFromCache(cacheEntryDesc)};
if (cover) if (cover)
@@ -361,9 +357,9 @@ Grabber::getFromTrack(Database::Session& dbSession, Database::IdType trackId, Im
} }
std::shared_ptr<IEncodedImage> std::shared_ptr<IEncodedImage>
Grabber::getFromRelease(Database::Session& session, Database::IdType releaseId, ImageSize width) Grabber::getFromRelease(Database::Session& session, Database::ReleaseId releaseId, ImageSize width)
{ {
const CacheEntryDesc cacheEntryDesc {CacheEntryDesc::Type::Release, releaseId, width}; const CacheEntryDesc cacheEntryDesc {releaseId, width};
std::shared_ptr<IEncodedImage> cover {loadFromCache(cacheEntryDesc)}; std::shared_ptr<IEncodedImage> cover {loadFromCache(cacheEntryDesc)};
if (cover) if (cover)
@@ -371,7 +367,7 @@ Grabber::getFromRelease(Database::Session& session, Database::IdType releaseId,
struct ReleaseInfo struct ReleaseInfo
{ {
Database::IdType firstTrackId; Database::TrackId firstTrackId;
std::filesystem::path releaseDirectory; std::filesystem::path releaseDirectory;
}; };
@@ -386,7 +382,7 @@ Grabber::getFromRelease(Database::Session& session, Database::IdType releaseId,
if (const auto firstTrack {release->getFirstTrack()}) if (const auto firstTrack {release->getFirstTrack()})
{ {
res = ReleaseInfo {}; res = ReleaseInfo {};
res->firstTrackId = firstTrack.id(); res->firstTrackId = firstTrack->getId();
res->releaseDirectory = firstTrack->getPath().parent_path(); res->releaseDirectory = firstTrack->getPath().parent_path();
} }
} }
+12 -15
View File
@@ -26,6 +26,7 @@
#include <shared_mutex> #include <shared_mutex>
#include <string_view> #include <string_view>
#include <unordered_map> #include <unordered_map>
#include <variant>
#include <vector> #include <vector>
#include "cover/ICoverArtGrabber.hpp" #include "cover/ICoverArtGrabber.hpp"
@@ -46,20 +47,12 @@ namespace CoverArt
{ {
struct CacheEntryDesc struct CacheEntryDesc
{ {
enum class Type std::variant<Database::TrackId, Database::ReleaseId> id;
{
Track,
Release,
};
Type type;
Database::IdType id;
std::size_t size; std::size_t size;
bool operator==(const CacheEntryDesc& other) const bool operator==(const CacheEntryDesc& other) const
{ {
return type == other.type return id == other.id
&& id == other.id
&& size == other.size; && size == other.size;
} }
}; };
@@ -75,8 +68,12 @@ namespace std
public: public:
size_t operator()(const CoverArt::CacheEntryDesc& e) const size_t operator()(const CoverArt::CacheEntryDesc& e) const
{ {
size_t h = std::hash<int>()(static_cast<int>(e.type)); size_t h {};
h ^= std::hash<Database::IdType>()(e.id) << 1; std::visit([&](auto id)
{
using IdType = std::decay_t<decltype(id)>;
h ^= std::hash<IdType>()(id);
}, e.id);
h ^= std::hash<std::size_t>()(e.size) << 1; h ^= std::hash<std::size_t>()(e.size) << 1;
return h; return h;
} }
@@ -101,11 +98,11 @@ namespace CoverArt
Grabber& operator=(Grabber&&) = delete; Grabber& operator=(Grabber&&) = delete;
private: private:
std::shared_ptr<IEncodedImage> getFromTrack(Database::Session& dbSession, Database::IdType trackId, ImageSize width) override; std::shared_ptr<IEncodedImage> getFromTrack(Database::Session& dbSession, Database::TrackId trackId, ImageSize width) override;
std::shared_ptr<IEncodedImage> getFromRelease(Database::Session& dbSession, Database::IdType releaseId, ImageSize width) override; std::shared_ptr<IEncodedImage> getFromRelease(Database::Session& dbSession, Database::ReleaseId releaseId, ImageSize width) override;
void flushCache() override; void flushCache() override;
std::shared_ptr<IEncodedImage> getFromTrack(Database::Session& dbSession, Database::IdType trackId, ImageSize width, bool allowReleaseFallback); std::shared_ptr<IEncodedImage> getFromTrack(Database::Session& dbSession, Database::TrackId trackId, ImageSize width, bool allowReleaseFallback);
std::unique_ptr<IEncodedImage> getFromAvMediaFile(const Av::IAudioFile& input, ImageSize width) const; std::unique_ptr<IEncodedImage> getFromAvMediaFile(const Av::IAudioFile& input, ImageSize width) const;
std::unique_ptr<IEncodedImage> getFromCoverFile(const std::filesystem::path& p, ImageSize width) const; std::unique_ptr<IEncodedImage> getFromCoverFile(const std::filesystem::path& p, ImageSize width) const;
@@ -37,8 +37,8 @@ namespace CoverArt
public: public:
virtual ~IGrabber() = default; virtual ~IGrabber() = default;
virtual std::shared_ptr<IEncodedImage> getFromTrack(Database::Session& dbSession, Database::IdType trackId, ImageSize width) = 0; virtual std::shared_ptr<IEncodedImage> getFromTrack(Database::Session& dbSession, Database::TrackId trackId, ImageSize width) = 0;
virtual std::shared_ptr<IEncodedImage> getFromRelease(Database::Session& dbSession, Database::IdType releaseId, ImageSize width) = 0; virtual std::shared_ptr<IEncodedImage> getFromRelease(Database::Session& dbSession, Database::ReleaseId releaseId, ImageSize width) = 0;
virtual void flushCache() = 0; virtual void flushCache() = 0;
}; };
+74 -92
View File
@@ -28,6 +28,7 @@
#include "utils/Logger.hpp" #include "utils/Logger.hpp"
#include "SqlQuery.hpp" #include "SqlQuery.hpp"
#include "Utils.hpp" #include "Utils.hpp"
#include "Traits.hpp"
namespace Database namespace Database
{ {
@@ -37,7 +38,6 @@ Artist::Artist(const std::string& name, const std::optional<UUID>& MBID)
_sortName {_name}, _sortName {_name},
_MBID {MBID ? MBID->getAsString() : ""} _MBID {MBID ? MBID->getAsString() : ""}
{ {
} }
std::vector<Artist::pointer> std::vector<Artist::pointer>
@@ -45,7 +45,7 @@ Artist::getByName(Session& session, const std::string& name)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<Artist::pointer> res = session.getDboSession().find<Artist>() Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> res = session.getDboSession().find<Artist>()
.where("name = ?").bind(std::string {name, 0, _maxNameLength}) .where("name = ?").bind(std::string {name, 0, _maxNameLength})
.orderBy("LENGTH(mbid) DESC"); // put mbid entries first .orderBy("LENGTH(mbid) DESC"); // put mbid entries first
@@ -56,14 +56,14 @@ Artist::pointer
Artist::getByMBID(Session& session, const UUID& mbid) Artist::getByMBID(Session& session, const UUID& mbid)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Artist>().where("mbid = ?").bind(std::string {mbid.getAsString()}); return session.getDboSession().find<Artist>().where("mbid = ?").bind(std::string {mbid.getAsString()}).resultValue();
} }
Artist::pointer Artist::pointer
Artist::getById(Session& session, IdType id) Artist::getById(Session& session, ArtistId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Artist>().where("id = ?").bind(id); return session.getDboSession().find<Artist>().where("id = ?").bind(id).resultValue();
} }
Artist::pointer Artist::pointer
@@ -82,7 +82,7 @@ static
Wt::Dbo::Query<T> Wt::Dbo::Query<T>
createQuery(Session& session, createQuery(Session& session,
const std::string& queryStr, const std::string& queryStr,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
const std::vector<std::string_view>& keywords, const std::vector<std::string_view>& keywords,
std::optional<TrackArtistLinkType> linkType) std::optional<TrackArtistLinkType> linkType)
{ {
@@ -125,7 +125,7 @@ createQuery(Session& session,
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id"; " INNER JOIN track_cluster t_c ON t_c.track_id = t.id";
WhereClause clusterClause; WhereClause clusterClause;
for (const IdType clusterId : clusterIds) for (const ClusterId clusterId : clusterIds)
{ {
clusterClause.Or(WhereClause("c.id = ?")); clusterClause.Or(WhereClause("c.id = ?"));
query.bind(clusterId); query.bind(clusterId);
@@ -145,7 +145,7 @@ Artist::getAll(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().find<Artist>(); Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> res = session.getDboSession().find<Artist>();
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
@@ -167,7 +167,7 @@ Artist::getAll(Session& session, SortMethod sortMethod)
break; break;
} }
Wt::Dbo::collection<pointer> res = query; Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> res = query;
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
@@ -177,7 +177,7 @@ Artist::getAll(Session& session, SortMethod sortMethod, std::optional<Range> ran
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Artist::pointer>(session, "SELECT a FROM Artist a", {}, {}, std::nullopt)}; auto query {createQuery<Wt::Dbo::ptr<Artist>>(session, "SELECT a FROM Artist a", {}, {}, std::nullopt)};
switch (sortMethod) switch (sortMethod)
{ {
@@ -191,11 +191,11 @@ Artist::getAll(Session& session, SortMethod sortMethod, std::optional<Range> ran
break; break;
} }
Wt::Dbo::collection<Artist::pointer> collection = query Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> collection = query
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1);
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<Artist::pointer> res (collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -207,27 +207,27 @@ Artist::getAll(Session& session, SortMethod sortMethod, std::optional<Range> ran
return res; return res;
} }
std::vector<IdType> std::vector<ArtistId>
Artist::getAllIds(Session& session) Artist::getAllIds(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType>("SELECT id FROM artist"); Wt::Dbo::collection<ArtistId> res = session.getDboSession().query<ArtistId>("SELECT id FROM artist");
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ArtistId>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<ArtistId>
Artist::getAllIdsRandom(Session& session, const std::set<IdType>& clusters, std::optional<TrackArtistLinkType> linkType, std::optional<std::size_t> size) Artist::getAllIdsRandom(Session& session, const std::vector<ClusterId>& clusters, std::optional<TrackArtistLinkType> linkType, std::optional<std::size_t> size)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<IdType>(session, "SELECT DISTINCT a.id from artist a", clusters, {}, linkType)}; auto query {createQuery<ArtistId>(session, "SELECT DISTINCT a.id from artist a", clusters, {}, linkType)};
Wt::Dbo::collection<IdType> res = query Wt::Dbo::collection<ArtistId> res = query
.orderBy("RANDOM()") .orderBy("RANDOM()")
.limit(size ? static_cast<int>(*size) : -1); .limit(size ? static_cast<int>(*size) : -1);
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ArtistId>(res.begin(), res.end());
} }
@@ -240,22 +240,22 @@ Artist::getAllOrphans(Session& session)
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<ArtistId>
Artist::getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit) Artist::getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType> Wt::Dbo::collection<ArtistId> res = session.getDboSession().query<ArtistId>
("SELECT DISTINCT a.id FROM artist a" ("SELECT DISTINCT a.id FROM artist a"
" INNER JOIN track t ON t.id = t_a_l.track_id INNER JOIN track_artist_link t_a_l ON t_a_l.artist_id = a.id" " INNER JOIN track t ON t.id = t_a_l.track_id INNER JOIN track_artist_link t_a_l ON t_a_l.artist_id = a.id"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id") " INNER JOIN track_cluster t_c ON t_c.track_id = t.id")
.limit(limit ? static_cast<int>(*limit) : -1); .limit(limit ? static_cast<int>(*limit) : -1);
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ArtistId>(res.begin(), res.end());
} }
std::vector<Artist::pointer> std::vector<Artist::pointer>
Artist::getByClusters(Session& session, const std::set<IdType>& clusters, SortMethod sortMethod) Artist::getByClusters(Session& session, const std::vector<ClusterId>& clusters, SortMethod sortMethod)
{ {
assert(!clusters.empty()); assert(!clusters.empty());
@@ -266,7 +266,7 @@ Artist::getByClusters(Session& session, const std::set<IdType>& clusters, SortMe
std::vector<Artist::pointer> std::vector<Artist::pointer>
Artist::getByFilter(Session& session, Artist::getByFilter(Session& session,
const std::set<IdType>& clusters, const std::vector<ClusterId>& clusters,
const std::vector<std::string_view>& keywords, const std::vector<std::string_view>& keywords,
std::optional<TrackArtistLinkType> linkType, std::optional<TrackArtistLinkType> linkType,
SortMethod sortMethod, SortMethod sortMethod,
@@ -275,7 +275,7 @@ Artist::getByFilter(Session& session,
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Artist::pointer>(session, "SELECT DISTINCT a from artist a", clusters, keywords, linkType)}; auto query {createQuery<Wt::Dbo::ptr<Artist>>(session, "SELECT DISTINCT a from artist a", clusters, keywords, linkType)};
switch (sortMethod) switch (sortMethod)
{ {
case Artist::SortMethod::None: case Artist::SortMethod::None:
@@ -288,11 +288,11 @@ Artist::getByFilter(Session& session,
break; break;
} }
Wt::Dbo::collection<Artist::pointer> collection = query Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> collection = query
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1);
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res (collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
@@ -308,23 +308,23 @@ Artist::getByFilter(Session& session,
std::vector<Artist::pointer> std::vector<Artist::pointer>
Artist::getLastWritten(Session& session, Artist::getLastWritten(Session& session,
std::optional<Wt::WDateTime> after, std::optional<Wt::WDateTime> after,
const std::set<IdType>& clusters, const std::vector<ClusterId>& clusters,
std::optional<TrackArtistLinkType> linkType, std::optional<TrackArtistLinkType> linkType,
std::optional<Range> range, bool& moreResults) std::optional<Range> range, bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Artist::pointer>(session, "SELECT DISTINCT a from artist a", clusters, {}, linkType)}; auto query {createQuery<Wt::Dbo::ptr<Artist>>(session, "SELECT DISTINCT a from artist a", clusters, {}, linkType)};
if (after) if (after)
query.where("t.file_last_write > ?").bind(*after); query.where("t.file_last_write > ?").bind(*after);
Wt::Dbo::collection<Artist::pointer> collection = query Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> collection = query
.orderBy("t.file_last_write DESC") .orderBy("t.file_last_write DESC")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1);
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res (collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
@@ -340,14 +340,14 @@ Artist::getLastWritten(Session& session,
std::vector<Artist::pointer> std::vector<Artist::pointer>
Artist::getStarred(Session& session, Artist::getStarred(Session& session,
User::pointer user, User::pointer user,
const std::set<IdType>& clusters, const std::vector<ClusterId>& clusters,
std::optional<TrackArtistLinkType> linkType, std::optional<TrackArtistLinkType> linkType,
SortMethod sortMethod, SortMethod sortMethod,
std::optional<Range> range, bool& moreResults) std::optional<Range> range, bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Artist::pointer>(session, "SELECT DISTINCT a from artist a", clusters, {}, linkType)}; auto query {createQuery<Wt::Dbo::ptr<Artist>>(session, "SELECT DISTINCT a from artist a", clusters, {}, linkType)};
{ {
std::ostringstream oss; std::ostringstream oss;
@@ -355,7 +355,7 @@ Artist::getStarred(Session& session,
" INNER JOIN user_artist_starred uas ON uas.artist_id = a.id" " INNER JOIN user_artist_starred uas ON uas.artist_id = a.id"
" INNER JOIN user u ON u.id = uas.user_id WHERE u.id = ?)"; " INNER JOIN user u ON u.id = uas.user_id WHERE u.id = ?)";
query.bind(user.id()); query.bind(user->getId());
query.where(oss.str()); query.where(oss.str());
} }
@@ -371,12 +371,12 @@ Artist::getStarred(Session& session,
break; break;
} }
Wt::Dbo::collection<Artist::pointer> collection = query Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> collection = query
.groupBy("a.id") .groupBy("a.id")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1);
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res (collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
@@ -389,11 +389,9 @@ Artist::getStarred(Session& session,
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<Wt::Dbo::ptr<Release>> std::vector<Release::pointer>
Artist::getReleases(const std::set<IdType>& clusterIds) const Artist::getReleases(const std::vector<ClusterId>& clusterIds) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
WhereClause where; WhereClause where;
@@ -409,12 +407,12 @@ Artist::getReleases(const std::set<IdType>& clusterIds) const
WhereClause clusterClause; WhereClause clusterClause;
for (auto id : clusterIds) for (auto id : clusterIds)
clusterClause.Or(WhereClause("c.id = ?")).bind(std::to_string(id)); clusterClause.Or(WhereClause("c.id = ?")).bind(id.toString());
where.And(clusterClause); where.And(clusterClause);
} }
where.And(WhereClause("a.id = ?")).bind(std::to_string(id())); where.And(WhereClause("a.id = ?")).bind(getId().toString());
oss << " " << where.get(); oss << " " << where.get();
@@ -423,56 +421,48 @@ Artist::getReleases(const std::set<IdType>& clusterIds) const
oss << " ORDER BY t.year DESC, r.name COLLATE NOCASE"; oss << " ORDER BY t.year DESC, r.name COLLATE NOCASE";
Wt::Dbo::Query<Release::pointer> query = session()->query<Release::pointer>( oss.str() ); auto query {session()->query<Wt::Dbo::ptr<Release>>(oss.str())};
for (const std::string& bindArg : where.getBindArgs()) for (const std::string& bindArg : where.getBindArgs())
query.bind(bindArg); query.bind(bindArg);
Wt::Dbo::collection<Wt::Dbo::ptr<Release>> res = query; auto res {query.resultList()};
return std::vector<Release::pointer>(res.begin(), res.end());
return std::vector<Wt::Dbo::ptr<Release>>(res.begin(), res.end());
} }
std::size_t std::size_t
Artist::getReleaseCount() const Artist::getReleaseCount() const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
int res = session()->query<int>("SELECT COUNT(DISTINCT r.id) FROM release r INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id INNER JOIN track t ON t.release_id = r.id") int res = session()->query<int>("SELECT COUNT(DISTINCT r.id) FROM release r INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id INNER JOIN track t ON t.release_id = r.id")
.where("a.id = ?").bind(self()->id()); .where("a.id = ?").bind(getId());
return res; return res;
} }
std::vector<Wt::Dbo::ptr<Track>> std::vector<Track::pointer>
Artist::getTracks(std::optional<TrackArtistLinkType> linkType) const Artist::getTracks(std::optional<TrackArtistLinkType> linkType) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
auto query {session()->query<Wt::Dbo::ptr<Track>>("SELECT DISTINCT t FROM track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id") auto query {session()->query<Wt::Dbo::ptr<Track>>("SELECT DISTINCT t FROM track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id")
.where("a.id = ?").bind(self()->id()) .where("a.id = ?").bind(getId())
.orderBy("t.year DESC,t.release_id,t.disc_number,t.track_number")}; .orderBy("t.year DESC,t.release_id,t.disc_number,t.track_number")};
if (linkType) if (linkType)
query.where("t_a_l.type = ?").bind(*linkType); query.where("t_a_l.type = ?").bind(*linkType);
Wt::Dbo::collection<Wt::Dbo::ptr<Track>> tracks {query.resultList()}; auto tracks {query.resultList()};
return std::vector<Track::pointer>(tracks.begin(), tracks.end());
return std::vector<Wt::Dbo::ptr<Track>>(tracks.begin(), tracks.end());
} }
std::vector<Wt::Dbo::ptr<Track>> std::vector<Track::pointer>
Artist::getNonReleaseTracks(std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const Artist::getNonReleaseTracks(std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
auto query {session()->query<Wt::Dbo::ptr<Track>>("SELECT t FROM track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id") auto query {session()->query<Wt::Dbo::ptr<Track>>("SELECT t FROM track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id")
.where("a.id = ?").bind(self()->id()) .where("a.id = ?").bind(getId())
.where("t.release_id is NULL") .where("t.release_id is NULL")
.orderBy("t.name") .orderBy("t.name")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
@@ -481,9 +471,8 @@ Artist::getNonReleaseTracks(std::optional<TrackArtistLinkType> linkType, std::op
if (linkType) if (linkType)
query.where("t_a_l.type = ?").bind(*linkType); query.where("t_a_l.type = ?").bind(*linkType);
Wt::Dbo::collection<Track::pointer> tracks {query.resultList()}; Wt::Dbo::collection<Wt::Dbo::ptr<Track>> tracks {query.resultList()};
std::vector<Track::pointer> res(tracks.begin(), tracks.end());
auto res {std::vector<Track::pointer>(tracks.begin(), tracks.end())};
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -499,37 +488,32 @@ bool
Artist::hasNonReleaseTracks(std::optional<TrackArtistLinkType> linkType) const Artist::hasNonReleaseTracks(std::optional<TrackArtistLinkType> linkType) const
{ {
auto query {session()->query<Wt::Dbo::ptr<Track>>("SELECT t FROM track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id") auto query {session()->query<Wt::Dbo::ptr<Track>>("SELECT t FROM track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id")
.where("a.id = ?").bind(self()->id()) .where("a.id = ?").bind(getId())
.where("t.release_id is NULL") .where("t.release_id is NULL")
.orderBy("t.name")}; .orderBy("t.name")};
if (linkType) if (linkType)
query.where("t_a_l.type = ?").bind(*linkType); query.where("t_a_l.type = ?").bind(*linkType);
Wt::Dbo::collection<Track::pointer> tracks {query.resultList()}; return !query.resultList().empty();
return !tracks.empty();
} }
std::vector<Wt::Dbo::ptr<Track>> std::vector<Track::pointer>
Artist::getRandomTracks(std::optional<std::size_t> count) const Artist::getRandomTracks(std::optional<std::size_t> count) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
Wt::Dbo::collection<Wt::Dbo::ptr<Track>> tracks {session()->query<Wt::Dbo::ptr<Track>>("SELECT t from track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id") Wt::Dbo::collection<Wt::Dbo::ptr<Track>> tracks {session()->query<Wt::Dbo::ptr<Track>>("SELECT t from track t INNER JOIN artist a ON a.id = t_a_l.artist_id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id")
.where("a.id = ?").bind(self()->id()) .where("a.id = ?").bind(getId())
.orderBy("RANDOM()") .orderBy("RANDOM()")
.limit(count ? static_cast<int>(*count) : -1)}; .limit(count ? static_cast<int>(*count) : -1)};
return std::vector<Wt::Dbo::ptr<Track>>(tracks.begin(), tracks.end()); return std::vector<Track::pointer>(tracks.begin(), tracks.end());
} }
std::vector<Wt::Dbo::ptr<Artist>> std::vector<Artist::pointer>
Artist::getSimilarArtists(EnumSet<TrackArtistLinkType> artistLinkTypes, std::optional<Range> range) const Artist::getSimilarArtists(EnumSet<TrackArtistLinkType> artistLinkTypes, std::optional<Range> range) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
std::ostringstream oss; std::ostringstream oss;
@@ -563,9 +547,9 @@ Artist::getSimilarArtists(EnumSet<TrackArtistLinkType> artistLinkTypes, std::opt
oss << ")"; oss << ")";
} }
Wt::Dbo::Query<pointer> query {session()->query<pointer>(oss.str()) Wt::Dbo::Query<Wt::Dbo::ptr<Artist>> query {session()->query<Wt::Dbo::ptr<Artist>>(oss.str())
.bind(self()->id()) .bind(getId())
.bind(self()->id()) .bind(getId())
.groupBy("a.id") .groupBy("a.id")
.orderBy("COUNT(*) DESC, RANDOM()") .orderBy("COUNT(*) DESC, RANDOM()")
.limit(range ? static_cast<int>(range->limit) : -1) .limit(range ? static_cast<int>(range->limit) : -1)
@@ -574,15 +558,13 @@ Artist::getSimilarArtists(EnumSet<TrackArtistLinkType> artistLinkTypes, std::opt
for (TrackArtistLinkType type : artistLinkTypes) for (TrackArtistLinkType type : artistLinkTypes)
query.bind(type); query.bind(type);
Wt::Dbo::collection<pointer> res = query; Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> res {query.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<std::vector<Wt::Dbo::ptr<Cluster>>> std::vector<std::vector<Cluster::pointer>>
Artist::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::size_t size) const Artist::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::size_t size) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
WhereClause where; WhereClause where;
@@ -590,34 +572,34 @@ Artist::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::si
std::ostringstream oss; std::ostringstream oss;
oss << "SELECT c FROM cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN cluster_type c_type ON c.cluster_type_id = c_type.id INNER JOIN artist a ON t_a_l.artist_id = a.id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id"; oss << "SELECT c FROM cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN cluster_type c_type ON c.cluster_type_id = c_type.id INNER JOIN artist a ON t_a_l.artist_id = a.id INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id";
where.And(WhereClause("a.id = ?")).bind(std::to_string(self()->id())); where.And(WhereClause("a.id = ?")).bind(getId().toString());
{ {
WhereClause clusterClause; WhereClause clusterClause;
for (auto clusterType : clusterTypes) for (auto clusterType : clusterTypes)
clusterClause.Or(WhereClause("c_type.id = ?")).bind(std::to_string(clusterType.id())); clusterClause.Or(WhereClause("c_type.id = ?")).bind(clusterType->getId().toString());
where.And(clusterClause); where.And(clusterClause);
} }
oss << " " << where.get(); oss << " " << where.get();
oss << "GROUP BY c.id ORDER BY COUNT(DISTINCT c.id) DESC"; oss << "GROUP BY c.id ORDER BY COUNT(DISTINCT c.id) DESC";
Wt::Dbo::Query<Cluster::pointer> query = session()->query<Cluster::pointer>( oss.str() ); Wt::Dbo::Query<Wt::Dbo::ptr<Cluster>> query = session()->query<Wt::Dbo::ptr<Cluster>>( oss.str() );
for (const std::string& bindArg : where.getBindArgs()) for (const std::string& bindArg : where.getBindArgs())
query.bind(bindArg); query.bind(bindArg);
Wt::Dbo::collection<Cluster::pointer> queryRes = query; Wt::Dbo::collection<Wt::Dbo::ptr<Cluster>> queryRes = query;
std::map<IdType, std::vector<Cluster::pointer>> clusters; std::map<ClusterTypeId, std::vector<Cluster::pointer>> clustersByType;
for (auto cluster : queryRes) for (const Cluster::pointer& cluster : queryRes)
{ {
if (clusters[cluster->getType().id()].size() < size) if (clustersByType[cluster->getType()->getId()].size() < size)
clusters[cluster->getType().id()].push_back(cluster); clustersByType[cluster->getType()->getId()].push_back(cluster);
} }
std::vector<std::vector<Cluster::pointer>> res; std::vector<std::vector<Cluster::pointer>> res;
for (auto cluster_list : clusters) for (const auto& [clusterTypeId, clusters] : clustersByType)
res.push_back(cluster_list.second); res.push_back(clusters);
return res; return res;
} }
+35 -47
View File
@@ -25,21 +25,18 @@
#include "database/Session.hpp" #include "database/Session.hpp"
#include "database/Track.hpp" #include "database/Track.hpp"
#include "SqlQuery.hpp" #include "SqlQuery.hpp"
#include "Traits.hpp"
namespace Database { namespace Database {
Cluster::Cluster() Cluster::Cluster(ObjectPtr<ClusterType> type, std::string_view name)
{ : _name {std::string {name, 0, _maxNameLength}},
} _clusterType {getDboPtr(type)}
Cluster::Cluster(Wt::Dbo::ptr<ClusterType> type, std::string_view name)
: _name(std::string {name, 0, _maxNameLength}),
_clusterType {type}
{ {
} }
Cluster::pointer Cluster::pointer
Cluster::create(Session& session, Wt::Dbo::ptr<ClusterType> type, std::string_view name) Cluster::create(Session& session, ObjectPtr<ClusterType> type, std::string_view name)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
@@ -54,8 +51,7 @@ Cluster::getAll(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<Cluster::pointer> res {session.getDboSession().find<Cluster>()}; Wt::Dbo::collection<Wt::Dbo::ptr<Cluster>> res {session.getDboSession().find<Cluster>()};
return std::vector<Cluster::pointer>(res.begin(), res.end()); return std::vector<Cluster::pointer>(res.begin(), res.end());
} }
@@ -63,67 +59,61 @@ std::vector<Cluster::pointer>
Cluster::getAllOrphans(Session& session) Cluster::getAllOrphans(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<Cluster::pointer> res {session.getDboSession().query<Cluster::pointer>("SELECT DISTINCT c FROM cluster c WHERE NOT EXISTS(SELECT 1 FROM track_cluster t_c WHERE t_c.cluster_id = c.id)")}; auto res {session.getDboSession().query<Wt::Dbo::ptr<Cluster>>("SELECT DISTINCT c FROM cluster c WHERE NOT EXISTS(SELECT 1 FROM track_cluster t_c WHERE t_c.cluster_id = c.id)").resultList()};
return std::vector<Cluster::pointer>(res.begin(), res.end()); return std::vector<Cluster::pointer>(res.begin(), res.end());
} }
Cluster::pointer Cluster::pointer
Cluster::getById(Session& session, IdType id) Cluster::getById(Session& session, ClusterId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Cluster>().where("id = ?").bind(id); return session.getDboSession().find<Cluster>().where("id = ?").bind(id).resultValue();
} }
void void
Cluster::addTrack(Wt::Dbo::ptr<Track> track) Cluster::addTrack(ObjectPtr<Track> track)
{ {
_tracks.insert(track); _tracks.insert(getDboPtr(track));
} }
std::vector<Wt::Dbo::ptr<Track>> std::vector<Track::pointer>
Cluster::getTracks(std::optional<std::size_t> offset, std::optional<std::size_t> limit) const Cluster::getTracks(std::optional<std::size_t> offset, std::optional<std::size_t> limit) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<Track::pointer> res auto res {session()->query<Wt::Dbo::ptr<Track>>("SELECT t FROM track t INNER JOIN cluster c ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id")
{session()->query<Track::pointer>("SELECT t FROM track t INNER JOIN cluster c ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id") .where("c.id = ?").bind(getId())
.where("c.id = ?").bind(self()->id())
.offset(offset ? static_cast<int>(*offset) : -1) .offset(offset ? static_cast<int>(*offset) : -1)
.limit(limit ? static_cast<int>(*limit) : -1)}; .limit(limit ? static_cast<int>(*limit) : -1)
.resultList()};
return std::vector<Wt::Dbo::ptr<Track>>(res.begin(), res.end()); return std::vector<Track::pointer>(res.begin(), res.end());
} }
std::set<IdType> std::vector<TrackId>
Cluster::getTrackIds() const Cluster::getTrackIds() const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<IdType> res = session()->query<IdType>("SELECT t_c.track_id FROM track_cluster t_c INNER JOIN cluster c ON c.id = t_c.cluster_id") Wt::Dbo::collection<TrackId> res = session()->query<TrackId>("SELECT t_c.track_id FROM track_cluster t_c INNER JOIN cluster c ON c.id = t_c.cluster_id")
.where("c.id = ?").bind(self()->id()); .where("c.id = ?").bind(getId());
return std::set<IdType>(res.begin(), res.end());
return std::vector<TrackId>(res.begin(), res.end());
} }
std::size_t std::size_t
Cluster::getReleasesCount() const Cluster::getReleasesCount() const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
return session()->query<int>("SELECT COUNT(DISTINCT r.id) FROM release r INNER JOIN track t on t.release_id = r.id INNER JOIN cluster c ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id") return session()->query<int>("SELECT COUNT(DISTINCT r.id) FROM release r INNER JOIN track t on t.release_id = r.id INNER JOIN cluster c ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id")
.where("c.id = ?").bind(self()->id()); .where("c.id = ?").bind(getId());
} }
ClusterType::ClusterType(std::string name) ClusterType::ClusterType(std::string_view name)
: _name(name) : _name {name}
{ {
} }
@@ -132,7 +122,7 @@ ClusterType::getAllOrphans(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().query<Wt::Dbo::ptr<ClusterType>>( Wt::Dbo::collection<Wt::Dbo::ptr<ClusterType>> res = session.getDboSession().query<Wt::Dbo::ptr<ClusterType>>(
"SELECT c_t from cluster_type c_t" "SELECT c_t from cluster_type c_t"
" LEFT OUTER JOIN cluster c ON c_t.id = c.cluster_type_id") " LEFT OUTER JOIN cluster c ON c_t.id = c.cluster_type_id")
.where("c.id IS NULL"); .where("c.id IS NULL");
@@ -145,7 +135,7 @@ ClusterType::getAllUsed(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().query<Wt::Dbo::ptr<ClusterType>>( Wt::Dbo::collection<Wt::Dbo::ptr<ClusterType>> res = session.getDboSession().query<Wt::Dbo::ptr<ClusterType>>(
"SELECT DISTINCT c_t from cluster_type c_t") "SELECT DISTINCT c_t from cluster_type c_t")
.join("cluster c ON c_t.id = c.cluster_type_id"); .join("cluster c ON c_t.id = c.cluster_type_id");
@@ -157,15 +147,15 @@ ClusterType::getByName(Session& session, const std::string& name)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<ClusterType>().where("name = ?").bind(name); return session.getDboSession().find<ClusterType>().where("name = ?").bind(name).resultValue();
} }
ClusterType::pointer ClusterType::pointer
ClusterType::getById(Session& session, IdType id) ClusterType::getById(Session& session, ClusterTypeId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<ClusterType>().where("id= ?").bind(id); return session.getDboSession().find<ClusterType>().where("id = ?").bind(id).resultValue();
} }
std::vector<ClusterType::pointer> std::vector<ClusterType::pointer>
@@ -173,8 +163,7 @@ ClusterType::getAll(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().find<ClusterType>(); auto res {session.getDboSession().find<ClusterType>().resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
@@ -193,24 +182,23 @@ Cluster::pointer
ClusterType::getCluster(const std::string& name) const ClusterType::getCluster(const std::string& name) const
{ {
assert(self()); assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
return session()->find<Cluster>() return session()->find<Cluster>()
.where("name = ?").bind(name) .where("name = ?").bind(name)
.where("cluster_type_id = ?").bind(self()->id()); .where("cluster_type_id = ?").bind(getId()).resultValue();
} }
std::vector<Cluster::pointer> std::vector<Cluster::pointer>
ClusterType::getClusters() const ClusterType::getClusters() const
{ {
assert(self()); assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
Wt::Dbo::collection<Cluster::pointer> res = session()->find<Cluster>() auto res = session()->find<Cluster>()
.where("cluster_type_id = ?").bind(self()->id()) .where("cluster_type_id = ?").bind(getId())
.orderBy("name"); .orderBy("name")
.resultList();
return std::vector<Cluster::pointer>(res.begin(), res.end()); return std::vector<Cluster::pointer>(res.begin(), res.end());
} }
+108 -114
View File
@@ -28,6 +28,7 @@
#include "database/User.hpp" #include "database/User.hpp"
#include "utils/Logger.hpp" #include "utils/Logger.hpp"
#include "SqlQuery.hpp" #include "SqlQuery.hpp"
#include "Traits.hpp"
#include "Utils.hpp" #include "Utils.hpp"
namespace Database namespace Database
@@ -38,7 +39,7 @@ static
Wt::Dbo::Query<T> Wt::Dbo::Query<T>
createQuery(Session& session, createQuery(Session& session,
const std::string& queryStr, const std::string& queryStr,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
const std::vector<std::string_view>& keywords) const std::vector<std::string_view>& keywords)
{ {
@@ -57,7 +58,7 @@ createQuery(Session& session,
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id"; " INNER JOIN track_cluster t_c ON t_c.track_id = t.id";
WhereClause clusterClause; WhereClause clusterClause;
for (const IdType clusterId : clusterIds) for (const ClusterId clusterId : clusterIds)
{ {
clusterClause.Or(WhereClause("c.id = ?")); clusterClause.Or(WhereClause("c.id = ?"));
query.bind(clusterId); query.bind(clusterId);
@@ -76,7 +77,6 @@ Release::Release(const std::string& name, const std::optional<UUID>& MBID)
: _name {std::string(name, 0 , _maxNameLength)}, : _name {std::string(name, 0 , _maxNameLength)},
_MBID {MBID ? MBID->getAsString() : ""} _MBID {MBID ? MBID->getAsString() : ""}
{ {
} }
std::vector<Release::pointer> std::vector<Release::pointer>
@@ -84,7 +84,11 @@ Release::getByName(Session& session, const std::string& name)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
Wt::Dbo::collection<Release::pointer> res = session.getDboSession().find<Release>().where("name = ?").bind( std::string(name, 0, _maxNameLength) ); auto res {session.getDboSession()
.find<Release>()
.where("name = ?").bind( std::string(name, 0, _maxNameLength) )
.resultList()};
return std::vector<Release::pointer>(res.begin(), res.end()); return std::vector<Release::pointer>(res.begin(), res.end());
} }
@@ -93,15 +97,21 @@ Release::getByMBID(Session& session, const UUID& mbid)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Release>().where("mbid = ?").bind(std::string {mbid.getAsString()}); return session.getDboSession()
.find<Release>()
.where("mbid = ?").bind(std::string {mbid.getAsString()})
.resultValue();;
} }
Release::pointer Release::pointer
Release::getById(Session& session, IdType id) Release::getById(Session& session, ReleaseId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Release>().where("id = ?").bind(id); return session.getDboSession()
.find<Release>()
.where("id = ?").bind(id)
.resultValue();
} }
Release::pointer Release::pointer
@@ -120,8 +130,7 @@ Release::getCount(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> releases {session.getDboSession().find<Release>()}; return session.getDboSession().find<Release>().resultList().size();
return releases.size();
} }
std::vector<Release::pointer> std::vector<Release::pointer>
@@ -129,21 +138,22 @@ Release::getAll(Session& session, std::optional<Range> range)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().find<Release>() auto res {session.getDboSession().find<Release>()
.offset(range ? static_cast<int>(range->offset) : -1) .offset(range ? static_cast<int>(range->offset) : -1)
.limit(range ? static_cast<int>(range->limit) : -1) .limit(range ? static_cast<int>(range->limit) : -1)
.orderBy("name COLLATE NOCASE"); .orderBy("name COLLATE NOCASE")
.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<ReleaseId>
Release::getAllIds(Session& session) Release::getAllIds(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType>("SELECT id FROM release"); Wt::Dbo::collection<ReleaseId> res = session.getDboSession().query<ReleaseId>("SELECT id FROM release");
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ReleaseId>(res.begin(), res.end());
} }
std::vector<Release::pointer> std::vector<Release::pointer>
@@ -151,44 +161,45 @@ Release::getAllOrderedByArtist(Session& session, std::optional<std::size_t> offs
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().query<Wt::Dbo::ptr<Release>>( auto res {session.getDboSession().query<Wt::Dbo::ptr<Release>>(
"SELECT DISTINCT r FROM release r" "SELECT DISTINCT r FROM release r"
" INNER JOIN track t ON r.id = t.release_id" " INNER JOIN track t ON r.id = t.release_id"
" INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id" " INNER JOIN track_artist_link t_a_l ON t_a_l.track_id = t.id"
" INNER JOIN artist a ON t_a_l.artist_id = a.id") " INNER JOIN artist a ON t_a_l.artist_id = a.id")
.offset(offset ? static_cast<int>(*offset) : -1) .offset(offset ? static_cast<int>(*offset) : -1)
.limit(size ? static_cast<int>(*size) : -1) .limit(size ? static_cast<int>(*size) : -1)
.orderBy("a.name COLLATE NOCASE, r.name COLLATE NOCASE"); .orderBy("a.name COLLATE NOCASE, r.name COLLATE NOCASE")
.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getAllRandom(Session& session, const std::set<IdType>& clusterIds, std::optional<std::size_t> size) Release::getAllRandom(Session& session, const std::vector<ClusterId>& clusterIds, std::optional<std::size_t> size)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Release::pointer>(session, "SELECT DISTINCT r from release r", clusterIds,{})}; auto query {createQuery<Wt::Dbo::ptr<Release>>(session, "SELECT DISTINCT r from release r", clusterIds, {})};
auto res {query
Wt::Dbo::collection<pointer> res = query
.orderBy("RANDOM()") .orderBy("RANDOM()")
.limit(size ? static_cast<int>(*size) : -1); .limit(size ? static_cast<int>(*size) : -1)
.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<ReleaseId>
Release::getAllIdsRandom(Session& session, const std::set<IdType>& clusterIds, std::optional<std::size_t> size) Release::getAllIdsRandom(Session& session, const std::vector<ClusterId>& clusterIds, std::optional<std::size_t> size)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<IdType>(session, "SELECT DISTINCT r.id from release r", clusterIds,{})}; auto query {createQuery<ReleaseId>(session, "SELECT DISTINCT r.id from release r", clusterIds, {})};
Wt::Dbo::collection<IdType> res = query Wt::Dbo::collection<ReleaseId> res = query
.orderBy("RANDOM()") .orderBy("RANDOM()")
.limit(size ? static_cast<int>(*size) : -1); .limit(size ? static_cast<int>(*size) : -1);
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ReleaseId>(res.begin(), res.end());
} }
@@ -197,31 +208,31 @@ Release::getAllOrphans(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<Release::pointer> res = session.getDboSession().query<Wt::Dbo::ptr<Release>>("select r from release r LEFT OUTER JOIN Track t ON r.id = t.release_id WHERE t.id IS NULL"); auto res {session.getDboSession().query<Wt::Dbo::ptr<Release>>("select r from release r LEFT OUTER JOIN Track t ON r.id = t.release_id WHERE t.id IS NULL").resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getLastWritten(Session& session, Release::getLastWritten(Session& session,
std::optional<Wt::WDateTime> after, std::optional<Wt::WDateTime> after,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
std::optional<Range> range, std::optional<Range> range,
bool& moreResults) bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Release::pointer>(session, "SELECT r from release r", clusterIds, {})}; auto query {createQuery<Wt::Dbo::ptr<Release>>(session, "SELECT r from release r", clusterIds, {})};
if (after) if (after)
query.where("t.file_last_write > ?").bind(after); query.where("t.file_last_write > ?").bind(after);
Wt::Dbo::collection<Release::pointer> collection = query auto collection {query
.orderBy("t.file_last_write DESC") .orderBy("t.file_last_write DESC")
.groupBy("r.id") .groupBy("r.id")
.offset(range ? static_cast<int>(range->offset) : -1) .offset(range ? static_cast<int>(range->offset) : -1)
.limit(range ? static_cast<int>(range->limit) + 1: -1); .limit(range ? static_cast<int>(range->limit) + 1: -1)
.resultList()};
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -236,13 +247,14 @@ Release::getLastWritten(Session& session,
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getByYear(Session& session, int yearFrom, int yearTo, std::optional<Range> range) Release::getByYear(Session& session, int yearFrom, int yearTo, std::optional<Range> range)
{ {
Wt::Dbo::collection<Release::pointer> res = session.getDboSession().query<Release::pointer> auto res {session.getDboSession().query<Wt::Dbo::ptr<Release>>
("SELECT DISTINCT r from release r INNER JOIN track t ON r.id = t.release_id") ("SELECT DISTINCT r from release r INNER JOIN track t ON r.id = t.release_id")
.where("t.year >= ?").bind(yearFrom) .where("t.year >= ?").bind(yearFrom)
.where("t.year <= ?").bind(yearTo) .where("t.year <= ?").bind(yearTo)
.orderBy("t.year, r.name COLLATE NOCASE") .orderBy("t.year, r.name COLLATE NOCASE")
.offset(range ? static_cast<int>(range->offset) : -1) .offset(range ? static_cast<int>(range->offset) : -1)
.limit(range ? static_cast<int>(range->limit) : -1); .limit(range ? static_cast<int>(range->limit) : -1)
.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
@@ -250,30 +262,31 @@ Release::getByYear(Session& session, int yearFrom, int yearTo, std::optional<Ran
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getStarred(Session& session, Release::getStarred(Session& session,
User::pointer user, User::pointer user,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
std::optional<Range> range, std::optional<Range> range,
bool& moreResults) bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Release::pointer>(session, "SELECT r from release r", clusterIds, {})}; auto query {createQuery<Wt::Dbo::ptr<Release>>(session, "SELECT r from release r", clusterIds, {})};
{ {
std::ostringstream oss; std::ostringstream oss;
oss << "r.id IN (SELECT DISTINCT r.id FROM release r" oss << "r.id IN (SELECT DISTINCT r.id FROM release r"
" INNER JOIN user_release_starred urs ON urs.release_id = r.id" " INNER JOIN user_release_starred urs ON urs.release_id = r.id"
" INNER JOIN user u ON u.id = urs.user_id WHERE u.id = ?)"; " INNER JOIN user u ON u.id = urs.user_id WHERE u.id = ?)";
query.bind(user.id()); query.bind(user->getId());
query.where(oss.str()); query.where(oss.str());
} }
Wt::Dbo::collection<Release::pointer> collection = query auto collection {query
.groupBy("r.id") .groupBy("r.id")
.orderBy("r.name COLLATE NOCASE") .orderBy("r.name COLLATE NOCASE")
.offset(range ? static_cast<int>(range->offset) : -1) .offset(range ? static_cast<int>(range->offset) : -1)
.limit(range ? static_cast<int>(range->limit) + 1: -1); .limit(range ? static_cast<int>(range->limit) + 1: -1)
.resultList()};
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -288,7 +301,7 @@ Release::getStarred(Session& session,
} }
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getByClusters(Session& session, const std::set<IdType>& clusters) Release::getByClusters(Session& session, const std::vector<ClusterId>& clusters)
{ {
assert(!clusters.empty()); assert(!clusters.empty());
@@ -300,21 +313,21 @@ Release::getByClusters(Session& session, const std::set<IdType>& clusters)
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getByFilter(Session& session, Release::getByFilter(Session& session,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
const std::vector<std::string_view>& keywords, const std::vector<std::string_view>& keywords,
std::optional<Range> range, std::optional<Range> range,
bool& moreResults) bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> collection = createQuery<Release::pointer>(session, "SELECT r from release r", clusterIds, keywords) auto collection {createQuery<Wt::Dbo::ptr<Release>>(session, "SELECT r from release r", clusterIds, keywords)
.groupBy("r.id") .groupBy("r.id")
.orderBy("r.name COLLATE NOCASE") .orderBy("r.name COLLATE NOCASE")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<pointer>(collection.begin(), collection.end())};
std::vector<pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -326,18 +339,18 @@ Release::getByFilter(Session& session,
return res; return res;
} }
std::vector<IdType> std::vector<ReleaseId>
Release::getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit) Release::getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType> Wt::Dbo::collection<ReleaseId> res = session.getDboSession().query<ReleaseId>
("SELECT DISTINCT r.id FROM release r" ("SELECT DISTINCT r.id FROM release r"
" INNER JOIN track t ON t.release_id = r.id" " INNER JOIN track t ON t.release_id = r.id"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id") " INNER JOIN track_cluster t_c ON t_c.track_id = t.id")
.limit(limit ? static_cast<int>(*limit) : -1); .limit(limit ? static_cast<int>(*limit) : -1);
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ReleaseId>(res.begin(), res.end());
} }
@@ -345,11 +358,10 @@ std::optional<std::size_t>
Release::getTotalTrack(void) const Release::getTotalTrack(void) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
int res = session()->query<int>("SELECT COALESCE(MAX(total_track),0) FROM track t INNER JOIN release r ON r.id = t.release_id") int res = session()->query<int>("SELECT COALESCE(MAX(total_track),0) FROM track t INNER JOIN release r ON r.id = t.release_id")
.where("r.id = ?") .where("r.id = ?")
.bind(this->id()); .bind(getId());
return (res > 0) ? std::make_optional<std::size_t>(res) : std::nullopt; return (res > 0) ? std::make_optional<std::size_t>(res) : std::nullopt;
} }
@@ -358,11 +370,10 @@ std::optional<std::size_t>
Release::getTotalDisc(void) const Release::getTotalDisc(void) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
int res = session()->query<int>("SELECT COALESCE(MAX(total_disc),0) FROM track t INNER JOIN release r ON r.id = t.release_id") int res = session()->query<int>("SELECT COALESCE(MAX(total_disc),0) FROM track t INNER JOIN release r ON r.id = t.release_id")
.where("r.id = ?") .where("r.id = ?")
.bind(this->id()); .bind(getId());
return (res > 0) ? std::make_optional<std::size_t>(res) : std::nullopt; return (res > 0) ? std::make_optional<std::size_t>(res) : std::nullopt;
} }
@@ -372,13 +383,13 @@ Release::getReleaseYear(bool original) const
{ {
assert(session()); assert(session());
const std::string field {original ? "original_year" : "year"}; const char* field {original ? "original_year" : "year"};
Wt::Dbo::collection<int> dates = session()->query<int>( Wt::Dbo::collection<int> dates = session()->query<int>(
std::string {"SELECT "} + "t." + field + " FROM track t INNER JOIN release r ON r.id = t.release_id") std::string {"SELECT "} + "t." + field + " FROM track t INNER JOIN release r ON r.id = t.release_id")
.where("r.id = ?") .where("r.id = ?")
.groupBy(field) .groupBy(field)
.bind(this->id()); .bind(getId());
// various dates => no date // various dates => no date
if (dates.empty() || dates.size() > 1) if (dates.empty() || dates.size() > 1)
@@ -388,8 +399,8 @@ Release::getReleaseYear(bool original) const
if (date > 0) if (date > 0)
return date; return date;
else
return std::nullopt; return std::nullopt;
} }
std::optional<std::string> std::optional<std::string>
@@ -401,7 +412,7 @@ Release::getCopyright() const
("SELECT copyright FROM track t INNER JOIN release r ON r.id = t.release_id") ("SELECT copyright FROM track t INNER JOIN release r ON r.id = t.release_id")
.where("r.id = ?") .where("r.id = ?")
.groupBy("copyright") .groupBy("copyright")
.bind(this->id()); .bind(getId());
std::vector<std::string> values(copyrights.begin(), copyrights.end()); std::vector<std::string> values(copyrights.begin(), copyrights.end());
@@ -421,7 +432,7 @@ Release::getCopyrightURL() const
("SELECT copyright_url FROM track t INNER JOIN release r ON r.id = t.release_id") ("SELECT copyright_url FROM track t INNER JOIN release r ON r.id = t.release_id")
.where("r.id = ?") .where("r.id = ?")
.groupBy("copyright_url") .groupBy("copyright_url")
.bind(this->id()); .bind(getId());
std::vector<std::string> values(copyrights.begin(), copyrights.end()); std::vector<std::string> values(copyrights.begin(), copyrights.end());
@@ -432,32 +443,29 @@ Release::getCopyrightURL() const
return values.front(); return values.front();
} }
std::vector<Wt::Dbo::ptr<Artist>> std::vector<Artist::pointer>
Release::getArtists(TrackArtistLinkType linkType) const Release::getArtists(TrackArtistLinkType linkType) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
Wt::Dbo::collection<Wt::Dbo::ptr<Artist>> res = session()->query<Wt::Dbo::ptr<Artist>>( auto res {session()->query<Wt::Dbo::ptr<Artist>>(
"SELECT DISTINCT a FROM artist a" "SELECT DISTINCT a FROM artist a"
" INNER JOIN track_artist_link t_a_l ON t_a_l.artist_id = a.id" " INNER JOIN track_artist_link t_a_l ON t_a_l.artist_id = a.id"
" INNER JOIN track t ON t.id = t_a_l.track_id" " INNER JOIN track t ON t.id = t_a_l.track_id"
" INNER JOIN release r ON r.id = t.release_id") " INNER JOIN release r ON r.id = t.release_id")
.where("r.id = ?").bind(self()->id()) .where("r.id = ?").bind(getId())
.where("t_a_l.type = ?").bind(linkType); .where("t_a_l.type = ?").bind(linkType)
.resultList()};
return std::vector<Wt::Dbo::ptr<Artist>>(res.begin(), res.end()); return std::vector<Artist::pointer>(res.begin(), res.end());
} }
std::vector<Release::pointer> std::vector<Release::pointer>
Release::getSimilarReleases(std::optional<std::size_t> offset, std::optional<std::size_t> count) const Release::getSimilarReleases(std::optional<std::size_t> offset, std::optional<std::size_t> count) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
Wt::Dbo::Query<pointer> query {session()->query<pointer>( auto res {session()->query<Wt::Dbo::ptr<Release>>(
"SELECT r FROM release r" "SELECT r FROM release r"
" INNER JOIN track t ON t.release_id = r.id" " INNER JOIN track t ON t.release_id = r.id"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id" " INNER JOIN track_cluster t_c ON t_c.track_id = t.id"
@@ -465,14 +473,14 @@ Release::getSimilarReleases(std::optional<std::size_t> offset, std::optional<std
" t_c.cluster_id IN (SELECT c.id from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN release r ON r.id = t.release_id WHERE r.id = ?)" " t_c.cluster_id IN (SELECT c.id from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN release r ON r.id = t.release_id WHERE r.id = ?)"
" AND r.id <> ?" " AND r.id <> ?"
) )
.bind(self()->id()) .bind(getId())
.bind(self()->id()) .bind(getId())
.groupBy("r.id") .groupBy("r.id")
.orderBy("COUNT(*) DESC, RANDOM()") .orderBy("COUNT(*) DESC, RANDOM()")
.limit(count ? static_cast<int>(*count) : -1) .limit(count ? static_cast<int>(*count) : -1)
.offset(offset ? static_cast<int>(*offset) : -1)}; .offset(offset ? static_cast<int>(*offset) : -1)
.resultList()};
Wt::Dbo::collection<pointer> res = query;
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
@@ -483,11 +491,9 @@ Release::hasVariousArtists() const
return getArtists().size() > 1; return getArtists().size() > 1;
} }
std::vector<Wt::Dbo::ptr<Track>> std::vector<Track::pointer>
Release::getTracks(const std::set<IdType>& clusterIds) const Release::getTracks(const std::vector<ClusterId>& clusterIds) const
{ {
assert(self());
assert(self()->id() != Wt::Dbo::dbo_traits<Release>::invalidId() );
assert(session()); assert(session());
WhereClause where; WhereClause where;
@@ -502,12 +508,12 @@ Release::getTracks(const std::set<IdType>& clusterIds) const
WhereClause clusterClause; WhereClause clusterClause;
for (auto id : clusterIds) for (auto id : clusterIds)
clusterClause.Or(WhereClause("c.id = ?")).bind(std::to_string(id)); clusterClause.Or(WhereClause("c.id = ?")).bind(id.toString());
where.And(clusterClause); where.And(clusterClause);
} }
where.And(WhereClause("r.id = ?")).bind(std::to_string(id())); where.And(WhereClause("r.id = ?")).bind(getId().toString());
oss << " " << where.get(); oss << " " << where.get();
@@ -516,16 +522,12 @@ Release::getTracks(const std::set<IdType>& clusterIds) const
oss << " ORDER BY t.disc_number,t.track_number"; oss << " ORDER BY t.disc_number,t.track_number";
Wt::Dbo::Query<Track::pointer> query = session()->query<Track::pointer>( oss.str() ); auto query {session()->query<Wt::Dbo::ptr<Track>>(oss.str())};
for (const std::string& bindArg : where.getBindArgs()) for (const std::string& bindArg : where.getBindArgs())
{
query.bind(bindArg); query.bind(bindArg);
}
Wt::Dbo::collection< Wt::Dbo::ptr<Track> > res = query; auto res {query.resultList()};
return std::vector<Track::pointer> (res.begin(), res.end());
return std::vector< Wt::Dbo::ptr<Track> > (res.begin(), res.end());
} }
std::size_t std::size_t
@@ -534,31 +536,28 @@ Release::getTracksCount() const
return _tracks.size(); return _tracks.size();
} }
Wt::Dbo::ptr<Track> Track::pointer
Release::getFirstTrack() const Release::getFirstTrack() const
{ {
assert(self());
assert(self()->id() != Wt::Dbo::dbo_traits<Artist>::invalidId());
assert(session()); assert(session());
return session()->query<Track::pointer>("SELECT t from track t") return session()->query<Wt::Dbo::ptr<Track>>("SELECT t from track t")
.join("release r ON t.release_id = r.id") .join("release r ON t.release_id = r.id")
.where("r.id = ?").bind(self()->id()) .where("r.id = ?").bind(getId())
.orderBy("t.disc_number,t.track_number") .orderBy("t.disc_number,t.track_number")
.limit(1); .limit(1)
.resultValue();
} }
std::chrono::milliseconds std::chrono::milliseconds
Release::getDuration() const Release::getDuration() const
{ {
assert(self());
assert(self()->id() != Wt::Dbo::dbo_traits<Artist>::invalidId());
assert(session()); assert(session());
using milli = std::chrono::duration<int, std::milli>; using milli = std::chrono::duration<int, std::milli>;
Wt::Dbo::Query<milli> query {session()->query<milli>("SELECT COALESCE(SUM(duration), 0) FROM track t INNER JOIN release r ON t.release_id = r.id") Wt::Dbo::Query<milli> query {session()->query<milli>("SELECT COALESCE(SUM(duration), 0) FROM track t INNER JOIN release r ON t.release_id = r.id")
.where("r.id = ?").bind(self()->id())}; .where("r.id = ?").bind(getId())};
return query.resultValue(); return query.resultValue();
} }
@@ -566,21 +565,17 @@ Release::getDuration() const
Wt::WDateTime Wt::WDateTime
Release::getLastWritten() const Release::getLastWritten() const
{ {
assert(self());
assert(self()->id() != Wt::Dbo::dbo_traits<Artist>::invalidId());
assert(session()); assert(session());
Wt::Dbo::Query<Wt::WDateTime> query {session()->query<Wt::WDateTime>("SELECT COALESCE(MAX(file_last_write), '1970-01-01T00:00:00') FROM track t INNER JOIN release r ON t.release_id = r.id") Wt::Dbo::Query<Wt::WDateTime> query {session()->query<Wt::WDateTime>("SELECT COALESCE(MAX(file_last_write), '1970-01-01T00:00:00') FROM track t INNER JOIN release r ON t.release_id = r.id")
.where("r.id = ?").bind(self()->id())}; .where("r.id = ?").bind(getId())};
return query.resultValue(); return query.resultValue();
} }
std::vector<std::vector<Wt::Dbo::ptr<Cluster>>> std::vector<std::vector<Cluster::pointer>>
Release::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::size_t size) const Release::getClusterGroups(const std::vector<ClusterType::pointer>& clusterTypes, std::size_t size) const
{ {
assert(self());
assert(self()->id() != Wt::Dbo::dbo_traits<Artist>::invalidId() );
assert(session()); assert(session());
WhereClause where; WhereClause where;
@@ -589,33 +584,32 @@ Release::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::s
oss << "SELECT c from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN cluster_type c_type ON c.cluster_type_id = c_type.id INNER JOIN release r ON t.release_id = r.id "; oss << "SELECT c from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN cluster_type c_type ON c.cluster_type_id = c_type.id INNER JOIN release r ON t.release_id = r.id ";
where.And(WhereClause("r.id = ?")).bind(std::to_string(self()->id())); where.And(WhereClause("r.id = ?")).bind(getId().toString());
{ {
WhereClause clusterClause; WhereClause clusterClause;
for (auto clusterType : clusterTypes) for (auto clusterType : clusterTypes)
clusterClause.Or(WhereClause("c_type.id = ?")).bind(std::to_string(clusterType.id())); clusterClause.Or(WhereClause("c_type.id = ?")).bind(clusterType->getId().toString());
where.And(clusterClause); where.And(clusterClause);
} }
oss << " " << where.get(); oss << " " << where.get();
oss << " GROUP BY c.id ORDER BY COUNT(c.id) DESC"; oss << " GROUP BY c.id ORDER BY COUNT(c.id) DESC";
Wt::Dbo::Query<Cluster::pointer> query = session()->query<Cluster::pointer>( oss.str() ); auto query {session()->query<Wt::Dbo::ptr<Cluster>>(oss.str())};
for (const std::string& bindArg : where.getBindArgs()) for (const std::string& bindArg : where.getBindArgs())
query.bind(bindArg); query.bind(bindArg);
Wt::Dbo::collection<Cluster::pointer> queryRes = query; auto queryRes {query.resultList()};
std::map<IdType, std::vector<Cluster::pointer>> clusters; std::map<ClusterTypeId, std::vector<Cluster::pointer>> clustersByType;
for (auto cluster : queryRes) for (const Wt::Dbo::ptr<Cluster>& cluster : queryRes)
{ {
if (clusters[cluster->getType().id()].size() < size) if (clustersByType[cluster->getType()->getId()].size() < size)
clusters[cluster->getType().id()].push_back(cluster); clustersByType[cluster->getType()->getId()].push_back(cluster);
} }
std::vector<std::vector<Cluster::pointer>> res; std::vector<std::vector<Cluster::pointer>> res;
for (auto cluster_list : clusters) for (const auto& [clusterTypeId, clusters] : clustersByType)
res.push_back(cluster_list.second); res.push_back(clusters);
return res; return res;
} }
+7 -7
View File
@@ -60,14 +60,14 @@ ScanSettings::get(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<ScanSettings>(); return session.getDboSession().find<ScanSettings>().resultValue();
} }
std::unordered_set<std::filesystem::path> std::vector<std::filesystem::path>
ScanSettings::getAudioFileExtensions() const ScanSettings::getAudioFileExtensions() const
{ {
auto extensions = StringUtils::splitString(_audioFileExtensions, " "); const auto extensions {StringUtils::splitString(_audioFileExtensions, " ")};
return std::unordered_set<std::filesystem::path>(std::cbegin(extensions), std::cend(extensions)); return std::vector<std::filesystem::path>(std::cbegin(extensions), std::cend(extensions));
} }
void void
@@ -111,19 +111,19 @@ ScanSettings::setClusterTypes(Session& session, const std::set<std::string>& clu
// Create any missing cluster type // Create any missing cluster type
for (const std::string& clusterTypeName : clusterTypeNames) for (const std::string& clusterTypeName : clusterTypeNames)
{ {
auto clusterType {ClusterType::getByName(session, clusterTypeName)}; ClusterType::pointer clusterType {ClusterType::getByName(session, clusterTypeName)};
if (!clusterType) if (!clusterType)
{ {
LMS_LOG(DB, INFO) << "Creating cluster type " << clusterTypeName; LMS_LOG(DB, INFO) << "Creating cluster type " << clusterTypeName;
clusterType = ClusterType::create(session, clusterTypeName); clusterType = ClusterType::create(session, clusterTypeName);
_clusterTypes.insert(clusterType); _clusterTypes.insert(getDboPtr(clusterType));
needRescan = true; needRescan = true;
} }
} }
// Delete no longer existing cluster types // Delete no longer existing cluster types
for (ClusterType::pointer& clusterType : _clusterTypes) for (Wt::Dbo::ptr<ClusterType> clusterType : _clusterTypes)
{ {
if (std::none_of(clusterTypeNames.begin(), clusterTypeNames.end(), if (std::none_of(clusterTypeNames.begin(), clusterTypeNames.end(),
[clusterType](const std::string& name) { return name == clusterType->getName(); })) [clusterType](const std::string& name) { return name == clusterType->getName(); }))
+108 -106
View File
@@ -27,10 +27,12 @@
#include "database/TrackArtistLink.hpp" #include "database/TrackArtistLink.hpp"
#include "database/TrackFeatures.hpp" #include "database/TrackFeatures.hpp"
#include "database/Session.hpp" #include "database/Session.hpp"
#include "database/User.hpp"
#include "utils/Logger.hpp" #include "utils/Logger.hpp"
#include "SqlQuery.hpp" #include "SqlQuery.hpp"
#include "StringViewTraits.hpp" #include "StringViewTraits.hpp"
#include "Traits.hpp"
#include "Utils.hpp" #include "Utils.hpp"
namespace Database { namespace Database {
@@ -40,7 +42,7 @@ static
Wt::Dbo::Query<T> Wt::Dbo::Query<T>
createQuery(Session& session, createQuery(Session& session,
const std::string& queryStr, const std::string& queryStr,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
const std::vector<std::string_view>& keywords) const std::vector<std::string_view>& keywords)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
@@ -58,7 +60,7 @@ createQuery(Session& session,
" INNER JOIN cluster c ON c.id = t_c.cluster_id"; " INNER JOIN cluster c ON c.id = t_c.cluster_id";
WhereClause clusterClause; WhereClause clusterClause;
for (const IdType clusterId : clusterIds) for (const ClusterId clusterId : clusterIds)
{ {
clusterClause.Or(WhereClause("c.id = ?")); clusterClause.Or(WhereClause("c.id = ?"));
query.bind(clusterId); query.bind(clusterId);
@@ -91,48 +93,49 @@ Track::getAll(Session& session, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<Track::pointer> res {session.getDboSession().find<Track>() auto res {session.getDboSession().find<Track>()
.limit(limit ? static_cast<int>(*limit) : -1)}; .limit(limit ? static_cast<int>(*limit) : -1)
.resultList()};
return std::vector<Track::pointer>(std::cbegin(res), std::cend(res)); return std::vector<Track::pointer>(std::cbegin(res), std::cend(res));
} }
std::vector<Track::pointer> std::vector<Track::pointer>
Track::getAllRandom(Session& session, const std::set<IdType>& clusterIds, std::optional<std::size_t> limit) Track::getAllRandom(Session& session, const std::vector<ClusterId>& clusterIds, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Track::pointer>(session, "SELECT t from track t", clusterIds, {})}; auto query {createQuery<Wt::Dbo::ptr<Track>>(session, "SELECT t from track t", clusterIds, {})};
auto collection {query
Wt::Dbo::collection<Track::pointer> collection = query
.orderBy("RANDOM()") .orderBy("RANDOM()")
.limit(limit ? static_cast<int>(*limit) + 1: -1); .limit(limit ? static_cast<int>(*limit) + 1: -1)
.resultList()};
return std::vector<pointer>(collection.begin(), collection.end()); return std::vector<pointer>(collection.begin(), collection.end());
} }
std::vector<Database::IdType> std::vector<TrackId>
Track::getAllIdsRandom(Session& session, const std::set<IdType>& clusterIds, std::optional<std::size_t> limit) Track::getAllIdsRandom(Session& session, const std::vector<ClusterId>& clusterIds, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<IdType>(session, "SELECT t.id from track t", clusterIds, {})}; auto query {createQuery<TrackId>(session, "SELECT t.id from track t", clusterIds, {})};
Wt::Dbo::collection<IdType> collection = query Wt::Dbo::collection<TrackId> collection = query
.orderBy("RANDOM()") .orderBy("RANDOM()")
.limit(limit ? static_cast<int>(*limit) + 1: -1); .limit(limit ? static_cast<int>(*limit) + 1: -1);
return std::vector<IdType>(collection.begin(), collection.end()); return std::vector<TrackId>(collection.begin(), collection.end());
} }
std::vector<IdType> std::vector<TrackId>
Track::getAllIds(Session& session) Track::getAllIds(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType>("SELECT id FROM track"); Wt::Dbo::collection<TrackId> res = session.getDboSession().query<TrackId>("SELECT id FROM track");
return std::vector<IdType>(res.begin(), res.end()); return std::vector<TrackId>(res.begin(), res.end());
} }
Track::pointer Track::pointer
@@ -140,16 +143,17 @@ Track::getByPath(Session& session, const std::filesystem::path& p)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Track>().where("file_path = ?").bind(p.string()); return session.getDboSession().find<Track>().where("file_path = ?").bind(p.string()).resultValue();
} }
Track::pointer Track::pointer
Track::getById(Session& session, IdType id) Track::getById(Session& session, TrackId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<Track>() return session.getDboSession().find<Track>()
.where("id = ?").bind(id); .where("id = ?").bind(id)
.resultValue();
} }
std::vector<Track::pointer> std::vector<Track::pointer>
@@ -157,8 +161,9 @@ Track::getByRecordingMBID(Session& session, const UUID& mbid)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<Track::pointer> res = session.getDboSession().find<Track>() auto res {session.getDboSession().find<Track>()
.where("recording_mbid = ?").bind(std::string {mbid.getAsString()}); .where("recording_mbid = ?").bind(std::string {mbid.getAsString()})
.resultList()};
return std::vector<Track::pointer>(res.begin(), res.end()); return std::vector<Track::pointer>(res.begin(), res.end());
} }
@@ -174,17 +179,17 @@ Track::create(Session& session, const std::filesystem::path& p)
return res; return res;
} }
std::vector<std::pair<IdType, std::filesystem::path>> std::vector<std::pair<TrackId, std::filesystem::path>>
Track::getAllPaths(Session& session, std::optional<std::size_t> offset, std::optional<std::size_t> size) Track::getAllPaths(Session& session, std::optional<std::size_t> offset, std::optional<std::size_t> size)
{ {
using QueryResultType = std::tuple<IdType, std::string>; using QueryResultType = std::tuple<TrackId, std::string>;
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<QueryResultType> queryRes = session.getDboSession().query<QueryResultType>("SELECT id,file_path FROM track") Wt::Dbo::collection<QueryResultType> queryRes = session.getDboSession().query<QueryResultType>("SELECT id,file_path FROM track")
.limit(size ? static_cast<int>(*size) + 1 : -1) .limit(size ? static_cast<int>(*size) + 1 : -1)
.offset(offset ? static_cast<int>(*offset) : -1); .offset(offset ? static_cast<int>(*offset) : -1);
std::vector<std::pair<IdType, std::filesystem::path>> result; std::vector<std::pair<TrackId, std::filesystem::path>> result;
result.reserve(queryRes.size()); result.reserve(queryRes.size());
std::transform(std::begin(queryRes), std::end(queryRes), std::back_inserter(result), std::transform(std::begin(queryRes), std::end(queryRes), std::back_inserter(result),
@@ -201,26 +206,29 @@ Track::getMBIDDuplicates(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().query<pointer>( "SELECT track FROM track WHERE mbid in (SELECT mbid FROM track WHERE mbid <> '' GROUP BY mbid HAVING COUNT (*) > 1)").orderBy("track.release_id,track.disc_number,track.track_number,track.mbid"); auto res {session.getDboSession().query<Wt::Dbo::ptr<Track>>( "SELECT track FROM track WHERE mbid in (SELECT mbid FROM track WHERE mbid <> '' GROUP BY mbid HAVING COUNT (*) > 1)")
.orderBy("track.release_id,track.disc_number,track.track_number,track.mbid")
.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<Track::pointer> std::vector<Track::pointer>
Track::getLastWritten(Session& session, std::optional<Wt::WDateTime> after, const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) Track::getLastWritten(Session& session, std::optional<Wt::WDateTime> after, const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Track::pointer>(session, "SELECT t from track t", clusterIds, {})}; auto query {createQuery<Wt::Dbo::ptr<Track>>(session, "SELECT t from track t", clusterIds, {})};
if (after) if (after)
query.where("t.file_last_write > ?").bind(after); query.where("t.file_last_write > ?").bind(after);
Wt::Dbo::collection<Track::pointer> collection = query auto collection {query
.orderBy("t.file_last_write DESC") .orderBy("t.file_last_write DESC")
.groupBy("t.id") .groupBy("t.id")
.offset(range ? static_cast<int>(range->offset) : -1) .offset(range ? static_cast<int>(range->offset) : -1)
.limit(range ? static_cast<int>(range->limit) + 1: -1); .limit(range ? static_cast<int>(range->limit) + 1: -1)
.resultList()};
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -237,63 +245,65 @@ Track::getAllWithRecordingMBIDAndMissingFeatures(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().query<pointer> auto res {session.getDboSession().query<Wt::Dbo::ptr<Track>>
("SELECT t FROM track t") ("SELECT t FROM track t")
.where("LENGTH(t.recording_mbid) > 0") .where("LENGTH(t.recording_mbid) > 0")
.where("NOT EXISTS (SELECT * FROM track_features t_f WHERE t_f.track_id = t.id)"); .where("NOT EXISTS (SELECT * FROM track_features t_f WHERE t_f.track_id = t.id)")
.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<TrackId>
Track::getAllIdsWithFeatures(Session& session, std::optional<std::size_t> limit) Track::getAllIdsWithFeatures(Session& session, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType> Wt::Dbo::collection<TrackId> res = session.getDboSession().query<TrackId>
("SELECT t.id FROM track t") ("SELECT t.id FROM track t")
.where("EXISTS (SELECT * from track_features t_f WHERE t_f.track_id = t.id)") .where("EXISTS (SELECT * from track_features t_f WHERE t_f.track_id = t.id)")
.limit(limit ? static_cast<int>(*limit) : -1); .limit(limit ? static_cast<int>(*limit) : -1);
return std::vector<IdType>(res.begin(), res.end()); return std::vector<TrackId>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<TrackId>
Track::getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit) Track::getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType> Wt::Dbo::collection<TrackId> res = session.getDboSession().query<TrackId>
("SELECT DISTINCT t.id FROM track t" ("SELECT DISTINCT t.id FROM track t"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id") " INNER JOIN track_cluster t_c ON t_c.track_id = t.id")
.limit(limit ? static_cast<int>(*limit) : -1); .limit(limit ? static_cast<int>(*limit) : -1);
return std::vector<IdType>(res.begin(), res.end()); return std::vector<TrackId>(res.begin(), res.end());
} }
std::vector<Track::pointer> std::vector<Track::pointer>
Track::getStarred(Session& session, Track::getStarred(Session& session,
Wt::Dbo::ptr<User> user, ObjectPtr<User> user,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
std::optional<Range> range, bool& moreResults) std::optional<Range> range, bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
auto query {createQuery<Track::pointer>(session, "SELECT t from track t", clusterIds, {})}; auto query {createQuery<Wt::Dbo::ptr<Track>>(session, "SELECT t from track t", clusterIds, {})};
{ {
std::ostringstream oss; std::ostringstream oss;
oss << "t.id IN (SELECT DISTINCT t.id FROM track t" oss << "t.id IN (SELECT DISTINCT t.id FROM track t"
" INNER JOIN user_track_starred uts ON uts.track_id = t.id" " INNER JOIN user_track_starred uts ON uts.track_id = t.id"
" INNER JOIN user u ON u.id = uts.user_id WHERE u.id = ?)"; " INNER JOIN user u ON u.id = uts.user_id WHERE u.id = ?)";
query.bind(user.id()); query.bind(user->getId().toString());
query.where(oss.str()); query.where(oss.str());
} }
Wt::Dbo::collection<Track::pointer> collection = query auto collection {query
.offset(range ? static_cast<int>(range->offset) : -1) .offset(range ? static_cast<int>(range->offset) : -1)
.limit(range ? static_cast<int>(range->limit) + 1: -1); .limit(range ? static_cast<int>(range->limit) + 1: -1)
.resultList()};
auto res {std::vector<pointer>(collection.begin(), collection.end())}; std::vector<pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -308,43 +318,41 @@ Track::getStarred(Session& session,
std::vector<Cluster::pointer> std::vector<Cluster::pointer>
Track::getClusters() const Track::getClusters() const
{ {
std::vector< Cluster::pointer > clusters; return std::vector<Cluster::pointer>(_clusters.begin(), _clusters.end());
std::copy(_clusters.begin(), _clusters.end(), std::back_inserter(clusters));
return clusters;
} }
std::vector<IdType> std::vector<ClusterId>
Track::getClusterIds() const Track::getClusterIds() const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
Wt::Dbo::collection<IdType> res = session()->query<IdType> auto res {session()->query<ClusterId>
("SELECT DISTINCT c.id FROM cluster c INNER JOIN track_cluster t_c ON t_c.cluster_id = c.id INNER JOIN track t ON t.id = t_c.track_id") ("SELECT DISTINCT c.id FROM cluster c INNER JOIN track_cluster t_c ON t_c.cluster_id = c.id INNER JOIN track t ON t.id = t_c.track_id")
.where("t.id = ?").bind(self()->id()); .where("t.id = ?").bind(getId())
.resultList()};
return std::vector<IdType>(res.begin(), res.end()); return std::vector<ClusterId>(res.begin(), res.end());
} }
bool bool
Track::hasTrackFeatures() const Track::hasTrackFeatures() const
{ {
return (_trackFeatures.lock() != Database::TrackFeatures::pointer()); return (_trackFeatures.lock() != Wt::Dbo::ptr<Database::TrackFeatures> {});
} }
std::vector<Track::pointer> std::vector<Track::pointer>
Track::getByFilter(Session& session, Track::getByFilter(Session& session,
const std::set<IdType>& clusterIds, const std::vector<ClusterId>& clusterIds,
const std::vector<std::string_view>& keywords, const std::vector<std::string_view>& keywords,
std::optional<Range> range, std::optional<Range> range,
bool& moreResults) bool& moreResults)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> collection = createQuery<Track::pointer>(session, "SELECT t from track t", clusterIds, keywords) auto collection {createQuery<Wt::Dbo::ptr<Track>>(session, "SELECT t from track t", clusterIds, keywords)
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
std::vector<pointer> res(collection.begin(), collection.end()); std::vector<pointer> res(collection.begin(), collection.end());
if (range && (res.size() == static_cast<std::size_t>(range->limit) + 1)) if (range && (res.size() == static_cast<std::size_t>(range->limit) + 1))
@@ -362,17 +370,18 @@ std::vector<Track::pointer>
Track::getByNameAndReleaseName(Session& session, std::string_view trackName, std::string_view releaseName) Track::getByNameAndReleaseName(Session& session, std::string_view trackName, std::string_view releaseName)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> collection = session.getDboSession().query<Track::pointer>("SELECT t from track t")
auto res {session.getDboSession().query<Wt::Dbo::ptr<Track>>("SELECT t from track t")
.join("release r ON t.release_id = r.id") .join("release r ON t.release_id = r.id")
.where("t.name = ?").bind(trackName) .where("t.name = ?").bind(trackName)
.where("r.name = ?").bind(releaseName); .where("r.name = ?").bind(releaseName)
.resultList()};
return std::vector<pointer>(collection.begin(), collection.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<Track::pointer> std::vector<Track::pointer>
Track::getSimilarTracks(Session& session, Track::getSimilarTracks(Session& session,
const std::unordered_set<IdType>& tracks, const std::vector<TrackId>& tracks,
std::optional<std::size_t> offset, std::optional<std::size_t> offset,
std::optional<std::size_t> size) std::optional<std::size_t> size)
{ {
@@ -387,7 +396,7 @@ Track::getSimilarTracks(Session& session,
oss << "?"; oss << "?";
} }
Wt::Dbo::Query<pointer> query {session.getDboSession().query<pointer>( auto query {session.getDboSession().query<Wt::Dbo::ptr<Track>>(
"SELECT t FROM track t" "SELECT t FROM track t"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id" " INNER JOIN track_cluster t_c ON t_c.track_id = t.id"
" AND t_c.cluster_id IN (SELECT c.id FROM cluster c INNER JOIN track_cluster t_c ON t_c.cluster_id = c.id WHERE t_c.track_id IN (" + oss.str() + "))" " AND t_c.cluster_id IN (SELECT c.id FROM cluster c INNER JOIN track_cluster t_c ON t_c.cluster_id = c.id WHERE t_c.track_id IN (" + oss.str() + "))"
@@ -397,19 +406,18 @@ Track::getSimilarTracks(Session& session,
.limit(size ? static_cast<int>(*size) : -1) .limit(size ? static_cast<int>(*size) : -1)
.offset(offset ? static_cast<int>(*offset) : -1)}; .offset(offset ? static_cast<int>(*offset) : -1)};
for (IdType trackId : tracks) for (TrackId trackId : tracks)
query.bind(trackId ); query.bind(trackId);
for (IdType trackId : tracks) for (TrackId trackId : tracks)
query.bind(trackId ); query.bind(trackId);
Wt::Dbo::collection<pointer> res = query; auto res {query.resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<Track::pointer> std::vector<Track::pointer>
Track::getByClusters(Session& session, Track::getByClusters(Session& session, const std::vector<ClusterId>& clusters)
const std::set<IdType>& clusters)
{ {
assert(!clusters.empty()); assert(!clusters.empty());
session.checkSharedLocked(); session.checkSharedLocked();
@@ -429,23 +437,23 @@ Track::clearArtistLinks()
} }
void void
Track::addArtistLink(const Wt::Dbo::ptr<TrackArtistLink>& artistLink) Track::addArtistLink(const ObjectPtr<TrackArtistLink>& artistLink)
{ {
_trackArtistLinks.insert(artistLink); _trackArtistLinks.insert(getDboPtr(artistLink));
} }
void void
Track::setClusters(const std::vector<Wt::Dbo::ptr<Cluster>>& clusters) Track::setClusters(const std::vector<ObjectPtr<Cluster>>& clusters)
{ {
_clusters.clear(); _clusters.clear();
for (const Wt::Dbo::ptr<Cluster>& cluster : clusters) for (const ObjectPtr<Cluster>& cluster : clusters)
_clusters.insert(cluster); _clusters.insert(getDboPtr(cluster));
} }
void void
Track::setFeatures(const Wt::Dbo::ptr<TrackFeatures>& features) Track::setFeatures(const ObjectPtr<TrackFeatures>& features)
{ {
_trackFeatures = features; _trackFeatures = getDboPtr(features);
} }
std::optional<std::size_t> std::optional<std::size_t>
@@ -496,11 +504,9 @@ Track::getCopyrightURL() const
return _copyrightURL != "" ? std::make_optional<std::string>(_copyrightURL) : std::nullopt; return _copyrightURL != "" ? std::make_optional<std::string>(_copyrightURL) : std::nullopt;
} }
std::vector<Wt::Dbo::ptr<Artist>> std::vector<Artist::pointer>
Track::getArtists(EnumSet<TrackArtistLinkType> linkTypes) const Track::getArtists(EnumSet<TrackArtistLinkType> linkTypes) const
{ {
assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
std::ostringstream oss; std::ostringstream oss;
@@ -525,22 +531,20 @@ Track::getArtists(EnumSet<TrackArtistLinkType> linkTypes) const
oss << ")"; oss << ")";
} }
Wt::Dbo::Query<Artist::pointer> query {session()->query<Artist::pointer>(oss.str())}; auto query {session()->query<Wt::Dbo::ptr<Artist>>(oss.str())};
for (TrackArtistLinkType type : linkTypes) for (TrackArtistLinkType type : linkTypes)
query.bind(type); query.bind(type);
query.where("t.id = ?").bind(self()->id()); query.where("t.id = ?").bind(getId());
Wt::Dbo::collection<Artist::pointer> res = query; auto res {query.resultList()};
return std::vector<Artist::pointer>(std::begin(res), std::end(res)); return std::vector<Artist::pointer>(std::begin(res), std::end(res));
} }
std::vector<IdType> std::vector<ArtistId>
Track::getArtistIds(EnumSet<TrackArtistLinkType> linkTypes) const Track::getArtistIds(EnumSet<TrackArtistLinkType> linkTypes) const
{ {
assert(self()); assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
std::ostringstream oss; std::ostringstream oss;
@@ -565,33 +569,32 @@ Track::getArtistIds(EnumSet<TrackArtistLinkType> linkTypes) const
oss << ")"; oss << ")";
} }
Wt::Dbo::Query<IdType> query {session()->query<IdType>(oss.str()) Wt::Dbo::Query<ArtistId> query {session()->query<ArtistId>(oss.str())
.where("t.id = ?").bind(self()->id())}; .where("t.id = ?").bind(getId())};
for (TrackArtistLinkType type : linkTypes) for (TrackArtistLinkType type : linkTypes)
query.bind(type); query.bind(type);
Wt::Dbo::collection<IdType> res = query; Wt::Dbo::collection<ArtistId> res = query;
return std::vector<IdType>(std::begin(res), std::end(res)); return std::vector<ArtistId>(std::begin(res), std::end(res));
} }
std::vector<Wt::Dbo::ptr<TrackArtistLink>> std::vector<TrackArtistLink::pointer>
Track::getArtistLinks() const Track::getArtistLinks() const
{ {
return std::vector<Wt::Dbo::ptr<TrackArtistLink>>(_trackArtistLinks.begin(), _trackArtistLinks.end()); return std::vector<TrackArtistLink::pointer>(_trackArtistLinks.begin(), _trackArtistLinks.end());
} }
Wt::Dbo::ptr<TrackFeatures> ObjectPtr<TrackFeatures>
Track::getTrackFeatures() const Track::getTrackFeatures() const
{ {
return _trackFeatures.lock(); return _trackFeatures.lock();
} }
std::vector<std::vector<Cluster::pointer>> std::vector<std::vector<Cluster::pointer>>
Track::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::size_t size) const Track::getClusterGroups(const std::vector<ClusterType::pointer>& clusterTypes, std::size_t size) const
{ {
assert(self()); assert(self());
assert(IdIsValid(self()->id()));
assert(session()); assert(session());
WhereClause where; WhereClause where;
@@ -600,28 +603,27 @@ Track::getClusterGroups(std::vector<ClusterType::pointer> clusterTypes, std::siz
oss << "SELECT c from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN cluster_type c_type ON c.cluster_type_id = c_type.id"; oss << "SELECT c from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN cluster_type c_type ON c.cluster_type_id = c_type.id";
where.And(WhereClause("t.id = ?")).bind(std::to_string(self()->id())); where.And(WhereClause("t.id = ?")).bind(getId().toString());
{ {
WhereClause clusterClause; WhereClause clusterClause;
for (auto clusterType : clusterTypes) for (auto clusterType : clusterTypes)
clusterClause.Or(WhereClause("c_type.id = ?")).bind(std::to_string(clusterType.id())); clusterClause.Or(WhereClause("c_type.id = ?")).bind(clusterType->getId().toString());
where.And(clusterClause); where.And(clusterClause);
} }
oss << " " << where.get(); oss << " " << where.get();
oss << " GROUP BY c.id ORDER BY COUNT(c.id) DESC"; oss << " GROUP BY c.id ORDER BY COUNT(c.id) DESC";
Wt::Dbo::Query<Cluster::pointer> query = session()->query<Cluster::pointer>( oss.str() ); auto query {session()->query<Wt::Dbo::ptr<Cluster>>(oss.str())};
for (const std::string& bindArg : where.getBindArgs()) for (const std::string& bindArg : where.getBindArgs())
query.bind(bindArg); query.bind(bindArg);
Wt::Dbo::collection<Cluster::pointer> queryRes = query; auto queryRes {query.resultList()};
std::map<IdType, std::vector<Cluster::pointer>> clusters; std::map<ClusterTypeId, std::vector<Cluster::pointer>> clusters;
for (auto cluster : queryRes) for (const Wt::Dbo::ptr<Cluster>& cluster : queryRes)
{ {
if (clusters[cluster->getType().id()].size() < size) if (clusters[cluster->getType()->getId()].size() < size)
clusters[cluster->getType().id()].push_back(cluster); clusters[cluster->getType()->getId()].push_back(cluster);
} }
std::vector<std::vector<Cluster::pointer>> res; std::vector<std::vector<Cluster::pointer>> res;
+8 -6
View File
@@ -23,17 +23,19 @@
#include "database/Session.hpp" #include "database/Session.hpp"
#include "database/Track.hpp" #include "database/Track.hpp"
#include "Traits.hpp"
namespace Database { namespace Database {
TrackArtistLink::TrackArtistLink(Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<Artist> artist, TrackArtistLinkType type) TrackArtistLink::TrackArtistLink(ObjectPtr<Track> track, ObjectPtr<Artist> artist, TrackArtistLinkType type)
: _type {type}, : _type {type},
_track {track}, _track {getDboPtr(track)},
_artist {artist} _artist {getDboPtr(artist)}
{ {
} }
TrackArtistLink::pointer TrackArtistLink::pointer
TrackArtistLink::create(Session& session, Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<Artist> artist, TrackArtistLinkType type) TrackArtistLink::create(Session& session, ObjectPtr<Track> track, ObjectPtr<Artist> artist, TrackArtistLinkType type)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
@@ -48,9 +50,9 @@ TrackArtistLink::getUsedTypes(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<TrackArtistLinkType> collection = session.getDboSession().query<TrackArtistLinkType>("SELECT DISTINCT type from track_artist_link"); auto res {session.getDboSession().query<TrackArtistLinkType>("SELECT DISTINCT type from track_artist_link").resultList()};
return EnumSet<TrackArtistLinkType>(std::begin(collection), std::end(collection)); return EnumSet<TrackArtistLinkType>(std::begin(res), std::end(res));
} }
} }
+17 -18
View File
@@ -22,18 +22,18 @@
#include "database/Session.hpp" #include "database/Session.hpp"
#include "database/Track.hpp" #include "database/Track.hpp"
#include "database/User.hpp" #include "database/User.hpp"
#include "Traits.hpp"
namespace Database { namespace Database {
TrackBookmark::TrackBookmark(Wt::Dbo::ptr<User> user, Wt::Dbo::ptr<Track> track) TrackBookmark::TrackBookmark(ObjectPtr<User> user, ObjectPtr<Track> track)
: _user {user}, : _user {getDboPtr(user)},
_track {track} _track {getDboPtr(track)}
{ {
} }
TrackBookmark::pointer TrackBookmark::pointer
TrackBookmark::create(Session& session, Wt::Dbo::ptr<User> user, Wt::Dbo::ptr<Track> track) TrackBookmark::create(Session& session, ObjectPtr<User> user, ObjectPtr<Track> track)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
@@ -48,42 +48,41 @@ TrackBookmark::getAll(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<TrackBookmark::pointer> res {session.getDboSession().find<TrackBookmark>()}; auto res {session.getDboSession().find<TrackBookmark>().resultList()};
return std::vector<TrackBookmark::pointer>(std::cbegin(res), std::cend(res)); return std::vector<TrackBookmark::pointer>(std::cbegin(res), std::cend(res));
} }
std::vector<TrackBookmark::pointer> std::vector<TrackBookmark::pointer>
TrackBookmark::getByUser(Session& session, Wt::Dbo::ptr<User> user) TrackBookmark::getByUser(Session& session, User::pointer user)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<TrackBookmark::pointer> res auto res {session.getDboSession().find<TrackBookmark>()
{ .where("user_id = ?").bind(user->getId())
session.getDboSession().find<TrackBookmark>() .resultList()};
.where("user_id = ?").bind(user.id())
};
return std::vector<TrackBookmark::pointer>(std::cbegin(res), std::cend(res)); return std::vector<TrackBookmark::pointer>(std::cbegin(res), std::cend(res));
} }
TrackBookmark::pointer TrackBookmark::pointer
TrackBookmark::getByUser(Session& session, Wt::Dbo::ptr<User> user, Wt::Dbo::ptr<Track> track) TrackBookmark::getByUser(Session& session, ObjectPtr<User> user, ObjectPtr<Track> track)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<TrackBookmark>() return session.getDboSession().find<TrackBookmark>()
.where("user_id = ?").bind(user.id()) .where("user_id = ?").bind(user->getId())
.where("track_id = ?").bind(track.id()); .where("track_id = ?").bind(track->getId())
.resultValue();
} }
TrackBookmark::pointer TrackBookmark::pointer
TrackBookmark::getById(Session& session, IdType id) TrackBookmark::getById(Session& session, TrackBookmarkId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<TrackBookmark>() return session.getDboSession().find<TrackBookmark>()
.where("id = ?").bind(id); .where("id = ?").bind(id)
.resultValue();
} }
+4 -4
View File
@@ -28,14 +28,14 @@
namespace Database { namespace Database {
TrackFeatures::TrackFeatures(Wt::Dbo::ptr<Track> track, const std::string& jsonEncodedFeatures) TrackFeatures::TrackFeatures(ObjectPtr<Track> track, const std::string& jsonEncodedFeatures)
: _data(jsonEncodedFeatures), : _data {jsonEncodedFeatures},
_track(track) _track {getDboPtr(track)}
{ {
} }
TrackFeatures::pointer TrackFeatures::pointer
TrackFeatures::create(Session& session, Wt::Dbo::ptr<Track> track, const std::string& jsonEncodedFeatures) TrackFeatures::create(Session& session, ObjectPtr<Track> track, const std::string& jsonEncodedFeatures)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
return session.getDboSession().add(std::make_unique<TrackFeatures>(track, jsonEncodedFeatures)); return session.getDboSession().add(std::make_unique<TrackFeatures>(track, jsonEncodedFeatures));
+110 -116
View File
@@ -30,32 +30,33 @@
#include "database/Track.hpp" #include "database/Track.hpp"
#include "SqlQuery.hpp" #include "SqlQuery.hpp"
#include "StringViewTraits.hpp" #include "StringViewTraits.hpp"
#include "Traits.hpp"
namespace Database { namespace Database {
TrackList::TrackList(std::string_view name, Type type, bool isPublic, Wt::Dbo::ptr<User> user) TrackList::TrackList(std::string_view name, Type type, bool isPublic, ObjectPtr<User> user)
: _name {name}, : _name {name},
_type {type}, _type {type},
_isPublic {isPublic}, _isPublic {isPublic},
_user {user} _user {getDboPtr(user)}
{ {
} }
TrackList::pointer TrackList::pointer
TrackList::create(Session& session, std::string_view name, Type type, bool isPublic, Wt::Dbo::ptr<User> user) TrackList::create(Session& session, std::string_view name, Type type, bool isPublic, ObjectPtr<User> user)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
assert(user); assert(user);
auto res = session.getDboSession().add( std::make_unique<TrackList>(name, type, isPublic, user) ); TrackList::pointer res {session.getDboSession().add( std::make_unique<TrackList>(name, type, isPublic, user) )};
session.getDboSession().flush(); session.getDboSession().flush();
return res; return res;
} }
TrackList::pointer TrackList::pointer
TrackList::get(Session& session, std::string_view name, Type type, Wt::Dbo::ptr<User> user) TrackList::get(Session& session, std::string_view name, Type type, ObjectPtr<User> user)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
assert(user); assert(user);
@@ -63,49 +64,51 @@ TrackList::get(Session& session, std::string_view name, Type type, Wt::Dbo::ptr<
return session.getDboSession().find<TrackList>() return session.getDboSession().find<TrackList>()
.where("name = ?").bind(name) .where("name = ?").bind(name)
.where("type = ?").bind(type) .where("type = ?").bind(type)
.where("user_id = ?").bind(user.id()); .where("user_id = ?").bind(user->getId()).resultValue();
} }
std::vector<TrackList::pointer> std::vector<TrackList::pointer>
TrackList::getAll(Session& session) TrackList::getAll(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<TrackList::pointer> res = session.getDboSession().find<TrackList>();
auto res = session.getDboSession().find<TrackList>().resultList();
return std::vector<TrackList::pointer>(res.begin(), res.end());
}
std::vector<TrackList::pointer>
TrackList::getAll(Session& session, ObjectPtr<User> user)
{
session.checkSharedLocked();
auto res {session.getDboSession().find<TrackList>()
.where("user_id = ?").bind(user->getId())
.orderBy("name COLLATE NOCASE")
.resultList()};
return std::vector<TrackList::pointer>(res.begin(), res.end()); return std::vector<TrackList::pointer>(res.begin(), res.end());
} }
std::vector<TrackList::pointer> std::vector<TrackList::pointer>
TrackList::getAll(Session& session, Wt::Dbo::ptr<User> user) TrackList::getAll(Session& session, ObjectPtr<User> user, Type type)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<TrackList::pointer> res = session.getDboSession().find<TrackList>() auto res {session.getDboSession().find<TrackList>()
.where("user_id = ?").bind(user.id()) .where("user_id = ?").bind(user->getId())
.orderBy("name COLLATE NOCASE");
return std::vector<TrackList::pointer>(res.begin(), res.end());
}
std::vector<TrackList::pointer>
TrackList::getAll(Session& session, Wt::Dbo::ptr<User> user, Type type)
{
session.checkSharedLocked();
Wt::Dbo::collection<TrackList::pointer> res = session.getDboSession().find<TrackList>()
.where("user_id = ?").bind(user.id())
.where("type = ?").bind(type) .where("type = ?").bind(type)
.orderBy("name COLLATE NOCASE"); .orderBy("name COLLATE NOCASE")
.resultList()};
return std::vector<TrackList::pointer>(res.begin(), res.end()); return std::vector<TrackList::pointer>(res.begin(), res.end());
} }
TrackList::pointer TrackList::pointer
TrackList::getById(Session& session, IdType id) TrackList::getById(Session& session, TrackListId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<TrackList>().where("id = ?").bind(id); return session.getDboSession().find<TrackList>().where("id = ?").bind(id).resultValue();
} }
bool bool
@@ -120,10 +123,10 @@ TrackList::getCount() const
return _entries.size(); return _entries.size();
} }
Wt::Dbo::ptr<TrackListEntry> TrackListEntry::pointer
TrackList::getEntry(std::size_t pos) const TrackList::getEntry(std::size_t pos) const
{ {
Wt::Dbo::ptr<TrackListEntry> res; TrackListEntry::pointer res;
auto entries = getEntries(pos, 1); auto entries = getEntries(pos, 1);
if (!entries.empty()) if (!entries.empty())
@@ -132,39 +135,39 @@ TrackList::getEntry(std::size_t pos) const
return res; return res;
} }
std::vector<Wt::Dbo::ptr<TrackListEntry>> std::vector<TrackListEntry::pointer>
TrackList::getEntries(std::optional<std::size_t> offset, std::optional<std::size_t> size) const TrackList::getEntries(std::optional<std::size_t> offset, std::optional<std::size_t> size) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<Wt::Dbo::ptr<TrackListEntry>> entries = auto entries {
session()->find<TrackListEntry>() session()->find<TrackListEntry>()
.where("tracklist_id = ?").bind(self().id()) .where("tracklist_id = ?").bind(getId())
.orderBy("id") .orderBy("id")
.limit(size ? static_cast<int>(*size) : -1) .limit(size ? static_cast<int>(*size) : -1)
.offset(offset ? static_cast<int>(*offset) : -1); .offset(offset ? static_cast<int>(*offset) : -1)
.resultList()};
return std::vector<Wt::Dbo::ptr<TrackListEntry>>(entries.begin(), entries.end()); return std::vector<TrackListEntry::pointer>(entries.begin(), entries.end());
} }
Wt::Dbo::ptr<TrackListEntry> TrackListEntry::pointer
TrackList::getEntryByTrackAndDateTime(Wt::Dbo::ptr<Track> track, const Wt::WDateTime& dateTime) const TrackList::getEntryByTrackAndDateTime(ObjectPtr<Track> track, const Wt::WDateTime& dateTime) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
return session()->find<TrackListEntry>() return session()->find<TrackListEntry>()
.where("tracklist_id = ?").bind(self().id()) .where("tracklist_id = ?").bind(getId())
.where("track_id = ?").bind(track.id()) .where("track_id = ?").bind(track->getId())
.where("date_time = ?").bind(Wt::WDateTime::fromTime_t(dateTime.toTime_t())); .where("date_time = ?").bind(Wt::WDateTime::fromTime_t(dateTime.toTime_t()))
.resultValue();
} }
static static
Wt::Dbo::Query<Artist::pointer> Wt::Dbo::Query<Wt::Dbo::ptr<Artist>>
createArtistsQuery(Wt::Dbo::Session& session, const std::string& queryStr, IdType tracklistId, const std::set<IdType>& clusterIds, std::optional<TrackArtistLinkType> linkType) createArtistsQuery(Wt::Dbo::Session& session, const std::string& queryStr, TrackListId tracklistId, const std::vector<ClusterId>& clusterIds, std::optional<TrackArtistLinkType> linkType)
{ {
auto query {session.query<Artist::pointer>(queryStr)}; auto query {session.query<Wt::Dbo::ptr<Artist>>(queryStr)};
query.join("track t ON t.id = t_a_l.track_id"); query.join("track t ON t.id = t_a_l.track_id");
query.join("track_artist_link t_a_l ON t_a_l.artist_id = a.id"); query.join("track_artist_link t_a_l ON t_a_l.artist_id = a.id");
query.join("tracklist_entry p_e ON p_e.track_id = t.id"); query.join("tracklist_entry p_e ON p_e.track_id = t.id");
@@ -201,10 +204,10 @@ createArtistsQuery(Wt::Dbo::Session& session, const std::string& queryStr, IdTyp
} }
static static
Wt::Dbo::Query<Release::pointer> Wt::Dbo::Query<Wt::Dbo::ptr<Release>>
createReleasesQuery(Wt::Dbo::Session& session, const std::string& queryStr, IdType tracklistId, const std::set<IdType>& clusterIds) createReleasesQuery(Wt::Dbo::Session& session, const std::string& queryStr, TrackListId tracklistId, const std::vector<ClusterId>& clusterIds)
{ {
auto query {session.query<Release::pointer>(queryStr)}; auto query {session.query<Wt::Dbo::ptr<Release>>(queryStr)};
query.join("track t ON t.release_id = r.id"); query.join("track t ON t.release_id = r.id");
query.join("tracklist_entry p_e ON p_e.track_id = t.id"); query.join("tracklist_entry p_e ON p_e.track_id = t.id");
query.join("tracklist p ON p.id = p_e.tracklist_id"); query.join("tracklist p ON p.id = p_e.tracklist_id");
@@ -220,7 +223,7 @@ createReleasesQuery(Wt::Dbo::Session& session, const std::string& queryStr, IdTy
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id"; " INNER JOIN track_cluster t_c ON t_c.track_id = t.id";
WhereClause clusterClause; WhereClause clusterClause;
for (auto id : clusterIds) for (ClusterId id : clusterIds)
{ {
clusterClause.Or(WhereClause("c.id = ?")); clusterClause.Or(WhereClause("c.id = ?"));
query.bind(id); query.bind(id);
@@ -236,10 +239,10 @@ createReleasesQuery(Wt::Dbo::Session& session, const std::string& queryStr, IdTy
} }
static static
Wt::Dbo::Query<Track::pointer> Wt::Dbo::Query<Wt::Dbo::ptr<Track>>
createTracksQuery(Wt::Dbo::Session& session, IdType tracklistId, const std::set<IdType>& clusterIds) createTracksQuery(Wt::Dbo::Session& session, TrackListId tracklistId, const std::vector<ClusterId>& clusterIds)
{ {
auto query {session.query<Track::pointer>("SELECT t from track t INNER JOIN tracklist_entry p_e ON p_e.track_id = t.id INNER JOIN tracklist p ON p.id = p_e.tracklist_id")}; auto query {session.query<Wt::Dbo::ptr<Track>>("SELECT t from track t INNER JOIN tracklist_entry p_e ON p_e.track_id = t.id INNER JOIN tracklist p ON p.id = p_e.tracklist_id")};
query.where("p.id = ?").bind(tracklistId); query.where("p.id = ?").bind(tracklistId);
@@ -253,7 +256,7 @@ createTracksQuery(Wt::Dbo::Session& session, IdType tracklistId, const std::set<
WhereClause clusterClause; WhereClause clusterClause;
for (auto id : clusterIds) for (auto id : clusterIds)
{ {
clusterClause.Or(WhereClause("c.id = ?")).bind(std::to_string(id)); clusterClause.Or(WhereClause("c.id = ?")).bind(id.toString());
query.bind(id); query.bind(id);
} }
@@ -267,16 +270,16 @@ createTracksQuery(Wt::Dbo::Session& session, IdType tracklistId, const std::set<
} }
std::vector<Artist::pointer> std::vector<Artist::pointer>
TrackList::getArtistsReverse(const std::set<IdType>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const TrackList::getArtistsReverse(const std::vector<ClusterId>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<Artist::pointer> collection = createArtistsQuery(*session(), "SELECT a from artist a", self()->id(), clusterIds, linkType) auto collection {createArtistsQuery(*session(), "SELECT a from artist a", getId(), clusterIds, linkType)
.groupBy("a.id").having("p_e.date_time = MAX(p_e.date_time)") .groupBy("a.id").having("p_e.date_time = MAX(p_e.date_time)")
.orderBy("p_e.date_time DESC") .orderBy("p_e.date_time DESC")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<Artist::pointer>(collection.begin(), collection.end())}; auto res {std::vector<Artist::pointer>(collection.begin(), collection.end())};
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
@@ -291,18 +294,18 @@ TrackList::getArtistsReverse(const std::set<IdType>& clusterIds, std::optional<T
} }
std::vector<Release::pointer> std::vector<Release::pointer>
TrackList::getReleasesReverse(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const TrackList::getReleasesReverse(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<Release::pointer> collection = createReleasesQuery(*session(), "SELECT r from release r", self()->id(), clusterIds) auto collection {createReleasesQuery(*session(), "SELECT r from release r", getId(), clusterIds)
.groupBy("r.id").having("p_e.date_time = MAX(p_e.date_time)") .groupBy("r.id").having("p_e.date_time = MAX(p_e.date_time)")
.orderBy("p_e.date_time DESC") .orderBy("p_e.date_time DESC")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<Release::pointer>(collection.begin(), collection.end())}; std::vector<Release::pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -315,18 +318,18 @@ TrackList::getReleasesReverse(const std::set<IdType>& clusterIds, std::optional<
} }
std::vector<Track::pointer> std::vector<Track::pointer>
TrackList::getTracksReverse(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const TrackList::getTracksReverse(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<Track::pointer> collection = createTracksQuery(*session(), self()->id(), clusterIds) auto collection {createTracksQuery(*session(), getId(), clusterIds)
.groupBy("t.id").having("p_e.date_time = MAX(p_e.date_time)") .groupBy("t.id").having("p_e.date_time = MAX(p_e.date_time)")
.orderBy("p_e.date_time DESC") .orderBy("p_e.date_time DESC")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<Track::pointer>(collection.begin(), collection.end())}; std::vector<Track::pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -338,29 +341,28 @@ TrackList::getTracksReverse(const std::set<IdType>& clusterIds, std::optional<Ra
return res; return res;
} }
std::vector<Wt::Dbo::ptr<Cluster>> std::vector<Cluster::pointer>
TrackList::getClusters() const TrackList::getClusters() const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<Cluster::pointer> res = session()->query<Cluster::pointer>("SELECT c from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN tracklist_entry p_e ON p_e.track_id = t.id INNER JOIN tracklist p ON p.id = p_e.tracklist_id") auto res {session()->query<Wt::Dbo::ptr<Cluster>>("SELECT c from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN tracklist_entry p_e ON p_e.track_id = t.id INNER JOIN tracklist p ON p.id = p_e.tracklist_id")
.where("p.id = ?").bind(self()->id()) .where("p.id = ?").bind(getId())
.groupBy("c.id") .groupBy("c.id")
.orderBy("COUNT(c.id) DESC"); .orderBy("COUNT(c.id) DESC")
.resultList()};
return std::vector<Wt::Dbo::ptr<Cluster>>(res.begin(), res.end()); return std::vector<Cluster::pointer>(res.begin(), res.end());
} }
bool bool
TrackList::hasTrack(IdType trackId) const TrackList::hasTrack(TrackId trackId) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<TrackListEntry::pointer> res = session()->query<TrackListEntry::pointer>("SELECT p_e from tracklist_entry p_e INNER JOIN tracklist p ON p_e.tracklist_id = p.id") Wt::Dbo::collection<TrackListEntry::pointer> res = session()->query<TrackListEntry::pointer>("SELECT p_e from tracklist_entry p_e INNER JOIN tracklist p ON p_e.tracklist_id = p.id")
.where("p_e.track_id = ?").bind(trackId) .where("p_e.track_id = ?").bind(trackId)
.where("p.id = ?").bind(self()->id()); .where("p.id = ?").bind(getId());
return res.size() > 0; return res.size() > 0;
} }
@@ -369,67 +371,64 @@ std::vector<Track::pointer>
TrackList::getSimilarTracks(std::optional<std::size_t> offset, std::optional<std::size_t> size) const TrackList::getSimilarTracks(std::optional<std::size_t> offset, std::optional<std::size_t> size) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::Query<Track::pointer> query {session()->query<Track::pointer>( auto res {session()->query<Wt::Dbo::ptr<Track>>(
"SELECT t FROM track t" "SELECT t FROM track t"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id" " INNER JOIN track_cluster t_c ON t_c.track_id = t.id"
" WHERE " " WHERE "
" (t_c.cluster_id IN (SELECT c.id from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN tracklist_entry p_e ON p_e.track_id = t.id INNER JOIN tracklist p ON p.id = p_e.tracklist_id WHERE p.id = ?)" " (t_c.cluster_id IN (SELECT c.id from cluster c INNER JOIN track t ON c.id = t_c.cluster_id INNER JOIN track_cluster t_c ON t_c.track_id = t.id INNER JOIN tracklist_entry p_e ON p_e.track_id = t.id INNER JOIN tracklist p ON p.id = p_e.tracklist_id WHERE p.id = ?)"
" AND t.id NOT IN (SELECT tracklist_t.id FROM track tracklist_t INNER JOIN tracklist_entry t_e ON t_e.track_id = tracklist_t.id WHERE t_e.tracklist_id = ?))" " AND t.id NOT IN (SELECT tracklist_t.id FROM track tracklist_t INNER JOIN tracklist_entry t_e ON t_e.track_id = tracklist_t.id WHERE t_e.tracklist_id = ?))"
) )
.bind(self()->id()) .bind(getId())
.bind(self()->id()) .bind(getId())
.groupBy("t.id") .groupBy("t.id")
.orderBy("COUNT(*) DESC, RANDOM()") .orderBy("COUNT(*) DESC, RANDOM()")
.limit(size ? static_cast<int>(*size) : -1) .limit(size ? static_cast<int>(*size) : -1)
.offset(offset ? static_cast<int>(*offset) : -1)}; .offset(offset ? static_cast<int>(*offset) : -1)
.resultList()};
Wt::Dbo::collection<Track::pointer> tracks = query; return std::vector<Track::pointer>(res.begin(), res.end());
return std::vector<Track::pointer>(tracks.begin(), tracks.end());
} }
std::vector<IdType> std::vector<TrackId>
TrackList::getTrackIds() const TrackList::getTrackIds() const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
Wt::Dbo::collection<IdType> res = session()->query<IdType>("SELECT p_e.track_id from tracklist_entry p_e INNER JOIN tracklist p ON p_e.tracklist_id = p.id") Wt::Dbo::collection<TrackId> res = session()->query<TrackId>("SELECT p_e.track_id from tracklist_entry p_e INNER JOIN tracklist p ON p_e.tracklist_id = p.id")
.where("p.id = ?").bind(self()->id()); .where("p.id = ?").bind(getId());
return std::vector<IdType>(res.begin(), res.end()); return std::vector<TrackId>(res.begin(), res.end());
} }
std::chrono::milliseconds std::chrono::milliseconds
TrackList::getDuration() const TrackList::getDuration() const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
using milli = std::chrono::duration<int, std::milli>; using milli = std::chrono::duration<int, std::milli>;
Wt::Dbo::Query<milli> query {session()->query<milli>("SELECT COALESCE(SUM(duration), 0) FROM track t INNER JOIN tracklist_entry p_e ON t.id = p_e.track_id") Wt::Dbo::Query<milli> query {session()->query<milli>("SELECT COALESCE(SUM(duration), 0) FROM track t INNER JOIN tracklist_entry p_e ON t.id = p_e.track_id")
.where("p_e.tracklist_id = ?").bind(self()->id())}; .where("p_e.tracklist_id = ?").bind(getId())};
return query.resultValue(); return query.resultValue();
} }
std::vector<Artist::pointer> std::vector<Artist::pointer>
TrackList::getTopArtists(const std::set<IdType>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const TrackList::getTopArtists(const std::vector<ClusterId>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
auto query {createArtistsQuery(*session(), "SELECT a from artist a", self()->id(), clusterIds, linkType)}; auto query {createArtistsQuery(*session(), "SELECT a from artist a", getId(), clusterIds, linkType)};
Wt::Dbo::collection<Artist::pointer> collection = query auto collection {query
.orderBy("COUNT(a.id) DESC") .orderBy("COUNT(a.id) DESC")
.groupBy("a.id") .groupBy("a.id")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<Artist::pointer>(collection.begin(), collection.end())}; std::vector<Artist::pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
@@ -439,26 +438,23 @@ TrackList::getTopArtists(const std::set<IdType>& clusterIds, std::optional<Track
else else
moreResults = false; moreResults = false;
return res; return res;
} }
std::vector<Release::pointer> std::vector<Release::pointer>
TrackList::getTopReleases(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const TrackList::getTopReleases(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
auto query {createReleasesQuery(*session(), "SELECT r from release r", self()->id(), clusterIds)}; auto query {createReleasesQuery(*session(), "SELECT r from release r", getId(), clusterIds)};
auto collection {query
Wt::Dbo::collection<Release::pointer> collection = query
.orderBy("COUNT(r.id) DESC") .orderBy("COUNT(r.id) DESC")
.groupBy("r.id") .groupBy("r.id")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<Release::pointer>(collection.begin(), collection.end())};
std::vector<Release::pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -471,21 +467,19 @@ TrackList::getTopReleases(const std::set<IdType>& clusterIds, std::optional<Rang
} }
std::vector<Track::pointer> std::vector<Track::pointer>
TrackList::getTopTracks(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const TrackList::getTopTracks(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const
{ {
assert(session()); assert(session());
assert(IdIsValid(self()->id()));
auto query {createTracksQuery(*session(), self()->id(), clusterIds)}; auto query {createTracksQuery(*session(), getId(), clusterIds)};
auto collection {query
Wt::Dbo::collection<Track::pointer> collection = query
.orderBy("COUNT(t.id) DESC") .orderBy("COUNT(t.id) DESC")
.groupBy("t.id") .groupBy("t.id")
.limit(range ? static_cast<int>(range->limit) + 1 : -1) .limit(range ? static_cast<int>(range->limit) + 1 : -1)
.offset(range ? static_cast<int>(range->offset) : -1); .offset(range ? static_cast<int>(range->offset) : -1)
.resultList()};
auto res {std::vector<Track::pointer>(collection.begin(), collection.end())};
std::vector<Track::pointer> res(collection.begin(), collection.end());
if (range && res.size() == static_cast<std::size_t>(range->limit) + 1) if (range && res.size() == static_cast<std::size_t>(range->limit) + 1)
{ {
moreResults = true; moreResults = true;
@@ -497,16 +491,16 @@ TrackList::getTopTracks(const std::set<IdType>& clusterIds, std::optional<Range>
return res; return res;
} }
TrackListEntry::TrackListEntry(Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<TrackList> tracklist, const Wt::WDateTime& dateTime) TrackListEntry::TrackListEntry(ObjectPtr<Track> track, ObjectPtr<TrackList> tracklist, const Wt::WDateTime& dateTime)
: _dateTime {Wt::WDateTime::fromTime_t(dateTime.toTime_t())} // force second resolution : _dateTime {Wt::WDateTime::fromTime_t(dateTime.toTime_t())} // force second resolution
, _track {track} , _track {getDboPtr(track)}
, _tracklist {tracklist} , _tracklist {getDboPtr(tracklist)}
{ {
assert(_dateTime.isValid()); assert(_dateTime.isValid());
} }
TrackListEntry::pointer TrackListEntry::pointer
TrackListEntry::create(Session& session, Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<TrackList> tracklist, const Wt::WDateTime& dateTime) TrackListEntry::create(Session& session, ObjectPtr<Track> track, ObjectPtr<TrackList> tracklist, const Wt::WDateTime& dateTime)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
assert(track); assert(track);
@@ -519,11 +513,11 @@ TrackListEntry::create(Session& session, Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr
} }
TrackListEntry::pointer TrackListEntry::pointer
TrackListEntry::getById(Session& session, IdType id) TrackListEntry::getById(Session& session, TrackListEntryId id)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<TrackListEntry>().where("id = ?").bind(id); return session.getDboSession().find<TrackListEntry>().where("id = ?").bind(id).resultValue();
} }
} // namespace Database } // namespace Database
+60
View File
@@ -0,0 +1,60 @@
/*
* 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 <http://www.gnu.org/licenses/>.
*/
#pragma once
#include <type_traits>
#include <Wt/Dbo/StdSqlTraits.h>
#include "database/Types.hpp"
namespace Wt::Dbo
{
template<typename T>
struct sql_value_traits<T, typename std::enable_if<std::is_base_of<Database::IdType, T>::value>::type>
{
static_assert(!std::is_same_v<Database::IdType, T>, "Cannot use IdType, use derived types");
static const bool specialized = true;
static std::string type(SqlConnection *conn, int size)
{
return sql_value_traits<typename T::ValueType, void>::type(conn, size);
}
static void bind(const T& v, SqlStatement *statement, int column, int size)
{
sql_value_traits<typename T::ValueType>::bind(v.getValue(), statement, column, size);
}
static bool read(T& v, SqlStatement *statement, int column, int size)
{
typename T::ValueType value;
if (sql_value_traits<typename T::ValueType>::read(value, statement, column, size))
{
v = value;
return true;
}
v = {};
return false;
}
};
}
+41 -40
View File
@@ -26,25 +26,25 @@
#include "database/TrackList.hpp" #include "database/TrackList.hpp"
#include "utils/Logger.hpp" #include "utils/Logger.hpp"
#include "StringViewTraits.hpp" #include "StringViewTraits.hpp"
#include "Traits.hpp"
namespace Database { namespace Database {
AuthToken::AuthToken(const std::string& value, const Wt::WDateTime& expiry, Wt::Dbo::ptr<User> user) AuthToken::AuthToken(const std::string& value, const Wt::WDateTime& expiry, ObjectPtr<User> user)
: _value {value} : _value {value}
, _expiry {expiry} , _expiry {expiry}
, _user {user} , _user {getDboPtr(user)}
{ {
} }
AuthToken::pointer AuthToken::pointer
AuthToken::create(Session& session, const std::string& value, const Wt::WDateTime& expiry, Wt::Dbo::ptr<User> user) AuthToken::create(Session& session, const std::string& value, const Wt::WDateTime& expiry, ObjectPtr<User> user)
{ {
session.checkUniqueLocked(); session.checkUniqueLocked();
auto res {session.getDboSession().add(std::make_unique<AuthToken>(value, expiry, user))}; AuthToken::pointer res {session.getDboSession().add(std::make_unique<AuthToken>(value, expiry, user))};
session.getDboSession().flush(); session.getDboSession().flush();
return res; return res;
@@ -65,7 +65,8 @@ AuthToken::getByValue(Session& session, const std::string& value)
session.checkSharedLocked(); session.checkSharedLocked();
return session.getDboSession().find<AuthToken>() return session.getDboSession().find<AuthToken>()
.where("value = ?").bind(value); .where("value = ?").bind(value)
.resultValue();
} }
static const std::string queuedListName {"__queued_tracks__"}; static const std::string queuedListName {"__queued_tracks__"};
@@ -80,17 +81,17 @@ User::getAll(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<pointer> res = session.getDboSession().find<User>(); auto res {session.getDboSession().find<User>().resultList()};
return std::vector<pointer>(res.begin(), res.end()); return std::vector<pointer>(res.begin(), res.end());
} }
std::vector<IdType> std::vector<UserId>
User::getAllIds(Session& session) User::getAllIds(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
Wt::Dbo::collection<IdType> res = session.getDboSession().query<IdType>("SELECT id FROM user"); auto res {session.getDboSession().query<UserId>("SELECT id FROM user").resultList()};
return std::vector<IdType>(res.begin(), res.end()); return std::vector<UserId>(res.begin(), res.end());
} }
User::pointer User::pointer
@@ -98,8 +99,7 @@ User::getDemo(Session& session)
{ {
session.checkSharedLocked(); session.checkSharedLocked();
pointer res = session.getDboSession().find<User>().where("type = ?").bind(UserType::DEMO); return session.getDboSession().find<User>().where("type = ?").bind(UserType::DEMO).resultValue();
return res;
} }
std::size_t std::size_t
@@ -125,16 +125,17 @@ User::create(Session& session, std::string_view loginName)
} }
User::pointer User::pointer
User::getById(Session& session, IdType id) User::getById(Session& session, UserId id)
{ {
return session.getDboSession().find<User>().where("id = ?").bind( id ); return session.getDboSession().find<User>().where("id = ?").bind(id).resultValue();
} }
User::pointer User::pointer
User::getByLoginName(Session& session, std::string_view name) User::getByLoginName(Session& session, std::string_view name)
{ {
return session.getDboSession().find<User>() return session.getDboSession().find<User>()
.where("login_name = ?").bind(name); .where("login_name = ?").bind(name)
.resultValue();
} }
void void
@@ -150,7 +151,7 @@ User::clearAuthTokens()
_authTokens.clear(); _authTokens.clear();
} }
Wt::Dbo::ptr<TrackList> TrackList::pointer
User::getQueuedTrackList(Session& session) const User::getQueuedTrackList(Session& session) const
{ {
assert(self()); assert(self());
@@ -160,63 +161,63 @@ User::getQueuedTrackList(Session& session) const
} }
void void
User::starArtist(Wt::Dbo::ptr<Artist> artist) User::starArtist(ObjectPtr<Artist> artist)
{ {
if (_starredArtists.count(artist) == 0) if (_starredArtists.count(getDboPtr(artist)) == 0)
_starredArtists.insert(artist); _starredArtists.insert(getDboPtr(artist));
} }
void void
User::unstarArtist(Wt::Dbo::ptr<Artist> artist) User::unstarArtist(ObjectPtr<Artist> artist)
{ {
if (_starredArtists.count(artist) != 0) if (_starredArtists.count(getDboPtr(artist)) != 0)
_starredArtists.erase(artist); _starredArtists.erase(getDboPtr(artist));
} }
bool bool
User::hasStarredArtist(Wt::Dbo::ptr<Artist> artist) const User::hasStarredArtist(ObjectPtr<Artist> artist) const
{ {
return _starredArtists.count(artist) != 0; return _starredArtists.count(getDboPtr(artist)) != 0;
} }
void void
User::starRelease(Wt::Dbo::ptr<Release> release) User::starRelease(ObjectPtr<Release> release)
{ {
if (_starredReleases.count(release) == 0) if (_starredReleases.count(getDboPtr(release)) == 0)
_starredReleases.insert(release); _starredReleases.insert(getDboPtr(release));
} }
void void
User::unstarRelease(Wt::Dbo::ptr<Release> release) User::unstarRelease(ObjectPtr<Release> release)
{ {
if (_starredReleases.count(release) != 0) if (_starredReleases.count(getDboPtr(release)) != 0)
_starredReleases.erase(release); _starredReleases.erase(getDboPtr(release));
} }
bool bool
User::hasStarredRelease(Wt::Dbo::ptr<Release> release) const User::hasStarredRelease(ObjectPtr<Release> release) const
{ {
return _starredReleases.count(release) != 0; return _starredReleases.count(getDboPtr(release)) != 0;
} }
void void
User::starTrack(Wt::Dbo::ptr<Track> track) User::starTrack(ObjectPtr<Track> track)
{ {
if (_starredTracks.count(track) == 0) if (_starredTracks.count(getDboPtr(track)) == 0)
_starredTracks.insert(track); _starredTracks.insert(getDboPtr(track));
} }
void void
User::unstarTrack(Wt::Dbo::ptr<Track> track) User::unstarTrack(ObjectPtr<Track> track)
{ {
if (_starredTracks.count(track) != 0) if (_starredTracks.count(getDboPtr(track)) != 0)
_starredTracks.erase(track); _starredTracks.erase(getDboPtr(track));
} }
bool bool
User::hasStarredTrack(Wt::Dbo::ptr<Track> track) const User::hasStarredTrack(ObjectPtr<Track> track) const
{ {
return _starredTracks.count(track) != 0; return _starredTracks.count(getDboPtr(track)) != 0;
} }
} // namespace Database } // namespace Database
+24 -29
View File
@@ -22,17 +22,15 @@
#include <optional> #include <optional>
#include <string> #include <string>
#include <string_view> #include <string_view>
#include <unordered_set>
#include <vector> #include <vector>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include "database/Types.hpp"
#include "utils/EnumSet.hpp" #include "utils/EnumSet.hpp"
#include "utils/UUID.hpp" #include "utils/UUID.hpp"
#include "Types.hpp"
namespace Database namespace Database
{ {
@@ -44,10 +42,9 @@ class Track;
class TrackArtistLink; class TrackArtistLink;
class User; class User;
class Artist : public Wt::Dbo::Dbo<Artist> class Artist : public Object<Artist, ArtistId>
{ {
public: public:
enum class SortMethod enum class SortMethod
{ {
None, None,
@@ -55,43 +52,41 @@ class Artist : public Wt::Dbo::Dbo<Artist>
BySortName, BySortName,
}; };
using pointer = Wt::Dbo::ptr<Artist>; Artist() = default;
Artist() {}
Artist(const std::string& name, const std::optional<UUID>& MBID = {}); Artist(const std::string& name, const std::optional<UUID>& MBID = {});
// Accessors // Accessors
static pointer getByMBID(Session& session, const UUID& MBID); static pointer getByMBID(Session& session, const UUID& MBID);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, ArtistId id);
static std::vector<pointer> getByName(Session& session, const std::string& name); // exact match on name field static std::vector<pointer> getByName(Session& session, const std::string& name); // exact match on name field
static std::vector<pointer> getByClusters(Session& session, static std::vector<pointer> getByClusters(Session& session,
const std::set<IdType>& clusters, // at least one track that belongs to these clusters const std::vector<ClusterId>& clusters, // at least one track that belongs to these clusters
SortMethod sortMethod SortMethod sortMethod
); );
static std::vector<pointer> getByFilter(Session& session, static std::vector<pointer> getByFilter(Session& session,
const std::set<IdType>& clusters, // if non empty, at least one artist that belongs to these clusters const std::vector<ClusterId>& clusters, // if non empty, at least one artist that belongs to these clusters
const std::vector<std::string_view>& keywords, // if non empty, name must match all of these keywords (name + sort name fields) const std::vector<std::string_view>& keywords, // if non empty, name must match all of these keywords (name + sort name fields)
std::optional<TrackArtistLinkType> linkType, // if set, only artists that have produced at least one track with this link type std::optional<TrackArtistLinkType> linkType, // if set, only artists that have produced at least one track with this link type
SortMethod sortMethod, SortMethod sortMethod,
std::optional<Range> range, std::optional<Range> range,
bool& moreExpected); bool& moreExpected);
static std::vector<pointer> getAll(Session& session); static std::vector<pointer> getAll(Session& session);
static std::vector<pointer> getAll(Session& session, SortMethod sortMethod); static std::vector<pointer> getAll(Session& session, SortMethod sortMethod);
static std::vector<pointer> getAll(Session& session, SortMethod sortMethod, std::optional<Range> range, bool& moreResults); static std::vector<pointer> getAll(Session& session, SortMethod sortMethod, std::optional<Range> range, bool& moreResults);
static std::vector<IdType> getAllIds(Session& session); static std::vector<ArtistId> getAllIds(Session& session);
static std::vector<IdType> getAllIdsRandom(Session& session, const std::set<IdType>& clusters, std::optional<TrackArtistLinkType> linkType, std::optional<std::size_t> size = {}); static std::vector<ArtistId> getAllIdsRandom(Session& session, const std::vector<ClusterId>& clusters, std::optional<TrackArtistLinkType> linkType, std::optional<std::size_t> size = {});
static std::vector<pointer> getAllOrphans(Session& session); // No track related static std::vector<pointer> getAllOrphans(Session& session); // No track related
static std::vector<pointer> getLastWritten(Session& session, static std::vector<pointer> getLastWritten(Session& session,
std::optional<Wt::WDateTime> after, std::optional<Wt::WDateTime> after,
const std::set<IdType>& clusters, const std::vector<ClusterId>& clusters,
std::optional<TrackArtistLinkType> linkType, // if set, only artists that have produced at least one track with this link type std::optional<TrackArtistLinkType> linkType, // if set, only artists that have produced at least one track with this link type
std::optional<Range>, std::optional<Range>,
bool& moreResults); bool& moreResults);
static std::vector<IdType> getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit = {}); static std::vector<ArtistId> getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit = {});
static std::vector<pointer> getStarred(Session& session, static std::vector<pointer> getStarred(Session& session,
Wt::Dbo::ptr<User> user, ObjectPtr<User> user,
const std::set<IdType>& clusters, const std::vector<ClusterId>& clusters,
std::optional<TrackArtistLinkType> linkType, // if set, only artists that have produced at least one track with this link type std::optional<TrackArtistLinkType> linkType, // if set, only artists that have produced at least one track with this link type
SortMethod sortMethod, SortMethod sortMethod,
std::optional<Range>, bool& moreResults); std::optional<Range>, bool& moreResults);
@@ -101,12 +96,12 @@ class Artist : public Wt::Dbo::Dbo<Artist>
const std::string& getSortName() const { return _sortName; } const std::string& getSortName() const { return _sortName; }
std::optional<UUID> getMBID() const { return UUID::fromString(_MBID); } std::optional<UUID> getMBID() const { return UUID::fromString(_MBID); }
std::vector<Wt::Dbo::ptr<Release>> getReleases(const std::set<IdType>& clusterIds = {}) const; // if non empty, get the releases that match all these clusters std::vector<ObjectPtr<Release>> getReleases(const std::vector<ClusterId>& clusterIds = {}) const; // if non empty, get the releases that match all these clusters
std::size_t getReleaseCount() const; std::size_t getReleaseCount() const;
std::vector<Wt::Dbo::ptr<Track>> getTracks(std::optional<TrackArtistLinkType> linkType = {}) const; std::vector<ObjectPtr<Track>> getTracks(std::optional<TrackArtistLinkType> linkType = {}) const;
bool hasNonReleaseTracks(std::optional<TrackArtistLinkType> linkType = std::nullopt) const; bool hasNonReleaseTracks(std::optional<TrackArtistLinkType> linkType = std::nullopt) const;
std::vector<Wt::Dbo::ptr<Track>> getNonReleaseTracks(std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Track>> getNonReleaseTracks(std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const;
std::vector<Wt::Dbo::ptr<Track>> getRandomTracks(std::optional<std::size_t> count) const; std::vector<ObjectPtr<Track>> getRandomTracks(std::optional<std::size_t> count) const;
// No artistLinkTypes means get them all // No artistLinkTypes means get them all
std::vector<pointer> getSimilarArtists(EnumSet<TrackArtistLinkType> artistLinkTypes = {}, std::optional<Range> range = std::nullopt) const; std::vector<pointer> getSimilarArtists(EnumSet<TrackArtistLinkType> artistLinkTypes = {}, std::optional<Range> range = std::nullopt) const;
@@ -114,14 +109,14 @@ class Artist : public Wt::Dbo::Dbo<Artist>
// Get the cluster of the tracks made by this artist // Get the cluster of the tracks made by this artist
// Each clusters are grouped by cluster type, sorted by the number of occurence // Each clusters are grouped by cluster type, sorted by the number of occurence
// size is the max number of cluster per cluster type // size is the max number of cluster per cluster type
std::vector<std::vector<Wt::Dbo::ptr<Cluster>>> getClusterGroups(std::vector<Wt::Dbo::ptr<ClusterType>> clusterTypes, std::size_t size) const; std::vector<std::vector<ObjectPtr<Cluster>>> getClusterGroups(std::vector<ObjectPtr<ClusterType>> clusterTypes, std::size_t size) const;
void setName(std::string_view name) { _name = name; } void setName(std::string_view name) { _name = name; }
void setMBID(const std::optional<UUID>& mbid) { _MBID = mbid ? mbid->getAsString() : ""; } void setMBID(const std::optional<UUID>& mbid) { _MBID = mbid ? mbid->getAsString() : ""; }
void setSortName(const std::string& sortName); void setSortName(const std::string& sortName);
// Create // Create
static pointer create(Session& session, const std::string& name, const std::optional<UUID>& UUID = {}); static pointer create(Session& session, const std::string& name, const std::optional<UUID>& UUID = {});
template<class Action> template<class Action>
void persist(Action& a) void persist(Action& a)
+15 -21
View File
@@ -24,10 +24,9 @@
#include <vector> #include <vector>
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include "Types.hpp" #include "database/Types.hpp"
namespace Database { namespace Database {
@@ -36,31 +35,29 @@ class ClusterType;
class ScanSettings; class ScanSettings;
class Session; class Session;
class Cluster : public Wt::Dbo::Dbo<Cluster> class Cluster : public Object<Cluster, ClusterId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<Cluster>; Cluster() = default;
Cluster(ObjectPtr<ClusterType> type, std::string_view name);
Cluster();
Cluster(Wt::Dbo::ptr<ClusterType> type, std::string_view name);
// Find utility // Find utility
static std::vector<pointer> getAll(Session& session); static std::vector<pointer> getAll(Session& session);
static std::vector<pointer> getAllOrphans(Session& session); static std::vector<pointer> getAllOrphans(Session& session);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, ClusterId id);
// Create utility // Create utility
static pointer create(Session& session, Wt::Dbo::ptr<ClusterType> type, std::string_view name); static pointer create(Session& session, ObjectPtr<ClusterType> type, std::string_view name);
// Accessors // Accessors
const std::string& getName() const { return _name; } const std::string& getName() const { return _name; }
Wt::Dbo::ptr<ClusterType> getType() const { return _clusterType; } ObjectPtr<ClusterType> getType() const { return _clusterType; }
std::size_t getTracksCount() const { return _tracks.size(); } std::size_t getTracksCount() const { return _tracks.size(); }
std::vector<Wt::Dbo::ptr<Track>> getTracks(std::optional<std::size_t> offset = {}, std::optional<std::size_t> limit = {}) const; std::vector<ObjectPtr<Track>> getTracks(std::optional<std::size_t> offset = {}, std::optional<std::size_t> limit = {}) const;
std::set<IdType> getTrackIds() const; std::vector<TrackId> getTrackIds() const;
std::size_t getReleasesCount() const; std::size_t getReleasesCount() const;
void addTrack(Wt::Dbo::ptr<Track> track); void addTrack(ObjectPtr<Track> track);
template<class Action> template<class Action>
void persist(Action& a) void persist(Action& a)
@@ -72,7 +69,6 @@ class Cluster : public Wt::Dbo::Dbo<Cluster>
} }
private: private:
static const std::size_t _maxNameLength = 128; static const std::size_t _maxNameLength = 128;
std::string _name; std::string _name;
@@ -82,19 +78,17 @@ class Cluster : public Wt::Dbo::Dbo<Cluster>
}; };
class ClusterType : public Wt::Dbo::Dbo<ClusterType> class ClusterType : public Object<ClusterType, ClusterTypeId>
{ {
public: public:
ClusterType() = default;
ClusterType(std::string_view name);
using pointer = Wt::Dbo::ptr<ClusterType>; // Getters
ClusterType() {}
ClusterType(std::string name);
static std::vector<pointer> getAllOrphans(Session& session); static std::vector<pointer> getAllOrphans(Session& session);
static std::vector<pointer> getAllUsed(Session& session); static std::vector<pointer> getAllUsed(Session& session);
static pointer getByName(Session& session, const std::string& name); static pointer getByName(Session& session, const std::string& name);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, ClusterTypeId id);
static std::vector<pointer> getAll(Session& session); static std::vector<pointer> getAll(Session& session);
static pointer create(Session& session, const std::string& name); static pointer create(Session& session, const std::string& name);
+20 -23
View File
@@ -20,13 +20,13 @@
#pragma once #pragma once
#include <optional> #include <optional>
#include <set> #include <vector>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include "database/Types.hpp"
#include "utils/UUID.hpp" #include "utils/UUID.hpp"
#include "Types.hpp"
namespace Database namespace Database
{ {
@@ -39,49 +39,46 @@ class Session;
class Track; class Track;
class User; class User;
class Release : public Wt::Dbo::Dbo<Release> class Release : public Object<Release, ReleaseId>
{ {
public: public:
Release() = default;
using pointer = Wt::Dbo::ptr<Release>;
Release() {}
Release(const std::string& name, const std::optional<UUID>& MBID = {}); Release(const std::string& name, const std::optional<UUID>& MBID = {});
// Accessors // Accessors
static std::size_t getCount(Session& session); static std::size_t getCount(Session& session);
static pointer getByMBID(Session& session, const UUID& MBID); static pointer getByMBID(Session& session, const UUID& MBID);
static std::vector<pointer> getByName(Session& session, const std::string& name); static std::vector<pointer> getByName(Session& session, const std::string& name);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, ReleaseId id);
static std::vector<pointer> getAllOrphans(Session& session); // no track related static std::vector<pointer> getAllOrphans(Session& session); // no track related
static std::vector<pointer> getAll(Session& session, std::optional<Range> range = std::nullopt); static std::vector<pointer> getAll(Session& session, std::optional<Range> range = std::nullopt);
static std::vector<IdType> getAllIds(Session& session); static std::vector<ReleaseId> getAllIds(Session& session);
static std::vector<pointer> getAllOrderedByArtist(Session& session, std::optional<std::size_t> offset = {}, std::optional<std::size_t> size = {}); static std::vector<pointer> getAllOrderedByArtist(Session& session, std::optional<std::size_t> offset = {}, std::optional<std::size_t> size = {});
static std::vector<pointer> getAllRandom(Session& session, const std::set<IdType>& clusters, std::optional<std::size_t> size = {}); static std::vector<pointer> getAllRandom(Session& session, const std::vector<ClusterId>& clusters, std::optional<std::size_t> size = {});
static std::vector<IdType> getAllIdsRandom(Session& session, const std::set<IdType>& clusters, std::optional<std::size_t> size = {}); static std::vector<ReleaseId> getAllIdsRandom(Session& session, const std::vector<ClusterId>& clusters, std::optional<std::size_t> size = {});
static std::vector<pointer> getLastWritten(Session& session, std::optional<Wt::WDateTime> after, const std::set<IdType>& clusters, std::optional<Range> range, bool& moreResults); static std::vector<pointer> getLastWritten(Session& session, std::optional<Wt::WDateTime> after, const std::vector<ClusterId>& clusters, std::optional<Range> range, bool& moreResults);
static std::vector<pointer> getByYear(Session& session, int yearFrom, int yearTo, std::optional<Range> range = std::nullopt); static std::vector<pointer> getByYear(Session& session, int yearFrom, int yearTo, std::optional<Range> range = std::nullopt);
static std::vector<pointer> getStarred(Session& session, Wt::Dbo::ptr<User> user, const std::set<IdType>& clusters, std::optional<Range> range, bool& moreResults); static std::vector<pointer> getStarred(Session& session, ObjectPtr<User> user, const std::vector<ClusterId>& clusters, std::optional<Range> range, bool& moreResults);
static std::vector<pointer> getByClusters(Session& session, const std::set<IdType>& clusters); static std::vector<pointer> getByClusters(Session& session, const std::vector<ClusterId>& clusters);
static std::vector<pointer> getByFilter(Session& session, static std::vector<pointer> getByFilter(Session& session,
const std::set<IdType>& clusters, // if non empty, at least one release that belongs to these clusters const std::vector<ClusterId>& clusters, // if non empty, at least one release that belongs to these clusters
const std::vector<std::string_view>& keywords, // if non empty, name must match all of these keywords const std::vector<std::string_view>& keywords, // if non empty, name must match all of these keywords
std::optional<Range> range, std::optional<Range> range,
bool& moreExpected); bool& moreExpected);
static std::vector<IdType> getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit = {}); static std::vector<ReleaseId> getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit = {});
std::vector<Wt::Dbo::ptr<Track>> getTracks(const std::set<IdType>& clusters = std::set<IdType>()) const; std::vector<ObjectPtr<Track>> getTracks(const std::vector<ClusterId>& clusters = {}) const;
std::size_t getTracksCount() const; std::size_t getTracksCount() const;
Wt::Dbo::ptr<Track> getFirstTrack() const; ObjectPtr<Track> getFirstTrack() const;
// Get the cluster of the tracks that belong to this release // Get the cluster of the tracks that belong to this release
// Each clusters are grouped by cluster type, sorted by the number of occurence (max to min) // Each clusters are grouped by cluster type, sorted by the number of occurence (max to min)
// size is the max number of cluster per cluster type // size is the max number of cluster per cluster type
std::vector<std::vector<Wt::Dbo::ptr<Cluster>>> getClusterGroups(std::vector<Wt::Dbo::ptr<ClusterType>> clusterTypes, std::size_t size) const; std::vector<std::vector<ObjectPtr<Cluster>>> getClusterGroups(const std::vector<ObjectPtr<ClusterType>>& clusterTypes, std::size_t size) const;
// Create // Create
static pointer create(Session& session, const std::string& name, const std::optional<UUID>& MBID = {}); static pointer create(Session& session, const std::string& name, const std::optional<UUID>& MBID = {});
// Utility functions // Utility functions
std::optional<int> getReleaseYear(bool originalDate = false) const; // 0 if unknown or various std::optional<int> getReleaseYear(bool originalDate = false) const; // 0 if unknown or various
@@ -97,8 +94,8 @@ class Release : public Wt::Dbo::Dbo<Release>
Wt::WDateTime getLastWritten() const; Wt::WDateTime getLastWritten() const;
// Get the artists of this release // Get the artists of this release
std::vector<Wt::Dbo::ptr<Artist> > getArtists(TrackArtistLinkType type = TrackArtistLinkType::Artist) const; std::vector<ObjectPtr<Artist> > getArtists(TrackArtistLinkType type = TrackArtistLinkType::Artist) const;
std::vector<Wt::Dbo::ptr<Artist> > getReleaseArtists() const { return getArtists(TrackArtistLinkType::ReleaseArtist); } std::vector<ObjectPtr<Artist> > getReleaseArtists() const { return getArtists(TrackArtistLinkType::ReleaseArtist); }
bool hasVariousArtists() const; bool hasVariousArtists() const;
std::vector<pointer> getSimilarReleases(std::optional<std::size_t> offset = {}, std::optional<std::size_t> count = {}) const; std::vector<pointer> getSimilarReleases(std::optional<std::size_t> offset = {}, std::optional<std::size_t> count = {}) const;
@@ -19,23 +19,22 @@
#pragma once #pragma once
#include <unordered_set> #include <filesystem>
#include <vector>
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include <Wt/WTime.h> #include <Wt/WTime.h>
#include "utils/Path.hpp" #include "database/Types.hpp"
namespace Database { namespace Database {
class ClusterType; class ClusterType;
class Session; class Session;
class ScanSettings : public Wt::Dbo::Dbo<ScanSettings> class ScanSettings : public Object<ScanSettings, ScanSettingsId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<ScanSettings>;
// Do not modify values (just add) // Do not modify values (just add)
enum class UpdatePeriod { enum class UpdatePeriod {
Never = 0, Never = 0,
@@ -61,8 +60,8 @@ class ScanSettings : public Wt::Dbo::Dbo<ScanSettings>
std::filesystem::path getMediaDirectory() const { return _mediaDirectory; } std::filesystem::path getMediaDirectory() const { return _mediaDirectory; }
Wt::WTime getUpdateStartTime() const { return _startTime; } Wt::WTime getUpdateStartTime() const { return _startTime; }
UpdatePeriod getUpdatePeriod() const { return _updatePeriod; } UpdatePeriod getUpdatePeriod() const { return _updatePeriod; }
std::vector<Wt::Dbo::ptr<ClusterType>> getClusterTypes() const; std::vector<ObjectPtr<ClusterType>> getClusterTypes() const;
std::unordered_set<std::filesystem::path> getAudioFileExtensions() const; std::vector<std::filesystem::path> getAudioFileExtensions() const;
RecommendationEngineType getRecommendationEngineType() const { return _recommendationEngineType; } RecommendationEngineType getRecommendationEngineType() const { return _recommendationEngineType; }
// Setters // Setters
+29 -32
View File
@@ -34,7 +34,7 @@
#include "utils/EnumSet.hpp" #include "utils/EnumSet.hpp"
#include "utils/UUID.hpp" #include "utils/UUID.hpp"
#include "Types.hpp" #include "database/Types.hpp"
namespace Database { namespace Database {
@@ -49,46 +49,43 @@ class TrackListEntry;
class TrackStats; class TrackStats;
class User; class User;
class Track : public Wt::Dbo::Dbo<Track> class Track : public Object<Track, TrackId>
{ {
public: public:
Track() = default;
using pointer = Wt::Dbo::ptr<Track>;
Track() {}
Track(const std::filesystem::path& p); Track(const std::filesystem::path& p);
// Find utility functions // Find utility functions
static std::size_t getCount(Session& session); static std::size_t getCount(Session& session);
static pointer getByPath(Session& session, const std::filesystem::path& p); static pointer getByPath(Session& session, const std::filesystem::path& p);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, TrackId id);
static std::vector<pointer> getByRecordingMBID(Session& session, const UUID& MBID); static std::vector<pointer> getByRecordingMBID(Session& session, const UUID& MBID);
static std::vector<pointer> getSimilarTracks(Session& session, static std::vector<pointer> getSimilarTracks(Session& session,
const std::unordered_set<IdType>& trackIds, const std::vector<TrackId>& trackIds,
std::optional<std::size_t> offset = {}, std::optional<std::size_t> offset = {},
std::optional<std::size_t> size = {}); std::optional<std::size_t> size = {});
static std::vector<pointer> getByClusters(Session& session, static std::vector<pointer> getByClusters(Session& session,
const std::set<IdType>& clusters); // tracks that belong to these clusters const std::vector<ClusterId>& clusters); // tracks that belong to these clusters
static std::vector<pointer> getByFilter(Session& session, static std::vector<pointer> getByFilter(Session& session,
const std::set<IdType>& clusters, // if non empty, tracks that belong to these clusters const std::vector<ClusterId>& clusters, // if non empty, tracks that belong to these clusters
const std::vector<std::string_view>& keywords, // if non empty, name must match all of these keywords const std::vector<std::string_view>& keywords, // if non empty, name must match all of these keywords
std::optional<Range> range, std::optional<Range> range,
bool& moreExpected); bool& moreExpected);
static std::vector<pointer> getByNameAndReleaseName(Session& session, std::string_view trackName, std::string_view releaseName); static std::vector<pointer> getByNameAndReleaseName(Session& session, std::string_view trackName, std::string_view releaseName);
static std::vector<pointer> getAll(Session& session, std::optional<std::size_t> limit = std::nullopt); static std::vector<pointer> getAll(Session& session, std::optional<std::size_t> limit = std::nullopt);
static std::vector<pointer> getAllRandom(Session& session, const std::set<IdType>& clusters, std::optional<std::size_t> limit = std::nullopt); static std::vector<pointer> getAllRandom(Session& session, const std::vector<ClusterId>& clusters, std::optional<std::size_t> limit = std::nullopt);
static std::vector<IdType> getAllIdsRandom(Session& session, const std::set<IdType>& clusters, std::optional<std::size_t> limit = std::nullopt); static std::vector<TrackId> getAllIdsRandom(Session& session, const std::vector<ClusterId>& clusters, std::optional<std::size_t> limit = std::nullopt);
static std::vector<IdType> getAllIds(Session& session); static std::vector<TrackId> getAllIds(Session& session);
static std::vector<std::pair<IdType, std::filesystem::path>> getAllPaths(Session& session, std::optional<std::size_t> offset = std::nullopt, std::optional<std::size_t> size = std::nullopt); static std::vector<std::pair<TrackId, std::filesystem::path>> getAllPaths(Session& session, std::optional<std::size_t> offset = std::nullopt, std::optional<std::size_t> size = std::nullopt);
static std::vector<pointer> getMBIDDuplicates(Session& session); static std::vector<pointer> getMBIDDuplicates(Session& session);
static std::vector<pointer> getLastWritten(Session& session, std::optional<Wt::WDateTime> after, const std::set<IdType>& clusters, std::optional<Range> range, bool& moreResults); static std::vector<pointer> getLastWritten(Session& session, std::optional<Wt::WDateTime> after, const std::vector<ClusterId>& clusters, std::optional<Range> range, bool& moreResults);
static std::vector<pointer> getAllWithRecordingMBIDAndMissingFeatures(Session& session); static std::vector<pointer> getAllWithRecordingMBIDAndMissingFeatures(Session& session);
static std::vector<IdType> getAllIdsWithFeatures(Session& session, std::optional<std::size_t> limit = {}); static std::vector<TrackId> getAllIdsWithFeatures(Session& session, std::optional<std::size_t> limit = {});
static std::vector<IdType> getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit = {}); static std::vector<TrackId> getAllIdsWithClusters(Session& session, std::optional<std::size_t> limit = {});
static std::vector<pointer> getStarred(Session& session, static std::vector<pointer> getStarred(Session& session,
Wt::Dbo::ptr<User> user, ObjectPtr<User> user,
const std::set<IdType>& clusters, const std::vector<ClusterId>& clusters,
std::optional<Range> range, bool& hasMore); std::optional<Range> range, bool& hasMore);
// Create utility // Create utility
@@ -115,10 +112,10 @@ class Track : public Wt::Dbo::Dbo<Track>
void setTrackReplayGain(float replayGain) { _trackReplayGain = replayGain; } void setTrackReplayGain(float replayGain) { _trackReplayGain = replayGain; }
void setReleaseReplayGain(float replayGain) { _releaseReplayGain = replayGain; } void setReleaseReplayGain(float replayGain) { _releaseReplayGain = replayGain; }
void clearArtistLinks(); void clearArtistLinks();
void addArtistLink(const Wt::Dbo::ptr<TrackArtistLink>& artistLink); void addArtistLink(const ObjectPtr<TrackArtistLink>& artistLink);
void setRelease(Wt::Dbo::ptr<Release> release) { _release = release; } void setRelease(ObjectPtr<Release> release) { _release = getDboPtr(release); }
void setClusters(const std::vector<Wt::Dbo::ptr<Cluster>>& clusters ); void setClusters(const std::vector<ObjectPtr<Cluster>>& clusters );
void setFeatures(const Wt::Dbo::ptr<TrackFeatures>& features); void setFeatures(const ObjectPtr<TrackFeatures>& features);
std::size_t getScanVersion() const { return _scanVersion; } std::size_t getScanVersion() const { return _scanVersion; }
std::optional<std::size_t> getTrackNumber() const; std::optional<std::size_t> getTrackNumber() const;
@@ -143,16 +140,16 @@ class Track : public Wt::Dbo::Dbo<Track>
std::optional<float> getReleaseReplayGain() const { return _releaseReplayGain; } std::optional<float> getReleaseReplayGain() const { return _releaseReplayGain; }
// no artistLinkTypes means get all // no artistLinkTypes means get all
std::vector<Wt::Dbo::ptr<Artist>> getArtists(EnumSet<TrackArtistLinkType> artistLinkTypes) const; std::vector<ObjectPtr<Artist>> getArtists(EnumSet<TrackArtistLinkType> artistLinkTypes) const;
std::vector<IdType> getArtistIds(EnumSet<TrackArtistLinkType> artistLinkTypes) const; std::vector<ArtistId> getArtistIds(EnumSet<TrackArtistLinkType> artistLinkTypes) const;
std::vector<Wt::Dbo::ptr<TrackArtistLink>> getArtistLinks() const; std::vector<ObjectPtr<TrackArtistLink>> getArtistLinks() const;
Wt::Dbo::ptr<Release> getRelease() const { return _release; } ObjectPtr<Release> getRelease() const { return _release; }
std::vector<Wt::Dbo::ptr<Cluster>> getClusters() const; std::vector<ObjectPtr<Cluster>> getClusters() const;
std::vector<IdType> getClusterIds() const; std::vector<ClusterId> getClusterIds() const;
bool hasTrackFeatures() const; bool hasTrackFeatures() const;
Wt::Dbo::ptr<TrackFeatures> getTrackFeatures() const; ObjectPtr<TrackFeatures> getTrackFeatures() const;
std::vector<std::vector<Wt::Dbo::ptr<Cluster>>> getClusterGroups(std::vector<Wt::Dbo::ptr<ClusterType>> clusterTypes, std::size_t size) const; std::vector<std::vector<ObjectPtr<Cluster>>> getClusterGroups(const std::vector<ObjectPtr<ClusterType>>& clusterTypes, std::size_t size) const;
template<class Action> template<class Action>
void persist(Action& a) void persist(Action& a)
@@ -23,7 +23,7 @@
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include "Types.hpp" #include "database/Types.hpp"
#include "utils/EnumSet.hpp" #include "utils/EnumSet.hpp"
namespace Database namespace Database
@@ -33,20 +33,18 @@ namespace Database
class Session; class Session;
class Track; class Track;
class TrackArtistLink class TrackArtistLink : public Object<TrackArtistLink, TrackArtistLinkId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<TrackArtistLink>;
TrackArtistLink() = default; TrackArtistLink() = default;
TrackArtistLink(Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<Artist> artist, TrackArtistLinkType type); TrackArtistLink(ObjectPtr<Track> track, ObjectPtr<Artist> artist, TrackArtistLinkType type);
static pointer create(Session& session, Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<Artist> artist, TrackArtistLinkType type); static pointer create(Session& session, ObjectPtr<Track> track, ObjectPtr<Artist> artist, TrackArtistLinkType type);
static EnumSet<TrackArtistLinkType> getUsedTypes(Session& session); static EnumSet<TrackArtistLinkType> getUsedTypes(Session& session);
Wt::Dbo::ptr<Track> getTrack() const { return _track; } ObjectPtr<Track> getTrack() const { return _track; }
Wt::Dbo::ptr<Artist> getArtist() const { return _artist; } ObjectPtr<Artist> getArtist() const { return _artist; }
TrackArtistLinkType getType() const { return _type; } TrackArtistLinkType getType() const { return _type; }
template<class Action> template<class Action>
@@ -23,7 +23,7 @@
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include "Types.hpp" #include "database/Types.hpp"
namespace Database { namespace Database {
@@ -31,22 +31,20 @@ class Session;
class Track; class Track;
class User; class User;
class TrackBookmark : public Wt::Dbo::Dbo<TrackBookmark> class TrackBookmark : public Object<TrackBookmark, TrackBookmarkId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<TrackBookmark>;
TrackBookmark () = default; TrackBookmark () = default;
TrackBookmark(Wt::Dbo::ptr<User> user, Wt::Dbo::ptr<Track> track); TrackBookmark(ObjectPtr<User> user, ObjectPtr<Track> track);
// utility // utility
static pointer create(Session& session, Wt::Dbo::ptr<User> user, Wt::Dbo::ptr<Track> track); static pointer create(Session& session, ObjectPtr<User> user, ObjectPtr<Track> track);
// Find utility functions // Find utility functions
static std::vector<pointer> getAll(Session& session); static std::vector<pointer> getAll(Session& session);
static std::vector<pointer> getByUser(Session& session, Wt::Dbo::ptr<User> user); static std::vector<pointer> getByUser(Session& session, ObjectPtr<User> user);
static pointer getByUser(Session& session, Wt::Dbo::ptr<User> user, Wt::Dbo::ptr<Track> track); static pointer getByUser(Session& session, ObjectPtr<User> user, ObjectPtr<Track> track);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, TrackBookmarkId id);
// Setters // Setters
void setOffset(std::chrono::milliseconds offset) { _offset = offset; } void setOffset(std::chrono::milliseconds offset) { _offset = offset; }
@@ -55,8 +53,8 @@ class TrackBookmark : public Wt::Dbo::Dbo<TrackBookmark>
// Getters // Getters
std::chrono::milliseconds getOffset() const { return _offset; } std::chrono::milliseconds getOffset() const { return _offset; }
std::string_view getComment() const { return _comment; } std::string_view getComment() const { return _comment; }
Wt::Dbo::ptr<Track> getTrack() const { return _track; } ObjectPtr<Track> getTrack() const { return _track; }
Wt::Dbo::ptr<User> getUser() const { return _user; } ObjectPtr<User> getUser() const { return _user; }
template<class Action> template<class Action>
void persist(Action& a) void persist(Action& a)
@@ -26,7 +26,7 @@
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include "Types.hpp" #include "database/Types.hpp"
namespace Database { namespace Database {
@@ -37,17 +37,14 @@ using FeatureName = std::string;
using FeatureValues = std::vector<double>; using FeatureValues = std::vector<double>;
using FeatureValuesMap = std::unordered_map<FeatureName, FeatureValues>; using FeatureValuesMap = std::unordered_map<FeatureName, FeatureValues>;
class TrackFeatures : public Wt::Dbo::Dbo<TrackFeatures> class TrackFeatures : public Object<TrackFeatures, TrackFeaturesId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<TrackFeatures>;
TrackFeatures() = default; TrackFeatures() = default;
TrackFeatures(Wt::Dbo::ptr<Track> track, const std::string& jsonEncodedFeatures); TrackFeatures(ObjectPtr<Track> track, const std::string& jsonEncodedFeatures);
// Create utility // Create utility
static pointer create(Session& session, Wt::Dbo::ptr<Track> track, const std::string& jsonEncodedFeatures); static pointer create(Session& session, ObjectPtr<Track> track, const std::string& jsonEncodedFeatures);
FeatureValues getFeatureValues(const FeatureName& feature) const; FeatureValues getFeatureValues(const FeatureName& feature) const;
FeatureValuesMap getFeatureValuesMap(const std::unordered_set<FeatureName>& featureNames) const; FeatureValuesMap getFeatureValuesMap(const std::unordered_set<FeatureName>& featureNames) const;
@@ -27,7 +27,7 @@
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include "Types.hpp" #include "database/Types.hpp"
namespace Database { namespace Database {
@@ -39,11 +39,9 @@ class Track;
class TrackListEntry; class TrackListEntry;
class User; class User;
class TrackList : public Wt::Dbo::Dbo<TrackList> class TrackList : public Object<TrackList, TrackListId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<TrackList>;
enum class Type enum class Type
{ {
Playlist, // user controlled playlists Playlist, // user controlled playlists
@@ -51,28 +49,28 @@ class TrackList : public Wt::Dbo::Dbo<TrackList>
}; };
TrackList() = default; TrackList() = default;
TrackList(std::string_view name, Type type, bool isPublic, Wt::Dbo::ptr<User> user); TrackList(std::string_view name, Type type, bool isPublic, ObjectPtr<User> user);
// Stats utility // Stats utility
std::vector<Wt::Dbo::ptr<Artist>> getTopArtists(const std::set<IdType>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Artist>> getTopArtists(const std::vector<ClusterId>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const;
std::vector<Wt::Dbo::ptr<Release>> getTopReleases(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Release>> getTopReleases(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const;
std::vector<Wt::Dbo::ptr<Track>> getTopTracks(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Track>> getTopTracks(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const;
// Search utility // Search utility
static pointer get(Session& session, std::string_view name, Type type, Wt::Dbo::ptr<User> user); static pointer get(Session& session, std::string_view name, Type type, ObjectPtr<User> user);
static pointer getById(Session& session, IdType tracklistId); static pointer getById(Session& session, TrackListId tracklistId);
static std::vector<pointer> getAll(Session& session); static std::vector<pointer> getAll(Session& session);
static std::vector<pointer> getAll(Session& session, Wt::Dbo::ptr<User> user); static std::vector<pointer> getAll(Session& session, ObjectPtr<User> user);
static std::vector<pointer> getAll(Session& session, Wt::Dbo::ptr<User> user, Type type); static std::vector<pointer> getAll(Session& session, ObjectPtr<User> user, Type type);
// Create utility // Create utility
static pointer create(Session& session, std::string_view name, Type type, bool isPublic, Wt::Dbo::ptr<User> user); static pointer create(Session& session, std::string_view name, Type type, bool isPublic, ObjectPtr<User> user);
// Accessors // Accessors
std::string getName() const { return _name; } std::string getName() const { return _name; }
bool isPublic() const { return _isPublic; } bool isPublic() const { return _isPublic; }
Type getType() const { return _type; } Type getType() const { return _type; }
Wt::Dbo::ptr<User> getUser() const { return _user; } ObjectPtr<User> getUser() const { return _user; }
// Modifiers // Modifiers
void setName(const std::string& name) { _name = name; } void setName(const std::string& name) { _name = name; }
@@ -80,29 +78,29 @@ class TrackList : public Wt::Dbo::Dbo<TrackList>
void clear() { _entries.clear(); } void clear() { _entries.clear(); }
// Get tracks, ordered by position // Get tracks, ordered by position
bool isEmpty() const; bool isEmpty() const;
std::size_t getCount() const; std::size_t getCount() const;
Wt::Dbo::ptr<TrackListEntry> getEntry(std::size_t pos) const; ObjectPtr<TrackListEntry> getEntry(std::size_t pos) const;
std::vector<Wt::Dbo::ptr<TrackListEntry>> getEntries(std::optional<std::size_t> offset = {}, std::optional<std::size_t> size = {}) const; std::vector<ObjectPtr<TrackListEntry>> getEntries(std::optional<std::size_t> offset = {}, std::optional<std::size_t> size = {}) const;
Wt::Dbo::ptr<TrackListEntry> getEntryByTrackAndDateTime(Wt::Dbo::ptr<Track> track, const Wt::WDateTime& dateTime) const; ObjectPtr<TrackListEntry> getEntryByTrackAndDateTime(ObjectPtr<Track> track, const Wt::WDateTime& dateTime) const;
// Get track bya // Get track bya
std::vector<Wt::Dbo::ptr<Artist>> getArtistsReverse(const std::set<IdType>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Artist>> getArtistsReverse(const std::vector<ClusterId>& clusterIds, std::optional<TrackArtistLinkType> linkType, std::optional<Range> range, bool& moreResults) const;
std::vector<Wt::Dbo::ptr<Release>> getReleasesReverse(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Release>> getReleasesReverse(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const;
std::vector<Wt::Dbo::ptr<Track>> getTracksReverse(const std::set<IdType>& clusterIds, std::optional<Range> range, bool& moreResults) const; std::vector<ObjectPtr<Track>> getTracksReverse(const std::vector<ClusterId>& clusterIds, std::optional<Range> range, bool& moreResults) const;
std::vector<IdType> getTrackIds() const; std::vector<TrackId> getTrackIds() const;
std::chrono::milliseconds getDuration() const; std::chrono::milliseconds getDuration() const;
// Get clusters, order by occurence // Get clusters, order by occurence
std::vector<Wt::Dbo::ptr<Cluster>> getClusters() const; std::vector<ObjectPtr<Cluster>> getClusters() const;
bool hasTrack(IdType trackId) const; bool hasTrack(TrackId trackId) const;
// Ordered from most clusters in common // Ordered from most clusters in common
std::vector<Wt::Dbo::ptr<Track>> getSimilarTracks(std::optional<std::size_t> offset = {}, std::optional<std::size_t> size = {}) const; std::vector<ObjectPtr<Track>> getSimilarTracks(std::optional<std::size_t> offset = {}, std::optional<std::size_t> size = {}) const;
template<class Action> template<class Action>
void persist(Action& a) void persist(Action& a)
@@ -122,27 +120,24 @@ class TrackList : public Wt::Dbo::Dbo<TrackList>
bool _isPublic {false}; bool _isPublic {false};
Wt::Dbo::ptr<User> _user; Wt::Dbo::ptr<User> _user;
Wt::Dbo::collection< Wt::Dbo::ptr<TrackListEntry> > _entries; Wt::Dbo::collection<Wt::Dbo::ptr<TrackListEntry>> _entries;
}; };
class TrackListEntry : public Wt::Dbo::Dbo<TrackListEntry> class TrackListEntry : public Object<TrackListEntry, TrackListEntryId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<TrackListEntry>;
TrackListEntry() = default; TrackListEntry() = default;
TrackListEntry(Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<TrackList> tracklist, const Wt::WDateTime& dateTime); TrackListEntry(ObjectPtr<Track> track, ObjectPtr<TrackList> tracklist, const Wt::WDateTime& dateTime);
// find utility // find utility
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, TrackListEntryId id);
// Create utility // Create utility
static pointer create(Session& session, Wt::Dbo::ptr<Track> track, Wt::Dbo::ptr<TrackList> tracklist, const Wt::WDateTime& dateTime = Wt::WDateTime::currentDateTime()); static pointer create(Session& session, ObjectPtr<Track> track, ObjectPtr<TrackList> tracklist, const Wt::WDateTime& dateTime = Wt::WDateTime::currentDateTime());
// Accessors // Accessors
Wt::Dbo::ptr<Track> getTrack() const { return _track; } ObjectPtr<Track> getTrack() const { return _track; }
const Wt::WDateTime& getDateTime() const { return _dateTime; } const Wt::WDateTime& getDateTime() const { return _dateTime; }
template<class Action> template<class Action>
+99 -5
View File
@@ -20,16 +20,32 @@
#pragma once #pragma once
#include <cstdint> #include <cstdint>
#include <cassert>
#include <functional>
#include <Wt/Dbo/ptr.h> #include <Wt/Dbo/ptr.h>
namespace Database namespace Database
{ {
using IdType = Wt::Dbo::dbo_default_traits::IdType; class IdType
static inline bool IdIsValid(IdType id)
{ {
return id != Wt::Dbo::dbo_default_traits::invalidId(); public:
} using ValueType = Wt::Dbo::dbo_default_traits::IdType;
IdType() = default;
IdType(ValueType id) : _id {id} { assert(isValid()); }
bool isValid() const { return _id != Wt::Dbo::dbo_default_traits::invalidId(); }
std::string toString() const { assert(isValid()); return std::to_string(_id); }
ValueType getValue() const { return _id; }
bool operator==(IdType other) const { return other._id == _id; }
bool operator!=(IdType other) const { return !(*this == other); }
bool operator<(IdType other) const { return other._id < _id; }
private:
Wt::Dbo::dbo_default_traits::IdType _id {Wt::Dbo::dbo_default_traits::invalidId()};
};
struct Range struct Range
{ {
@@ -78,5 +94,83 @@ namespace Database
ADMIN = 1, ADMIN = 1,
DEMO = 2, DEMO = 2,
}; };
template <typename T>
class ObjectPtr
{
public:
ObjectPtr() = default;
ObjectPtr(Wt::Dbo::ptr<T> obj) : _obj {obj} {}
const T* operator->() const { return _obj.get(); }
operator bool() const { return _obj.get(); }
bool operator!() const { return !_obj.get(); }
auto modify() { return _obj.modify(); }
void remove() { _obj.remove(); }
private:
template <typename, typename> friend class Object;
Wt::Dbo::ptr<T> _obj;
};
template <typename T, typename ObjectIdType>
class Object : public Wt::Dbo::Dbo<T>
{
static_assert(std::is_base_of_v<Database::IdType, ObjectIdType>);
static_assert(!std::is_same_v<Database::IdType, ObjectIdType>);
public:
using pointer = ObjectPtr<T>;
using IdType = ObjectIdType;
IdType getId() const { return Wt::Dbo::Dbo<T>::self()->Wt::Dbo::Dbo<T>::id(); }
// catch some misuses
typename Wt::Dbo::dbo_traits<T>::IdType id() const = delete;
protected:
// Can get raw dbo ptr only from Objects
template <typename SomeObject>
static
Wt::Dbo::ptr<SomeObject> getDboPtr(ObjectPtr<SomeObject> ptr) { return ptr._obj; }
};
} }
// TODO factorize hash with std::enable_if
#define LMS_DECLARE_IDTYPE(name) \
namespace Database { \
class name : public IdType \
{ \
public: \
using IdType::IdType; \
};\
} \
namespace std \
{ \
template<> \
class hash<Database::name> \
{ \
public: \
size_t operator()(Database::name id) const \
{ \
return std::hash<Database::name::ValueType>()(id.getValue()); \
} \
}; \
} // ns std
LMS_DECLARE_IDTYPE(ArtistId)
LMS_DECLARE_IDTYPE(AuthTokenId)
LMS_DECLARE_IDTYPE(ClusterId)
LMS_DECLARE_IDTYPE(ClusterTypeId)
LMS_DECLARE_IDTYPE(ReleaseId)
LMS_DECLARE_IDTYPE(ScanSettingsId)
LMS_DECLARE_IDTYPE(TrackArtistLinkId)
LMS_DECLARE_IDTYPE(TrackBookmarkId)
LMS_DECLARE_IDTYPE(TrackFeaturesId)
LMS_DECLARE_IDTYPE(TrackId)
LMS_DECLARE_IDTYPE(TrackListId)
LMS_DECLARE_IDTYPE(TrackListEntryId)
LMS_DECLARE_IDTYPE(UserId)
+19 -25
View File
@@ -26,8 +26,8 @@
#include <Wt/Dbo/Dbo.h> #include <Wt/Dbo/Dbo.h>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include "database/Types.hpp"
#include "utils/UUID.hpp" #include "utils/UUID.hpp"
#include "Types.hpp"
namespace Database { namespace Database {
@@ -39,24 +39,21 @@ class TrackList;
class Track; class Track;
class User; class User;
class AuthToken class AuthToken : public Object<AuthToken, AuthTokenId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<AuthToken>;
AuthToken() = default; AuthToken() = default;
AuthToken(const std::string& value, const Wt::WDateTime& expiry, Wt::Dbo::ptr<User> user); AuthToken(const std::string& value, const Wt::WDateTime& expiry, ObjectPtr<User> user);
// Utility // Utility
static pointer create(Session& session, const std::string& value, const Wt::WDateTime&expiry, Wt::Dbo::ptr<User> user); static pointer create(Session& session, const std::string& value, const Wt::WDateTime&expiry, ObjectPtr<User> user);
static void removeExpiredTokens(Session& session, const Wt::WDateTime& now); static void removeExpiredTokens(Session& session, const Wt::WDateTime& now);
static pointer getByValue(Session& session, const std::string& value); static pointer getByValue(Session& session, const std::string& value);
static pointer getById(Session& session, IdType tokenId); static pointer getById(Session& session, AuthTokenId tokenId);
// Accessors // Accessors
const Wt::WDateTime& getExpiry() const { return _expiry; } const Wt::WDateTime& getExpiry() const { return _expiry; }
Wt::Dbo::ptr<User> getUser() const { return _user; } ObjectPtr<User> getUser() const { return _user; }
const std::string& getValue() const { return _value; } const std::string& getValue() const { return _value; }
template<class Action> template<class Action>
@@ -75,11 +72,9 @@ class AuthToken
Wt::Dbo::ptr<User> _user; Wt::Dbo::ptr<User> _user;
}; };
class User : public Wt::Dbo::Dbo<User> class User : public Object<User, UserId>
{ {
public: public:
using pointer = Wt::Dbo::ptr<User>;
struct PasswordHash struct PasswordHash
{ {
std::string salt; std::string salt;
@@ -120,17 +115,16 @@ class User : public Wt::Dbo::Dbo<User>
static inline const SubsonicArtistListMode defaultSubsonicArtistListMode {SubsonicArtistListMode::AllArtists}; static inline const SubsonicArtistListMode defaultSubsonicArtistListMode {SubsonicArtistListMode::AllArtists};
static inline const Scrobbler defaultScrobbler {Scrobbler::Internal}; static inline const Scrobbler defaultScrobbler {Scrobbler::Internal};
User() = default; User() = default;
User(std::string_view loginName); User(std::string_view loginName);
// utility // utility
static pointer create(Session& session, std::string_view loginName); static pointer create(Session& session, std::string_view loginName);
static pointer getById(Session& session, IdType id); static pointer getById(Session& session, UserId id);
static pointer getByLoginName(Session& session, std::string_view loginName); static pointer getByLoginName(Session& session, std::string_view loginName);
static std::vector<pointer> getAll(Session& session); static std::vector<pointer> getAll(Session& session);
static std::vector<IdType> getAllIds(Session& session); static std::vector<UserId> getAllIds(Session& session);
static pointer getDemo(Session& session); static pointer getDemo(Session& session);
static std::size_t getCount(Session& session); static std::size_t getCount(Session& session);
@@ -171,20 +165,20 @@ class User : public Wt::Dbo::Dbo<User>
Scrobbler getScrobbler() const { return _scrobbler; } Scrobbler getScrobbler() const { return _scrobbler; }
std::optional<UUID> getListenBrainzToken() const { return UUID::fromString(_listenbrainzToken); } std::optional<UUID> getListenBrainzToken() const { return UUID::fromString(_listenbrainzToken); }
Wt::Dbo::ptr<TrackList> getQueuedTrackList(Session& session) const; ObjectPtr<TrackList> getQueuedTrackList(Session& session) const;
void starArtist(Wt::Dbo::ptr<Artist> artist); void starArtist(ObjectPtr<Artist> artist);
void unstarArtist(Wt::Dbo::ptr<Artist> artist); void unstarArtist(ObjectPtr<Artist> artist);
bool hasStarredArtist(Wt::Dbo::ptr<Artist> artist) const; bool hasStarredArtist(ObjectPtr<Artist> artist) const;
void starRelease(Wt::Dbo::ptr<Release> release); void starRelease(ObjectPtr<Release> release);
void unstarRelease(Wt::Dbo::ptr<Release> release); void unstarRelease(ObjectPtr<Release> release);
bool hasStarredRelease(Wt::Dbo::ptr<Release> release) const; bool hasStarredRelease(ObjectPtr<Release> release) const;
// Stars // Stars
void starTrack(Wt::Dbo::ptr<Track> track); void starTrack(ObjectPtr<Track> track);
void unstarTrack(Wt::Dbo::ptr<Track> track); void unstarTrack(ObjectPtr<Track> track);
bool hasStarredTrack(Wt::Dbo::ptr<Track> track) const; bool hasStarredTrack(ObjectPtr<Track> track) const;
template<class Action> template<class Action>
void persist(Action& a) void persist(Action& a)
+2 -2
View File
@@ -1,8 +1,8 @@
add_library(lmsrecommendation SHARED add_library(lmsrecommendation SHARED
impl/clusters/ClustersClassifier.cpp impl/clusters/ClustersClassifier.cpp
impl/features/FeaturesClassifierCache.cpp impl/features/FeaturesEngineCache.cpp
impl/features/FeaturesClassifier.cpp impl/features/FeaturesEngine.cpp
impl/features/FeaturesDefs.cpp impl/features/FeaturesDefs.cpp
impl/Engine.cpp impl/Engine.cpp
) )
+14 -17
View File
@@ -23,7 +23,7 @@
#include <vector> #include <vector>
#include "ClustersClassifierCreator.hpp" #include "ClustersClassifierCreator.hpp"
#include "FeaturesClassifierCreator.hpp" #include "FeaturesEngineCreator.hpp"
#include "database/Db.hpp" #include "database/Db.hpp"
#include "database/Session.hpp" #include "database/Session.hpp"
@@ -45,7 +45,7 @@ createClassifier(ClassifierType type)
break; break;
case ClassifierType::Features: case ClassifierType::Features:
return createFeaturesClassifier(); return createFeaturesEngine();
break; break;
} }
@@ -63,10 +63,10 @@ Engine::Engine(Database::Db& db)
{ {
} }
std::unordered_set<Database::IdType> Engine::TrackContainer
Engine::getSimilarTracksFromTrackList(Database::Session& session, Database::IdType trackListId, std::size_t maxCount) Engine::getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId trackListId, std::size_t maxCount)
{ {
std::unordered_set<Database::IdType> res; TrackContainer res;
std::shared_lock lock {_classifiersMutex}; std::shared_lock lock {_classifiersMutex};
for (const auto& classifierName : _classifierPriorities) for (const auto& classifierName : _classifierPriorities)
@@ -83,10 +83,10 @@ Engine::getSimilarTracksFromTrackList(Database::Session& session, Database::IdTy
return res; return res;
} }
std::unordered_set<Database::IdType> Engine::TrackContainer
Engine::getSimilarTracks(Database::Session& dbSession, const std::unordered_set<Database::IdType>& trackIds, std::size_t maxCount) Engine::getSimilarTracks(Database::Session& dbSession, const std::vector<Database::TrackId>& trackIds, std::size_t maxCount)
{ {
std::unordered_set<Database::IdType> res; TrackContainer res;
std::shared_lock lock {_classifiersMutex}; std::shared_lock lock {_classifiersMutex};
for (ClassifierType classifierType : _classifierPriorities) for (ClassifierType classifierType : _classifierPriorities)
@@ -107,10 +107,10 @@ Engine::getSimilarTracks(Database::Session& dbSession, const std::unordered_set<
return res; return res;
} }
std::unordered_set<Database::IdType> Engine::ReleaseContainer
Engine::getSimilarReleases(Database::Session& dbSession, Database::IdType releaseId, std::size_t maxCount) Engine::getSimilarReleases(Database::Session& dbSession, Database::ReleaseId releaseId, std::size_t maxCount)
{ {
std::unordered_set<Database::IdType> res; ReleaseContainer res;
std::shared_lock lock {_classifiersMutex}; std::shared_lock lock {_classifiersMutex};
for (ClassifierType classifierType : _classifierPriorities) for (ClassifierType classifierType : _classifierPriorities)
@@ -131,13 +131,10 @@ Engine::getSimilarReleases(Database::Session& dbSession, Database::IdType releas
return res; return res;
} }
std::unordered_set<Database::IdType> Engine::ArtistContainer
Engine::getSimilarArtists(Database::Session& dbSession, Engine::getSimilarArtists(Database::Session& dbSession, Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount)
Database::IdType artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes,
std::size_t maxCount)
{ {
std::unordered_set<Database::IdType> res; ArtistContainer res;
std::shared_lock lock {_classifiersMutex}; std::shared_lock lock {_classifiersMutex};
for (ClassifierType classifierType : _classifierPriorities) for (ClassifierType classifierType : _classifierPriorities)
+6 -6
View File
@@ -55,12 +55,13 @@ namespace Recommendation
private: private:
void load(bool forceReload, const ProgressCallback& progressCallback) override; void load(bool forceReload, const ProgressCallback& progressCallback) override;
void cancelLoad() override; void cancelLoad() override;
void requestCancelLoad() override {};
ResultContainer getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) override; ResultContainer<Database::TrackId> getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId tracklistId, std::size_t maxCount) override;
ResultContainer getSimilarTracks(Database::Session& session, const std::unordered_set<Database::IdType>& tracksId, std::size_t maxCount) override; ResultContainer<Database::TrackId> getSimilarTracks(Database::Session& session, const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) override;
ResultContainer getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) override; ResultContainer<Database::ReleaseId> getSimilarReleases(Database::Session& session, Database::ReleaseId releaseId, std::size_t maxCount) override;
ResultContainer getSimilarArtists(Database::Session& session, ResultContainer<Database::ArtistId> getSimilarArtists(Database::Session& session,
Database::IdType artistId, Database::ArtistId artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes, EnumSet<Database::TrackArtistLinkType> linkTypes,
std::size_t maxCount) override; std::size_t maxCount) override;
@@ -79,7 +80,6 @@ namespace Recommendation
using ClassifierContainer = std::unordered_map<ClassifierType, std::unique_ptr<IClassifier>>; using ClassifierContainer = std::unordered_map<ClassifierType, std::unique_ptr<IClassifier>>;
ClassifierContainer _classifiers; ClassifierContainer _classifiers;
std::vector<ClassifierType> _classifierPriorities; // ordered by priority std::vector<ClassifierType> _classifierPriorities; // ordered by priority
}; };
} // ns Recommendation } // ns Recommendation
@@ -24,6 +24,6 @@
namespace Recommendation namespace Recommendation
{ {
std::unique_ptr<IClassifier> createFeaturesClassifier(); std::unique_ptr<IClassifier> createFeaturesEngine();
} }
+10 -8
View File
@@ -21,9 +21,10 @@
#include <functional> #include <functional>
#include <string_view> #include <string_view>
#include <unordered_set> #include <vector>
#include "database/Types.hpp" #include "database/Types.hpp"
#include "recommendation/IRecommendation.hpp"
#include "utils/EnumSet.hpp" #include "utils/EnumSet.hpp"
namespace Database namespace Database
@@ -34,7 +35,7 @@ namespace Database
namespace Recommendation namespace Recommendation
{ {
class IClassifier class IClassifier : public IRecommendation
{ {
public: public:
virtual ~IClassifier() = default; virtual ~IClassifier() = default;
@@ -50,13 +51,14 @@ namespace Recommendation
virtual bool load(Database::Session& session, bool forceReload, const ProgressCallback& progressCallback) = 0; virtual bool load(Database::Session& session, bool forceReload, const ProgressCallback& progressCallback) = 0;
virtual void requestCancelLoad() = 0; virtual void requestCancelLoad() = 0;
using ResultContainer = std::unordered_set<Database::IdType>; template <typename IdType>
using ResultContainer = std::vector<IdType>;
virtual ResultContainer getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const = 0; virtual ResultContainer<Database::TrackId> getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId tracklistId, std::size_t maxCount) const = 0;
virtual ResultContainer getSimilarTracks(Database::Session& session, const std::unordered_set<Database::IdType>& tracksId, std::size_t maxCount) const = 0; virtual ResultContainer<Database::TrackId> getSimilarTracks(Database::Session& session, const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const = 0;
virtual ResultContainer getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) const = 0; virtual ResultContainer<Database::ReleaseId> getSimilarReleases(Database::Session& session, Database::ReleaseId releaseId, std::size_t maxCount) const = 0;
virtual ResultContainer getSimilarArtists(Database::Session& session, virtual ResultContainer<Database::ArtistId> getSimilarArtists(Database::Session& session,
Database::IdType artistId, Database::ArtistId artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const = 0; EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const = 0;
}; };
@@ -33,23 +33,23 @@ std::unique_ptr<IClassifier> createClustersClassifier()
return std::make_unique<ClusterClassifier>(); return std::make_unique<ClusterClassifier>();
} }
std::unordered_set<Database::IdType> IClassifier::ResultContainer<Database::TrackId>
ClusterClassifier::getSimilarTracks(Database::Session& dbSession, const std::unordered_set<Database::IdType>& trackIds, std::size_t maxCount) const ClusterClassifier::getSimilarTracks(Database::Session& dbSession, const std::vector<Database::TrackId>& trackIds, std::size_t maxCount) const
{ {
ResultContainer<Database::TrackId> res;
auto transaction {dbSession.createSharedTransaction()}; auto transaction {dbSession.createSharedTransaction()};
const auto tracks {Database::Track::getSimilarTracks(dbSession, trackIds, 0, maxCount)}; const auto tracks {Database::Track::getSimilarTracks(dbSession, trackIds, 0, maxCount)};
std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const auto& track) { return track->getId(); });
std::unordered_set<Database::IdType> res;
std::transform(std::cbegin(tracks), std::cend(tracks), std::inserter(res, std::end(res)),
[](const auto& track) { return track.id(); });
return res; return res;
} }
std::unordered_set<Database::IdType> IClassifier::ResultContainer<Database::TrackId>
ClusterClassifier::getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const ClusterClassifier::getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId tracklistId, std::size_t maxCount) const
{ {
std::unordered_set<Database::IdType> res; ResultContainer<Database::TrackId> res;
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
@@ -58,16 +58,15 @@ ClusterClassifier::getSimilarTracksFromTrackList(Database::Session& session, Dat
return res; return res;
const auto tracks {trackList->getSimilarTracks(0, maxCount)}; const auto tracks {trackList->getSimilarTracks(0, maxCount)};
std::transform(std::cbegin(tracks), std::cend(tracks), std::inserter(res, std::end(res)), std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const auto& track) { return track->getId(); });
[](const Database::Track::pointer& track) { return track.id(); });
return res; return res;
} }
std::unordered_set<Database::IdType> IClassifier::ResultContainer<Database::ReleaseId>
ClusterClassifier::getSimilarReleases(Database::Session& dbSession, Database::IdType releaseId, std::size_t maxCount) const ClusterClassifier::getSimilarReleases(Database::Session& dbSession, Database::ReleaseId releaseId, std::size_t maxCount) const
{ {
std::unordered_set<Database::IdType> res; ResultContainer<Database::ReleaseId> res;
auto transaction {dbSession.createSharedTransaction()}; auto transaction {dbSession.createSharedTransaction()};
@@ -76,19 +75,18 @@ ClusterClassifier::getSimilarReleases(Database::Session& dbSession, Database::Id
return res; return res;
const auto releases {release->getSimilarReleases(0, maxCount)}; const auto releases {release->getSimilarReleases(0, maxCount)};
std::transform(std::cbegin(releases), std::cend(releases), std::inserter(res, std::end(res)), std::transform(std::cbegin(releases), std::cend(releases), std::back_inserter(res), [](const auto& release) { return release->getId(); });
[](const auto& release) { return release.id(); });
return res; return res;
} }
std::unordered_set<Database::IdType> IClassifier::ResultContainer<Database::ArtistId>
ClusterClassifier::getSimilarArtists(Database::Session& dbSession, ClusterClassifier::getSimilarArtists(Database::Session& dbSession,
Database::IdType artistId, Database::ArtistId artistId,
EnumSet<Database::TrackArtistLinkType> artistLinkTypes, EnumSet<Database::TrackArtistLinkType> artistLinkTypes,
std::size_t maxCount) const std::size_t maxCount) const
{ {
std::unordered_set<Database::IdType> res; ResultContainer<Database::ArtistId> res;
auto transaction {dbSession.createSharedTransaction()}; auto transaction {dbSession.createSharedTransaction()};
@@ -97,8 +95,7 @@ ClusterClassifier::getSimilarArtists(Database::Session& dbSession,
return res; return res;
const auto artists {artist->getSimilarArtists(artistLinkTypes, Database::Range {0, maxCount})}; const auto artists {artist->getSimilarArtists(artistLinkTypes, Database::Range {0, maxCount})};
std::transform(std::cbegin(artists), std::cend(artists), std::inserter(res, std::end(res)), std::transform(std::cbegin(artists), std::cend(artists), std::back_inserter(res), [](const auto& artist) { return artist->getId(); });
[](const auto& artist) { return artist.id(); });
return res; return res;
} }
@@ -40,11 +40,11 @@ namespace Recommendation
bool load(Database::Session&, bool, const ProgressCallback&) override { return true; } bool load(Database::Session&, bool, const ProgressCallback&) override { return true; }
void requestCancelLoad() override {} void requestCancelLoad() override {}
ResultContainer getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const override; ResultContainer<Database::TrackId> getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId tracklistId, std::size_t maxCount) const override;
ResultContainer getSimilarTracks(Database::Session& session, const std::unordered_set<Database::IdType>& tracksId, std::size_t maxCount) const override; ResultContainer<Database::TrackId> getSimilarTracks(Database::Session& session, const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
ResultContainer getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) const override; ResultContainer<Database::ReleaseId> getSimilarReleases(Database::Session& session, Database::ReleaseId releaseId, std::size_t maxCount) const override;
ResultContainer getSimilarArtists(Database::Session& session, ResultContainer<Database::ArtistId> getSimilarArtists(Database::Session& session,
Database::IdType artistId, Database::ArtistId artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes, EnumSet<Database::TrackArtistLinkType> linkTypes,
std::size_t maxCount) const override; std::size_t maxCount) const override;
}; };
@@ -1,117 +0,0 @@
/*
* Copyright (C) 2018 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 <http://www.gnu.org/licenses/>.
*/
#pragma once
#include <functional>
#include <unordered_map>
#include <optional>
#include <string>
#include "som/DataNormalizer.hpp"
#include "som/Network.hpp"
#include "FeaturesClassifierCache.hpp"
#include "FeaturesDefs.hpp"
#include "IClassifier.hpp"
namespace Database
{
class Session;
}
namespace Recommendation {
using FeatureWeight = double;
class FeaturesClassifier : public IClassifier
{
public:
FeaturesClassifier() = default;
FeaturesClassifier(const FeaturesClassifier&) = delete;
FeaturesClassifier(FeaturesClassifier&&) = delete;
FeaturesClassifier& operator=(const FeaturesClassifier&) = delete;
FeaturesClassifier& operator=(FeaturesClassifier&&) = delete;
using FeaturesFetchFunc = std::function<std::optional<std::unordered_map<std::string, std::vector<double>>>(Database::IdType /*trackId*/, const std::unordered_set<std::string>& /*features*/)>;
// Default is to retrieve the features from the database (may be slow).
// Use this only if you want to train different searchers with some cached data
static void setFeaturesFetchFunc(FeaturesFetchFunc func) { _featuresFetchFunc = func; }
static const FeatureSettingsMap& getDefaultTrainFeatureSettings();
private:
std::string_view getName() const override { return "Features"; }
bool load(Database::Session& session, bool forceReload, const ProgressCallback& progressCallback) override;
void requestCancelLoad() override;
std::unordered_set<Database::IdType> getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const override;
std::unordered_set<Database::IdType> getSimilarTracks(Database::Session& session, const std::unordered_set<Database::IdType>& tracksId, std::size_t maxCount) const override;
std::unordered_set<Database::IdType> getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) const override;
std::unordered_set<Database::IdType> getSimilarArtists(Database::Session& session,
Database::IdType artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes,
std::size_t maxCount) const override;
bool loadFromCache(Database::Session& session, const FeaturesClassifierCache& cache);
// Use training (may be very slow)
struct TrainSettings
{
std::size_t iterationCount {10};
float sampleCountPerNeuron {4};
FeatureSettingsMap featureSettingsMap;
};
bool loadFromTraining(Database::Session& session, const TrainSettings& trainSettings, const ProgressCallback& progressCallback);
using ObjectPositions = std::unordered_map<Database::IdType, std::unordered_set<SOM::Position>>;
using MatrixOfObjects = SOM::Matrix<std::unordered_set<Database::IdType>>;
bool load(Database::Session& session,
SOM::Network network,
const ObjectPositions& tracksPosition);
FeaturesClassifierCache toCache() const;
static std::unordered_set<SOM::Position> getMatchingRefVectorsPosition(const std::unordered_set<Database::IdType>& ids, const ObjectPositions& objectPositions);
static std::unordered_set<Database::IdType> getObjectsIds(const std::unordered_set<SOM::Position>& positionSet, const MatrixOfObjects& objectsMap);
std::unordered_set<Database::IdType> getSimilarObjects(const std::unordered_set<Database::IdType>& ids,
const SOM::Matrix<std::unordered_set<Database::IdType>>& objectsMap,
const ObjectPositions& objectPosition,
std::size_t maxCount) const;
bool _loadCancelled {};
std::unique_ptr<SOM::Network> _network;
double _networkRefVectorsDistanceMedian {};
ObjectPositions _artistPositions;
std::unordered_map<Database::TrackArtistLinkType, MatrixOfObjects> _artistsMap;
MatrixOfObjects _releasesMap;
ObjectPositions _releasePositions;
MatrixOfObjects _tracksMap;
ObjectPositions _trackPositions;
static inline FeaturesFetchFunc _featuresFetchFunc;
};
} // ns Recommendation
@@ -17,7 +17,7 @@
* along with LMS. If not, see <http://www.gnu.org/licenses/>. * along with LMS. If not, see <http://www.gnu.org/licenses/>.
*/ */
#include "FeaturesClassifier.hpp" #include "FeaturesEngine.hpp"
#include <numeric> #include <numeric>
@@ -35,13 +35,13 @@
namespace Recommendation { namespace Recommendation {
std::unique_ptr<IClassifier> createFeaturesClassifier() std::unique_ptr<IClassifier> createFeaturesEngine()
{ {
return std::make_unique<FeaturesClassifier>(); return std::make_unique<FeaturesEngine>();
} }
const FeatureSettingsMap& const FeatureSettingsMap&
FeaturesClassifier::getDefaultTrainFeatureSettings() FeaturesEngine::getDefaultTrainFeatureSettings()
{ {
static const FeatureSettingsMap defaultTrainFeatureSettings static const FeatureSettingsMap defaultTrainFeatureSettings
{ {
@@ -57,16 +57,16 @@ FeaturesClassifier::getDefaultTrainFeatureSettings()
static static
std::optional<FeatureValuesMap> std::optional<FeatureValuesMap>
getTrackFeatureValues(FeaturesClassifier::FeaturesFetchFunc func, Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames) getTrackFeatureValues(FeaturesEngine::FeaturesFetchFunc func, Database::TrackId trackId, const std::unordered_set<FeatureName>& featureNames)
{ {
return func(trackId, featureNames); return func(trackId, featureNames);
} }
static static
std::optional<FeatureValuesMap> std::optional<FeatureValuesMap>
getTrackFeatureValuesFromDb(Database::Session& session, Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames) getTrackFeatureValuesFromDb(Database::Session& session, Database::TrackId trackId, const std::unordered_set<FeatureName>& featureNames)
{ {
auto func = [&](Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames) auto func = [&](Database::TrackId trackId, const std::unordered_set<FeatureName>& featureNames)
{ {
std::optional<FeatureValuesMap> res; std::optional<FeatureValuesMap> res;
@@ -128,7 +128,7 @@ getInputVectorWeights(const FeatureSettingsMap& featureSettingsMap, std::size_t
} }
bool bool
FeaturesClassifier::loadFromTraining(Database::Session& session, const TrainSettings& trainSettings, const ProgressCallback& progressCallback) FeaturesEngine::loadFromTraining(Database::Session& session, const TrainSettings& trainSettings, const ProgressCallback& progressCallback)
{ {
LMS_LOG(RECOMMENDATION, INFO) << "Constructing features classifier..."; LMS_LOG(RECOMMENDATION, INFO) << "Constructing features classifier...";
@@ -141,7 +141,7 @@ FeaturesClassifier::loadFromTraining(Database::Session& session, const TrainSett
LMS_LOG(RECOMMENDATION, DEBUG) << "Features dimension = " << nbDimensions; LMS_LOG(RECOMMENDATION, DEBUG) << "Features dimension = " << nbDimensions;
std::vector<Database::IdType> trackIds; std::vector<Database::TrackId> trackIds;
{ {
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
@@ -151,13 +151,13 @@ FeaturesClassifier::loadFromTraining(Database::Session& session, const TrainSett
} }
std::vector<SOM::InputVector> samples; std::vector<SOM::InputVector> samples;
std::vector<Database::IdType> samplesTrackIds; std::vector<Database::TrackId> samplesTrackIds;
samples.reserve(trackIds.size()); samples.reserve(trackIds.size());
samplesTrackIds.reserve(trackIds.size()); samplesTrackIds.reserve(trackIds.size());
LMS_LOG(RECOMMENDATION, DEBUG) << "Extracting features..."; LMS_LOG(RECOMMENDATION, DEBUG) << "Extracting features...";
for (Database::IdType trackId : trackIds) for (Database::TrackId trackId : trackIds)
{ {
if (_loadCancelled) if (_loadCancelled)
return false; return false;
@@ -223,7 +223,7 @@ FeaturesClassifier::loadFromTraining(Database::Session& session, const TrainSett
return false; return false;
LMS_LOG(RECOMMENDATION, DEBUG) << "Classifying tracks..."; LMS_LOG(RECOMMENDATION, DEBUG) << "Classifying tracks...";
ObjectPositions trackPositions; TrackPositions trackPositions;
for (std::size_t i {}; i < samples.size(); ++i) for (std::size_t i {}; i < samples.size(); ++i)
{ {
if (_loadCancelled) if (_loadCancelled)
@@ -231,7 +231,7 @@ FeaturesClassifier::loadFromTraining(Database::Session& session, const TrainSett
const SOM::Position position {network.getClosestRefVectorPosition(samples[i])}; const SOM::Position position {network.getClosestRefVectorPosition(samples[i])};
trackPositions[samplesTrackIds[i]].insert(position); trackPositions[samplesTrackIds[i]].push_back(position);
} }
LMS_LOG(RECOMMENDATION, DEBUG) << "Classifying tracks DONE"; LMS_LOG(RECOMMENDATION, DEBUG) << "Classifying tracks DONE";
@@ -240,28 +240,25 @@ FeaturesClassifier::loadFromTraining(Database::Session& session, const TrainSett
} }
bool bool
FeaturesClassifier::loadFromCache(Database::Session& session, const FeaturesClassifierCache& cache) FeaturesEngine::loadFromCache(Database::Session& session, const FeaturesEngineCache& cache)
{ {
LMS_LOG(RECOMMENDATION, INFO) << "Constructing features classifier from cache..."; LMS_LOG(RECOMMENDATION, INFO) << "Constructing features classifier from cache...";
return load(session, std::move(cache._network), cache._trackPositions); return load(session, std::move(cache._network), cache._trackPositions);
} }
std::unordered_set<Database::IdType> IClassifier::ResultContainer<Database::TrackId>
FeaturesClassifier::getSimilarTracksFromTrackList(Database::Session& session, Database::IdType trackListId, std::size_t maxCount) const FeaturesEngine::getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId trackListId, std::size_t maxCount) const
{ {
const std::unordered_set<Database::IdType> trackIds {[&] const std::vector<Database::TrackId> trackIds {[&]
{ {
std::unordered_set<Database::IdType> res; std::vector<Database::TrackId> res;
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
const Database::TrackList::pointer trackList {Database::TrackList::getById(session, trackListId)}; const Database::TrackList::pointer trackList {Database::TrackList::getById(session, trackListId)};
if (trackList) if (trackList)
{ res = trackList->getTrackIds();
const std::vector<Database::IdType> orderedTrackIds {trackList->getTrackIds()};
res = std::unordered_set<Database::IdType>(std::cbegin(orderedTrackIds), std::cend(orderedTrackIds));
}
return res; return res;
}()}; }()};
@@ -269,72 +266,64 @@ FeaturesClassifier::getSimilarTracksFromTrackList(Database::Session& session, Da
return getSimilarTracks(session, trackIds, maxCount); return getSimilarTracks(session, trackIds, maxCount);
} }
std::unordered_set<Database::IdType> std::vector<Database::TrackId>
FeaturesClassifier::getSimilarTracks(Database::Session& session, const std::unordered_set<Database::IdType>& tracksIds, std::size_t maxCount) const FeaturesEngine::getSimilarTracks(Database::Session& session, const std::vector<Database::TrackId>& tracksIds, std::size_t maxCount) const
{ {
auto similarTrackIds {getSimilarObjects(tracksIds, _tracksMap, _trackPositions, maxCount)}; auto similarTrackIds {getSimilarObjects(tracksIds, _trackMatrix, _trackPositions, maxCount)};
if (!similarTrackIds.empty())
{ {
// Report only existing ids // Report only existing ids, as tracks may have been removed a long time ago (refreshing the SOM takes some time)
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
for (auto it {std::begin(similarTrackIds)}; it != std::end(similarTrackIds);) similarTrackIds.erase(std::remove_if(std::begin(similarTrackIds), std::end(similarTrackIds),
{ [&](Database::TrackId trackId)
const Database::IdType trackId {*it}; {
if (!Database::Track::getById(session, trackId)) return Database::Track::getById(session, trackId); // TODO exists
it = similarTrackIds.erase(it); }), std::end(similarTrackIds));
else
it++;
}
} }
return similarTrackIds; return similarTrackIds;
} }
std::unordered_set<Database::IdType> std::vector<Database::ReleaseId>
FeaturesClassifier::getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) const FeaturesEngine::getSimilarReleases(Database::Session& session, Database::ReleaseId releaseId, std::size_t maxCount) const
{ {
auto similarReleaseIds {getSimilarObjects({releaseId}, _releasesMap, _releasePositions, maxCount)}; auto similarReleaseIds {getSimilarObjects<Database::ReleaseId>({releaseId}, _releaseMatrix, _releasePositions, maxCount)};
if (!similarReleaseIds.empty())
{ {
// Report only existing ids // Report only existing ids
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
for (auto it {std::begin(similarReleaseIds)}; it != std::end(similarReleaseIds);) similarReleaseIds.erase(std::remove_if(std::begin(similarReleaseIds), std::end(similarReleaseIds),
{ [&](Database::ReleaseId releaseId)
const Database::IdType similarReleaseId {*it}; {
if (!Database::Release::getById(session, similarReleaseId)) return Database::Release::getById(session, releaseId); // TODO exists
it = similarReleaseIds.erase(it); }), std::end(similarReleaseIds));
else
it++;
}
} }
return similarReleaseIds; return similarReleaseIds;
} }
std::unordered_set<Database::IdType> std::vector<Database::ArtistId>
FeaturesClassifier::getSimilarArtists(Database::Session& session, FeaturesEngine::getSimilarArtists(Database::Session& session,
Database::IdType artistId, Database::ArtistId artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes, EnumSet<Database::TrackArtistLinkType> linkTypes,
std::size_t maxCount) const std::size_t maxCount) const
{ {
auto getSimilarArtistIdsForLinkType {[&] (Database::TrackArtistLinkType linkType) auto getSimilarArtistIdsForLinkType {[&] (Database::TrackArtistLinkType linkType)
{ {
std::unordered_set<Database::IdType> similarArtistIds; std::vector<Database::ArtistId> similarArtistIds;
const auto itArtists {_artistsMap.find(linkType)}; const auto itArtists {_artistMatrix.find(linkType)};
if (itArtists == std::cend(_artistsMap)) if (itArtists == std::cend(_artistMatrix))
{ {
return similarArtistIds; return similarArtistIds;
} }
similarArtistIds = getSimilarObjects({artistId}, itArtists->second, _artistPositions, maxCount); return getSimilarObjects({artistId}, itArtists->second, _artistPositions, maxCount);
return similarArtistIds;
}}; }};
std::unordered_set<Database::IdType> similarArtistIds; std::unordered_set<Database::ArtistId> similarArtistIds;
for (Database::TrackArtistLinkType linkType : linkTypes) for (Database::TrackArtistLinkType linkType : linkTypes)
{ {
@@ -342,44 +331,42 @@ FeaturesClassifier::getSimilarArtists(Database::Session& session,
similarArtistIds.insert(std::begin(similarArtistIdsForLinkType), std::end(similarArtistIdsForLinkType)); similarArtistIds.insert(std::begin(similarArtistIdsForLinkType), std::end(similarArtistIdsForLinkType));
} }
if (!similarArtistIds.empty()) std::vector<Database::ArtistId> res(std::cbegin(similarArtistIds), std::cend(similarArtistIds));
{ {
// Report only existing ids // Report only existing ids
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
for (auto it {std::begin(similarArtistIds)}; it != std::end(similarArtistIds);) res.erase(std::remove_if(std::begin(res), std::end(res),
{ [&](Database::ArtistId artistId)
const Database::IdType similarArtistId {*it}; {
if (!Database::Artist::getById(session, similarArtistId)) return Database::Artist::getById(session, artistId); // TODO exists
it = similarArtistIds.erase(it); }), std::end(res));
else
it++;
}
} }
while (similarArtistIds.size() > maxCount) while (res.size() > maxCount)
similarArtistIds.erase(Random::pickRandom(similarArtistIds)); res.erase(Random::pickRandom(res));
return similarArtistIds; return res;
} }
FeaturesClassifierCache FeaturesEngineCache
FeaturesClassifier::toCache() const FeaturesEngine::toCache() const
{ {
return FeaturesClassifierCache {*_network, _trackPositions}; return FeaturesEngineCache {*_network, _trackPositions};
} }
bool bool
FeaturesClassifier::load(Database::Session& session, bool forceReload, const ProgressCallback& progressCallback) FeaturesEngine::load(Database::Session& session, bool forceReload, const ProgressCallback& progressCallback)
{ {
if (forceReload) if (forceReload)
{ {
FeaturesClassifierCache::invalidate(); FeaturesEngineCache::invalidate();
} }
else else
{ {
const std::optional<FeaturesClassifierCache> cache {FeaturesClassifierCache::read()}; const std::optional<FeaturesEngineCache> cache {FeaturesEngineCache::read()};
if (cache) if (cache)
return loadFromCache(session, *cache); return loadFromCache(session, *cache);
} }
@@ -395,64 +382,65 @@ FeaturesClassifier::load(Database::Session& session, bool forceReload, const Pro
} }
void void
FeaturesClassifier::requestCancelLoad() FeaturesEngine::requestCancelLoad()
{ {
LMS_LOG(RECOMMENDATION, DEBUG) << "Requesting init cancellation"; LMS_LOG(RECOMMENDATION, DEBUG) << "Requesting init cancellation";
_loadCancelled = true; _loadCancelled = true;
} }
bool bool
FeaturesClassifier::load(Database::Session& session, FeaturesEngine::load(Database::Session& session,
SOM::Network network, SOM::Network network,
const ObjectPositions& tracksPosition) const TrackPositions& trackPositions)
{ {
using namespace Database;
_networkRefVectorsDistanceMedian = network.computeRefVectorsDistanceMedian(); _networkRefVectorsDistanceMedian = network.computeRefVectorsDistanceMedian();
LMS_LOG(RECOMMENDATION, DEBUG) << "Median distance betweend ref vectors = " << _networkRefVectorsDistanceMedian; LMS_LOG(RECOMMENDATION, DEBUG) << "Median distance betweend ref vectors = " << _networkRefVectorsDistanceMedian;
const SOM::Coordinate width {network.getWidth()}; const SOM::Coordinate width {network.getWidth()};
const SOM::Coordinate height {network.getHeight()}; const SOM::Coordinate height {network.getHeight()};
_releasesMap = MatrixOfObjects {width, height}; _releaseMatrix = ReleaseMatrix {width, height};
_tracksMap = MatrixOfObjects {width, height}; _trackMatrix = TrackMatrix {width, height};
LMS_LOG(RECOMMENDATION, DEBUG) << "Constructing maps..."; LMS_LOG(RECOMMENDATION, DEBUG) << "Constructing maps...";
for (auto itTrackCoord : tracksPosition) for (const auto& [trackId, positions] : trackPositions)
{ {
if (_loadCancelled) if (_loadCancelled)
return false; return false;
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
Database::IdType trackId {itTrackCoord.first}; const Track::pointer track {Database::Track::getById(session, trackId)};
const std::unordered_set<SOM::Position>& positionSet {itTrackCoord.second};
const Database::Track::pointer track {Database::Track::getById(session, trackId)};
if (!track) if (!track)
continue; continue;
for (const SOM::Position& position : positionSet) for (const SOM::Position& position : positions)
{ {
_tracksMap[position].insert(trackId); Utils::push_back_if_not_present(_trackPositions[trackId], position);
_trackPositions[trackId].insert(position); Utils::push_back_if_not_present(_trackMatrix[position], trackId);
if (track->getRelease()) if (Release::pointer release {track->getRelease()})
{ {
_releasePositions[track->getRelease().id()].insert(position); const ReleaseId releaseId {release->getId()};
_releasesMap[position].insert(track->getRelease().id()); Utils::push_back_if_not_present(_releasePositions[releaseId], position);
Utils::push_back_if_not_present(_releaseMatrix[position], releaseId);
} }
for (const auto& artistLink : track->getArtistLinks()) for (const TrackArtistLink::pointer& artistLink : track->getArtistLinks())
{ {
_artistPositions[artistLink->getArtist().id()].insert(position); const ArtistId artistId {artistLink->getArtist()->getId()};
auto itArtists {_artistsMap.find(artistLink->getType())};
if (itArtists == std::cend(_artistsMap)) Utils::push_back_if_not_present(_artistPositions[artistId], position);
auto itArtists {_artistMatrix.find(artistLink->getType())};
if (itArtists == std::cend(_artistMatrix))
{ {
auto [it, inserted] = _artistsMap.try_emplace(artistLink->getType(), MatrixOfObjects {}); auto [it, inserted] = _artistMatrix.try_emplace(artistLink->getType(), ArtistMatrix {width, height});
assert(inserted); assert(inserted);
itArtists = it; itArtists = it;
itArtists->second = MatrixOfObjects {width, height};
} }
itArtists->second[position].insert(artistLink->getArtist().id()); Utils::push_back_if_not_present(itArtists->second[position], artistId);
} }
} }
} }
@@ -464,83 +452,4 @@ FeaturesClassifier::load(Database::Session& session,
return true; return true;
} }
std::unordered_set<SOM::Position>
FeaturesClassifier::getMatchingRefVectorsPosition(const std::unordered_set<Database::IdType>& ids, const ObjectPositions& objectPositions)
{
std::unordered_set<SOM::Position> res;
if (ids.empty())
return res;
for (auto id : ids)
{
auto it = objectPositions.find(id);
if (it == objectPositions.end())
continue;
for (const auto& position : it->second)
res.insert(position);
}
return res;
}
std::unordered_set<Database::IdType>
FeaturesClassifier::getObjectsIds(const std::unordered_set<SOM::Position>& positionSet, const MatrixOfObjects& objectsMap)
{
std::unordered_set<Database::IdType> res;
for (const auto& position : positionSet)
{
for (auto id : objectsMap.get(position))
res.insert(id);
}
return res;
}
std::unordered_set<Database::IdType>
FeaturesClassifier::getSimilarObjects(const std::unordered_set<Database::IdType>& ids,
const MatrixOfObjects& objectsMap,
const ObjectPositions& objectPosition,
std::size_t maxCount) const
{
std::unordered_set<Database::IdType> res;
std::unordered_set<SOM::Position> searchedRefVectorsPosition {getMatchingRefVectorsPosition(ids, objectPosition)};
if (searchedRefVectorsPosition.empty())
return res;
while (1)
{
std::unordered_set<Database::IdType> closestObjectIds {getObjectsIds(searchedRefVectorsPosition, objectsMap)};
// Remove objects that are already in input or already reported
for (auto id : ids)
closestObjectIds.erase(id);
for (auto it {std::cbegin(closestObjectIds)}; it != std::cend(closestObjectIds); ++it)
{
if (res.size() == maxCount)
break;
res.insert(*it);
}
if (res.size() == maxCount)
break;
// If there is not enough objects, try again with closest neighbour until there is too much distance
const std::optional<SOM::Position> closestRefVectorPosition {_network->getClosestRefVectorPosition(searchedRefVectorsPosition, _networkRefVectorsDistanceMedian * 0.75)};
if (!closestRefVectorPosition)
break;
searchedRefVectorsPosition.insert(closestRefVectorPosition.value());
}
return res;
}
} // ns Recommendation } // ns Recommendation
@@ -0,0 +1,216 @@
/*
* Copyright (C) 2018 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 <http://www.gnu.org/licenses/>.
*/
#pragma once
#include <algorithm>
#include <functional>
#include <unordered_map>
#include <optional>
#include <string>
#include <vector>
#include "som/DataNormalizer.hpp"
#include "som/Network.hpp"
#include "utils/Utils.hpp"
#include "FeaturesEngineCache.hpp"
#include "FeaturesDefs.hpp"
#include "IClassifier.hpp"
namespace Database
{
class Session;
}
namespace Recommendation {
using FeatureWeight = double;
class FeaturesEngine : public IClassifier
{
public:
FeaturesEngine() = default;
FeaturesEngine(const FeaturesEngine&) = delete;
FeaturesEngine(FeaturesEngine&&) = delete;
FeaturesEngine& operator=(const FeaturesEngine&) = delete;
FeaturesEngine& operator=(FeaturesEngine&&) = delete;
using FeaturesFetchFunc = std::function<std::optional<std::unordered_map<std::string, std::vector<double>>>(Database::TrackId, const std::unordered_set<std::string>& /*features*/)>;
// Default is to retrieve the features from the database (may be slow).
// Use this only if you want to train different searchers with some cached data
static void setFeaturesFetchFunc(FeaturesFetchFunc func) { _featuresFetchFunc = func; }
static const FeatureSettingsMap& getDefaultTrainFeatureSettings();
private:
std::string_view getName() const override { return "Features"; }
bool load(Database::Session& session, bool forceReload, const ProgressCallback& progressCallback) override;
void requestCancelLoad() override;
ResultContainer<Database::TrackId> getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId tracklistId, std::size_t maxCount) const override;
ResultContainer<Database::TrackId> getSimilarTracks(Database::Session& session, const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
ResultContainer<Database::ReleaseId> getSimilarReleases(Database::Session& session, Database::ReleaseId releaseId, std::size_t maxCount) const override;
ResultContainer<Database::ArtistId> getSimilarArtists(Database::Session& session,
Database::ArtistId artistId,
EnumSet<Database::TrackArtistLinkType> linkTypes,
std::size_t maxCount) const override;
bool loadFromCache(Database::Session& session, const FeaturesEngineCache& cache);
// Use training (may be very slow)
struct TrainSettings
{
std::size_t iterationCount {10};
float sampleCountPerNeuron {4};
FeatureSettingsMap featureSettingsMap;
};
bool loadFromTraining(Database::Session& session, const TrainSettings& trainSettings, const ProgressCallback& progressCallback);
template <typename IdType>
using ObjectPositions = std::unordered_map<IdType, std::vector<SOM::Position>>;
using ArtistPositions = ObjectPositions<Database::ArtistId>;
using ReleasePositions = ObjectPositions<Database::ReleaseId>;
using TrackPositions = ObjectPositions<Database::TrackId>;
template <typename IdType>
using ObjectMatrix = SOM::Matrix<std::vector<IdType>>;
using ArtistMatrix = ObjectMatrix<Database::ArtistId>;
using ReleaseMatrix = ObjectMatrix<Database::ReleaseId>;
using TrackMatrix = ObjectMatrix<Database::TrackId>;
bool load(Database::Session& session, SOM::Network network, const TrackPositions& tracksPosition);
FeaturesEngineCache toCache() const;
template <typename IdType>
static std::vector<SOM::Position> getMatchingRefVectorsPosition(const std::vector<IdType>& ids, const ObjectPositions<IdType>& objectPositions);
template <typename IdType>
static std::vector<IdType> getObjectsIds(const std::vector<SOM::Position>& positions, const ObjectMatrix<IdType>& objectsMatrix);
template <typename IdType>
std::vector<IdType> getSimilarObjects(const std::vector<IdType>& ids,
const ObjectMatrix<IdType>& objectMatrix,
const ObjectPositions<IdType>& objectPositions,
std::size_t maxCount) const;
bool _loadCancelled {};
std::unique_ptr<SOM::Network> _network;
double _networkRefVectorsDistanceMedian {};
ArtistPositions _artistPositions;
std::unordered_map<Database::TrackArtistLinkType, ArtistMatrix> _artistMatrix;
ReleasePositions _releasePositions;
ReleaseMatrix _releaseMatrix;
TrackPositions _trackPositions;
TrackMatrix _trackMatrix;
static inline FeaturesFetchFunc _featuresFetchFunc;
};
template <typename IdType>
std::vector<SOM::Position>
FeaturesEngine::getMatchingRefVectorsPosition(const std::vector<IdType>& ids, const ObjectPositions<IdType>& objectPositions)
{
std::vector<SOM::Position> res;
if (ids.empty())
return res;
for (const IdType id : ids)
{
auto it = objectPositions.find(id);
if (it == objectPositions.end())
continue;
for (const SOM::Position& position : it->second)
Utils::push_back_if_not_present(res, position);
}
return res;
}
template <typename IdType>
std::vector<IdType>
FeaturesEngine::getObjectsIds(const std::vector<SOM::Position>& positions, const ObjectMatrix<IdType>& objectMatrix)
{
std::vector<IdType> res;
for (const SOM::Position& position : positions)
{
for (const IdType id : objectMatrix.get(position))
Utils::push_back_if_not_present(res, id);
}
return res;
}
template <typename IdType>
std::vector<IdType>
FeaturesEngine::getSimilarObjects(const std::vector<IdType>& ids,
const ObjectMatrix<IdType>& objectMatrix,
const ObjectPositions<IdType>& objectPositions,
std::size_t maxCount) const
{
std::vector<IdType> res;
std::vector<SOM::Position> searchedRefVectorsPosition {getMatchingRefVectorsPosition(ids, objectPositions)};
if (searchedRefVectorsPosition.empty())
return res;
while (1)
{
std::vector<IdType> closestObjectIds {getObjectsIds(searchedRefVectorsPosition, objectMatrix)};
// Remove objects that are already in input or already reported
closestObjectIds.erase(std::remove_if(std::begin(closestObjectIds), std::end(closestObjectIds),
[&](IdType id)
{
return std::find(std::cbegin(ids), std::cend(ids), id) != std::cend(ids);
})
, std::end(closestObjectIds));
for (IdType id : closestObjectIds)
{
if (res.size() == maxCount)
break;
Utils::push_back_if_not_present(res, id);
}
if (res.size() == maxCount)
break;
// If there is not enough objects, try again with closest neighbour until there is too much distance
const std::optional<SOM::Position> closestRefVectorPosition {_network->getClosestRefVectorPosition(searchedRefVectorsPosition, _networkRefVectorsDistanceMedian * 0.75)};
if (!closestRefVectorPosition)
break;
Utils::push_back_if_not_present(searchedRefVectorsPosition, closestRefVectorPosition.value());
}
return res;
}
} // ns Recommendation
@@ -17,7 +17,7 @@
* along with LMS. If not, see <http://www.gnu.org/licenses/>. * along with LMS. If not, see <http://www.gnu.org/licenses/>.
*/ */
#include "FeaturesClassifierCache.hpp" #include "FeaturesEngineCache.hpp"
#include <boost/property_tree/ptree.hpp> #include <boost/property_tree/ptree.hpp>
#include <boost/property_tree/xml_parser.hpp> #include <boost/property_tree/xml_parser.hpp>
@@ -90,7 +90,7 @@ networkToCacheFile(const SOM::Network& network, std::filesystem::path path)
} }
std::optional<SOM::Network> std::optional<SOM::Network>
FeaturesClassifierCache::createNetworkFromCacheFile(const std::filesystem::path& path) FeaturesEngineCache::createNetworkFromCacheFile(const std::filesystem::path& path)
{ {
if (!std::filesystem::exists(path)) if (!std::filesystem::exists(path))
return std::nullopt; return std::nullopt;
@@ -143,19 +143,19 @@ FeaturesClassifierCache::createNetworkFromCacheFile(const std::filesystem::path&
} }
bool bool
FeaturesClassifierCache::objectPositionToCacheFile(const ObjectPositions& objectsPosition, const std::filesystem::path& path) FeaturesEngineCache::objectPositionToCacheFile(const TrackPositions& trackPositions, const std::filesystem::path& path)
{ {
try try
{ {
boost::property_tree::ptree root; boost::property_tree::ptree root;
for (const auto& objectPosition : objectsPosition) for (const auto& [id, positions] : trackPositions)
{ {
boost::property_tree::ptree node; boost::property_tree::ptree node;
node.put("id", objectPosition.first); node.put("id", id.getValue());
for (const auto& position : objectPosition.second) for (const SOM::Position& position : positions)
{ {
boost::property_tree::ptree positionNode; boost::property_tree::ptree positionNode;
positionNode.put("x", position.x); positionNode.put("x", position.x);
@@ -177,8 +177,8 @@ FeaturesClassifierCache::objectPositionToCacheFile(const ObjectPositions& object
} }
} }
std::optional<FeaturesClassifierCache::ObjectPositions> std::optional<FeaturesEngineCache::TrackPositions>
FeaturesClassifierCache::createObjectPositionsFromCacheFile(const std::filesystem::path& path) FeaturesEngineCache::createObjectPositionsFromCacheFile(const std::filesystem::path& path)
{ {
try try
{ {
@@ -188,17 +188,17 @@ FeaturesClassifierCache::createObjectPositionsFromCacheFile(const std::filesyste
boost::property_tree::read_xml(path.string(), root); boost::property_tree::read_xml(path.string(), root);
ObjectPositions res; TrackPositions res;
for (const auto& object : root.get_child("objects")) for (const auto& object : root.get_child("objects"))
{ {
auto id = object.second.get<Database::IdType>("id"); const Database::TrackId id {object.second.get<Database::IdType::ValueType>("id")};
for (const auto& position : object.second.get_child("position")) for (const auto& position : object.second.get_child("position"))
{ {
auto x = position.second.get<SOM::Coordinate>("x"); auto x = position.second.get<SOM::Coordinate>("x");
auto y = position.second.get<SOM::Coordinate>("y"); auto y = position.second.get<SOM::Coordinate>("y");
res[id].insert({x, y}); res[id].push_back({x, y});
} }
} }
@@ -214,14 +214,14 @@ FeaturesClassifierCache::createObjectPositionsFromCacheFile(const std::filesyste
} }
void void
FeaturesClassifierCache::invalidate() FeaturesEngineCache::invalidate()
{ {
std::filesystem::remove(getCacheNetworkFilePath()); std::filesystem::remove(getCacheNetworkFilePath());
std::filesystem::remove(getCacheTrackPositionsFilePath()); std::filesystem::remove(getCacheTrackPositionsFilePath());
} }
std::optional<FeaturesClassifierCache> std::optional<FeaturesEngineCache>
FeaturesClassifierCache::read() FeaturesEngineCache::read()
{ {
auto network{createNetworkFromCacheFile(getCacheNetworkFilePath())}; auto network{createNetworkFromCacheFile(getCacheNetworkFilePath())};
if (!network) if (!network)
@@ -231,11 +231,11 @@ FeaturesClassifierCache::read()
if (!trackPositions) if (!trackPositions)
return std::nullopt; return std::nullopt;
return FeaturesClassifierCache {std::move(*network), std::move(*trackPositions)}; return FeaturesEngineCache {std::move(*network), std::move(*trackPositions)};
} }
void void
FeaturesClassifierCache::write() const FeaturesEngineCache::write() const
{ {
std::filesystem::create_directories(Service<IConfig>::get()->getPath("working-dir") / "cache" / "features"); std::filesystem::create_directories(Service<IConfig>::get()->getPath("working-dir") / "cache" / "features");
@@ -246,7 +246,7 @@ FeaturesClassifierCache::write() const
} }
} }
FeaturesClassifierCache::FeaturesClassifierCache(SOM::Network network, ObjectPositions trackPositions) FeaturesEngineCache::FeaturesEngineCache(SOM::Network network, TrackPositions trackPositions)
: _network {std::move(network)}, : _network {std::move(network)},
_trackPositions {std::move(trackPositions)} _trackPositions {std::move(trackPositions)}
{ {
@@ -28,27 +28,27 @@
namespace Recommendation { namespace Recommendation {
class FeaturesClassifierCache class FeaturesEngineCache
{ {
public: public:
static void invalidate(); static void invalidate();
static std::optional<FeaturesClassifierCache> read(); static std::optional<FeaturesEngineCache> read();
void write() const; void write() const;
private: private:
using ObjectPositions = std::unordered_map<Database::IdType, std::unordered_set<SOM::Position>>; using TrackPositions = std::unordered_map<Database::TrackId, std::vector<SOM::Position>>;
FeaturesClassifierCache(SOM::Network network, ObjectPositions trackPositions); FeaturesEngineCache(SOM::Network network, TrackPositions trackPositions);
static std::optional<SOM::Network> createNetworkFromCacheFile(const std::filesystem::path& path); static std::optional<SOM::Network> createNetworkFromCacheFile(const std::filesystem::path& path);
static std::optional<ObjectPositions> createObjectPositionsFromCacheFile(const std::filesystem::path& path); static std::optional<TrackPositions> createObjectPositionsFromCacheFile(const std::filesystem::path& path);
static bool objectPositionToCacheFile(const ObjectPositions& objectsPosition, const std::filesystem::path& path); static bool objectPositionToCacheFile(const TrackPositions& trackPositions, const std::filesystem::path& path);
friend class FeaturesClassifier; friend class FeaturesEngine;
SOM::Network _network; SOM::Network _network;
ObjectPositions _trackPositions; TrackPositions _trackPositions;
}; };
} // namespace Recommendation } // namespace Recommendation
@@ -20,9 +20,7 @@
#pragma once #pragma once
#include <functional> #include <functional>
#include <optional> #include <memory>
#include <unordered_set>
#include "database/Types.hpp" #include "database/Types.hpp"
#include "utils/EnumSet.hpp" #include "utils/EnumSet.hpp"
@@ -48,15 +46,20 @@ namespace Recommendation
virtual void load(bool forceReload, const ProgressCallback& progressCallback = {}) = 0; virtual void load(bool forceReload, const ProgressCallback& progressCallback = {}) = 0;
virtual void cancelLoad() = 0; virtual void cancelLoad() = 0;
using ResultContainer = std::unordered_set<Database::IdType>; template <typename IdType>
using ResultContainer = std::vector<IdType>;
virtual ResultContainer getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) = 0; using ArtistContainer = ResultContainer<Database::ArtistId>;
virtual ResultContainer getSimilarTracks(Database::Session& session, const std::unordered_set<Database::IdType>& tracksId, std::size_t maxCount) = 0; using ReleaseContainer = ResultContainer<Database::ReleaseId>;
virtual ResultContainer getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) = 0; using TrackContainer = ResultContainer<Database::TrackId>;
virtual ResultContainer getSimilarArtists(Database::Session& session,
Database::IdType artistId, virtual TrackContainer getSimilarTracksFromTrackList(Database::Session& session, Database::TrackListId tracklistId, std::size_t maxCount) = 0;
EnumSet<Database::TrackArtistLinkType> linkTypes, virtual TrackContainer getSimilarTracks(Database::Session& session, const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) = 0;
std::size_t maxCount) = 0; virtual ReleaseContainer getSimilarReleases(Database::Session& session, Database::ReleaseId releaseId, std::size_t maxCount) = 0;
virtual ArtistContainer getSimilarArtists(Database::Session& session, Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) = 0;
protected:
virtual void requestCancelLoad() = 0;
}; };
std::unique_ptr<IEngine> createEngine(Database::Db& db); std::unique_ptr<IEngine> createEngine(Database::Db& db);
@@ -1,5 +1,5 @@
/* /*
* Copyright (C) 2020 Emeric Poupon * Copyright (C) 2021 Emeric Poupon
* *
* This file is part of LMS. * This file is part of LMS.
* *
@@ -17,27 +17,22 @@
* along with LMS. If not, see <http://www.gnu.org/licenses/>. * along with LMS. If not, see <http://www.gnu.org/licenses/>.
*/ */
#include "ParameterParsing.hpp" #pragma once
namespace StringUtils #include "database/Types.hpp"
#include "utils/EnumSet.hpp"
namespace Database
{ {
template<> class Db;
std::optional<API::Subsonic::Id> class Session;
readAs(std::string_view str)
{
return API::Subsonic::IdFromString(str);
}
template<>
std::optional<bool>
readAs(std::string_view str)
{
if (str == "true")
return true;
else if (str == "false")
return false;
return {};
}
} }
namespace Recommendation
{
class IRecommendation
{
};
} // ns Recommendation
+10 -10
View File
@@ -111,7 +111,7 @@ createArtist(Session& session, const MetaData::Artist& artistInfo)
static static
void void
updateArtistIfNeeded(const Artist::pointer& artist, const MetaData::Artist& artistInfo) updateArtistIfNeeded(Artist::pointer artist, const MetaData::Artist& artistInfo)
{ {
// Name may have been updated // Name may have been updated
if (artist->getName() != artistInfo.name) if (artist->getName() != artistInfo.name)
@@ -543,7 +543,7 @@ Scanner::scan(bool forceScan)
} }
bool bool
Scanner::fetchTrackFeatures(Database::IdType trackId, const UUID& recordingMBID) Scanner::fetchTrackFeatures(Database::TrackId trackId, const UUID& recordingMBID)
{ {
std::map<std::string, double> features; std::map<std::string, double> features;
@@ -551,14 +551,14 @@ Scanner::fetchTrackFeatures(Database::IdType trackId, const UUID& recordingMBID)
const std::string data {AcousticBrainz::extractLowLevelFeatures(recordingMBID)}; const std::string data {AcousticBrainz::extractLowLevelFeatures(recordingMBID)};
if (data.empty()) if (data.empty())
{ {
LMS_LOG(DBUPDATER, ERROR) << "Track " << trackId << ", recording MBID = '" << recordingMBID.getAsString() << "': cannot extract features using AcousticBrainz"; LMS_LOG(DBUPDATER, ERROR) << "Track " << trackId.getValue() << ", recording MBID = '" << recordingMBID.getAsString() << "': cannot extract features using AcousticBrainz";
return false; return false;
} }
{ {
auto uniqueTransaction {_dbSession.createUniqueTransaction()}; auto uniqueTransaction {_dbSession.createUniqueTransaction()};
Wt::Dbo::ptr<Database::Track> track {Database::Track::getById(_dbSession, trackId)}; Database::Track::pointer track {Database::Track::getById(_dbSession, trackId)};
if (!track) if (!track)
return false; return false;
@@ -580,7 +580,7 @@ Scanner::fetchTrackFeatures(ScanStats& stats)
struct TrackInfo struct TrackInfo
{ {
Database::IdType id; Database::TrackId id;
UUID recordingMBID; UUID recordingMBID;
}; };
@@ -592,7 +592,7 @@ Scanner::fetchTrackFeatures(ScanStats& stats)
auto tracks {Database::Track::getAllWithRecordingMBIDAndMissingFeatures(_dbSession)}; auto tracks {Database::Track::getAllWithRecordingMBIDAndMissingFeatures(_dbSession)};
for (const auto& track : tracks) for (const auto& track : tracks)
res.emplace_back(TrackInfo {track.id(), *track->getRecordingMBID()}); res.emplace_back(TrackInfo {track->getId(), *track->getRecordingMBID()});
return res; return res;
}()}; }()};
@@ -922,8 +922,8 @@ Scanner::removeMissingTracks(ScanStats& stats)
stepStats.totalElems = trackCount; stepStats.totalElems = trackCount;
notifyInProgress(stepStats); notifyInProgress(stepStats);
std::vector<std::pair<Database::IdType, std::filesystem::path>> trackPaths; std::vector<std::pair<Database::TrackId, std::filesystem::path>> trackPaths;
std::vector<IdType> tracksToRemove; std::vector<TrackId> tracksToRemove;
for (std::size_t i {trackCount < batchSize ? 0 : trackCount - batchSize}; ; i -= (i > batchSize ? batchSize : i)) for (std::size_t i {trackCount < batchSize ? 0 : trackCount - batchSize}; ; i -= (i > batchSize ? batchSize : i))
{ {
@@ -950,7 +950,7 @@ Scanner::removeMissingTracks(ScanStats& stats)
{ {
auto transaction {_dbSession.createUniqueTransaction()}; auto transaction {_dbSession.createUniqueTransaction()};
for (const IdType trackId : tracksToRemove) for (const TrackId trackId : tracksToRemove)
{ {
Track::pointer track {Track::getById(_dbSession, trackId)}; Track::pointer track {Track::getById(_dbSession, trackId)};
if (track) if (track)
@@ -1026,7 +1026,7 @@ Scanner::checkDuplicatedAudioFiles(ScanStats& stats)
if (auto trackMBID {track->getTrackMBID()}) if (auto trackMBID {track->getTrackMBID()})
{ {
LMS_LOG(DBUPDATER, INFO) << "Found duplicated Track MBID [" << trackMBID->getAsString() << "], file: " << track->getPath().string() << " - " << track->getName(); LMS_LOG(DBUPDATER, INFO) << "Found duplicated Track MBID [" << trackMBID->getAsString() << "], file: " << track->getPath().string() << " - " << track->getName();
stats.duplicates.emplace_back(ScanDuplicate {track.id(), DuplicateReason::SameMBID}); stats.duplicates.emplace_back(ScanDuplicate {track->getId(), DuplicateReason::SameMBID});
} }
} }
+5 -4
View File
@@ -22,6 +22,7 @@
#include <chrono> #include <chrono>
#include <shared_mutex> #include <shared_mutex>
#include <optional> #include <optional>
#include <unordered_set>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include <Wt/WIOService.h> #include <Wt/WIOService.h>
@@ -34,6 +35,7 @@
#include "database/Session.hpp" #include "database/Session.hpp"
#include "metadata/IParser.hpp" #include "metadata/IParser.hpp"
#include "scanner/IScanner.hpp" #include "scanner/IScanner.hpp"
#include "utils/Path.hpp"
class UUID; class UUID;
@@ -75,7 +77,7 @@ class Scanner : public IScanner
void scan(bool force); void scan(bool force);
void scanMediaDirectory( const std::filesystem::path& mediaDirectory, bool forceScan, ScanStats& stats); void scanMediaDirectory( const std::filesystem::path& mediaDirectory, bool forceScan, ScanStats& stats);
bool fetchTrackFeatures(Database::IdType trackId, const UUID& MBID); bool fetchTrackFeatures(Database::TrackId trackId, const UUID& MBID);
void fetchTrackFeatures(ScanStats& stats); void fetchTrackFeatures(ScanStats& stats);
// Helpers // Helpers
@@ -86,7 +88,6 @@ class Scanner : public IScanner
void removeOrphanEntries(); void removeOrphanEntries();
void checkDuplicatedAudioFiles(ScanStats& stats); void checkDuplicatedAudioFiles(ScanStats& stats);
void scanAudioFile(const std::filesystem::path& file, bool forceScan, ScanStats& stats); void scanAudioFile(const std::filesystem::path& file, bool forceScan, ScanStats& stats);
Database::IdType doScanAudioFile(const std::filesystem::path& file, ScanStats& stats);
void notifyInProgressIfNeeded(const ScanStepStats& stats); void notifyInProgressIfNeeded(const ScanStepStats& stats);
void notifyInProgress(const ScanStepStats& stats); void notifyInProgress(const ScanStepStats& stats);
void reloadSimilarityEngine(ScanStats& stats); void reloadSimilarityEngine(ScanStats& stats);
@@ -112,8 +113,8 @@ class Scanner : public IScanner
std::size_t _scanVersion {}; std::size_t _scanVersion {};
Wt::WTime _startTime; Wt::WTime _startTime;
Database::ScanSettings::UpdatePeriod _updatePeriod {Database::ScanSettings::UpdatePeriod::Never}; Database::ScanSettings::UpdatePeriod _updatePeriod {Database::ScanSettings::UpdatePeriod::Never};
std::unordered_set<std::filesystem::path> _fileExtensions; std::unordered_set<std::filesystem::path> _fileExtensions;
std::filesystem::path _mediaDirectory; std::filesystem::path _mediaDirectory;
Database::ScanSettings::RecommendationEngineType _recommendationEngineType; Database::ScanSettings::RecommendationEngineType _recommendationEngineType;
}; };
@@ -53,7 +53,7 @@ namespace Scanner {
struct ScanDuplicate struct ScanDuplicate
{ {
Database::IdType trackId; Database::TrackId trackId;
DuplicateReason reason; DuplicateReason reason;
}; };
+1 -2
View File
@@ -29,7 +29,6 @@
namespace Database namespace Database
{ {
class Db;
class Session; class Session;
class TrackList; class TrackList;
class User; class User;
@@ -48,7 +47,7 @@ namespace Scrobbling
virtual void addTimedListen(const TimedListen& listen) = 0; virtual void addTimedListen(const TimedListen& listen) = 0;
virtual Wt::Dbo::ptr<Database::TrackList> getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user) = 0; virtual Database::ObjectPtr<Database::TrackList> getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user) = 0;
}; };
std::unique_ptr<IScrobbler> createScrobbler(std::string_view backendName); std::unique_ptr<IScrobbler> createScrobbler(std::string_view backendName);
+33 -33
View File
@@ -64,7 +64,7 @@ namespace Scrobbling
} }
std::optional<Database::Scrobbler> std::optional<Database::Scrobbler>
Scrobbling::getUserScrobbler(Database::IdType userId) Scrobbling::getUserScrobbler(Database::UserId userId)
{ {
std::optional<Database::Scrobbler> scrobbler; std::optional<Database::Scrobbler> scrobbler;
@@ -76,49 +76,49 @@ namespace Scrobbling
return scrobbler; return scrobbler;
} }
std::vector<Wt::Dbo::ptr<Database::Artist>> std::vector<Database::ObjectPtr<Database::Artist>>
Scrobbling::getRecentArtists(Database::Session& session, Scrobbling::getRecentArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType, std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) bool& moreResults)
{ {
const Wt::Dbo::ptr<Database::TrackList> history {getListensTrackList(session, user)}; const Database::ObjectPtr<Database::TrackList> history {getListensTrackList(session, user)};
std::vector<Wt::Dbo::ptr<Database::Artist>> res; std::vector<Database::ObjectPtr<Database::Artist>> res;
if (history) if (history)
res = history->getArtistsReverse(clusterIds, linkType, range, moreResults); res = history->getArtistsReverse(clusterIds, linkType, range, moreResults);
return res; return res;
} }
std::vector<Wt::Dbo::ptr<Database::Release>> std::vector<Database::ObjectPtr<Database::Release>>
Scrobbling::getRecentReleases(Database::Session& session, Scrobbling::getRecentReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) bool& moreResults)
{ {
const Wt::Dbo::ptr<Database::TrackList> history {getListensTrackList(session, user)}; const Database::ObjectPtr<Database::TrackList> history {getListensTrackList(session, user)};
std::vector<Wt::Dbo::ptr<Database::Release>> res; std::vector<Database::ObjectPtr<Database::Release>> res;
if (history) if (history)
res = history->getReleasesReverse(clusterIds, range, moreResults); res = history->getReleasesReverse(clusterIds, range, moreResults);
return res; return res;
} }
std::vector<Wt::Dbo::ptr<Database::Track>> std::vector<Database::ObjectPtr<Database::Track>>
Scrobbling::getRecentTracks(Database::Session& session, Scrobbling::getRecentTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) bool& moreResults)
{ {
const Wt::Dbo::ptr<Database::TrackList> history {getListensTrackList(session, user)}; const Database::ObjectPtr<Database::TrackList> history {getListensTrackList(session, user)};
std::vector<Wt::Dbo::ptr<Database::Track>> res; std::vector<Database::ObjectPtr<Database::Track>> res;
if (history) if (history)
res = history->getTracksReverse(clusterIds, range, moreResults); res = history->getTracksReverse(clusterIds, range, moreResults);
@@ -127,57 +127,57 @@ namespace Scrobbling
// Top // Top
std::vector<Wt::Dbo::ptr<Database::Artist>> std::vector<Database::ObjectPtr<Database::Artist>>
Scrobbling::getTopArtists(Database::Session& session, Scrobbling::getTopArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType, std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) bool& moreResults)
{ {
const Wt::Dbo::ptr<Database::TrackList> history {getListensTrackList(session, user)}; const Database::ObjectPtr<Database::TrackList> history {getListensTrackList(session, user)};
std::vector<Wt::Dbo::ptr<Database::Artist>> res; std::vector<Database::ObjectPtr<Database::Artist>> res;
if (history) if (history)
res = history->getTopArtists(clusterIds, linkType, range, moreResults); res = history->getTopArtists(clusterIds, linkType, range, moreResults);
return res; return res;
} }
std::vector<Wt::Dbo::ptr<Database::Release>> std::vector<Database::ObjectPtr<Database::Release>>
Scrobbling::getTopReleases(Database::Session& session, Scrobbling::getTopReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) bool& moreResults)
{ {
const Wt::Dbo::ptr<Database::TrackList> history {getListensTrackList(session, user)}; const Database::ObjectPtr<Database::TrackList> history {getListensTrackList(session, user)};
std::vector<Wt::Dbo::ptr<Database::Release>> res; std::vector<Database::ObjectPtr<Database::Release>> res;
if (history) if (history)
res = history->getTopReleases(clusterIds, range, moreResults); res = history->getTopReleases(clusterIds, range, moreResults);
return res; return res;
} }
std::vector<Wt::Dbo::ptr<Database::Track>> std::vector<Database::ObjectPtr<Database::Track>>
Scrobbling::getTopTracks(Database::Session& session, Scrobbling::getTopTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) bool& moreResults)
{ {
const Wt::Dbo::ptr<Database::TrackList> history {getListensTrackList(session, user)}; const Database::ObjectPtr<Database::TrackList> history {getListensTrackList(session, user)};
std::vector<Wt::Dbo::ptr<Database::Track>> res; std::vector<Database::ObjectPtr<Database::Track>> res;
if (history) if (history)
res = history->getTopTracks(clusterIds, range, moreResults); res = history->getTopTracks(clusterIds, range, moreResults);
return res; return res;
} }
Wt::Dbo::ptr<Database::TrackList> Database::ObjectPtr<Database::TrackList>
Scrobbling::getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user) Scrobbling::getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user)
{ {
return _scrobblers[user->getScrobbler()]->getListensTrackList(session, user); return _scrobblers[user->getScrobbler()]->getListensTrackList(session, user);
} }
+20 -20
View File
@@ -38,47 +38,47 @@ namespace Scrobbling
void listenFinished(const Listen& listen, std::optional<std::chrono::seconds> duration) override; void listenFinished(const Listen& listen, std::optional<std::chrono::seconds> duration) override;
void addTimedListen(const TimedListen& listen) override; void addTimedListen(const TimedListen& listen) override;
std::vector<Wt::Dbo::ptr<Database::Artist>> getRecentArtists(Database::Session& session, std::vector<Database::ObjectPtr<Database::Artist>> getRecentArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType, std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) override; bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Release>> getRecentReleases(Database::Session& session, std::vector<Database::ObjectPtr<Database::Release>> getRecentReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) override; bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Track>> getRecentTracks(Database::Session& session, std::vector<Database::ObjectPtr<Database::Track>> getRecentTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) override; bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Artist>> getTopArtists(Database::Session& session, std::vector<Database::ObjectPtr<Database::Artist>> getTopArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType, std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) override; bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Release>> getTopReleases(Database::Session& session, std::vector<Database::ObjectPtr<Database::Release>> getTopReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) override; bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Track>> getTopTracks(Database::Session& session, std::vector<Database::ObjectPtr<Database::Track>> getTopTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) override; bool& moreResults) override;
Wt::Dbo::ptr<Database::TrackList> getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user); Database::ObjectPtr<Database::TrackList> getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user);
std::optional<Database::Scrobbler> getUserScrobbler(Database::IdType userId); std::optional<Database::Scrobbler> getUserScrobbler(Database::UserId userId);
Database::Db& _db; Database::Db& _db;
std::unordered_map<Database::Scrobbler, std::unique_ptr<IScrobbler>> _scrobblers; std::unordered_map<Database::Scrobbler, std::unique_ptr<IScrobbler>> _scrobblers;
@@ -61,7 +61,7 @@ namespace Scrobbling
if (!user) if (!user)
return; return;
Wt::Dbo::ptr<Database::TrackList> tracklist {getListensTrackList(session, user)}; Database::TrackList::pointer tracklist {getListensTrackList(session, user)};
if (!tracklist) if (!tracklist)
tracklist = Database::TrackList::create(session, historyTracklistName, Database::TrackList::Type::Internal, false, user); tracklist = Database::TrackList::create(session, historyTracklistName, Database::TrackList::Type::Internal, false, user);
@@ -72,8 +72,8 @@ namespace Scrobbling
Database::TrackListEntry::create(session, track, getListensTrackList(session, user), listen.listenedAt); Database::TrackListEntry::create(session, track, getListensTrackList(session, user), listen.listenedAt);
} }
Wt::Dbo::ptr<Database::TrackList> Database::TrackList::pointer
InternalScrobbler::getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user) InternalScrobbler::getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user)
{ {
return Database::TrackList::get(session, historyTracklistName, Database::TrackList::Type::Internal, user); return Database::TrackList::get(session, historyTracklistName, Database::TrackList::Type::Internal, user);
} }
@@ -21,6 +21,11 @@
#include "IScrobbler.hpp" #include "IScrobbler.hpp"
namespace Database
{
class Db;
}
namespace Scrobbling namespace Scrobbling
{ {
class InternalScrobbler final : public IScrobbler class InternalScrobbler final : public IScrobbler
@@ -34,7 +39,7 @@ namespace Scrobbling
void addTimedListen(const TimedListen& listen) override; void addTimedListen(const TimedListen& listen) override;
Wt::Dbo::ptr<Database::TrackList> getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user) override; Database::ObjectPtr<Database::TrackList> getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user) override;
Database::Db& _db; Database::Db& _db;
}; };
@@ -41,7 +41,7 @@
namespace namespace
{ {
bool bool
canBeScrobbled(Database::Session& session, Database::IdType trackId, std::chrono::seconds duration) canBeScrobbled(Database::Session& session, Database::TrackId trackId, std::chrono::seconds duration)
{ {
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
@@ -50,7 +50,7 @@ namespace Scrobbling::ListenBrainz
void listenStarted(const Listen& listen) override; void listenStarted(const Listen& listen) override;
void listenFinished(const Listen& listen, std::optional<std::chrono::seconds> duration) override; void listenFinished(const Listen& listen, std::optional<std::chrono::seconds> duration) override;
void addTimedListen(const TimedListen& listen) override; void addTimedListen(const TimedListen& listen) override;
Wt::Dbo::ptr<Database::TrackList> getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user) override; Database::ObjectPtr<Database::TrackList> getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user) override;
// Submit listens // Submit listens
void enqueListen(const Listen& listen, const Wt::WDateTime& timePoint); void enqueListen(const Listen& listen, const Wt::WDateTime& timePoint);
@@ -195,7 +195,7 @@ namespace
std::vector<Scrobbling::TimedListen> matchedListens; std::vector<Scrobbling::TimedListen> matchedListens;
}; };
ParseGetListensResult ParseGetListensResult
parseGetListens(Database::Session& session, std::string_view msgBody, Database::IdType userId) parseGetListens(Database::Session& session, std::string_view msgBody, Database::UserId userId)
{ {
ParseGetListensResult result; ParseGetListensResult result;
@@ -233,7 +233,7 @@ namespace
result.oldestEntry = listenedAt; result.oldestEntry = listenedAt;
if (const Database::Track::pointer track {tryMatchListen(session, metadata)}) if (const Database::Track::pointer track {tryMatchListen(session, metadata)})
result.matchedListens.emplace_back(Scrobbling::TimedListen {userId, track.id(), listenedAt}); result.matchedListens.emplace_back(Scrobbling::TimedListen {userId, track->getId(), listenedAt});
} }
} }
catch (const Wt::WException& error) catch (const Wt::WException& error)
@@ -285,7 +285,7 @@ namespace Scrobbling::ListenBrainz
} }
ListensSynchronizer::UserContext& ListensSynchronizer::UserContext&
ListensSynchronizer::getUserContext(Database::IdType userId) ListensSynchronizer::getUserContext(Database::UserId userId)
{ {
auto itContext {_userContexts.find(userId)}; auto itContext {_userContexts.find(userId)};
if (itContext == std::cend(_userContexts)) if (itContext == std::cend(_userContexts))
@@ -338,14 +338,14 @@ namespace Scrobbling::ListenBrainz
assert(!isFetching()); assert(!isFetching());
std::vector<Database::IdType> userIds; std::vector<Database::UserId> userIds;
{ {
Database::Session& session {_db.getTLSSession()}; Database::Session& session {_db.getTLSSession()};
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
userIds = Database::User::getAllIds(_db.getTLSSession()); userIds = Database::User::getAllIds(_db.getTLSSession());
} }
for (const Database::IdType userId : userIds) for (const Database::UserId userId : userIds)
{ {
if (Utils::getListenBrainzToken(_db.getTLSSession(), userId)) if (Utils::getListenBrainzToken(_db.getTLSSession(), userId))
startGetListens(getUserContext(userId)); startGetListens(getUserContext(userId));
@@ -373,7 +373,7 @@ namespace Scrobbling::ListenBrainz
{ {
_strand.dispatch([this, &context] _strand.dispatch([this, &context]
{ {
LOG(DEBUG) << "Fetch done for user " << context.userId << ", fetched: " << context.fetchedListenCount << ", matched: " << context.matchedListenCount << ", imported: " << context.importedListenCount; LOG(DEBUG) << "Fetch done for user " << context.userId.getValue() << ", fetched: " << context.fetchedListenCount << ", matched: " << context.matchedListenCount << ", imported: " << context.importedListenCount;
context.fetching = false; context.fetching = false;
if (!isFetching()) if (!isFetching())
@@ -473,7 +473,7 @@ namespace Scrobbling::ListenBrainz
} }
std::optional<SendQueue::RequestData> std::optional<SendQueue::RequestData>
ListensSynchronizer::createValidateTokenRequestData(Database::IdType userId) ListensSynchronizer::createValidateTokenRequestData(Database::UserId userId)
{ {
Database::Session& session {_db.getTLSSession()}; Database::Session& session {_db.getTLSSession()};
@@ -50,14 +50,14 @@ namespace Scrobbling::ListenBrainz
private: private:
struct UserContext struct UserContext
{ {
UserContext(Database::IdType id) : userId {id} {} UserContext(Database::UserId id) : userId {id} {}
UserContext(const UserContext&) = delete; UserContext(const UserContext&) = delete;
UserContext(UserContext&&) = delete; UserContext(UserContext&&) = delete;
UserContext& operator=(const UserContext&) = delete; UserContext& operator=(const UserContext&) = delete;
UserContext& operator=(UserContext&&) = delete; UserContext& operator=(UserContext&&) = delete;
const Database::IdType userId; const Database::UserId userId;
bool fetching {}; bool fetching {};
std::optional<std::size_t> listenCount {}; std::optional<std::size_t> listenCount {};
@@ -69,7 +69,7 @@ namespace Scrobbling::ListenBrainz
std::size_t importedListenCount{}; std::size_t importedListenCount{};
}; };
UserContext& getUserContext(Database::IdType userId); UserContext& getUserContext(Database::UserId userId);
bool isFetching() const; bool isFetching() const;
void scheduleGetListens(std::chrono::seconds fromNow); void scheduleGetListens(std::chrono::seconds fromNow);
void startGetListens(); void startGetListens();
@@ -78,7 +78,7 @@ namespace Scrobbling::ListenBrainz
void enqueValidateToken(UserContext& context); void enqueValidateToken(UserContext& context);
void enqueGetListenCount(UserContext& context); void enqueGetListenCount(UserContext& context);
void enqueGetListens(UserContext& context); void enqueGetListens(UserContext& context);
std::optional<SendQueue::RequestData> createValidateTokenRequestData(Database::IdType userId); std::optional<SendQueue::RequestData> createValidateTokenRequestData(Database::UserId userId);
std::optional<SendQueue::RequestData> createGetListensRequestData(std::string_view listenBrainzUserName, const Wt::WDateTime& maxDateTime); std::optional<SendQueue::RequestData> createGetListensRequestData(std::string_view listenBrainzUserName, const Wt::WDateTime& maxDateTime);
void processGetListensResponse(std::string_view body, UserContext& context); void processGetListensResponse(std::string_view body, UserContext& context);
@@ -88,7 +88,7 @@ namespace Scrobbling::ListenBrainz
SendQueue& _sendQueue; SendQueue& _sendQueue;
boost::asio::steady_timer _getListensTimer {_ioContext}; boost::asio::steady_timer _getListensTimer {_ioContext};
std::unordered_map<Database::IdType, UserContext> _userContexts; std::unordered_map<Database::UserId, UserContext> _userContexts;
const std::size_t _maxSyncListenCount; const std::size_t _maxSyncListenCount;
const std::chrono::hours _syncListensPeriod; const std::chrono::hours _syncListensPeriod;
@@ -30,7 +30,7 @@ static constexpr std::string_view historyTracklistName {"__scrobbler_listenbrain
namespace Scrobbling::ListenBrainz::Utils namespace Scrobbling::ListenBrainz::Utils
{ {
std::optional<UUID> std::optional<UUID>
getListenBrainzToken(Database::Session& session, Database::IdType userId) getListenBrainzToken(Database::Session& session, Database::UserId userId)
{ {
auto transaction {session.createSharedTransaction()}; auto transaction {session.createSharedTransaction()};
@@ -21,6 +21,7 @@
#include <Wt/Dbo/ptr.h> #include <Wt/Dbo/ptr.h>
#include "utils/UUID.hpp" #include "utils/UUID.hpp"
#include "database/Types.hpp" #include "database/Types.hpp"
namespace Database namespace Database
@@ -32,7 +33,7 @@ namespace Database
namespace Scrobbling::ListenBrainz::Utils namespace Scrobbling::ListenBrainz::Utils
{ {
std::optional<UUID> getListenBrainzToken(Database::Session& session, Database::IdType userId); std::optional<UUID> getListenBrainzToken(Database::Session& session, Database::UserId userId);
Wt::Dbo::ptr<Database::TrackList> getOrCreateListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user); Database::ObjectPtr<Database::TrackList> getOrCreateListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user);
Wt::Dbo::ptr<Database::TrackList> getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user); Database::ObjectPtr<Database::TrackList> getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user);
} }
@@ -24,12 +24,12 @@
#include <chrono> #include <chrono>
#include <memory> #include <memory>
#include <optional> #include <optional>
#include <set>
#include <vector> #include <vector>
#include <Wt/WDateTime.h> #include <Wt/WDateTime.h>
#include "scrobbling/Listen.hpp" #include "scrobbling/Listen.hpp"
#include "database/Types.hpp"
namespace Database namespace Database
{ {
@@ -57,42 +57,42 @@ namespace Scrobbling
// Stats // Stats
// From most recent to oldest // From most recent to oldest
virtual std::vector<Wt::Dbo::ptr<Database::Artist>> getRecentArtists(Database::Session& session, virtual std::vector<Database::ObjectPtr<Database::Artist>> getRecentArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType, std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) = 0; bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Release>> getRecentReleases(Database::Session& session, virtual std::vector<Database::ObjectPtr<Database::Release>> getRecentReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) = 0; bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Track>> getRecentTracks(Database::Session& session, virtual std::vector<Database::ObjectPtr<Database::Track>> getRecentTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) = 0; bool& moreResults) = 0;
// Top // Top
virtual std::vector<Wt::Dbo::ptr<Database::Artist>> getTopArtists(Database::Session& session, virtual std::vector<Database::ObjectPtr<Database::Artist>> getTopArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType, std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) = 0; bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Release>> getTopReleases(Database::Session& session, virtual std::vector<Database::ObjectPtr<Database::Release>> getTopReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) = 0; bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Track>> getTopTracks(Database::Session& session, virtual std::vector<Database::ObjectPtr<Database::Track>> getTopTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user, Database::ObjectPtr<Database::User> user,
const std::set<Database::IdType>& clusterIds, const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range, std::optional<Database::Range> range,
bool& moreResults) = 0; bool& moreResults) = 0;
}; };
@@ -27,8 +27,8 @@ namespace Scrobbling
{ {
struct Listen struct Listen
{ {
Database::IdType userId {}; Database::UserId userId {};
Database::IdType trackId {}; Database::TrackId trackId {};
}; };
struct TimedListen : public Listen struct TimedListen : public Listen
+2 -1
View File
@@ -24,6 +24,7 @@
#include <cmath> #include <cmath>
#include <random> #include <random>
#include <sstream> #include <sstream>
#include <unordered_set>
#include "utils/Logger.hpp" #include "utils/Logger.hpp"
#include "utils/Random.hpp" #include "utils/Random.hpp"
@@ -197,7 +198,7 @@ Network::getClosestRefVectorPosition(const InputVector& data, InputVector::Dista
} }
std::optional<Position> std::optional<Position>
Network::getClosestRefVectorPosition(const std::unordered_set<Position>& refVectorsPosition, InputVector::Distance maxDistance) const Network::getClosestRefVectorPosition(const std::vector<Position>& refVectorsPosition, InputVector::Distance maxDistance) const
{ {
std::unordered_set<Position> neighboursPosition; std::unordered_set<Position> neighboursPosition;
for (const Position& refVectorPosition : refVectorsPosition) for (const Position& refVectorPosition : refVectorsPosition)
+1 -2
View File
@@ -20,7 +20,6 @@
#pragma once #pragma once
#include <vector> #include <vector>
#include <unordered_set>
#include <optional> #include <optional>
#include <ostream> #include <ostream>
#include <functional> #include <functional>
@@ -70,7 +69,7 @@ class Network
Position getClosestRefVectorPosition(const InputVector& data) const; Position getClosestRefVectorPosition(const InputVector& data) const;
std::optional<Position> getClosestRefVectorPosition(const InputVector& data, InputVector::Distance maxDistance) const; std::optional<Position> getClosestRefVectorPosition(const InputVector& data, InputVector::Distance maxDistance) const;
std::optional<Position> getClosestRefVectorPosition(const std::unordered_set<Position>& refVectorsPosition, InputVector::Distance maxDistance) const; std::optional<Position> getClosestRefVectorPosition(const std::vector<Position>& refVectorsPosition, InputVector::Distance maxDistance) const;
InputVector::Distance getRefVectorsDistance(const Position& position1, const Position& position2) const; InputVector::Distance getRefVectorsDistance(const Position& position1, const Position& position2) const;
-1
View File
@@ -1,6 +1,5 @@
add_library(lmssubsonic SHARED add_library(lmssubsonic SHARED
impl/ParameterParsing.cpp
impl/Scan.cpp impl/Scan.cpp
impl/Stream.cpp impl/Stream.cpp
impl/SubsonicId.cpp impl/SubsonicId.cpp
+1 -14
View File
@@ -20,8 +20,8 @@
#include <Wt/Http/Request.h> #include <Wt/Http/Request.h>
#include "database/Types.hpp"
#include "utils/String.hpp" #include "utils/String.hpp"
#include "SubsonicId.hpp"
#include "SubsonicResponse.hpp" #include "SubsonicResponse.hpp"
namespace API::Subsonic namespace API::Subsonic
@@ -82,18 +82,5 @@ namespace API::Subsonic
return *res; return *res;
} }
}
namespace StringUtils
{
template<>
std::optional<API::Subsonic::Id>
readAs(std::string_view str);
template<>
std::optional<bool>
readAs(std::string_view str);
} }
+1 -1
View File
@@ -36,7 +36,7 @@ namespace API::Subsonic
{ {
const Wt::Http::ParameterMap& parameters; const Wt::Http::ParameterMap& parameters;
Database::Session& dbSession; Database::Session& dbSession;
Database::IdType userId; Database::UserId userId;
std::string clientName; std::string clientName;
}; };
} }
+5 -5
View File
@@ -63,7 +63,7 @@ StreamParameters
getStreamParameters(RequestContext& context) getStreamParameters(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; const TrackId id {getMandatoryParameterAs<TrackId>(context.parameters, "id")};
// Optional params // Optional params
std::optional<std::size_t> maxBitRate {getParameterAs<std::size_t>(context.parameters, "maxBitRate")}; std::optional<std::size_t> maxBitRate {getParameterAs<std::size_t>(context.parameters, "maxBitRate")};
@@ -74,7 +74,7 @@ getStreamParameters(RequestContext& context)
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
{ {
auto track {Track::getById(context.dbSession, id.value)}; auto track {Track::getById(context.dbSession, id)};
if (!track) if (!track)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -94,7 +94,7 @@ getStreamParameters(RequestContext& context)
// "If set to zero, no limit is imposed" // "If set to zero, no limit is imposed"
if (maxBitRate && *maxBitRate != 0) if (maxBitRate && *maxBitRate != 0)
bitRate = clamp(*maxBitRate, std::size_t {48}, bitRate); bitRate = Utils::clamp(*maxBitRate, std::size_t {48}, bitRate);
Av::TranscodeParameters transcodeParameters; Av::TranscodeParameters transcodeParameters;
@@ -118,13 +118,13 @@ handleDownload(RequestContext& context, const Wt::Http::Request& request, Wt::Ht
if (!continuation) if (!continuation)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; Database::TrackId id {getMandatoryParameterAs<Database::TrackId>(context.parameters, "id")};
std::filesystem::path trackPath; std::filesystem::path trackPath;
{ {
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
auto track {Track::getById(context.dbSession, id.value)}; auto track {Track::getById(context.dbSession, id)};
if (!track) if (!track)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
+119 -56
View File
@@ -26,64 +26,127 @@
namespace API::Subsonic namespace API::Subsonic
{ {
std::string
std::optional<Id> idToString(Database::ArtistId id)
IdFromString(const std::string_view id)
{
if (id == "root")
return Id {Id::Type::Root};
std::vector<std::string_view> values {StringUtils::splitString(id, "-")};
if (values.size() != 2)
return std::nullopt;
Id res;
const std::string type {std::move(values[0])};
if (type == "ar")
res.type = Id::Type::Artist;
else if (type == "al")
res.type = Id::Type::Release;
else if (type == "tr")
res.type = Id::Type::Track;
else if (type == "pl")
res.type = Id::Type::Playlist;
else
return std::nullopt;
auto optId {StringUtils::readAs<Database::IdType>(values[1])};
if (!optId)
return std::nullopt;
res.value = *optId;
return res;
}
std::string
IdToString(const Id& id)
{
std::string res;
switch (id.type)
{ {
case Id::Type::Root: return "ar-" + id.toString();
return "root";
case Id::Type::Artist:
res = "ar-";
break;
case Id::Type::Release:
res = "al-";
break;
case Id::Type::Track:
res = "tr-";
break;
case Id::Type::Playlist:
res = "pl-";
break;
} }
return res + std::to_string(id.value); std::string
idToString(Database::ReleaseId id)
{
return "al-" + id.toString();
}
std::string
idToString(RootId)
{
return "root";
}
std::string
idToString(Database::TrackId id)
{
return "tr-" + id.toString();
}
std::string
idToString(Database::TrackListId id)
{
return "pl-" + id.toString();
}
} // namespace API::Subsonic
namespace StringUtils
{
template<>
std::optional<Database::ArtistId>
readAs(std::string_view str)
{
std::vector<std::string_view> values {StringUtils::splitString(str, "-")};
if (values.size() != 2)
return std::nullopt;
if (values[0] != "ar")
return std::nullopt;
if (const auto value {StringUtils::readAs<Database::ArtistId::ValueType>(values[1])})
return Database::ArtistId {*value};
return std::nullopt;
}
template<>
std::optional<Database::ReleaseId>
readAs(std::string_view str)
{
std::vector<std::string_view> values {StringUtils::splitString(str, "-")};
if (values.size() != 2)
return std::nullopt;
if (values[0] != "al")
return std::nullopt;
if (const auto value {StringUtils::readAs<Database::ReleaseId::ValueType>(values[1])})
return Database::ReleaseId {*value};
return std::nullopt;
}
template<>
std::optional<API::Subsonic::RootId>
readAs(std::string_view str)
{
if (str == "root")
return API::Subsonic::RootId {};
return std::nullopt;
}
template<>
std::optional<Database::TrackId>
readAs(std::string_view str)
{
std::vector<std::string_view> values {StringUtils::splitString(str, "-")};
if (values.size() != 2)
return std::nullopt;
if (values[0] != "tr")
return std::nullopt;
if (const auto value {StringUtils::readAs<Database::TrackId::ValueType>(values[1])})
return Database::TrackId {*value};
return std::nullopt;
}
template<>
std::optional<Database::TrackListId>
readAs(std::string_view str)
{
std::vector<std::string_view> values {StringUtils::splitString(str, "-")};
if (values.size() != 2)
return std::nullopt;
if (values[0] != "pl")
return std::nullopt;
if (const auto value {StringUtils::readAs<Database::TrackListId::ValueType>(values[1])})
return Database::TrackListId {*value};
return std::nullopt;
}
template<>
std::optional<bool>
readAs(std::string_view str)
{
if (str == "true")
return true;
else if (str == "false")
return false;
return {};
}
} }
} // namespace API::Subsonic
+36 -20
View File
@@ -19,29 +19,45 @@
#pragma once #pragma once
#include <optional>
#include "database/Types.hpp" #include "database/Types.hpp"
#include "utils/String.hpp"
namespace API::Subsonic namespace API::Subsonic
{ {
struct RootId {};
struct Id std::string idToString(Database::ArtistId id);
{ std::string idToString(Database::ReleaseId id);
enum class Type std::string idToString(Database::TrackId id);
{ std::string idToString(Database::TrackListId id);
Root, // Where all artists artistless albums reside std::string idToString(RootId);
Track,
Release,
Artist,
Playlist,
};
Type type;
Database::IdType value {};
};
std::optional<Id> IdFromString(std::string_view id);
std::string IdToString(const Id& id);
} // namespace API::Subsonic } // namespace API::Subsonic
// Used to parse parameters
namespace StringUtils
{
template<>
std::optional<API::Subsonic::RootId>
readAs(std::string_view str);
template<>
std::optional<Database::ArtistId>
readAs(std::string_view str);
template<>
std::optional<Database::ReleaseId>
readAs(std::string_view str);
template<>
std::optional<Database::TrackId>
readAs(std::string_view str);
template<>
std::optional<Database::TrackListId>
readAs(std::string_view str);
template<>
std::optional<bool>
readAs(std::string_view str);
}
+141 -206
View File
@@ -48,6 +48,7 @@
#include "RequestContext.hpp" #include "RequestContext.hpp"
#include "Scan.hpp" #include "Scan.hpp"
#include "Stream.hpp" #include "Stream.hpp"
#include "SubsonicId.hpp"
#include "SubsonicResponse.hpp" #include "SubsonicResponse.hpp"
using namespace Database; using namespace Database;
@@ -321,7 +322,7 @@ trackToResponseNode(const Track::pointer& track, Session& dbSession, const User:
{ {
Response::Node trackResponse; Response::Node trackResponse;
trackResponse.setAttribute("id", IdToString({Id::Type::Track, track.id()})); trackResponse.setAttribute("id", idToString(track->getId()));
trackResponse.setAttribute("isDir", false); trackResponse.setAttribute("isDir", false);
trackResponse.setAttribute("title", track->getName()); trackResponse.setAttribute("title", track->getName());
if (track->getTrackNumber()) if (track->getTrackNumber())
@@ -348,7 +349,7 @@ trackToResponseNode(const Track::pointer& track, Session& dbSession, const User:
if (user->getSubsonicTranscodeEnable()) if (user->getSubsonicTranscodeEnable())
trackResponse.setAttribute("transcodedSuffix", formatToSuffix(user->getSubsonicTranscodeFormat())); trackResponse.setAttribute("transcodedSuffix", formatToSuffix(user->getSubsonicTranscodeFormat()));
trackResponse.setAttribute("coverArt", IdToString({Id::Type::Track, track.id()})); trackResponse.setAttribute("coverArt", idToString(track->getId()));
const std::vector<Artist::pointer>& artists {track->getArtists({TrackArtistLinkType::Artist})}; const std::vector<Artist::pointer>& artists {track->getArtists({TrackArtistLinkType::Artist})};
if (!artists.empty()) if (!artists.empty())
@@ -356,14 +357,14 @@ trackToResponseNode(const Track::pointer& track, Session& dbSession, const User:
trackResponse.setAttribute("artist", getArtistNames(artists)); trackResponse.setAttribute("artist", getArtistNames(artists));
if (artists.size() == 1) if (artists.size() == 1)
trackResponse.setAttribute("artistId", IdToString({Id::Type::Artist, artists.front().id()})); trackResponse.setAttribute("artistId", idToString(artists.front()->getId()));
} }
if (track->getRelease()) if (track->getRelease())
{ {
trackResponse.setAttribute("album", track->getRelease()->getName()); trackResponse.setAttribute("album", track->getRelease()->getName());
trackResponse.setAttribute("albumId", IdToString({Id::Type::Release, track->getRelease().id()})); trackResponse.setAttribute("albumId", idToString(track->getRelease()->getId()));
trackResponse.setAttribute("parent", IdToString({Id::Type::Release, track->getRelease().id()})); trackResponse.setAttribute("parent", idToString(track->getRelease()->getId()));
} }
trackResponse.setAttribute("duration", std::chrono::duration_cast<std::chrono::seconds>(track->getDuration()).count()); trackResponse.setAttribute("duration", std::chrono::duration_cast<std::chrono::seconds>(track->getDuration()).count());
@@ -420,8 +421,8 @@ releaseToResponseNode(const Release::pointer& release, Session& dbSession, const
} }
albumNode.setAttribute("created", dateTimeToCreatedString(release->getLastWritten())); albumNode.setAttribute("created", dateTimeToCreatedString(release->getLastWritten()));
albumNode.setAttribute("id", IdToString({Id::Type::Release, release.id()})); albumNode.setAttribute("id", idToString(release->getId()));
albumNode.setAttribute("coverArt", IdToString({Id::Type::Release, release.id()})); albumNode.setAttribute("coverArt", idToString(release->getId()));
auto releaseYear {release->getReleaseYear()}; auto releaseYear {release->getReleaseYear()};
if (releaseYear) if (releaseYear)
albumNode.setAttribute("year", *releaseYear); albumNode.setAttribute("year", *releaseYear);
@@ -432,7 +433,7 @@ releaseToResponseNode(const Release::pointer& release, Session& dbSession, const
if (artists.empty() && !id3) if (artists.empty() && !id3)
{ {
albumNode.setAttribute("parent", IdToString({Id::Type::Root})); albumNode.setAttribute("parent", idToString(RootId {}));
} }
else if (!artists.empty()) else if (!artists.empty())
{ {
@@ -440,15 +441,12 @@ releaseToResponseNode(const Release::pointer& release, Session& dbSession, const
if (artists.size() == 1) if (artists.size() == 1)
{ {
if (id3) albumNode.setAttribute(id3 ? "artistId" : "parent", idToString(artists.front()->getId()));
albumNode.setAttribute("artistId", IdToString({Id::Type::Artist, artists.front().id()}));
else
albumNode.setAttribute("parent", IdToString({Id::Type::Artist, artists.front().id()}));
} }
else else
{ {
if (!id3) if (!id3)
albumNode.setAttribute("parent", IdToString({Id::Type::Root})); albumNode.setAttribute("parent", idToString(RootId {}));
} }
} }
@@ -476,7 +474,7 @@ artistToResponseNode(const User::pointer& user, const Artist::pointer& artist, b
{ {
Response::Node artistNode; Response::Node artistNode;
artistNode.setAttribute("id", IdToString({Id::Type::Artist, artist.id()})); artistNode.setAttribute("id", idToString(artist->getId()));
artistNode.setAttribute("name", artist->getName()); artistNode.setAttribute("name", artist->getName());
if (id3) if (id3)
@@ -544,7 +542,7 @@ handleChangePassword(RequestContext& context)
try try
{ {
Database::IdType userId; Database::UserId userId;
{ {
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
@@ -554,7 +552,7 @@ handleChangePassword(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
userId = user.id(); userId = user->getId();
} }
Service<Auth::IPasswordService>::get()->setPassword(context.dbSession, userId, password); Service<Auth::IPasswordService>::get()->setPassword(context.dbSession, userId, password);
@@ -580,15 +578,10 @@ Response
handleCreatePlaylistRequest(RequestContext& context) handleCreatePlaylistRequest(RequestContext& context)
{ {
// Optional params // Optional params
auto id {getParameterAs<Id>(context.parameters, "playlistId")}; const auto id {getParameterAs<TrackListId>(context.parameters, "playlistId")};
if (id && id->type != Id::Type::Playlist)
throw BadParameterGenericError {"playlistId"};
auto name {getParameterAs<std::string>(context.parameters, "name")}; auto name {getParameterAs<std::string>(context.parameters, "name")};
std::vector<Id> trackIds {getMultiParametersAs<Id>(context.parameters, "songId")}; std::vector<TrackId> trackIds {getMultiParametersAs<TrackId>(context.parameters, "songId")};
if (!std::all_of(std::cbegin(trackIds), std::cend(trackIds ), [](const Id& id) { return id.type == Id::Type::Track; }))
throw BadParameterGenericError {"songId"};
if (!name && !id) if (!name && !id)
throw RequiredParameterMissingError {"name or id"}; throw RequiredParameterMissingError {"name or id"};
@@ -602,7 +595,7 @@ handleCreatePlaylistRequest(RequestContext& context)
TrackList::pointer tracklist; TrackList::pointer tracklist;
if (id) if (id)
{ {
tracklist = TrackList::getById(context.dbSession, id->value); tracklist = TrackList::getById(context.dbSession, *id);
if (!tracklist if (!tracklist
|| tracklist->getUser() != user || tracklist->getUser() != user
|| tracklist->getType() != TrackList::Type::Playlist) || tracklist->getType() != TrackList::Type::Playlist)
@@ -618,9 +611,9 @@ handleCreatePlaylistRequest(RequestContext& context)
tracklist = TrackList::create(context.dbSession, *name, TrackList::Type::Playlist, false, user); tracklist = TrackList::create(context.dbSession, *name, TrackList::Type::Playlist, false, user);
} }
for (const Id& trackId : trackIds) for (const TrackId trackId : trackIds)
{ {
Track::pointer track {Track::getById(context.dbSession, trackId.value)}; Track::pointer track {Track::getById(context.dbSession, trackId)};
if (!track) if (!track)
continue; continue;
@@ -638,7 +631,7 @@ handleCreateUserRequest(RequestContext& context)
std::string password {decodePasswordIfNeeded(getMandatoryParameterAs<std::string>(context.parameters, "password"))}; std::string password {decodePasswordIfNeeded(getMandatoryParameterAs<std::string>(context.parameters, "password"))};
// Just ignore all the other fields as we don't handle them // Just ignore all the other fields as we don't handle them
Database::IdType userId; Database::UserId userId;
{ {
auto transaction {context.dbSession.createUniqueTransaction()}; auto transaction {context.dbSession.createUniqueTransaction()};
@@ -647,7 +640,7 @@ handleCreateUserRequest(RequestContext& context)
throw UserAlreadyExistsGenericError {}; throw UserAlreadyExistsGenericError {};
user = User::create(context.dbSession, username); user = User::create(context.dbSession, username);
userId = user.id(); userId = user->getId();
} }
auto removeCreatedUser {[&]() auto removeCreatedUser {[&]()
@@ -685,9 +678,7 @@ static
Response Response
handleDeletePlaylistRequest(RequestContext& context) handleDeletePlaylistRequest(RequestContext& context)
{ {
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; TrackListId id {getMandatoryParameterAs<TrackListId>(context.parameters, "id")};
if (id.type != Id::Type::Playlist)
throw BadParameterGenericError {"id"};
auto transaction {context.dbSession.createUniqueTransaction()}; auto transaction {context.dbSession.createUniqueTransaction()};
@@ -695,7 +686,7 @@ handleDeletePlaylistRequest(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
TrackList::pointer tracklist {TrackList::getById(context.dbSession, id.value)}; TrackList::pointer tracklist {TrackList::getById(context.dbSession, id)};
if (!tracklist if (!tracklist
|| tracklist->getUser() != user || tracklist->getUser() != user
|| tracklist->getType() != TrackList::Type::Playlist) || tracklist->getType() != TrackList::Type::Playlist)
@@ -721,7 +712,7 @@ handleDeleteUserRequest(RequestContext& context)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
// cannot delete ourself // cannot delete ourself
if (user.id() == context.userId) if (user->getId() == context.userId)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
user.remove(); user.remove();
@@ -809,7 +800,7 @@ handleGetAlbumListRequestCommon(const RequestContext& context, bool id3)
if (cluster) if (cluster)
{ {
bool more; bool more;
releases = Release::getByFilter(context.dbSession, {cluster.id()}, {}, range, more); releases = Release::getByFilter(context.dbSession, {cluster->getId()}, {}, range, more);
} }
} }
} }
@@ -876,14 +867,11 @@ Response
handleGetAlbumRequest(RequestContext& context) handleGetAlbumRequest(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; ReleaseId id {getMandatoryParameterAs<ReleaseId>(context.parameters, "id")};
if (id.type != Id::Type::Release)
throw BadParameterGenericError {"id"};
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
Release::pointer release {Release::getById(context.dbSession, id.value)}; Release::pointer release {Release::getById(context.dbSession, id)};
if (!release) if (!release)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -908,14 +896,11 @@ Response
handleGetArtistRequest(RequestContext& context) handleGetArtistRequest(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; ArtistId id {getMandatoryParameterAs<ArtistId>(context.parameters, "id")};
if (id.type != Id::Type::Artist)
throw BadParameterGenericError {"id"};
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
Artist::pointer artist {Artist::getById(context.dbSession, id.value)}; Artist::pointer artist {Artist::getById(context.dbSession, id)};
if (!artist) if (!artist)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -940,9 +925,7 @@ Response
handleGetArtistInfoRequestCommon(RequestContext& context, bool id3) handleGetArtistInfoRequestCommon(RequestContext& context, bool id3)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; ArtistId id {getMandatoryParameterAs<ArtistId>(context.parameters, "id")};
if (id.type != Id::Type::Artist)
throw BadParameterGenericError {"id"};
// Optional params // Optional params
std::size_t count {getParameterAs<std::size_t>(context.parameters, "count").value_or(20)}; std::size_t count {getParameterAs<std::size_t>(context.parameters, "count").value_or(20)};
@@ -953,7 +936,7 @@ handleGetArtistInfoRequestCommon(RequestContext& context, bool id3)
{ {
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
Artist::pointer artist {Artist::getById(context.dbSession, id.value)}; Artist::pointer artist {Artist::getById(context.dbSession, id)};
if (!artist) if (!artist)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -963,7 +946,7 @@ handleGetArtistInfoRequestCommon(RequestContext& context, bool id3)
} }
auto similarArtistsId {Service<Recommendation::IEngine>::get()->getSimilarArtists(context.dbSession, auto similarArtistsId {Service<Recommendation::IEngine>::get()->getSimilarArtists(context.dbSession,
id.value, id,
{TrackArtistLinkType::Artist, TrackArtistLinkType::ReleaseArtist}, {TrackArtistLinkType::Artist, TrackArtistLinkType::ReleaseArtist},
count)}; count)};
@@ -974,7 +957,7 @@ handleGetArtistInfoRequestCommon(RequestContext& context, bool id3)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
for ( const auto& similarArtistId : similarArtistsId ) for ( const ArtistId similarArtistId : similarArtistsId )
{ {
Artist::pointer similarArtist {Artist::getById(context.dbSession, similarArtistId)}; Artist::pointer similarArtist {Artist::getById(context.dbSession, similarArtistId)};
if (similarArtist) if (similarArtist)
@@ -1049,66 +1032,63 @@ Response
handleGetMusicDirectoryRequest(RequestContext& context) handleGetMusicDirectoryRequest(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; const auto artistId {getParameterAs<ArtistId>(context.parameters, "id")};
const auto releaseId {getParameterAs<ReleaseId>(context.parameters, "id")};
const auto trackId {getParameterAs<TrackId>(context.parameters, "id")};
const auto root {getParameterAs<RootId>(context.parameters, "id")};
if (!root && !artistId && !releaseId && !trackId)
throw BadParameterGenericError {"id"};
Response response {Response::createOkResponse(context)}; Response response {Response::createOkResponse(context)};
Response::Node& directoryNode {response.createNode("directory")}; Response::Node& directoryNode {response.createNode("directory")};
directoryNode.setAttribute("id", IdToString(id));
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
User::pointer user {User::getById(context.dbSession, context.userId)}; User::pointer user {User::getById(context.dbSession, context.userId)};
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
switch (id.type) if (root)
{ {
case Id::Type::Root: directoryNode.setAttribute("id", idToString(RootId {}));
{ directoryNode.setAttribute("name", "Music");
directoryNode.setAttribute("name", "Music");
bool moreResults{}; bool moreResults{};
auto artists {Artist::getAll(context.dbSession, Artist::SortMethod::BySortName, std::nullopt, moreResults)}; auto artists {Artist::getAll(context.dbSession, Artist::SortMethod::BySortName, std::nullopt, moreResults)};
for (const Artist::pointer& artist : artists) for (const Artist::pointer& artist : artists)
directoryNode.addArrayChild("child", artistToResponseNode(user, artist, false /* no id3 */)); directoryNode.addArrayChild("child", artistToResponseNode(user, artist, false /* no id3 */));
break;
}
case Id::Type::Artist:
{
auto artist {Artist::getById(context.dbSession, id.value)};
if (!artist)
throw RequestedDataNotFoundError {};
directoryNode.setAttribute("name", makeNameFilesystemCompatible(artist->getName()));
auto releases {artist->getReleases()};
for (const Release::pointer& release : releases)
directoryNode.addArrayChild("child", releaseToResponseNode(release, context.dbSession, user, false /* no id3 */));
break;
}
case Id::Type::Release:
{
auto release {Release::getById(context.dbSession, id.value)};
if (!release)
throw RequestedDataNotFoundError {};
directoryNode.setAttribute("name", makeNameFilesystemCompatible(release->getName()));
auto tracks {release->getTracks()};
for (const Track::pointer& track : tracks)
directoryNode.addArrayChild("child", trackToResponseNode(track, context.dbSession, user));
break;
}
default:
throw BadParameterGenericError {"id"};
} }
else if (artistId)
{
directoryNode.setAttribute("id", idToString(*artistId));
auto artist {Artist::getById(context.dbSession, *artistId)};
if (!artist)
throw RequestedDataNotFoundError {};
directoryNode.setAttribute("name", makeNameFilesystemCompatible(artist->getName()));
auto releases {artist->getReleases()};
for (const Release::pointer& release : releases)
directoryNode.addArrayChild("child", releaseToResponseNode(release, context.dbSession, user, false /* no id3 */));
}
else if (releaseId)
{
directoryNode.setAttribute("id", idToString(*releaseId));
auto release {Release::getById(context.dbSession, *releaseId)};
if (!release)
throw RequestedDataNotFoundError {};
directoryNode.setAttribute("name", makeNameFilesystemCompatible(release->getName()));
auto tracks {release->getTracks()};
for (const Track::pointer& track : tracks)
directoryNode.addArrayChild("child", trackToResponseNode(track, context.dbSession, user));
}
else
throw BadParameterGenericError {"id"};
return response; return response;
} }
@@ -1199,21 +1179,19 @@ Response
handleGetSimilarSongsRequestCommon(RequestContext& context, bool id3) handleGetSimilarSongsRequestCommon(RequestContext& context, bool id3)
{ {
// Mandatory params // Mandatory params
const Id artistId {getMandatoryParameterAs<Id>(context.parameters, "id")}; const ArtistId artistId {getMandatoryParameterAs<ArtistId>(context.parameters, "id")};
if (artistId.type != Id::Type::Artist)
throw BadParameterGenericError {"id"};
// Optional params // Optional params
std::size_t count {getParameterAs<std::size_t>(context.parameters, "count").value_or(50)}; std::size_t count {getParameterAs<std::size_t>(context.parameters, "count").value_or(50)};
auto similarArtistIds {Service<Recommendation::IEngine>::get()->getSimilarArtists(context.dbSession, const auto similarArtistIds {Service<Recommendation::IEngine>::get()->getSimilarArtists(context.dbSession,
artistId.value, artistId,
{TrackArtistLinkType::Artist, TrackArtistLinkType::ReleaseArtist}, {TrackArtistLinkType::Artist, TrackArtistLinkType::ReleaseArtist},
5)}; 5)};
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
const Artist::pointer artist {Artist::getById(context.dbSession, artistId.value)}; const Artist::pointer artist {Artist::getById(context.dbSession, artistId)};
if (!artist) if (!artist)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -1223,7 +1201,7 @@ handleGetSimilarSongsRequestCommon(RequestContext& context, bool id3)
// "Returns a random collection of songs from the given artist and similar artists" // "Returns a random collection of songs from the given artist and similar artists"
auto tracks {artist->getRandomTracks(count / 2)}; auto tracks {artist->getRandomTracks(count / 2)};
for (const Database::IdType similarArtistId : similarArtistIds) for (const ArtistId similarArtistId : similarArtistIds)
{ {
const Artist::pointer similarArtist {Artist::getById(context.dbSession, similarArtistId)}; const Artist::pointer similarArtist {Artist::getById(context.dbSession, similarArtistId)};
if (!similarArtist) if (!similarArtist)
@@ -1318,7 +1296,7 @@ tracklistToResponseNode(const TrackList::pointer& tracklist, Session&)
{ {
Response::Node playlistNode; Response::Node playlistNode;
playlistNode.setAttribute("id", IdToString({Id::Type::Playlist, tracklist.id()})); playlistNode.setAttribute("id", idToString(tracklist->getId()));
playlistNode.setAttribute("name", tracklist->getName()); playlistNode.setAttribute("name", tracklist->getName());
playlistNode.setAttribute("songCount", tracklist->getCount()); playlistNode.setAttribute("songCount", tracklist->getCount());
playlistNode.setAttribute("duration", std::chrono::duration_cast<std::chrono::seconds>(tracklist->getDuration()).count()); playlistNode.setAttribute("duration", std::chrono::duration_cast<std::chrono::seconds>(tracklist->getDuration()).count());
@@ -1334,9 +1312,7 @@ Response
handleGetPlaylistRequest(RequestContext& context) handleGetPlaylistRequest(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; TrackListId trackListId {getMandatoryParameterAs<TrackListId>(context.parameters, "id")};
if (id.type != Id::Type::Playlist)
throw BadParameterGenericError {"id"};
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
@@ -1344,7 +1320,7 @@ handleGetPlaylistRequest(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
TrackList::pointer tracklist {TrackList::getById(context.dbSession, id.value)}; TrackList::pointer tracklist {TrackList::getById(context.dbSession, trackListId)};
if (!tracklist) if (!tracklist)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -1411,7 +1387,7 @@ handleGetSongsByGenreRequest(RequestContext& context)
Response::Node& songsByGenreNode {response.createNode("songsByGenre")}; Response::Node& songsByGenreNode {response.createNode("songsByGenre")};
bool more; bool more;
auto tracks {Track::getByFilter(context.dbSession, {cluster.id()}, {}, Range {offset, size}, more)}; auto tracks {Track::getByFilter(context.dbSession, {cluster->getId()}, {}, Range {offset, size}, more)};
for (const Track::pointer& track : tracks) for (const Track::pointer& track : tracks)
songsByGenreNode.addArrayChild("song", trackToResponseNode(track, context.dbSession, user)); songsByGenreNode.addArrayChild("song", trackToResponseNode(track, context.dbSession, user));
@@ -1504,9 +1480,9 @@ handleSearchRequestCommon(RequestContext& context, bool id3)
struct StarParameters struct StarParameters
{ {
std::vector<Id> artistIds; std::vector<ArtistId> artistIds;
std::vector<Id> releaseIds; std::vector<ReleaseId> releaseIds;
std::vector<Id> trackIds; std::vector<TrackId> trackIds;
}; };
static static
@@ -1515,34 +1491,10 @@ getStarParameters(const Wt::Http::ParameterMap& parameters)
{ {
StarParameters res; StarParameters res;
std::vector<Id> ids {getMultiParametersAs<Id>(parameters, "id")}; // TODO handle parameters for legacy file browsing
res.artistIds = getMultiParametersAs<Id>(parameters, "artistId"); res.trackIds = getMultiParametersAs<TrackId>(parameters, "id");
res.releaseIds = getMultiParametersAs<Id>(parameters, "albumId"); res.artistIds = getMultiParametersAs<ArtistId>(parameters, "artistId");
res.releaseIds = getMultiParametersAs<ReleaseId>(parameters, "albumId");
if (!std::all_of(std::cbegin(res.releaseIds ), std::cend(res.releaseIds ), [](const Id& id) { return id.type == Id::Type::Release; }))
throw BadParameterGenericError {"albumId"};
if (!std::all_of(std::cbegin(res.artistIds ), std::cend(res.artistIds ), [](const Id& id) { return id.type == Id::Type::Artist; }))
throw BadParameterGenericError {"artistId"};
// Redispatch the old "id" parameter in new lists
for (const Id& id : ids)
{
switch (id.type)
{
case Id::Type::Artist:
res.artistIds.emplace_back(id);
break;
case Id::Type::Release:
res.releaseIds.emplace_back(id);
break;
case Id::Type::Track:
res.trackIds.emplace_back(id);
break;
default:
throw BadParameterGenericError {"id"};
}
}
return res; return res;
} }
@@ -1559,27 +1511,27 @@ handleStarRequest(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
for (const Id& id : params.artistIds) for (const ArtistId id : params.artistIds)
{ {
Artist::pointer artist {Artist::getById(context.dbSession, id.value)}; Artist::pointer artist {Artist::getById(context.dbSession, id)};
if (!artist) if (!artist)
continue; continue;
user.modify()->starArtist(artist); user.modify()->starArtist(artist);
} }
for (const Id& id : params.releaseIds) for (const ReleaseId id : params.releaseIds)
{ {
Release::pointer release {Release::getById(context.dbSession, id.value)}; Release::pointer release {Release::getById(context.dbSession, id)};
if (!release) if (!release)
continue; continue;
user.modify()->starRelease(release); user.modify()->starRelease(release);
} }
for (const Id& id : params.trackIds) for (const TrackId id : params.trackIds)
{ {
Track::pointer track {Track::getById(context.dbSession, id.value)}; Track::pointer track {Track::getById(context.dbSession, id)};
if (!track) if (!track)
continue; continue;
@@ -1615,27 +1567,27 @@ handleUnstarRequest(RequestContext& context)
if (!user) if (!user)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
for (const Id& id : params.artistIds) for (const ArtistId id : params.artistIds)
{ {
Artist::pointer artist {Artist::getById(context.dbSession, id.value)}; Artist::pointer artist {Artist::getById(context.dbSession, id)};
if (!artist) if (!artist)
continue; continue;
user.modify()->unstarArtist(artist); user.modify()->unstarArtist(artist);
} }
for (const Id& id : params.releaseIds) for (const ReleaseId id : params.releaseIds)
{ {
Release::pointer release {Release::getById(context.dbSession, id.value)}; Release::pointer release {Release::getById(context.dbSession, id)};
if (!release) if (!release)
continue; continue;
user.modify()->unstarRelease(release); user.modify()->unstarRelease(release);
} }
for (const Id& id : params.trackIds) for (const TrackId id : params.trackIds)
{ {
Track::pointer track {Track::getById(context.dbSession, id.value)}; Track::pointer track {Track::getById(context.dbSession, id)};
if (!track) if (!track)
continue; continue;
@@ -1650,14 +1602,10 @@ static
Response Response
handleScrobble(RequestContext& context) handleScrobble(RequestContext& context)
{ {
const std::vector<Id> ids {getMandatoryMultiParametersAs<Id>(context.parameters, "id")}; const std::vector<TrackId> ids {getMandatoryMultiParametersAs<TrackId>(context.parameters, "id")};
const std::vector<unsigned long> times {getMultiParametersAs<unsigned long>(context.parameters, "time")}; const std::vector<unsigned long> times {getMultiParametersAs<unsigned long>(context.parameters, "time")};
const bool submission{getParameterAs<bool>(context.parameters, "submission").value_or(true)}; const bool submission{getParameterAs<bool>(context.parameters, "submission").value_or(true)};
// only for tracks
if (!std::all_of(std::cbegin(ids), std::cend(ids), [](const Id& id) { return id.type == Id::Type::Track; }))
throw BadParameterGenericError {"id"};
// playing now => no time to be provided // playing now => no time to be provided
if (!submission && !times.empty()) if (!submission && !times.empty())
throw BadParameterGenericError {"time"}; throw BadParameterGenericError {"time"};
@@ -1672,19 +1620,19 @@ handleScrobble(RequestContext& context)
if (!submission) if (!submission)
{ {
Service<Scrobbling::IScrobbling>::get()->listenStarted({context.userId, ids.front().value}); Service<Scrobbling::IScrobbling>::get()->listenStarted({context.userId, ids.front()});
} }
else else
{ {
if (times.empty()) if (times.empty())
{ {
Service<Scrobbling::IScrobbling>::get()->listenFinished({context.userId, ids.front().value}); Service<Scrobbling::IScrobbling>::get()->listenFinished({context.userId, ids.front()});
} }
else else
{ {
for (std::size_t i {}; i < ids.size(); ++i) for (std::size_t i {}; i < ids.size(); ++i)
{ {
const Database::IdType trackId {ids[i].value}; const TrackId trackId {ids[i]};
const unsigned long time {times[i]}; const unsigned long time {times[i]};
Service<Scrobbling::IScrobbling>::get()->addTimedListen({context.userId, trackId, Wt::WDateTime::fromTime_t(static_cast<std::time_t>(time / 1000))}); Service<Scrobbling::IScrobbling>::get()->addTimedListen({context.userId, trackId, Wt::WDateTime::fromTime_t(static_cast<std::time_t>(time / 1000))});
} }
@@ -1701,7 +1649,7 @@ handleUpdateUserRequest(RequestContext& context)
std::string username {getMandatoryParameterAs<std::string>(context.parameters, "username")}; std::string username {getMandatoryParameterAs<std::string>(context.parameters, "username")};
std::optional<std::string> password {getParameterAs<std::string>(context.parameters, "password")}; std::optional<std::string> password {getParameterAs<std::string>(context.parameters, "password")};
Database::IdType userId; UserId userId;
{ {
auto transaction {context.dbSession.createSharedTransaction()}; auto transaction {context.dbSession.createSharedTransaction()};
@@ -1709,7 +1657,7 @@ handleUpdateUserRequest(RequestContext& context)
if (!user) if (!user)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
userId = user.id(); userId = user->getId();
} }
if (password) if (password)
@@ -1742,18 +1690,13 @@ Response
handleUpdatePlaylistRequest(RequestContext& context) handleUpdatePlaylistRequest(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "playlistId")}; TrackListId id {getMandatoryParameterAs<TrackListId>(context.parameters, "playlistId")};
if (id.type != Id::Type::Playlist)
throw BadParameterGenericError {"playlistId"};
// Optional parameters // Optional parameters
auto name {getParameterAs<std::string>(context.parameters, "name")}; auto name {getParameterAs<std::string>(context.parameters, "name")};
auto isPublic {getParameterAs<bool>(context.parameters, "public")}; auto isPublic {getParameterAs<bool>(context.parameters, "public")};
std::vector<Id> trackIdsToAdd {getMultiParametersAs<Id>(context.parameters, "songIdToAdd")}; std::vector<TrackId> trackIdsToAdd {getMultiParametersAs<TrackId>(context.parameters, "songIdToAdd")};
if (!std::all_of(std::cbegin(trackIdsToAdd), std::cend(trackIdsToAdd), [](const Id& id) { return id.type == Id::Type::Track; }))
throw BadParameterGenericError {"songIdToAdd"};
std::vector<std::size_t> trackPositionsToRemove {getMultiParametersAs<std::size_t>(context.parameters, "songIndexToRemove")}; std::vector<std::size_t> trackPositionsToRemove {getMultiParametersAs<std::size_t>(context.parameters, "songIndexToRemove")};
auto transaction {context.dbSession.createUniqueTransaction()}; auto transaction {context.dbSession.createUniqueTransaction()};
@@ -1762,7 +1705,7 @@ handleUpdatePlaylistRequest(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
TrackList::pointer tracklist {TrackList::getById(context.dbSession, id.value)}; TrackList::pointer tracklist {TrackList::getById(context.dbSession, id)};
if (!tracklist if (!tracklist
|| tracklist->getUser() != user || tracklist->getUser() != user
|| tracklist->getType() != TrackList::Type::Playlist) || tracklist->getType() != TrackList::Type::Playlist)
@@ -1789,13 +1732,13 @@ handleUpdatePlaylistRequest(RequestContext& context)
} }
// Add tracks // Add tracks
for (const Id& trackIdToAdd : trackIdsToAdd) for (const TrackId trackIdToAdd : trackIdsToAdd)
{ {
Track::pointer track {Track::getById(context.dbSession, trackIdToAdd.value)}; Track::pointer track {Track::getById(context.dbSession, trackIdToAdd)};
if (!track) if (!track)
continue; continue;
TrackListEntry::create(context.dbSession, track, tracklist ); TrackListEntry::create(context.dbSession, track, tracklist);
} }
return Response::createOkResponse(context); return Response::createOkResponse(context);
@@ -1832,10 +1775,7 @@ Response
handleCreateBookmark(RequestContext& context) handleCreateBookmark(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; TrackId id {getMandatoryParameterAs<TrackId>(context.parameters, "id")};
if (id.type != Id::Type::Track)
throw BadParameterGenericError {"id"};
unsigned long position {getMandatoryParameterAs<unsigned long>(context.parameters, "position")}; unsigned long position {getMandatoryParameterAs<unsigned long>(context.parameters, "position")};
const std::optional<std::string> comment {getParameterAs<std::string>(context.parameters, "comment")}; const std::optional<std::string> comment {getParameterAs<std::string>(context.parameters, "comment")};
@@ -1845,7 +1785,7 @@ handleCreateBookmark(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
const Track::pointer track {Track::getById(context.dbSession, id.value)}; const Track::pointer track {Track::getById(context.dbSession, id)};
if (!track) if (!track)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -1866,9 +1806,7 @@ Response
handleDeleteBookmark(RequestContext& context) handleDeleteBookmark(RequestContext& context)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; TrackId id {getMandatoryParameterAs<TrackId>(context.parameters, "id")};
if (id.type != Id::Type::Track)
throw BadParameterGenericError {"id"};
auto transaction {context.dbSession.createUniqueTransaction()}; auto transaction {context.dbSession.createUniqueTransaction()};
@@ -1876,7 +1814,7 @@ handleDeleteBookmark(RequestContext& context)
if (!user) if (!user)
throw UserNotAuthorizedError {}; throw UserNotAuthorizedError {};
const Track::pointer track {Track::getById(context.dbSession, id.value)}; const Track::pointer track {Track::getById(context.dbSession, id)};
if (!track) if (!track)
throw RequestedDataNotFoundError {}; throw RequestedDataNotFoundError {};
@@ -1901,23 +1839,20 @@ void
handleGetCoverArt(RequestContext& context, const Wt::Http::Request& /*request*/, Wt::Http::Response& response) handleGetCoverArt(RequestContext& context, const Wt::Http::Request& /*request*/, Wt::Http::Response& response)
{ {
// Mandatory params // Mandatory params
Id id {getMandatoryParameterAs<Id>(context.parameters, "id")}; const auto trackId {getParameterAs<TrackId>(context.parameters, "id")};
const auto releaseId {getParameterAs<ReleaseId>(context.parameters, "id")};
if (!trackId && !releaseId)
throw BadParameterGenericError {"id"};
std::size_t size {getParameterAs<std::size_t>(context.parameters, "size").value_or(256)}; std::size_t size {getParameterAs<std::size_t>(context.parameters, "size").value_or(256)};
size = clamp(size, std::size_t {32}, std::size_t {1024}); size = Utils::clamp(size, std::size_t {32}, std::size_t {1024});
std::shared_ptr<CoverArt::IEncodedImage> cover; std::shared_ptr<CoverArt::IEncodedImage> cover;
switch (id.type) if (trackId)
{ cover = Service<CoverArt::IGrabber>::get()->getFromTrack(context.dbSession, *trackId, size);
case Id::Type::Track: else if (releaseId)
cover = Service<CoverArt::IGrabber>::get()->getFromTrack(context.dbSession, id.value, size); cover = Service<CoverArt::IGrabber>::get()->getFromRelease(context.dbSession, *releaseId, size);
break;
case Id::Type::Release:
cover = Service<CoverArt::IGrabber>::get()->getFromRelease(context.dbSession, id.value, size);
break;
default:
throw BadParameterGenericError {"id"};
}
response.out().write(reinterpret_cast<const char*>(cover->getData()), cover->getDataSize()); response.out().write(reinterpret_cast<const char*>(cover->getData()), cover->getDataSize());
response.setMimeType(std::string {cover->getMimeType()}); response.setMimeType(std::string {cover->getMimeType()});
@@ -1928,7 +1863,7 @@ using CheckImplementedFunc = std::function<void()>;
struct RequestEntryPointInfo struct RequestEntryPointInfo
{ {
RequestHandlerFunc func; RequestHandlerFunc func;
EnumSet<Database::UserType> allowedUserTypes {Database::UserType::DEMO, Database::UserType::REGULAR, Database::UserType::ADMIN}; EnumSet<UserType> allowedUserTypes {UserType::DEMO, UserType::REGULAR, UserType::ADMIN};
CheckImplementedFunc checkFunc {}; CheckImplementedFunc checkFunc {};
}; };
@@ -2019,11 +1954,11 @@ static std::unordered_map<std::string, RequestEntryPointInfo> requestEntryPoints
// User management // User management
{"getUser", {handleGetUserRequest}}, {"getUser", {handleGetUserRequest}},
{"getUsers", {handleGetUsersRequest, {Database::UserType::ADMIN}}}, {"getUsers", {handleGetUsersRequest, {UserType::ADMIN}}},
{"createUser", {handleCreateUserRequest, {Database::UserType::ADMIN}, &checkSetPasswordImplemented}}, {"createUser", {handleCreateUserRequest, {UserType::ADMIN}, &checkSetPasswordImplemented}},
{"updateUser", {handleUpdateUserRequest, {Database::UserType::ADMIN}}}, {"updateUser", {handleUpdateUserRequest, {UserType::ADMIN}}},
{"deleteUser", {handleDeleteUserRequest, {Database::UserType::ADMIN}}}, {"deleteUser", {handleDeleteUserRequest, {UserType::ADMIN}}},
{"changePassword", {handleChangePassword, {Database::UserType::REGULAR, Database::UserType::ADMIN}, &checkSetPasswordImplemented}}, {"changePassword", {handleChangePassword, {UserType::REGULAR, UserType::ADMIN}, &checkSetPasswordImplemented}},
// Bookmarks // Bookmarks
{"getBookmarks", {handleGetBookmarks}}, {"getBookmarks", {handleGetBookmarks}},
@@ -2033,8 +1968,8 @@ static std::unordered_map<std::string, RequestEntryPointInfo> requestEntryPoints
{"savePlayQueue", {handleNotImplemented}}, {"savePlayQueue", {handleNotImplemented}},
// Media library scanning // Media library scanning
{"getScanStatus", {Scan::handleGetScanStatus, {Database::UserType::ADMIN}}}, {"getScanStatus", {Scan::handleGetScanStatus, {UserType::ADMIN}}},
{"startScan", {Scan::handleStartScan, {Database::UserType::ADMIN}}}, {"startScan", {Scan::handleStartScan, {UserType::ADMIN}}},
}; };
using MediaRetrievalHandlerFunc = std::function<void(RequestContext&, const Wt::Http::Request&, Wt::Http::Response&)>; using MediaRetrievalHandlerFunc = std::function<void(RequestContext&, const Wt::Http::Request&, Wt::Http::Response&)>;
@@ -2047,7 +1982,7 @@ static std::unordered_map<std::string, MediaRetrievalHandlerFunc> mediaRetrieval
}; };
static static
Database::IdType Database::UserId
authenticateUser(const Wt::Http::Request &request, const ClientInfo& clientInfo, Session& dbSession) authenticateUser(const Wt::Http::Request &request, const ClientInfo& clientInfo, Session& dbSession)
{ {
if (auto *authEnvService {Service<::Auth::IEnvService>::get()}) if (auto *authEnvService {Service<::Auth::IEnvService>::get()})
@@ -2108,7 +2043,7 @@ SubsonicResource::handleRequest(const Wt::Http::Request &request, Wt::Http::Resp
Session& dbSession {_db.getTLSSession()}; Session& dbSession {_db.getTLSSession()};
const Database::IdType userId {authenticateUser(request, clientInfo, dbSession)}; const Database::UserId userId {authenticateUser(request, clientInfo, dbSession)};
RequestContext requestContext {parameters, dbSession, userId, clientInfo.name}; RequestContext requestContext {parameters, dbSession, userId, clientInfo.name};
auto itEntryPoint {requestEntryPoints.find(requestPath)}; auto itEntryPoint {requestEntryPoints.find(requestPath)};
+18 -5
View File
@@ -19,12 +19,25 @@
#pragma once #pragma once
#include <algorithm>
#include <functional> #include <functional>
template<class T, class Compare = std::less<>> namespace Utils
constexpr T clamp(T v, T lo, T hi, Compare comp = {})
{ {
assert(!comp(hi, lo));
return comp(v, lo) ? lo : comp(hi, v) ? hi : v;
}
template<class T, class Compare = std::less<>>
constexpr T clamp(T v, T lo, T hi, Compare comp = {})
{
assert(!comp(hi, lo));
return comp(v, lo) ? lo : comp(hi, v) ? hi : v;
}
template <typename Container, typename T>
void
push_back_if_not_present(Container& container, const T& val)
{
if (std::find(std::cbegin(container), std::cend(container), val) == std::cend(container))
container.push_back(val);
}
}
+5 -5
View File
@@ -47,7 +47,7 @@ static const std::string authCookieName {"LmsAuth"};
static static
void void
createAuthToken(Database::IdType userId, const Wt::WDateTime& expiry) createAuthToken(Database::UserId userId, const Wt::WDateTime& expiry)
{ {
const std::string secret {Service<::Auth::IAuthTokenService>::get()->createAuthToken(LmsApp->getDbSession(), userId, expiry)}; const std::string secret {Service<::Auth::IAuthTokenService>::get()->createAuthToken(LmsApp->getDbSession(), userId, expiry)};
@@ -60,7 +60,7 @@ createAuthToken(Database::IdType userId, const Wt::WDateTime& expiry)
} }
std::optional<Database::IdType> std::optional<Database::UserId>
processAuthToken(const Wt::WEnvironment& env) processAuthToken(const Wt::WEnvironment& env)
{ {
const std::string* authCookie {env.getCookie(authCookieName)}; const std::string* authCookie {env.getCookie(authCookieName)};
@@ -111,7 +111,7 @@ class AuthModel : public Wt::WFormModel
Database::User::pointer user {Database::User::getByLoginName(LmsApp->getDbSession(), valueText(LoginNameField).toUTF8())}; Database::User::pointer user {Database::User::getByLoginName(LmsApp->getDbSession(), valueText(LoginNameField).toUTF8())};
user.modify()->setLastLogin(Wt::WDateTime::currentDateTime()); user.modify()->setLastLogin(Wt::WDateTime::currentDateTime());
_userId = user.id(); _userId = user->getId();
isDemo = user->isDemo(); isDemo = user->isDemo();
} }
@@ -158,11 +158,11 @@ class AuthModel : public Wt::WFormModel
return (validation(field).state() == Wt::ValidationState::Valid); return (validation(field).state() == Wt::ValidationState::Valid);
} }
std::optional<Database::IdType> getUserId() const { return _userId; } std::optional<Database::UserId> getUserId() const { return _userId; }
private: private:
std::optional<Database::IdType> _userId; std::optional<Database::UserId> _userId;
}; };
const AuthModel::Field AuthModel::LoginNameField {"login-name"}; const AuthModel::Field AuthModel::LoginNameField {"login-name"};
+2 -2
View File
@@ -25,7 +25,7 @@
namespace UserInterface { namespace UserInterface {
std::optional<Database::IdType> std::optional<Database::UserId>
processAuthToken(const Wt::WEnvironment& env); processAuthToken(const Wt::WEnvironment& env);
class Auth : public Wt::WTemplateFormView class Auth : public Wt::WTemplateFormView
@@ -33,7 +33,7 @@ class Auth : public Wt::WTemplateFormView
public: public:
Auth(); Auth();
Wt::Signal<Database::IdType /*userId*/> userLoggedIn; Wt::Signal<Database::UserId> userLoggedIn;
}; };
} // namespace UserInterface } // namespace UserInterface
+14 -14
View File
@@ -95,7 +95,7 @@ LmsApplication::getDbSession()
return _db.getTLSSession(); return _db.getTLSSession();
} }
Wt::Dbo::ptr<Database::User> Database::User::pointer
LmsApplication::getUser() LmsApplication::getUser()
{ {
if (!_authenticatedUser) if (!_authenticatedUser)
@@ -104,7 +104,7 @@ LmsApplication::getUser()
return Database::User::getById(getDbSession(), _authenticatedUser->userId); return Database::User::getById(getDbSession(), _authenticatedUser->userId);
} }
Database::IdType Database::UserId
LmsApplication::getUserId() LmsApplication::getUserId()
{ {
return _authenticatedUser->userId; return _authenticatedUser->userId;
@@ -135,7 +135,7 @@ LmsApplication::getUserLoginName()
LmsApplication::LmsApplication(const Wt::WEnvironment& env, LmsApplication::LmsApplication(const Wt::WEnvironment& env,
Database::Db& db, Database::Db& db,
LmsApplicationManager& appManager, LmsApplicationManager& appManager,
std::optional<Database::IdType> userId) std::optional<Database::UserId> userId)
: Wt::WApplication {env} : Wt::WApplication {env}
, _db {db} , _db {db}
, _appManager {appManager} , _appManager {appManager}
@@ -208,7 +208,7 @@ void
LmsApplication::processPasswordAuth() LmsApplication::processPasswordAuth()
{ {
{ {
std::optional<Database::IdType> userId {processAuthToken(environment())}; std::optional<Database::UserId> userId {processAuthToken(environment())};
if (userId) if (userId)
{ {
LMS_LOG(UI, DEBUG) << "User authenticated using Auth token!"; LMS_LOG(UI, DEBUG) << "User authenticated using Auth token!";
@@ -236,7 +236,7 @@ LmsApplication::processPasswordAuth()
else else
{ {
Auth* auth {root()->addNew<Auth>()}; Auth* auth {root()->addNew<Auth>()};
auth->userLoggedIn.connect(this, [this](Database::IdType userId) auth->userLoggedIn.connect(this, [this](Database::UserId userId)
{ {
_authenticatedUser = {userId, true}; _authenticatedUser = {userId, true};
onUserLoggedIn(); onUserLoggedIn();
@@ -272,7 +272,7 @@ LmsApplication::createArtistLink(Database::Artist::pointer artist)
if (const auto mbid {artist->getMBID()}) if (const auto mbid {artist->getMBID()})
return Wt::WLink {Wt::LinkType::InternalPath, "/artist/mbid/" + std::string {mbid->getAsString()}}; return Wt::WLink {Wt::LinkType::InternalPath, "/artist/mbid/" + std::string {mbid->getAsString()}};
else else
return Wt::WLink {Wt::LinkType::InternalPath, "/artist/" + std::to_string(artist.id())}; return Wt::WLink {Wt::LinkType::InternalPath, "/artist/" + artist->getId().toString()};
} }
std::unique_ptr<Wt::WAnchor> std::unique_ptr<Wt::WAnchor>
@@ -296,7 +296,7 @@ LmsApplication::createReleaseLink(Database::Release::pointer release)
if (const auto mbid {release->getMBID()}) if (const auto mbid {release->getMBID()})
return Wt::WLink {Wt::LinkType::InternalPath, "/release/mbid/" + std::string {mbid->getAsString()}}; return Wt::WLink {Wt::LinkType::InternalPath, "/release/mbid/" + std::string {mbid->getAsString()}};
else else
return Wt::WLink {Wt::LinkType::InternalPath, "/release/" + std::to_string(release.id())}; return Wt::WLink {Wt::LinkType::InternalPath, "/release/" + release->getId().toString()};
} }
std::unique_ptr<Wt::WAnchor> std::unique_ptr<Wt::WAnchor>
@@ -320,7 +320,7 @@ LmsApplication::createCluster(Database::Cluster::pointer cluster, bool canDelete
{ {
auto getStyleClass = [](const Database::Cluster::pointer cluster) auto getStyleClass = [](const Database::Cluster::pointer cluster)
{ {
switch (cluster->getType().id() % 6) switch (cluster->getType()->getId().getValue() % 6)
{ {
case 0: return "label-info"; case 0: return "label-info";
case 1: return "label-warning"; case 1: return "label-warning";
@@ -529,7 +529,7 @@ LmsApplication::createHome()
mainStack->addNew<UserView>(); mainStack->addNew<UserView>();
} }
explore->tracksAction.connect([this] (PlayQueueAction action, const std::vector<Database::IdType>& trackIds) explore->tracksAction.connect([this] (PlayQueueAction action, const std::vector<Database::TrackId>& trackIds)
{ {
_playQueue->processTracks(action, trackIds); _playQueue->processTracks(action, trackIds);
}); });
@@ -544,15 +544,15 @@ LmsApplication::createHome()
_playQueue->playPrevious(); _playQueue->playPrevious();
}); });
_mediaPlayer->scrobbleListenNow.connect([this](Database::IdType trackId) _mediaPlayer->scrobbleListenNow.connect([this](Database::TrackId trackId)
{ {
LMS_LOG(UI, DEBUG) << "Received ScrobbleListenNow from player for trackId = " << trackId; LMS_LOG(UI, DEBUG) << "Received ScrobbleListenNow from player for trackId = " << trackId.toString();
const Scrobbling::Listen listen {getUserId(), trackId}; const Scrobbling::Listen listen {getUserId(), trackId};
Service<Scrobbling::IScrobbling>::get()->listenStarted(listen); Service<Scrobbling::IScrobbling>::get()->listenStarted(listen);
}); });
_mediaPlayer->scrobbleListenFinished.connect([this](Database::IdType trackId, unsigned durationMs) _mediaPlayer->scrobbleListenFinished.connect([this](Database::TrackId trackId, unsigned durationMs)
{ {
LMS_LOG(UI, DEBUG) << "Received ScrobbleListenFinished from player for trackId = " << trackId << ", duration = " << (durationMs / 1000) << "s"; LMS_LOG(UI, DEBUG) << "Received ScrobbleListenFinished from player for trackId = " << trackId.toString() << ", duration = " << (durationMs / 1000) << "s";
const std::chrono::milliseconds duration {durationMs}; const std::chrono::milliseconds duration {durationMs};
const Scrobbling::Listen listen {getUserId(), trackId}; const Scrobbling::Listen listen {getUserId(), trackId};
Service<Scrobbling::IScrobbling>::get()->listenFinished(listen, std::chrono::duration_cast<std::chrono::seconds>(duration)); Service<Scrobbling::IScrobbling>::get()->listenFinished(listen, std::chrono::duration_cast<std::chrono::seconds>(duration));
@@ -563,7 +563,7 @@ LmsApplication::createHome()
_playQueue->playNext(); _playQueue->playNext();
}); });
_playQueue->trackSelected.connect([this] (Database::IdType trackId, bool play, float replayGain) _playQueue->trackSelected.connect([this] (Database::TrackId trackId, bool play, float replayGain)
{ {
_mediaPlayer->loadTrack(trackId, play, replayGain); _mediaPlayer->loadTrack(trackId, play, replayGain);
}); });
+9 -9
View File
@@ -51,7 +51,7 @@ class LmsApplication : public Wt::WApplication
{ {
public: public:
LmsApplication(const Wt::WEnvironment& env, Database::Db& db, LmsApplicationManager& appManager, std::optional<Database::IdType> userId = std::nullopt); LmsApplication(const Wt::WEnvironment& env, Database::Db& db, LmsApplicationManager& appManager, std::optional<Database::UserId> userId = std::nullopt);
~LmsApplication(); ~LmsApplication();
static std::unique_ptr<Wt::WApplication> create(const Wt::WEnvironment& env, Database::Db& db, LmsApplicationManager& appManager); static std::unique_ptr<Wt::WApplication> create(const Wt::WEnvironment& env, Database::Db& db, LmsApplicationManager& appManager);
@@ -62,8 +62,8 @@ class LmsApplication : public Wt::WApplication
std::shared_ptr<CoverResource> getCoverResource() { return _coverResource; } std::shared_ptr<CoverResource> getCoverResource() { return _coverResource; }
Database::Session& getDbSession(); // always thread safe Database::Session& getDbSession(); // always thread safe
Wt::Dbo::ptr<Database::User> getUser(); Database::ObjectPtr<Database::User> getUser();
Database::IdType getUserId(); Database::UserId getUserId();
bool isUserAuthStrong() const; // user must be logged in prior this call bool isUserAuthStrong() const; // user must be logged in prior this call
Database::UserType getUserType(); // user must be logged in prior this call Database::UserType getUserType(); // user must be logged in prior this call
std::string getUserLoginName(); // user must be logged in prior this call std::string getUserLoginName(); // user must be logged in prior this call
@@ -84,11 +84,11 @@ class LmsApplication : public Wt::WApplication
}; };
void notifyMsg(MsgType type, const Wt::WString& message, std::chrono::milliseconds duration = std::chrono::milliseconds {4000}); void notifyMsg(MsgType type, const Wt::WString& message, std::chrono::milliseconds duration = std::chrono::milliseconds {4000});
static Wt::WLink createArtistLink(Wt::Dbo::ptr<Database::Artist> artist); static Wt::WLink createArtistLink(Database::ObjectPtr<Database::Artist> artist);
static std::unique_ptr<Wt::WAnchor> createArtistAnchor(Wt::Dbo::ptr<Database::Artist> artist, bool addText = true); static std::unique_ptr<Wt::WAnchor> createArtistAnchor(Database::ObjectPtr<Database::Artist> artist, bool addText = true);
static Wt::WLink createReleaseLink(Wt::Dbo::ptr<Database::Release> release); static Wt::WLink createReleaseLink(Database::ObjectPtr<Database::Release> release);
static std::unique_ptr<Wt::WAnchor> createReleaseAnchor(Wt::Dbo::ptr<Database::Release> release, bool addText = true); static std::unique_ptr<Wt::WAnchor> createReleaseAnchor(Database::ObjectPtr<Database::Release> release, bool addText = true);
static std::unique_ptr<Wt::WText> createCluster(Wt::Dbo::ptr<Database::Cluster> cluster, bool canDelete = false); static std::unique_ptr<Wt::WText> createCluster(Database::ObjectPtr<Database::Cluster> cluster, bool canDelete = false);
Wt::WPopupMenu* createPopupMenu(); Wt::WPopupMenu* createPopupMenu();
MediaPlayer& getMediaPlayer() const { return *_mediaPlayer; } MediaPlayer& getMediaPlayer() const { return *_mediaPlayer; }
@@ -119,7 +119,7 @@ class LmsApplication : public Wt::WApplication
Scanner::Events _scannerEvents; Scanner::Events _scannerEvents;
struct UserAuthInfo struct UserAuthInfo
{ {
Database::IdType userId; Database::UserId userId;
bool strongAuth {}; bool strongAuth {};
}; };
std::optional<UserAuthInfo> _authenticatedUser; std::optional<UserAuthInfo> _authenticatedUser;
+1 -1
View File
@@ -43,6 +43,6 @@ namespace UserInterface
void unregisterApplication(LmsApplication& application); void unregisterApplication(LmsApplication& application);
std::mutex _mutex; std::mutex _mutex;
std::unordered_map<Database::IdType /* user */, std::unordered_set<LmsApplication*>> m_applications; std::unordered_map<Database::UserId, std::unordered_set<LmsApplication*>> m_applications;
}; };
} // UserInterface } // UserInterface
+4 -4
View File
@@ -154,7 +154,7 @@ replayGainPreAmpGainFromString(const std::string& str)
if (!value) if (!value)
return std::nullopt; return std::nullopt;
return clamp(*value, (double)MediaPlayer::Settings::ReplayGain::minPreAmpGain, (double)MediaPlayer::Settings::ReplayGain::maxPreAmpGain); return Utils::clamp(*value, (double)MediaPlayer::Settings::ReplayGain::minPreAmpGain, (double)MediaPlayer::Settings::ReplayGain::maxPreAmpGain);
} }
static MediaPlayer::Settings settingsfromJSString(const std::string& strSettings) static MediaPlayer::Settings settingsfromJSString(const std::string& strSettings)
@@ -233,9 +233,9 @@ MediaPlayer::MediaPlayer()
} }
void void
MediaPlayer::loadTrack(Database::IdType trackId, bool play, float replayGain) MediaPlayer::loadTrack(Database::TrackId trackId, bool play, float replayGain)
{ {
LMS_LOG(UI, DEBUG) << "Playing track ID = " << trackId; LMS_LOG(UI, DEBUG) << "Playing track ID = " << trackId.toString();
std::ostringstream oss; std::ostringstream oss;
{ {
@@ -252,7 +252,7 @@ MediaPlayer::loadTrack(Database::IdType trackId, bool play, float replayGain)
oss oss
<< "var params = {" << "var params = {"
<< " trackId :\"" << trackId << "\"," << " trackId :\"" << trackId.toString() << "\","
<< " nativeResource: \"" << nativeResource << "\"," << " nativeResource: \"" << nativeResource << "\","
<< " transcodeResource: \"" << transcodeResource << "\"," << " transcodeResource: \"" << transcodeResource << "\","
<< " duration: " << std::chrono::duration_cast<std::chrono::seconds>(track->getDuration()).count() << "," << " duration: " << std::chrono::duration_cast<std::chrono::seconds>(track->getDuration()).count() << ","
+9 -9
View File
@@ -94,22 +94,22 @@ class MediaPlayer : public Wt::WTemplate
MediaPlayer& operator=(const MediaPlayer&) = delete; MediaPlayer& operator=(const MediaPlayer&) = delete;
MediaPlayer& operator=(MediaPlayer&&) = delete; MediaPlayer& operator=(MediaPlayer&&) = delete;
std::optional<Database::IdType> getTrackLoaded() const { return _trackIdLoaded; } std::optional<Database::TrackId> getTrackLoaded() const { return _trackIdLoaded; }
void loadTrack(Database::IdType trackId, bool play, float replayGain); void loadTrack(Database::TrackId trackId, bool play, float replayGain);
void stop(); void stop();
std::optional<Settings> getSettings() const { return _settings; } std::optional<Settings> getSettings() const { return _settings; }
void setSettings(const Settings& settings); void setSettings(const Settings& settings);
// Signals // Signals
Wt::JSignal<> playPrevious; Wt::JSignal<> playPrevious;
Wt::JSignal<> playNext; Wt::JSignal<> playNext;
Wt::Signal<Database::IdType> trackLoaded; Wt::Signal<Database::TrackId> trackLoaded;
Wt::Signal<> settingsLoaded; Wt::Signal<> settingsLoaded;
Wt::JSignal<Database::IdType> scrobbleListenNow; Wt::JSignal<Database::TrackId::ValueType> scrobbleListenNow;
Wt::JSignal<Database::IdType, unsigned /* ms */> scrobbleListenFinished; Wt::JSignal<Database::TrackId::ValueType, unsigned /* ms */> scrobbleListenFinished;
Wt::JSignal<> playbackEnded; Wt::JSignal<> playbackEnded;
@@ -117,7 +117,7 @@ class MediaPlayer : public Wt::WTemplate
std::unique_ptr<AudioFileResource> _audioFileResource; std::unique_ptr<AudioFileResource> _audioFileResource;
std::unique_ptr<AudioTranscodeResource> _audioTranscodeResource; std::unique_ptr<AudioTranscodeResource> _audioTranscodeResource;
std::optional<Database::IdType> _trackIdLoaded; std::optional<Database::TrackId> _trackIdLoaded;
std::optional<Settings> _settings; std::optional<Settings> _settings;
Wt::JSignal<std::string> _settingsLoaded; Wt::JSignal<std::string> _settingsLoaded;
+14 -13
View File
@@ -23,6 +23,7 @@
#include <Wt/WText.h> #include <Wt/WText.h>
#include "database/Cluster.hpp" #include "database/Cluster.hpp"
#include "database/Release.hpp"
#include "database/Session.hpp" #include "database/Session.hpp"
#include "database/Track.hpp" #include "database/Track.hpp"
#include "database/TrackList.hpp" #include "database/TrackList.hpp"
@@ -129,7 +130,7 @@ PlayQueue::PlayQueue()
if (LmsApp->getUser()->isDemo()) if (LmsApp->getUser()->isDemo())
{ {
LMS_LOG(UI, DEBUG) << "Removing tracklist id " << _tracklistId; LMS_LOG(UI, DEBUG) << "Removing tracklist id " << _tracklistId.toString();
auto tracklist = Database::TrackList::getById(LmsApp->getDbSession(), _tracklistId); auto tracklist = Database::TrackList::getById(LmsApp->getDbSession(), _tracklistId);
if (tracklist) if (tracklist)
tracklist.remove(); tracklist.remove();
@@ -167,7 +168,7 @@ PlayQueue::PlayQueue()
trackList = Database::TrackList::create(LmsApp->getDbSession(), currentPlayQueueName, Database::TrackList::Type::Internal, false, LmsApp->getUser()); trackList = Database::TrackList::create(LmsApp->getDbSession(), currentPlayQueueName, Database::TrackList::Type::Internal, false, LmsApp->getUser());
} }
_tracklistId = trackList.id(); _tracklistId = trackList->getId();
} }
updateInfo(); updateInfo();
@@ -226,7 +227,7 @@ PlayQueue::loadTrack(std::size_t pos, bool play)
{ {
updateCurrentTrack(false); updateCurrentTrack(false);
Database::IdType trackId {}; Database::TrackId trackId {};
bool addRadioTrack {}; bool addRadioTrack {};
std::optional<float> replayGain {}; std::optional<float> replayGain {};
{ {
@@ -253,7 +254,7 @@ PlayQueue::loadTrack(std::size_t pos, bool play)
_trackPos = pos; _trackPos = pos;
auto track = tracklist->getEntry(*_trackPos)->getTrack(); auto track = tracklist->getEntry(*_trackPos)->getTrack();
trackId = track.id(); trackId = track->getId();
replayGain = getReplayGain(pos, track); replayGain = getReplayGain(pos, track);
@@ -313,7 +314,7 @@ PlayQueue::updateCurrentTrack(bool selected)
} }
std::size_t std::size_t
PlayQueue::enqueueTracks(const std::vector<Database::IdType>& trackIds) PlayQueue::enqueueTracks(const std::vector<Database::TrackId>& trackIds)
{ {
std::size_t nbTracksQueued {}; std::size_t nbTracksQueued {};
@@ -323,7 +324,7 @@ PlayQueue::enqueueTracks(const std::vector<Database::IdType>& trackIds)
auto tracklist {getTrackList()}; auto tracklist {getTrackList()};
std::size_t nbTracksToEnqueue {tracklist->getCount() + trackIds.size() > _nbMaxEntries ? _nbMaxEntries - tracklist->getCount() : trackIds.size()}; std::size_t nbTracksToEnqueue {tracklist->getCount() + trackIds.size() > _nbMaxEntries ? _nbMaxEntries - tracklist->getCount() : trackIds.size()};
for (Database::IdType trackId : trackIds) for (const Database::TrackId trackId : trackIds)
{ {
Database::Track::pointer track {Database::Track::getById(LmsApp->getDbSession(), trackId)}; Database::Track::pointer track {Database::Track::getById(LmsApp->getDbSession(), trackId)};
if (!track) if (!track)
@@ -344,7 +345,7 @@ PlayQueue::enqueueTracks(const std::vector<Database::IdType>& trackIds)
} }
void void
PlayQueue::processTracks(PlayQueueAction action, const std::vector<Database::IdType>& trackIds) PlayQueue::processTracks(PlayQueueAction action, const std::vector<Database::TrackId>& trackIds)
{ {
std::size_t nbAddedTracks {}; std::size_t nbAddedTracks {};
@@ -367,7 +368,7 @@ PlayQueue::processTracks(PlayQueueAction action, const std::vector<Database::IdT
{ {
clearTracks(); clearTracks();
{ {
std::vector<Database::IdType> shuffledTrackIds {trackIds}; std::vector<Database::TrackId> shuffledTrackIds {trackIds};
Random::shuffleContainer(shuffledTrackIds); Random::shuffleContainer(shuffledTrackIds);
nbAddedTracks = enqueueTracks(shuffledTrackIds); nbAddedTracks = enqueueTracks(shuffledTrackIds);
} }
@@ -402,9 +403,9 @@ PlayQueue::addSome()
void void
PlayQueue::addEntry(const Database::TrackListEntry::pointer& tracklistEntry) PlayQueue::addEntry(const Database::TrackListEntry::pointer& tracklistEntry)
{ {
const auto tracklistEntryId {tracklistEntry.id()}; const Database::TrackListEntryId tracklistEntryId {tracklistEntry->getId()};
const auto track {tracklistEntry->getTrack()}; const auto track {tracklistEntry->getTrack()};
const Database::IdType trackId {track->id()}; const Database::TrackId trackId {track->getId()};
Wt::WTemplate* entry = _entriesContainer->addNew<Wt::WTemplate>(Wt::WString::tr("Lms.PlayQueue.template.entry")); Wt::WTemplate* entry = _entriesContainer->addNew<Wt::WTemplate>(Wt::WString::tr("Lms.PlayQueue.template.entry"));
@@ -434,7 +435,7 @@ PlayQueue::addEntry(const Database::TrackListEntry::pointer& tracklistEntry)
{ {
Wt::WAnchor* anchor = entry->bindWidget("cover", LmsApplication::createReleaseAnchor(release, false)); Wt::WAnchor* anchor = entry->bindWidget("cover", LmsApplication::createReleaseAnchor(release, false));
auto cover = std::make_unique<Wt::WImage>(); auto cover = std::make_unique<Wt::WImage>();
cover->setImageLink(LmsApp->getCoverResource()->getReleaseUrl(release.id(), CoverResource::Size::Large)); cover->setImageLink(LmsApp->getCoverResource()->getReleaseUrl(release->getId(), CoverResource::Size::Large));
cover->setStyleClass("Lms-cover"); cover->setStyleClass("Lms-cover");
cover->setAttributeValue("onload", LmsApp->javaScriptClass() + ".onLoadCover(this)"); cover->setAttributeValue("onload", LmsApp->javaScriptClass() + ".onLoadCover(this)");
anchor->setImage(std::move(cover)); anchor->setImage(std::move(cover));
@@ -443,7 +444,7 @@ PlayQueue::addEntry(const Database::TrackListEntry::pointer& tracklistEntry)
else else
{ {
auto cover = entry->bindNew<Wt::WImage>("cover"); auto cover = entry->bindNew<Wt::WImage>("cover");
cover->setImageLink(LmsApp->getCoverResource()->getTrackUrl(track.id(), CoverResource::Size::Large)); cover->setImageLink(LmsApp->getCoverResource()->getTrackUrl(track->getId(), CoverResource::Size::Large));
cover->setStyleClass("Lms-cover"); cover->setStyleClass("Lms-cover");
cover->setAttributeValue("onload", LmsApp->javaScriptClass() + ".onLoadCover(this)"); cover->setAttributeValue("onload", LmsApp->javaScriptClass() + ".onLoadCover(this)");
} }
@@ -520,7 +521,7 @@ PlayQueue::enqueueRadioTracks()
{ {
const auto similarTrackIds {Service<Recommendation::IEngine>::get()->getSimilarTracksFromTrackList(LmsApp->getDbSession(), _tracklistId, 3)}; const auto similarTrackIds {Service<Recommendation::IEngine>::get()->getSimilarTracksFromTrackList(LmsApp->getDbSession(), _tracklistId, 3)};
std::vector<Database::IdType> trackToAddIds(std::cbegin(similarTrackIds), std::cend(similarTrackIds)); std::vector<Database::TrackId> trackToAddIds(std::cbegin(similarTrackIds), std::cend(similarTrackIds));
Random::shuffleContainer(trackToAddIds); Random::shuffleContainer(trackToAddIds);
enqueueTracks(trackToAddIds); enqueueTracks(trackToAddIds);
} }
+7 -7
View File
@@ -49,7 +49,7 @@ class PlayQueue : public Wt::WTemplate
public: public:
PlayQueue(); PlayQueue();
void processTracks(PlayQueueAction action, const std::vector<Database::IdType>& trackIds); void processTracks(PlayQueueAction action, const std::vector<Database::TrackId>& trackIds);
// play the next track in the queue // play the next track in the queue
void playNext(); void playNext();
@@ -58,19 +58,19 @@ class PlayQueue : public Wt::WTemplate
void playPrevious(); void playPrevious();
// Signal emitted when a track is to be load(and optionally played) // Signal emitted when a track is to be load(and optionally played)
Wt::Signal<Database::IdType /*trackId*/, bool /*play*/, float /* replayGain */> trackSelected; Wt::Signal<Database::TrackId, bool /*play*/, float /* replayGain */> trackSelected;
// Signal emitted when track is unselected (has to be stopped) // Signal emitted when track is unselected (has to be stopped)
Wt::Signal<> trackUnselected; Wt::Signal<> trackUnselected;
private: private:
Wt::Dbo::ptr<Database::TrackList> getTrackList() const; Database::ObjectPtr<Database::TrackList> getTrackList() const;
bool isFull() const; bool isFull() const;
void clearTracks(); void clearTracks();
std::size_t enqueueTracks(const std::vector<Database::IdType>& trackIds); std::size_t enqueueTracks(const std::vector<Database::TrackId>& trackIds);
void addSome(); void addSome();
void addEntry(const Wt::Dbo::ptr<Database::TrackListEntry>& entry); void addEntry(const Database::ObjectPtr<Database::TrackListEntry>& entry);
void enqueueRadioTracks(); void enqueueRadioTracks();
void updateInfo(); void updateInfo();
void updateCurrentTrack(bool selected); void updateCurrentTrack(bool selected);
@@ -82,7 +82,7 @@ class PlayQueue : public Wt::WTemplate
void addRadioTrackFromSimilarity(std::shared_ptr<Similarity::Finder> similarityFinder); void addRadioTrackFromSimilarity(std::shared_ptr<Similarity::Finder> similarityFinder);
void addRadioTrackFromClusters(); void addRadioTrackFromClusters();
std::optional<float> getReplayGain(std::size_t pos, const Wt::Dbo::ptr<Database::Track>& track) const; std::optional<float> getReplayGain(std::size_t pos, const Database::ObjectPtr<Database::Track>& track) const;
static inline constexpr std::size_t _nbMaxEntries {1000}; static inline constexpr std::size_t _nbMaxEntries {1000};
static inline constexpr std::size_t _batchSize {12}; static inline constexpr std::size_t _batchSize {12};
@@ -90,7 +90,7 @@ class PlayQueue : public Wt::WTemplate
bool _repeatAll {}; bool _repeatAll {};
bool _radioMode {}; bool _radioMode {};
bool _mediaPlayerSettingsLoaded {}; bool _mediaPlayerSettingsLoaded {};
Database::IdType _tracklistId {}; Database::TrackListId _tracklistId {};
InfiniteScrollingContainer* _entriesContainer {}; InfiniteScrollingContainer* _entriesContainer {};
Wt::WText* _nbTracks {}; Wt::WText* _nbTracks {};
Wt::WText* _repeatBtn {}; Wt::WText* _repeatBtn {};
+3 -2
View File
@@ -34,7 +34,8 @@ namespace UserInterface
PlayShuffled, PlayShuffled,
}; };
using PlayQueueActionSignal = Wt::Signal<PlayQueueAction, const std::vector<Database::IdType>&>; using PlayQueueActionArtistSignal = Wt::Signal<PlayQueueAction, const std::vector<Database::ArtistId>&>;
using PlayQueueActionReleaseSignal = Wt::Signal<PlayQueueAction, const std::vector<Database::ReleaseId>&>;
using PlayQueueActionTrackSignal = Wt::Signal<PlayQueueAction, const std::vector<Database::TrackId>&>;
} }
+1 -1
View File
@@ -197,7 +197,7 @@ class SettingsModel : public Wt::WFormModel
if (_authPasswordService && !valueText(PasswordField).empty()) if (_authPasswordService && !valueText(PasswordField).empty())
{ {
_authPasswordService->setPassword(LmsApp->getDbSession(), user.id(), valueText(PasswordField).toUTF8()); _authPasswordService->setPassword(LmsApp->getDbSession(), user->getId(), valueText(PasswordField).toUTF8());
} }
} }
+1 -1
View File
@@ -70,7 +70,7 @@ class InitWizardModel : public Wt::WFormModel
Database::User::pointer user {Database::User::create(LmsApp->getDbSession(), valueText(AdminLoginField).toUTF8())}; Database::User::pointer user {Database::User::create(LmsApp->getDbSession(), valueText(AdminLoginField).toUTF8())};
user.modify()->setType(Database::UserType::ADMIN); user.modify()->setType(Database::UserType::ADMIN);
Service<::Auth::IPasswordService>::get()->setPassword(LmsApp->getDbSession(), user.id(), valueText(PasswordField).toUTF8()); Service<::Auth::IPasswordService>::get()->setPassword(LmsApp->getDbSession(), user->getId(), valueText(PasswordField).toUTF8());
} }
bool validateField(Field field) bool validateField(Field field)
+5 -5
View File
@@ -53,7 +53,7 @@ class UserModel : public Wt::WFormModel
static inline const Field PasswordField {"password"}; static inline const Field PasswordField {"password"};
static inline const Field DemoField {"demo"}; static inline const Field DemoField {"demo"};
UserModel(std::optional<Database::IdType> userId, ::Auth::IPasswordService* authPasswordService) UserModel(std::optional<Database::UserId> userId, ::Auth::IPasswordService* authPasswordService)
: _userId {userId} : _userId {userId}
, _authPasswordService {authPasswordService} , _authPasswordService {authPasswordService}
{ {
@@ -87,7 +87,7 @@ class UserModel : public Wt::WFormModel
throw UserNotFoundException {}; throw UserNotFoundException {};
if (_authPasswordService && !valueText(PasswordField).empty()) if (_authPasswordService && !valueText(PasswordField).empty())
_authPasswordService->setPassword(LmsApp->getDbSession(), user.id(), valueText(PasswordField).toUTF8()); _authPasswordService->setPassword(LmsApp->getDbSession(), user->getId(), valueText(PasswordField).toUTF8());
} }
else else
{ {
@@ -103,7 +103,7 @@ class UserModel : public Wt::WFormModel
user.modify()->setType(Database::UserType::DEMO); user.modify()->setType(Database::UserType::DEMO);
if (_authPasswordService) if (_authPasswordService)
_authPasswordService->setPassword(LmsApp->getDbSession(), user.id(), valueText(PasswordField).toUTF8()); _authPasswordService->setPassword(LmsApp->getDbSession(), user->getId(), valueText(PasswordField).toUTF8());
} }
} }
@@ -176,7 +176,7 @@ class UserModel : public Wt::WFormModel
return false; return false;
} }
std::optional<Database::IdType> _userId; std::optional<Database::UserId> _userId;
::Auth::IPasswordService* _authPasswordService {}; ::Auth::IPasswordService* _authPasswordService {};
}; };
@@ -196,7 +196,7 @@ UserView::refreshView()
if (!wApp->internalPathMatches("/admin/user")) if (!wApp->internalPathMatches("/admin/user"))
return; return;
auto userId = StringUtils::readAs<Database::IdType>(wApp->internalPathNextPart("/admin/user/")); const std::optional<Database::UserId> userId {StringUtils::readAs<Database::UserId::ValueType>(wApp->internalPathNextPart("/admin/user/"))};
clear(); clear();
+2 -2
View File
@@ -72,7 +72,7 @@ UsersView::refreshView()
auto users = Database::User::getAll(LmsApp->getDbSession()); auto users = Database::User::getAll(LmsApp->getDbSession());
for (const auto& user : users) for (const auto& user : users)
{ {
const Database::IdType userId {user.id()}; const Database::UserId userId {user->getId()};
Wt::WTemplate* entry {_container->addNew<Wt::WTemplate>(Wt::WString::tr("Lms.Admin.Users.template.entry"))}; Wt::WTemplate* entry {_container->addNew<Wt::WTemplate>(Wt::WString::tr("Lms.Admin.Users.template.entry"))};
@@ -93,7 +93,7 @@ UsersView::refreshView()
Wt::WPushButton* editBtn = entry->bindNew<Wt::WPushButton>("edit-btn", Wt::WString::tr("Lms.Admin.Users.edit")); Wt::WPushButton* editBtn = entry->bindNew<Wt::WPushButton>("edit-btn", Wt::WString::tr("Lms.Admin.Users.edit"));
editBtn->clicked().connect([=]() editBtn->clicked().connect([=]()
{ {
LmsApp->setInternalPath("/admin/user/" + std::to_string(userId), true); LmsApp->setInternalPath("/admin/user/" + userId.toString(), true);
}); });
Wt::WPushButton* delBtn = entry->bindNew<Wt::WPushButton>("del-btn", Wt::WString::tr("Lms.Admin.Users.del")); Wt::WPushButton* delBtn = entry->bindNew<Wt::WPushButton>("del-btn", Wt::WString::tr("Lms.Admin.Users.del"));
+1 -1
View File
@@ -31,7 +31,7 @@ namespace UserInterface
{ {
using namespace Database; using namespace Database;
std::vector<Wt::Dbo::ptr<Database::Artist>> std::vector<Database::ObjectPtr<Database::Artist>>
ArtistCollector::get(std::optional<Database::Range> range, bool& moreResults) ArtistCollector::get(std::optional<Database::Range> range, bool& moreResults)
{ {
range = getActualRange(range); range = getActualRange(range);
+3 -3
View File
@@ -38,13 +38,13 @@ namespace UserInterface
public: public:
using DatabaseCollectorBase::DatabaseCollectorBase; using DatabaseCollectorBase::DatabaseCollectorBase;
std::vector<Wt::Dbo::ptr<Database::Artist>> get(std::optional<Database::Range> range, bool& moreResults); std::vector<Database::ObjectPtr<Database::Artist>> get(std::optional<Database::Range> range, bool& moreResults);
void reset() { _randomArtists.clear(); } void reset() { _randomArtists.clear(); }
void setArtistLinkType(std::optional<Database::TrackArtistLinkType> linkType) { _linkType = linkType; } void setArtistLinkType(std::optional<Database::TrackArtistLinkType> linkType) { _linkType = linkType; }
private: private:
std::vector<Wt::Dbo::ptr<Database::Artist>> getRandomArtists(std::optional<Range> range, bool& moreResults); std::vector<Database::ObjectPtr<Database::Artist>> getRandomArtists(std::optional<Range> range, bool& moreResults);
std::vector<Database::IdType> _randomArtists; std::vector<Database::ArtistId> _randomArtists;
std::optional<Database::TrackArtistLinkType> _linkType; std::optional<Database::TrackArtistLinkType> _linkType;
}; };
} // ns UserInterface } // ns UserInterface
+2 -2
View File
@@ -26,7 +26,7 @@
namespace UserInterface::ArtistListHelpers namespace UserInterface::ArtistListHelpers
{ {
std::unique_ptr<Wt::WTemplate> std::unique_ptr<Wt::WTemplate>
createEntry(const Wt::Dbo::ptr<Database::Artist>& artist) createEntry(const Database::ObjectPtr<Database::Artist>& artist)
{ {
auto res {std::make_unique<Wt::WTemplate>(Wt::WString::tr("Lms.Explore.Artists.template.entry"))}; auto res {std::make_unique<Wt::WTemplate>(Wt::WString::tr("Lms.Explore.Artists.template.entry"))};
res->bindWidget("name", LmsApplication::createArtistAnchor(artist)); res->bindWidget("name", LmsApplication::createArtistAnchor(artist));
@@ -35,7 +35,7 @@ namespace UserInterface::ArtistListHelpers
} }
std::unique_ptr<Wt::WTemplate> std::unique_ptr<Wt::WTemplate>
createEntrySmall(const Wt::Dbo::ptr<Database::Artist>& artist) createEntrySmall(const Database::ObjectPtr<Database::Artist>& artist)
{ {
auto res {std::make_unique<Wt::WTemplate>(Wt::WString::tr("Lms.Explore.Artists.template.entry-small"))}; auto res {std::make_unique<Wt::WTemplate>(Wt::WString::tr("Lms.Explore.Artists.template.entry-small"))};
res->bindWidget("name", LmsApplication::createArtistAnchor(artist)); res->bindWidget("name", LmsApplication::createArtistAnchor(artist));
+2 -2
View File
@@ -32,7 +32,7 @@ namespace Database
namespace UserInterface::ArtistListHelpers namespace UserInterface::ArtistListHelpers
{ {
std::unique_ptr<Wt::WTemplate> createEntry(const Wt::Dbo::ptr<Database::Artist>& artist); std::unique_ptr<Wt::WTemplate> createEntry(const Database::ObjectPtr<Database::Artist>& artist);
std::unique_ptr<Wt::WTemplate> createEntrySmall(const Wt::Dbo::ptr<Database::Artist>& artist); std::unique_ptr<Wt::WTemplate> createEntrySmall(const Database::ObjectPtr<Database::Artist>& artist);
} }
+9 -8
View File
@@ -26,6 +26,7 @@
#include <Wt/WText.h> #include <Wt/WText.h>
#include "database/Artist.hpp" #include "database/Artist.hpp"
#include "database/Cluster.hpp"
#include "database/Release.hpp" #include "database/Release.hpp"
#include "database/ScanSettings.hpp" #include "database/ScanSettings.hpp"
#include "database/Session.hpp" #include "database/Session.hpp"
@@ -68,7 +69,7 @@ Artist::Artist(Filters* filters)
} }
static static
std::optional<IdType> std::optional<ArtistId>
extractArtistIdFromInternalPath() extractArtistIdFromInternalPath()
{ {
if (wApp->internalPathMatches("/artist/mbid/")) if (wApp->internalPathMatches("/artist/mbid/"))
@@ -78,13 +79,13 @@ extractArtistIdFromInternalPath()
{ {
auto transaction {LmsApp->getDbSession().createSharedTransaction()}; auto transaction {LmsApp->getDbSession().createSharedTransaction()};
if (const Database::Artist::pointer artist {Database::Artist::getByMBID(LmsApp->getDbSession(), *mbid)}) if (const Database::Artist::pointer artist {Database::Artist::getByMBID(LmsApp->getDbSession(), *mbid)})
return artist.id(); return artist->getId();
} }
return std::nullopt; return std::nullopt;
} }
return StringUtils::readAs<Database::IdType>(wApp->internalPathNextPart("/artist/")); return StringUtils::readAs<Database::ArtistId::ValueType>(wApp->internalPathNextPart("/artist/"));
} }
void void
@@ -129,7 +130,7 @@ Artist::refreshView()
{ {
for (auto cluster : clusters) for (auto cluster : clusters)
{ {
auto clusterId = cluster.id(); auto clusterId = cluster->getId();
auto entry = clusterContainers->addWidget(LmsApp->createCluster(cluster)); auto entry = clusterContainers->addWidget(LmsApp->createCluster(cluster));
entry->clicked().connect([=] entry->clicked().connect([=]
{ {
@@ -197,7 +198,7 @@ Artist::refreshView()
} }
void void
Artist::refreshReleases(const Wt::Dbo::ptr<Database::Artist>& artist) Artist::refreshReleases(const Database::ObjectPtr<Database::Artist>& artist)
{ {
const auto releases {artist->getReleases(_filters->getClusterIds())}; const auto releases {artist->getReleases(_filters->getClusterIds())};
if (releases.empty()) if (releases.empty())
@@ -213,7 +214,7 @@ Artist::refreshReleases(const Wt::Dbo::ptr<Database::Artist>& artist)
} }
void void
Artist::refreshNonReleaseTracks(const Wt::Dbo::ptr<Database::Artist>& artist) Artist::refreshNonReleaseTracks(const Database::ObjectPtr<Database::Artist>& artist)
{ {
if (!artist->hasNonReleaseTracks()) if (!artist->hasNonReleaseTracks())
return; return;
@@ -229,7 +230,7 @@ Artist::refreshNonReleaseTracks(const Wt::Dbo::ptr<Database::Artist>& artist)
} }
void void
Artist::refreshSimilarArtists(const std::unordered_set<Database::IdType>& similarArtistsId) Artist::refreshSimilarArtists(const std::vector<Database::ArtistId>& similarArtistsId)
{ {
if (similarArtistsId.empty()) if (similarArtistsId.empty())
return; return;
@@ -237,7 +238,7 @@ Artist::refreshSimilarArtists(const std::unordered_set<Database::IdType>& simila
setCondition("if-has-similar-artists", true); setCondition("if-has-similar-artists", true);
Wt::WContainerWidget* similarArtistsContainer {bindNew<Wt::WContainerWidget>("similar-artists")}; Wt::WContainerWidget* similarArtistsContainer {bindNew<Wt::WContainerWidget>("similar-artists")};
for (Database::IdType artistId : similarArtistsId) for (const Database::ArtistId artistId : similarArtistsId)
{ {
const Database::Artist::pointer similarArtist{Database::Artist::getById(LmsApp->getDbSession(), artistId)}; const Database::Artist::pointer similarArtist{Database::Artist::getById(LmsApp->getDbSession(), artistId)};
if (!similarArtist) if (!similarArtist)
+7 -7
View File
@@ -45,15 +45,15 @@ namespace UserInterface
public: public:
Artist(Filters* filters); Artist(Filters* filters);
PlayQueueActionSignal artistsAction; PlayQueueActionArtistSignal artistsAction;
PlayQueueActionSignal tracksAction; PlayQueueActionTrackSignal tracksAction;
private: private:
void refreshView(); void refreshView();
void refreshReleases(const Wt::Dbo::ptr<Database::Artist>& artist); void refreshReleases(const Database::ObjectPtr<Database::Artist>& artist);
void refreshNonReleaseTracks(const Wt::Dbo::ptr<Database::Artist>& artist); void refreshNonReleaseTracks(const Database::ObjectPtr<Database::Artist>& artist);
void refreshSimilarArtists(const std::unordered_set<Database::IdType>& similarArtistsId); void refreshSimilarArtists(const std::vector<Database::ArtistId>& similarArtistsId);
void refreshLinks(const Wt::Dbo::ptr<Database::Artist>& artist); void refreshLinks(const Database::ObjectPtr<Database::Artist>& artist);
void addSomeNonReleaseTracks(); void addSomeNonReleaseTracks();
static constexpr std::size_t _tracksBatchSize {6}; static constexpr std::size_t _tracksBatchSize {6};
@@ -61,7 +61,7 @@ namespace UserInterface
Filters* _filters {}; Filters* _filters {};
InfiniteScrollingContainer* _trackContainer {}; InfiniteScrollingContainer* _trackContainer {};
Database::IdType _artistId {}; Database::ArtistId _artistId {};
}; };
} // namespace UserInterface } // namespace UserInterface
+15 -15
View File
@@ -129,16 +129,16 @@ Explore::search(const Wt::WString& searchText)
} }
static static
std::vector<Database::IdType> std::vector<Database::TrackId>
getArtistsTracks(Database::Session& session, const std::vector<Database::IdType>& artistsId, const std::set<Database::IdType>&) getArtistsTracks(Database::Session& session, const std::vector<Database::ArtistId>& artistsId, const std::vector<Database::ClusterId>&)
{ {
std::vector<Database::IdType> res; std::vector<Database::TrackId> res;
auto transaction {LmsApp->getDbSession().createSharedTransaction()}; auto transaction {LmsApp->getDbSession().createSharedTransaction()};
for (Database::IdType artistId : artistsId) for (const Database::ArtistId artistId : artistsId)
{ {
Database::Artist::pointer artist {Database::Artist::getById(session, artistId)}; const Database::Artist::pointer artist {Database::Artist::getById(session, artistId)};
if (!artist) if (!artist)
continue; continue;
@@ -146,49 +146,49 @@ getArtistsTracks(Database::Session& session, const std::vector<Database::IdType>
const std::vector<Database::Track::pointer> tracks {artist->getTracks()}; const std::vector<Database::Track::pointer> tracks {artist->getTracks()};
res.reserve(res.size() + tracks.size()); res.reserve(res.size() + tracks.size());
std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const Database::Track::pointer& track) { return track.id(); }); std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const Database::Track::pointer& track) { return track->getId(); });
} }
return res; return res;
} }
static static
std::vector<Database::IdType> std::vector<Database::TrackId>
getReleasesTracks(Database::Session& session, const std::vector<Database::IdType>& releasesId, const std::set<Database::IdType>& clusters) getReleasesTracks(Database::Session& session, const std::vector<Database::ReleaseId>& releasesId, const std::vector<Database::ClusterId>& clusters)
{ {
std::vector<Database::IdType> res; std::vector<Database::TrackId> res;
auto transaction {LmsApp->getDbSession().createSharedTransaction()}; auto transaction {LmsApp->getDbSession().createSharedTransaction()};
for (Database::IdType releaseId : releasesId) for (const Database::ReleaseId releaseId : releasesId)
{ {
Database::Release::pointer release {Database::Release::getById(session, releaseId)}; const Database::Release::pointer release {Database::Release::getById(session, releaseId)};
if (!release) if (!release)
continue; continue;
const std::vector<Database::Track::pointer> tracks {release->getTracks(clusters)}; const std::vector<Database::Track::pointer> tracks {release->getTracks(clusters)};
res.reserve(res.size() + tracks.size()); res.reserve(res.size() + tracks.size());
std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const Database::Track::pointer& track) { return track.id(); }); std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const Database::Track::pointer& track) { return track->getId(); });
} }
return res; return res;
} }
void void
Explore::handleArtistsAction(PlayQueueAction action, const std::vector<Database::IdType>& artistsId) Explore::handleArtistsAction(PlayQueueAction action, const std::vector<Database::ArtistId>& artistsId)
{ {
tracksAction.emit(action, getArtistsTracks(LmsApp->getDbSession(), artistsId, _filters->getClusterIds())); tracksAction.emit(action, getArtistsTracks(LmsApp->getDbSession(), artistsId, _filters->getClusterIds()));
} }
void void
Explore::handleReleasesAction(PlayQueueAction action, const std::vector<Database::IdType>& releasesId) Explore::handleReleasesAction(PlayQueueAction action, const std::vector<Database::ReleaseId>& releasesId)
{ {
tracksAction.emit(action, getReleasesTracks(LmsApp->getDbSession(), releasesId, _filters->getClusterIds())); tracksAction.emit(action, getReleasesTracks(LmsApp->getDbSession(), releasesId, _filters->getClusterIds()));
} }
void void
Explore::handleTracksAction(PlayQueueAction action, const std::vector<Database::IdType>& tracksId) Explore::handleTracksAction(PlayQueueAction action, const std::vector<Database::TrackId>& tracksId)
{ {
tracksAction.emit(action, tracksId); tracksAction.emit(action, tracksId);
} }

Some files were not shown because too many files have changed in this diff Show More