Restored recommendations based on acoustic similarities (using musicnn), fixes #301

This commit is contained in:
emeric
2026-06-02 08:32:43 +02:00
parent 1524106124
commit eb7f65878f
227 changed files with 10324 additions and 4673 deletions
+19 -1
View File
@@ -35,7 +35,7 @@ namespace lms::db
{
namespace
{
static constexpr Version LMS_DATABASE_VERSION{ 103 };
static constexpr Version LMS_DATABASE_VERSION{ 104 };
}
VersionInfo::VersionInfo()
@@ -1706,6 +1706,23 @@ FROM track)");
utils::executeCommand(*session.getDboSession(), "UPDATE scan_settings SET audio_scan_version = audio_scan_version + 1");
}
void migrateFromV103(Session& session)
{
// Drop previous track_audio_features with a brand new table dedicated to embeddings
utils::executeCommand(*session.getDboSession(), R"(DROP TABLE track_features)");
utils::executeCommand(*session.getDboSession(), R"(CREATE TABLE IF NOT EXISTS "track_musicnn_embeddings" (
"id" integer primary key autoincrement,
"version" integer not null,
"data" blob not null,
"track_id" bigint,
constraint "fk_track_musicnn_embeddings_track" foreign key ("track_id") references "track" ("id") on delete cascade deferrable initially deferred
))");
utils::executeCommand(*session.getDboSession(), "ALTER TABLE scan_settings RENAME COLUMN similarity_engine_type TO recommendation_engine_type");
utils::executeCommand(*session.getDboSession(), "ALTER TABLE scan_settings ADD COLUMN musicnn_model_identifier TEXT NOT NULL DEFAULT ''");
}
bool doDbMigration(Session& session)
{
constexpr std::string_view outdatedMsg{ "Outdated database, please rebuild it (delete the .db file and restart)" };
@@ -1785,6 +1802,7 @@ FROM track)");
{ 100, migrateFromV100 },
{ 101, migrateFromV101 },
{ 102, migrateFromV102 },
{ 103, migrateFromV103 },
};
bool migrationPerformed{};
+3 -3
View File
@@ -52,9 +52,9 @@
#include "database/objects/TrackBookmark.hpp"
#include "database/objects/TrackEmbeddedImage.hpp"
#include "database/objects/TrackEmbeddedImageLink.hpp"
#include "database/objects/TrackFeatures.hpp"
#include "database/objects/TrackList.hpp"
#include "database/objects/TrackLyrics.hpp"
#include "database/objects/TrackMusicNNEmbeddings.hpp"
#include "database/objects/UIState.hpp"
#include "database/objects/User.hpp"
@@ -106,7 +106,7 @@ namespace lms::db
_session.mapClass<TrackArtistLink>("track_artist_link");
_session.mapClass<TrackEmbeddedImage>("track_embedded_image");
_session.mapClass<TrackEmbeddedImageLink>("track_embedded_image_link");
_session.mapClass<TrackFeatures>("track_features");
_session.mapClass<TrackMusicNNEmbeddings>("track_musicnn_embeddings");
_session.mapClass<TrackList>("tracklist");
_session.mapClass<TrackListEntry>("tracklist_entry");
_session.mapClass<TrackLyrics>("track_lyrics");
@@ -302,7 +302,7 @@ namespace lms::db
utils::executeCommand(_session, "CREATE INDEX IF NOT EXISTS track_artist_link_track_artist_idx ON track_artist_link(track_id, artist_id)");
utils::executeCommand(_session, "CREATE INDEX IF NOT EXISTS track_artist_link_track_type_idx ON track_artist_link(track_id, type)");
utils::executeCommand(_session, "CREATE INDEX IF NOT EXISTS track_features_track_idx ON track_features(track_id)");
utils::executeCommand(_session, "CREATE INDEX IF NOT EXISTS track_musicnn_embeddings_track_idx ON track_musicnn_embeddings(track_id)");
utils::executeCommand(_session, "CREATE INDEX IF NOT EXISTS track_lyrics_id_idx ON track_lyrics(id)");
utils::executeCommand(_session, "CREATE INDEX IF NOT EXISTS track_lyrics_absolute_file_path_idx ON track_lyrics(absolute_file_path)");
+13 -21
View File
@@ -29,10 +29,9 @@
#include <Wt/WDateTime.h>
#include "core/ITraceLogger.hpp"
#include "core/Service.hpp"
#include "database/Types.hpp"
#include "QueryPlanRecorder.hpp"
#include "profiling/ScopedQueryProfiler.hpp"
namespace lms::db::utils
{
@@ -42,16 +41,6 @@ namespace lms::db::utils
Wt::WDateTime normalizeDateTime(const Wt::WDateTime& dateTime);
namespace detail
{
template<typename Query>
void recordQueryPlanIfNeeded(const Query& query)
{
if (IQueryPlanRecorder * recorder{ core::Service<IQueryPlanRecorder>::get() })
static_cast<QueryPlanRecorder*>(recorder)->recordQueryPlanIfNeeded(query.session(), query.asString());
}
} // namespace detail
template<typename Query>
void applyRange(Query& query, std::optional<Range> range)
{
@@ -100,20 +89,22 @@ namespace lms::db::utils
template<typename Query, typename UnaryFunc>
void forEachQueryResult(const Query& query, UnaryFunc&& func)
{
detail::recordQueryPlanIfNeeded(query);
LMS_SCOPED_TRACE_DETAILED_WITH_ARG("Database", "ForEachQueryResult", "Query", query.asString());
forEachResult(query.resultList(), std::forward<UnaryFunc>(func));
ScopedQueryProfiler queryProfiler{ query };
forEachResult(query.resultList(), [&](const auto& result) {
queryProfiler.suspend();
func(result);
queryProfiler.resume();
});
}
template<typename T, typename Query>
std::vector<T> fetchQueryResults(const Query& query)
{
detail::recordQueryPlanIfNeeded(query);
LMS_SCOPED_TRACE_DETAILED_WITH_ARG("Database", "FetchQueryResults", "Query", query.asString());
ScopedQueryProfiler queryProfiler{ query };
auto collection{ query.resultList() };
return std::vector<T>(collection.begin(), collection.end());
}
@@ -121,10 +112,9 @@ namespace lms::db::utils
template<typename Query>
std::vector<typename QueryResultType<Query>::type> fetchQueryResults(const Query& query)
{
detail::recordQueryPlanIfNeeded(query);
LMS_SCOPED_TRACE_DETAILED_WITH_ARG("Database", "FetchQueryResults", "Query", query.asString());
ScopedQueryProfiler queryProfiler{ query };
auto collection{ query.resultList() };
return std::vector<typename QueryResultType<Query>::type>(collection.begin(), collection.end());
}
@@ -132,9 +122,8 @@ namespace lms::db::utils
template<typename Query>
auto fetchQuerySingleResult(const Query& query)
{
detail::recordQueryPlanIfNeeded(query);
LMS_SCOPED_TRACE_DETAILED_WITH_ARG("Database", "FetchQuerySingleResult", "Query", query.asString());
ScopedQueryProfiler queryProfiler{ query };
return query.resultValue();
}
@@ -184,6 +173,7 @@ namespace lms::db::utils
moreResults = false;
std::size_t count{};
ScopedQueryProfiler queryProfiler{ query };
const auto collection{ query.resultList() };
auto it{ fetchFirstResult(collection) };
while (it != collection.end())
@@ -194,7 +184,9 @@ namespace lms::db::utils
break;
}
queryProfiler.suspend();
func(*it);
queryProfiler.resume();
fetchNextResult<ResultType>(it);
}
}
-42
View File
@@ -403,48 +403,6 @@ AND NOT EXISTS (
return _preferredArtwork.id();
}
RangeResults<ArtistId> Artist::findSimilarArtistIds(core::EnumSet<TrackArtistLinkType> artistLinkTypes, std::optional<Range> range) const
{
assert(session());
std::ostringstream oss;
oss << "SELECT a.id FROM artist a"
" 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_cluster t_c ON t_c.track_id = t.id"
" WHERE "
" t_c.cluster_id IN (SELECT DISTINCT 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 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 = ?)"
" AND a.id <> ?";
if (!artistLinkTypes.empty())
{
oss << " AND t_a_l.type IN (";
bool first{ true };
for (TrackArtistLinkType type : artistLinkTypes)
{
(void)type;
if (!first)
oss << ", ";
oss << "?";
first = false;
}
oss << ")";
}
auto query{ session()->query<ArtistId>(oss.str()).bind(getId()).bind(getId()).groupBy("a.id").orderBy("COUNT(*) DESC, RANDOM()") };
for (const TrackArtistLinkType type : artistLinkTypes)
query.bind(type);
return utils::execRangeQuery<ArtistId>(query, range);
}
std::vector<std::vector<Cluster::pointer>> Artist::getClusterGroups(std::span<const ClusterTypeId> clusterTypeIds, std::size_t size) const
{
assert(session());
@@ -749,33 +749,6 @@ namespace lms::db
return utils::fetchQueryResults(query);
}
std::vector<Release::pointer> Release::getSimilarReleases(std::optional<std::size_t> offset, std::optional<std::size_t> count) const
{
assert(session());
// Select the similar releases using the 5 most used clusters of the release
auto query{ session()->query<Wt::Dbo::ptr<Release>>(
"SELECT r FROM release r"
" INNER JOIN track t ON t.release_id = r.id"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id"
" WHERE "
" t_c.cluster_id IN "
"(SELECT DISTINCT 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 <> ?")
.bind(getId())
.bind(getId())
.groupBy("r.id")
.orderBy("COUNT(*) DESC, RANDOM()")
.limit(count ? static_cast<int>(*count) : -1)
.offset(offset ? static_cast<int>(*offset) : -1) };
return utils::fetchQueryResults<Release::pointer>(query);
}
ObjectPtr<Artwork> Release::getPreferredArtwork() const
{
return ObjectPtr<Artwork>{ _preferredArtwork };
+30 -38
View File
@@ -36,7 +36,6 @@
#include "database/objects/TrackArtistLink.hpp"
#include "database/objects/TrackEmbeddedImage.hpp"
#include "database/objects/TrackEmbeddedImageLink.hpp"
#include "database/objects/TrackFeatures.hpp"
#include "database/objects/TrackLyrics.hpp"
#include "database/objects/User.hpp"
@@ -191,6 +190,20 @@ namespace lms::db
if (params.fileSize.has_value())
query.where("t.file_size = ?").bind(static_cast<long long>(params.fileSize.value()));
if (params.hasMusicNNEmbeddings.has_value())
{
if (*params.hasMusicNNEmbeddings)
query.where("EXISTS (SELECT t_m_e.track_id FROM track_musicnn_embeddings t_m_e WHERE t_m_e.track_id = t.id)");
else
query.where("NOT EXISTS (SELECT t_m_e.track_id FROM track_musicnn_embeddings t_m_e WHERE t_m_e.track_id = t.id)");
}
if (params.lastTrackId.isValid())
{
assert(params.sortMethod == TrackSortMethod::Id);
query.where("t.id > ?").bind(params.lastTrackId);
}
if (params.embeddedImageId.isValid())
{
query.join("track_embedded_image_link t_e_i_l ON t_e_i_l.track_id = t.id");
@@ -322,7 +335,7 @@ namespace lms::db
});
}
void Track::findAbsoluteFilePath(Session& session, TrackId& lastRetrievedId, std::size_t count, const std::function<void(TrackId trackId, const std::filesystem::path& absoluteFilePath)>& func)
void Track::findAbsoluteFilePath(Session& session, TrackId& lastRetrievedId, std::size_t count, const TrackLocationVisitor& func)
{
session.checkReadTransaction();
@@ -334,6 +347,19 @@ namespace lms::db
});
}
void Track::findAbsoluteFilePath(Session& session, const FindParameters& params, const TrackLocationVisitor& func)
{
session.checkReadTransaction();
std::string_view itemToSelect{ "t.id, t.absolute_file_path" };
auto query{ createQuery<std::tuple<TrackId, std::filesystem::path>>(session, itemToSelect, params) };
utils::forEachQueryRangeResult(query, params.range, [&](const auto& res) {
func(std::get<0>(res), std::get<1>(res));
});
}
void Track::find(Session& session, const IdRange<TrackId>& idRange, const std::function<void(const Track::pointer&)>& func)
{
assert(idRange.isValid());
@@ -385,15 +411,6 @@ namespace lms::db
return utils::execRangeQuery<TrackId>(query, range);
}
RangeResults<TrackId> Track::findIdsWithRecordingMBIDAndMissingFeatures(Session& session, std::optional<Range> range)
{
session.checkReadTransaction();
auto query{ session.getDboSession()->query<TrackId>("SELECT t.id FROM track t").where("LENGTH(t.recording_mbid) > 0").where("NOT EXISTS (SELECT * FROM track_features t_f WHERE t_f.track_id = t.id)") };
return utils::execRangeQuery<TrackId>(query, range);
}
void Track::updatePreferredArtwork(Session& session, TrackId trackId, ArtworkId artworkId)
{
session.checkWriteTransaction();
@@ -490,36 +507,11 @@ namespace lms::db
utils::forEachQueryRangeResult(query, params.range, moreResults, func);
}
RangeResults<TrackId> Track::findSimilarTrackIds(Session& session, const std::vector<TrackId>& tracks, std::optional<Range> range)
std::size_t Track::getCount(Session& session, const FindParameters& params)
{
assert(!tracks.empty());
session.checkReadTransaction();
std::ostringstream oss;
for (std::size_t i{}; i < tracks.size(); ++i)
{
if (!oss.str().empty())
oss << ", ";
oss << "?";
}
auto query{ session.getDboSession()->query<TrackId>(
"SELECT t.id FROM track t"
" INNER JOIN track_cluster t_c ON t_c.track_id = t.id"
" AND t_c.cluster_id IN (SELECT DISTINCT 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.id NOT IN ("
+ oss.str() + ")")
.groupBy("t.id")
.orderBy("COUNT(*) DESC, RANDOM()") };
for (TrackId trackId : tracks)
query.bind(trackId);
for (TrackId trackId : tracks)
query.bind(trackId);
return utils::execRangeQuery<TrackId>(query, range);
return utils::fetchQuerySingleResult(createQuery<int>(session, "COUNT(*)", params));
}
void Track::setAbsoluteFilePath(const std::filesystem::path& filePath)
@@ -1,123 +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/>.
*/
#include "database/objects/TrackFeatures.hpp"
#include <Wt/Dbo/Impl.h>
#include <boost/property_tree/json_parser.hpp>
#include <boost/property_tree/ptree.hpp>
#include "core/ILogger.hpp"
#include "database/Session.hpp"
#include "database/objects/Directory.hpp"
#include "database/objects/Track.hpp"
#include "Utils.hpp"
#include "traits/IdTypeTraits.hpp"
DBO_INSTANTIATE_TEMPLATES(lms::db::TrackFeatures)
namespace lms::db
{
TrackFeatures::TrackFeatures(ObjectPtr<Track> track, const std::string& jsonEncodedFeatures)
: _data{ jsonEncodedFeatures }
, _track{ getDboPtr(track) }
{
}
TrackFeatures::pointer TrackFeatures::create(Session& session, ObjectPtr<Track> track, const std::string& jsonEncodedFeatures)
{
return session.getDboSession()->add(std::unique_ptr<TrackFeatures>{ new TrackFeatures{ track, jsonEncodedFeatures } });
}
std::size_t TrackFeatures::getCount(Session& session)
{
session.checkReadTransaction();
return utils::fetchQuerySingleResult(session.getDboSession()->query<int>("SELECT COUNT(*) FROM track_features"));
}
TrackFeatures::pointer TrackFeatures::find(Session& session, TrackFeaturesId id)
{
session.checkReadTransaction();
return utils::fetchQuerySingleResult(session.getDboSession()->find<TrackFeatures>().where("id = ?").bind(id));
}
TrackFeatures::pointer TrackFeatures::find(Session& session, TrackId trackId)
{
session.checkReadTransaction();
return utils::fetchQuerySingleResult(session.getDboSession()->find<TrackFeatures>().where("track_id = ?").bind(trackId));
}
RangeResults<TrackFeaturesId> TrackFeatures::find(Session& session, std::optional<Range> range)
{
session.checkReadTransaction();
auto query{ session.getDboSession()->query<TrackFeaturesId>("SELECT id from track_features") };
return utils::execRangeQuery<TrackFeaturesId>(query, range);
}
FeatureValues TrackFeatures::getFeatureValues(const FeatureName& featureNode) const
{
FeatureValuesMap featuresValuesMap{ getFeatureValuesMap({ featureNode }) };
return std::move(featuresValuesMap[featureNode]);
}
FeatureValuesMap TrackFeatures::getFeatureValuesMap(const std::unordered_set<FeatureName>& featureNames) const
{
FeatureValuesMap res;
try
{
std::istringstream iss{ _data };
boost::property_tree::ptree root;
boost::property_tree::read_json(iss, root);
for (const FeatureName& featureName : featureNames)
{
FeatureValues& featureValues{ res[featureName] };
auto node{ root.get_child(featureName) };
bool hasChildren = false;
for (const auto& child : node.get_child(""))
{
hasChildren = true;
featureValues.push_back(child.second.get_value<double>());
}
if (!hasChildren)
featureValues.push_back(node.get_value<double>());
}
}
catch (boost::property_tree::ptree_error& error)
{
LMS_LOG(DB, ERROR, "Track " << _track.id() << ": ptree exception: " << error.what());
res.clear();
}
return res;
}
} // namespace lms::db
@@ -0,0 +1,99 @@
/*
* Copyright (C) 2026 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/>.
*/
#include "database/objects/TrackMusicNNEmbeddings.hpp"
#include <Wt/Dbo/Impl.h>
#include "database/Session.hpp"
#include "database/objects/Track.hpp"
#include "Utils.hpp"
#include "traits/IdTypeTraits.hpp"
DBO_INSTANTIATE_TEMPLATES(lms::db::TrackMusicNNEmbeddings)
namespace lms::db
{
TrackMusicNNEmbeddings::TrackMusicNNEmbeddings(ObjectPtr<Track> track)
: _track{ getDboPtr(track) }
{
}
TrackMusicNNEmbeddings::pointer TrackMusicNNEmbeddings::create(Session& session, ObjectPtr<Track> track)
{
return session.getDboSession()->add(std::unique_ptr<TrackMusicNNEmbeddings>{ new TrackMusicNNEmbeddings{ track } });
}
std::size_t TrackMusicNNEmbeddings::getCount(Session& session)
{
session.checkReadTransaction();
return utils::fetchQuerySingleResult(session.getDboSession()->query<int>("SELECT COUNT(*) FROM track_musicnn_embeddings"));
}
TrackMusicNNEmbeddings::pointer TrackMusicNNEmbeddings::find(Session& session, TrackMusicNNEmbeddingsId id)
{
session.checkReadTransaction();
return utils::fetchQuerySingleResult(session.getDboSession()->find<TrackMusicNNEmbeddings>().where("id = ?").bind(id));
}
TrackMusicNNEmbeddings::pointer TrackMusicNNEmbeddings::find(Session& session, TrackId trackId)
{
session.checkReadTransaction();
return utils::fetchQuerySingleResult(session.getDboSession()->find<TrackMusicNNEmbeddings>().where("track_id = ?").bind(trackId));
}
RangeResults<TrackMusicNNEmbeddingsId> TrackMusicNNEmbeddings::find(Session& session, std::optional<Range> range)
{
session.checkReadTransaction();
auto query{ session.getDboSession()->query<TrackMusicNNEmbeddingsId>("SELECT id from track_musicnn_embeddings") };
return utils::execRangeQuery<TrackMusicNNEmbeddingsId>(query, range);
}
void TrackMusicNNEmbeddings::find(Session& session, std::function<void(const pointer&)> func)
{
auto query{ session.getDboSession()->find<TrackMusicNNEmbeddings>() };
utils::forEachQueryResult(query, [&](const TrackMusicNNEmbeddings::pointer& embeddings) {
func(embeddings);
});
}
void TrackMusicNNEmbeddings::removeAll(Session& session)
{
session.checkWriteTransaction();
utils::executeCommand(*session.getDboSession(), "DELETE FROM track_musicnn_embeddings");
}
std::span<const std::byte> TrackMusicNNEmbeddings::getData() const
{
return std::span<const std::byte>{ reinterpret_cast<const std::byte*>(_data.data()), _data.size() };
}
void TrackMusicNNEmbeddings::setData(std::span<const std::byte> data)
{
const auto* start{ reinterpret_cast<const unsigned char*>(data.data()) };
_data.assign(start, start + data.size());
}
} // namespace lms::db
@@ -17,7 +17,7 @@
* along with LMS. If not, see <http://www.gnu.org/licenses/>.
*/
#include "QueryPlanRecorder.hpp"
#include "profiling/QueryProfiler.hpp"
#include <memory>
#include <mutex>
@@ -29,35 +29,38 @@
namespace lms::db
{
std::unique_ptr<IQueryPlanRecorder> createQueryPlanRecorder()
std::unique_ptr<IQueryProfiler> createQueryProfiler()
{
return std::make_unique<QueryPlanRecorder>();
return std::make_unique<QueryProfiler>();
}
QueryPlanRecorder::QueryPlanRecorder()
QueryProfiler::QueryProfiler()
{
LMS_LOG(DB, INFO, "Recording database query plans");
LMS_LOG(DB, INFO, "Recording database queries");
}
QueryPlanRecorder::~QueryPlanRecorder() = default;
QueryProfiler::~QueryProfiler() = default;
void QueryPlanRecorder::visitQueryPlans(const QueryPlanVisitor& visitor) const
void QueryProfiler::visitQueries(const QueryVisitor& visitor) const
{
const std::shared_lock lock{ _mutex };
for (const auto& [query, plan] : _queryPlans)
visitor(query, plan);
for (const auto& [query, data] : _queries)
{
const QueryStats stats{
.query = query,
.plan = data.plan,
.callCount = data.timeStats.getCount(),
.totalTime = std::chrono::microseconds{ static_cast<long long>(data.timeStats.getMean() * static_cast<double>(data.timeStats.getCount())) },
.meanTime = std::chrono::microseconds{ static_cast<long long>(data.timeStats.getMean()) },
.stdDevTime = std::chrono::microseconds{ static_cast<long long>(data.timeStats.getSampleStdDev()) },
};
visitor(stats);
}
}
void QueryPlanRecorder::recordQueryPlanIfNeeded(Wt::Dbo::Session& session, const std::string& query)
void QueryProfiler::recordQueryPlan(Wt::Dbo::Session& session, const std::string& query)
{
{
const std::shared_lock lock{ _mutex };
if (_queryPlans.contains(query))
return;
}
Wt::Dbo::Transaction transaction{ session };
Wt::Dbo::SqlConnection* connection{ transaction.connection() };
@@ -106,7 +109,23 @@ namespace lms::db
{
const std::unique_lock lock{ _mutex };
_queryPlans.try_emplace(query, std::move(result));
_queries[query].plan = std::move(result);
}
}
void QueryProfiler::recordQueryExecution(Wt::Dbo::Session& session, const std::string& query, Clock::duration elapsed)
{
bool needQueryPlan{};
const double elapsedUs{ std::chrono::duration_cast<std::chrono::duration<double, std::micro>>(elapsed).count() };
{
std::unique_lock lock{ _mutex };
auto& queryStats{ _queries[query] };
queryStats.timeStats.add(elapsedUs);
needQueryPlan = queryStats.plan.empty();
}
if (needQueryPlan)
recordQueryPlan(session, query);
}
} // namespace lms::db
@@ -25,24 +25,33 @@
#include <Wt/Dbo/Session.h>
#include "database/IQueryPlanRecorder.hpp"
#include "database/profiling/IQueryProfiler.hpp"
#include "math/StatsAccumulator.hpp"
namespace lms::db
{
class QueryPlanRecorder : public IQueryPlanRecorder
class QueryProfiler : public IQueryProfiler
{
public:
QueryPlanRecorder();
~QueryPlanRecorder() override;
QueryPlanRecorder(const QueryPlanRecorder&) = delete;
QueryPlanRecorder& operator=(const QueryPlanRecorder&) = delete;
QueryProfiler();
~QueryProfiler() override;
QueryProfiler(const QueryProfiler&) = delete;
QueryProfiler& operator=(const QueryProfiler&) = delete;
void visitQueryPlans(const QueryPlanVisitor& visitor) const override;
void visitQueries(const QueryVisitor& visitor) const override;
void recordQueryPlanIfNeeded(Wt::Dbo::Session& session, const std::string& query);
void recordQueryExecution(Wt::Dbo::Session& session, const std::string& query, Clock::duration elapsed);
private:
void recordQueryPlan(Wt::Dbo::Session& session, const std::string& query);
struct QueryData
{
std::string plan;
math::StatsAccumulator<double> timeStats; // in Us
};
mutable std::shared_mutex _mutex;
std::map<std::string, std::string> _queryPlans;
std::map<std::string, QueryData> _queries;
};
} // namespace lms::db
@@ -0,0 +1,86 @@
/*
* Copyright (C) 2025 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 <cassert>
#include <string>
#include "core/Service.hpp"
#include "database/profiling/IQueryProfiler.hpp"
#include "profiling/QueryProfiler.hpp"
namespace lms::db::utils
{
template<typename Query>
class ScopedQueryProfiler
{
public:
explicit ScopedQueryProfiler(const Query& query)
: _recorder{ static_cast<QueryProfiler*>(core::Service<IQueryProfiler>::get()) }
{
if (_recorder)
{
_query = &query;
_start = IQueryProfiler::Clock::now();
}
}
~ScopedQueryProfiler()
{
if (_recorder)
{
if (_active)
_elapsed += IQueryProfiler::Clock::now() - _start;
_recorder->recordQueryExecution(_query->session(), _query->asString(), _elapsed);
}
}
ScopedQueryProfiler(const ScopedQueryProfiler&) = delete;
ScopedQueryProfiler& operator=(const ScopedQueryProfiler&) = delete;
void suspend()
{
if (_recorder)
{
assert(_active);
_elapsed += IQueryProfiler::Clock::now() - _start;
_active = false;
}
}
void resume()
{
if (_recorder)
{
assert(!_active);
_start = IQueryProfiler::Clock::now();
_active = true;
}
}
private:
QueryProfiler* _recorder{};
const Query* _query{};
IQueryProfiler::Clock::time_point _start;
IQueryProfiler::Clock::duration _elapsed{};
bool _active{ true };
};
} // namespace lms::db::utils