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
+1 -2
View File
@@ -29,7 +29,6 @@
namespace Database
{
class Db;
class Session;
class TrackList;
class User;
@@ -48,7 +47,7 @@ namespace Scrobbling
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);
+33 -33
View File
@@ -64,7 +64,7 @@ namespace Scrobbling
}
std::optional<Database::Scrobbler>
Scrobbling::getUserScrobbler(Database::IdType userId)
Scrobbling::getUserScrobbler(Database::UserId userId)
{
std::optional<Database::Scrobbler> scrobbler;
@@ -76,49 +76,49 @@ namespace Scrobbling
return scrobbler;
}
std::vector<Wt::Dbo::ptr<Database::Artist>>
std::vector<Database::ObjectPtr<Database::Artist>>
Scrobbling::getRecentArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range,
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)
res = history->getArtistsReverse(clusterIds, linkType, range, moreResults);
return res;
}
std::vector<Wt::Dbo::ptr<Database::Release>>
std::vector<Database::ObjectPtr<Database::Release>>
Scrobbling::getRecentReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
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)
res = history->getReleasesReverse(clusterIds, range, moreResults);
return res;
}
std::vector<Wt::Dbo::ptr<Database::Track>>
std::vector<Database::ObjectPtr<Database::Track>>
Scrobbling::getRecentTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
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)
res = history->getTracksReverse(clusterIds, range, moreResults);
@@ -127,57 +127,57 @@ namespace Scrobbling
// Top
std::vector<Wt::Dbo::ptr<Database::Artist>>
std::vector<Database::ObjectPtr<Database::Artist>>
Scrobbling::getTopArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range,
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)
res = history->getTopArtists(clusterIds, linkType, range, moreResults);
return res;
}
std::vector<Wt::Dbo::ptr<Database::Release>>
std::vector<Database::ObjectPtr<Database::Release>>
Scrobbling::getTopReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
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)
res = history->getTopReleases(clusterIds, range, moreResults);
return res;
}
std::vector<Wt::Dbo::ptr<Database::Track>>
std::vector<Database::ObjectPtr<Database::Track>>
Scrobbling::getTopTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
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)
res = history->getTopTracks(clusterIds, range, moreResults);
return res;
}
Wt::Dbo::ptr<Database::TrackList>
Scrobbling::getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user)
Database::ObjectPtr<Database::TrackList>
Scrobbling::getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> 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 addTimedListen(const TimedListen& listen) override;
std::vector<Wt::Dbo::ptr<Database::Artist>> getRecentArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
std::vector<Database::ObjectPtr<Database::Artist>> getRecentArtists(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range,
bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Release>> getRecentReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
std::vector<Database::ObjectPtr<Database::Release>> getRecentReleases(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Track>> getRecentTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
std::vector<Database::ObjectPtr<Database::Track>> getRecentTracks(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Artist>> getTopArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
std::vector<Database::ObjectPtr<Database::Artist>> getTopArtists(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range,
bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Release>> getTopReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
std::vector<Database::ObjectPtr<Database::Release>> getTopReleases(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) override;
std::vector<Wt::Dbo::ptr<Database::Track>> getTopTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
std::vector<Database::ObjectPtr<Database::Track>> getTopTracks(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
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;
std::unordered_map<Database::Scrobbler, std::unique_ptr<IScrobbler>> _scrobblers;
@@ -61,7 +61,7 @@ namespace Scrobbling
if (!user)
return;
Wt::Dbo::ptr<Database::TrackList> tracklist {getListensTrackList(session, user)};
Database::TrackList::pointer tracklist {getListensTrackList(session, user)};
if (!tracklist)
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);
}
Wt::Dbo::ptr<Database::TrackList>
InternalScrobbler::getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user)
Database::TrackList::pointer
InternalScrobbler::getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user)
{
return Database::TrackList::get(session, historyTracklistName, Database::TrackList::Type::Internal, user);
}
@@ -21,6 +21,11 @@
#include "IScrobbler.hpp"
namespace Database
{
class Db;
}
namespace Scrobbling
{
class InternalScrobbler final : public IScrobbler
@@ -34,7 +39,7 @@ namespace Scrobbling
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;
};
@@ -41,7 +41,7 @@
namespace
{
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()};
@@ -50,7 +50,7 @@ namespace Scrobbling::ListenBrainz
void listenStarted(const Listen& listen) override;
void listenFinished(const Listen& listen, std::optional<std::chrono::seconds> duration) 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
void enqueListen(const Listen& listen, const Wt::WDateTime& timePoint);
@@ -195,7 +195,7 @@ namespace
std::vector<Scrobbling::TimedListen> matchedListens;
};
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;
@@ -233,7 +233,7 @@ namespace
result.oldestEntry = listenedAt;
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)
@@ -285,7 +285,7 @@ namespace Scrobbling::ListenBrainz
}
ListensSynchronizer::UserContext&
ListensSynchronizer::getUserContext(Database::IdType userId)
ListensSynchronizer::getUserContext(Database::UserId userId)
{
auto itContext {_userContexts.find(userId)};
if (itContext == std::cend(_userContexts))
@@ -338,14 +338,14 @@ namespace Scrobbling::ListenBrainz
assert(!isFetching());
std::vector<Database::IdType> userIds;
std::vector<Database::UserId> userIds;
{
Database::Session& session {_db.getTLSSession()};
auto transaction {session.createSharedTransaction()};
userIds = Database::User::getAllIds(_db.getTLSSession());
}
for (const Database::IdType userId : userIds)
for (const Database::UserId userId : userIds)
{
if (Utils::getListenBrainzToken(_db.getTLSSession(), userId))
startGetListens(getUserContext(userId));
@@ -373,7 +373,7 @@ namespace Scrobbling::ListenBrainz
{
_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;
if (!isFetching())
@@ -473,7 +473,7 @@ namespace Scrobbling::ListenBrainz
}
std::optional<SendQueue::RequestData>
ListensSynchronizer::createValidateTokenRequestData(Database::IdType userId)
ListensSynchronizer::createValidateTokenRequestData(Database::UserId userId)
{
Database::Session& session {_db.getTLSSession()};
@@ -50,14 +50,14 @@ namespace Scrobbling::ListenBrainz
private:
struct UserContext
{
UserContext(Database::IdType id) : userId {id} {}
UserContext(Database::UserId id) : userId {id} {}
UserContext(const UserContext&) = delete;
UserContext(UserContext&&) = delete;
UserContext& operator=(const UserContext&) = delete;
UserContext& operator=(UserContext&&) = delete;
const Database::IdType userId;
const Database::UserId userId;
bool fetching {};
std::optional<std::size_t> listenCount {};
@@ -69,7 +69,7 @@ namespace Scrobbling::ListenBrainz
std::size_t importedListenCount{};
};
UserContext& getUserContext(Database::IdType userId);
UserContext& getUserContext(Database::UserId userId);
bool isFetching() const;
void scheduleGetListens(std::chrono::seconds fromNow);
void startGetListens();
@@ -78,7 +78,7 @@ namespace Scrobbling::ListenBrainz
void enqueValidateToken(UserContext& context);
void enqueGetListenCount(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);
void processGetListensResponse(std::string_view body, UserContext& context);
@@ -88,7 +88,7 @@ namespace Scrobbling::ListenBrainz
SendQueue& _sendQueue;
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::chrono::hours _syncListensPeriod;
@@ -30,7 +30,7 @@ static constexpr std::string_view historyTracklistName {"__scrobbler_listenbrain
namespace Scrobbling::ListenBrainz::Utils
{
std::optional<UUID>
getListenBrainzToken(Database::Session& session, Database::IdType userId)
getListenBrainzToken(Database::Session& session, Database::UserId userId)
{
auto transaction {session.createSharedTransaction()};
@@ -21,6 +21,7 @@
#include <Wt/Dbo/ptr.h>
#include "utils/UUID.hpp"
#include "database/Types.hpp"
namespace Database
@@ -32,7 +33,7 @@ namespace Database
namespace Scrobbling::ListenBrainz::Utils
{
std::optional<UUID> getListenBrainzToken(Database::Session& session, Database::IdType userId);
Wt::Dbo::ptr<Database::TrackList> getOrCreateListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user);
Wt::Dbo::ptr<Database::TrackList> getListensTrackList(Database::Session& session, Wt::Dbo::ptr<Database::User> user);
std::optional<UUID> getListenBrainzToken(Database::Session& session, Database::UserId userId);
Database::ObjectPtr<Database::TrackList> getOrCreateListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user);
Database::ObjectPtr<Database::TrackList> getListensTrackList(Database::Session& session, Database::ObjectPtr<Database::User> user);
}
@@ -24,12 +24,12 @@
#include <chrono>
#include <memory>
#include <optional>
#include <set>
#include <vector>
#include <Wt/WDateTime.h>
#include "scrobbling/Listen.hpp"
#include "database/Types.hpp"
namespace Database
{
@@ -57,42 +57,42 @@ namespace Scrobbling
// Stats
// From most recent to oldest
virtual std::vector<Wt::Dbo::ptr<Database::Artist>> getRecentArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
virtual std::vector<Database::ObjectPtr<Database::Artist>> getRecentArtists(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range,
bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Release>> getRecentReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
virtual std::vector<Database::ObjectPtr<Database::Release>> getRecentReleases(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Track>> getRecentTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
virtual std::vector<Database::ObjectPtr<Database::Track>> getRecentTracks(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) = 0;
// Top
virtual std::vector<Wt::Dbo::ptr<Database::Artist>> getTopArtists(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
virtual std::vector<Database::ObjectPtr<Database::Artist>> getTopArtists(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::TrackArtistLinkType> linkType,
std::optional<Database::Range> range,
bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Release>> getTopReleases(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
virtual std::vector<Database::ObjectPtr<Database::Release>> getTopReleases(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) = 0;
virtual std::vector<Wt::Dbo::ptr<Database::Track>> getTopTracks(Database::Session& session,
Wt::Dbo::ptr<Database::User> user,
const std::set<Database::IdType>& clusterIds,
virtual std::vector<Database::ObjectPtr<Database::Track>> getTopTracks(Database::Session& session,
Database::ObjectPtr<Database::User> user,
const std::vector<Database::ClusterId>& clusterIds,
std::optional<Database::Range> range,
bool& moreResults) = 0;
};
@@ -27,8 +27,8 @@ namespace Scrobbling
{
struct Listen
{
Database::IdType userId {};
Database::IdType trackId {};
Database::UserId userId {};
Database::TrackId trackId {};
};
struct TimedListen : public Listen