package database_test import ( "context" "database/sql" "errors" "os" "testing" "time" "sneak.berlin/go/vaultik/internal/database" ) // 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 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) } } 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") } }