package persistence import ( "context" "database/sql" "errors" "fmt" "iter" "reflect" "regexp" "strings" "time" . "github.com/Masterminds/squirrel" "github.com/deluan/rest" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" id2 "github.com/navidrome/navidrome/model/id" "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/utils/hasher" "github.com/navidrome/navidrome/utils/slice" "github.com/pocketbase/dbx" "github.com/zeebo/xxh3" ) // sqlRepository is the base repository for all SQL repositories. It provides common functions to interact with the DB. // When creating a new repository using this base, you must: // // - Embed this struct. // - Set the db field. // - Call registerModel with the model instance and any possible filters. // - If the model has a different table name than the default (lowercase of the model name), it should be set manually // using the tableName field. // - Sort mappings must be set with setSortMappings method. If a sort field is not in the map, it will be used as the name of the column. // // All fields in filters and sortMappings must be in snake_case. Only sorts and filters based on real field names or // defined in the mappings will be allowed. type sqlRepository struct { tableName string db dbx.Builder // Do not set these fields manually, they are set by the registerModel method filterMappings map[string]filterFunc isFieldWhiteListed fieldWhiteListedFunc // Do not set this field manually, it is set by the setSortMappings method sortMappings map[string]string } const invalidUserId = "-1" func loggedUser(ctx context.Context) *model.User { if user, ok := request.UserFrom(ctx); !ok { return &model.User{ID: invalidUserId} } else { return &user } } // ownerFilter returns the predicate restricting access to rows owned by the logged-in user, for // tables with a user_id column. It returns nil for admins and for headless/system contexts (invalid // user), meaning "no ownership restriction". Callers should skip the WHERE clause when it is nil. // // The predicate uses an unqualified user_id, so it only works on queries where that column is // unambiguous (no join introducing a second user_id). func (r sqlRepository) ownerFilter(ctx context.Context) Sqlizer { if usr := loggedUser(ctx); !usr.IsAdmin && usr.ID != invalidUserId { return Eq{"user_id": usr.ID} } return nil } // addRestriction combines an optional caller predicate with the ownership filter, producing the // WHERE clause for owner-scoped reads. For admins and headless contexts ownerFilter() is nil and // only the caller's predicate (if any) remains. func (r sqlRepository) addRestriction(ctx context.Context, sql ...Sqlizer) Sqlizer { s := And{} if len(sql) > 0 { s = append(s, sql[0]) } if owner := r.ownerFilter(ctx); owner != nil { s = append(s, owner) } return s } // writeAccess says who may change a row in a table with a user_id column. type writeAccess int const ( ownerOrAdmin writeAccess = iota // admins may write any row ownerOnly // even admins may only write their own rows ) // ownedRow matches the row rowID only if the logged-in user may write it under access. func (r sqlRepository) ownedRow(ctx context.Context, rowID string, access writeAccess) Sqlizer { if access == ownerOnly { return And{Eq{"id": rowID}, Eq{"user_id": loggedUser(ctx).ID}} } return r.addRestriction(ctx, Eq{"id": rowID}) } func (r *sqlRepository) registerModel(instance any, filters map[string]filterFunc) { if r.tableName == "" { r.tableName = strings.TrimPrefix(reflect.TypeOf(instance).String(), "*model.") r.tableName = toSnakeCase(r.tableName) } r.tableName = strings.ToLower(r.tableName) r.isFieldWhiteListed = registerModelWhiteList(instance) r.filterMappings = filters } // setSortMappings sets the mappings for the sort fields. If the sort field is not in the map, it will be used as is. // This applies per comma-separated part, so a key added here also defines that bare name wherever a // caller uses it inside a sort list. // // If PreferSortTags is enabled, it will map the order fields to the corresponding sort expression, // which gives precedence to sort tags. // Ex: order_title => (coalesce(nullif(sort_title,”),order_title) collate nocase) // To avoid performance issues, indexes should be created for these sort expressions // // NOTE: if an individual item has spaces, it should be wrapped in parentheses. For example, // you should write "(lyrics != '[]')". This prevents the item being split unexpectedly. // Without parentheses, "lyrics != '[]'" would be mapped as simply "lyrics" func (r *sqlRepository) setSortMappings(mappings map[string]string, tableName ...string) { tn := r.tableName if len(tableName) > 0 { tn = tableName[0] } if conf.Server.PreferSortTags || conf.Server.EnableNaturalSorting { for k, v := range mappings { mappings[k] = mapSortOrder(tn, v) } } r.sortMappings = mappings } func (r sqlRepository) newSelect(ctx context.Context, options ...model.QueryOptions) SelectBuilder { sq := Select().From(r.tableName) if len(options) > 0 { r.resetSeededRandom(ctx, options) sq = r.applyOptions(sq, options...) sq = r.applyFilters(sq, options...) } return sq } func (r sqlRepository) applyOptions(sq SelectBuilder, options ...model.QueryOptions) SelectBuilder { if len(options) > 0 { if options[0].Max > 0 { sq = sq.Limit(uint64(options[0].Max)) } if options[0].Offset > 0 { sq = sq.Offset(uint64(options[0].Offset)) } if options[0].Sort != "" { sq = sq.OrderBy(r.buildSortOrder(options[0].Sort, options[0].Order)) } } return sq } // TODO Change all sortMappings to have a consistent case func (r sqlRepository) sortMapping(sort string) string { if mapping, _, ok := r.lookupSortMapping(sort); ok { return mapping } // Each part of a comma list is resolved on its own, so a mix of mapped keys and plain columns // keeps the mappings the recognized parts have. parts := strings.FieldsFunc(sort, splitFunc(',')) mapped := make([]string, 0, len(parts)) for _, part := range parts { part = strings.TrimSpace(part) if partMapping, _, ok := r.lookupSortMapping(part); ok { part = partMapping } else { part = toSnakeCase(part) } mapped = append(mapped, part) } return strings.Join(mapped, ", ") } // lookupSortMapping also returns the snake_case form when it had to derive one, so a caller's // fallback doesn't recompute it: toSnakeCase runs two regexps. func (r sqlRepository) lookupSortMapping(sort string) (mapping, snakeCased string, ok bool) { if mapping, ok = r.sortMappings[sort]; ok { return mapping, sort, true } if mapping, ok = r.sortMappings[toCamelCase(sort)]; ok { return mapping, "", true } snakeCased = toSnakeCase(sort) mapping, ok = r.sortMappings[snakeCased] return mapping, snakeCased, ok } func (r sqlRepository) buildSortOrder(sort, order string) string { sort = r.sortMapping(sort) order = strings.ToLower(strings.TrimSpace(order)) var reverseOrder string if order == "desc" { reverseOrder = "asc" } else { order = "asc" reverseOrder = "desc" } parts := strings.FieldsFunc(sort, splitFunc(',')) newSort := make([]string, 0, len(parts)) for _, p := range parts { f := strings.FieldsFunc(p, splitFunc(' ')) newField := make([]string, 1, len(f)) newField[0] = f[0] if len(f) == 1 { newField = append(newField, order) } else { if f[1] == "asc" { newField = append(newField, order) } else { newField = append(newField, reverseOrder) } } newSort = append(newSort, strings.Join(newField, " ")) } return strings.Join(newSort, ", ") } func splitFunc(delimiter rune) func(c rune) bool { open := 0 return func(c rune) bool { if c == '(' { open++ return false } if open > 0 { if c == ')' { open-- } return false } return c == delimiter } } func (r sqlRepository) applyFilters(sq SelectBuilder, options ...model.QueryOptions) SelectBuilder { if len(options) > 0 && options[0].Filters != nil { sq = sq.Where(options[0].Filters) } return sq } // libraryIdFilter is a filter function to be added to resources that have a library_id column. func libraryIdFilter(_ string, value any) Sqlizer { return Eq{"library_id": value} } // applyLibraryFilter adds library filtering to queries for tables that have a library_id column // This ensures users only see content from libraries they have access to func (r sqlRepository) applyLibraryFilter(ctx context.Context, sq SelectBuilder, tableName ...string) SelectBuilder { user := loggedUser(ctx) // If the user is an admin, or the user ID is invalid (e.g., when no user is logged in), skip the library filter if user.IsAdmin || user.ID == invalidUserId { return sq } // A non-admin granted every library sees everything the subquery would return, so applying it is // pure overhead. Skip it in that case (same fast path admins get). if visible, err := r.visibleLibraryIDs(ctx); err == nil && r.userSeesAllLibraries(ctx, visible) { return sq } table := r.tableName if len(tableName) > 0 { table = tableName[0] } // Get user's accessible library IDs // Use subquery to filter by user's library access return sq.Where(Expr(table+".library_id IN ("+ "SELECT ul.library_id FROM user_library ul WHERE ul.user_id = ?)", user.ID)) } // userSeesAllLibraries reports whether the visible set already covers every library, so a // library filter would exclude nothing. func (r sqlRepository) userSeesAllLibraries(ctx context.Context, visible []int) bool { user := loggedUser(ctx) if user.IsAdmin || user.ID == invalidUserId { return true // visible is the whole library table } var res struct{ Count int64 } err := r.queryOne(ctx, Select("count(*) as count").From("library"), &res) if err != nil || res.Count == 0 { return false } return int64(len(visible)) == res.Count } // visibleLibraryIDs returns the libraries the current user can see: all libraries for admin and // headless processes, otherwise the user's granted libraries. func (r sqlRepository) visibleLibraryIDs(ctx context.Context) ([]int, error) { user := loggedUser(ctx) if user.IsAdmin || user.ID == invalidUserId { var ids []int err := r.queryAllSlice(ctx, Select("id").From("library"), &ids) return ids, err } return slice.Map(user.Libraries, func(lib model.Library) int { return lib.ID }), nil } func (r sqlRepository) seedKey(ctx context.Context) string { // Seed keys must be all lowercase, or else SQLite3 will encode it, making it not match the seed // used in the query. Hashing the user ID and converting it to a hex string will do the trick userIDHash := xxh3.Hash([]byte(loggedUser(ctx).ID)) return fmt.Sprintf("%s|%016x", r.tableName, userIDHash) } func (r sqlRepository) resetSeededRandom(ctx context.Context, options []model.QueryOptions) { if len(options) == 0 || options[0].Sort != "random" { return } // CAST: playlist_tracks.id is an INTEGER (unlike other tables' TEXT ids); passing it to // SEEDEDRAND's string param uncast silently drops every row (go-sqlite3 binding gotcha). options[0].Sort = fmt.Sprintf("SEEDEDRAND('%s', CAST(%s.id AS TEXT))", r.seedKey(ctx), r.tableName) if options[0].Seed != "" { hasher.SetSeed(r.seedKey(ctx), options[0].Seed) return } if options[0].Offset == 0 { hasher.Reseed(r.seedKey(ctx)) } } func (r sqlRepository) executeSQL(ctx context.Context, sq Sqlizer) (int64, error) { query, args, err := r.toSQL(sq) if err != nil { return 0, err } start := time.Now() var c int64 res, err := r.db.NewQuery(query).Bind(args).WithContext(ctx).Execute() if res != nil { c, _ = res.RowsAffected() } r.logSQL(ctx, query, args, err, c, start) if err != nil { if err.Error() != "LastInsertId is not supported by this driver" { return 0, err } } return c, err } var placeholderRegex = regexp.MustCompile(`\?`) func (r sqlRepository) toSQL(sq Sqlizer) (string, dbx.Params, error) { query, args, err := sq.ToSql() if err != nil { return "", nil, err } // Replace query placeholders with named params params := make(dbx.Params, len(args)) counter := 0 result := placeholderRegex.ReplaceAllStringFunc(query, func(_ string) string { p := fmt.Sprintf("p%d", counter) params[p] = args[counter] counter++ return "{:" + p + "}" }) return result, params, nil } func (r sqlRepository) queryOne(ctx context.Context, sq Sqlizer, response any) error { query, args, err := r.toSQL(sq) if err != nil { return err } start := time.Now() err = r.db.NewQuery(query).Bind(args).WithContext(ctx).One(response) if errors.Is(err, sql.ErrNoRows) { r.logSQL(ctx, query, args, nil, 0, start) return model.ErrNotFound } r.logSQL(ctx, query, args, err, 1, start) return err } // wrapCursor adapts a cursor over db rows into one over their models. toModel pulls out the row's // embedded model, which a type parameter can't reach on its own. func wrapCursor[D, T any](cursor iter.Seq2[D, error], toModel func(D) *T) iter.Seq2[T, error] { return func(yield func(T, error) bool) { for row, err := range cursor { m := toModel(row) if m == nil { // Don't format row: its String() derefs the nil model (golang/go#81238). var zero T yield(zero, fmt.Errorf("unexpected nil %T: %w", zero, err)) return } if !yield(*m, err) || err != nil { return } } } } // queryWithStableResults is a helper function to execute a query and return an iterator that will yield its results // from a cursor, guaranteeing that the results will be stable, even if the underlying data changes. func queryWithStableResults[T any](ctx context.Context, r sqlRepository, sq SelectBuilder, options ...model.QueryOptions) (iter.Seq2[T, error], error) { if len(options) > 0 && options[0].Offset > 0 { sq = r.optimizePagination(sq, options[0]) } query, args, err := r.toSQL(sq) if err != nil { return nil, err } start := time.Now() rows, err := r.db.NewQuery(query).Bind(args).WithContext(ctx).Rows() r.logSQL(ctx, query, args, err, -1, start) if err != nil { return nil, err } return func(yield func(T, error) bool) { defer rows.Close() for rows.Next() { var row T err := rows.ScanStruct(&row) if !yield(row, err) || err != nil { return } } if err := rows.Err(); err != nil { var empty T yield(empty, err) } }, nil } func (r sqlRepository) queryAll(ctx context.Context, sq SelectBuilder, response any, options ...model.QueryOptions) error { if len(options) > 0 && options[0].Offset > 0 { sq = r.optimizePagination(sq, options[0]) } query, args, err := r.toSQL(sq) if err != nil { return err } start := time.Now() err = r.db.NewQuery(query).Bind(args).WithContext(ctx).All(response) if errors.Is(err, sql.ErrNoRows) { r.logSQL(ctx, query, args, nil, -1, start) return model.ErrNotFound } r.logSQL(ctx, query, args, err, int64(reflect.ValueOf(response).Elem().Len()), start) return err } // queryAllSlice is a helper function to query a single column and return the result in a slice func (r sqlRepository) queryAllSlice(ctx context.Context, sq SelectBuilder, response any) error { query, args, err := r.toSQL(sq) if err != nil { return err } start := time.Now() err = r.db.NewQuery(query).Bind(args).WithContext(ctx).Column(response) if errors.Is(err, sql.ErrNoRows) { r.logSQL(ctx, query, args, nil, -1, start) return model.ErrNotFound } r.logSQL(ctx, query, args, err, int64(reflect.ValueOf(response).Elem().Len()), start) return err } // optimizePagination uses a less inefficient pagination, by not using OFFSET. // See https://gist.github.com/ssokolow/262503 func (r sqlRepository) optimizePagination(sq SelectBuilder, options model.QueryOptions) SelectBuilder { if options.Offset > conf.Server.DevOffsetOptimize { sq = sq.RemoveOffset() rowidSq := sq.RemoveColumns().Columns(r.tableName + ".rowid") rowidSq = rowidSq.Limit(uint64(options.Offset)) rowidSql, args, _ := rowidSq.ToSql() sq = sq.Where(r.tableName+".rowid not in ("+rowidSql+")", args...) } return sq } func (r sqlRepository) exists(ctx context.Context, cond Sqlizer) (bool, error) { existsQuery := Select("count(*) as exist").From(r.tableName).Where(cond) var res struct{ Exist int64 } err := r.queryOne(ctx, existsQuery, &res) return res.Exist > 0, err } // updateOwned performs an atomic, ownership-restricted update of the row identified by id, for // repositories whose table has a user_id column. Non-admins can only update rows they own: the // ownership predicate is part of the UPDATE's WHERE clause, so a row owned by another user simply // does not match and no write happens. Ownership itself is immutable here: user_id is never written, // so no caller (admin included) can reassign a row to a different owner via an update. Unlike put, // it never falls through to an INSERT, so a non-matching id never creates a row. // // When the update matches no row it classifies the failure: if the row exists but is owned by // another user it returns rest.ErrPermissionDenied, otherwise rest.ErrNotFound. The write itself is // still atomic; the extra lookup happens only on the failure path (count == 0), where no write // occurred, so there is no TOCTOU on the update. func (r sqlRepository) updateOwned(ctx context.Context, id string, m any, colsToUpdate ...string) error { values, err := toSQLArgs(m) if err != nil { return fmt.Errorf("error preparing values to write to DB: %w", err) } updateValues := filterUpdateValues(values, id, colsToUpdate...) delete(updateValues, "user_id") // ownership is immutable on update return r.updateOwnedRow(ctx, id, ownerOrAdmin, updateValues) } // updateOwnedRow sets values on the row rowID if the logged-in user may write it under access. func (r sqlRepository) updateOwnedRow(ctx context.Context, rowID string, access writeAccess, values map[string]any) error { return r.runRowWrite(ctx, rowID, Update(r.tableName).SetMap(values).Where(r.ownedRow(ctx, rowID, access))) } // deleteOwned performs an atomic, ownership-restricted delete of the row identified by id, for // repositories whose table has a user_id column. Non-admins can only delete rows they own: the // ownership predicate is part of the DELETE's WHERE clause, so a row owned by another user simply // does not match and is left untouched. The failure path mirrors updateOwned (see // classifyOwnedWriteMiss), so there is no TOCTOU on the delete. func (r sqlRepository) deleteOwned(ctx context.Context, id string) error { return r.runRowWrite(ctx, id, Delete(r.tableName).Where(r.ownedRow(ctx, id, ownerOrAdmin))) } // runRowWrite executes q, a write already filtered by ownedRow(rowID, …), and classifies a miss. func (r sqlRepository) runRowWrite(ctx context.Context, rowID string, q Sqlizer) error { count, err := r.executeSQL(ctx, q) if err != nil { return err } if count == 0 { return r.classifyOwnedWriteMiss(ctx, rowID) } return nil } func (r sqlRepository) deleteOwnedAll(ctx context.Context, ids ...string) error { for _, id := range ids { if err := r.deleteOwned(ctx, id); err != nil { return err } } return nil } // classifyOwnedWriteMiss explains why an ownership-filtered write (updateOwned/deleteOwned) matched // no row: rest.ErrPermissionDenied if the row exists but is owned by another user, otherwise // rest.ErrNotFound. It runs only on the failure path (count == 0), where no write occurred. func (r sqlRepository) classifyOwnedWriteMiss(ctx context.Context, id string) error { exists, err := r.exists(ctx, Eq{"id": id}) if err != nil { return err } if exists { return rest.ErrPermissionDenied } return rest.ErrNotFound } func (r sqlRepository) count(ctx context.Context, countQuery SelectBuilder, options ...model.QueryOptions) (int64, error) { countQuery = countQuery. RemoveColumns().Columns("count(distinct " + r.tableName + ".id) as count"). RemoveOffset().RemoveLimit(). OrderBy(r.tableName + ".id"). // To remove any ORDER BY clause that could slow down the query From(r.tableName) countQuery = r.applyFilters(countQuery, options...) var res struct{ Count int64 } err := r.queryOne(ctx, countQuery, &res) return res.Count, err } func (r sqlRepository) putByMatch(ctx context.Context, filter Sqlizer, id string, m any, colsToUpdate ...string) (string, error) { if id != "" { return r.put(ctx, id, m, colsToUpdate...) } existsQuery := r.newSelect(ctx).Columns("id").From(r.tableName).Where(filter) var res struct{ ID string } err := r.queryOne(ctx, existsQuery, &res) if err != nil && !errors.Is(err, model.ErrNotFound) { return "", err } return r.put(ctx, res.ID, m, colsToUpdate...) } // selectUpdateColumns keeps only the requested colsToUpdate (or all columns when none are // specified), dropping columns that must never be overwritten on update (created_at, birth_time). func selectUpdateColumns(values map[string]any, colsToUpdate ...string) map[string]any { updateValues := map[string]any{} c2upd := slice.ToMap(colsToUpdate, func(s string) (string, struct{}) { return toSnakeCase(s), struct{}{} }) for k, v := range values { if _, found := c2upd[k]; len(c2upd) == 0 || found { updateValues[k] = v } } delete(updateValues, "created_at") // To avoid updating the media_file birth_time on each scan. Not the best solution, but it works for now // TODO move to mediafile_repository when each repo has its own upsert method delete(updateValues, "birth_time") return updateValues } func filterUpdateValues(values map[string]any, id string, colsToUpdate ...string) map[string]any { updateValues := selectUpdateColumns(values, colsToUpdate...) updateValues["id"] = id return updateValues } func (r sqlRepository) put(ctx context.Context, id string, m any, colsToUpdate ...string) (newId string, err error) { values, err := toSQLArgs(m) if err != nil { return "", fmt.Errorf("error preparing values to write to DB: %w", err) } // If there's an ID, try to update first if id != "" { update := Update(r.tableName).Where(Eq{"id": id}).SetMap(filterUpdateValues(values, id, colsToUpdate...)) count, err := r.executeSQL(ctx, update) if err != nil { return "", err } if count > 0 { return id, nil } } // If it does not have an ID OR the ID was not found (when it is a new record with predefined id) if id == "" { id = id2.NewRandom() values["id"] = id } insert := Insert(r.tableName).SetMap(values) _, err = r.executeSQL(ctx, insert) return id, err } func (r sqlRepository) delete(ctx context.Context, cond Sqlizer) error { _, err := r.executeSQL(ctx, Delete(r.tableName).Where(cond)) return err } // deleteByID is for single-item deletes that must report a missing row; delete succeeds silently. func (r sqlRepository) deleteByID(ctx context.Context, id string) error { count, err := r.executeSQL(ctx, Delete(r.tableName).Where(Eq{"id": id})) if err != nil { return err } if count == 0 { return model.ErrNotFound } return nil } func (r sqlRepository) logSQL(ctx context.Context, sql string, args dbx.Params, err error, rowsAffected int64, start time.Time) { elapsed := time.Since(start) fields := []any{ctx, "SQL: `" + sql + "`", "args", args, "rowsAffected", rowsAffected, "elapsedTime", elapsed} if err == nil || errors.Is(err, context.Canceled) { log.Trace(append(fields, err)...) return } // The result codes separate errors that share a message, notably SQLITE_BUSY from // SQLITE_BUSY_SNAPSHOT, which no busy_timeout can retry. if code, extended, ok := db.ErrorCodes(err); ok { fields = append(fields, "sqliteCode", code, "sqliteExtended", extended) } if db.IsBusy(err) && hasBusyRetry(ctx) { log.Warn(append(fields, err)...) return } log.Error(append(fields, err)...) }