The ten sqlclosecheck findings were not leaks: every one of these queries already deferred a close through the package-local CloseRows helper. sqlclosecheck only recognises a Close call on the rows value in the function that produced it (directly deferred, or inside a deferred closure), so a call that hands rows to a helper reads as unhandled. Rather than keep a helper the linter cannot see through, drop CloseRows and defer a closure that calls rows.Close() directly at each of the eighteen call sites, keeping the existing fatal-on-close-error behaviour byte for byte. The close still runs exactly once, on function exit, after the rows have been read. Fatalf stays; it is still used by the transaction helpers.
442 lines
10 KiB
Go
442 lines
10 KiB
Go
package database
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
"sneak.berlin/go/vaultik/internal/log"
|
|
"sneak.berlin/go/vaultik/internal/types"
|
|
)
|
|
|
|
// FileRepository provides access to the files table, which stores file
|
|
// metadata (path, times, permissions, ownership, symlink targets).
|
|
type FileRepository struct {
|
|
db *DB
|
|
}
|
|
|
|
// NewFileRepository creates a FileRepository backed by db.
|
|
func NewFileRepository(db *DB) *FileRepository {
|
|
return &FileRepository{db: db}
|
|
}
|
|
|
|
// Create inserts or updates a file row (upsert on path), using tx when
|
|
// non-nil. The file's ID is generated when zero and updated from the
|
|
// database's RETURNING clause.
|
|
func (r *FileRepository) Create(ctx context.Context, tx *sql.Tx, file *File) error {
|
|
// Generate UUID if not provided
|
|
if file.ID.IsZero() {
|
|
file.ID = types.NewFileID()
|
|
}
|
|
|
|
query := `
|
|
INSERT INTO files (id, path, source_path, mtime, size, mode, uid, gid, link_target)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(path) DO UPDATE SET
|
|
source_path = excluded.source_path,
|
|
mtime = excluded.mtime,
|
|
size = excluded.size,
|
|
mode = excluded.mode,
|
|
uid = excluded.uid,
|
|
gid = excluded.gid,
|
|
link_target = excluded.link_target
|
|
RETURNING id
|
|
`
|
|
|
|
var (
|
|
idStr string
|
|
err error
|
|
)
|
|
|
|
if tx != nil {
|
|
LogSQL("Execute", query,
|
|
file.ID.String(), file.Path.String(), file.SourcePath.String(),
|
|
file.MTime.Unix(), file.Size, file.Mode, file.UID, file.GID,
|
|
file.LinkTarget.String())
|
|
err = tx.QueryRowContext(ctx, query,
|
|
file.ID.String(), file.Path.String(), file.SourcePath.String(),
|
|
file.MTime.Unix(), file.Size, file.Mode, file.UID, file.GID,
|
|
file.LinkTarget.String()).Scan(&idStr)
|
|
} else {
|
|
err = r.db.QueryRowWithLog(ctx, query,
|
|
file.ID.String(), file.Path.String(), file.SourcePath.String(),
|
|
file.MTime.Unix(), file.Size, file.Mode, file.UID, file.GID,
|
|
file.LinkTarget.String()).Scan(&idStr)
|
|
}
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("inserting file: %w", err)
|
|
}
|
|
|
|
// Parse the returned ID
|
|
file.ID, err = types.ParseFileID(idStr)
|
|
if err != nil {
|
|
return fmt.Errorf("parsing file ID: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// GetByPath returns the file at the given path, or nil if the path is not
|
|
// in the index.
|
|
func (r *FileRepository) GetByPath(ctx context.Context, path string) (*File, error) {
|
|
query := `
|
|
SELECT id, path, source_path, mtime, size, mode, uid, gid, link_target
|
|
FROM files
|
|
WHERE path = ?
|
|
`
|
|
|
|
file, err := r.scanFile(r.db.conn.QueryRowContext(ctx, query, path))
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil //nolint:nilnil // nil,nil signals not-found; callers check nil
|
|
}
|
|
|
|
if err != nil {
|
|
return nil, fmt.Errorf("querying file: %w", err)
|
|
}
|
|
|
|
return file, nil
|
|
}
|
|
|
|
// GetByID retrieves a file by its UUID
|
|
func (r *FileRepository) GetByID(ctx context.Context, id types.FileID) (*File, error) {
|
|
query := `
|
|
SELECT id, path, source_path, mtime, size, mode, uid, gid, link_target
|
|
FROM files
|
|
WHERE id = ?
|
|
`
|
|
|
|
file, err := r.scanFile(r.db.conn.QueryRowContext(ctx, query, id.String()))
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil //nolint:nilnil // nil,nil signals not-found; callers check nil
|
|
}
|
|
|
|
if err != nil {
|
|
return nil, fmt.Errorf("querying file: %w", err)
|
|
}
|
|
|
|
return file, nil
|
|
}
|
|
|
|
// GetByPathTx returns the file at the given path within a transaction, or
|
|
// nil if the path is not in the index.
|
|
func (r *FileRepository) GetByPathTx(
|
|
ctx context.Context, tx *sql.Tx, path string,
|
|
) (*File, error) {
|
|
query := `
|
|
SELECT id, path, source_path, mtime, size, mode, uid, gid, link_target
|
|
FROM files
|
|
WHERE path = ?
|
|
`
|
|
|
|
LogSQL("GetByPathTx QueryRowContext", query, path)
|
|
file, err := r.scanFile(tx.QueryRowContext(ctx, query, path))
|
|
LogSQL("GetByPathTx Scan complete", query, path)
|
|
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil //nolint:nilnil // nil,nil signals not-found; callers check nil
|
|
}
|
|
|
|
if err != nil {
|
|
return nil, fmt.Errorf("querying file: %w", err)
|
|
}
|
|
|
|
return file, nil
|
|
}
|
|
|
|
// fileRowScanner abstracts *sql.Row and *sql.Rows for scanning a file row.
|
|
type fileRowScanner interface {
|
|
Scan(dest ...any) error
|
|
}
|
|
|
|
// ListModifiedSince returns all files whose recorded mtime is at or after
|
|
// since, ordered by path.
|
|
func (r *FileRepository) ListModifiedSince(
|
|
ctx context.Context, since time.Time,
|
|
) ([]*File, error) {
|
|
query := `
|
|
SELECT id, path, source_path, mtime, size, mode, uid, gid, link_target
|
|
FROM files
|
|
WHERE mtime >= ?
|
|
ORDER BY path
|
|
`
|
|
|
|
rows, err := r.db.conn.QueryContext(ctx, query, since.Unix())
|
|
if err != nil {
|
|
return nil, fmt.Errorf("querying files: %w", err)
|
|
}
|
|
|
|
defer func() {
|
|
err := rows.Close()
|
|
if err != nil {
|
|
Fatalf("failed to close rows: %v", err)
|
|
}
|
|
}()
|
|
|
|
var files []*File
|
|
|
|
for rows.Next() {
|
|
file, err := r.scanFileRows(rows)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("scanning file: %w", err)
|
|
}
|
|
|
|
files = append(files, file)
|
|
}
|
|
|
|
return files, rows.Err()
|
|
}
|
|
|
|
// Delete removes the file row at the given path, using tx when non-nil.
|
|
func (r *FileRepository) Delete(ctx context.Context, tx *sql.Tx, path string) error {
|
|
query := `DELETE FROM files WHERE path = ?`
|
|
|
|
var err error
|
|
if tx != nil {
|
|
_, err = tx.ExecContext(ctx, query, path)
|
|
} else {
|
|
_, err = r.db.ExecWithLog(ctx, query, path)
|
|
}
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("deleting file: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteByID deletes a file by its UUID
|
|
func (r *FileRepository) DeleteByID(
|
|
ctx context.Context, tx *sql.Tx, id types.FileID,
|
|
) error {
|
|
query := `DELETE FROM files WHERE id = ?`
|
|
|
|
var err error
|
|
if tx != nil {
|
|
_, err = tx.ExecContext(ctx, query, id.String())
|
|
} else {
|
|
_, err = r.db.ExecWithLog(ctx, query, id.String())
|
|
}
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("deleting file: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// ListByPrefix returns all files whose path starts with prefix, ordered by
|
|
// path.
|
|
func (r *FileRepository) ListByPrefix(
|
|
ctx context.Context, prefix string,
|
|
) ([]*File, error) {
|
|
query := `
|
|
SELECT id, path, source_path, mtime, size, mode, uid, gid, link_target
|
|
FROM files
|
|
WHERE path LIKE ? || '%'
|
|
ORDER BY path
|
|
`
|
|
|
|
rows, err := r.db.conn.QueryContext(ctx, query, prefix)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("querying files: %w", err)
|
|
}
|
|
|
|
defer func() {
|
|
err := rows.Close()
|
|
if err != nil {
|
|
Fatalf("failed to close rows: %v", err)
|
|
}
|
|
}()
|
|
|
|
var files []*File
|
|
|
|
for rows.Next() {
|
|
file, err := r.scanFileRows(rows)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("scanning file: %w", err)
|
|
}
|
|
|
|
files = append(files, file)
|
|
}
|
|
|
|
return files, rows.Err()
|
|
}
|
|
|
|
// ListAll returns all files in the database
|
|
func (r *FileRepository) ListAll(ctx context.Context) ([]*File, error) {
|
|
query := `
|
|
SELECT id, path, source_path, mtime, size, mode, uid, gid, link_target
|
|
FROM files
|
|
ORDER BY path
|
|
`
|
|
|
|
rows, err := r.db.conn.QueryContext(ctx, query)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("querying files: %w", err)
|
|
}
|
|
|
|
defer func() {
|
|
err := rows.Close()
|
|
if err != nil {
|
|
Fatalf("failed to close rows: %v", err)
|
|
}
|
|
}()
|
|
|
|
var files []*File
|
|
|
|
for rows.Next() {
|
|
file, err := r.scanFileRows(rows)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("scanning file: %w", err)
|
|
}
|
|
|
|
files = append(files, file)
|
|
}
|
|
|
|
return files, rows.Err()
|
|
}
|
|
|
|
// CreateBatch inserts or updates multiple files in a single statement for efficiency.
|
|
// File IDs must be pre-generated before calling this method.
|
|
func (r *FileRepository) CreateBatch(
|
|
ctx context.Context, tx *sql.Tx, files []*File,
|
|
) error {
|
|
if len(files) == 0 {
|
|
return nil
|
|
}
|
|
|
|
// Each files row binds this many SQL variables.
|
|
const fileCols = 9
|
|
|
|
// Batch at 100 rows to be safe with SQLite's variable limit.
|
|
const batchSize = 100
|
|
|
|
for i := 0; i < len(files); i += batchSize {
|
|
end := min(i+batchSize, len(files))
|
|
|
|
batch := files[i:end]
|
|
|
|
query := `INSERT INTO files
|
|
(id, path, source_path, mtime, size, mode, uid, gid, link_target)
|
|
VALUES `
|
|
|
|
args := make([]any, 0, len(batch)*fileCols)
|
|
|
|
var querySb325 strings.Builder
|
|
|
|
for j, f := range batch {
|
|
if j > 0 {
|
|
querySb325.WriteString(", ")
|
|
}
|
|
|
|
querySb325.WriteString("(?, ?, ?, ?, ?, ?, ?, ?, ?)")
|
|
|
|
args = append(args,
|
|
f.ID.String(), f.Path.String(), f.SourcePath.String(),
|
|
f.MTime.Unix(), f.Size, f.Mode, f.UID, f.GID,
|
|
f.LinkTarget.String())
|
|
}
|
|
|
|
query += querySb325.String() //nolint:gosec // G202: appends "?" placeholders only
|
|
|
|
query += ` ON CONFLICT(path) DO UPDATE SET
|
|
source_path = excluded.source_path,
|
|
mtime = excluded.mtime,
|
|
size = excluded.size,
|
|
mode = excluded.mode,
|
|
uid = excluded.uid,
|
|
gid = excluded.gid,
|
|
link_target = excluded.link_target`
|
|
|
|
var err error
|
|
if tx != nil {
|
|
_, err = tx.ExecContext(ctx, query, args...)
|
|
} else {
|
|
_, err = r.db.ExecWithLog(ctx, query, args...)
|
|
}
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("batch inserting files: %w", err)
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteOrphaned deletes files that are not referenced by any snapshot
|
|
func (r *FileRepository) DeleteOrphaned(ctx context.Context) error {
|
|
query := `
|
|
DELETE FROM files
|
|
WHERE NOT EXISTS (
|
|
SELECT 1 FROM snapshot_files
|
|
WHERE snapshot_files.file_id = files.id
|
|
)
|
|
`
|
|
|
|
result, err := r.db.ExecWithLog(ctx, query)
|
|
if err != nil {
|
|
return fmt.Errorf("deleting orphaned files: %w", err)
|
|
}
|
|
|
|
rowsAffected, _ := result.RowsAffected()
|
|
if rowsAffected > 0 {
|
|
log.Debug("Deleted orphaned files", "count", rowsAffected)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// scanFile is a helper that scans a single file row
|
|
func (r *FileRepository) scanFile(row *sql.Row) (*File, error) {
|
|
return r.scanFileFrom(row)
|
|
}
|
|
|
|
// scanFileRows is a helper that scans a file row from rows iterator
|
|
func (r *FileRepository) scanFileRows(rows *sql.Rows) (*File, error) {
|
|
return r.scanFileFrom(rows)
|
|
}
|
|
|
|
// scanFileFrom scans one file row from any row scanner.
|
|
func (r *FileRepository) scanFileFrom(row fileRowScanner) (*File, error) {
|
|
var (
|
|
file File
|
|
idStr, pathStr, sourcePathStr string
|
|
mtimeUnix int64
|
|
linkTarget sql.NullString
|
|
)
|
|
|
|
err := row.Scan(
|
|
&idStr,
|
|
&pathStr,
|
|
&sourcePathStr,
|
|
&mtimeUnix,
|
|
&file.Size,
|
|
&file.Mode,
|
|
&file.UID,
|
|
&file.GID,
|
|
&linkTarget,
|
|
)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
file.ID, err = types.ParseFileID(idStr)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("parsing file ID: %w", err)
|
|
}
|
|
|
|
file.Path = types.FilePath(pathStr)
|
|
file.SourcePath = types.SourcePath(sourcePathStr)
|
|
|
|
file.MTime = time.Unix(mtimeUnix, 0).UTC()
|
|
if linkTarget.Valid {
|
|
file.LinkTarget = types.FilePath(linkTarget.String)
|
|
}
|
|
|
|
return &file, nil
|
|
}
|