diff --git a/balancer/autosharding/autosharding.go b/balancer/autosharding/autosharding.go new file mode 100644 index 000000000000..f26106853b34 --- /dev/null +++ b/balancer/autosharding/autosharding.go @@ -0,0 +1,79 @@ +/* + * + * Copyright 2026 gRPC authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + * + */ + +// Package autosharding implements the autosharding load balancing policy. +package autosharding + +import ( + "encoding/json" + "fmt" + "time" + + "google.golang.org/grpc/balancer" + iserviceconfig "google.golang.org/grpc/internal/serviceconfig" + "google.golang.org/grpc/serviceconfig" +) + +// Name is the name of the autosharding balancer. +const Name = "autosharding_experimental" + +func init() { + balancer.Register(bb{}) +} + +// lbConfig is the balancer config for the autosharding balancer. +type lbConfig struct { + serviceconfig.LoadBalancingConfig `json:"-"` + + ChannelFactoryKey string `json:"channelFactoryKey,omitempty"` + AutoShardingTarget string `json:"autoshardingTarget,omitempty"` + KeyHeaderName string `json:"keyHeaderName,omitempty"` + EnableFallback bool `json:"enableFallback,omitempty"` + InitialAssignmentTimeout iserviceconfig.Duration `json:"initialAssignmentTimeout,omitempty"` +} + +type bb struct{} + +func (bb) Name() string { + return Name +} + +func (bb) ParseConfig(s json.RawMessage) (serviceconfig.LoadBalancingConfig, error) { + lbConfig := &lbConfig{InitialAssignmentTimeout: iserviceconfig.Duration(60 * time.Second)} + if err := json.Unmarshal(s, lbConfig); err != nil { + return nil, fmt.Errorf("autosharding: unable to unmarshal LBConfig: %v", err) + } + if lbConfig.ChannelFactoryKey == "" { + return nil, fmt.Errorf("autosharding: channelFactoryKey field is required") + } + if lbConfig.AutoShardingTarget == "" { + return nil, fmt.Errorf("autosharding: autoshardingTarget field is required") + } + if lbConfig.KeyHeaderName == "" { + return nil, fmt.Errorf("autosharding: keyHeaderName field is required") + } + return lbConfig, nil +} + +func (bb) Build(balancer.ClientConn, balancer.BuildOptions) balancer.Balancer { + return &autoshardingBalancer{} +} + +type autoshardingBalancer struct { + balancer.Balancer +} diff --git a/balancer/autosharding/autosharding_test.go b/balancer/autosharding/autosharding_test.go new file mode 100644 index 000000000000..2a4c1e46f5de --- /dev/null +++ b/balancer/autosharding/autosharding_test.go @@ -0,0 +1,144 @@ +/* + * + * Copyright 2026 gRPC authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + * + */ + +package autosharding + +import ( + "encoding/json" + "testing" + "time" + + "github.com/google/go-cmp/cmp" + "google.golang.org/grpc/internal/grpctest" + iserviceconfig "google.golang.org/grpc/internal/serviceconfig" + "google.golang.org/grpc/serviceconfig" +) + +type s struct { + grpctest.Tester +} + +func Test(t *testing.T) { + grpctest.RunSubTests(t, s{}) +} + +func (s) TestParseConfig_Success(t *testing.T) { + parser := bb{} + tests := []struct { + name string + input string + wantCfg serviceconfig.LoadBalancingConfig + }{ + { + name: "all-fields", + input: `{ + "channelFactoryKey": "factory-key", + "autoshardingTarget": "target", + "keyHeaderName": "header", + "enableFallback": true, + "initialAssignmentTimeout": "30s" + }`, + wantCfg: &lbConfig{ + ChannelFactoryKey: "factory-key", + AutoShardingTarget: "target", + KeyHeaderName: "header", + EnableFallback: true, + InitialAssignmentTimeout: iserviceconfig.Duration(30 * time.Second), + }, + }, + { + name: "default-timeout", + input: `{ + "channelFactoryKey": "factory-key", + "autoshardingTarget": "target", + "keyHeaderName": "header" + }`, + wantCfg: &lbConfig{ + ChannelFactoryKey: "factory-key", + AutoShardingTarget: "target", + KeyHeaderName: "header", + InitialAssignmentTimeout: iserviceconfig.Duration(60 * time.Second), + }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + gotCfg, err := parser.ParseConfig(json.RawMessage(test.input)) + if err != nil { + t.Fatalf("ParseConfig() error = %v, want nil", err) + } + if diff := cmp.Diff(test.wantCfg, gotCfg); diff != "" { + t.Errorf("ParseConfig() config diff (-want +got):\n%s", diff) + } + }) + } +} + +func (s) TestParseConfig_Failure(t *testing.T) { + parser := bb{} + tests := []struct { + name string + input string + }{ + { + name: "invalid-json", + input: "{{invalidjson{{", + }, + { + name: "invalid-duration", + input: `{ + "channelFactoryKey": "factory-key", + "autoshardingTarget": "target", + "keyHeaderName": "header", + "initialAssignmentTimeout": "invalid" + }`, + }, + { + name: "missing-channel-factory-key", + input: `{ + "autoshardingTarget": "target", + "keyHeaderName": "header" + }`, + }, + { + name: "missing-autosharding-target", + input: `{ + "channelFactoryKey": "factory-key", + "keyHeaderName": "header" + }`, + }, + { + name: "missing-key-header-name", + input: `{ + "channelFactoryKey": "factory-key", + "autoshardingTarget": "target" + }`, + }, + { + name: "empty-config", + input: `{}`, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if _, err := parser.ParseConfig(json.RawMessage(test.input)); err == nil { + t.Fatalf("ParseConfig() succeeded, want error") + } + }) + } +}