primeiro commit
This commit is contained in:
@@ -0,0 +1,94 @@
|
||||
// Package daemon coordinates periodic sync cycles and change hooks.
|
||||
package daemon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"s3watch/internal/syncer"
|
||||
)
|
||||
|
||||
// Config controls daemon scheduling behavior.
|
||||
type Config struct {
|
||||
Interval time.Duration
|
||||
Once bool
|
||||
}
|
||||
|
||||
// Syncer synchronizes remote state into local state.
|
||||
type Syncer interface {
|
||||
Sync(context.Context) (syncer.Result, error)
|
||||
}
|
||||
|
||||
// HookRunner runs after a sync cycle applies changes.
|
||||
type HookRunner interface {
|
||||
Run(context.Context, syncer.Result) error
|
||||
}
|
||||
|
||||
// Run executes sync cycles until the context is canceled.
|
||||
func Run(ctx context.Context, cfg Config, syncService Syncer, runner HookRunner, logger *slog.Logger) error {
|
||||
if cfg.Interval <= 0 {
|
||||
return fmt.Errorf("interval must be greater than zero")
|
||||
}
|
||||
if syncService == nil {
|
||||
return fmt.Errorf("syncer is required")
|
||||
}
|
||||
if logger == nil {
|
||||
logger = slog.Default()
|
||||
}
|
||||
|
||||
for {
|
||||
if err := runCycle(ctx, syncService, runner, logger, cfg.Once); err != nil {
|
||||
if cfg.Once {
|
||||
return err
|
||||
}
|
||||
logger.Error("sync cycle failed", "error", err)
|
||||
}
|
||||
|
||||
if cfg.Once {
|
||||
return nil
|
||||
}
|
||||
|
||||
timer := time.NewTimer(cfg.Interval)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if !timer.Stop() {
|
||||
select {
|
||||
case <-timer.C:
|
||||
default:
|
||||
}
|
||||
}
|
||||
return nil
|
||||
case <-timer.C:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func runCycle(ctx context.Context, syncService Syncer, runner HookRunner, logger *slog.Logger, strict bool) error {
|
||||
result, err := syncService.Sync(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("syncing: %w", err)
|
||||
}
|
||||
|
||||
logger.Info("sync cycle completed",
|
||||
"downloaded", result.Downloaded,
|
||||
"uploaded", result.Uploaded,
|
||||
"updated", result.Updated,
|
||||
"deleted", result.Deleted,
|
||||
"unchanged", result.Unchanged,
|
||||
)
|
||||
|
||||
if !result.Changed() || runner == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := runner.Run(ctx, result); err != nil {
|
||||
if strict {
|
||||
return fmt.Errorf("running hook: %w", err)
|
||||
}
|
||||
logger.Error("hook failed", "error", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package daemon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"s3watch/internal/syncer"
|
||||
)
|
||||
|
||||
func TestRunOnceRunsHookOnlyWhenChanged(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
result syncer.Result
|
||||
wantCalls int
|
||||
}{
|
||||
{
|
||||
name: "changed",
|
||||
result: syncer.Result{Downloaded: 1},
|
||||
wantCalls: 1,
|
||||
},
|
||||
{
|
||||
name: "unchanged",
|
||||
result: syncer.Result{Unchanged: 1},
|
||||
wantCalls: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
hook := &fakeHook{}
|
||||
err := Run(context.Background(), Config{Interval: time.Second, Once: true}, fakeSyncer{result: tt.result}, hook, slog.New(slog.DiscardHandler))
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v", err)
|
||||
}
|
||||
if hook.calls != tt.wantCalls {
|
||||
t.Fatalf("hook calls = %d, want %d", hook.calls, tt.wantCalls)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunOnceReturnsHookError(t *testing.T) {
|
||||
wantErr := errors.New("hook failed")
|
||||
err := Run(
|
||||
context.Background(),
|
||||
Config{Interval: time.Second, Once: true},
|
||||
fakeSyncer{result: syncer.Result{Updated: 1}},
|
||||
&fakeHook{err: wantErr},
|
||||
slog.New(slog.DiscardHandler),
|
||||
)
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("Run() error = %v, want %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeSyncer struct {
|
||||
result syncer.Result
|
||||
err error
|
||||
}
|
||||
|
||||
func (f fakeSyncer) Sync(context.Context) (syncer.Result, error) {
|
||||
return f.result, f.err
|
||||
}
|
||||
|
||||
type fakeHook struct {
|
||||
calls int
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeHook) Run(context.Context, syncer.Result) error {
|
||||
f.calls++
|
||||
return f.err
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
// Package hook runs post-sync hook scripts.
|
||||
package hook
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strconv"
|
||||
|
||||
"s3watch/internal/syncer"
|
||||
)
|
||||
|
||||
// Runner executes a configured hook script after changes.
|
||||
type Runner struct {
|
||||
path string
|
||||
dir string
|
||||
}
|
||||
|
||||
// NewRunner creates a hook runner for an executable script path.
|
||||
func NewRunner(path string, dir string) Runner {
|
||||
return Runner{path: path, dir: dir}
|
||||
}
|
||||
|
||||
// Run executes the hook and passes change counts through environment variables.
|
||||
func (r Runner) Run(ctx context.Context, result syncer.Result) error {
|
||||
if r.path == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
cmd := exec.CommandContext(ctx, r.path)
|
||||
cmd.Dir = r.dir
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
cmd.Env = append(os.Environ(),
|
||||
"S3WATCH_CHANGED="+strconv.FormatBool(result.Changed()),
|
||||
"S3WATCH_DOWNLOADED="+strconv.Itoa(result.Downloaded),
|
||||
"S3WATCH_UPLOADED="+strconv.Itoa(result.Uploaded),
|
||||
"S3WATCH_UPDATED="+strconv.Itoa(result.Updated),
|
||||
"S3WATCH_DELETED="+strconv.Itoa(result.Deleted),
|
||||
"S3WATCH_UNCHANGED="+strconv.Itoa(result.Unchanged),
|
||||
)
|
||||
|
||||
if err := cmd.Run(); err != nil {
|
||||
return fmt.Errorf("executing %q: %w", r.path, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package hook
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"s3watch/internal/syncer"
|
||||
)
|
||||
|
||||
func TestRunnerRunPassesChangeEnvironment(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
script := filepath.Join(t.TempDir(), "hook.sh")
|
||||
if err := os.WriteFile(script, []byte("#!/bin/sh\nprintf '%s,%s,%s,%s,%s,%s' \"$S3WATCH_CHANGED\" \"$S3WATCH_DOWNLOADED\" \"$S3WATCH_UPLOADED\" \"$S3WATCH_UPDATED\" \"$S3WATCH_DELETED\" \"$S3WATCH_UNCHANGED\" > hook.out\n"), 0o755); err != nil {
|
||||
t.Fatalf("writing hook script: %v", err)
|
||||
}
|
||||
|
||||
result := syncer.Result{Downloaded: 2, Uploaded: 3, Updated: 4, Deleted: 5, Unchanged: 6}
|
||||
if err := NewRunner(script, dir).Run(context.Background(), result); err != nil {
|
||||
t.Fatalf("Run() error = %v", err)
|
||||
}
|
||||
|
||||
output, err := os.ReadFile(filepath.Join(dir, "hook.out"))
|
||||
if err != nil {
|
||||
t.Fatalf("reading hook output: %v", err)
|
||||
}
|
||||
if got, want := strings.TrimSpace(string(output)), "true,2,3,4,5,6"; got != want {
|
||||
t.Fatalf("hook output = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,592 @@
|
||||
// Package syncer mirrors files between S3 objects and a local directory.
|
||||
package syncer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3/types"
|
||||
)
|
||||
|
||||
const (
|
||||
stateDirName = ".s3watch"
|
||||
stateFileName = "state.json"
|
||||
tempFileGlob = ".s3watch-*"
|
||||
)
|
||||
|
||||
// Config controls S3/local sync behavior.
|
||||
type Config struct {
|
||||
Bucket string
|
||||
Prefix string
|
||||
Dir string
|
||||
Prune bool
|
||||
}
|
||||
|
||||
// Result describes changes applied during one sync cycle.
|
||||
type Result struct {
|
||||
Downloaded int
|
||||
Uploaded int
|
||||
Updated int
|
||||
Deleted int
|
||||
Unchanged int
|
||||
}
|
||||
|
||||
// Changed reports whether a sync cycle changed either side.
|
||||
func (r Result) Changed() bool {
|
||||
return r.Downloaded > 0 || r.Uploaded > 0 || r.Updated > 0 || r.Deleted > 0
|
||||
}
|
||||
|
||||
// S3API is the subset of S3 operations required by Syncer.
|
||||
type S3API interface {
|
||||
ListObjectsV2(context.Context, *s3.ListObjectsV2Input, ...func(*s3.Options)) (*s3.ListObjectsV2Output, error)
|
||||
GetObject(context.Context, *s3.GetObjectInput, ...func(*s3.Options)) (*s3.GetObjectOutput, error)
|
||||
PutObject(context.Context, *s3.PutObjectInput, ...func(*s3.Options)) (*s3.PutObjectOutput, error)
|
||||
HeadObject(context.Context, *s3.HeadObjectInput, ...func(*s3.Options)) (*s3.HeadObjectOutput, error)
|
||||
DeleteObject(context.Context, *s3.DeleteObjectInput, ...func(*s3.Options)) (*s3.DeleteObjectOutput, error)
|
||||
}
|
||||
|
||||
// Syncer synchronizes S3 objects and local files.
|
||||
type Syncer struct {
|
||||
client S3API
|
||||
cfg Config
|
||||
}
|
||||
|
||||
// New creates a Syncer.
|
||||
func New(client S3API, cfg Config) *Syncer {
|
||||
return &Syncer{
|
||||
client: client,
|
||||
cfg: cfg,
|
||||
}
|
||||
}
|
||||
|
||||
// Sync reconciles configured S3 objects and the configured local directory.
|
||||
func (s *Syncer) Sync(ctx context.Context) (Result, error) {
|
||||
if s.client == nil {
|
||||
return Result{}, fmt.Errorf("s3 client is required")
|
||||
}
|
||||
if s.cfg.Bucket == "" {
|
||||
return Result{}, fmt.Errorf("bucket is required")
|
||||
}
|
||||
if s.cfg.Dir == "" {
|
||||
return Result{}, fmt.Errorf("dir is required")
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(s.cfg.Dir, 0o755); err != nil {
|
||||
return Result{}, fmt.Errorf("creating sync directory: %w", err)
|
||||
}
|
||||
|
||||
previous, err := loadState(s.statePath())
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
remoteFiles, err := s.listRemoteFiles(ctx)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
localFiles, err := scanLocalFiles(s.cfg.Dir)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
result := Result{}
|
||||
for _, rel := range unionKeys(localFiles, remoteFiles, previous.Files) {
|
||||
local, localOK := localFiles[rel]
|
||||
remote, remoteOK := remoteFiles[rel]
|
||||
prior, priorOK := previous.Files[rel]
|
||||
|
||||
switch {
|
||||
case localOK && remoteOK:
|
||||
changed, err := s.syncExisting(ctx, rel, local, remote, localFiles, remoteFiles)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
if changed {
|
||||
result.Updated++
|
||||
} else {
|
||||
result.Unchanged++
|
||||
}
|
||||
case localOK:
|
||||
if s.shouldDeleteLocal(local, prior, priorOK) {
|
||||
if err := os.Remove(local.Path); err != nil {
|
||||
return Result{}, fmt.Errorf("deleting local %s: %w", local.Path, err)
|
||||
}
|
||||
delete(localFiles, rel)
|
||||
result.Deleted++
|
||||
continue
|
||||
}
|
||||
if err := s.uploadFile(ctx, rel, local, localFiles, remoteFiles); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
result.Uploaded++
|
||||
case remoteOK:
|
||||
if s.shouldDeleteRemote(remote, prior, priorOK) {
|
||||
if err := s.deleteRemote(ctx, rel); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
delete(remoteFiles, rel)
|
||||
result.Deleted++
|
||||
continue
|
||||
}
|
||||
if err := s.downloadFile(ctx, rel, remote, localFiles); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
result.Downloaded++
|
||||
}
|
||||
}
|
||||
|
||||
if err := saveState(s.statePath(), buildState(localFiles, remoteFiles)); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *Syncer) syncExisting(ctx context.Context, rel string, local localFile, remote remoteFile, localFiles map[string]localFile, remoteFiles map[string]remoteFile) (bool, error) {
|
||||
if local.Snapshot.Equal(remote.Snapshot) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if local.Snapshot.ModTimeUnix >= remote.Snapshot.ModTimeUnix {
|
||||
if err := s.uploadFile(ctx, rel, local, localFiles, remoteFiles); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
if err := s.downloadFile(ctx, rel, remote, localFiles); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *Syncer) shouldDeleteLocal(local localFile, prior stateEntry, priorOK bool) bool {
|
||||
return s.cfg.Prune && priorOK && prior.Remote != nil && prior.Local != nil && local.Snapshot.Equal(*prior.Local)
|
||||
}
|
||||
|
||||
func (s *Syncer) shouldDeleteRemote(remote remoteFile, prior stateEntry, priorOK bool) bool {
|
||||
return s.cfg.Prune && priorOK && prior.Local != nil && prior.Remote != nil && remote.Snapshot.Equal(*prior.Remote)
|
||||
}
|
||||
|
||||
func (s *Syncer) listRemoteFiles(ctx context.Context) (map[string]remoteFile, error) {
|
||||
files := make(map[string]remoteFile)
|
||||
var token *string
|
||||
|
||||
prefix := normalizePrefix(s.cfg.Prefix)
|
||||
for {
|
||||
output, err := s.client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{
|
||||
Bucket: aws.String(s.cfg.Bucket),
|
||||
Prefix: aws.String(prefix),
|
||||
ContinuationToken: token,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("listing s3://%s/%s: %w", s.cfg.Bucket, prefix, err)
|
||||
}
|
||||
|
||||
for _, object := range output.Contents {
|
||||
if object.Key == nil || strings.HasSuffix(*object.Key, "/") {
|
||||
continue
|
||||
}
|
||||
rel, err := relativeKey(*object.Key, s.cfg.Prefix)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := safeLocalPath(s.cfg.Dir, rel); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
files[rel] = remoteFile{
|
||||
Key: *object.Key,
|
||||
Snapshot: snapshotFromObject(object),
|
||||
}
|
||||
}
|
||||
|
||||
if output.IsTruncated == nil || !*output.IsTruncated {
|
||||
return files, nil
|
||||
}
|
||||
if output.NextContinuationToken == nil {
|
||||
return nil, fmt.Errorf("listing s3://%s/%s: truncated response missing continuation token", s.cfg.Bucket, prefix)
|
||||
}
|
||||
token = output.NextContinuationToken
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Syncer) downloadFile(ctx context.Context, rel string, remote remoteFile, localFiles map[string]localFile) error {
|
||||
localPath, err := safeLocalPath(s.cfg.Dir, rel)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(localPath), 0o755); err != nil {
|
||||
return fmt.Errorf("creating parent directory for %s: %w", localPath, err)
|
||||
}
|
||||
|
||||
output, err := s.client.GetObject(ctx, &s3.GetObjectInput{
|
||||
Bucket: aws.String(s.cfg.Bucket),
|
||||
Key: aws.String(remote.Key),
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("getting s3://%s/%s: %w", s.cfg.Bucket, remote.Key, err)
|
||||
}
|
||||
defer output.Body.Close()
|
||||
|
||||
temp, err := os.CreateTemp(filepath.Dir(localPath), tempFileGlob)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating temporary file for %s: %w", localPath, err)
|
||||
}
|
||||
tempPath := temp.Name()
|
||||
removeTemp := true
|
||||
defer func() {
|
||||
if removeTemp {
|
||||
_ = os.Remove(tempPath)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := io.Copy(temp, output.Body); err != nil {
|
||||
_ = temp.Close()
|
||||
return fmt.Errorf("writing temporary file for %s: %w", localPath, err)
|
||||
}
|
||||
if err := temp.Close(); err != nil {
|
||||
return fmt.Errorf("closing temporary file for %s: %w", localPath, err)
|
||||
}
|
||||
|
||||
if err := os.Chtimes(tempPath, unixTime(remote.Snapshot.ModTimeUnix), unixTime(remote.Snapshot.ModTimeUnix)); err != nil {
|
||||
return fmt.Errorf("setting timestamp on %s: %w", tempPath, err)
|
||||
}
|
||||
|
||||
if err := os.Rename(tempPath, localPath); err != nil {
|
||||
return fmt.Errorf("replacing %s: %w", localPath, err)
|
||||
}
|
||||
removeTemp = false
|
||||
|
||||
localFiles[rel] = localFile{
|
||||
Path: localPath,
|
||||
Snapshot: remote.Snapshot,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Syncer) uploadFile(ctx context.Context, rel string, local localFile, localFiles map[string]localFile, remoteFiles map[string]remoteFile) error {
|
||||
file, err := os.Open(local.Path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("opening local %s: %w", local.Path, err)
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
key := s.remoteKey(rel)
|
||||
if _, err := s.client.PutObject(ctx, &s3.PutObjectInput{
|
||||
Bucket: aws.String(s.cfg.Bucket),
|
||||
Key: aws.String(key),
|
||||
Body: file,
|
||||
ContentLength: aws.Int64(local.Snapshot.Size),
|
||||
}); err != nil {
|
||||
return fmt.Errorf("putting s3://%s/%s: %w", s.cfg.Bucket, key, err)
|
||||
}
|
||||
|
||||
head, err := s.client.HeadObject(ctx, &s3.HeadObjectInput{
|
||||
Bucket: aws.String(s.cfg.Bucket),
|
||||
Key: aws.String(key),
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("heading uploaded s3://%s/%s: %w", s.cfg.Bucket, key, err)
|
||||
}
|
||||
|
||||
remoteSnapshot := snapshotFromHead(head, local.Snapshot)
|
||||
if err := os.Chtimes(local.Path, unixTime(remoteSnapshot.ModTimeUnix), unixTime(remoteSnapshot.ModTimeUnix)); err != nil {
|
||||
return fmt.Errorf("setting timestamp on %s: %w", local.Path, err)
|
||||
}
|
||||
|
||||
localFiles[rel] = localFile{
|
||||
Path: local.Path,
|
||||
Snapshot: remoteSnapshot,
|
||||
}
|
||||
remoteFiles[rel] = remoteFile{
|
||||
Key: key,
|
||||
Snapshot: remoteSnapshot,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Syncer) deleteRemote(ctx context.Context, rel string) error {
|
||||
key := s.remoteKey(rel)
|
||||
if _, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{
|
||||
Bucket: aws.String(s.cfg.Bucket),
|
||||
Key: aws.String(key),
|
||||
}); err != nil {
|
||||
return fmt.Errorf("deleting s3://%s/%s: %w", s.cfg.Bucket, key, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Syncer) remoteKey(rel string) string {
|
||||
prefix := normalizePrefix(s.cfg.Prefix)
|
||||
return prefix + strings.TrimLeft(path.Clean(rel), "/")
|
||||
}
|
||||
|
||||
func (s *Syncer) statePath() string {
|
||||
return filepath.Join(s.cfg.Dir, stateDirName, stateFileName)
|
||||
}
|
||||
|
||||
type localFile struct {
|
||||
Path string
|
||||
Snapshot fileSnapshot
|
||||
}
|
||||
|
||||
type remoteFile struct {
|
||||
Key string
|
||||
Snapshot fileSnapshot
|
||||
}
|
||||
|
||||
type fileSnapshot struct {
|
||||
Size int64 `json:"size"`
|
||||
ModTimeUnix int64 `json:"mod_time_unix"`
|
||||
}
|
||||
|
||||
// Equal reports whether two snapshots describe the same file state.
|
||||
func (s fileSnapshot) Equal(other fileSnapshot) bool {
|
||||
return s.Size == other.Size && s.ModTimeUnix == other.ModTimeUnix
|
||||
}
|
||||
|
||||
type syncState struct {
|
||||
Version int `json:"version"`
|
||||
Files map[string]stateEntry `json:"files"`
|
||||
}
|
||||
|
||||
type stateEntry struct {
|
||||
Local *fileSnapshot `json:"local,omitempty"`
|
||||
Remote *fileSnapshot `json:"remote,omitempty"`
|
||||
}
|
||||
|
||||
func loadState(path string) (syncState, error) {
|
||||
state := syncState{
|
||||
Version: 1,
|
||||
Files: map[string]stateEntry{},
|
||||
}
|
||||
|
||||
file, err := os.Open(path)
|
||||
if os.IsNotExist(err) {
|
||||
return state, nil
|
||||
}
|
||||
if err != nil {
|
||||
return syncState{}, fmt.Errorf("opening sync state: %w", err)
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
if err := json.NewDecoder(file).Decode(&state); err != nil {
|
||||
return syncState{}, fmt.Errorf("decoding sync state: %w", err)
|
||||
}
|
||||
if state.Files == nil {
|
||||
state.Files = map[string]stateEntry{}
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func saveState(path string, state syncState) error {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return fmt.Errorf("creating sync state directory: %w", err)
|
||||
}
|
||||
|
||||
temp, err := os.CreateTemp(filepath.Dir(path), tempFileGlob)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating sync state temporary file: %w", err)
|
||||
}
|
||||
tempPath := temp.Name()
|
||||
removeTemp := true
|
||||
defer func() {
|
||||
if removeTemp {
|
||||
_ = os.Remove(tempPath)
|
||||
}
|
||||
}()
|
||||
|
||||
encoder := json.NewEncoder(temp)
|
||||
encoder.SetIndent("", " ")
|
||||
if err := encoder.Encode(state); err != nil {
|
||||
_ = temp.Close()
|
||||
return fmt.Errorf("encoding sync state: %w", err)
|
||||
}
|
||||
if err := temp.Close(); err != nil {
|
||||
return fmt.Errorf("closing sync state temporary file: %w", err)
|
||||
}
|
||||
if err := os.Rename(tempPath, path); err != nil {
|
||||
return fmt.Errorf("replacing sync state: %w", err)
|
||||
}
|
||||
removeTemp = false
|
||||
return nil
|
||||
}
|
||||
|
||||
func scanLocalFiles(root string) (map[string]localFile, error) {
|
||||
files := make(map[string]localFile)
|
||||
rootAbs, err := filepath.Abs(root)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("resolving root directory: %w", err)
|
||||
}
|
||||
stateDir := filepath.Join(rootAbs, stateDirName)
|
||||
|
||||
if err := filepath.WalkDir(rootAbs, func(current string, entry os.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("walking %s: %w", current, err)
|
||||
}
|
||||
if current == stateDir && entry.IsDir() {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
if entry.IsDir() {
|
||||
return nil
|
||||
}
|
||||
if entry.Type()&os.ModeSymlink != 0 {
|
||||
return nil
|
||||
}
|
||||
if strings.HasPrefix(entry.Name(), strings.TrimSuffix(tempFileGlob, "*")) {
|
||||
return nil
|
||||
}
|
||||
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
return fmt.Errorf("statting %s: %w", current, err)
|
||||
}
|
||||
rel, err := filepath.Rel(rootAbs, current)
|
||||
if err != nil {
|
||||
return fmt.Errorf("relativizing %s: %w", current, err)
|
||||
}
|
||||
rel = filepath.ToSlash(rel)
|
||||
files[rel] = localFile{
|
||||
Path: current,
|
||||
Snapshot: fileSnapshot{
|
||||
Size: info.Size(),
|
||||
ModTimeUnix: info.ModTime().Unix(),
|
||||
},
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return files, nil
|
||||
}
|
||||
|
||||
func buildState(localFiles map[string]localFile, remoteFiles map[string]remoteFile) syncState {
|
||||
state := syncState{
|
||||
Version: 1,
|
||||
Files: make(map[string]stateEntry),
|
||||
}
|
||||
for _, rel := range unionKeys(localFiles, remoteFiles, nil) {
|
||||
entry := stateEntry{}
|
||||
if local, ok := localFiles[rel]; ok {
|
||||
snapshot := local.Snapshot
|
||||
entry.Local = &snapshot
|
||||
}
|
||||
if remote, ok := remoteFiles[rel]; ok {
|
||||
snapshot := remote.Snapshot
|
||||
entry.Remote = &snapshot
|
||||
}
|
||||
state.Files[rel] = entry
|
||||
}
|
||||
return state
|
||||
}
|
||||
|
||||
func unionKeys(localFiles map[string]localFile, remoteFiles map[string]remoteFile, previous map[string]stateEntry) []string {
|
||||
keys := make(map[string]struct{}, len(localFiles)+len(remoteFiles)+len(previous))
|
||||
for key := range localFiles {
|
||||
keys[key] = struct{}{}
|
||||
}
|
||||
for key := range remoteFiles {
|
||||
keys[key] = struct{}{}
|
||||
}
|
||||
for key := range previous {
|
||||
keys[key] = struct{}{}
|
||||
}
|
||||
|
||||
out := make([]string, 0, len(keys))
|
||||
for key := range keys {
|
||||
out = append(out, key)
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func snapshotFromObject(object types.Object) fileSnapshot {
|
||||
snapshot := fileSnapshot{}
|
||||
if object.Size != nil {
|
||||
snapshot.Size = *object.Size
|
||||
}
|
||||
if object.LastModified != nil {
|
||||
snapshot.ModTimeUnix = object.LastModified.Unix()
|
||||
}
|
||||
return snapshot
|
||||
}
|
||||
|
||||
func snapshotFromHead(output *s3.HeadObjectOutput, fallback fileSnapshot) fileSnapshot {
|
||||
snapshot := fallback
|
||||
if output.ContentLength != nil {
|
||||
snapshot.Size = *output.ContentLength
|
||||
}
|
||||
if output.LastModified != nil {
|
||||
snapshot.ModTimeUnix = output.LastModified.Unix()
|
||||
}
|
||||
return snapshot
|
||||
}
|
||||
|
||||
func normalizePrefix(prefix string) string {
|
||||
prefix = strings.Trim(prefix, "/")
|
||||
if prefix == "" {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSuffix(prefix, "/") + "/"
|
||||
}
|
||||
|
||||
func relativeKey(key string, prefix string) (string, error) {
|
||||
normalizedPrefix := normalizePrefix(prefix)
|
||||
if normalizedPrefix != "" {
|
||||
if !strings.HasPrefix(key, normalizedPrefix) {
|
||||
return "", fmt.Errorf("s3 key %q does not match prefix %q", key, normalizedPrefix)
|
||||
}
|
||||
key = strings.TrimPrefix(key, normalizedPrefix)
|
||||
}
|
||||
if key == "" {
|
||||
return "", fmt.Errorf("s3 key resolves to empty local path")
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func safeLocalPath(root string, rel string) (string, error) {
|
||||
if rel == "" {
|
||||
return "", fmt.Errorf("empty relative path")
|
||||
}
|
||||
if strings.HasPrefix(rel, "/") {
|
||||
return "", fmt.Errorf("unsafe absolute s3 key %q", rel)
|
||||
}
|
||||
|
||||
clean := path.Clean(rel)
|
||||
for _, segment := range strings.Split(clean, "/") {
|
||||
if segment == "." || segment == ".." || segment == "" {
|
||||
return "", fmt.Errorf("unsafe s3 key %q", rel)
|
||||
}
|
||||
}
|
||||
|
||||
localPath := filepath.Join(root, filepath.FromSlash(clean))
|
||||
rootAbs, err := filepath.Abs(root)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolving root directory: %w", err)
|
||||
}
|
||||
localAbs, err := filepath.Abs(localPath)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolving local path: %w", err)
|
||||
}
|
||||
if localAbs != rootAbs && !strings.HasPrefix(localAbs, rootAbs+string(os.PathSeparator)) {
|
||||
return "", fmt.Errorf("unsafe s3 key %q", rel)
|
||||
}
|
||||
|
||||
return localAbs, nil
|
||||
}
|
||||
|
||||
func unixTime(seconds int64) time.Time {
|
||||
return time.Unix(seconds, 0)
|
||||
}
|
||||
@@ -0,0 +1,311 @@
|
||||
package syncer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3/types"
|
||||
)
|
||||
|
||||
func TestSyncDownloadsObjects(t *testing.T) {
|
||||
modified := time.Unix(100, 0)
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"config/app.yml": {body: "port: 8080\n", modified: modified},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
|
||||
result, err := New(client, Config{Bucket: "bucket", Prefix: "config", Dir: dir}).Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Downloaded: 1}) {
|
||||
t.Fatalf("result = %+v, want one download", result)
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(filepath.Join(dir, "app.yml"))
|
||||
if err != nil {
|
||||
t.Fatalf("reading downloaded file: %v", err)
|
||||
}
|
||||
if string(content) != "port: 8080\n" {
|
||||
t.Fatalf("downloaded content = %q", content)
|
||||
}
|
||||
|
||||
info, err := os.Stat(filepath.Join(dir, "app.yml"))
|
||||
if err != nil {
|
||||
t.Fatalf("stat downloaded file: %v", err)
|
||||
}
|
||||
if info.ModTime().Unix() != modified.Unix() {
|
||||
t.Fatalf("mod time = %s, want %s", info.ModTime(), modified)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncSkipsUnchangedObjects(t *testing.T) {
|
||||
modified := time.Unix(200, 0)
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"app.yml": {body: "same", modified: modified},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
service := New(client, Config{Bucket: "bucket", Dir: dir})
|
||||
|
||||
if _, err := service.Sync(context.Background()); err != nil {
|
||||
t.Fatalf("first Sync() error = %v", err)
|
||||
}
|
||||
client.getCalls = 0
|
||||
client.putCalls = 0
|
||||
|
||||
result, err := service.Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("second Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Unchanged: 1}) {
|
||||
t.Fatalf("result = %+v, want one unchanged", result)
|
||||
}
|
||||
if client.getCalls != 0 {
|
||||
t.Fatalf("GetObject calls = %d, want 0", client.getCalls)
|
||||
}
|
||||
if client.putCalls != 0 {
|
||||
t.Fatalf("PutObject calls = %d, want 0", client.putCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncDownloadsRemoteNewerObject(t *testing.T) {
|
||||
firstModified := time.Unix(300, 0)
|
||||
secondModified := time.Unix(301, 0)
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"app.yml": {body: "old", modified: firstModified},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
service := New(client, Config{Bucket: "bucket", Dir: dir})
|
||||
|
||||
if _, err := service.Sync(context.Background()); err != nil {
|
||||
t.Fatalf("first Sync() error = %v", err)
|
||||
}
|
||||
|
||||
client.objects["app.yml"] = fakeObject{body: "new", modified: secondModified}
|
||||
result, err := service.Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("second Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Updated: 1}) {
|
||||
t.Fatalf("result = %+v, want one update", result)
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(filepath.Join(dir, "app.yml"))
|
||||
if err != nil {
|
||||
t.Fatalf("reading updated file: %v", err)
|
||||
}
|
||||
if string(content) != "new" {
|
||||
t.Fatalf("updated content = %q, want new", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncUploadsLocalOnlyFile(t *testing.T) {
|
||||
client := newFakeS3(map[string]fakeObject{})
|
||||
dir := t.TempDir()
|
||||
localPath := filepath.Join(dir, "local.txt")
|
||||
writeFileAt(t, localPath, "local content", time.Unix(600, 0))
|
||||
|
||||
result, err := New(client, Config{Bucket: "bucket", Prefix: "config", Dir: dir}).Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Uploaded: 1}) {
|
||||
t.Fatalf("result = %+v, want one upload", result)
|
||||
}
|
||||
if got := client.objects["config/local.txt"].body; got != "local content" {
|
||||
t.Fatalf("uploaded body = %q, want local content", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncUploadsLocalNewerObject(t *testing.T) {
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"app.yml": {body: "remote", modified: time.Unix(700, 0)},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
service := New(client, Config{Bucket: "bucket", Dir: dir})
|
||||
|
||||
if _, err := service.Sync(context.Background()); err != nil {
|
||||
t.Fatalf("first Sync() error = %v", err)
|
||||
}
|
||||
localPath := filepath.Join(dir, "app.yml")
|
||||
writeFileAt(t, localPath, "local", time.Unix(800, 0))
|
||||
|
||||
result, err := service.Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("second Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Updated: 1}) {
|
||||
t.Fatalf("result = %+v, want one update", result)
|
||||
}
|
||||
if got := client.objects["app.yml"].body; got != "local" {
|
||||
t.Fatalf("uploaded body = %q, want local", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncPruneDeletesUnchangedLocalAfterRemoteDelete(t *testing.T) {
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"stale.txt": {body: "remote", modified: time.Unix(400, 0)},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
service := New(client, Config{Bucket: "bucket", Dir: dir, Prune: true})
|
||||
|
||||
if _, err := service.Sync(context.Background()); err != nil {
|
||||
t.Fatalf("first Sync() error = %v", err)
|
||||
}
|
||||
delete(client.objects, "stale.txt")
|
||||
|
||||
result, err := service.Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("second Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Deleted: 1}) {
|
||||
t.Fatalf("result = %+v, want one delete", result)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, "stale.txt")); !os.IsNotExist(err) {
|
||||
t.Fatalf("stale file still exists or unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncPruneDeletesUnchangedRemoteAfterLocalDelete(t *testing.T) {
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"stale.txt": {body: "remote", modified: time.Unix(450, 0)},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
service := New(client, Config{Bucket: "bucket", Dir: dir, Prune: true})
|
||||
|
||||
if _, err := service.Sync(context.Background()); err != nil {
|
||||
t.Fatalf("first Sync() error = %v", err)
|
||||
}
|
||||
if err := os.Remove(filepath.Join(dir, "stale.txt")); err != nil {
|
||||
t.Fatalf("removing local file: %v", err)
|
||||
}
|
||||
|
||||
result, err := service.Sync(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("second Sync() error = %v", err)
|
||||
}
|
||||
if result != (Result{Deleted: 1}) {
|
||||
t.Fatalf("result = %+v, want one delete", result)
|
||||
}
|
||||
if _, ok := client.objects["stale.txt"]; ok {
|
||||
t.Fatal("remote object still exists")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncRejectsUnsafeKeys(t *testing.T) {
|
||||
client := newFakeS3(map[string]fakeObject{
|
||||
"../secret.txt": {body: "secret", modified: time.Unix(500, 0)},
|
||||
})
|
||||
dir := t.TempDir()
|
||||
|
||||
_, err := New(client, Config{Bucket: "bucket", Dir: dir}).Sync(context.Background())
|
||||
if err == nil {
|
||||
t.Fatal("Sync() error = nil, want unsafe key error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "unsafe") {
|
||||
t.Fatalf("Sync() error = %v, want unsafe key error", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(filepath.Dir(dir), "secret.txt")); !os.IsNotExist(err) {
|
||||
t.Fatalf("unsafe file was written or unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeObject struct {
|
||||
body string
|
||||
modified time.Time
|
||||
}
|
||||
|
||||
type fakeS3 struct {
|
||||
objects map[string]fakeObject
|
||||
getCalls int
|
||||
putCalls int
|
||||
delCalls int
|
||||
now time.Time
|
||||
}
|
||||
|
||||
func newFakeS3(objects map[string]fakeObject) *fakeS3 {
|
||||
return &fakeS3{
|
||||
objects: objects,
|
||||
now: time.Unix(900, 0),
|
||||
}
|
||||
}
|
||||
|
||||
func (f *fakeS3) ListObjectsV2(_ context.Context, input *s3.ListObjectsV2Input, _ ...func(*s3.Options)) (*s3.ListObjectsV2Output, error) {
|
||||
var contents []types.Object
|
||||
prefix := aws.ToString(input.Prefix)
|
||||
for key, object := range f.objects {
|
||||
if !strings.HasPrefix(key, prefix) {
|
||||
continue
|
||||
}
|
||||
size := int64(len(object.body))
|
||||
modified := object.modified
|
||||
contents = append(contents, types.Object{
|
||||
Key: aws.String(key),
|
||||
LastModified: &modified,
|
||||
Size: &size,
|
||||
})
|
||||
}
|
||||
truncated := false
|
||||
return &s3.ListObjectsV2Output{
|
||||
Contents: contents,
|
||||
IsTruncated: &truncated,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeS3) GetObject(_ context.Context, input *s3.GetObjectInput, _ ...func(*s3.Options)) (*s3.GetObjectOutput, error) {
|
||||
f.getCalls++
|
||||
object := f.objects[aws.ToString(input.Key)]
|
||||
return &s3.GetObjectOutput{
|
||||
Body: io.NopCloser(strings.NewReader(object.body)),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeS3) PutObject(_ context.Context, input *s3.PutObjectInput, _ ...func(*s3.Options)) (*s3.PutObjectOutput, error) {
|
||||
f.putCalls++
|
||||
body, err := io.ReadAll(input.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
f.objects[aws.ToString(input.Key)] = fakeObject{
|
||||
body: string(body),
|
||||
modified: f.now,
|
||||
}
|
||||
return &s3.PutObjectOutput{}, nil
|
||||
}
|
||||
|
||||
func (f *fakeS3) HeadObject(_ context.Context, input *s3.HeadObjectInput, _ ...func(*s3.Options)) (*s3.HeadObjectOutput, error) {
|
||||
object := f.objects[aws.ToString(input.Key)]
|
||||
size := int64(len(object.body))
|
||||
modified := object.modified
|
||||
return &s3.HeadObjectOutput{
|
||||
ContentLength: &size,
|
||||
LastModified: &modified,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeS3) DeleteObject(_ context.Context, input *s3.DeleteObjectInput, _ ...func(*s3.Options)) (*s3.DeleteObjectOutput, error) {
|
||||
f.delCalls++
|
||||
delete(f.objects, aws.ToString(input.Key))
|
||||
return &s3.DeleteObjectOutput{}, nil
|
||||
}
|
||||
|
||||
func writeFileAt(t *testing.T, path string, content string, modified time.Time) {
|
||||
t.Helper()
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
t.Fatalf("creating parent dir: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
|
||||
t.Fatalf("writing file: %v", err)
|
||||
}
|
||||
if err := os.Chtimes(path, modified, modified); err != nil {
|
||||
t.Fatalf("setting mtime: %v", err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user