From 8c493c1199653a6d43e20a68f6026641fa4ade01 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 05:01:51 +0900 Subject: [PATCH 1/9] Add snapshot offload substrate --- ...oposed_physical_snapshot_object_offload.md | 19 +- internal/snapshotoffload/manifest.go | 140 +++++++++ internal/snapshotoffload/offload_test.go | 182 ++++++++++++ internal/snapshotoffload/publish.go | 195 +++++++++++++ internal/snapshotoffload/restore.go | 165 +++++++++++ internal/snapshotoffload/store.go | 276 ++++++++++++++++++ 6 files changed, 975 insertions(+), 2 deletions(-) create mode 100644 internal/snapshotoffload/manifest.go create mode 100644 internal/snapshotoffload/offload_test.go create mode 100644 internal/snapshotoffload/publish.go create mode 100644 internal/snapshotoffload/restore.go create mode 100644 internal/snapshotoffload/store.go diff --git a/docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md b/docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md index 7980e47a2..32ee1d594 100644 --- a/docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md +++ b/docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md @@ -1,8 +1,9 @@ # Physical Snapshot Object Offload -Status: Proposed +Status: Proposed — M0 implemented; M1 object-store-neutral substrate partial Author: bootjp Date: 2026-07-19 +Updated: 2026-07-23 ## 1. Scope @@ -30,6 +31,20 @@ only the independent substrate: No object client or runtime scheduling is enabled by this first slice. +The M1 object-store-neutral substrate now adds: + +- a `snapshotoffload.ObjectStore` interface with a local filesystem + implementation for deterministic tests and offline drills; +- the v1 JSON manifest schema and content-addressed payload key layout; +- payload-first publish from `OpenPersistedSnapshotExport`, including exact + byte count and SHA-256 verification before manifest commit; +- manifest-driven restore that downloads the opaque payload, checks exact + length and SHA-256, then calls `PreparePhysicalSnapshotRestore` with + operator-supplied target membership. + +The S3-compatible client, operator CLI, runtime scheduler, and retention/GC +remain pending. + ## 2. Safety boundary The exporter consumes only snapshots already committed by the etcd engine: @@ -143,7 +158,7 @@ permissions below the configured prefix. | Milestone | Scope | Status | |---|---|---| | M0 | Persisted snapshot export handle, complete-payload restore preparation, focused design | Implemented in the first substrate PR | -| M1 | Object client interface, S3-compatible implementation, immutable payload/manifest publication, download verification, operator CLI | Pending; depends on stable #1130 contract | +| M1 | Object client interface, S3-compatible implementation, immutable payload/manifest publication, download verification, operator CLI | Partial: object-store interface, local implementation, manifest schema, payload-first publish, and verified restore are implemented; S3 client and CLI pending | | M2 | Leader-only per-group scheduler, metrics, jitter, concurrency bounds, cancellation and restart idempotency | Pending | | M3 | Retention/GC, restore drills, corruption tests, multi-node acceptance, operational documentation | Pending | diff --git a/internal/snapshotoffload/manifest.go b/internal/snapshotoffload/manifest.go new file mode 100644 index 000000000..b963d937d --- /dev/null +++ b/internal/snapshotoffload/manifest.go @@ -0,0 +1,140 @@ +package snapshotoffload + +import ( + "encoding/json" + "fmt" + "path" + "strings" + "time" + + "github.com/cockroachdb/errors" + raftpb "go.etcd.io/raft/v3/raftpb" +) + +const ( + ManifestSchemaVersion = 1 + payloadObjectSuffix = ".fsm" + manifestObjectSuffix = ".json" +) + +var ( + ErrInvalidOptions = errors.New("snapshot offload: invalid options") + ErrIntegrity = errors.New("snapshot offload: integrity check failed") + ErrObjectNotFound = errors.New("snapshot offload: object not found") +) + +type Manifest struct { + SchemaVersion int `json:"schema_version"` + CreatedAt time.Time `json:"created_at"` + SourceCluster string `json:"source_cluster,omitempty"` + GroupID uint64 `json:"group_id"` + SnapshotIndex uint64 `json:"snapshot_index"` + SnapshotTerm uint64 `json:"snapshot_term"` + ConfState ManifestConfState `json:"conf_state"` + Payload PayloadDescriptor `json:"payload"` + BinaryVersion string `json:"binary_version,omitempty"` + ManifestKey string `json:"manifest_key"` + ManifestSHA256 string `json:"manifest_sha256,omitempty"` +} + +type ManifestConfState struct { + Voters []uint64 `json:"voters,omitempty"` + Learners []uint64 `json:"learners,omitempty"` + VotersOutgoing []uint64 `json:"voters_outgoing,omitempty"` + LearnersNext []uint64 `json:"learners_next,omitempty"` + AutoLeave bool `json:"auto_leave,omitempty"` +} + +type PayloadDescriptor struct { + Key string `json:"key"` + Bytes int64 `json:"bytes"` + SHA256 string `json:"sha256"` + SourceCRC32C uint32 `json:"source_crc32c"` +} + +func (m Manifest) MarshalCanonical() ([]byte, string, error) { + m.ManifestSHA256 = "" + data, err := json.MarshalIndent(m, "", " ") + if err != nil { + return nil, "", errors.WithStack(err) + } + sum := hexSHA256Bytes(data) + m.ManifestSHA256 = sum + final, err := json.MarshalIndent(m, "", " ") + if err != nil { + return nil, "", errors.WithStack(err) + } + return append(final, '\n'), sum, nil +} + +func DecodeManifest(data []byte) (Manifest, error) { + var manifest Manifest + if err := json.Unmarshal(data, &manifest); err != nil { + return Manifest{}, errors.WithStack(err) + } + if err := validateManifest(manifest); err != nil { + return Manifest{}, err + } + return manifest, nil +} + +func validateManifest(manifest Manifest) error { + switch { + case manifest.SchemaVersion != ManifestSchemaVersion: + return errors.Wrapf(ErrInvalidOptions, "unsupported manifest schema version %d", manifest.SchemaVersion) + case manifest.GroupID == 0 && strings.TrimSpace(manifest.SourceCluster) == "": + return errors.Wrap(ErrInvalidOptions, "source cluster is required for group 0 manifests") + case manifest.SnapshotIndex == 0: + return errors.Wrap(ErrInvalidOptions, "snapshot index must be > 0") + case manifest.SnapshotTerm == 0: + return errors.Wrap(ErrInvalidOptions, "snapshot term must be > 0") + case manifest.Payload.Key == "": + return errors.Wrap(ErrInvalidOptions, "payload key is required") + case manifest.Payload.Bytes < 0: + return errors.Wrap(ErrInvalidOptions, "payload byte count must be >= 0") + case !isSHA256Hex(manifest.Payload.SHA256): + return errors.Wrap(ErrInvalidOptions, "payload sha256 must be 64 lowercase hex characters") + case manifest.ManifestKey == "": + return errors.Wrap(ErrInvalidOptions, "manifest key is required") + } + return nil +} + +func payloadKey(prefix, sha256Hex string) (string, error) { + if !isSHA256Hex(sha256Hex) { + return "", errors.Wrap(ErrInvalidOptions, "payload sha256 must be 64 lowercase hex characters") + } + prefix = cleanObjectPrefix(prefix) + return path.Join(prefix, "v1", "payloads", "sha256", sha256Hex[:2], sha256Hex+payloadObjectSuffix), nil +} + +func manifestKey(prefix string, groupID, index, term uint64) (string, error) { + if index == 0 || term == 0 { + return "", errors.Wrap(ErrInvalidOptions, "snapshot index and term must be > 0") + } + prefix = cleanObjectPrefix(prefix) + return path.Join(prefix, "v1", "groups", fmt.Sprintf("%d", groupID), "snapshots", + fmt.Sprintf("%020d-%020d%s", index, term, manifestObjectSuffix)), nil +} + +func cleanObjectPrefix(prefix string) string { + prefix = strings.TrimSpace(prefix) + prefix = strings.Trim(prefix, "/") + if prefix == "" { + return "." + } + return path.Clean(prefix) +} + +func manifestConfState(confState *raftpb.ConfState) ManifestConfState { + if confState == nil { + return ManifestConfState{} + } + return ManifestConfState{ + Voters: append([]uint64(nil), confState.GetVoters()...), + Learners: append([]uint64(nil), confState.GetLearners()...), + VotersOutgoing: append([]uint64(nil), confState.GetVotersOutgoing()...), + LearnersNext: append([]uint64(nil), confState.GetLearnersNext()...), + AutoLeave: confState.GetAutoLeave(), + } +} diff --git a/internal/snapshotoffload/offload_test.go b/internal/snapshotoffload/offload_test.go new file mode 100644 index 000000000..d437045a9 --- /dev/null +++ b/internal/snapshotoffload/offload_test.go @@ -0,0 +1,182 @@ +package snapshotoffload + +import ( + "bytes" + "context" + "io" + "os" + "path/filepath" + "testing" + "time" + + etcdraftengine "github.com/bootjp/elastickv/internal/raftengine/etcd" + "github.com/stretchr/testify/require" +) + +func TestPublishAndRestorePhysicalSnapshotRoundTrip(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1opaque-physical-snapshot-payload") + sourceDataDir := seedPhysicalSnapshot(t, root, "source", payload, 42, 7, []etcdraftengine.Peer{ + {NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}, + {NodeID: 2, ID: "n2", Address: "127.0.0.1:12002"}, + }) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 7, + SourceCluster: "cluster-a", + BinaryVersion: "test-version", + CreatedAt: time.Unix(100, 0).UTC(), + }) + require.NoError(t, err) + require.Equal(t, uint64(42), manifest.SnapshotIndex) + require.Equal(t, uint64(7), manifest.SnapshotTerm) + require.Equal(t, int64(len(payload)), manifest.Payload.Bytes) + require.Equal(t, []uint64{1, 2}, manifest.ConfState.Voters) + require.Contains(t, manifest.Payload.Key, "/v1/payloads/sha256/") + require.Contains(t, manifest.ManifestKey, "/v1/groups/7/snapshots/") + + storedManifest := loadTestManifest(t, ctx, store, manifest.ManifestKey) + require.Equal(t, manifest.Payload, storedManifest.Payload) + + payloadPath, err := store.pathForKey(manifest.Payload.Key) + require.NoError(t, err) + storedPayload, err := os.ReadFile(payloadPath) + require.NoError(t, err) + require.Equal(t, payload, storedPayload) + + restoreDataDir := filepath.Join(root, "restored") + result, err := RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: store, + ManifestKey: manifest.ManifestKey, + DataDir: restoreDataDir, + Peers: []etcdraftengine.Peer{ + {NodeID: 9, ID: "n9", Address: "127.0.0.1:19009"}, + }, + }) + require.NoError(t, err) + require.Equal(t, int64(len(payload)), result.PayloadBytes) + require.Equal(t, manifest.Payload.SHA256, result.PayloadSHA256) + + export, ok, err := etcdraftengine.OpenPersistedSnapshotExport(restoreDataDir) + require.NoError(t, err) + require.True(t, ok) + defer func() { require.NoError(t, export.Close()) }() + require.Equal(t, []uint64{9}, export.Metadata().ConfState.GetVoters()) + var restored bytes.Buffer + n, err := export.WriteTo(&restored) + require.NoError(t, err) + require.Equal(t, int64(len(payload)), n) + require.Equal(t, payload, restored.Bytes()) +} + +func TestRestoreRejectsCorruptPayloadAndLeavesDestinationAbsent(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-before-corruption") + sourceDataDir := seedPhysicalSnapshot(t, root, "source", payload, 10, 3, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + payloadPath, err := store.pathForKey(manifest.Payload.Key) + require.NoError(t, err) + require.NoError(t, os.WriteFile(payloadPath, []byte("corrupt"), 0o600)) + + restoreDataDir := filepath.Join(root, "restored") + _, err = RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: store, + ManifestKey: manifest.ManifestKey, + DataDir: restoreDataDir, + Peers: singlePeer(), + }) + require.ErrorIs(t, err, ErrIntegrity) + _, statErr := os.Stat(restoreDataDir) + require.True(t, os.IsNotExist(statErr)) +} + +func TestPublishRejectsExistingPayloadMismatchBeforeManifest(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-with-known-sha") + sourceDataDir := seedPhysicalSnapshot(t, root, "source", payload, 11, 4, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + payloadKey, err := payloadKey("cluster-a", hexSHA256Bytes(payload)) + require.NoError(t, err) + bad := []byte("different-payload") + _, err = store.PutObject(ctx, payloadKey, bytes.NewReader(bad), PutOptions{ + Size: int64(len(bad)), + SHA256: hexSHA256Bytes(bad), + }) + require.NoError(t, err) + + _, err = PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.ErrorIs(t, err, ErrIntegrity) + manifestKey, err := manifestKey("cluster-a", 1, 11, 4) + require.NoError(t, err) + _, ok, err := store.HeadObject(ctx, manifestKey) + require.NoError(t, err) + require.False(t, ok) +} + +func seedPhysicalSnapshot( + t *testing.T, + root string, + name string, + payload []byte, + index uint64, + term uint64, + peers []etcdraftengine.Peer, +) string { + t.Helper() + input := filepath.Join(root, name+".fsm") + require.NoError(t, os.WriteFile(input, payload, 0o600)) + dataDir := filepath.Join(root, name+"-raft") + _, err := etcdraftengine.PreparePhysicalSnapshotRestore(etcdraftengine.PhysicalSnapshotRestoreOptions{ + InputFSMPath: input, + DataDir: dataDir, + Index: index, + Term: term, + Peers: peers, + }) + require.NoError(t, err) + return dataDir +} + +func newTestLocalStore(t *testing.T, root string) *LocalStore { + t.Helper() + store, err := NewLocalStore(root) + require.NoError(t, err) + return store +} + +func loadTestManifest(t *testing.T, ctx context.Context, store ObjectStore, key string) Manifest { + t.Helper() + body, _, err := store.GetObject(ctx, key) + require.NoError(t, err) + defer func() { require.NoError(t, body.Close()) }() + raw, err := io.ReadAll(body) + require.NoError(t, err) + manifest, err := DecodeManifest(raw) + require.NoError(t, err) + return manifest +} + +func singlePeer() []etcdraftengine.Peer { + return []etcdraftengine.Peer{{NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}} +} diff --git a/internal/snapshotoffload/publish.go b/internal/snapshotoffload/publish.go new file mode 100644 index 000000000..81379ee4b --- /dev/null +++ b/internal/snapshotoffload/publish.go @@ -0,0 +1,195 @@ +package snapshotoffload + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "io" + "os" + "strings" + "time" + + etcdraftengine "github.com/bootjp/elastickv/internal/raftengine/etcd" + "github.com/cockroachdb/errors" +) + +type PublishOptions struct { + Store ObjectStore + DataDir string + Prefix string + GroupID uint64 + SourceCluster string + BinaryVersion string + CreatedAt time.Time +} + +func PublishPersistedSnapshot(ctx context.Context, opts PublishOptions) (*Manifest, error) { + if err := validatePublishOptions(opts); err != nil { + return nil, err + } + export, err := openPublishExport(opts.DataDir) + if err != nil { + return nil, err + } + defer func() { _ = export.Close() }() + + metadata := export.Metadata() + payloadFile, payloadSHA, payloadBytes, err := spoolExport(ctx, export) + if err != nil { + return nil, err + } + defer func() { + _ = payloadFile.Close() + _ = os.Remove(payloadFile.Name()) + }() + if payloadBytes != metadata.PayloadBytes { + return nil, errors.Wrapf(ErrIntegrity, "exported %d bytes, metadata expected %d", payloadBytes, metadata.PayloadBytes) + } + payloadObjectKey, err := payloadKey(opts.Prefix, payloadSHA) + if err != nil { + return nil, err + } + if err := putPayload(ctx, opts.Store, payloadObjectKey, payloadFile, payloadBytes, payloadSHA); err != nil { + return nil, err + } + manifest, err := buildManifest(opts, metadata, payloadObjectKey, payloadSHA) + if err != nil { + return nil, err + } + if err := putManifest(ctx, opts.Store, manifest); err != nil { + return nil, err + } + return manifest, nil +} + +func openPublishExport(dataDir string) (*etcdraftengine.PersistedSnapshotExport, error) { + export, ok, err := etcdraftengine.OpenPersistedSnapshotExport(dataDir) + if err != nil { + return nil, errors.Wrap(err, "open persisted snapshot export") + } + if !ok { + return nil, errors.Wrap(ErrObjectNotFound, "no persisted snapshot available") + } + return export, nil +} + +func buildManifest( + opts PublishOptions, + metadata etcdraftengine.PersistedSnapshotExportMetadata, + payloadObjectKey string, + payloadSHA string, +) (*Manifest, error) { + manifestObjectKey, err := manifestKey(opts.Prefix, opts.GroupID, metadata.Index, metadata.Term) + if err != nil { + return nil, err + } + createdAt := opts.CreatedAt + if createdAt.IsZero() { + createdAt = time.Now().UTC() + } + return &Manifest{ + SchemaVersion: ManifestSchemaVersion, + CreatedAt: createdAt.UTC(), + SourceCluster: stringsTrim(opts.SourceCluster), + GroupID: opts.GroupID, + SnapshotIndex: metadata.Index, + SnapshotTerm: metadata.Term, + ConfState: manifestConfState(metadata.ConfState), + Payload: PayloadDescriptor{ + Key: payloadObjectKey, + Bytes: metadata.PayloadBytes, + SHA256: payloadSHA, + SourceCRC32C: metadata.CRC32C, + }, + BinaryVersion: stringsTrim(opts.BinaryVersion), + ManifestKey: manifestObjectKey, + }, nil +} + +func putManifest(ctx context.Context, store ObjectStore, manifest *Manifest) error { + data, manifestSHA, err := manifest.MarshalCanonical() + if err != nil { + return err + } + if _, err := store.PutObject(ctx, manifest.ManifestKey, bytes.NewReader(data), PutOptions{ + Size: int64(len(data)), + SHA256: hexSHA256Bytes(data), + ContentType: "application/json", + }); err != nil { + return errors.Wrap(err, "put snapshot manifest") + } + manifest.ManifestSHA256 = manifestSHA + return nil +} + +func validatePublishOptions(opts PublishOptions) error { + switch { + case opts.Store == nil: + return errors.Wrap(ErrInvalidOptions, "object store is required") + case stringsTrim(opts.DataDir) == "": + return errors.Wrap(ErrInvalidOptions, "data dir is required") + default: + return nil + } +} + +func spoolExport(ctx context.Context, export *etcdraftengine.PersistedSnapshotExport) (*os.File, string, int64, error) { + tmp, err := os.CreateTemp("", "elastickv-snapshot-offload-*.fsm") + if err != nil { + return nil, "", 0, errors.WithStack(err) + } + keep := false + defer func() { + if !keep { + _ = tmp.Close() + _ = os.Remove(tmp.Name()) + } + }() + hash := sha256.New() + n, err := export.WriteTo(io.MultiWriter(tmp, hash)) + if err != nil { + return nil, "", n, errors.Wrap(err, "spool persisted snapshot export") + } + if err := ctx.Err(); err != nil { + return nil, "", n, errors.WithStack(err) + } + if err := tmp.Sync(); err != nil { + return nil, "", n, errors.WithStack(err) + } + if _, err := tmp.Seek(0, io.SeekStart); err != nil { + return nil, "", n, errors.WithStack(err) + } + keep = true + return tmp, hex.EncodeToString(hash.Sum(nil)), n, nil +} + +func putPayload(ctx context.Context, store ObjectStore, key string, file *os.File, size int64, sha string) error { + if info, ok, err := store.HeadObject(ctx, key); err != nil { + return errors.Wrap(err, "head snapshot payload") + } else if ok { + if info.Size == size && info.SHA256 == sha { + return nil + } + return errors.Wrapf(ErrIntegrity, "payload object %s already exists with different content", key) + } + if _, err := file.Seek(0, io.SeekStart); err != nil { + return errors.WithStack(err) + } + info, err := store.PutObject(ctx, key, file, PutOptions{ + Size: size, + SHA256: sha, + ContentType: "application/octet-stream", + }) + if err != nil { + return errors.Wrap(err, "put snapshot payload") + } + if info.Size != size || info.SHA256 != sha { + return errors.Wrapf(ErrIntegrity, "payload object %s remote integrity mismatch", key) + } + return nil +} + +func stringsTrim(v string) string { + return strings.TrimSpace(v) +} diff --git a/internal/snapshotoffload/restore.go b/internal/snapshotoffload/restore.go new file mode 100644 index 000000000..85dea5539 --- /dev/null +++ b/internal/snapshotoffload/restore.go @@ -0,0 +1,165 @@ +package snapshotoffload + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "io" + "os" + "path/filepath" + + etcdraftengine "github.com/bootjp/elastickv/internal/raftengine/etcd" + "github.com/cockroachdb/errors" +) + +type RestoreOptions struct { + Store ObjectStore + ManifestKey string + Manifest *Manifest + DataDir string + Peers []etcdraftengine.Peer +} + +func LoadManifest(ctx context.Context, store ObjectStore, key string) (Manifest, error) { + if store == nil { + return Manifest{}, errors.Wrap(ErrInvalidOptions, "object store is required") + } + body, _, err := store.GetObject(ctx, key) + if err != nil { + return Manifest{}, errors.Wrap(err, "get snapshot manifest") + } + defer func() { _ = body.Close() }() + var buf bytes.Buffer + if _, err := io.Copy(&buf, contextReader{ctx: ctx, reader: body}); err != nil { + return Manifest{}, errors.WithStack(err) + } + manifest, err := DecodeManifest(buf.Bytes()) + if err != nil { + return Manifest{}, err + } + if normalizeObjectKey(key) != normalizeObjectKey(manifest.ManifestKey) { + return Manifest{}, errors.Wrapf(ErrIntegrity, "manifest key mismatch: loaded %s, body says %s", key, manifest.ManifestKey) + } + return manifest, nil +} + +func RestorePhysicalSnapshot(ctx context.Context, opts RestoreOptions) (*etcdraftengine.ExternalSnapshotRestoreResult, error) { + if err := validateRestoreOptions(opts); err != nil { + return nil, err + } + manifest, err := restoreManifest(ctx, opts) + if err != nil { + return nil, err + } + if err := validateManifest(manifest); err != nil { + return nil, err + } + downloadDir, err := os.MkdirTemp(filepath.Dir(filepath.Clean(opts.DataDir)), ".snapshot-offload-restore-*") + if err != nil { + return nil, errors.WithStack(err) + } + defer func() { _ = os.RemoveAll(downloadDir) }() + payloadPath := filepath.Join(downloadDir, "payload.fsm") + if err := downloadVerifiedPayload(ctx, opts.Store, manifest, payloadPath); err != nil { + return nil, err + } + result, err := etcdraftengine.PreparePhysicalSnapshotRestore(etcdraftengine.PhysicalSnapshotRestoreOptions{ + InputFSMPath: payloadPath, + DataDir: opts.DataDir, + Index: manifest.SnapshotIndex, + Term: manifest.SnapshotTerm, + Peers: opts.Peers, + ExpectedPayloadSHA256: manifest.Payload.SHA256, + }) + if err != nil { + return nil, errors.Wrap(err, "prepare physical snapshot restore") + } + return result, nil +} + +func validateRestoreOptions(opts RestoreOptions) error { + switch { + case opts.Store == nil: + return errors.Wrap(ErrInvalidOptions, "object store is required") + case opts.Manifest == nil && stringsTrim(opts.ManifestKey) == "": + return errors.Wrap(ErrInvalidOptions, "manifest key is required") + case stringsTrim(opts.DataDir) == "": + return errors.Wrap(ErrInvalidOptions, "data dir is required") + case len(opts.Peers) == 0: + return errors.Wrap(ErrInvalidOptions, "restore peers are required") + default: + return nil + } +} + +func restoreManifest(ctx context.Context, opts RestoreOptions) (Manifest, error) { + if opts.Manifest != nil { + return *opts.Manifest, nil + } + return LoadManifest(ctx, opts.Store, opts.ManifestKey) +} + +func downloadVerifiedPayload(ctx context.Context, store ObjectStore, manifest Manifest, finalPath string) error { + body, info, err := store.GetObject(ctx, manifest.Payload.Key) + if err != nil { + return errors.Wrap(err, "get snapshot payload") + } + defer func() { _ = body.Close() }() + if err := validatePayloadInfo(manifest, info); err != nil { + return err + } + tmpPath, err := writeDownloadedPayloadTemp(ctx, filepath.Dir(finalPath), manifest, body) + if err != nil { + return err + } + defer func() { + _ = os.Remove(tmpPath) + }() + if err := os.Rename(tmpPath, finalPath); err != nil { + return errors.WithStack(err) + } + return syncDir(filepath.Dir(finalPath)) +} + +func validatePayloadInfo(manifest Manifest, info ObjectInfo) error { + if info.Size != manifest.Payload.Bytes || info.SHA256 != manifest.Payload.SHA256 { + return errors.Wrapf(ErrIntegrity, "payload head mismatch for %s", manifest.Payload.Key) + } + return nil +} + +func writeDownloadedPayloadTemp(ctx context.Context, dir string, manifest Manifest, body io.Reader) (string, error) { + tmp, err := os.CreateTemp(dir, ".payload-*") + if err != nil { + return "", errors.WithStack(err) + } + tmpPath := tmp.Name() + keep := false + defer func() { + if !keep { + _ = tmp.Close() + _ = os.Remove(tmpPath) + } + }() + hash := sha256.New() + n, err := io.Copy(io.MultiWriter(tmp, hash), contextReader{ctx: ctx, reader: body}) + if err != nil { + return "", errors.WithStack(err) + } + gotSHA := hex.EncodeToString(hash.Sum(nil)) + if n != manifest.Payload.Bytes { + return "", errors.Wrapf(ErrIntegrity, "payload %s downloaded %d bytes, expected %d", manifest.Payload.Key, n, manifest.Payload.Bytes) + } + if gotSHA != manifest.Payload.SHA256 { + return "", errors.Wrapf(ErrIntegrity, "payload %s sha256 %s, expected %s", manifest.Payload.Key, gotSHA, manifest.Payload.SHA256) + } + if err := tmp.Sync(); err != nil { + return "", errors.WithStack(err) + } + if err := tmp.Close(); err != nil { + return "", errors.WithStack(err) + } + keep = true + return tmpPath, nil +} diff --git a/internal/snapshotoffload/store.go b/internal/snapshotoffload/store.go new file mode 100644 index 000000000..d091e923d --- /dev/null +++ b/internal/snapshotoffload/store.go @@ -0,0 +1,276 @@ +package snapshotoffload + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "io" + "os" + "path" + "path/filepath" + "strings" + + "github.com/cockroachdb/errors" +) + +type ObjectStore interface { + PutObject(ctx context.Context, key string, body io.Reader, opts PutOptions) (ObjectInfo, error) + GetObject(ctx context.Context, key string) (io.ReadCloser, ObjectInfo, error) + HeadObject(ctx context.Context, key string) (ObjectInfo, bool, error) +} + +type PutOptions struct { + Size int64 + SHA256 string + ContentType string +} + +type ObjectInfo struct { + Key string + Size int64 + SHA256 string +} + +type LocalStore struct { + root string +} + +const localStoreDirPerm = 0o755 + +func NewLocalStore(root string) (*LocalStore, error) { + if strings.TrimSpace(root) == "" { + return nil, errors.Wrap(ErrInvalidOptions, "local store root is required") + } + return &LocalStore{root: filepath.Clean(root)}, nil +} + +func (s *LocalStore) PutObject(ctx context.Context, key string, body io.Reader, opts PutOptions) (ObjectInfo, error) { + if err := validatePutOptions(opts); err != nil { + return ObjectInfo{}, err + } + finalPath, err := s.pathForKey(key) + if err != nil { + return ObjectInfo{}, err + } + if err := os.MkdirAll(filepath.Dir(finalPath), localStoreDirPerm); err != nil { + return ObjectInfo{}, errors.WithStack(err) + } + tmpPath, info, err := writeLocalObjectTemp(ctx, filepath.Dir(finalPath), key, body, opts) + if err != nil { + return ObjectInfo{}, err + } + defer func() { _ = os.Remove(tmpPath) }() + return s.commitTempObject(key, tmpPath, finalPath, info) +} + +func (s *LocalStore) GetObject(ctx context.Context, key string) (io.ReadCloser, ObjectInfo, error) { + if err := ctx.Err(); err != nil { + return nil, ObjectInfo{}, errors.WithStack(err) + } + objectPath, err := s.pathForKey(key) + if err != nil { + return nil, ObjectInfo{}, err + } + info, err := s.objectInfoForPath(key, objectPath) + if err != nil { + return nil, ObjectInfo{}, err + } + file, err := os.Open(objectPath) + if err != nil { + if os.IsNotExist(err) { + return nil, ObjectInfo{}, errors.Wrapf(ErrObjectNotFound, "object %s", key) + } + return nil, ObjectInfo{}, errors.WithStack(err) + } + return file, info, nil +} + +func (s *LocalStore) HeadObject(ctx context.Context, key string) (ObjectInfo, bool, error) { + if err := ctx.Err(); err != nil { + return ObjectInfo{}, false, errors.WithStack(err) + } + objectPath, err := s.pathForKey(key) + if err != nil { + return ObjectInfo{}, false, err + } + info, err := s.objectInfoForPath(key, objectPath) + if err != nil { + if errors.Is(err, ErrObjectNotFound) { + return ObjectInfo{}, false, nil + } + return ObjectInfo{}, false, err + } + return info, true, nil +} + +func (s *LocalStore) pathForKey(key string) (string, error) { + if s == nil { + return "", errors.Wrap(ErrInvalidOptions, "object store is required") + } + normalized := normalizeObjectKey(key) + if normalized == "" || normalized == "." || strings.HasPrefix(normalized, "../") { + return "", errors.Wrapf(ErrInvalidOptions, "invalid object key %q", key) + } + return filepath.Join(s.root, filepath.FromSlash(normalized)), nil +} + +func (s *LocalStore) objectInfoForPath(key, objectPath string) (ObjectInfo, error) { + file, err := os.Open(objectPath) + if err != nil { + if os.IsNotExist(err) { + return ObjectInfo{}, errors.Wrapf(ErrObjectNotFound, "object %s", key) + } + return ObjectInfo{}, errors.WithStack(err) + } + defer func() { _ = file.Close() }() + stat, err := file.Stat() + if err != nil { + return ObjectInfo{}, errors.WithStack(err) + } + if !stat.Mode().IsRegular() { + return ObjectInfo{}, errors.Wrapf(ErrInvalidOptions, "object %s is not a regular file", key) + } + sum := sha256.New() + if _, err := io.Copy(sum, file); err != nil { + return ObjectInfo{}, errors.WithStack(err) + } + return ObjectInfo{ + Key: normalizeObjectKey(key), + Size: stat.Size(), + SHA256: hex.EncodeToString(sum.Sum(nil)), + }, nil +} + +func writeLocalObjectTemp( + ctx context.Context, + dir string, + key string, + body io.Reader, + opts PutOptions, +) (string, ObjectInfo, error) { + tmp, err := os.CreateTemp(dir, ".put-*") + if err != nil { + return "", ObjectInfo{}, errors.WithStack(err) + } + tmpPath := tmp.Name() + keep := false + defer func() { + if !keep { + _ = tmp.Close() + _ = os.Remove(tmpPath) + } + }() + sum := sha256.New() + n, err := io.Copy(io.MultiWriter(tmp, sum), contextReader{ctx: ctx, reader: body}) + if err != nil { + return "", ObjectInfo{}, errors.WithStack(err) + } + gotSHA := hex.EncodeToString(sum.Sum(nil)) + if n != opts.Size { + return "", ObjectInfo{}, errors.Wrapf(ErrIntegrity, "object %s wrote %d bytes, expected %d", key, n, opts.Size) + } + if gotSHA != opts.SHA256 { + return "", ObjectInfo{}, errors.Wrapf(ErrIntegrity, "object %s wrote sha256 %s, expected %s", key, gotSHA, opts.SHA256) + } + if err := tmp.Sync(); err != nil { + return "", ObjectInfo{}, errors.WithStack(err) + } + if err := tmp.Close(); err != nil { + return "", ObjectInfo{}, errors.WithStack(err) + } + keep = true + return tmpPath, ObjectInfo{Key: normalizeObjectKey(key), Size: opts.Size, SHA256: opts.SHA256}, nil +} + +func (s *LocalStore) commitTempObject(key, tmpPath, finalPath string, expected ObjectInfo) (ObjectInfo, error) { + if err := os.Link(tmpPath, finalPath); err != nil { + if !os.IsExist(err) { + return ObjectInfo{}, errors.WithStack(err) + } + return s.verifyExistingObject(key, finalPath, expected) + } + if err := syncDir(filepath.Dir(finalPath)); err != nil { + return ObjectInfo{}, err + } + return expected, nil +} + +func (s *LocalStore) verifyExistingObject(key, finalPath string, expected ObjectInfo) (ObjectInfo, error) { + info, err := s.objectInfoForPath(key, finalPath) + if err != nil { + return ObjectInfo{}, err + } + if info.Size == expected.Size && info.SHA256 == expected.SHA256 { + return info, nil + } + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "object %s already exists with different content", key) +} + +func validatePutOptions(opts PutOptions) error { + switch { + case opts.Size < 0: + return errors.Wrap(ErrInvalidOptions, "object size must be >= 0") + case !isSHA256Hex(opts.SHA256): + return errors.Wrap(ErrInvalidOptions, "object sha256 must be 64 lowercase hex characters") + default: + return nil + } +} + +func normalizeObjectKey(key string) string { + key = strings.TrimSpace(key) + key = strings.TrimPrefix(key, "/") + return path.Clean(key) +} + +type contextReader struct { + ctx context.Context + reader io.Reader +} + +func (r contextReader) Read(p []byte) (int, error) { + if r.ctx != nil { + if err := r.ctx.Err(); err != nil { + return 0, errors.WithStack(err) + } + } + n, err := r.reader.Read(p) + if err != nil { + if errors.Is(err, io.EOF) { + return n, err //nolint:wrapcheck // io.Reader must return io.EOF unwrapped so io.Copy treats it as normal completion. + } + return n, errors.WithStack(err) + } + if r.ctx != nil { + if ctxErr := r.ctx.Err(); ctxErr != nil { + return n, errors.WithStack(ctxErr) + } + } + return n, nil +} + +func syncDir(dir string) error { + f, err := os.Open(dir) + if err != nil { + return errors.WithStack(err) + } + defer func() { _ = f.Close() }() + return errors.WithStack(f.Sync()) +} + +func hexSHA256Bytes(data []byte) string { + sum := sha256.Sum256(data) + return hex.EncodeToString(sum[:]) +} + +func isSHA256Hex(s string) bool { + if len(s) != sha256.Size*2 { + return false + } + for _, c := range s { + if (c < '0' || c > '9') && (c < 'a' || c > 'f') { + return false + } + } + return true +} From b96976e8aabd9a6419663aa43523a961b8691917 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 05:28:11 +0900 Subject: [PATCH 2/9] Harden snapshot offload validation --- .../2026_06_12_proposed_scaling_roadmap.md | 2 +- ...rtial_physical_snapshot_object_offload.md} | 4 +- internal/snapshotoffload/manifest.go | 17 ++ internal/snapshotoffload/offload_test.go | 215 +++++++++++++++++- internal/snapshotoffload/publish.go | 76 ++++++- internal/snapshotoffload/restore.go | 86 ++++++- internal/snapshotoffload/store.go | 50 +++- 7 files changed, 417 insertions(+), 33 deletions(-) rename docs/design/{2026_07_19_proposed_physical_snapshot_object_offload.md => 2026_07_19_partial_physical_snapshot_object_offload.md} (98%) diff --git a/docs/design/2026_06_12_proposed_scaling_roadmap.md b/docs/design/2026_06_12_proposed_scaling_roadmap.md index f2ad409b2..76a179915 100644 --- a/docs/design/2026_06_12_proposed_scaling_roadmap.md +++ b/docs/design/2026_06_12_proposed_scaling_roadmap.md @@ -409,7 +409,7 @@ TBD).** reuses the existing TTL helper path. **M4 — Disaster-recovery snapshot offload -([`2026_07_19_proposed_physical_snapshot_object_offload.md`](2026_07_19_proposed_physical_snapshot_object_offload.md)).** +([`2026_07_19_partial_physical_snapshot_object_offload.md`](2026_07_19_partial_physical_snapshot_object_offload.md)).** - Periodic per-shard Pebble snapshot uploaded to an S3-compatible bucket (the S3 adapter already speaks the protocol). - Restore is `s3 fetch → pebble.Ingest`; combined with M1's diff --git a/docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md b/docs/design/2026_07_19_partial_physical_snapshot_object_offload.md similarity index 98% rename from docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md rename to docs/design/2026_07_19_partial_physical_snapshot_object_offload.md index 32ee1d594..e70269a18 100644 --- a/docs/design/2026_07_19_proposed_physical_snapshot_object_offload.md +++ b/docs/design/2026_07_19_partial_physical_snapshot_object_offload.md @@ -1,6 +1,6 @@ # Physical Snapshot Object Offload -Status: Proposed — M0 implemented; M1 object-store-neutral substrate partial +Status: Partial — M0 implemented; M1 object-store-neutral substrate partial Author: bootjp Date: 2026-07-19 Updated: 2026-07-23 @@ -162,7 +162,7 @@ permissions below the configured prefix. | M2 | Leader-only per-group scheduler, metrics, jitter, concurrency bounds, cancellation and restart idempotency | Pending | | M3 | Retention/GC, restore drills, corruption tests, multi-node acceptance, operational documentation | Pending | -The filename and header remain `proposed` until M1-M3 complete the central +The filename and header remain `partial` until M1-M3 complete the central object-offload subsystem. At that point the completion PR must use `git mv` to rename this file to `2026_07_19_implemented_physical_snapshot_object_offload.md`, change the header status, update every reference, and verify with `rg` that the diff --git a/internal/snapshotoffload/manifest.go b/internal/snapshotoffload/manifest.go index b963d937d..c6a091f66 100644 --- a/internal/snapshotoffload/manifest.go +++ b/internal/snapshotoffload/manifest.go @@ -75,6 +75,9 @@ func DecodeManifest(data []byte) (Manifest, error) { if err := validateManifest(manifest); err != nil { return Manifest{}, err } + if err := verifyManifestSelfHash(manifest); err != nil { + return Manifest{}, err + } return manifest, nil } @@ -100,6 +103,20 @@ func validateManifest(manifest Manifest) error { return nil } +func verifyManifestSelfHash(manifest Manifest) error { + if !isSHA256Hex(manifest.ManifestSHA256) { + return errors.Wrap(ErrInvalidOptions, "manifest sha256 must be 64 lowercase hex characters") + } + _, sum, err := manifest.MarshalCanonical() + if err != nil { + return err + } + if sum != manifest.ManifestSHA256 { + return errors.Wrapf(ErrIntegrity, "manifest sha256 %s, expected %s", manifest.ManifestSHA256, sum) + } + return nil +} + func payloadKey(prefix, sha256Hex string) (string, error) { if !isSHA256Hex(sha256Hex) { return "", errors.Wrap(ErrInvalidOptions, "payload sha256 must be 64 lowercase hex characters") diff --git a/internal/snapshotoffload/offload_test.go b/internal/snapshotoffload/offload_test.go index d437045a9..698e30258 100644 --- a/internal/snapshotoffload/offload_test.go +++ b/internal/snapshotoffload/offload_test.go @@ -3,6 +3,7 @@ package snapshotoffload import ( "bytes" "context" + "encoding/json" "io" "os" "path/filepath" @@ -17,7 +18,7 @@ func TestPublishAndRestorePhysicalSnapshotRoundTrip(t *testing.T) { ctx := context.Background() root := t.TempDir() payload := []byte("EKVTHLC1opaque-physical-snapshot-payload") - sourceDataDir := seedPhysicalSnapshot(t, root, "source", payload, 42, 7, []etcdraftengine.Peer{ + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 42, 7, []etcdraftengine.Peer{ {NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}, {NodeID: 2, ID: "n2", Address: "127.0.0.1:12002"}, }) @@ -78,7 +79,7 @@ func TestRestoreRejectsCorruptPayloadAndLeavesDestinationAbsent(t *testing.T) { ctx := context.Background() root := t.TempDir() payload := []byte("EKVTHLC1payload-before-corruption") - sourceDataDir := seedPhysicalSnapshot(t, root, "source", payload, 10, 3, singlePeer()) + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 10, 3, singlePeer()) store := newTestLocalStore(t, filepath.Join(root, "objects")) manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ Store: store, @@ -90,7 +91,7 @@ func TestRestoreRejectsCorruptPayloadAndLeavesDestinationAbsent(t *testing.T) { require.NoError(t, err) payloadPath, err := store.pathForKey(manifest.Payload.Key) require.NoError(t, err) - require.NoError(t, os.WriteFile(payloadPath, []byte("corrupt"), 0o600)) + require.NoError(t, os.WriteFile(payloadPath, bytes.Repeat([]byte("x"), len(payload)), 0o600)) restoreDataDir := filepath.Join(root, "restored") _, err = RestorePhysicalSnapshot(ctx, RestoreOptions{ @@ -108,7 +109,7 @@ func TestPublishRejectsExistingPayloadMismatchBeforeManifest(t *testing.T) { ctx := context.Background() root := t.TempDir() payload := []byte("EKVTHLC1payload-with-known-sha") - sourceDataDir := seedPhysicalSnapshot(t, root, "source", payload, 11, 4, singlePeer()) + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 11, 4, singlePeer()) store := newTestLocalStore(t, filepath.Join(root, "objects")) payloadKey, err := payloadKey("cluster-a", hexSHA256Bytes(payload)) require.NoError(t, err) @@ -134,19 +135,219 @@ func TestPublishRejectsExistingPayloadMismatchBeforeManifest(t *testing.T) { require.False(t, ok) } +func TestPublishRejectsExistingManifestMismatch(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-with-conflicting-manifest") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 15, 8, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + key, err := manifestKey("cluster-a", 1, 15, 8) + require.NoError(t, err) + bad := []byte(`{"schema_version":1}`) + _, err = store.PutObject(ctx, key, bytes.NewReader(bad), PutOptions{ + Size: int64(len(bad)), + SHA256: hexSHA256Bytes(bad), + ContentType: "application/json", + }) + require.NoError(t, err) + + _, err = PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + CreatedAt: time.Unix(200, 0).UTC(), + }) + require.ErrorIs(t, err, ErrIntegrity) +} + +func TestPublishReusesExistingObjectsWithoutHeadChecksum(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-reused-by-retry") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 16, 9, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + opts := PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + CreatedAt: time.Unix(300, 0).UTC(), + } + + first, err := PublishPersistedSnapshot(ctx, opts) + require.NoError(t, err) + second, err := PublishPersistedSnapshot(ctx, opts) + require.NoError(t, err) + require.Equal(t, first.ManifestKey, second.ManifestKey) + require.Equal(t, first.Payload.Key, second.Payload.Key) +} + +func TestLoadManifestRejectsStaleSelfHash(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-for-stale-manifest-hash") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 12, 5, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + manifestPath, err := store.pathForKey(manifest.ManifestKey) + require.NoError(t, err) + raw, err := os.ReadFile(manifestPath) + require.NoError(t, err) + var body map[string]any + require.NoError(t, json.Unmarshal(raw, &body)) + body["snapshot_index"] = float64(99) + tampered, err := json.MarshalIndent(body, "", " ") + require.NoError(t, err) + require.NoError(t, os.WriteFile(manifestPath, append(tampered, '\n'), 0o600)) + + _, err = LoadManifest(ctx, store, manifest.ManifestKey) + require.ErrorIs(t, err, ErrIntegrity) +} + +func TestLoadManifestRejectsOversizedBody(t *testing.T) { + ctx := context.Background() + store := newTestLocalStore(t, filepath.Join(t.TempDir(), "objects")) + key := "cluster-a/v1/groups/1/snapshots/oversized.json" + body := bytes.Repeat([]byte("x"), maxManifestBytes+1) + _, err := store.PutObject(ctx, key, bytes.NewReader(body), PutOptions{ + Size: int64(len(body)), + SHA256: hexSHA256Bytes(body), + ContentType: "application/json", + }) + require.NoError(t, err) + + _, err = LoadManifest(ctx, store, key) + require.ErrorIs(t, err, ErrInvalidOptions) +} + +func TestLocalStoreRejectsParentDirectoryKeys(t *testing.T) { + ctx := context.Background() + store := newTestLocalStore(t, filepath.Join(t.TempDir(), "objects")) + emptySHA := hexSHA256Bytes(nil) + + _, err := store.PutObject(ctx, "..", bytes.NewReader(nil), PutOptions{SHA256: emptySHA}) + require.ErrorIs(t, err, ErrInvalidOptions) + _, err = store.PutObject(ctx, "a/../..", bytes.NewReader(nil), PutOptions{SHA256: emptySHA}) + require.ErrorIs(t, err, ErrInvalidOptions) +} + +func TestLocalStoreHeadObjectUsesStatMetadata(t *testing.T) { + ctx := context.Background() + store := newTestLocalStore(t, filepath.Join(t.TempDir(), "objects")) + body := []byte("metadata-only-head") + key := "cluster-a/v1/payloads/test.fsm" + putInfo, err := store.PutObject(ctx, key, bytes.NewReader(body), PutOptions{ + Size: int64(len(body)), + SHA256: hexSHA256Bytes(body), + ContentType: "application/octet-stream", + }) + require.NoError(t, err) + require.Equal(t, hexSHA256Bytes(body), putInfo.SHA256) + + headInfo, ok, err := store.HeadObject(ctx, key) + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, int64(len(body)), headInfo.Size) + require.Empty(t, headInfo.SHA256) + + reader, getInfo, err := store.GetObject(ctx, key) + require.NoError(t, err) + defer func() { require.NoError(t, reader.Close()) }() + require.Equal(t, int64(len(body)), getInfo.Size) + require.Empty(t, getInfo.SHA256) + got, err := io.ReadAll(reader) + require.NoError(t, err) + require.Equal(t, body, got) +} + +func TestPublishHonorsCancelledContextWhileSpooling(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-cancelled-before-spool") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 13, 6, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + + _, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.ErrorIs(t, err, context.Canceled) +} + +func TestRestorePreflightsExistingDestinationBeforePayloadDownload(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-before-existing-destination") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 14, 7, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + payloadPath, err := store.pathForKey(manifest.Payload.Key) + require.NoError(t, err) + require.NoError(t, os.WriteFile(payloadPath, []byte("corrupt"), 0o600)) + restoreDataDir := filepath.Join(root, "restored") + require.NoError(t, os.Mkdir(restoreDataDir, 0o755)) + + _, err = RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: store, + ManifestKey: manifest.ManifestKey, + DataDir: restoreDataDir, + Peers: singlePeer(), + }) + require.ErrorIs(t, err, etcdraftengine.ErrExternalSnapshotRestoreExists) +} + +func TestPrepareRestoreDownloadDirCreatesParentAndCleansStaleDirs(t *testing.T) { + root := t.TempDir() + dataDir := filepath.Join(root, "missing-parent", "restored") + + downloadDir, err := prepareRestoreDownloadDir(dataDir) + require.NoError(t, err) + require.DirExists(t, filepath.Dir(dataDir)) + require.Contains(t, downloadDir, filepath.Dir(dataDir)) + require.NoError(t, os.RemoveAll(downloadDir)) + + staleDir := filepath.Join(filepath.Dir(dataDir), ".snapshot-offload-restore-stale") + require.NoError(t, os.MkdirAll(staleDir, 0o755)) + downloadDir, err = prepareRestoreDownloadDir(dataDir) + require.NoError(t, err) + defer func() { require.NoError(t, os.RemoveAll(downloadDir)) }() + _, err = os.Stat(staleDir) + require.True(t, os.IsNotExist(err)) +} + func seedPhysicalSnapshot( t *testing.T, root string, - name string, payload []byte, index uint64, term uint64, peers []etcdraftengine.Peer, ) string { t.Helper() - input := filepath.Join(root, name+".fsm") + input := filepath.Join(root, "source.fsm") require.NoError(t, os.WriteFile(input, payload, 0o600)) - dataDir := filepath.Join(root, name+"-raft") + dataDir := filepath.Join(root, "source-raft") _, err := etcdraftengine.PreparePhysicalSnapshotRestore(etcdraftengine.PhysicalSnapshotRestoreOptions{ InputFSMPath: input, DataDir: dataDir, diff --git a/internal/snapshotoffload/publish.go b/internal/snapshotoffload/publish.go index 81379ee4b..367d52b2b 100644 --- a/internal/snapshotoffload/publish.go +++ b/internal/snapshotoffload/publish.go @@ -112,13 +112,24 @@ func putManifest(ctx context.Context, store ObjectStore, manifest *Manifest) err if err != nil { return err } - if _, err := store.PutObject(ctx, manifest.ManifestKey, bytes.NewReader(data), PutOptions{ + objectSHA := hexSHA256Bytes(data) + if exists, err := verifyExistingStoreObject(ctx, store, manifest.ManifestKey, int64(len(data)), objectSHA); err != nil { + return errors.Wrap(err, "verify existing snapshot manifest") + } else if exists { + manifest.ManifestSHA256 = manifestSHA + return nil + } + info, err := store.PutObject(ctx, manifest.ManifestKey, bytes.NewReader(data), PutOptions{ Size: int64(len(data)), - SHA256: hexSHA256Bytes(data), + SHA256: objectSHA, ContentType: "application/json", - }); err != nil { + }) + if err != nil { return errors.Wrap(err, "put snapshot manifest") } + if info.Size != int64(len(data)) || (info.SHA256 != "" && info.SHA256 != objectSHA) { + return errors.Wrapf(ErrIntegrity, "manifest object %s remote integrity mismatch", manifest.ManifestKey) + } manifest.ManifestSHA256 = manifestSHA return nil } @@ -147,7 +158,10 @@ func spoolExport(ctx context.Context, export *etcdraftengine.PersistedSnapshotEx } }() hash := sha256.New() - n, err := export.WriteTo(io.MultiWriter(tmp, hash)) + n, err := export.WriteTo(contextWriter{ + ctx: ctx, + writer: io.MultiWriter(tmp, hash), + }) if err != nil { return nil, "", n, errors.Wrap(err, "spool persisted snapshot export") } @@ -165,13 +179,10 @@ func spoolExport(ctx context.Context, export *etcdraftengine.PersistedSnapshotEx } func putPayload(ctx context.Context, store ObjectStore, key string, file *os.File, size int64, sha string) error { - if info, ok, err := store.HeadObject(ctx, key); err != nil { - return errors.Wrap(err, "head snapshot payload") - } else if ok { - if info.Size == size && info.SHA256 == sha { - return nil - } - return errors.Wrapf(ErrIntegrity, "payload object %s already exists with different content", key) + if exists, err := verifyExistingStoreObject(ctx, store, key, size, sha); err != nil { + return errors.Wrap(err, "verify existing snapshot payload") + } else if exists { + return nil } if _, err := file.Seek(0, io.SeekStart); err != nil { return errors.WithStack(err) @@ -184,12 +195,53 @@ func putPayload(ctx context.Context, store ObjectStore, key string, file *os.Fil if err != nil { return errors.Wrap(err, "put snapshot payload") } - if info.Size != size || info.SHA256 != sha { + if info.Size != size || (info.SHA256 != "" && info.SHA256 != sha) { return errors.Wrapf(ErrIntegrity, "payload object %s remote integrity mismatch", key) } return nil } +func verifyExistingStoreObject(ctx context.Context, store ObjectStore, key string, size int64, sha string) (bool, error) { + info, ok, err := store.HeadObject(ctx, key) + if err != nil { + return false, errors.Wrap(err, "head existing object") + } + if !ok { + return false, nil + } + if info.Size != size { + return true, errors.Wrapf(ErrIntegrity, "object %s already exists with different size", key) + } + if info.SHA256 != "" { + if info.SHA256 == sha { + return true, nil + } + return true, errors.Wrapf(ErrIntegrity, "object %s already exists with different sha256", key) + } + gotSize, gotSHA, err := hashExistingStoreObject(ctx, store, key) + if err != nil { + return true, err + } + if gotSize == size && gotSHA == sha { + return true, nil + } + return true, errors.Wrapf(ErrIntegrity, "object %s already exists with different content", key) +} + +func hashExistingStoreObject(ctx context.Context, store ObjectStore, key string) (int64, string, error) { + body, _, err := store.GetObject(ctx, key) + if err != nil { + return 0, "", errors.Wrap(err, "get existing object") + } + defer func() { _ = body.Close() }() + sum := sha256.New() + n, err := io.Copy(sum, contextReader{ctx: ctx, reader: body}) + if err != nil { + return 0, "", errors.WithStack(err) + } + return n, hex.EncodeToString(sum.Sum(nil)), nil +} + func stringsTrim(v string) string { return strings.TrimSpace(v) } diff --git a/internal/snapshotoffload/restore.go b/internal/snapshotoffload/restore.go index 85dea5539..2e238b436 100644 --- a/internal/snapshotoffload/restore.go +++ b/internal/snapshotoffload/restore.go @@ -21,6 +21,12 @@ type RestoreOptions struct { Peers []etcdraftengine.Peer } +const ( + maxManifestBytes = 1 << 20 + restoreTempDirMode = 0o755 + restoreTempDirPattern = ".snapshot-offload-restore-*" +) + func LoadManifest(ctx context.Context, store ObjectStore, key string) (Manifest, error) { if store == nil { return Manifest{}, errors.Wrap(ErrInvalidOptions, "object store is required") @@ -30,11 +36,11 @@ func LoadManifest(ctx context.Context, store ObjectStore, key string) (Manifest, return Manifest{}, errors.Wrap(err, "get snapshot manifest") } defer func() { _ = body.Close() }() - var buf bytes.Buffer - if _, err := io.Copy(&buf, contextReader{ctx: ctx, reader: body}); err != nil { - return Manifest{}, errors.WithStack(err) + data, err := readLimitedManifest(ctx, body) + if err != nil { + return Manifest{}, err } - manifest, err := DecodeManifest(buf.Bytes()) + manifest, err := DecodeManifest(data) if err != nil { return Manifest{}, err } @@ -55,9 +61,12 @@ func RestorePhysicalSnapshot(ctx context.Context, opts RestoreOptions) (*etcdraf if err := validateManifest(manifest); err != nil { return nil, err } - downloadDir, err := os.MkdirTemp(filepath.Dir(filepath.Clean(opts.DataDir)), ".snapshot-offload-restore-*") + if err := ensureRestoreDestinationAbsent(opts.DataDir); err != nil { + return nil, err + } + downloadDir, err := prepareRestoreDownloadDir(opts.DataDir) if err != nil { - return nil, errors.WithStack(err) + return nil, err } defer func() { _ = os.RemoveAll(downloadDir) }() payloadPath := filepath.Join(downloadDir, "payload.fsm") @@ -78,6 +87,66 @@ func RestorePhysicalSnapshot(ctx context.Context, opts RestoreOptions) (*etcdraf return result, nil } +func readLimitedManifest(ctx context.Context, body io.Reader) ([]byte, error) { + var buf bytes.Buffer + limited := io.LimitReader(contextReader{ctx: ctx, reader: body}, maxManifestBytes+1) + if _, err := io.Copy(&buf, limited); err != nil { + return nil, errors.WithStack(err) + } + if buf.Len() > maxManifestBytes { + return nil, errors.Wrapf(ErrInvalidOptions, "manifest exceeds %d bytes", maxManifestBytes) + } + return buf.Bytes(), nil +} + +func ensureRestoreDestinationAbsent(dataDir string) error { + cleaned := filepath.Clean(dataDir) + if _, err := os.Stat(cleaned); err == nil { + return errors.Wrapf(etcdraftengine.ErrExternalSnapshotRestoreExists, "destination exists: %s", cleaned) + } else if !os.IsNotExist(err) { + return errors.WithStack(err) + } + return nil +} + +func prepareRestoreDownloadDir(dataDir string) (string, error) { + parent := filepath.Dir(filepath.Clean(dataDir)) + if err := os.MkdirAll(parent, restoreTempDirMode); err != nil { + return "", errors.WithStack(err) + } + if err := cleanupRestoreDownloadDirs(parent); err != nil { + return "", err + } + downloadDir, err := os.MkdirTemp(parent, restoreTempDirPattern) + if err != nil { + return "", errors.WithStack(err) + } + return downloadDir, nil +} + +func cleanupRestoreDownloadDirs(parent string) error { + matches, err := filepath.Glob(filepath.Join(parent, restoreTempDirPattern)) + if err != nil { + return errors.WithStack(err) + } + for _, match := range matches { + info, err := os.Lstat(match) + if err != nil { + if os.IsNotExist(err) { + continue + } + return errors.WithStack(err) + } + if !info.IsDir() { + continue + } + if err := os.RemoveAll(match); err != nil { + return errors.WithStack(err) + } + } + return nil +} + func validateRestoreOptions(opts RestoreOptions) error { switch { case opts.Store == nil: @@ -123,7 +192,10 @@ func downloadVerifiedPayload(ctx context.Context, store ObjectStore, manifest Ma } func validatePayloadInfo(manifest Manifest, info ObjectInfo) error { - if info.Size != manifest.Payload.Bytes || info.SHA256 != manifest.Payload.SHA256 { + if info.Size != manifest.Payload.Bytes { + return errors.Wrapf(ErrIntegrity, "payload head mismatch for %s", manifest.Payload.Key) + } + if info.SHA256 != "" && info.SHA256 != manifest.Payload.SHA256 { return errors.Wrapf(ErrIntegrity, "payload head mismatch for %s", manifest.Payload.Key) } return nil diff --git a/internal/snapshotoffload/store.go b/internal/snapshotoffload/store.go index d091e923d..9c37baf11 100644 --- a/internal/snapshotoffload/store.go +++ b/internal/snapshotoffload/store.go @@ -26,8 +26,10 @@ type PutOptions struct { } type ObjectInfo struct { - Key string - Size int64 + Key string + Size int64 + // SHA256 is optional for metadata-only Head/Get paths; PutObject returns it + // when the writer verified the committed content. SHA256 string } @@ -108,13 +110,30 @@ func (s *LocalStore) pathForKey(key string) (string, error) { return "", errors.Wrap(ErrInvalidOptions, "object store is required") } normalized := normalizeObjectKey(key) - if normalized == "" || normalized == "." || strings.HasPrefix(normalized, "../") { + if normalized == "" || normalized == "." || normalized == ".." || strings.HasPrefix(normalized, "../") { return "", errors.Wrapf(ErrInvalidOptions, "invalid object key %q", key) } return filepath.Join(s.root, filepath.FromSlash(normalized)), nil } func (s *LocalStore) objectInfoForPath(key, objectPath string) (ObjectInfo, error) { + stat, err := os.Stat(objectPath) + if err != nil { + if os.IsNotExist(err) { + return ObjectInfo{}, errors.Wrapf(ErrObjectNotFound, "object %s", key) + } + return ObjectInfo{}, errors.WithStack(err) + } + if !stat.Mode().IsRegular() { + return ObjectInfo{}, errors.Wrapf(ErrInvalidOptions, "object %s is not a regular file", key) + } + return ObjectInfo{ + Key: normalizeObjectKey(key), + Size: stat.Size(), + }, nil +} + +func (s *LocalStore) hashedObjectInfoForPath(key, objectPath string) (ObjectInfo, error) { file, err := os.Open(objectPath) if err != nil { if os.IsNotExist(err) { @@ -196,7 +215,7 @@ func (s *LocalStore) commitTempObject(key, tmpPath, finalPath string, expected O } func (s *LocalStore) verifyExistingObject(key, finalPath string, expected ObjectInfo) (ObjectInfo, error) { - info, err := s.objectInfoForPath(key, finalPath) + info, err := s.hashedObjectInfoForPath(key, finalPath) if err != nil { return ObjectInfo{}, err } @@ -249,6 +268,29 @@ func (r contextReader) Read(p []byte) (int, error) { return n, nil } +type contextWriter struct { + ctx context.Context + writer io.Writer +} + +func (w contextWriter) Write(p []byte) (int, error) { + if w.ctx != nil { + if err := w.ctx.Err(); err != nil { + return 0, errors.WithStack(err) + } + } + n, err := w.writer.Write(p) + if err != nil { + return n, errors.WithStack(err) + } + if w.ctx != nil { + if ctxErr := w.ctx.Err(); ctxErr != nil { + return n, errors.WithStack(ctxErr) + } + } + return n, nil +} + func syncDir(dir string) error { f, err := os.Open(dir) if err != nil { From 0dc76d41a29f054f98c407cd1d71d1abb40c7d47 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 05:58:23 +0900 Subject: [PATCH 3/9] Harden snapshot offload retries --- internal/snapshotoffload/offload_test.go | 75 +++++++++++++++++++- internal/snapshotoffload/publish.go | 90 +++++++++++++++++++++--- internal/snapshotoffload/restore.go | 12 +++- 3 files changed, 164 insertions(+), 13 deletions(-) diff --git a/internal/snapshotoffload/offload_test.go b/internal/snapshotoffload/offload_test.go index 698e30258..83b1ed76b 100644 --- a/internal/snapshotoffload/offload_test.go +++ b/internal/snapshotoffload/offload_test.go @@ -185,6 +185,70 @@ func TestPublishReusesExistingObjectsWithoutHeadChecksum(t *testing.T) { require.Equal(t, first.Payload.Key, second.Payload.Key) } +func TestPublishReusesExistingManifestWhenCreatedAtOmitted(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-retry-with-implicit-created-at") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 17, 10, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + opts := PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + } + + first, err := PublishPersistedSnapshot(ctx, opts) + require.NoError(t, err) + time.Sleep(time.Millisecond) + second, err := PublishPersistedSnapshot(ctx, opts) + require.NoError(t, err) + require.Equal(t, first.ManifestKey, second.ManifestKey) + require.Equal(t, first.ManifestSHA256, second.ManifestSHA256) + require.Equal(t, first.CreatedAt, second.CreatedAt) +} + +func TestPublishRejectsGroupZeroWithoutSourceClusterBeforePayload(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-invalid-group-zero") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 18, 11, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + + _, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 0, + }) + require.ErrorIs(t, err, ErrInvalidOptions) + + payloadKey, err := payloadKey("cluster-a", hexSHA256Bytes(payload)) + require.NoError(t, err) + _, ok, err := store.HeadObject(ctx, payloadKey) + require.NoError(t, err) + require.False(t, ok) +} + +func TestPublishUsesDataDirLocalSpoolByDefault(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-spooled-near-data-dir") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 19, 12, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + + _, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + require.DirExists(t, filepath.Join(filepath.Dir(sourceDataDir), ".snapshot-offload-spool")) +} + func TestLoadManifestRejectsStaleSelfHash(t *testing.T) { ctx := context.Background() root := t.TempDir() @@ -317,7 +381,7 @@ func TestRestorePreflightsExistingDestinationBeforePayloadDownload(t *testing.T) require.ErrorIs(t, err, etcdraftengine.ErrExternalSnapshotRestoreExists) } -func TestPrepareRestoreDownloadDirCreatesParentAndCleansStaleDirs(t *testing.T) { +func TestPrepareRestoreDownloadDirCreatesParentAndCleansOnlyStaleDirs(t *testing.T) { root := t.TempDir() dataDir := filepath.Join(root, "missing-parent", "restored") @@ -327,11 +391,18 @@ func TestPrepareRestoreDownloadDirCreatesParentAndCleansStaleDirs(t *testing.T) require.Contains(t, downloadDir, filepath.Dir(dataDir)) require.NoError(t, os.RemoveAll(downloadDir)) - staleDir := filepath.Join(filepath.Dir(dataDir), ".snapshot-offload-restore-stale") + parent := filepath.Dir(dataDir) + activeDir := filepath.Join(parent, ".snapshot-offload-restore-active") + require.NoError(t, os.MkdirAll(activeDir, 0o755)) + staleDir := filepath.Join(parent, ".snapshot-offload-restore-stale") require.NoError(t, os.MkdirAll(staleDir, 0o755)) + oldTime := time.Now().Add(-restoreTempDirStaleAfter - time.Hour) + require.NoError(t, os.Chtimes(staleDir, oldTime, oldTime)) downloadDir, err = prepareRestoreDownloadDir(dataDir) require.NoError(t, err) defer func() { require.NoError(t, os.RemoveAll(downloadDir)) }() + require.DirExists(t, activeDir) + require.NoError(t, os.RemoveAll(activeDir)) _, err = os.Stat(staleDir) require.True(t, os.IsNotExist(err)) } diff --git a/internal/snapshotoffload/publish.go b/internal/snapshotoffload/publish.go index 367d52b2b..d18b69258 100644 --- a/internal/snapshotoffload/publish.go +++ b/internal/snapshotoffload/publish.go @@ -7,6 +7,8 @@ import ( "encoding/hex" "io" "os" + "path/filepath" + "reflect" "strings" "time" @@ -22,6 +24,7 @@ type PublishOptions struct { SourceCluster string BinaryVersion string CreatedAt time.Time + SpoolDir string } func PublishPersistedSnapshot(ctx context.Context, opts PublishOptions) (*Manifest, error) { @@ -35,7 +38,7 @@ func PublishPersistedSnapshot(ctx context.Context, opts PublishOptions) (*Manife defer func() { _ = export.Close() }() metadata := export.Metadata() - payloadFile, payloadSHA, payloadBytes, err := spoolExport(ctx, export) + payloadFile, payloadSHA, payloadBytes, err := spoolExport(ctx, export, publishSpoolDir(opts)) if err != nil { return nil, err } @@ -57,7 +60,10 @@ func PublishPersistedSnapshot(ctx context.Context, opts PublishOptions) (*Manife if err != nil { return nil, err } - if err := putManifest(ctx, opts.Store, manifest); err != nil { + if err := validateManifest(*manifest); err != nil { + return nil, err + } + if err := putManifest(ctx, opts.Store, manifest, opts.CreatedAt.IsZero()); err != nil { return nil, err } return manifest, nil @@ -107,16 +113,18 @@ func buildManifest( }, nil } -func putManifest(ctx context.Context, store ObjectStore, manifest *Manifest) error { +func putManifest(ctx context.Context, store ObjectStore, manifest *Manifest, reuseExistingCreatedAt bool) error { data, manifestSHA, err := manifest.MarshalCanonical() if err != nil { return err } objectSHA := hexSHA256Bytes(data) - if exists, err := verifyExistingStoreObject(ctx, store, manifest.ManifestKey, int64(len(data)), objectSHA); err != nil { - return errors.Wrap(err, "verify existing snapshot manifest") + if exists, err := verifyExistingManifest(ctx, store, manifest, int64(len(data)), objectSHA, reuseExistingCreatedAt); err != nil { + return err } else if exists { - manifest.ManifestSHA256 = manifestSHA + if manifest.ManifestSHA256 == "" { + manifest.ManifestSHA256 = manifestSHA + } return nil } info, err := store.PutObject(ctx, manifest.ManifestKey, bytes.NewReader(data), PutOptions{ @@ -134,19 +142,85 @@ func putManifest(ctx context.Context, store ObjectStore, manifest *Manifest) err return nil } +func verifyExistingManifest( + ctx context.Context, + store ObjectStore, + manifest *Manifest, + size int64, + sha string, + reuseExistingCreatedAt bool, +) (bool, error) { + info, ok, err := store.HeadObject(ctx, manifest.ManifestKey) + if err != nil { + return false, errors.Wrap(err, "head existing snapshot manifest") + } + if !ok { + return false, nil + } + if matches, err := existingStoreObjectMatches(ctx, store, manifest.ManifestKey, info, size, sha); err != nil { + return true, errors.Wrap(err, "verify existing snapshot manifest") + } else if matches { + return true, nil + } + if !reuseExistingCreatedAt { + return true, errors.Wrapf(ErrIntegrity, "manifest object %s already exists with different content", manifest.ManifestKey) + } + existing, err := LoadManifest(ctx, store, manifest.ManifestKey) + if err != nil { + return true, errors.Wrap(err, "load existing snapshot manifest") + } + if !sameManifestExceptCreation(existing, *manifest) { + return true, errors.Wrapf(ErrIntegrity, "manifest object %s already exists with different content", manifest.ManifestKey) + } + *manifest = existing + return true, nil +} + +func existingStoreObjectMatches(ctx context.Context, store ObjectStore, key string, info ObjectInfo, size int64, sha string) (bool, error) { + if info.Size != size { + return false, nil + } + if info.SHA256 != "" { + return info.SHA256 == sha, nil + } + gotSize, gotSHA, err := hashExistingStoreObject(ctx, store, key) + if err != nil { + return false, err + } + return gotSize == size && gotSHA == sha, nil +} + +func sameManifestExceptCreation(existing Manifest, candidate Manifest) bool { + candidate.CreatedAt = existing.CreatedAt + candidate.ManifestSHA256 = existing.ManifestSHA256 + return reflect.DeepEqual(existing, candidate) +} + func validatePublishOptions(opts PublishOptions) error { switch { case opts.Store == nil: return errors.Wrap(ErrInvalidOptions, "object store is required") case stringsTrim(opts.DataDir) == "": return errors.Wrap(ErrInvalidOptions, "data dir is required") + case opts.GroupID == 0 && stringsTrim(opts.SourceCluster) == "": + return errors.Wrap(ErrInvalidOptions, "source cluster is required for group 0 manifests") default: return nil } } -func spoolExport(ctx context.Context, export *etcdraftengine.PersistedSnapshotExport) (*os.File, string, int64, error) { - tmp, err := os.CreateTemp("", "elastickv-snapshot-offload-*.fsm") +func publishSpoolDir(opts PublishOptions) string { + if stringsTrim(opts.SpoolDir) != "" { + return filepath.Clean(opts.SpoolDir) + } + return filepath.Join(filepath.Dir(filepath.Clean(opts.DataDir)), ".snapshot-offload-spool") +} + +func spoolExport(ctx context.Context, export *etcdraftengine.PersistedSnapshotExport, spoolDir string) (*os.File, string, int64, error) { + if err := os.MkdirAll(spoolDir, localStoreDirPerm); err != nil { + return nil, "", 0, errors.WithStack(err) + } + tmp, err := os.CreateTemp(spoolDir, "elastickv-snapshot-offload-*.fsm") if err != nil { return nil, "", 0, errors.WithStack(err) } diff --git a/internal/snapshotoffload/restore.go b/internal/snapshotoffload/restore.go index 2e238b436..8f979b820 100644 --- a/internal/snapshotoffload/restore.go +++ b/internal/snapshotoffload/restore.go @@ -8,6 +8,7 @@ import ( "io" "os" "path/filepath" + "time" etcdraftengine "github.com/bootjp/elastickv/internal/raftengine/etcd" "github.com/cockroachdb/errors" @@ -22,9 +23,10 @@ type RestoreOptions struct { } const ( - maxManifestBytes = 1 << 20 - restoreTempDirMode = 0o755 - restoreTempDirPattern = ".snapshot-offload-restore-*" + maxManifestBytes = 1 << 20 + restoreTempDirMode = 0o755 + restoreTempDirPattern = ".snapshot-offload-restore-*" + restoreTempDirStaleAfter = 7 * 24 * time.Hour ) func LoadManifest(ctx context.Context, store ObjectStore, key string) (Manifest, error) { @@ -129,6 +131,7 @@ func cleanupRestoreDownloadDirs(parent string) error { if err != nil { return errors.WithStack(err) } + staleBefore := time.Now().Add(-restoreTempDirStaleAfter) for _, match := range matches { info, err := os.Lstat(match) if err != nil { @@ -140,6 +143,9 @@ func cleanupRestoreDownloadDirs(parent string) error { if !info.IsDir() { continue } + if info.ModTime().After(staleBefore) { + continue + } if err := os.RemoveAll(match); err != nil { return errors.WithStack(err) } From 0288808b86572e82d5d6040cb1f9b46498b21c21 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 07:13:30 +0900 Subject: [PATCH 4/9] raft: avoid snapshot cleanup blocking ready loop --- internal/raftengine/etcd/engine.go | 4 +++- .../etcd/engine_applied_index_test.go | 24 +++++++++++++++++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/internal/raftengine/etcd/engine.go b/internal/raftengine/etcd/engine.go index 7017caceb..dfff8a8d6 100644 --- a/internal/raftengine/etcd/engine.go +++ b/internal/raftengine/etcd/engine.go @@ -3365,7 +3365,9 @@ func (e *Engine) releaseProtectedReceivedFSMSnapshotsUpTo(index uint64) { if index == 0 { return } - e.snapshotMu.Lock() + if !e.snapshotMu.TryLock() { + return + } defer e.snapshotMu.Unlock() e.releaseProtectedReceivedFSMSnapshotsUpToLocked(index) } diff --git a/internal/raftengine/etcd/engine_applied_index_test.go b/internal/raftengine/etcd/engine_applied_index_test.go index 7111f348b..2e67c2836 100644 --- a/internal/raftengine/etcd/engine_applied_index_test.go +++ b/internal/raftengine/etcd/engine_applied_index_test.go @@ -162,6 +162,30 @@ func TestPersistReadyWithSnapshotHoldsSnapshotMuThroughSaveSnap(t *testing.T) { require.Empty(t, e.protectedReceivedFSMSnaps) } +func TestReleaseProtectedReceivedFSMSnapshotsUpToDoesNotBlockSnapshotMu(t *testing.T) { + e := &Engine{ + protectedReceivedFSMSnaps: map[uint64]int{7: 1}, + } + e.snapshotMu.Lock() + + done := make(chan struct{}) + go func() { + e.releaseProtectedReceivedFSMSnapshotsUpTo(7) + close(done) + }() + + select { + case <-done: + case <-time.After(100 * time.Millisecond): + t.Fatal("releaseProtectedReceivedFSMSnapshotsUpTo blocked on snapshotMu") + } + require.Equal(t, map[uint64]int{7: 1}, e.protectedReceivedFSMSnaps) + + e.snapshotMu.Unlock() + e.releaseProtectedReceivedFSMSnapshotsUpTo(7) + require.Empty(t, e.protectedReceivedFSMSnaps) +} + func TestProtectReceivedFSMSnapshotRechecksAppliedIndexUnderLock(t *testing.T) { e := &Engine{} e.snapshotMu.Lock() From bd92cd78f6adca118252c3d21cd17e8eecf5acc5 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 07:13:50 +0900 Subject: [PATCH 5/9] redis: reduce idle blocking poll pressure --- adapter/redis.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/adapter/redis.go b/adapter/redis.go index a2b9b8c39..2ac41a182 100644 --- a/adapter/redis.go +++ b/adapter/redis.go @@ -120,9 +120,9 @@ const ( // blocking-command wait loops when no in-process write signal arrives. // Signals cover normal XADD / ZADD / ZINCRBY wakeups immediately; this // interval only bounds missed-signal, wrong-type, and BLOCK-deadline checks. - // Keep it below common sub-second BLOCK budgets while reducing idle - // fallback scans versus the previous 100ms cadence. - defaultRedisBlockWaitFallback = 250 * time.Millisecond + // Keep normal producer wakeups event-driven while reducing idle fallback scans + // from blocked consumers waiting on empty collections. + defaultRedisBlockWaitFallback = time.Second redisFlushLegacyTimeout = 10 * time.Minute redisRelayPublishTimeout = 2 * time.Second redisTraceArgLimit = 6 From 958d56da2707aeef6291dfcae62b3f4cd9c8db91 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 09:25:27 +0900 Subject: [PATCH 6/9] redis: raise proxy connection capacity --- adapter/redis_peer_limiter.go | 2 +- adapter/redis_peer_limiter_test.go | 10 ++- cmd/redis-proxy/main_test.go | 8 +-- deploy/redis-proxy/docker-compose.ha.yml | 4 +- docs/redis-proxy-deployment.md | 20 +++--- kv/coordinator.go | 29 ++++++-- kv/keyviz_label.go | 26 ++++++++ kv/lease_read_test.go | 18 ++++- kv/lease_warmup_test.go | 20 ++++++ kv/sharded_coordinator.go | 85 ++++++++++++++++++++---- kv/tso.go | 72 ++++++++++++++++++-- kv/tso_test.go | 65 +++++++++++++++++- proxy/backend.go | 9 +-- proxy/proxy_test.go | 2 +- 14 files changed, 320 insertions(+), 50 deletions(-) diff --git a/adapter/redis_peer_limiter.go b/adapter/redis_peer_limiter.go index d1bace22c..c9d44329a 100644 --- a/adapter/redis_peer_limiter.go +++ b/adapter/redis_peer_limiter.go @@ -10,7 +10,7 @@ import ( const ( redisPerPeerLimitEnv = "ELASTICKV_REDIS_PER_PEER_CONNECTIONS" - defaultRedisPerPeerConnectionCap = 8 + defaultRedisPerPeerConnectionCap = 64 redisPeerLimitError = "ERR max connections per client exceeded" unknownRedisPeer = "unknown" ) diff --git a/adapter/redis_peer_limiter_test.go b/adapter/redis_peer_limiter_test.go index ea631375f..74312d7fb 100644 --- a/adapter/redis_peer_limiter_test.go +++ b/adapter/redis_peer_limiter_test.go @@ -7,9 +7,17 @@ import ( ) const ( - testPeerLimit = 2 + testPeerLimit = 2 + defaultElasticKVProxyPoolSizeForTest = 64 ) +func TestRedisPeerLimiterDefaultMatchesProxyPool(t *testing.T) { + t.Setenv(redisPerPeerLimitEnv, "") + limiter := newDefaultRedisPeerLimiter() + require.NotNil(t, limiter) + require.Equal(t, defaultElasticKVProxyPoolSizeForTest, limiter.limit) +} + func TestRedisPeerLimiterRejectsAndReleases(t *testing.T) { server := NewRedisServer(nil, "", nil, nil, nil, nil, WithRedisPerPeerConnectionLimit(testPeerLimit)) c1 := &remoteCommandRecorder{remote: "192.168.0.64:10001"} diff --git a/cmd/redis-proxy/main_test.go b/cmd/redis-proxy/main_test.go index 0330ae8eb..7188637f9 100644 --- a/cmd/redis-proxy/main_test.go +++ b/cmd/redis-proxy/main_test.go @@ -72,10 +72,10 @@ func TestDeriveSecondaryConcurrency(t *testing.T) { name: "dual write derives from ElasticKV pool", mode: proxy.ModeDualWrite, primaryPoolSize: 128, - elasticKVPoolSize: 4, - wantWriteConcurrency: 2, - wantScriptConcurrency: 1, - wantBlockingConcurrency: 2, + elasticKVPoolSize: 64, + wantWriteConcurrency: 32, + wantScriptConcurrency: 16, + wantBlockingConcurrency: 20, }, { name: "shadow mode derives from ElasticKV pool", diff --git a/deploy/redis-proxy/docker-compose.ha.yml b/deploy/redis-proxy/docker-compose.ha.yml index c150fd9d1..1376c0a84 100644 --- a/deploy/redis-proxy/docker-compose.ha.yml +++ b/deploy/redis-proxy/docker-compose.ha.yml @@ -26,7 +26,7 @@ services: - -listen=:6379 - -primary=${REDIS_PROXY_PRIMARY:-redis:6379} - -secondary=${REDIS_PROXY_SECONDARY:-elastickv:6380} - - -elastickv-pool-size=${REDIS_PROXY_ELASTICKV_POOL_SIZE:-4} + - -elastickv-pool-size=${REDIS_PROXY_ELASTICKV_POOL_SIZE:-64} - -mode=${REDIS_PROXY_MODE:-dual-write-shadow} - -metrics=:9191 networks: @@ -46,7 +46,7 @@ services: - -listen=:6379 - -primary=${REDIS_PROXY_PRIMARY:-redis:6379} - -secondary=${REDIS_PROXY_SECONDARY:-elastickv:6380} - - -elastickv-pool-size=${REDIS_PROXY_ELASTICKV_POOL_SIZE:-4} + - -elastickv-pool-size=${REDIS_PROXY_ELASTICKV_POOL_SIZE:-64} - -mode=${REDIS_PROXY_MODE:-dual-write-shadow} - -metrics=:9191 networks: diff --git a/docs/redis-proxy-deployment.md b/docs/redis-proxy-deployment.md index 003597f2d..743fc5b8d 100644 --- a/docs/redis-proxy-deployment.md +++ b/docs/redis-proxy-deployment.md @@ -35,7 +35,7 @@ go build -o redis-proxy ./cmd/redis-proxy/ | `-secondary-db` | `0` | Secondary Redis DB number | | `-secondary-password` | (empty) | Secondary Redis password | | `-primary-pool-size` | `128` | Primary Redis backend connection pool size | -| `-elastickv-pool-size` | `4` | ElasticKV backend connection pool size | +| `-elastickv-pool-size` | `64` | ElasticKV backend connection pool size | | `-secondary-write-concurrency` | `0` | Shared maximum for all asynchronous secondary writes, including scripts. `0` derives half of the secondary backend pool size, minimum `1` | | `-secondary-script-concurrency` | `0` | Lua-script sublimit within `-secondary-write-concurrency`. `0` derives half of the shared write limit, minimum `1` | | `-secondary-write-queue-size` | `0` | Bounded queue for non-script secondary writes. `0` derives `64 * concurrency`, clamped to `64..8192` | @@ -94,9 +94,9 @@ docker run --rm \ -primary redis.internal:6379 \ -primary-password "${REDIS_PASSWORD}" \ -secondary elastickv.internal:6380 \ - -elastickv-pool-size 4 \ - -secondary-write-concurrency 2 \ - -secondary-script-concurrency 1 \ + -elastickv-pool-size 64 \ + -secondary-write-concurrency 32 \ + -secondary-script-concurrency 16 \ -mode dual-write-shadow \ -secondary-timeout 5s \ -shadow-timeout 3s \ @@ -118,9 +118,9 @@ services: - -listen=:6479 - -primary=redis:6379 - -secondary=elastickv:6380 - - -elastickv-pool-size=4 - - -secondary-write-concurrency=2 - - -secondary-script-concurrency=1 + - -elastickv-pool-size=64 + - -secondary-write-concurrency=32 + - -secondary-script-concurrency=16 - -mode=dual-write-shadow - -metrics=:9191 depends_on: @@ -213,7 +213,7 @@ Override backend wiring via env vars before `docker compose up`: ```bash REDIS_PROXY_PRIMARY=redis.prod.internal:6379 \ REDIS_PROXY_SECONDARY=elastickv-1.prod.internal:6380,elastickv-2.prod.internal:6380,elastickv-3.prod.internal:6380 \ -REDIS_PROXY_ELASTICKV_POOL_SIZE=4 \ +REDIS_PROXY_ELASTICKV_POOL_SIZE=64 \ REDIS_PROXY_MODE=dual-write-shadow \ docker compose -f docker-compose.ha.yml up -d ``` @@ -414,7 +414,7 @@ groups: | Parameter | Value | Description | |-----------|-------|-------------| | Redis connection pool size | 128 | Default go-redis pool size for Redis | -| ElasticKV connection pool size | 4 | Default per-leader pool; keep within the server per-peer connection limit | +| ElasticKV connection pool size | 64 | Default per-leader pool; keep within the server per-peer connection limit | | Dial timeout | 5s | Backend connection timeout | | Read timeout | 3s | Backend read timeout | | Write timeout | 3s | Backend write timeout | @@ -439,7 +439,7 @@ Recommended shutdown order: `redis-proxy -> application -> Redis / ElasticKV`. ### Secondary writes are falling behind - Check `proxy_async_queue_depth`, `proxy_async_queue_delay_seconds`, and `proxy_async_drops_by_queue_total` before increasing concurrency. - Check `proxy_backend_pool_pending_requests` and the `waits`/`timeouts` pool events. Pool waits mean concurrency is too high for the configured pool. -- Increase the ElasticKV pool only together with `ELASTICKV_REDIS_PER_PEER_CONNECTIONS`; keep `-secondary-write-concurrency` at or below the pool size. +- Keep `ELASTICKV_REDIS_PER_PEER_CONNECTIONS` at least as high as `-elastickv-pool-size`; keep `-secondary-write-concurrency` at or below the pool size. - A sustained `expired` rate means secondary throughput is below ingress. Increasing queue size only delays the loss; profile ElasticKV before raising concurrency. ### High divergence count diff --git a/kv/coordinator.go b/kv/coordinator.go index 7fa2ff18f..093d77686 100644 --- a/kv/coordinator.go +++ b/kv/coordinator.go @@ -7,6 +7,7 @@ import ( "log/slog" "reflect" "strings" + "sync" "time" "github.com/bootjp/elastickv/internal/monoclock" @@ -216,6 +217,7 @@ type Coordinate struct { // hlcRenewalBlocked lets startup code temporarily suppress background HLC // lease proposals while another startup-only Raft mutation must run first. hlcRenewalBlocked func() bool + hlcRecoveryMu sync.Mutex // deregisterLeaseCb removes the leader-loss callback registered // against engine at construction. Long-lived Coordinates don't // need to call it (the engine will be closed after them), but @@ -813,6 +815,27 @@ func (c *Coordinate) ProposeHLCLease(ctx context.Context, ceilingMs int64) error return nil } +func (c *Coordinate) RecoverHLCLease(ctx context.Context) error { + if c == nil || c.clock == nil { + return errors.WithStack(errHLCLeaseRecoveryUnavailable) + } + c.hlcRecoveryMu.Lock() + defer c.hlcRecoveryMu.Unlock() + if !hlcCeilingExpired(c.clock) { + return nil + } + if !c.IsLeaderAcceptingWrites() { + return errors.WithStack(errHLCLeaseRecoveryUnavailable) + } + if c.hlcRenewalBlocked != nil && c.hlcRenewalBlocked() { + return errors.WithStack(errHLCLeaseRecoveryBlocked) + } + rctx, cancel := hlcRecoveryContext(ctx) + defer cancel() + ceilingMs := time.Now().UnixMilli() + hlcPhysicalWindowMs + return errors.Wrap(c.ProposeHLCLease(rctx, ceilingMs), "recover hlc lease") +} + // extendLeaseAfterRenewal warms the read lease after a successful HLC // ceiling propose. It is the renewal-path counterpart of // refreshLeaseAfterDispatch's success branch: the propose was a @@ -1026,11 +1049,7 @@ func (c *Coordinate) allocateTimestampAfter(ctx context.Context, label string, m if min > 0 { c.clock.Observe(min) } - ts, err := c.clock.NextFenced() - if err != nil { - return 0, errors.Wrap(err, label) - } - return ts, nil + return nextFencedWithRecovery(ctx, c.clock, c, label) } // Next makes Coordinate usable as a TimestampAllocator for adapter helpers. diff --git a/kv/keyviz_label.go b/kv/keyviz_label.go index 1c056eae9..80d727638 100644 --- a/kv/keyviz_label.go +++ b/kv/keyviz_label.go @@ -90,6 +90,32 @@ func (c keyVizLabeledCoordinator) RaftLeaderForKey(key []byte) string { func (c keyVizLabeledCoordinator) Clock() *HLC { return c.inner.Clock() } +func (c keyVizLabeledCoordinator) Next(ctx context.Context) (uint64, error) { + ctx = contextWithKeyVizLabel(ctx, c.label) + if alloc, ok := c.inner.(TimestampAllocator); ok { + ts, err := alloc.Next(ctx) + return ts, errors.WithStack(err) + } + return NextTimestampThrough(ctx, c.inner, "allocate keyviz-labeled timestamp") +} + +func (c keyVizLabeledCoordinator) NextAfter(ctx context.Context, min uint64) (uint64, error) { + ctx = contextWithKeyVizLabel(ctx, c.label) + if after, ok := c.inner.(TimestampAfterAllocator); ok { + ts, err := after.NextAfter(ctx, min) + return ts, errors.WithStack(err) + } + return NextTimestampAfterThrough(ctx, c.inner, min, "allocate keyviz-labeled timestamp after observed ts") +} + +func (c keyVizLabeledCoordinator) RecoverHLCLease(ctx context.Context) error { + ctx = contextWithKeyVizLabel(ctx, c.label) + if recoverer, ok := c.inner.(hlcLeaseRecoverer); ok { + return errors.WithStack(recoverer.RecoverHLCLease(ctx)) + } + return errors.WithStack(errHLCLeaseRecoveryUnavailable) +} + func (c keyVizLabeledCoordinator) LeaseRead(ctx context.Context) (uint64, error) { if lr, ok := c.inner.(LeaseReadableCoordinator); ok { idx, err := lr.LeaseRead(ctx) diff --git a/kv/lease_read_test.go b/kv/lease_read_test.go index d62c103d2..19bb598f2 100644 --- a/kv/lease_read_test.go +++ b/kv/lease_read_test.go @@ -2,6 +2,7 @@ package kv import ( "context" + "encoding/binary" "errors" "sync" "sync/atomic" @@ -24,6 +25,7 @@ type fakeLeaseEngine struct { proposeErr error // when set, Propose returns it (warm-up failure tests) proposeCalls atomic.Int32 proposeHook func() // invoked inside Propose before returning (race injection) + proposeApply func([]byte) // invoked after a successful propose (FSM apply simulation) state atomic.Value // stores raftengine.State; default Leader lastQuorumAckMonoNs atomic.Int64 // 0 = no ack yet. Updated by setQuorumAck(). leaderLossCallbacksMu sync.Mutex @@ -62,7 +64,7 @@ func (e *fakeLeaseEngine) Status() raftengine.Status { func (e *fakeLeaseEngine) Configuration(context.Context) (raftengine.Configuration, error) { return raftengine.Configuration{}, nil } -func (e *fakeLeaseEngine) Propose(context.Context, []byte) (*raftengine.ProposalResult, error) { +func (e *fakeLeaseEngine) Propose(_ context.Context, data []byte) (*raftengine.ProposalResult, error) { e.proposeCalls.Add(1) if e.proposeHook != nil { e.proposeHook() @@ -70,6 +72,9 @@ func (e *fakeLeaseEngine) Propose(context.Context, []byte) (*raftengine.Proposal if e.proposeErr != nil { return nil, e.proposeErr } + if e.proposeApply != nil { + e.proposeApply(data) + } return &raftengine.ProposalResult{}, nil } func (e *fakeLeaseEngine) ProposeAdmin(ctx context.Context, data []byte) (*raftengine.ProposalResult, error) { @@ -183,6 +188,17 @@ func (e *fakeLeaseEngine) setQuorumAck(i monoclock.Instant) { e.lastQuorumAckMonoNs.Store(i.Nanos()) } +func applyHLCLeaseEntryToClock(t testing.TB, clock *HLC) func([]byte) { + t.Helper() + return func(data []byte) { + if len(data) != hlcLeaseEntryLen || data[0] != raftEncodeHLCLease { + return + } + ceilingMs := int64(binary.BigEndian.Uint64(data[1:])) //nolint:gosec // encoded from a positive Unix ms timestamp. + clock.SetPhysicalCeiling(ceilingMs) + } +} + // --- Coordinate.LeaseRead ----------------------------------------------- // TestCoordinate_LeaseRead_EngineAckFastPath covers the engine-driven diff --git a/kv/lease_warmup_test.go b/kv/lease_warmup_test.go index 5558a913c..206875bf6 100644 --- a/kv/lease_warmup_test.go +++ b/kv/lease_warmup_test.go @@ -294,6 +294,26 @@ func TestShardedCoordinator_RenewHLCLeases_ProposesToEveryLedGroup(t *testing.T) "the non-default group lease must be warmed by all-group renewal") } +func TestShardedCoordinator_RecoverHLCLease_ProposesToEveryLedGroup(t *testing.T) { + t.Parallel() + clock := NewHLC() + clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) + eng1 := newShardedLeaseEngine(100) + eng2 := newShardedLeaseEngine(200) + eng1.proposeApply = applyHLCLeaseEntryToClock(t, clock) + eng2.proposeApply = applyHLCLeaseEntryToClock(t, clock) + coord := mustShardedLeaseCoord(t, eng1, eng2) + coord.clock = clock + + require.NoError(t, coord.RecoverHLCLease(context.Background())) + require.Equal(t, int32(1), eng1.proposeCalls.Load()) + require.Equal(t, int32(1), eng2.proposeCalls.Load()) + + got, err := clock.NextFenced() + require.NoError(t, err) + require.NotZero(t, got) +} + func TestShardedCoordinator_RenewHLCLeases_SkipsNonLeaders(t *testing.T) { t.Parallel() eng1 := newShardedLeaseEngine(100) diff --git a/kv/sharded_coordinator.go b/kv/sharded_coordinator.go index da532dcb2..d3a705a5b 100644 --- a/kv/sharded_coordinator.go +++ b/kv/sharded_coordinator.go @@ -408,6 +408,7 @@ type ShardedCoordinator struct { // hlcRenewalBlocked lets startup code temporarily suppress background HLC // lease proposals while another startup-only Raft mutation must run first. hlcRenewalBlocked func() bool + hlcRecoveryMu sync.Mutex // hlcRenewalInFlight prevents a slow or quorum-stalled group from stacking // another background HLC lease proposal for the same group on the next tick. hlcRenewalMu sync.Mutex @@ -1509,11 +1510,7 @@ func (c *ShardedCoordinator) allocateTimestampAfter(ctx context.Context, label s if min > 0 { c.clock.Observe(min) } - ts, err := c.clock.NextFenced() - if err != nil { - return 0, errors.Wrap(err, label) - } - return ts, nil + return nextFencedWithRecovery(ctx, c.clock, c, label) } // Next makes ShardedCoordinator usable as a TimestampAllocator for adapter @@ -2352,7 +2349,7 @@ func (c *ShardedCoordinator) renewHLCLeases(ctx context.Context) <-chan struct{} if group == nil || group.Engine == nil { continue } - if group.Engine.State() != raftengine.StateLeader { + if !shardGroupAcceptingWrites(group) { continue } if !c.startHLCLeaseRenewal(gid) { @@ -2416,8 +2413,76 @@ func (c *ShardedCoordinator) finishHLCLeaseRenewal(gid uint64) { // the callback latency window. Non-leadership errors (no quorum, // validation) are NOT leadership signals and must not tear down a warm // lease -- doing so would force every read onto the slow path. +func (c *ShardedCoordinator) RecoverHLCLease(ctx context.Context) error { + if c == nil || c.clock == nil { + return errors.WithStack(errHLCLeaseRecoveryUnavailable) + } + c.hlcRecoveryMu.Lock() + defer c.hlcRecoveryMu.Unlock() + if !hlcCeilingExpired(c.clock) { + return nil + } + if c.hlcRenewalBlocked != nil && c.hlcRenewalBlocked() { + return errors.WithStack(errHLCLeaseRecoveryBlocked) + } + targets := c.hlcLeaseRecoveryTargets() + if len(targets) == 0 { + return errors.WithStack(errHLCLeaseRecoveryUnavailable) + } + rctx, cancel := hlcRecoveryContext(ctx) + defer cancel() + ceilingMs := time.Now().UnixMilli() + hlcPhysicalWindowMs + for _, target := range targets { + if err := c.proposeHLCLeaseForGroup(rctx, target.group, ceilingMs); err != nil { + return errors.Wrapf(err, "recover hlc lease group %d", target.gid) + } + } + return nil +} + +type hlcLeaseRecoveryTarget struct { + gid uint64 + group *ShardGroup +} + +func (c *ShardedCoordinator) hlcLeaseRecoveryTargets() []hlcLeaseRecoveryTarget { + if c == nil { + return nil + } + if c.timestampGroupConfigured { + group := c.groups[c.timestampGroup] + if !shardGroupAcceptingWrites(group) { + return nil + } + return []hlcLeaseRecoveryTarget{{gid: c.timestampGroup, group: group}} + } + ids := c.timestampBridgeCandidateGroupIDs() + targets := make([]hlcLeaseRecoveryTarget, 0, len(ids)) + for _, gid := range ids { + group := c.groups[gid] + if shardGroupAcceptingWrites(group) { + targets = append(targets, hlcLeaseRecoveryTarget{gid: gid, group: group}) + } + } + return targets +} + +func shardGroupAcceptingWrites(group *ShardGroup) bool { + return group != nil && group.Engine != nil && isLeaderAcceptingWrites(group.Engine) +} + func (c *ShardedCoordinator) renewHLCLease(ctx context.Context, gid uint64, group *ShardGroup) { ceilingMs := time.Now().UnixMilli() + hlcPhysicalWindowMs + if err := c.proposeHLCLeaseForGroup(ctx, group, ceilingMs); err != nil { + c.logger().WarnContext(ctx, "hlc lease renewal failed", + slog.Uint64("group_id", gid), + slog.Int64("ceiling_ms", ceilingMs), + slog.Any("err", err), + ) + } +} + +func (c *ShardedCoordinator) proposeHLCLeaseForGroup(ctx context.Context, group *ShardGroup, ceilingMs int64) error { start := monoclock.Now() expectedGen := group.lease.generation() // Route through the ShardGroup's wrap-aware proposer chain — NOT a @@ -2430,16 +2495,12 @@ func (c *ShardedCoordinator) renewHLCLease(ctx context.Context, gid uint64, grou if isLeadershipLossError(err) { group.lease.invalidate() } - c.logger().WarnContext(ctx, "hlc lease renewal failed", - slog.Uint64("group_id", gid), - slog.Int64("ceiling_ms", ceilingMs), - slog.Any("err", err), - ) - return + return errors.WithStack(err) } if lp, ok := group.Engine.(raftengine.LeaseProvider); ok { group.lease.extend(start.Add(lp.LeaseDuration()), expectedGen) } + return nil } func keyMutations(muts []*pb.Mutation) []*pb.Mutation { diff --git a/kv/tso.go b/kv/tso.go index e07cb34f9..292c5d852 100644 --- a/kv/tso.go +++ b/kv/tso.go @@ -16,6 +16,9 @@ var ( ErrTSOCoordinatorNil = errors.New("tso: coordinator is required") ErrTSOClockNil = errors.New("tso: coordinator clock is nil") ErrInvalidTSOBatchSize = errors.New("tso: invalid batch size") + + errHLCLeaseRecoveryUnavailable = errors.New("hlc lease recovery unavailable") + errHLCLeaseRecoveryBlocked = errors.New("hlc lease recovery blocked") ) // TSOAllocator issues globally monotonic timestamps. NextBatch returns the @@ -42,6 +45,10 @@ type tsoBatchAfterAllocator interface { NextBatchAfter(ctx context.Context, n int, min uint64) (uint64, error) } +type hlcLeaseRecoverer interface { + RecoverHLCLease(context.Context) error +} + type timestampIssuer interface { IsTimestampLeader() bool } @@ -66,11 +73,8 @@ func NextTimestampThrough(ctx context.Context, coord Coordinator, label string) if coord == nil || coord.Clock() == nil { return 1, nil } - ts, err := coord.Clock().NextFenced() - if err != nil { - return 0, errors.Wrap(err, label) - } - return ts, nil + recoverer, _ := coord.(hlcLeaseRecoverer) + return nextFencedWithRecovery(ctx, coord.Clock(), recoverer, label) } // NextTimestampAfterThrough allocates a timestamp strictly greater than @@ -169,6 +173,60 @@ func nextTimestampAfterFallback(startTS uint64) (uint64, error) { return nextTS, nil } +func nextFencedWithRecovery(ctx context.Context, clock *HLC, recoverer hlcLeaseRecoverer, label string) (uint64, error) { + ts, err := clock.NextFenced() + if err == nil { + return ts, nil + } + if !errors.Is(err, ErrCeilingExpired) || recoverer == nil { + return 0, errors.Wrap(err, label) + } + if recoverErr := recoverer.RecoverHLCLease(ctx); recoverErr != nil { + return 0, errors.Wrapf(err, "%s: on-demand HLC lease renewal failed: %v", label, recoverErr) + } + ts, err = clock.NextFenced() + if err != nil { + return 0, errors.Wrap(err, label) + } + return ts, nil +} + +func nextBatchFencedWithRecovery(ctx context.Context, clock *HLC, n int, recoverer hlcLeaseRecoverer, label string) (uint64, error) { + base, err := clock.NextBatchFenced(n) + if err == nil { + return base, nil + } + if !errors.Is(err, ErrCeilingExpired) || recoverer == nil { + return 0, errors.Wrap(err, label) + } + if recoverErr := recoverer.RecoverHLCLease(ctx); recoverErr != nil { + return 0, errors.Wrapf(err, "%s: on-demand HLC lease renewal failed: %v", label, recoverErr) + } + base, err = clock.NextBatchFenced(n) + if err != nil { + return 0, errors.Wrap(err, label) + } + return base, nil +} + +func hlcRecoveryContext(ctx context.Context) (context.Context, context.CancelFunc) { + if ctx == nil { + ctx = context.Background() + } + if _, ok := ctx.Deadline(); ok { + return ctx, func() {} + } + return context.WithTimeout(ctx, dispatchLeaderRetryBudget) +} + +func hlcCeilingExpired(clock *HLC) bool { + if clock == nil { + return false + } + ceiling := clock.PhysicalCeiling() + return ceiling > 0 && time.Now().UnixMilli() >= ceiling +} + type tsoCoordinator interface { IsLeader() bool Clock() *HLC @@ -240,8 +298,8 @@ func (a *LocalTSOAllocator) nextBatchAfter(ctx context.Context, n int, min uint6 if min > 0 { clock.Observe(min) } - base, err := clock.NextBatchFenced(n) - return base, errors.Wrap(err, "tso next batch") + recoverer, _ := a.coord.(hlcLeaseRecoverer) + return nextBatchFencedWithRecovery(ctx, clock, n, recoverer, "tso next batch") } func (a *LocalTSOAllocator) IsLeader() bool { diff --git a/kv/tso_test.go b/kv/tso_test.go index ef4e58510..694f23261 100644 --- a/kv/tso_test.go +++ b/kv/tso_test.go @@ -7,6 +7,7 @@ import ( "time" "github.com/bootjp/elastickv/distribution" + "github.com/bootjp/elastickv/keyviz" "github.com/cockroachdb/errors" "github.com/stretchr/testify/require" ) @@ -291,6 +292,56 @@ func TestNextTimestampAfterThroughUsesAllocatorFloor(t *testing.T) { require.EqualValues(t, testTSOInitialBase+6, got) } +func TestKeyVizLabeledCoordinatorRecoversExpiredHLCCeiling(t *testing.T) { + t.Parallel() + clock := NewHLC() + clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) + eng := &fakeLeaseEngine{ + applied: 11, + leaseDur: time.Hour, + proposeApply: applyHLCLeaseEntryToClock(t, clock), + } + coord := WithKeyVizLabel(NewCoordinatorWithEngine(nil, eng, WithHLC(clock)), keyviz.LabelRedis) + + got, err := NextTimestampThrough(context.Background(), coord, "redis allocate ts") + require.NoError(t, err) + require.NotZero(t, got) + require.Equal(t, int32(1), eng.proposeCalls.Load()) +} + +func TestCoordinateTimestampRecoveryKeepsFailClosedWhenRenewalFails(t *testing.T) { + t.Parallel() + sentinel := errors.New("propose rejected: no quorum") + clock := NewHLC() + clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) + eng := &fakeLeaseEngine{applied: 11, leaseDur: time.Hour, proposeErr: sentinel} + coord := WithKeyVizLabel(NewCoordinatorWithEngine(nil, eng, WithHLC(clock)), keyviz.LabelRedis) + + _, err := NextTimestampThrough(context.Background(), coord, "redis allocate ts") + require.ErrorIs(t, err, ErrCeilingExpired) + require.Contains(t, err.Error(), "on-demand HLC lease renewal failed") + require.Equal(t, int32(1), eng.proposeCalls.Load()) +} + +func TestLocalTSOAllocatorRecoversExpiredHLCCeiling(t *testing.T) { + t.Parallel() + clock := NewHLC() + clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) + coord := &fakeTSOCoordinator{clock: clock} + coord.leader.Store(true) + coord.recoverHLCLease = func(context.Context) error { + clock.SetPhysicalCeiling(time.Now().Add(testTSOFutureCeiling).UnixMilli()) + return nil + } + alloc, err := NewLocalTSOAllocator(coord, WithTSOLeaderPollInterval(testTSOPollInterval)) + require.NoError(t, err) + + got, err := alloc.Next(context.Background()) + require.NoError(t, err) + require.NotZero(t, got) + require.EqualValues(t, 1, coord.recoverCalls.Load()) +} + func TestCoordinateUsesTSOAllocatorForIssuedTimestamps(t *testing.T) { t.Parallel() @@ -423,8 +474,10 @@ func (b *blockingTimestampAllocator) Next(ctx context.Context) (uint64, error) { } type fakeTSOCoordinator struct { - leader atomic.Bool - clock *HLC + leader atomic.Bool + clock *HLC + recoverCalls atomic.Uint64 + recoverHLCLease func(context.Context) error } func (f *fakeTSOCoordinator) IsLeader() bool { @@ -435,6 +488,14 @@ func (f *fakeTSOCoordinator) Clock() *HLC { return f.clock } +func (f *fakeTSOCoordinator) RecoverHLCLease(ctx context.Context) error { + f.recoverCalls.Add(1) + if f.recoverHLCLease != nil { + return f.recoverHLCLease(ctx) + } + return errors.WithStack(errHLCLeaseRecoveryUnavailable) +} + type fakeTimestampLeaderCoordinator struct { leader atomic.Bool timestampLeader atomic.Bool diff --git a/proxy/backend.go b/proxy/backend.go index 694155dba..9b6b27cd9 100644 --- a/proxy/backend.go +++ b/proxy/backend.go @@ -11,7 +11,7 @@ import ( const ( defaultPoolSize = 128 - defaultElasticKVPoolSize = 4 + defaultElasticKVPoolSize = 64 defaultDialTimeout = 5 * time.Second defaultReadTimeout = 3 * time.Second defaultWriteTimeout = 3 * time.Second @@ -72,9 +72,10 @@ func DefaultBackendOptions() BackendOptions { } // DefaultElasticKVBackendOptions returns defaults for proxy backends that -// connect to ElasticKV's Redis adapter. ElasticKV limits concurrent Redis -// connections per peer by default, so keep the pool below that cap unless the -// operator also raises ELASTICKV_REDIS_PER_PEER_CONNECTIONS on the cluster. +// connect to ElasticKV's Redis adapter. Production dual-write deployments +// should run the cluster with ELASTICKV_REDIS_PER_PEER_CONNECTIONS at least as +// high as this pool size; lower the proxy pool instead for clusters that keep +// the server-side per-peer cap below the default. func DefaultElasticKVBackendOptions() BackendOptions { opts := DefaultBackendOptions() opts.PoolSize = defaultElasticKVPoolSize diff --git a/proxy/proxy_test.go b/proxy/proxy_test.go index 2a5cd808f..94cac2cf4 100644 --- a/proxy/proxy_test.go +++ b/proxy/proxy_test.go @@ -1334,7 +1334,7 @@ func TestDefaultBackendOptions(t *testing.T) { func TestDefaultElasticKVBackendOptions(t *testing.T) { opts := DefaultElasticKVBackendOptions() - assert.Equal(t, 4, opts.PoolSize) + assert.Equal(t, 64, opts.PoolSize) assert.Equal(t, 5*time.Second, opts.DialTimeout) assert.Equal(t, 3*time.Second, opts.ReadTimeout) assert.Equal(t, 3*time.Second, opts.WriteTimeout) From 4d895c18949b4684cb11266ab1759a9e677040d0 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 09:55:48 +0900 Subject: [PATCH 7/9] snapshot: complete offload m1 fixes --- adapter/redis_compat_commands_stream_test.go | 46 +++ adapter/redis_stream_cmds.go | 36 +- cmd/elastickv-snapshot-offload/main.go | 314 ++++++++++++++++ cmd/elastickv-snapshot-offload/main_test.go | 90 +++++ ...artial_physical_snapshot_object_offload.md | 13 +- go.mod | 14 +- go.sum | 28 +- .../etcd/external_snapshot_restore.go | 143 ++++++-- .../etcd/external_snapshot_restore_test.go | 49 +++ internal/raftengine/etcd/wal_store.go | 2 + internal/snapshotoffload/restore.go | 4 + internal/snapshotoffload/s3_store.go | 347 ++++++++++++++++++ internal/snapshotoffload/s3_store_test.go | 244 ++++++++++++ internal/snapshotoffload/store.go | 8 +- 14 files changed, 1293 insertions(+), 45 deletions(-) create mode 100644 cmd/elastickv-snapshot-offload/main.go create mode 100644 cmd/elastickv-snapshot-offload/main_test.go create mode 100644 internal/snapshotoffload/s3_store.go create mode 100644 internal/snapshotoffload/s3_store_test.go diff --git a/adapter/redis_compat_commands_stream_test.go b/adapter/redis_compat_commands_stream_test.go index 10c38d965..a0cc0000e 100644 --- a/adapter/redis_compat_commands_stream_test.go +++ b/adapter/redis_compat_commands_stream_test.go @@ -185,6 +185,52 @@ func TestRedis_StreamXReadShortBlockReturnsNullNotError(t *testing.T) { } } +func TestRedis_StreamXReadBlockChecksWrongTypeAtDeadline(t *testing.T) { + t.Parallel() + nodes, _, _ := createNode(t, 3) + defer shutdown(nodes) + + rdbReader := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdbReader.Close() }() + rdbWriter := redis.NewClient(&redis.Options{Addr: nodes[0].redisAddress}) + defer func() { _ = rdbWriter.Close() }() + ctx := context.Background() + + key := "stream-block-wrongtype" + _, err := rdbWriter.XAdd(ctx, &redis.XAddArgs{ + Stream: key, + ID: "1-0", + Values: []string{"k", "v"}, + }).Result() + require.NoError(t, err) + + type readResult struct { + streams []redis.XStream + err error + } + resultCh := make(chan readResult, 1) + go func() { + streams, err := rdbReader.XRead(ctx, &redis.XReadArgs{ + Streams: []string{key, "$"}, + Count: 1, + Block: 100 * time.Millisecond, + }).Result() + resultCh <- readResult{streams: streams, err: err} + }() + + time.Sleep(20 * time.Millisecond) + require.NoError(t, rdbWriter.Set(ctx, key, "now-a-string", 0).Err()) + + select { + case res := <-resultCh: + require.Error(t, res.err) + require.Contains(t, res.err.Error(), "WRONGTYPE") + require.Empty(t, res.streams) + case <-time.After(2 * time.Second): + t.Fatal("XREAD BLOCK did not return after wrong-type overwrite") + } +} + // TestRedis_StreamCommandsRejectWrongType locks down the wrongType // detection on the stream fast path: keyTypeAtExpect short-circuits to // the slow path when the expected (stream) prefixes return empty, so diff --git a/adapter/redis_stream_cmds.go b/adapter/redis_stream_cmds.go index 099ecf672..ea8139986 100644 --- a/adapter/redis_stream_cmds.go +++ b/adapter/redis_stream_cmds.go @@ -1275,6 +1275,38 @@ func isXReadIterCtxError(err error) bool { } } +func (r *RedisServer) xreadFinalCheck(conn redcon.Conn, req xreadRequest) bool { + ctx, cancel := context.WithTimeout(r.handlerContext(), redisDispatchTimeout) + defer cancel() + var results []xreadResult + var err error + if ok := r.runWithHeavyCommandSlot(func() { + results, err = r.xreadOnce(ctx, req) + }); !ok { + conn.WriteError(errRedisHeavyCommandPoolFull.Error()) + return true + } + if err != nil { + if isXReadIterCtxError(err) { + return false + } + writeRedisError(conn, err) + return true + } + if len(results) > 0 { + writeXReadResults(conn, results) + return true + } + return false +} + +func (r *RedisServer) writeXReadFinalOrNull(conn redcon.Conn, req xreadRequest) { + if r.xreadFinalCheck(conn, req) { + return + } + conn.WriteNull() +} + func (r *RedisServer) xread(conn redcon.Conn, cmd redcon.Command) { req, err := parseXReadRequest(cmd.Args) if err != nil { @@ -1350,7 +1382,7 @@ func (r *RedisServer) xreadBusyPoll(conn redcon.Conn, req xreadRequest, deadline // return DeadlineExceeded, which we'd then surface as an error. iterTimeout := time.Until(deadline) if iterTimeout <= 0 { - conn.WriteNull() + r.writeXReadFinalOrNull(conn, req) return } // Cap each iteration at redisDispatchTimeout to avoid holding @@ -1394,7 +1426,7 @@ func (r *RedisServer) xreadBusyPoll(conn redcon.Conn, req xreadRequest, deadline } if !time.Now().Before(deadline) { - conn.WriteNull() + r.writeXReadFinalOrNull(conn, req) return } waitForBlockedCommandUpdate(handlerCtx, w, deadline, r.blockWaitFallback) diff --git a/cmd/elastickv-snapshot-offload/main.go b/cmd/elastickv-snapshot-offload/main.go new file mode 100644 index 000000000..e1db6abb5 --- /dev/null +++ b/cmd/elastickv-snapshot-offload/main.go @@ -0,0 +1,314 @@ +// Command elastickv-snapshot-offload publishes and restores physical Raft/FSM +// snapshots through an object store. +package main + +import ( + "context" + "flag" + "fmt" + "io" + "log/slog" + "os" + "sort" + "strings" + + "github.com/bootjp/elastickv/internal/raftengine/etcd" + "github.com/bootjp/elastickv/internal/snapshotoffload" + "github.com/cockroachdb/errors" +) + +const ( + exitSuccess = 0 + exitUserErr = 1 + exitDataErr = 2 + + commandPublish = "publish" + commandRestore = "restore" + storeLocal = "local" + storeS3 = "s3" +) + +type storeFlags struct { + storeKind string + localRoot string + s3Bucket string + s3Region string + s3Endpoint string + s3Profile string + s3PathStyle bool + s3ServerSideEncryption string + s3KMSKeyID string + s3DisableChecksumHeaders bool +} + +type publishConfig struct { + store storeFlags + dataDir string + prefix string + groupID uint64 + sourceCluster string + binaryVersion string + spoolDir string +} + +type restoreConfig struct { + store storeFlags + manifestKey string + dataDir string + peerCSV string +} + +func main() { + logger := slog.New(slog.NewTextHandler(os.Stderr, nil)) + code, err := run(context.Background(), os.Args[1:], os.Stdout, logger) + if err != nil { + logger.Error("elastickv-snapshot-offload", "err", err) + } + os.Exit(code) +} + +func run(ctx context.Context, argv []string, stdout io.Writer, logger *slog.Logger) (int, error) { + if len(argv) == 0 { + return exitUserErr, errors.New("subcommand required: publish or restore") + } + switch argv[0] { + case commandPublish: + cfg, err := parsePublishFlags(argv[1:]) + if err != nil { + return exitUserErr, err + } + if err := runPublish(ctx, cfg, stdout, logger); err != nil { + return classifyError(err), err + } + return exitSuccess, nil + case commandRestore: + cfg, err := parseRestoreFlags(argv[1:]) + if err != nil { + return exitUserErr, err + } + if err := runRestore(ctx, cfg, logger); err != nil { + return classifyError(err), err + } + return exitSuccess, nil + default: + return exitUserErr, errors.Errorf("unknown subcommand %q", argv[0]) + } +} + +func classifyError(err error) int { + switch { + case errors.Is(err, snapshotoffload.ErrIntegrity), + errors.Is(err, snapshotoffload.ErrObjectNotFound), + errors.Is(err, etcd.ErrExternalSnapshotRestoreInvalid), + errors.Is(err, etcd.ErrExternalSnapshotRestoreSHA256): + return exitDataErr + default: + return exitUserErr + } +} + +func parsePublishFlags(argv []string) (*publishConfig, error) { + cfg := &publishConfig{} + fs := flag.NewFlagSet("elastickv-snapshot-offload publish", flag.ContinueOnError) + fs.SetOutput(io.Discard) + addStoreFlags(fs, &cfg.store) + fs.StringVar(&cfg.dataDir, "data-dir", "", "Source raft data directory containing a persisted snapshot (required)") + fs.StringVar(&cfg.prefix, "prefix", "", "Object key prefix for published snapshot objects") + fs.Uint64Var(&cfg.groupID, "group-id", 0, "Raft group ID recorded in the manifest") + fs.StringVar(&cfg.sourceCluster, "source-cluster", "", "Source cluster identifier (required for group 0 manifests)") + fs.StringVar(&cfg.binaryVersion, "binary-version", "", "Binary version recorded in the manifest") + fs.StringVar(&cfg.spoolDir, "spool-dir", "", "Temporary spool directory for the payload stream") + if err := fs.Parse(argv); err != nil { + return nil, errors.WithStack(err) + } + if strings.TrimSpace(cfg.dataDir) == "" { + return nil, errors.New("--data-dir is required") + } + if cfg.groupID == 0 && strings.TrimSpace(cfg.sourceCluster) == "" { + return nil, errors.New("--source-cluster is required when --group-id is 0") + } + if err := validateStoreFlags(cfg.store); err != nil { + return nil, err + } + return cfg, nil +} + +func parseRestoreFlags(argv []string) (*restoreConfig, error) { + cfg := &restoreConfig{} + fs := flag.NewFlagSet("elastickv-snapshot-offload restore", flag.ContinueOnError) + fs.SetOutput(io.Discard) + addStoreFlags(fs, &cfg.store) + fs.StringVar(&cfg.manifestKey, "manifest-key", "", "Object key of the snapshot manifest to restore (required)") + fs.StringVar(&cfg.dataDir, "data-dir", "", "Fresh target raft data directory to create (required; must not already exist)") + fs.StringVar(&cfg.peerCSV, "peers", "", "Comma-separated raft peers id=addr,id=addr (required)") + if err := fs.Parse(argv); err != nil { + return nil, errors.WithStack(err) + } + if strings.TrimSpace(cfg.manifestKey) == "" { + return nil, errors.New("--manifest-key is required") + } + if strings.TrimSpace(cfg.dataDir) == "" { + return nil, errors.New("--data-dir is required") + } + if strings.TrimSpace(cfg.peerCSV) == "" { + return nil, errors.New("--peers is required") + } + if err := validateStoreFlags(cfg.store); err != nil { + return nil, err + } + return cfg, nil +} + +func addStoreFlags(fs *flag.FlagSet, cfg *storeFlags) { + cfg.storeKind = storeLocal + cfg.s3Region = "us-east-1" + cfg.s3PathStyle = true + fs.StringVar(&cfg.storeKind, "store", storeLocal, "Object store backend: local or s3") + fs.StringVar(&cfg.localRoot, "local-root", "", "Local object store root when --store=local") + fs.StringVar(&cfg.s3Bucket, "s3-bucket", "", "S3 bucket when --store=s3") + fs.StringVar(&cfg.s3Region, "s3-region", cfg.s3Region, "S3 signing region") + fs.StringVar(&cfg.s3Endpoint, "s3-endpoint", "", "S3-compatible endpoint URL") + fs.StringVar(&cfg.s3Profile, "s3-profile", "", "AWS shared config profile") + fs.BoolVar(&cfg.s3PathStyle, "s3-path-style", cfg.s3PathStyle, "Use path-style S3 addressing") + fs.StringVar(&cfg.s3ServerSideEncryption, "s3-sse", "", "Server-side encryption algorithm for uploaded objects, for example AES256 or aws:kms") + fs.StringVar(&cfg.s3KMSKeyID, "s3-kms-key-id", "", "KMS key ID when --s3-sse=aws:kms") + fs.BoolVar(&cfg.s3DisableChecksumHeaders, "s3-disable-checksum-headers", false, "Do not send S3 checksum headers; keep metadata and restore-time verification") +} + +func validateStoreFlags(cfg storeFlags) error { + switch cfg.storeKind { + case storeLocal: + if strings.TrimSpace(cfg.localRoot) == "" { + return errors.New("--local-root is required when --store=local") + } + case storeS3: + if strings.TrimSpace(cfg.s3Bucket) == "" { + return errors.New("--s3-bucket is required when --store=s3") + } + default: + return errors.Errorf("unknown --store %q", cfg.storeKind) + } + if strings.TrimSpace(cfg.s3KMSKeyID) != "" && strings.TrimSpace(cfg.s3ServerSideEncryption) == "" { + return errors.New("--s3-kms-key-id requires --s3-sse") + } + return nil +} + +func runPublish(ctx context.Context, cfg *publishConfig, stdout io.Writer, logger *slog.Logger) error { + store, err := openObjectStore(ctx, cfg.store) + if err != nil { + return err + } + manifest, err := snapshotoffload.PublishPersistedSnapshot(ctx, snapshotoffload.PublishOptions{ + Store: store, + DataDir: cfg.dataDir, + Prefix: cfg.prefix, + GroupID: cfg.groupID, + SourceCluster: cfg.sourceCluster, + BinaryVersion: cfg.binaryVersion, + SpoolDir: cfg.spoolDir, + }) + if err != nil { + return errors.Wrap(err, "publish persisted snapshot") + } + raw, _, err := manifest.MarshalCanonical() + if err != nil { + return errors.Wrap(err, "marshal snapshot manifest") + } + if _, err := stdout.Write(raw); err != nil { + return errors.WithStack(err) + } + logger.Info("snapshot offload manifest published", + "manifest_key", manifest.ManifestKey, + "payload_key", manifest.Payload.Key, + "group_id", manifest.GroupID, + "index", manifest.SnapshotIndex, + "term", manifest.SnapshotTerm, + ) + return nil +} + +func runRestore(ctx context.Context, cfg *restoreConfig, logger *slog.Logger) error { + store, err := openObjectStore(ctx, cfg.store) + if err != nil { + return err + } + peers, err := parsePeers(cfg.peerCSV) + if err != nil { + return err + } + result, err := snapshotoffload.RestorePhysicalSnapshot(ctx, snapshotoffload.RestoreOptions{ + Store: store, + ManifestKey: cfg.manifestKey, + DataDir: cfg.dataDir, + Peers: peers, + }) + if err != nil { + return errors.Wrap(err, "restore physical snapshot") + } + logger.Info("snapshot offload restore data dir prepared", + "data_dir", result.DataDir, + "fsm", result.FSMPath, + "snap", result.SnapPath, + "crc32c", fmt.Sprintf("%08x", result.CRC32C), + "payload_sha256", result.PayloadSHA256, + "payload_bytes", result.PayloadBytes, + "peers", result.Peers, + ) + return nil +} + +func openObjectStore(ctx context.Context, cfg storeFlags) (snapshotoffload.ObjectStore, error) { + switch cfg.storeKind { + case storeLocal: + store, err := snapshotoffload.NewLocalStore(cfg.localRoot) + if err != nil { + return nil, errors.Wrap(err, "open local object store") + } + return store, nil + case storeS3: + store, err := snapshotoffload.NewS3Store(ctx, snapshotoffload.S3StoreConfig{ + Bucket: cfg.s3Bucket, + Region: cfg.s3Region, + Endpoint: cfg.s3Endpoint, + Profile: cfg.s3Profile, + ForcePathStyle: cfg.s3PathStyle, + ServerSideEncryption: cfg.s3ServerSideEncryption, + SSEKMSKeyID: cfg.s3KMSKeyID, + DisableChecksumHeaders: cfg.s3DisableChecksumHeaders, + }) + if err != nil { + return nil, errors.Wrap(err, "open s3 object store") + } + return store, nil + default: + return nil, errors.Errorf("unknown --store %q", cfg.storeKind) + } +} + +func parsePeers(raw string) ([]etcd.Peer, error) { + peers, err := etcd.ParsePeers(raw) + if err != nil { + return nil, errors.Wrap(err, "parse peers") + } + if len(peers) == 0 { + return nil, errors.New("--peers selected zero peers") + } + seen := make(map[uint64]struct{}, len(peers)) + for i := range peers { + if peers[i].NodeID == 0 { + peers[i].NodeID = etcd.DeriveNodeID(peers[i].ID) + } + if peers[i].NodeID == 0 { + return nil, errors.Errorf("peer %q derived node id 0", peers[i].ID) + } + if _, ok := seen[peers[i].NodeID]; ok { + return nil, errors.Errorf("duplicate peer node id %d", peers[i].NodeID) + } + seen[peers[i].NodeID] = struct{}{} + } + sort.Slice(peers, func(i, j int) bool { + return peers[i].NodeID < peers[j].NodeID + }) + return peers, nil +} diff --git a/cmd/elastickv-snapshot-offload/main_test.go b/cmd/elastickv-snapshot-offload/main_test.go new file mode 100644 index 000000000..dc4700595 --- /dev/null +++ b/cmd/elastickv-snapshot-offload/main_test.go @@ -0,0 +1,90 @@ +package main + +import ( + "bytes" + "context" + "io" + "log/slog" + "os" + "path/filepath" + "testing" + + "github.com/bootjp/elastickv/internal/raftengine/etcd" + "github.com/bootjp/elastickv/internal/snapshotoffload" + "github.com/stretchr/testify/require" +) + +func TestSnapshotOffloadCLIPublishAndRestoreLocal(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1cli-physical-snapshot-payload") + sourceDataDir := seedCLISnapshot(t, root, payload, 60, 9) + objectRoot := filepath.Join(root, "objects") + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + + var stdout bytes.Buffer + code, err := run(ctx, []string{ + commandPublish, + "--store", storeLocal, + "--local-root", objectRoot, + "--data-dir", sourceDataDir, + "--prefix", "cluster-cli", + "--group-id", "2", + "--source-cluster", "cluster-cli", + "--binary-version", "test-version", + }, &stdout, logger) + require.NoError(t, err) + require.Equal(t, exitSuccess, code) + + manifest, err := snapshotoffload.DecodeManifest(stdout.Bytes()) + require.NoError(t, err) + require.Equal(t, uint64(60), manifest.SnapshotIndex) + require.Equal(t, int64(len(payload)), manifest.Payload.Bytes) + + restoreDataDir := filepath.Join(root, "restored") + code, err = run(ctx, []string{ + commandRestore, + "--store", storeLocal, + "--local-root", objectRoot, + "--manifest-key", manifest.ManifestKey, + "--data-dir", restoreDataDir, + "--peers", "n2=127.0.0.1:12002", + }, io.Discard, logger) + require.NoError(t, err) + require.Equal(t, exitSuccess, code) + + export, ok, err := etcd.OpenPersistedSnapshotExport(restoreDataDir) + require.NoError(t, err) + require.True(t, ok) + defer func() { require.NoError(t, export.Close()) }() + require.Equal(t, []uint64{etcd.DeriveNodeID("n2")}, export.Metadata().ConfState.GetVoters()) +} + +func TestSnapshotOffloadCLIRequiresLocalRoot(t *testing.T) { + code, err := run(context.Background(), []string{ + commandPublish, + "--store", storeLocal, + "--data-dir", "data", + "--group-id", "1", + }, io.Discard, slog.New(slog.NewTextHandler(io.Discard, nil))) + require.ErrorContains(t, err, "--local-root is required") + require.Equal(t, exitUserErr, code) +} + +func seedCLISnapshot(t *testing.T, root string, payload []byte, index uint64, term uint64) string { + t.Helper() + input := filepath.Join(root, "source.fsm") + require.NoError(t, os.WriteFile(input, payload, 0o600)) + dataDir := filepath.Join(root, "source-raft") + _, err := etcd.PreparePhysicalSnapshotRestore(etcd.PhysicalSnapshotRestoreOptions{ + InputFSMPath: input, + DataDir: dataDir, + Index: index, + Term: term, + Peers: []etcd.Peer{ + {NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}, + }, + }) + require.NoError(t, err) + return dataDir +} diff --git a/docs/design/2026_07_19_partial_physical_snapshot_object_offload.md b/docs/design/2026_07_19_partial_physical_snapshot_object_offload.md index e70269a18..7bf1765a6 100644 --- a/docs/design/2026_07_19_partial_physical_snapshot_object_offload.md +++ b/docs/design/2026_07_19_partial_physical_snapshot_object_offload.md @@ -1,6 +1,6 @@ # Physical Snapshot Object Offload -Status: Partial — M0 implemented; M1 object-store-neutral substrate partial +Status: Partial — M0/M1 implemented; M2/M3 pending Author: bootjp Date: 2026-07-19 Updated: 2026-07-23 @@ -35,15 +35,20 @@ The M1 object-store-neutral substrate now adds: - a `snapshotoffload.ObjectStore` interface with a local filesystem implementation for deterministic tests and offline drills; +- an S3-compatible implementation backed by AWS SDK v2, including optional + path-style endpoints, SDK credential-chain loading, server-side encryption + headers, immutable `If-None-Match: *` writes, SHA-256 metadata, and optional + checksum headers; - the v1 JSON manifest schema and content-addressed payload key layout; - payload-first publish from `OpenPersistedSnapshotExport`, including exact byte count and SHA-256 verification before manifest commit; - manifest-driven restore that downloads the opaque payload, checks exact length and SHA-256, then calls `PreparePhysicalSnapshotRestore` with operator-supplied target membership. +- `cmd/elastickv-snapshot-offload publish` and `restore` for local and + S3-backed operator workflows. -The S3-compatible client, operator CLI, runtime scheduler, and retention/GC -remain pending. +The runtime scheduler and retention/GC remain pending. ## 2. Safety boundary @@ -158,7 +163,7 @@ permissions below the configured prefix. | Milestone | Scope | Status | |---|---|---| | M0 | Persisted snapshot export handle, complete-payload restore preparation, focused design | Implemented in the first substrate PR | -| M1 | Object client interface, S3-compatible implementation, immutable payload/manifest publication, download verification, operator CLI | Partial: object-store interface, local implementation, manifest schema, payload-first publish, and verified restore are implemented; S3 client and CLI pending | +| M1 | Object client interface, S3-compatible implementation, immutable payload/manifest publication, download verification, operator CLI | Implemented: local and S3 stores, manifest schema, payload-first publish, verified restore, and publish/restore CLI | | M2 | Leader-only per-group scheduler, metrics, jitter, concurrency bounds, cancellation and restart idempotency | Pending | | M3 | Retention/GC, restore drills, corruption tests, multi-node acceptance, operational documentation | Pending | diff --git a/go.mod b/go.mod index b9c9df389..a21bfe432 100644 --- a/go.mod +++ b/go.mod @@ -6,10 +6,11 @@ toolchain go1.26.5 require ( github.com/Jille/grpc-multi-resolver v1.3.0 - github.com/aws/aws-sdk-go-v2 v1.42.1 + github.com/aws/aws-sdk-go-v2 v1.43.0 github.com/aws/aws-sdk-go-v2/config v1.32.29 github.com/aws/aws-sdk-go-v2/credentials v1.19.28 github.com/aws/aws-sdk-go-v2/service/dynamodb v1.60.0 + github.com/aws/aws-sdk-go-v2/service/s3 v1.106.0 github.com/aws/smithy-go v1.27.3 github.com/cockroachdb/errors v1.14.0 github.com/cockroachdb/pebble/v2 v2.1.6 @@ -43,13 +44,16 @@ require ( github.com/DataDog/zstd v1.5.7 // indirect github.com/RaduBerinde/axisds v0.1.0 // indirect github.com/RaduBerinde/btreemap v0.0.0-20250419174037-3d62b7205d54 // indirect + github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.14 // indirect github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.30 // indirect - github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30 // indirect - github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30 // indirect - github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.31 // indirect + github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.31 // indirect + github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.31 // indirect + github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.32 // indirect github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.13 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.24 // indirect github.com/aws/aws-sdk-go-v2/service/internal/endpoint-discovery v1.12.7 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.30 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.31 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.32 // indirect github.com/aws/aws-sdk-go-v2/service/signin v1.4.0 // indirect github.com/aws/aws-sdk-go-v2/service/sso v1.32.0 // indirect github.com/aws/aws-sdk-go-v2/service/ssooidc v1.37.0 // indirect diff --git a/go.sum b/go.sum index 1203c7d30..9528f532d 100644 --- a/go.sum +++ b/go.sum @@ -8,28 +8,36 @@ github.com/RaduBerinde/btreemap v0.0.0-20250419174037-3d62b7205d54 h1:bsU8Tzxr/P github.com/RaduBerinde/btreemap v0.0.0-20250419174037-3d62b7205d54/go.mod h1:0tr7FllbE9gJkHq7CVeeDDFAFKQVy5RnCSSNBOvdqbc= github.com/aclements/go-perfevent v0.0.0-20240301234650-f7843625020f h1:JjxwchlOepwsUWcQwD2mLUAGE9aCp0/ehy6yCHFBOvo= github.com/aclements/go-perfevent v0.0.0-20240301234650-f7843625020f/go.mod h1:tMDTce/yLLN/SK8gMOxQfnyeMeCg8KGzp0D1cbECEeo= -github.com/aws/aws-sdk-go-v2 v1.42.1 h1:9eOTgu1z/dVtYpNZ3/8/XbbaX0x/BqE3HUzAzs6K0ek= -github.com/aws/aws-sdk-go-v2 v1.42.1/go.mod h1:5pKeft2eJj+gElQ38Jqg4ibCqh+/AK33/0X3hip7IjM= +github.com/aws/aws-sdk-go-v2 v1.43.0 h1:fharf/WhbRAVZ1du0QL7roNFxZ6T/sWr+4Ni617bwSI= +github.com/aws/aws-sdk-go-v2 v1.43.0/go.mod h1:5pKeft2eJj+gElQ38Jqg4ibCqh+/AK33/0X3hip7IjM= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.14 h1:3IZY0XAJquT3aHzbkHfPzy4ACPcEjVG0x87KOwtpqGY= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.14/go.mod h1:zwM6veDkhGgQFqkBy+uT28AAYpLu+uFMlPl+rCg/73E= github.com/aws/aws-sdk-go-v2/config v1.32.29 h1:BcMHHnpiWKogf+gGfpj3K1w+Sktz29XDo/cPSAPO3FU= github.com/aws/aws-sdk-go-v2/config v1.32.29/go.mod h1:+Kbhn8Es4kPUph3F/0W7avykytc+Jh2Ld9/msv9ljV4= github.com/aws/aws-sdk-go-v2/credentials v1.19.28 h1:zTXJSsNcoO91/mTXsZoYf0AK8dvNPiA58/VtyGXR+wM= github.com/aws/aws-sdk-go-v2/credentials v1.19.28/go.mod h1:Kd9E0JzDBW/q1xbsHFrev/GnbAf5J0Ng8xoyc7HZ91Q= github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.30 h1:/hi1JADLEW9YYryEz1w4GQu0EtP23pP553Cf9KgsDV4= github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.30/go.mod h1:/3AOgy4K17Dm4ucMZVC/MJkzy5kmfKUcINRHZyo0koQ= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30 h1:xM/Is9cKMHa8Jj8zkvWhvrFkZsXJV9E+BB4g0HW0duQ= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30/go.mod h1:WueJeNDZvK1fMYEWJIkcivBfEzUkTpBhzlrUKKY8EuA= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30 h1:jn46zC9LdsVR/ZpMIJqMqb8hHv31BlLx3ulVqNspUOk= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30/go.mod h1:1hTMsAgbdS/AtUi4bw8+gUuh1pceo+eXRLfpSuSQj3M= -github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.31 h1:3GUprIsfmGcC5SACIyB0e7E0BM1O1b3Erl5CePYIAeQ= -github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.31/go.mod h1:7PuV1yl5e2xnUbm+RqvVg5i2iBM8EyijZNoI9wsOoOc= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.31 h1:Z8F3hfCY33IGpJjFAnv0wvtv1FIKj1GHmRDEYqy64tw= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.31/go.mod h1:aVyUoytEyOViR6jhq6jula0xkc5NfBE2hgeF6BvOrao= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.31 h1:hyOxUyXdh3AyjE93gBgsfziJag9ACwcs+ZpDBLzi8mw= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.31/go.mod h1:OERqI9k0draSLB8O8woxY3q25ZWTELRK4RRoLMuMZFo= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.32 h1:0MrUL35H/Y4kdFfItoR5jCgtDQ4Z/8LudAoIHRfA4hE= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.32/go.mod h1:2tNZkuWz54arj8mHVf+8Y7cKkcD8Wr/fBpENgEXpjLc= github.com/aws/aws-sdk-go-v2/service/dynamodb v1.60.0 h1:wnyC01GGkqvRvVI73+xRTr8qDyIZMzqjHlsFqWUUVSo= github.com/aws/aws-sdk-go-v2/service/dynamodb v1.60.0/go.mod h1:HnWoC3m6VmjUSg+kBL6OgQsXdyRAGzBYWb7B3J2f+JM= github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.13 h1:mbRIur/BiHK6SKPjoBIXSE/hJ6g6JGRLuxQy1jGjlN4= github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.13/go.mod h1:ITg9em2KbJx1s0y4aqRX5OYWG6HBZ5TVR//OdpEZ2CQ= +github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.24 h1:mdPwDQPqxlw9Sc62Nt15yjEcARaDbPXkjRYtXsUripo= +github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.24/go.mod h1:ls5ytnwLTcQaUu32fMYXFI3MjpKuTwL840PAm9iqyEg= github.com/aws/aws-sdk-go-v2/service/internal/endpoint-discovery v1.12.7 h1:uqsKxr7kJp9DXVj2m8KbVeZcYMuwsNEwvoVrYl2Vpf8= github.com/aws/aws-sdk-go-v2/service/internal/endpoint-discovery v1.12.7/go.mod h1:Js/P8Zbwe1mRejnD+OpFLyQiJ8ioQlo3GMAg7Dfxk7w= -github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.30 h1:/Z5jmNrKsSD7EmDjzAPsm/3L9IuOkzaynklJZ1qX7S4= -github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.30/go.mod h1:lEzEZnOosE7zi8Z6royW1cFJTD9fpab4Ul1SBrllewk= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.31 h1:w2SIhW92DZPFrSL4ksVCr8IYff5OZwIcxg8+95tzvAI= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.31/go.mod h1:wAhpCQbkov+IcvjozJbd2xRCoZybUEHNkcFunssNACg= +github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.32 h1:jWXtZdCnhXa9sGFixRaU2AxT4DIVse9HS4E2f+/KwV0= +github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.32/go.mod h1:9JS1UpfVvyD/ZPX8GsKb/Pq8scEM+7GP5fqh9SwH7po= +github.com/aws/aws-sdk-go-v2/service/s3 v1.106.0 h1:7QZWVJZWzHivHWIa+5TELLaBBkbuoj0GPwQtMlJ0sqk= +github.com/aws/aws-sdk-go-v2/service/s3 v1.106.0/go.mod h1:fcvq5L7dK+5cQFicEJwpI6e6Wn8NY2i6yT5wRLYVc7s= github.com/aws/aws-sdk-go-v2/service/signin v1.4.0 h1:sLzmJGCMv+C8KqiJgEqDLB6vxaJGmobRh4rr//ZpA3w= github.com/aws/aws-sdk-go-v2/service/signin v1.4.0/go.mod h1:mxC0nT/C8wMMS97DemZPzvUZxvIt+2Iq+eS3JdFZGgg= github.com/aws/aws-sdk-go-v2/service/sso v1.32.0 h1:qjMmry/cBDee1E/2gyvel0uRYCi3mwRZ2hf6N+GAodo= diff --git a/internal/raftengine/etcd/external_snapshot_restore.go b/internal/raftengine/etcd/external_snapshot_restore.go index 71364c401..dc3497594 100644 --- a/internal/raftengine/etcd/external_snapshot_restore.go +++ b/internal/raftengine/etcd/external_snapshot_restore.go @@ -2,6 +2,7 @@ package etcd import ( "bufio" + "context" "crypto/sha256" "encoding/binary" "encoding/hex" @@ -27,6 +28,7 @@ var ( ) type ExternalSnapshotRestoreOptions struct { + Context context.Context InputFSMPath string DataDir string Index uint64 @@ -50,6 +52,7 @@ type ExternalSnapshotRestoreResult struct { // emitted by the Raft state machine. Unlike ExternalSnapshotRestoreOptions, the // input already contains the KV snapshot header and must not be wrapped again. type PhysicalSnapshotRestoreOptions struct { + Context context.Context InputFSMPath string DataDir string Index uint64 @@ -77,6 +80,7 @@ func PrepareExternalSnapshotRestore(opts ExternalSnapshotRestoreOptions) (*Exter // payloads all use the same path without this package parsing their internals. func PreparePhysicalSnapshotRestore(opts PhysicalSnapshotRestoreOptions) (*ExternalSnapshotRestoreResult, error) { externalOpts := ExternalSnapshotRestoreOptions{ + Context: opts.Context, InputFSMPath: opts.InputFSMPath, DataDir: opts.DataDir, Index: opts.Index, @@ -91,12 +95,22 @@ func PreparePhysicalSnapshotRestore(opts PhysicalSnapshotRestoreOptions) (*Exter return prepareExternalSnapshotRestore(normalized, writePhysicalFSMSnapshotFile) } -type externalSnapshotFileWriter func(inputPath, fsmSnapDir string, index uint64, ceilingMs uint64) (uint32, int64, string, error) +type externalSnapshotFileWriter func( + ctx context.Context, + inputPath string, + fsmSnapDir string, + index uint64, + ceilingMs uint64, +) (uint32, int64, string, error) func prepareExternalSnapshotRestore( opts ExternalSnapshotRestoreOptions, writeSnapshot externalSnapshotFileWriter, ) (*ExternalSnapshotRestoreResult, error) { + ctx := externalSnapshotRestoreContext(opts.Context) + if err := ctx.Err(); err != nil { + return nil, errors.WithStack(err) + } destDataDir, tempDir, err := prepareExternalSnapshotRestoreDest(opts.DataDir) if err != nil { return nil, err @@ -109,7 +123,7 @@ func prepareExternalSnapshotRestore( }() fsmSnapDir := filepath.Join(tempDir, fsmSnapDirName) - crc, payloadBytes, payloadSHA, err := writeSnapshot(opts.InputFSMPath, fsmSnapDir, opts.Index, opts.SnapshotCeilingMs) + crc, payloadBytes, payloadSHA, err := writeSnapshot(ctx, opts.InputFSMPath, fsmSnapDir, opts.Index, opts.SnapshotCeilingMs) if err != nil { return nil, err } @@ -118,13 +132,10 @@ func prepareExternalSnapshotRestore( "copied payload has %s, expected %s", payloadSHA, opts.ExpectedPayloadSHA256) } token := encodeSnapshotToken(opts.Index, crc) - if err := seedExternalSnapshotRestoreDir(tempDir, opts, token); err != nil { + if err := seedExternalSnapshotRestoreDir(ctx, tempDir, opts, token); err != nil { return nil, err } - if err := finalizeMigrationDir(tempDir, destDataDir); err != nil { - if errors.Is(err, errMigrationDestinationExists) { - return nil, errors.Wrapf(ErrExternalSnapshotRestoreExists, "destination exists: %s", destDataDir) - } + if err := finalizeExternalSnapshotRestoreDir(ctx, tempDir, destDataDir); err != nil { return nil, err } committed = true @@ -140,6 +151,26 @@ func prepareExternalSnapshotRestore( }, nil } +func finalizeExternalSnapshotRestoreDir(ctx context.Context, tempDir string, destDataDir string) error { + if err := ctx.Err(); err != nil { + return errors.WithStack(err) + } + if err := finalizeMigrationDir(tempDir, destDataDir); err != nil { + if errors.Is(err, errMigrationDestinationExists) { + return errors.Wrapf(ErrExternalSnapshotRestoreExists, "destination exists: %s", destDataDir) + } + return err + } + return nil +} + +func externalSnapshotRestoreContext(ctx context.Context) context.Context { + if ctx == nil { + return context.Background() + } + return ctx +} + func validateExternalSnapshotRestoreOptions(opts ExternalSnapshotRestoreOptions) error { switch { case opts.InputFSMPath == "": @@ -247,16 +278,37 @@ func ensureExternalRestorePathAbsent(path string, kind string) error { return nil } -func writeExternalFSMSnapshotFile(inputPath, fsmSnapDir string, index uint64, ceilingMs uint64) (uint32, int64, string, error) { +func writeExternalFSMSnapshotFile( + ctx context.Context, + inputPath string, + fsmSnapDir string, + index uint64, + ceilingMs uint64, +) (uint32, int64, string, error) { header := encodeExternalRestoreSnapshotHeader(ceilingMs) - return writePreparedFSMSnapshotFile(inputPath, fsmSnapDir, index, header[:]) + return writePreparedFSMSnapshotFile(ctx, inputPath, fsmSnapDir, index, header[:]) } -func writePhysicalFSMSnapshotFile(inputPath, fsmSnapDir string, index uint64, _ uint64) (uint32, int64, string, error) { - return writePreparedFSMSnapshotFile(inputPath, fsmSnapDir, index, nil) +func writePhysicalFSMSnapshotFile( + ctx context.Context, + inputPath string, + fsmSnapDir string, + index uint64, + _ uint64, +) (uint32, int64, string, error) { + return writePreparedFSMSnapshotFile(ctx, inputPath, fsmSnapDir, index, nil) } -func writePreparedFSMSnapshotFile(inputPath, fsmSnapDir string, index uint64, prefix []byte) (uint32, int64, string, error) { +func writePreparedFSMSnapshotFile( + ctx context.Context, + inputPath string, + fsmSnapDir string, + index uint64, + prefix []byte, +) (uint32, int64, string, error) { + if err := ctx.Err(); err != nil { + return 0, 0, "", errors.WithStack(err) + } if err := os.MkdirAll(fsmSnapDir, defaultDirPerm); err != nil { return 0, 0, "", errors.WithStack(err) } @@ -280,21 +332,34 @@ func writePreparedFSMSnapshotFile(inputPath, fsmSnapDir string, index uint64, pr _ = os.Remove(tmpPath) }() - crc, bytesWritten, payloadSHA, err := copySnapshotPayloadWithPrefixAndFooter(in, tmpFile, prefix) + crc, bytesWritten, payloadSHA, err := copySnapshotPayloadWithPrefixAndFooter(ctx, in, tmpFile, prefix) if err != nil { return 0, 0, "", err } - closed = true - if err := tmpFile.Close(); err != nil { + if err := ctx.Err(); err != nil { return 0, 0, "", errors.WithStack(err) } - if err := os.Rename(tmpPath, finalPath); err != nil { - return 0, 0, "", errors.WithStack(err) + if err := closeAndInstallPreparedFSMSnapshot(ctx, tmpFile, finalPath, fsmSnapDir); err != nil { + return 0, 0, "", err + } + closed = true + return crc, bytesWritten, payloadSHA, nil +} + +func closeAndInstallPreparedFSMSnapshot(ctx context.Context, file *os.File, finalPath string, fsmSnapDir string) error { + if err := file.Close(); err != nil { + return errors.WithStack(err) + } + if err := ctx.Err(); err != nil { + return errors.WithStack(err) + } + if err := os.Rename(file.Name(), finalPath); err != nil { + return errors.WithStack(err) } if err := syncDir(fsmSnapDir); err != nil { - return 0, 0, "", errors.WithStack(err) + return errors.WithStack(err) } - return crc, bytesWritten, payloadSHA, nil + return nil } func openRegularExternalFSMInput(inputPath string) (*os.File, error) { @@ -323,11 +388,14 @@ func requireRegularFile(f *os.File, path string) error { return nil } -func copySnapshotPayloadWithPrefixAndFooter(in io.Reader, out *os.File, prefix []byte) (uint32, int64, string, error) { +func copySnapshotPayloadWithPrefixAndFooter(ctx context.Context, in io.Reader, out *os.File, prefix []byte) (uint32, int64, string, error) { bw := bufio.NewWriterSize(out, fsmWriteBufSize) crcHash := crc32.New(crc32cTable) payloadHash := sha256.New() if len(prefix) != 0 { + if err := ctx.Err(); err != nil { + return 0, 0, "", errors.WithStack(err) + } if _, err := bw.Write(prefix); err != nil { return 0, 0, "", errors.WithStack(err) } @@ -335,10 +403,16 @@ func copySnapshotPayloadWithPrefixAndFooter(in io.Reader, out *os.File, prefix [ return 0, 0, "", errors.WithStack(err) } } - n, err := io.Copy(io.MultiWriter(bw, crcHash, payloadHash), in) + n, err := io.Copy(io.MultiWriter(bw, crcHash, payloadHash), externalSnapshotRestoreContextReader{ + ctx: ctx, + reader: in, + }) if err != nil { return 0, n, "", errors.WithStack(err) } + if err := ctx.Err(); err != nil { + return 0, n, "", errors.WithStack(err) + } crc := crcHash.Sum32() if err := binary.Write(bw, binary.BigEndian, crc); err != nil { return 0, n, "", errors.WithStack(err) @@ -352,6 +426,25 @@ func copySnapshotPayloadWithPrefixAndFooter(in io.Reader, out *os.File, prefix [ return crc, n, hex.EncodeToString(payloadHash.Sum(nil)), nil } +type externalSnapshotRestoreContextReader struct { + ctx context.Context + reader io.Reader +} + +func (r externalSnapshotRestoreContextReader) Read(p []byte) (int, error) { + if err := r.ctx.Err(); err != nil { + return 0, errors.WithStack(err) + } + n, err := r.reader.Read(p) + if err != nil { + if errors.Is(err, io.EOF) { + return n, err //nolint:wrapcheck // io.Reader must return io.EOF unwrapped so io.Copy treats it as normal completion. + } + return n, errors.WithStack(err) + } + return n, nil +} + func encodeExternalRestoreSnapshotHeader(ceilingMs uint64) [externalRestoreSnapshotHeaderLen]byte { var header [externalRestoreSnapshotHeaderLen]byte copy(header[:externalRestoreSnapshotMagicLen], externalRestoreSnapshotMagic[:]) @@ -366,7 +459,10 @@ const ( externalRestoreSnapshotHeaderLen = 16 ) -func seedExternalSnapshotRestoreDir(tempDir string, opts ExternalSnapshotRestoreOptions, token []byte) error { +func seedExternalSnapshotRestoreDir(ctx context.Context, tempDir string, opts ExternalSnapshotRestoreOptions, token []byte) error { + if err := ctx.Err(); err != nil { + return errors.WithStack(err) + } confState := confStateForSnapshotRestorePeers(opts.Peers) state := persistedState{ HardState: raftpb.HardState{ @@ -397,6 +493,9 @@ func seedExternalSnapshotRestoreDir(tempDir string, opts ExternalSnapshotRestore if err := closePersist(disk.Persist); err != nil { return err } + if err := ctx.Err(); err != nil { + return errors.WithStack(err) + } return savePersistedPeers(tempDir, opts.Index, opts.Peers) } diff --git a/internal/raftengine/etcd/external_snapshot_restore_test.go b/internal/raftengine/etcd/external_snapshot_restore_test.go index 553fb392f..0421ecc80 100644 --- a/internal/raftengine/etcd/external_snapshot_restore_test.go +++ b/internal/raftengine/etcd/external_snapshot_restore_test.go @@ -1,6 +1,7 @@ package etcd import ( + "context" "crypto/sha256" "encoding/binary" "encoding/hex" @@ -70,6 +71,54 @@ func TestPrepareExternalSnapshotRestoreSeedsRuntimeFiles(t *testing.T) { }, peers) } +func TestPrepareExternalSnapshotRestoreHonorsCanceledContext(t *testing.T) { + root := t.TempDir() + input := filepath.Join(root, "encoded.fsm") + require.NoError(t, os.WriteFile(input, []byte("payload"), 0o600)) + dataDir := filepath.Join(root, "raft") + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := PrepareExternalSnapshotRestore(ExternalSnapshotRestoreOptions{ + Context: ctx, + InputFSMPath: input, + DataDir: dataDir, + Index: 1, + Term: 1, + Peers: []Peer{{NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}}, + }) + require.ErrorIs(t, err, context.Canceled) + _, statErr := os.Stat(dataDir) + require.True(t, os.IsNotExist(statErr)) + _, statErr = os.Stat(dataDir + ".restore-prep") + require.True(t, os.IsNotExist(statErr)) +} + +func TestPrepareExternalSnapshotRestoreCleansTempWhenContextCanceledAfterPayload(t *testing.T) { + root := t.TempDir() + dataDir := filepath.Join(root, "raft") + ctx, cancel := context.WithCancel(context.Background()) + opts, err := normalizeExternalSnapshotRestoreOptions(ExternalSnapshotRestoreOptions{ + Context: ctx, + InputFSMPath: filepath.Join(root, "encoded.fsm"), + DataDir: dataDir, + Index: 1, + Term: 1, + Peers: []Peer{{NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}}, + }) + require.NoError(t, err) + + _, err = prepareExternalSnapshotRestore(opts, func(context.Context, string, string, uint64, uint64) (uint32, int64, string, error) { + cancel() + return 1, 0, hex.EncodeToString(sha256.New().Sum(nil)), nil + }) + require.ErrorIs(t, err, context.Canceled) + _, statErr := os.Stat(dataDir) + require.True(t, os.IsNotExist(statErr)) + _, statErr = os.Stat(dataDir + ".restore-prep") + require.True(t, os.IsNotExist(statErr)) +} + func TestPrepareExternalSnapshotRestoreFailures(t *testing.T) { testCases := []struct { name string diff --git a/internal/raftengine/etcd/wal_store.go b/internal/raftengine/etcd/wal_store.go index c69982f90..0612ba5f0 100644 --- a/internal/raftengine/etcd/wal_store.go +++ b/internal/raftengine/etcd/wal_store.go @@ -917,7 +917,9 @@ func persistLocalSnapshotPayload(storage *etcdraft.MemoryStorage, persist etcdst const defaultMaxSnapFiles = 3 func buildLocalSnapshot(storage *etcdraft.MemoryStorage, applied uint64, payload []byte) (raftpb.Snapshot, error) { + storage.Lock() _, confState, err := storage.InitialState() + storage.Unlock() if err != nil { return raftpb.Snapshot{}, errors.WithStack(err) } diff --git a/internal/snapshotoffload/restore.go b/internal/snapshotoffload/restore.go index 8f979b820..5a6683650 100644 --- a/internal/snapshotoffload/restore.go +++ b/internal/snapshotoffload/restore.go @@ -75,7 +75,11 @@ func RestorePhysicalSnapshot(ctx context.Context, opts RestoreOptions) (*etcdraf if err := downloadVerifiedPayload(ctx, opts.Store, manifest, payloadPath); err != nil { return nil, err } + if err := ctx.Err(); err != nil { + return nil, errors.WithStack(err) + } result, err := etcdraftengine.PreparePhysicalSnapshotRestore(etcdraftengine.PhysicalSnapshotRestoreOptions{ + Context: ctx, InputFSMPath: payloadPath, DataDir: opts.DataDir, Index: manifest.SnapshotIndex, diff --git a/internal/snapshotoffload/s3_store.go b/internal/snapshotoffload/s3_store.go new file mode 100644 index 000000000..5450849cb --- /dev/null +++ b/internal/snapshotoffload/s3_store.go @@ -0,0 +1,347 @@ +package snapshotoffload + +import ( + "context" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "io" + "strings" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" + "github.com/aws/smithy-go" + "github.com/cockroachdb/errors" +) + +const ( + s3MetadataSHA256 = "elastickv-sha256" + s3LoadConfigOptionCapHint = 4 +) + +type S3ObjectClient interface { + PutObject(context.Context, *s3.PutObjectInput, ...func(*s3.Options)) (*s3.PutObjectOutput, error) + GetObject(context.Context, *s3.GetObjectInput, ...func(*s3.Options)) (*s3.GetObjectOutput, error) + HeadObject(context.Context, *s3.HeadObjectInput, ...func(*s3.Options)) (*s3.HeadObjectOutput, error) +} + +type S3StoreConfig struct { + Client S3ObjectClient + Bucket string + Region string + Endpoint string + Profile string + ForcePathStyle bool + AccessKeyID string + SecretAccessKey string + SessionToken string + ServerSideEncryption string + SSEKMSKeyID string + DisableChecksumHeaders bool +} + +type S3Store struct { + client S3ObjectClient + bucket string + serverSideEncryption string + sseKMSKeyID string + disableChecksumHeaders bool +} + +func NewS3Store(ctx context.Context, cfg S3StoreConfig) (*S3Store, error) { + if stringsTrim(cfg.Bucket) == "" { + return nil, errors.Wrap(ErrInvalidOptions, "s3 bucket is required") + } + client := cfg.Client + if client == nil { + awsCfg, err := loadS3AWSConfig(ctx, cfg) + if err != nil { + return nil, err + } + client = s3.NewFromConfig(awsCfg, func(o *s3.Options) { + o.UsePathStyle = cfg.ForcePathStyle + if stringsTrim(cfg.Endpoint) != "" { + o.BaseEndpoint = aws.String(stringsTrim(cfg.Endpoint)) + } + }) + } + return &S3Store{ + client: client, + bucket: stringsTrim(cfg.Bucket), + serverSideEncryption: stringsTrim(cfg.ServerSideEncryption), + sseKMSKeyID: stringsTrim(cfg.SSEKMSKeyID), + disableChecksumHeaders: cfg.DisableChecksumHeaders, + }, nil +} + +func loadS3AWSConfig(ctx context.Context, cfg S3StoreConfig) (aws.Config, error) { + optFns := make([]func(*config.LoadOptions) error, 0, s3LoadConfigOptionCapHint) + if stringsTrim(cfg.Region) != "" { + optFns = append(optFns, config.WithRegion(stringsTrim(cfg.Region))) + } + if stringsTrim(cfg.Profile) != "" { + optFns = append(optFns, config.WithSharedConfigProfile(stringsTrim(cfg.Profile))) + } + if stringsTrim(cfg.AccessKeyID) != "" || stringsTrim(cfg.SecretAccessKey) != "" || stringsTrim(cfg.SessionToken) != "" { + if stringsTrim(cfg.AccessKeyID) == "" || stringsTrim(cfg.SecretAccessKey) == "" { + return aws.Config{}, errors.Wrap(ErrInvalidOptions, "both s3 access key id and secret access key are required") + } + optFns = append(optFns, config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider( + stringsTrim(cfg.AccessKeyID), + stringsTrim(cfg.SecretAccessKey), + stringsTrim(cfg.SessionToken), + ))) + } + awsCfg, err := config.LoadDefaultConfig(ctx, optFns...) + if err != nil { + return aws.Config{}, errors.Wrap(err, "load s3 config") + } + return awsCfg, nil +} + +func (s *S3Store) PutObject(ctx context.Context, key string, body io.Reader, opts PutOptions) (ObjectInfo, error) { + if err := validatePutOptions(opts); err != nil { + return ObjectInfo{}, err + } + normalized, err := validateStoreObjectKey(key) + if err != nil { + return ObjectInfo{}, err + } + input, err := s.putObjectInput(normalized, body, opts) + if err != nil { + return ObjectInfo{}, err + } + if _, err := s.client.PutObject(ctx, input); err != nil { + if !isS3PreconditionFailed(err) { + return ObjectInfo{}, errors.Wrap(err, "put s3 object") + } + return s.verifyS3ExistingObject(ctx, normalized, opts) + } + return s.verifyS3PutObject(ctx, normalized, opts) +} + +func (s *S3Store) verifyS3PutObject(ctx context.Context, key string, opts PutOptions) (ObjectInfo, error) { + info, ok, err := s.HeadObject(ctx, key) + if err != nil { + return ObjectInfo{}, errors.Wrap(err, "head s3 object after put") + } + if !ok { + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s missing after put", key) + } + if info.Size != opts.Size || (info.SHA256 != "" && info.SHA256 != opts.SHA256) { + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s remote integrity mismatch", key) + } + if info.SHA256 == "" { + info.SHA256 = opts.SHA256 + } + return info, nil +} + +func (s *S3Store) putObjectInput(key string, body io.Reader, opts PutOptions) (*s3.PutObjectInput, error) { + input := &s3.PutObjectInput{ + Bucket: aws.String(s.bucket), + Key: aws.String(key), + Body: body, + ContentLength: aws.Int64(opts.Size), + IfNoneMatch: aws.String("*"), + Metadata: map[string]string{ + s3MetadataSHA256: opts.SHA256, + }, + } + if stringsTrim(opts.ContentType) != "" { + input.ContentType = aws.String(stringsTrim(opts.ContentType)) + } + if s.serverSideEncryption != "" { + input.ServerSideEncryption = types.ServerSideEncryption(s.serverSideEncryption) + } + if s.sseKMSKeyID != "" { + input.SSEKMSKeyId = aws.String(s.sseKMSKeyID) + } + if !s.disableChecksumHeaders { + checksum, err := sha256HexToBase64(opts.SHA256) + if err != nil { + return nil, err + } + input.ChecksumAlgorithm = types.ChecksumAlgorithmSha256 + input.ChecksumSHA256 = aws.String(checksum) + } + return input, nil +} + +func (s *S3Store) GetObject(ctx context.Context, key string) (io.ReadCloser, ObjectInfo, error) { + normalized, err := validateStoreObjectKey(key) + if err != nil { + return nil, ObjectInfo{}, err + } + out, err := s.client.GetObject(ctx, &s3.GetObjectInput{ + Bucket: aws.String(s.bucket), + Key: aws.String(normalized), + }) + if err != nil { + if isS3NotFound(err) { + return nil, ObjectInfo{}, errors.Wrapf(ErrObjectNotFound, "object %s", normalized) + } + return nil, ObjectInfo{}, errors.Wrap(err, "get s3 object") + } + if out.Body == nil { + return nil, ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s returned empty body", normalized) + } + info, err := s3ObjectInfo(normalized, out.ContentLength, out.Metadata, out.ChecksumSHA256) + if err != nil { + _ = out.Body.Close() + return nil, ObjectInfo{}, err + } + return out.Body, info, nil +} + +func (s *S3Store) HeadObject(ctx context.Context, key string) (ObjectInfo, bool, error) { + normalized, err := validateStoreObjectKey(key) + if err != nil { + return ObjectInfo{}, false, err + } + out, err := s.client.HeadObject(ctx, &s3.HeadObjectInput{ + Bucket: aws.String(s.bucket), + Key: aws.String(normalized), + }) + if err != nil { + if isS3NotFound(err) { + return ObjectInfo{}, false, nil + } + return ObjectInfo{}, false, errors.Wrap(err, "head s3 object") + } + info, err := s3ObjectInfo(normalized, out.ContentLength, out.Metadata, out.ChecksumSHA256) + if err != nil { + return ObjectInfo{}, false, err + } + return info, true, nil +} + +func (s *S3Store) verifyS3ExistingObject(ctx context.Context, key string, opts PutOptions) (ObjectInfo, error) { + info, ok, err := s.HeadObject(ctx, key) + if err != nil { + return ObjectInfo{}, err + } + if !ok { + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s conflicted but is not visible", key) + } + if info.Size == opts.Size && (info.SHA256 == "" || info.SHA256 == opts.SHA256) { + if info.SHA256 == "" { + info.SHA256 = opts.SHA256 + } + return info, nil + } + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s already exists with different content", key) +} + +func validateStoreObjectKey(key string) (string, error) { + normalized := normalizeObjectKey(key) + if normalized == "" || normalized == "." || normalized == ".." || strings.HasPrefix(normalized, "../") { + return "", errors.Wrapf(ErrInvalidOptions, "invalid object key %q", key) + } + return normalized, nil +} + +func s3ObjectInfo(key string, contentLength *int64, metadata map[string]string, checksumSHA256 *string) (ObjectInfo, error) { + if contentLength == nil || *contentLength < 0 { + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s missing content length", key) + } + sha, err := s3ObjectSHA256(key, metadata, checksumSHA256) + if err != nil { + return ObjectInfo{}, err + } + return ObjectInfo{Key: key, Size: *contentLength, SHA256: sha}, nil +} + +func s3ObjectSHA256(key string, metadata map[string]string, checksumSHA256 *string) (string, error) { + metadataSHA, err := s3MetadataSHA(metadata) + if err != nil { + return "", err + } + checksumSHA, err := s3ChecksumSHA(checksumSHA256) + if err != nil { + return "", err + } + switch { + case metadataSHA != "" && checksumSHA != "" && metadataSHA != checksumSHA: + return "", errors.Wrapf(ErrIntegrity, "s3 object %s sha256 metadata/checksum mismatch", key) + case metadataSHA != "": + return metadataSHA, nil + case checksumSHA != "": + return checksumSHA, nil + default: + return "", nil + } +} + +func s3MetadataSHA(metadata map[string]string) (string, error) { + if len(metadata) == 0 { + return "", nil + } + for key, value := range metadata { + if strings.EqualFold(key, s3MetadataSHA256) { + sha := stringsTrim(value) + if sha == "" { + return "", nil + } + if !isSHA256Hex(sha) { + return "", errors.Wrap(ErrIntegrity, "s3 object sha256 metadata is invalid") + } + return sha, nil + } + } + return "", nil +} + +func s3ChecksumSHA(checksum *string) (string, error) { + raw := stringsTrim(aws.ToString(checksum)) + if raw == "" { + return "", nil + } + decoded, err := base64.StdEncoding.DecodeString(raw) + if err != nil { + return "", errors.Wrap(ErrIntegrity, "s3 object sha256 checksum is invalid base64") + } + if len(decoded) != sha256.Size { + return "", errors.Wrapf(ErrIntegrity, "s3 object sha256 checksum has %d bytes", len(decoded)) + } + return hex.EncodeToString(decoded), nil +} + +func sha256HexToBase64(sha string) (string, error) { + decoded, err := hex.DecodeString(sha) + if err != nil || len(decoded) != sha256.Size { + return "", errors.Wrap(ErrInvalidOptions, "object sha256 must be 64 lowercase hex characters") + } + return base64.StdEncoding.EncodeToString(decoded), nil +} + +func isS3NotFound(err error) bool { + var notFound *types.NotFound + if errors.As(err, ¬Found) { + return true + } + var apiErr smithy.APIError + if errors.As(err, &apiErr) { + switch apiErr.ErrorCode() { + case "NotFound", "NoSuchKey", "404": + return true + } + } + return false +} + +func isS3PreconditionFailed(err error) bool { + var apiErr smithy.APIError + if !errors.As(err, &apiErr) { + return false + } + switch apiErr.ErrorCode() { + case "PreconditionFailed", "ConditionalRequestConflict": + return true + default: + return false + } +} diff --git a/internal/snapshotoffload/s3_store_test.go b/internal/snapshotoffload/s3_store_test.go new file mode 100644 index 000000000..5e4b04a78 --- /dev/null +++ b/internal/snapshotoffload/s3_store_test.go @@ -0,0 +1,244 @@ +package snapshotoffload + +import ( + "bytes" + "context" + "io" + "strings" + "sync" + "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" + "github.com/aws/smithy-go" + etcdraftengine "github.com/bootjp/elastickv/internal/raftengine/etcd" + "github.com/stretchr/testify/require" +) + +func TestPublishAndRestorePhysicalSnapshotRoundTripWithS3Store(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1opaque-s3-physical-snapshot-payload") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 50, 8, singlePeer()) + store := newTestS3Store(t, newFakeS3Client()) + + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-s3", + GroupID: 3, + SourceCluster: "cluster-s3", + BinaryVersion: "test-version", + CreatedAt: time.Unix(500, 0).UTC(), + }) + require.NoError(t, err) + require.Equal(t, uint64(50), manifest.SnapshotIndex) + require.Equal(t, int64(len(payload)), manifest.Payload.Bytes) + + restoreDataDir := root + "/restored-s3" + result, err := RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: store, + ManifestKey: manifest.ManifestKey, + DataDir: restoreDataDir, + Peers: []etcdraftengine.Peer{ + {NodeID: 2, ID: "n2", Address: "127.0.0.1:12002"}, + }, + }) + require.NoError(t, err) + require.Equal(t, manifest.Payload.SHA256, result.PayloadSHA256) +} + +func TestS3StorePutHeadGetPreservesIntegrityMetadata(t *testing.T) { + ctx := context.Background() + fake := newFakeS3Client() + store := newTestS3Store(t, fake) + body := []byte("s3-body") + sha := hexSHA256Bytes(body) + + info, err := store.PutObject(ctx, "snapshots/body.fsm", bytes.NewReader(body), PutOptions{ + Size: int64(len(body)), + SHA256: sha, + ContentType: "application/octet-stream", + }) + require.NoError(t, err) + require.Equal(t, sha, info.SHA256) + require.Equal(t, types.ChecksumAlgorithmSha256, fake.lastPutChecksumAlgorithm()) + + head, ok, err := store.HeadObject(ctx, "snapshots/body.fsm") + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, int64(len(body)), head.Size) + require.Equal(t, sha, head.SHA256) + + reader, gotInfo, err := store.GetObject(ctx, "snapshots/body.fsm") + require.NoError(t, err) + defer func() { require.NoError(t, reader.Close()) }() + require.Equal(t, sha, gotInfo.SHA256) + gotBody, err := io.ReadAll(reader) + require.NoError(t, err) + require.Equal(t, body, gotBody) +} + +func TestS3StoreReusesIdenticalExistingObject(t *testing.T) { + ctx := context.Background() + fake := newFakeS3Client() + store := newTestS3Store(t, fake) + body := []byte("same-body") + sha := hexSHA256Bytes(body) + _, err := store.PutObject(ctx, "snapshots/body.fsm", bytes.NewReader(body), PutOptions{ + Size: int64(len(body)), + SHA256: sha, + }) + require.NoError(t, err) + + info, err := store.PutObject(ctx, "snapshots/body.fsm", strings.NewReader("not-read-on-conflict"), PutOptions{ + Size: int64(len(body)), + SHA256: sha, + }) + require.NoError(t, err) + require.Equal(t, sha, info.SHA256) + require.Equal(t, 2, fake.putAttempts()) +} + +func TestS3StoreRejectsExistingObjectMismatch(t *testing.T) { + ctx := context.Background() + fake := newFakeS3Client() + store := newTestS3Store(t, fake) + existing := []byte("existing") + _, err := store.PutObject(ctx, "snapshots/body.fsm", bytes.NewReader(existing), PutOptions{ + Size: int64(len(existing)), + SHA256: hexSHA256Bytes(existing), + }) + require.NoError(t, err) + + candidate := []byte("candidate") + _, err = store.PutObject(ctx, "snapshots/body.fsm", bytes.NewReader(candidate), PutOptions{ + Size: int64(len(candidate)), + SHA256: hexSHA256Bytes(candidate), + }) + require.ErrorIs(t, err, ErrIntegrity) +} + +func TestS3StoreRejectsParentDirectoryKeys(t *testing.T) { + ctx := context.Background() + store := newTestS3Store(t, newFakeS3Client()) + emptySHA := hexSHA256Bytes(nil) + + _, err := store.PutObject(ctx, "..", bytes.NewReader(nil), PutOptions{SHA256: emptySHA}) + require.ErrorIs(t, err, ErrInvalidOptions) + _, _, err = store.GetObject(ctx, "a/../..") + require.ErrorIs(t, err, ErrInvalidOptions) + _, _, err = store.HeadObject(ctx, "../payload") + require.ErrorIs(t, err, ErrInvalidOptions) +} + +func newTestS3Store(t *testing.T, client *fakeS3Client) *S3Store { + t.Helper() + store, err := NewS3Store(context.Background(), S3StoreConfig{ + Client: client, + Bucket: "backup-bucket", + ForcePathStyle: true, + }) + require.NoError(t, err) + return store +} + +type fakeS3Client struct { + mu sync.Mutex + objects map[string]fakeS3Object + lastPut types.ChecksumAlgorithm + attempts int +} + +type fakeS3Object struct { + body []byte + metadata map[string]string + checksum *string +} + +func newFakeS3Client() *fakeS3Client { + return &fakeS3Client{objects: make(map[string]fakeS3Object)} +} + +func (c *fakeS3Client) PutObject(_ context.Context, input *s3.PutObjectInput, _ ...func(*s3.Options)) (*s3.PutObjectOutput, error) { + key := fakeS3ClientKey(input.Bucket, input.Key) + c.mu.Lock() + defer c.mu.Unlock() + c.attempts++ + c.lastPut = input.ChecksumAlgorithm + if _, ok := c.objects[key]; ok && aws.ToString(input.IfNoneMatch) == "*" { + return nil, &smithy.GenericAPIError{Code: "PreconditionFailed", Message: "exists"} + } + body, err := io.ReadAll(input.Body) + if err != nil { + return nil, err + } + if input.ContentLength != nil && int64(len(body)) != *input.ContentLength { + return nil, &smithy.GenericAPIError{Code: "BadDigest", Message: "content length mismatch"} + } + metadata := make(map[string]string, len(input.Metadata)) + for k, v := range input.Metadata { + metadata[k] = v + } + c.objects[key] = fakeS3Object{ + body: append([]byte(nil), body...), + metadata: metadata, + checksum: input.ChecksumSHA256, + } + return &s3.PutObjectOutput{}, nil +} + +func (c *fakeS3Client) HeadObject(_ context.Context, input *s3.HeadObjectInput, _ ...func(*s3.Options)) (*s3.HeadObjectOutput, error) { + c.mu.Lock() + defer c.mu.Unlock() + obj, ok := c.objects[fakeS3ClientKey(input.Bucket, input.Key)] + if !ok { + return nil, &types.NotFound{} + } + metadata := make(map[string]string, len(obj.metadata)) + for k, v := range obj.metadata { + metadata[k] = v + } + return &s3.HeadObjectOutput{ + ContentLength: aws.Int64(int64(len(obj.body))), + Metadata: metadata, + ChecksumSHA256: obj.checksum, + }, nil +} + +func (c *fakeS3Client) GetObject(_ context.Context, input *s3.GetObjectInput, _ ...func(*s3.Options)) (*s3.GetObjectOutput, error) { + c.mu.Lock() + defer c.mu.Unlock() + obj, ok := c.objects[fakeS3ClientKey(input.Bucket, input.Key)] + if !ok { + return nil, &types.NotFound{} + } + metadata := make(map[string]string, len(obj.metadata)) + for k, v := range obj.metadata { + metadata[k] = v + } + return &s3.GetObjectOutput{ + Body: io.NopCloser(bytes.NewReader(obj.body)), + ContentLength: aws.Int64(int64(len(obj.body))), + Metadata: metadata, + ChecksumSHA256: obj.checksum, + }, nil +} + +func (c *fakeS3Client) lastPutChecksumAlgorithm() types.ChecksumAlgorithm { + c.mu.Lock() + defer c.mu.Unlock() + return c.lastPut +} + +func (c *fakeS3Client) putAttempts() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.attempts +} + +func fakeS3ClientKey(bucket *string, key *string) string { + return aws.ToString(bucket) + "/" + aws.ToString(key) +} diff --git a/internal/snapshotoffload/store.go b/internal/snapshotoffload/store.go index 9c37baf11..d8fd4beb9 100644 --- a/internal/snapshotoffload/store.go +++ b/internal/snapshotoffload/store.go @@ -208,8 +208,12 @@ func (s *LocalStore) commitTempObject(key, tmpPath, finalPath string, expected O } return s.verifyExistingObject(key, finalPath, expected) } - if err := syncDir(filepath.Dir(finalPath)); err != nil { - return ObjectInfo{}, err + finalDir := filepath.Dir(finalPath) + if err := syncDir(finalDir); err != nil { + if removeErr := os.Remove(finalPath); removeErr != nil && !os.IsNotExist(removeErr) { + err = errors.CombineErrors(err, errors.WithStack(removeErr)) + } + return ObjectInfo{}, errors.WithStack(err) } return expected, nil } From 33250285fe335149f3f4cbb347f0cebf13105a70 Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 10:23:00 +0900 Subject: [PATCH 8/9] redis: stabilize xread final check --- adapter/redis_compat_commands_stream_test.go | 18 ++++++-- adapter/redis_stream_cmds.go | 3 +- cmd/elastickv-snapshot-offload/main.go | 23 ++++++++++- cmd/elastickv-snapshot-offload/main_test.go | 43 ++++++++++++++++++++ 4 files changed, 80 insertions(+), 7 deletions(-) diff --git a/adapter/redis_compat_commands_stream_test.go b/adapter/redis_compat_commands_stream_test.go index a0cc0000e..e116d3c6a 100644 --- a/adapter/redis_compat_commands_stream_test.go +++ b/adapter/redis_compat_commands_stream_test.go @@ -213,12 +213,12 @@ func TestRedis_StreamXReadBlockChecksWrongTypeAtDeadline(t *testing.T) { streams, err := rdbReader.XRead(ctx, &redis.XReadArgs{ Streams: []string{key, "$"}, Count: 1, - Block: 100 * time.Millisecond, + Block: 2 * time.Second, }).Result() resultCh <- readResult{streams: streams, err: err} }() - time.Sleep(20 * time.Millisecond) + requireStreamWaiterRegistered(t, nodes[0].redisServer.streamWaiters, key) require.NoError(t, rdbWriter.Set(ctx, key, "now-a-string", 0).Err()) select { @@ -226,11 +226,23 @@ func TestRedis_StreamXReadBlockChecksWrongTypeAtDeadline(t *testing.T) { require.Error(t, res.err) require.Contains(t, res.err.Error(), "WRONGTYPE") require.Empty(t, res.streams) - case <-time.After(2 * time.Second): + case <-time.After(4 * time.Second): t.Fatal("XREAD BLOCK did not return after wrong-type overwrite") } } +func requireStreamWaiterRegistered(t *testing.T, reg *keyWaiterRegistry, key string) { + t.Helper() + require.Eventually(t, func() bool { + if reg == nil { + return false + } + reg.mu.Lock() + defer reg.mu.Unlock() + return len(reg.waiters[key]) > 0 + }, 2*time.Second, 10*time.Millisecond) +} + // TestRedis_StreamCommandsRejectWrongType locks down the wrongType // detection on the stream fast path: keyTypeAtExpect short-circuits to // the slow path when the expected (stream) prefixes return empty, so diff --git a/adapter/redis_stream_cmds.go b/adapter/redis_stream_cmds.go index ea8139986..15e12a580 100644 --- a/adapter/redis_stream_cmds.go +++ b/adapter/redis_stream_cmds.go @@ -1283,8 +1283,7 @@ func (r *RedisServer) xreadFinalCheck(conn redcon.Conn, req xreadRequest) bool { if ok := r.runWithHeavyCommandSlot(func() { results, err = r.xreadOnce(ctx, req) }); !ok { - conn.WriteError(errRedisHeavyCommandPoolFull.Error()) - return true + return false } if err != nil { if isXReadIterCtxError(err) { diff --git a/cmd/elastickv-snapshot-offload/main.go b/cmd/elastickv-snapshot-offload/main.go index e1db6abb5..7f8c11462 100644 --- a/cmd/elastickv-snapshot-offload/main.go +++ b/cmd/elastickv-snapshot-offload/main.go @@ -26,6 +26,7 @@ const ( commandRestore = "restore" storeLocal = "local" storeS3 = "s3" + s3SSEAWSKMS = "aws:kms" ) type storeFlags struct { @@ -121,6 +122,9 @@ func parsePublishFlags(argv []string) (*publishConfig, error) { if err := fs.Parse(argv); err != nil { return nil, errors.WithStack(err) } + if err := rejectPositionalArgs(fs); err != nil { + return nil, err + } if strings.TrimSpace(cfg.dataDir) == "" { return nil, errors.New("--data-dir is required") } @@ -144,6 +148,9 @@ func parseRestoreFlags(argv []string) (*restoreConfig, error) { if err := fs.Parse(argv); err != nil { return nil, errors.WithStack(err) } + if err := rejectPositionalArgs(fs); err != nil { + return nil, err + } if strings.TrimSpace(cfg.manifestKey) == "" { return nil, errors.New("--manifest-key is required") } @@ -159,6 +166,13 @@ func parseRestoreFlags(argv []string) (*restoreConfig, error) { return cfg, nil } +func rejectPositionalArgs(fs *flag.FlagSet) error { + if fs.NArg() == 0 { + return nil + } + return errors.Errorf("unexpected positional argument %q", fs.Arg(0)) +} + func addStoreFlags(fs *flag.FlagSet, cfg *storeFlags) { cfg.storeKind = storeLocal cfg.s3Region = "us-east-1" @@ -188,8 +202,13 @@ func validateStoreFlags(cfg storeFlags) error { default: return errors.Errorf("unknown --store %q", cfg.storeKind) } - if strings.TrimSpace(cfg.s3KMSKeyID) != "" && strings.TrimSpace(cfg.s3ServerSideEncryption) == "" { - return errors.New("--s3-kms-key-id requires --s3-sse") + if strings.TrimSpace(cfg.s3KMSKeyID) != "" { + if strings.TrimSpace(cfg.s3ServerSideEncryption) == "" { + return errors.New("--s3-kms-key-id requires --s3-sse") + } + if strings.TrimSpace(cfg.s3ServerSideEncryption) != s3SSEAWSKMS { + return errors.New("--s3-kms-key-id requires --s3-sse=aws:kms") + } } return nil } diff --git a/cmd/elastickv-snapshot-offload/main_test.go b/cmd/elastickv-snapshot-offload/main_test.go index dc4700595..d2991b401 100644 --- a/cmd/elastickv-snapshot-offload/main_test.go +++ b/cmd/elastickv-snapshot-offload/main_test.go @@ -71,6 +71,49 @@ func TestSnapshotOffloadCLIRequiresLocalRoot(t *testing.T) { require.Equal(t, exitUserErr, code) } +func TestSnapshotOffloadCLIRejectsPositionalArgs(t *testing.T) { + _, err := parsePublishFlags([]string{ + "--store", storeLocal, + "--local-root", "objects", + "--data-dir", "data", + "--group-id", "1", + "extra", + }) + require.ErrorContains(t, err, "unexpected positional argument") + + _, err = parseRestoreFlags([]string{ + "--store", storeLocal, + "--local-root", "objects", + "--manifest-key", "manifest.json", + "--data-dir", "data", + "--peers", "n1=127.0.0.1:12001", + "extra", + }) + require.ErrorContains(t, err, "unexpected positional argument") +} + +func TestSnapshotOffloadCLIS3KMSRequiresAWSKMS(t *testing.T) { + _, err := parsePublishFlags([]string{ + "--store", storeS3, + "--s3-bucket", "bucket", + "--s3-sse", s3SSEAWSKMS, + "--s3-kms-key-id", "key-id", + "--data-dir", "data", + "--group-id", "1", + }) + require.NoError(t, err) + + _, err = parsePublishFlags([]string{ + "--store", storeS3, + "--s3-bucket", "bucket", + "--s3-sse", "AES256", + "--s3-kms-key-id", "key-id", + "--data-dir", "data", + "--group-id", "1", + }) + require.ErrorContains(t, err, "--s3-kms-key-id requires --s3-sse=aws:kms") +} + func seedCLISnapshot(t *testing.T, root string, payload []byte, index uint64, term uint64) string { t.Helper() input := filepath.Join(root, "source.fsm") From a48f7acfe65026d7d72b9b8beb9c1cb877c60f8b Mon Sep 17 00:00:00 2001 From: bootjp Date: Thu, 23 Jul 2026 10:39:45 +0900 Subject: [PATCH 9/9] redis: raise proxy headroom --- adapter/redis.go | 1 + adapter/redis_peer_limiter.go | 10 +- adapter/redis_peer_limiter_test.go | 6 +- adapter/redis_stream_cmds.go | 27 ++-- docs/redis-proxy-deployment.md | 4 +- internal/raftengine/etcd/engine.go | 37 +++++ .../etcd/engine_applied_index_test.go | 9 +- internal/snapshotoffload/offload_test.go | 147 +++++++++++++++++ internal/snapshotoffload/publish.go | 93 ++++++++--- internal/snapshotoffload/restore.go | 153 +++++++++++++++--- internal/snapshotoffload/s3_store.go | 37 ++++- internal/snapshotoffload/s3_store_test.go | 57 +++++++ kv/lease_warmup_test.go | 38 +++++ kv/sharded_coordinator.go | 30 +++- kv/tso.go | 11 +- kv/tso_test.go | 25 +++ proxy/backend.go | 7 +- 17 files changed, 611 insertions(+), 81 deletions(-) diff --git a/adapter/redis.go b/adapter/redis.go index 2ac41a182..01a7cdf96 100644 --- a/adapter/redis.go +++ b/adapter/redis.go @@ -123,6 +123,7 @@ const ( // Keep normal producer wakeups event-driven while reducing idle fallback scans // from blocked consumers waiting on empty collections. defaultRedisBlockWaitFallback = time.Second + redisFinalTypeCheckTimeout = 100 * time.Millisecond redisFlushLegacyTimeout = 10 * time.Minute redisRelayPublishTimeout = 2 * time.Second redisTraceArgLimit = 6 diff --git a/adapter/redis_peer_limiter.go b/adapter/redis_peer_limiter.go index c9d44329a..88a3077fa 100644 --- a/adapter/redis_peer_limiter.go +++ b/adapter/redis_peer_limiter.go @@ -9,10 +9,12 @@ import ( ) const ( - redisPerPeerLimitEnv = "ELASTICKV_REDIS_PER_PEER_CONNECTIONS" - defaultRedisPerPeerConnectionCap = 64 - redisPeerLimitError = "ERR max connections per client exceeded" - unknownRedisPeer = "unknown" + redisPerPeerLimitEnv = "ELASTICKV_REDIS_PER_PEER_CONNECTIONS" + defaultRedisProxyPoolPeerCap = 64 + defaultRedisDedicatedPeerHeadroom = 64 + defaultRedisPerPeerConnectionCap = defaultRedisProxyPoolPeerCap + defaultRedisDedicatedPeerHeadroom + redisPeerLimitError = "ERR max connections per client exceeded" + unknownRedisPeer = "unknown" ) type redisPeerLimiter struct { diff --git a/adapter/redis_peer_limiter_test.go b/adapter/redis_peer_limiter_test.go index 74312d7fb..d40f23b58 100644 --- a/adapter/redis_peer_limiter_test.go +++ b/adapter/redis_peer_limiter_test.go @@ -3,19 +3,19 @@ package adapter import ( "testing" + "github.com/bootjp/elastickv/proxy" "github.com/stretchr/testify/require" ) const ( - testPeerLimit = 2 - defaultElasticKVProxyPoolSizeForTest = 64 + testPeerLimit = 2 ) func TestRedisPeerLimiterDefaultMatchesProxyPool(t *testing.T) { t.Setenv(redisPerPeerLimitEnv, "") limiter := newDefaultRedisPeerLimiter() require.NotNil(t, limiter) - require.Equal(t, defaultElasticKVProxyPoolSizeForTest, limiter.limit) + require.Equal(t, proxy.DefaultElasticKVBackendOptions().PoolSize+defaultRedisDedicatedPeerHeadroom, limiter.limit) } func TestRedisPeerLimiterRejectsAndReleases(t *testing.T) { diff --git a/adapter/redis_stream_cmds.go b/adapter/redis_stream_cmds.go index 15e12a580..52efed53d 100644 --- a/adapter/redis_stream_cmds.go +++ b/adapter/redis_stream_cmds.go @@ -1275,13 +1275,12 @@ func isXReadIterCtxError(err error) bool { } } -func (r *RedisServer) xreadFinalCheck(conn redcon.Conn, req xreadRequest) bool { - ctx, cancel := context.WithTimeout(r.handlerContext(), redisDispatchTimeout) +func (r *RedisServer) xreadFinalTypeCheck(conn redcon.Conn, req xreadRequest) bool { + ctx, cancel := context.WithTimeout(r.handlerContext(), redisFinalTypeCheckTimeout) defer cancel() - var results []xreadResult var err error if ok := r.runWithHeavyCommandSlot(func() { - results, err = r.xreadOnce(ctx, req) + err = r.xreadCheckTypes(ctx, req) }); !ok { return false } @@ -1292,15 +1291,25 @@ func (r *RedisServer) xreadFinalCheck(conn redcon.Conn, req xreadRequest) bool { writeRedisError(conn, err) return true } - if len(results) > 0 { - writeXReadResults(conn, results) - return true - } return false } +func (r *RedisServer) xreadCheckTypes(ctx context.Context, req xreadRequest) error { + for _, key := range req.keys { + readTS := r.readTS() + typ, err := r.keyTypeAtExpect(ctx, key, readTS, redisTypeStream) + if err != nil { + return err + } + if typ != redisTypeNone && typ != redisTypeStream { + return wrongTypeError() + } + } + return nil +} + func (r *RedisServer) writeXReadFinalOrNull(conn redcon.Conn, req xreadRequest) { - if r.xreadFinalCheck(conn, req) { + if r.xreadFinalTypeCheck(conn, req) { return } conn.WriteNull() diff --git a/docs/redis-proxy-deployment.md b/docs/redis-proxy-deployment.md index 743fc5b8d..1f0468936 100644 --- a/docs/redis-proxy-deployment.md +++ b/docs/redis-proxy-deployment.md @@ -414,7 +414,7 @@ groups: | Parameter | Value | Description | |-----------|-------|-------------| | Redis connection pool size | 128 | Default go-redis pool size for Redis | -| ElasticKV connection pool size | 64 | Default per-leader pool; keep within the server per-peer connection limit | +| ElasticKV connection pool size | 64 | Default per-leader command pool; leave server per-peer headroom for dedicated PubSub connections | | Dial timeout | 5s | Backend connection timeout | | Read timeout | 3s | Backend read timeout | | Write timeout | 3s | Backend write timeout | @@ -439,7 +439,7 @@ Recommended shutdown order: `redis-proxy -> application -> Redis / ElasticKV`. ### Secondary writes are falling behind - Check `proxy_async_queue_depth`, `proxy_async_queue_delay_seconds`, and `proxy_async_drops_by_queue_total` before increasing concurrency. - Check `proxy_backend_pool_pending_requests` and the `waits`/`timeouts` pool events. Pool waits mean concurrency is too high for the configured pool. -- Keep `ELASTICKV_REDIS_PER_PEER_CONNECTIONS` at least as high as `-elastickv-pool-size`; keep `-secondary-write-concurrency` at or below the pool size. +- Keep `ELASTICKV_REDIS_PER_PEER_CONNECTIONS` above `-elastickv-pool-size`; PubSub and shadow PubSub use dedicated connections outside the command pool. Keep `-secondary-write-concurrency` at or below the pool size. - A sustained `expired` rate means secondary throughput is below ingress. Increasing queue size only delays the loss; profile ElasticKV before raising concurrency. ### High divergence count diff --git a/internal/raftengine/etcd/engine.go b/internal/raftengine/etcd/engine.go index dfff8a8d6..7a8680909 100644 --- a/internal/raftengine/etcd/engine.go +++ b/internal/raftengine/etcd/engine.go @@ -409,6 +409,8 @@ type Engine struct { snapshotMu sync.Mutex protectedReceivedFSMSnaps map[uint64]int pendingReceivedFSMSnapshotStep map[uint64]int + pendingReceivedFSMReleaseIndex atomic.Uint64 + pendingReceivedFSMReleaseRun atomic.Bool dispatchDropCount atomic.Uint64 dispatchErrorCount atomic.Uint64 @@ -3366,12 +3368,47 @@ func (e *Engine) releaseProtectedReceivedFSMSnapshotsUpTo(index uint64) { return } if !e.snapshotMu.TryLock() { + e.scheduleProtectedReceivedFSMSnapshotRelease(index) return } defer e.snapshotMu.Unlock() e.releaseProtectedReceivedFSMSnapshotsUpToLocked(index) } +func (e *Engine) scheduleProtectedReceivedFSMSnapshotRelease(index uint64) { + for { + prev := e.pendingReceivedFSMReleaseIndex.Load() + if index <= prev { + break + } + if e.pendingReceivedFSMReleaseIndex.CompareAndSwap(prev, index) { + break + } + } + if e.pendingReceivedFSMReleaseRun.CompareAndSwap(false, true) { + go e.drainPendingProtectedReceivedFSMSnapshotRelease() + } +} + +func (e *Engine) drainPendingProtectedReceivedFSMSnapshotRelease() { + for { + index := e.pendingReceivedFSMReleaseIndex.Swap(0) + if index != 0 { + e.snapshotMu.Lock() + e.releaseProtectedReceivedFSMSnapshotsUpToLocked(index) + e.snapshotMu.Unlock() + } + + e.pendingReceivedFSMReleaseRun.Store(false) + if e.pendingReceivedFSMReleaseIndex.Load() == 0 { + return + } + if !e.pendingReceivedFSMReleaseRun.CompareAndSwap(false, true) { + return + } + } +} + func (e *Engine) unprotectReceivedFSMSnapshotTokenIfApplied(msg raftpb.Message) { index, ok := receivedFSMSnapshotTokenIndex(msg) if !ok || index > e.appliedIndex.Load() { diff --git a/internal/raftengine/etcd/engine_applied_index_test.go b/internal/raftengine/etcd/engine_applied_index_test.go index 2e67c2836..6a47d431e 100644 --- a/internal/raftengine/etcd/engine_applied_index_test.go +++ b/internal/raftengine/etcd/engine_applied_index_test.go @@ -162,7 +162,7 @@ func TestPersistReadyWithSnapshotHoldsSnapshotMuThroughSaveSnap(t *testing.T) { require.Empty(t, e.protectedReceivedFSMSnaps) } -func TestReleaseProtectedReceivedFSMSnapshotsUpToDoesNotBlockSnapshotMu(t *testing.T) { +func TestReleaseProtectedReceivedFSMSnapshotsUpToRetriesAfterSnapshotMuUnlock(t *testing.T) { e := &Engine{ protectedReceivedFSMSnaps: map[uint64]int{7: 1}, } @@ -182,8 +182,11 @@ func TestReleaseProtectedReceivedFSMSnapshotsUpToDoesNotBlockSnapshotMu(t *testi require.Equal(t, map[uint64]int{7: 1}, e.protectedReceivedFSMSnaps) e.snapshotMu.Unlock() - e.releaseProtectedReceivedFSMSnapshotsUpTo(7) - require.Empty(t, e.protectedReceivedFSMSnaps) + require.Eventually(t, func() bool { + e.snapshotMu.Lock() + defer e.snapshotMu.Unlock() + return len(e.protectedReceivedFSMSnaps) == 0 + }, time.Second, 10*time.Millisecond) } func TestProtectReceivedFSMSnapshotRechecksAppliedIndexUnderLock(t *testing.T) { diff --git a/internal/snapshotoffload/offload_test.go b/internal/snapshotoffload/offload_test.go index 83b1ed76b..f90bae9d3 100644 --- a/internal/snapshotoffload/offload_test.go +++ b/internal/snapshotoffload/offload_test.go @@ -209,6 +209,46 @@ func TestPublishReusesExistingManifestWhenCreatedAtOmitted(t *testing.T) { require.Equal(t, first.CreatedAt, second.CreatedAt) } +func TestPutManifestReusesExistingManifestAfterCreateConflict(t *testing.T) { + ctx := context.Background() + store := newTestLocalStore(t, filepath.Join(t.TempDir(), "objects")) + key, err := manifestKey("cluster-a", 1, 20, 13) + require.NoError(t, err) + payloadSHA := hexSHA256Bytes([]byte("payload")) + existing := Manifest{ + SchemaVersion: ManifestSchemaVersion, + CreatedAt: time.Unix(400, 0).UTC(), + SourceCluster: "cluster-a", + GroupID: 1, + SnapshotIndex: 20, + SnapshotTerm: 13, + ConfState: ManifestConfState{ + Voters: []uint64{1}, + }, + Payload: PayloadDescriptor{ + Key: "cluster-a/v1/payloads/test.fsm", + Bytes: int64(len("payload")), + SHA256: payloadSHA, + }, + ManifestKey: key, + } + data, _, err := existing.MarshalCanonical() + require.NoError(t, err) + _, err = store.PutObject(ctx, key, bytes.NewReader(data), PutOptions{ + Size: int64(len(data)), + SHA256: hexSHA256Bytes(data), + ContentType: "application/json", + }) + require.NoError(t, err) + + candidate := existing + candidate.CreatedAt = time.Unix(401, 0).UTC() + racingStore := &headMissOnceStore{ObjectStore: store, key: key} + require.NoError(t, putManifest(ctx, racingStore, &candidate, true)) + require.Equal(t, existing.CreatedAt, candidate.CreatedAt) + require.NotEmpty(t, candidate.ManifestSHA256) +} + func TestPublishRejectsGroupZeroWithoutSourceClusterBeforePayload(t *testing.T) { ctx := context.Background() root := t.TempDir() @@ -278,6 +318,34 @@ func TestLoadManifestRejectsStaleSelfHash(t *testing.T) { require.ErrorIs(t, err, ErrIntegrity) } +func TestRestoreInlineManifestRejectsStaleSelfHashBeforePayloadDownload(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-for-inline-stale-manifest-hash") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 21, 14, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + tampered := *manifest + tampered.SnapshotIndex++ + tracked := &countingObjectStore{ObjectStore: store} + + _, err = RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: tracked, + Manifest: &tampered, + DataDir: filepath.Join(root, "restored"), + Peers: singlePeer(), + }) + require.ErrorIs(t, err, ErrIntegrity) + require.Zero(t, tracked.getObjectCalls) +} + func TestLoadManifestRejectsOversizedBody(t *testing.T) { ctx := context.Background() store := newTestLocalStore(t, filepath.Join(t.TempDir(), "objects")) @@ -381,6 +449,61 @@ func TestRestorePreflightsExistingDestinationBeforePayloadDownload(t *testing.T) require.ErrorIs(t, err, etcdraftengine.ErrExternalSnapshotRestoreExists) } +func TestRestoreRejectsInvalidPeersBeforePayloadDownload(t *testing.T) { + ctx := context.Background() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-invalid-restore-peers") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 22, 15, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + manifest, err := PublishPersistedSnapshot(ctx, PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + tracked := &countingObjectStore{ObjectStore: store} + + _, err = RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: tracked, + Manifest: manifest, + DataDir: filepath.Join(root, "restored"), + Peers: []etcdraftengine.Peer{ + {NodeID: 0, ID: "n0", Address: "127.0.0.1:12000"}, + }, + }) + require.ErrorIs(t, err, ErrInvalidOptions) + require.Zero(t, tracked.getObjectCalls) +} + +func TestRestoreHonorsCancelledContextBeforePayloadDownload(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + root := t.TempDir() + payload := []byte("EKVTHLC1payload-cancelled-before-restore-download") + sourceDataDir := seedPhysicalSnapshot(t, root, payload, 23, 16, singlePeer()) + store := newTestLocalStore(t, filepath.Join(root, "objects")) + manifest, err := PublishPersistedSnapshot(context.Background(), PublishOptions{ + Store: store, + DataDir: sourceDataDir, + Prefix: "cluster-a", + GroupID: 1, + SourceCluster: "cluster-a", + }) + require.NoError(t, err) + tracked := &countingObjectStore{ObjectStore: store} + + _, err = RestorePhysicalSnapshot(ctx, RestoreOptions{ + Store: tracked, + Manifest: manifest, + DataDir: filepath.Join(root, "restored"), + Peers: singlePeer(), + }) + require.ErrorIs(t, err, context.Canceled) + require.Zero(t, tracked.getObjectCalls) +} + func TestPrepareRestoreDownloadDirCreatesParentAndCleansOnlyStaleDirs(t *testing.T) { root := t.TempDir() dataDir := filepath.Join(root, "missing-parent", "restored") @@ -449,6 +572,30 @@ func loadTestManifest(t *testing.T, ctx context.Context, store ObjectStore, key return manifest } +type countingObjectStore struct { + ObjectStore + getObjectCalls int +} + +func (s *countingObjectStore) GetObject(ctx context.Context, key string) (io.ReadCloser, ObjectInfo, error) { + s.getObjectCalls++ + return s.ObjectStore.GetObject(ctx, key) +} + +type headMissOnceStore struct { + ObjectStore + key string + miss bool +} + +func (s *headMissOnceStore) HeadObject(ctx context.Context, key string) (ObjectInfo, bool, error) { + if !s.miss && normalizeObjectKey(key) == normalizeObjectKey(s.key) { + s.miss = true + return ObjectInfo{}, false, nil + } + return s.ObjectStore.HeadObject(ctx, key) +} + func singlePeer() []etcdraftengine.Peer { return []etcdraftengine.Peer{{NodeID: 1, ID: "n1", Address: "127.0.0.1:12001"}} } diff --git a/internal/snapshotoffload/publish.go b/internal/snapshotoffload/publish.go index d18b69258..734359d20 100644 --- a/internal/snapshotoffload/publish.go +++ b/internal/snapshotoffload/publish.go @@ -118,27 +118,76 @@ func putManifest(ctx context.Context, store ObjectStore, manifest *Manifest, reu if err != nil { return err } + size := int64(len(data)) objectSHA := hexSHA256Bytes(data) - if exists, err := verifyExistingManifest(ctx, store, manifest, int64(len(data)), objectSHA, reuseExistingCreatedAt); err != nil { + if exists, err := verifyExistingManifest(ctx, store, manifest, size, objectSHA, reuseExistingCreatedAt); err != nil { return err } else if exists { - if manifest.ManifestSHA256 == "" { - manifest.ManifestSHA256 = manifestSHA - } return nil } + if err := createManifestObject(ctx, store, manifest, data, size, objectSHA, reuseExistingCreatedAt); err != nil { + return err + } + manifest.ManifestSHA256 = manifestSHA + return verifyCommittedManifest(ctx, store, manifest, size, objectSHA, reuseExistingCreatedAt) +} + +func createManifestObject( + ctx context.Context, + store ObjectStore, + manifest *Manifest, + data []byte, + size int64, + objectSHA string, + reuseExistingCreatedAt bool, +) error { info, err := store.PutObject(ctx, manifest.ManifestKey, bytes.NewReader(data), PutOptions{ - Size: int64(len(data)), + Size: size, SHA256: objectSHA, ContentType: "application/json", }) if err != nil { - return errors.Wrap(err, "put snapshot manifest") + return handleManifestPutError(ctx, store, manifest, size, objectSHA, reuseExistingCreatedAt, err) } - if info.Size != int64(len(data)) || (info.SHA256 != "" && info.SHA256 != objectSHA) { + if info.Size != size || (info.SHA256 != "" && info.SHA256 != objectSHA) { return errors.Wrapf(ErrIntegrity, "manifest object %s remote integrity mismatch", manifest.ManifestKey) } - manifest.ManifestSHA256 = manifestSHA + return nil +} + +func handleManifestPutError( + ctx context.Context, + store ObjectStore, + manifest *Manifest, + size int64, + objectSHA string, + reuseExistingCreatedAt bool, + err error, +) error { + if !errors.Is(err, ErrIntegrity) { + return errors.Wrap(err, "put snapshot manifest") + } + if exists, verifyErr := verifyExistingManifest(ctx, store, manifest, size, objectSHA, reuseExistingCreatedAt); verifyErr != nil { + return errors.Wrap(verifyErr, "verify conflicting snapshot manifest") + } else if exists { + return nil + } + return errors.Wrap(err, "put snapshot manifest") +} + +func verifyCommittedManifest( + ctx context.Context, + store ObjectStore, + manifest *Manifest, + size int64, + objectSHA string, + reuseExistingCreatedAt bool, +) error { + if exists, err := verifyExistingManifest(ctx, store, manifest, size, objectSHA, reuseExistingCreatedAt); err != nil { + return errors.Wrap(err, "verify committed snapshot manifest") + } else if !exists { + return errors.Wrapf(ErrIntegrity, "manifest object %s missing after put", manifest.ManifestKey) + } return nil } @@ -157,37 +206,29 @@ func verifyExistingManifest( if !ok { return false, nil } - if matches, err := existingStoreObjectMatches(ctx, store, manifest.ManifestKey, info, size, sha); err != nil { - return true, errors.Wrap(err, "verify existing snapshot manifest") - } else if matches { - return true, nil + if info.Size != size && !reuseExistingCreatedAt { + return true, errors.Wrapf(ErrIntegrity, "manifest object %s already exists with different size", manifest.ManifestKey) } - if !reuseExistingCreatedAt { - return true, errors.Wrapf(ErrIntegrity, "manifest object %s already exists with different content", manifest.ManifestKey) + if info.SHA256 != "" && info.SHA256 != sha && !reuseExistingCreatedAt { + return true, errors.Wrapf(ErrIntegrity, "manifest object %s already exists with different sha256", manifest.ManifestKey) } existing, err := LoadManifest(ctx, store, manifest.ManifestKey) if err != nil { return true, errors.Wrap(err, "load existing snapshot manifest") } - if !sameManifestExceptCreation(existing, *manifest) { + if !manifestMatchesCandidate(existing, *manifest, reuseExistingCreatedAt) { return true, errors.Wrapf(ErrIntegrity, "manifest object %s already exists with different content", manifest.ManifestKey) } *manifest = existing return true, nil } -func existingStoreObjectMatches(ctx context.Context, store ObjectStore, key string, info ObjectInfo, size int64, sha string) (bool, error) { - if info.Size != size { - return false, nil - } - if info.SHA256 != "" { - return info.SHA256 == sha, nil +func manifestMatchesCandidate(existing Manifest, candidate Manifest, reuseExistingCreatedAt bool) bool { + if reuseExistingCreatedAt { + return sameManifestExceptCreation(existing, candidate) } - gotSize, gotSHA, err := hashExistingStoreObject(ctx, store, key) - if err != nil { - return false, err - } - return gotSize == size && gotSHA == sha, nil + candidate.ManifestSHA256 = existing.ManifestSHA256 + return reflect.DeepEqual(existing, candidate) } func sameManifestExceptCreation(existing Manifest, candidate Manifest) bool { diff --git a/internal/snapshotoffload/restore.go b/internal/snapshotoffload/restore.go index 5a6683650..0a45e819b 100644 --- a/internal/snapshotoffload/restore.go +++ b/internal/snapshotoffload/restore.go @@ -30,6 +30,7 @@ const ( ) func LoadManifest(ctx context.Context, store ObjectStore, key string) (Manifest, error) { + ctx = restoreContext(ctx) if store == nil { return Manifest{}, errors.Wrap(ErrInvalidOptions, "object store is required") } @@ -53,31 +54,15 @@ func LoadManifest(ctx context.Context, store ObjectStore, key string) (Manifest, } func RestorePhysicalSnapshot(ctx context.Context, opts RestoreOptions) (*etcdraftengine.ExternalSnapshotRestoreResult, error) { + ctx = restoreContext(ctx) if err := validateRestoreOptions(opts); err != nil { return nil, err } - manifest, err := restoreManifest(ctx, opts) - if err != nil { - return nil, err - } - if err := validateManifest(manifest); err != nil { - return nil, err - } - if err := ensureRestoreDestinationAbsent(opts.DataDir); err != nil { - return nil, err - } - downloadDir, err := prepareRestoreDownloadDir(opts.DataDir) + manifest, payloadPath, cleanup, err := prepareRestorePayload(ctx, opts) if err != nil { return nil, err } - defer func() { _ = os.RemoveAll(downloadDir) }() - payloadPath := filepath.Join(downloadDir, "payload.fsm") - if err := downloadVerifiedPayload(ctx, opts.Store, manifest, payloadPath); err != nil { - return nil, err - } - if err := ctx.Err(); err != nil { - return nil, errors.WithStack(err) - } + defer cleanup() result, err := etcdraftengine.PreparePhysicalSnapshotRestore(etcdraftengine.PhysicalSnapshotRestoreOptions{ Context: ctx, InputFSMPath: payloadPath, @@ -93,6 +78,64 @@ func RestorePhysicalSnapshot(ctx context.Context, opts RestoreOptions) (*etcdraf return result, nil } +func prepareRestorePayload(ctx context.Context, opts RestoreOptions) (Manifest, string, func(), error) { + if err := checkRestoreContext(ctx); err != nil { + return Manifest{}, "", nil, err + } + manifest, err := restoreManifest(ctx, opts) + if err != nil { + return Manifest{}, "", nil, err + } + if err := validateManifest(manifest); err != nil { + return Manifest{}, "", nil, err + } + if err := checkRestorePreflight(ctx, opts.DataDir); err != nil { + return Manifest{}, "", nil, err + } + downloadDir, err := prepareRestoreDownloadDir(opts.DataDir) + if err != nil { + return Manifest{}, "", nil, err + } + cleanup := func() { _ = os.RemoveAll(downloadDir) } + payloadPath := filepath.Join(downloadDir, "payload.fsm") + if err := downloadRestorePayload(ctx, opts.Store, manifest, payloadPath); err != nil { + cleanup() + return Manifest{}, "", nil, err + } + return manifest, payloadPath, cleanup, nil +} + +func checkRestorePreflight(ctx context.Context, dataDir string) error { + if err := checkRestoreContext(ctx); err != nil { + return err + } + return ensureRestoreDestinationAbsent(dataDir) +} + +func downloadRestorePayload(ctx context.Context, store ObjectStore, manifest Manifest, payloadPath string) error { + if err := checkRestoreContext(ctx); err != nil { + return err + } + if err := downloadVerifiedPayload(ctx, store, manifest, payloadPath); err != nil { + return err + } + return checkRestoreContext(ctx) +} + +func checkRestoreContext(ctx context.Context) error { + if err := ctx.Err(); err != nil { + return errors.WithStack(err) + } + return nil +} + +func restoreContext(ctx context.Context) context.Context { + if ctx == nil { + return context.Background() + } + return ctx +} + func readLimitedManifest(ctx context.Context, body io.Reader) ([]byte, error) { var buf bytes.Buffer limited := io.LimitReader(contextReader{ctx: ctx, reader: body}, maxManifestBytes+1) @@ -167,14 +210,84 @@ func validateRestoreOptions(opts RestoreOptions) error { return errors.Wrap(ErrInvalidOptions, "data dir is required") case len(opts.Peers) == 0: return errors.Wrap(ErrInvalidOptions, "restore peers are required") + default: + return validateRestorePeers(opts.Peers) + } +} + +func validateRestorePeers(peers []etcdraftengine.Peer) error { + seenNodeIDs := make(map[uint64]struct{}, len(peers)) + seenIDs := make(map[string]struct{}, len(peers)) + voters := 0 + for i, peer := range peers { + isVoter, err := validateRestorePeer(i, peer, seenNodeIDs, seenIDs) + if err != nil { + return err + } + if isVoter { + voters++ + } + } + if voters == 0 { + return errors.Wrap(ErrInvalidOptions, "restore peers require at least one voter") + } + return nil +} + +func validateRestorePeer( + i int, + peer etcdraftengine.Peer, + seenNodeIDs map[uint64]struct{}, + seenIDs map[string]struct{}, +) (bool, error) { + if err := validateRestorePeerShape(i, peer); err != nil { + return false, err + } + if _, ok := seenNodeIDs[peer.NodeID]; ok { + return false, errors.Wrapf(ErrInvalidOptions, "restore peer[%d] has duplicate node id %d", i, peer.NodeID) + } + seenNodeIDs[peer.NodeID] = struct{}{} + peerID := restorePeerIdentity(peer) + if _, ok := seenIDs[peerID]; ok { + return false, errors.Wrapf(ErrInvalidOptions, "restore peer[%d] has duplicate id %q", i, peerID) + } + seenIDs[peerID] = struct{}{} + return peer.Suffrage != etcdraftengine.SuffrageLearner, nil +} + +func validateRestorePeerShape(i int, peer etcdraftengine.Peer) error { + switch { + case peer.NodeID == 0: + return errors.Wrapf(ErrInvalidOptions, "restore peer[%d] has zero node id", i) + case stringsTrim(peer.Address) == "": + return errors.Wrapf(ErrInvalidOptions, "restore peer[%d] has empty address", i) + case peer.Suffrage != "" && + peer.Suffrage != etcdraftengine.SuffrageVoter && + peer.Suffrage != etcdraftengine.SuffrageLearner: + return errors.Wrapf(ErrInvalidOptions, "restore peer[%d] has invalid suffrage %q", i, peer.Suffrage) default: return nil } } +func restorePeerIdentity(peer etcdraftengine.Peer) string { + peerID := stringsTrim(peer.ID) + if peerID != "" { + return peerID + } + return stringsTrim(peer.Address) +} + func restoreManifest(ctx context.Context, opts RestoreOptions) (Manifest, error) { if opts.Manifest != nil { - return *opts.Manifest, nil + manifest := *opts.Manifest + if err := validateManifest(manifest); err != nil { + return Manifest{}, err + } + if err := verifyManifestSelfHash(manifest); err != nil { + return Manifest{}, err + } + return manifest, nil } return LoadManifest(ctx, opts.Store, opts.ManifestKey) } diff --git a/internal/snapshotoffload/s3_store.go b/internal/snapshotoffload/s3_store.go index 5450849cb..98ee9963f 100644 --- a/internal/snapshotoffload/s3_store.go +++ b/internal/snapshotoffload/s3_store.go @@ -135,7 +135,11 @@ func (s *S3Store) verifyS3PutObject(ctx context.Context, key string, opts PutOpt return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s remote integrity mismatch", key) } if info.SHA256 == "" { - info.SHA256 = opts.SHA256 + verified, err := s.verifyS3ObjectBytes(ctx, key, opts) + if err != nil { + return ObjectInfo{}, err + } + info = verified } return info, nil } @@ -227,15 +231,38 @@ func (s *S3Store) verifyS3ExistingObject(ctx context.Context, key string, opts P if !ok { return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s conflicted but is not visible", key) } - if info.Size == opts.Size && (info.SHA256 == "" || info.SHA256 == opts.SHA256) { - if info.SHA256 == "" { - info.SHA256 = opts.SHA256 - } + if info.Size != opts.Size { + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s already exists with different content", key) + } + if info.SHA256 == opts.SHA256 { return info, nil } + if info.SHA256 == "" { + return s.verifyS3ObjectBytes(ctx, key, opts) + } return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s already exists with different content", key) } +func (s *S3Store) verifyS3ObjectBytes(ctx context.Context, key string, opts PutOptions) (ObjectInfo, error) { + body, info, err := s.GetObject(ctx, key) + if err != nil { + return ObjectInfo{}, errors.Wrap(err, "get s3 object for integrity verification") + } + defer func() { _ = body.Close() }() + sum := sha256.New() + n, err := io.Copy(sum, contextReader{ctx: ctx, reader: body}) + if err != nil { + return ObjectInfo{}, errors.WithStack(err) + } + gotSHA := hex.EncodeToString(sum.Sum(nil)) + if n != opts.Size || gotSHA != opts.SHA256 { + return ObjectInfo{}, errors.Wrapf(ErrIntegrity, "s3 object %s already exists with different content", key) + } + info.Size = n + info.SHA256 = gotSHA + return info, nil +} + func validateStoreObjectKey(key string) (string, error) { normalized := normalizeObjectKey(key) if normalized == "" || normalized == "." || normalized == ".." || strings.HasPrefix(normalized, "../") { diff --git a/internal/snapshotoffload/s3_store_test.go b/internal/snapshotoffload/s3_store_test.go index 5e4b04a78..5feaa4c9e 100644 --- a/internal/snapshotoffload/s3_store_test.go +++ b/internal/snapshotoffload/s3_store_test.go @@ -121,6 +121,41 @@ func TestS3StoreRejectsExistingObjectMismatch(t *testing.T) { require.ErrorIs(t, err, ErrIntegrity) } +func TestS3StoreConflictHashesExistingObjectWithoutIntegrityMetadata(t *testing.T) { + ctx := context.Background() + fake := newFakeS3Client() + store := newTestS3Store(t, fake) + key := "snapshots/body.fsm" + body := []byte("same-body") + sha := hexSHA256Bytes(body) + fake.putRawObject("backup-bucket", key, body, nil, nil) + + info, err := store.PutObject(ctx, key, strings.NewReader("not-read-on-conflict"), PutOptions{ + Size: int64(len(body)), + SHA256: sha, + }) + require.NoError(t, err) + require.Equal(t, sha, info.SHA256) + require.Equal(t, 1, fake.getAttempts()) +} + +func TestS3StoreConflictRejectsExistingObjectWithoutIntegrityMetadataMismatch(t *testing.T) { + ctx := context.Background() + fake := newFakeS3Client() + store := newTestS3Store(t, fake) + key := "snapshots/body.fsm" + existing := []byte("aaaa") + candidate := []byte("bbbb") + fake.putRawObject("backup-bucket", key, existing, nil, nil) + + _, err := store.PutObject(ctx, key, strings.NewReader("not-read-on-conflict"), PutOptions{ + Size: int64(len(candidate)), + SHA256: hexSHA256Bytes(candidate), + }) + require.ErrorIs(t, err, ErrIntegrity) + require.Equal(t, 1, fake.getAttempts()) +} + func TestS3StoreRejectsParentDirectoryKeys(t *testing.T) { ctx := context.Background() store := newTestS3Store(t, newFakeS3Client()) @@ -150,6 +185,7 @@ type fakeS3Client struct { objects map[string]fakeS3Object lastPut types.ChecksumAlgorithm attempts int + gets int } type fakeS3Object struct { @@ -211,6 +247,7 @@ func (c *fakeS3Client) HeadObject(_ context.Context, input *s3.HeadObjectInput, func (c *fakeS3Client) GetObject(_ context.Context, input *s3.GetObjectInput, _ ...func(*s3.Options)) (*s3.GetObjectOutput, error) { c.mu.Lock() defer c.mu.Unlock() + c.gets++ obj, ok := c.objects[fakeS3ClientKey(input.Bucket, input.Key)] if !ok { return nil, &types.NotFound{} @@ -239,6 +276,26 @@ func (c *fakeS3Client) putAttempts() int { return c.attempts } +func (c *fakeS3Client) getAttempts() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.gets +} + +func (c *fakeS3Client) putRawObject(bucket string, key string, body []byte, metadata map[string]string, checksum *string) { + c.mu.Lock() + defer c.mu.Unlock() + clonedMetadata := make(map[string]string, len(metadata)) + for k, v := range metadata { + clonedMetadata[k] = v + } + c.objects[bucket+"/"+key] = fakeS3Object{ + body: append([]byte(nil), body...), + metadata: clonedMetadata, + checksum: checksum, + } +} + func fakeS3ClientKey(bucket *string, key *string) string { return aws.ToString(bucket) + "/" + aws.ToString(key) } diff --git a/kv/lease_warmup_test.go b/kv/lease_warmup_test.go index 206875bf6..22385c7ea 100644 --- a/kv/lease_warmup_test.go +++ b/kv/lease_warmup_test.go @@ -314,6 +314,44 @@ func TestShardedCoordinator_RecoverHLCLease_ProposesToEveryLedGroup(t *testing.T require.NotZero(t, got) } +func TestShardedCoordinator_RecoverHLCLease_SucceedsWhenAnyTargetAdvancesCeiling(t *testing.T) { + t.Parallel() + for _, tc := range []struct { + name string + failFirst bool + failSecond bool + }{ + {name: "first fails then second advances", failFirst: true}, + {name: "first advances then second fails", failSecond: true}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + clock := NewHLC() + clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) + eng1 := newShardedLeaseEngine(100) + eng2 := newShardedLeaseEngine(200) + eng1.proposeApply = applyHLCLeaseEntryToClock(t, clock) + eng2.proposeApply = applyHLCLeaseEntryToClock(t, clock) + if tc.failFirst { + eng1.proposeErr = errors.New("group 1 unavailable") + } + if tc.failSecond { + eng2.proposeErr = errors.New("group 2 unavailable") + } + coord := mustShardedLeaseCoord(t, eng1, eng2) + coord.clock = clock + + require.NoError(t, coord.RecoverHLCLease(context.Background())) + require.Equal(t, int32(1), eng1.proposeCalls.Load()) + require.Equal(t, int32(1), eng2.proposeCalls.Load()) + + got, err := clock.NextFenced() + require.NoError(t, err) + require.NotZero(t, got) + }) + } +} + func TestShardedCoordinator_RenewHLCLeases_SkipsNonLeaders(t *testing.T) { t.Parallel() eng1 := newShardedLeaseEngine(100) diff --git a/kv/sharded_coordinator.go b/kv/sharded_coordinator.go index d3a705a5b..8c6a81ad4 100644 --- a/kv/sharded_coordinator.go +++ b/kv/sharded_coordinator.go @@ -2431,13 +2431,35 @@ func (c *ShardedCoordinator) RecoverHLCLease(ctx context.Context) error { } rctx, cancel := hlcRecoveryContext(ctx) defer cancel() - ceilingMs := time.Now().UnixMilli() + hlcPhysicalWindowMs + return c.recoverHLCLeaseTargets(rctx, targets, time.Now().UnixMilli()+hlcPhysicalWindowMs) +} + +func (c *ShardedCoordinator) recoverHLCLeaseTargets(ctx context.Context, targets []hlcLeaseRecoveryTarget, ceilingMs int64) error { + var recoveryErr error for _, target := range targets { - if err := c.proposeHLCLeaseForGroup(rctx, target.group, ceilingMs); err != nil { - return errors.Wrapf(err, "recover hlc lease group %d", target.gid) + if err := c.proposeHLCLeaseForGroup(ctx, target.group, ceilingMs); err != nil { + recoveryErr = combineHLCRecoveryErrors(recoveryErr, errors.Wrapf(err, "recover hlc lease group %d", target.gid)) } } - return nil + if !hlcCeilingExpired(c.clock) { + return nil + } + return hlcRecoveryFailedError(recoveryErr) +} + +func combineHLCRecoveryErrors(first error, next error) error { + if first == nil { + return next + } + return errors.Wrap(errors.CombineErrors(first, next), "combine hlc recovery errors") +} + +func hlcRecoveryFailedError(recoveryErr error) error { + expiredErr := errors.Wrap(ErrCeilingExpired, "recover hlc lease did not advance ceiling") + if recoveryErr != nil { + return errors.Wrap(errors.CombineErrors(expiredErr, recoveryErr), "recover hlc lease failed") + } + return expiredErr } type hlcLeaseRecoveryTarget struct { diff --git a/kv/tso.go b/kv/tso.go index 292c5d852..48e800c33 100644 --- a/kv/tso.go +++ b/kv/tso.go @@ -182,7 +182,7 @@ func nextFencedWithRecovery(ctx context.Context, clock *HLC, recoverer hlcLeaseR return 0, errors.Wrap(err, label) } if recoverErr := recoverer.RecoverHLCLease(ctx); recoverErr != nil { - return 0, errors.Wrapf(err, "%s: on-demand HLC lease renewal failed: %v", label, recoverErr) + return 0, wrapHLCLeaseRecoveryFailure(err, recoverErr, label) } ts, err = clock.NextFenced() if err != nil { @@ -200,7 +200,7 @@ func nextBatchFencedWithRecovery(ctx context.Context, clock *HLC, n int, recover return 0, errors.Wrap(err, label) } if recoverErr := recoverer.RecoverHLCLease(ctx); recoverErr != nil { - return 0, errors.Wrapf(err, "%s: on-demand HLC lease renewal failed: %v", label, recoverErr) + return 0, wrapHLCLeaseRecoveryFailure(err, recoverErr, label) } base, err = clock.NextBatchFenced(n) if err != nil { @@ -209,6 +209,13 @@ func nextBatchFencedWithRecovery(ctx context.Context, clock *HLC, n int, recover return base, nil } +func wrapHLCLeaseRecoveryFailure(err, recoverErr error, label string) error { + return errors.Join( + errors.Wrapf(err, "%s: on-demand HLC lease renewal failed", label), + errors.Wrap(recoverErr, "hlc lease recovery"), + ) +} + func hlcRecoveryContext(ctx context.Context) (context.Context, context.CancelFunc) { if ctx == nil { ctx = context.Background() diff --git a/kv/tso_test.go b/kv/tso_test.go index 694f23261..46e4c31c2 100644 --- a/kv/tso_test.go +++ b/kv/tso_test.go @@ -323,6 +323,31 @@ func TestCoordinateTimestampRecoveryKeepsFailClosedWhenRenewalFails(t *testing.T require.Equal(t, int32(1), eng.proposeCalls.Load()) } +func TestCoordinateTimestampRecoveryPreservesContextErrorWhenRenewalFails(t *testing.T) { + t.Parallel() + for _, tc := range []struct { + name string + err error + }{ + {name: "canceled", err: context.Canceled}, + {name: "deadline", err: context.DeadlineExceeded}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + clock := NewHLC() + clock.SetPhysicalCeiling(time.Now().Add(-time.Millisecond).UnixMilli()) + eng := &fakeLeaseEngine{applied: 11, leaseDur: time.Hour, proposeErr: tc.err} + coord := WithKeyVizLabel(NewCoordinatorWithEngine(nil, eng, WithHLC(clock)), keyviz.LabelRedis) + + _, err := NextTimestampThrough(context.Background(), coord, "redis allocate ts") + require.ErrorIs(t, err, ErrCeilingExpired) + require.ErrorIs(t, err, tc.err) + require.Contains(t, err.Error(), "on-demand HLC lease renewal failed") + require.Equal(t, int32(1), eng.proposeCalls.Load()) + }) + } +} + func TestLocalTSOAllocatorRecoversExpiredHLCCeiling(t *testing.T) { t.Parallel() clock := NewHLC() diff --git a/proxy/backend.go b/proxy/backend.go index 9b6b27cd9..599e995dc 100644 --- a/proxy/backend.go +++ b/proxy/backend.go @@ -73,9 +73,10 @@ func DefaultBackendOptions() BackendOptions { // DefaultElasticKVBackendOptions returns defaults for proxy backends that // connect to ElasticKV's Redis adapter. Production dual-write deployments -// should run the cluster with ELASTICKV_REDIS_PER_PEER_CONNECTIONS at least as -// high as this pool size; lower the proxy pool instead for clusters that keep -// the server-side per-peer cap below the default. +// should run the cluster with ELASTICKV_REDIS_PER_PEER_CONNECTIONS above this +// pool size because PubSub and shadow PubSub use dedicated connections outside +// the go-redis command pool. Lower the proxy pool instead for clusters that +// keep the server-side per-peer cap below the default. func DefaultElasticKVBackendOptions() BackendOptions { opts := DefaultBackendOptions() opts.PoolSize = defaultElasticKVPoolSize