diff --git a/internal/controller/shardsplitjob_copy.go b/internal/controller/shardsplitjob_copy.go index 44e0d42b..bc70f267 100644 --- a/internal/controller/shardsplitjob_copy.go +++ b/internal/controller/shardsplitjob_copy.go @@ -69,12 +69,15 @@ func sourceShardPodDNS(cluster, ns, shardID string) (string, error) { if !strings.HasPrefix(shardID, "shard-") { return "", fmt.Errorf("source shard %q is not ordinal (want shard-N)", shardID) } - ord, err := strconv.Atoi(strings.TrimPrefix(shardID, "shard-")) + + // bitSize 32 — int32 를 넘는 ordinal 은 잘려 다른 shard pod 를 가리키므로 거부한다. + v, err := strconv.ParseInt(strings.TrimPrefix(shardID, "shard-"), 10, 32) if err != nil { return "", fmt.Errorf("source shard %q: %w", shardID, err) } - pod := ShardStatefulSetName(cluster, int32(ord)) + "-0" - return fmt.Sprintf("%s.%s.%s.svc.cluster.local", pod, ShardServiceName(cluster, int32(ord)), ns), nil + ord := int32(v) + pod := ShardStatefulSetName(cluster, ord) + "-0" + return fmt.Sprintf("%s.%s.%s.svc.cluster.local", pod, ShardServiceName(cluster, ord), ns), nil } // targetShardPodDNS 는 resharding target shard 의 primary pod(-0) 안정 DNS 를 만든다. diff --git a/internal/controller/shardsplitjob_copy_bounds_test.go b/internal/controller/shardsplitjob_copy_bounds_test.go new file mode 100644 index 00000000..28431bf9 --- /dev/null +++ b/internal/controller/shardsplitjob_copy_bounds_test.go @@ -0,0 +1,45 @@ +/* +Copyright 2026 keiailab. + +Licensed under the MIT License. See the LICENSE file for details. +*/ + +package controller + +import ( + "strings" + "testing" +) + +// TestSourceShardPodDNS_OrdinalBounds 는 shard ordinal 이 int32 범위를 벗어나면 +// 잘라서 엉뚱한 pod 를 가리키지 않고 에러를 내는지 검증한다. +func TestSourceShardPodDNS_OrdinalBounds(t *testing.T) { + tests := []struct { + shard string + wantErr bool + wantPod string + }{ + {shard: "shard-0", wantPod: "c-shard-0-0."}, + {shard: "shard-2147483647", wantPod: "c-shard-2147483647-0."}, + {shard: "shard-2147483648", wantErr: true}, + {shard: "shard-4294967296", wantErr: true}, + {shard: "shard-x", wantErr: true}, + {shard: "s-1", wantErr: true}, + } + for _, tt := range tests { + got, err := sourceShardPodDNS("c", "ns", tt.shard) + if tt.wantErr { + if err == nil { + t.Errorf("%s: want error, got %q", tt.shard, got) + } + continue + } + if err != nil { + t.Errorf("%s: unexpected error: %v", tt.shard, err) + continue + } + if !strings.HasPrefix(got, tt.wantPod) { + t.Errorf("%s: got %q, want prefix %q", tt.shard, got, tt.wantPod) + } + } +} diff --git a/internal/router/vindex.go b/internal/router/vindex.go index 4223b764..40a63e38 100644 --- a/internal/router/vindex.go +++ b/internal/router/vindex.go @@ -108,11 +108,12 @@ func hashKey(fn v1alpha1.VindexHashFunction, key string) (uint32, error) { } // parseHashBound 는 hex 문자열 ("0x..." 또는 "ffffffff") 또는 10진수를 uint32 로 해석한다. +// uint32 를 넘는 값은 잘려 다른 범위가 되므로 bitSize 32 로 거부한다. func parseHashBound(s string) (uint32, error) { - if v, err := strconv.ParseUint(s, 0, 64); err == nil { + if v, err := strconv.ParseUint(s, 0, 32); err == nil { return uint32(v), nil } - if v, err := strconv.ParseUint(s, 16, 64); err == nil { + if v, err := strconv.ParseUint(s, 16, 32); err == nil { return uint32(v), nil } return 0, fmt.Errorf("invalid hash bound %q (expected hex or decimal)", s) diff --git a/internal/router/vindex_bound_test.go b/internal/router/vindex_bound_test.go new file mode 100644 index 00000000..1cc1675d --- /dev/null +++ b/internal/router/vindex_bound_test.go @@ -0,0 +1,46 @@ +/* +Copyright 2026 keiailab. + +Licensed under the MIT License. See the LICENSE file for details. +*/ + +package router + +import "testing" + +// TestParseHashBound_Bounds 는 uint32 를 넘는 bound 를 잘라 엉뚱한 범위로 +// 해석하지 않고 에러를 내는지 검증한다 (예: 0x100000000 → 0 이 되면 안 된다). +func TestParseHashBound_Bounds(t *testing.T) { + tests := []struct { + in string + want uint32 + wantErr bool + }{ + {in: "0", want: 0}, + {in: "0x00000000", want: 0}, + {in: "0xffffffff", want: 0xffffffff}, + {in: "ffffffff", want: 0xffffffff}, + {in: "4294967295", want: 0xffffffff}, + {in: "4294967296", wantErr: true}, + {in: "0x100000000", wantErr: true}, + {in: "100000000f", wantErr: true}, + {in: "-1", wantErr: true}, + {in: "zz", wantErr: true}, + } + for _, tt := range tests { + got, err := parseHashBound(tt.in) + if tt.wantErr { + if err == nil { + t.Errorf("%q: want error, got %#x", tt.in, got) + } + continue + } + if err != nil { + t.Errorf("%q: unexpected error: %v", tt.in, err) + continue + } + if got != tt.want { + t.Errorf("%q: got %#x, want %#x", tt.in, got, tt.want) + } + } +}