package database_test import ( "context" "database/sql" "errors" "os" "slices" "testing" "time" "sneak.berlin/go/vaultik/internal/database" "sneak.berlin/go/vaultik/internal/types" ) // errTestRollback is the sentinel returned from transaction bodies to // force a rollback in tests. var errTestRollback = errors.New("test rollback") func TestFileRepository(t *testing.T) { t.Parallel() db, cleanup := setupTestDB(t) defer cleanup() ctx := context.Background() repo := database.NewFileRepository(db) // Test Create file := &database.File{ Path: testFileTxt, MTime: time.Now().Truncate(time.Second), Size: 1024, Mode: 0644, UID: 1000, GID: 1000, LinkTarget: "", } err := repo.Create(ctx, nil, file) if err != nil { t.Fatalf("failed to create file: %v", err) } // Test GetByPath retrieved, err := repo.GetByPath(ctx, file.Path.String()) if err != nil { t.Fatalf("failed to get file: %v", err) } if retrieved == nil { t.Fatal("expected file, got nil") } if retrieved.Path != file.Path { t.Errorf("path mismatch: got %s, want %s", retrieved.Path, file.Path) } if !retrieved.MTime.Equal(file.MTime) { t.Errorf("mtime mismatch: got %v, want %v", retrieved.MTime, file.MTime) } if retrieved.Size != file.Size { t.Errorf("size mismatch: got %d, want %d", retrieved.Size, file.Size) } if retrieved.Mode != file.Mode { t.Errorf("mode mismatch: got %o, want %o", retrieved.Mode, file.Mode) } // Test Update (upsert) file.Size = 2048 file.MTime = time.Now().Truncate(time.Second) err = repo.Create(ctx, nil, file) if err != nil { t.Fatalf("failed to update file: %v", err) } retrieved, err = repo.GetByPath(ctx, file.Path.String()) if err != nil { t.Fatalf("failed to get updated file: %v", err) } if retrieved.Size != 2048 { t.Errorf("size not updated: got %d, want %d", retrieved.Size, 2048) } } func TestFileRepositoryListDelete(t *testing.T) { t.Parallel() db, cleanup := setupTestDB(t) defer cleanup() ctx := context.Background() repo := database.NewFileRepository(db) file := &database.File{ Path: testFileTxt, MTime: time.Now().Truncate(time.Second), Size: 1024, Mode: 0644, UID: 1000, GID: 1000, } err := repo.Create(ctx, nil, file) if err != nil { t.Fatalf("failed to create file: %v", err) } // Test ListModifiedSince files, err := repo.ListModifiedSince(ctx, time.Now().Add(-1*time.Hour)) if err != nil { t.Fatalf("failed to list files: %v", err) } if len(files) != 1 { t.Errorf("expected 1 file, got %d", len(files)) } // Test Delete err = repo.Delete(ctx, nil, file.Path.String()) if err != nil { t.Fatalf("failed to delete file: %v", err) } retrieved, err := repo.GetByPath(ctx, file.Path.String()) if err != nil { t.Fatalf("error getting deleted file: %v", err) } if retrieved != nil { t.Error("expected nil for deleted file") } } func TestFileRepositoryListUnderPath(t *testing.T) { t.Parallel() db, cleanup := setupTestDB(t) defer cleanup() ctx := context.Background() repo := database.NewFileRepository(db) const ( docDir = "/home/u/doc" docFile = "/home/u/doc/a.txt" ) // In path order, so the root case can expect all of them as listed. paths := []string{ "/home/u/50%/x.txt", "/home/u/50percent/y.txt", "/home/u/DOC/c.txt", "/home/u/a_b/x.txt", "/home/u/axb/y.txt", docDir, "/home/u/doc.txt.bak", docFile, "/home/u/doc/sub/b.txt", "/home/u/doc2/b.txt", } for _, path := range paths { err := repo.Create(ctx, nil, &database.File{ Path: types.FilePath(path), MTime: time.Now().Truncate(time.Second), Mode: 0644, }) if err != nil { t.Fatalf("failed to create %s: %v", path, err) } } docTree := []string{docDir, docFile, "/home/u/doc/sub/b.txt"} tests := []struct { name string path string want []string }{ {"directory", docDir, docTree}, {"directory with trailing slash", docDir + "/", docTree}, {"directory differing only in case", "/home/u/DOC", []string{"/home/u/DOC/c.txt"}}, {"file", docFile, []string{docFile}}, {"underscore is literal", "/home/u/a_b", []string{"/home/u/a_b/x.txt"}}, {"percent is literal", "/home/u/50%", []string{"/home/u/50%/x.txt"}}, {"root", "/", paths}, } for _, tt := range tests { files, err := repo.ListUnderPath(ctx, tt.path) if err != nil { t.Fatalf("%s: failed to list files: %v", tt.name, err) } got := make([]string, 0, len(files)) for _, f := range files { got = append(got, f.Path.String()) } if !slices.Equal(got, tt.want) { t.Errorf("%s: listing %q got %q, want %q", tt.name, tt.path, got, tt.want) } } } func TestFileRepositorySymlink(t *testing.T) { t.Parallel() db, cleanup := setupTestDB(t) defer cleanup() ctx := context.Background() repo := database.NewFileRepository(db) // Test symlink symlink := &database.File{ Path: "/test/link", MTime: time.Now().Truncate(time.Second), Size: 0, Mode: uint32(0777 | os.ModeSymlink), UID: 1000, GID: 1000, LinkTarget: "/test/target", } err := repo.Create(ctx, nil, symlink) if err != nil { t.Fatalf("failed to create symlink: %v", err) } retrieved, err := repo.GetByPath(ctx, symlink.Path.String()) if err != nil { t.Fatalf("failed to get symlink: %v", err) } if !retrieved.IsSymlink() { t.Error("expected IsSymlink() to be true") } if retrieved.LinkTarget != symlink.LinkTarget { t.Errorf("link target mismatch: got %s, want %s", retrieved.LinkTarget, symlink.LinkTarget) } } // An mtime after 2262 or before 1678 does not fit in int64 nanoseconds // since the epoch, and must still come back from the database unchanged. func TestFileRepositoryMTimeOutsideInt64NanosecondRange(t *testing.T) { t.Parallel() db, cleanup := setupTestDB(t) defer cleanup() ctx := context.Background() repo := database.NewFileRepository(db) mtimes := []time.Time{ time.Date(2300, time.January, 1, 0, 0, 0, 123456789, time.UTC), time.Date(1601, time.January, 1, 0, 0, 0, 987654321, time.UTC), } for _, mtime := range mtimes { created := &database.File{ Path: types.FilePath("/created-" + mtime.Format(time.RFC3339Nano)), MTime: mtime, } err := repo.Create(ctx, nil, created) if err != nil { t.Fatalf("failed to create file: %v", err) } batched := &database.File{ ID: types.NewFileID(), Path: types.FilePath("/batched-" + mtime.Format(time.RFC3339Nano)), MTime: mtime, } err = repo.CreateBatch(ctx, nil, []*database.File{batched}) if err != nil { t.Fatalf("failed to batch create file: %v", err) } for _, path := range []types.FilePath{created.Path, batched.Path} { retrieved, err := repo.GetByPath(ctx, path.String()) if err != nil { t.Fatalf("failed to get file: %v", err) } if !retrieved.MTime.Equal(mtime) { t.Errorf("%s: mtime got %v, want %v", path, retrieved.MTime, mtime) } } } } func TestFileRepositoryTransaction(t *testing.T) { t.Parallel() db, cleanup := setupTestDB(t) defer cleanup() ctx := context.Background() repos := database.NewRepositories(db) // Test transaction rollback err := repos.WithTx(ctx, func(ctx context.Context, tx *sql.Tx) error { file := &database.File{ Path: testTxFile, MTime: time.Now().Truncate(time.Second), Size: 1024, Mode: 0644, UID: 1000, GID: 1000, } err := repos.Files.Create(ctx, tx, file) if err != nil { return err } // Return error to trigger rollback return errTestRollback }) if !errors.Is(err, errTestRollback) { t.Fatalf("expected rollback error, got: %v", err) } // Verify file was not created retrieved, err := repos.Files.GetByPath(ctx, testTxFile) if err != nil { t.Fatalf("error checking for file: %v", err) } if retrieved != nil { t.Error("file should not exist after rollback") } }