Skip to content

Commit 5b233fd

Browse files
authored
[refactor] Unify sidecar Config and Options (llm-d#751)
* sidecar: embed Config in Options, move port and target URL into Config Config now holds the complete runtime configuration including port, target URL, SSRF fields, and renamed TLS fields. Options embeds Config so Complete() populates it directly. NewProxy takes a single Config. Start constructs AllowlistValidator from config internally. main.go no longer builds Config manually or calls NewAllowlistValidator directly. Signed-off-by: Etai Lev Ran <elevran@gmail.com> * sidecar: update tests for new NewProxy and Start signatures Signed-off-by: Etai Lev Ran <elevran@gmail.com> * config: add String via MarshalJSON Signed-off-by: Etai Lev Ran <elevran@gmail.com> * simplify code Signed-off-by: Etai Lev Ran <elevran@gmail.com> * Port and URL are already in Options, removed from Server - Removed port and decoderURL fields from Server struct (they duplicated config.Port and config.TargetURL) - Removed the two lines in NewProxy (L186-187) that initialized them - Removed them from Clone() as well — the config copy already carries the values - Replaced all s.port / s.decoderURL / clone.port / clone.decoderURL references with s.config.Port / s.config.TargetURL across proxy.go, proxy_helpers.go, data_parallel.go, connector_sglang.go, connector_nixlv2.go, and the test file Signed-off-by: Etai Lev Ran <elevran@gmail.com> * change TargetURL to DecodeURL; make Options fields private Signed-off-by: Etai Lev Ran <elevran@gmail.com> * fix race in sidecar, reorder imports in test file Signed-off-by: Etai Lev Ran <elevran@gmail.com> --------- Signed-off-by: Etai Lev Ran <elevran@gmail.com>
1 parent 62bb5ed commit 5b233fd

16 files changed

Lines changed: 267 additions & 293 deletions

cmd/pd-sidecar/main.go

Lines changed: 5 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -16,12 +16,9 @@ limitations under the License.
1616
package main
1717

1818
import (
19-
"net/url"
20-
2119
"github.com/spf13/pflag"
2220
ctrl "sigs.k8s.io/controller-runtime"
2321
"sigs.k8s.io/controller-runtime/pkg/log"
24-
"sigs.k8s.io/controller-runtime/pkg/log/zap"
2522

2623
"github.com/llm-d/llm-d-inference-scheduler/pkg/sidecar/proxy"
2724
"github.com/llm-d/llm-d-inference-scheduler/pkg/sidecar/version"
@@ -36,7 +33,7 @@ func main() {
3633
opts.AddFlags(pflag.CommandLine)
3734
pflag.Parse()
3835

39-
logger := zap.New(zap.UseFlagOptions(&opts.LoggingOptions))
36+
logger := opts.NewLogger()
4037
log.SetLogger(logger)
4138

4239
ctx := ctrl.SetupSignalHandler()
@@ -56,7 +53,7 @@ func main() {
5653
}()
5754
}
5855

59-
// Complete options (handles migration from deprecated flags)
56+
// Complete options (handles migration from deprecated flags, populates Config)
6057
if err := opts.Complete(); err != nil {
6158
logger.Error(err, "Failed to complete configuration")
6259
return
@@ -69,56 +66,10 @@ func main() {
6966
}
7067

7168
logger.Info("Proxy starting", "Built on", version.BuildRef, "From Git SHA", version.CommitSHA)
69+
logger.Info("Proxy configuration", "config", opts.Config)
7270

73-
// Parse target URL
74-
targetURL, err := url.Parse(opts.TargetURL)
75-
if err != nil {
76-
logger.Error(err, "failed to parse targetURL")
77-
return
78-
}
79-
80-
config := proxy.Config{
81-
KVConnector: opts.KVConnector,
82-
ECConnector: opts.ECConnector,
83-
PrefillerUseTLS: opts.UseTLSForPrefiller,
84-
EncoderUseTLS: opts.UseTLSForEncoder,
85-
PrefillerInsecureSkipVerify: opts.InsecureSkipVerifyForPrefiller,
86-
EncoderInsecureSkipVerify: opts.InsecureSkipVerifyForEncoder,
87-
DecoderInsecureSkipVerify: opts.InsecureSkipVerifyForDecoder,
88-
DataParallelSize: opts.DataParallelSize,
89-
EnablePrefillerSampling: opts.EnablePrefillerSampling,
90-
SecureServing: opts.SecureProxy,
91-
CertPath: opts.CertPath,
92-
}
93-
94-
logger.Info("Proxy configuration",
95-
"port", opts.Port,
96-
"targetURL", opts.TargetURL,
97-
"kvConnector", config.KVConnector,
98-
"ecConnector", config.ECConnector,
99-
"dataParallelSize", config.DataParallelSize,
100-
"prefillerUseTLS", config.PrefillerUseTLS,
101-
"prefillerInsecureSkipVerify", config.PrefillerInsecureSkipVerify,
102-
"decoderInsecureSkipVerify", config.DecoderInsecureSkipVerify,
103-
"enablePrefillerSampling", config.EnablePrefillerSampling,
104-
"secureServing", config.SecureServing,
105-
"certPath", config.CertPath,
106-
"enableSSRFProtection", opts.EnableSSRFProtection,
107-
"inferencePoolNamespace", opts.InferencePoolNamespace,
108-
"inferencePoolName", opts.InferencePoolName,
109-
"poolGroup", opts.PoolGroup,
110-
)
111-
112-
// Create SSRF protection validator
113-
validator, err := proxy.NewAllowlistValidator(opts.EnableSSRFProtection, opts.PoolGroup, opts.InferencePoolNamespace, opts.InferencePoolName)
114-
if err != nil {
115-
logger.Error(err, "failed to create SSRF protection validator")
116-
return
117-
}
118-
119-
proxyServer := proxy.NewProxy(opts.Port, targetURL, config)
120-
121-
if err := proxyServer.Start(ctx, validator); err != nil {
71+
proxyServer := proxy.NewProxy(opts.Config)
72+
if err := proxyServer.Start(ctx); err != nil {
12273
logger.Error(err, "failed to start proxy server")
12374
}
12475
}

pkg/plugins/scorer/precise_prefix_cache_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,8 +17,8 @@ import (
1717
"github.com/llm-d/llm-d-kv-cache/pkg/tokenization"
1818
"github.com/stretchr/testify/assert"
1919
"github.com/stretchr/testify/require"
20-
"k8s.io/apimachinery/pkg/util/sets"
2120
k8stypes "k8s.io/apimachinery/pkg/types"
21+
"k8s.io/apimachinery/pkg/util/sets"
2222
fwkdl "sigs.k8s.io/gateway-api-inference-extension/pkg/epp/framework/interface/datalayer"
2323
"sigs.k8s.io/gateway-api-inference-extension/pkg/epp/framework/interface/plugin"
2424
"sigs.k8s.io/gateway-api-inference-extension/pkg/epp/framework/interface/scheduling"

pkg/sidecar/proxy/chat_completions_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,7 @@ func TestServer_chatCompletionsHandler(t *testing.T) {
113113

114114
for i := 0; i < maxAttempts; i++ {
115115
t.Run(fmt.Sprintf("%s_%d", tt.name, i), func(t *testing.T) {
116-
s := NewProxy("8000", nil, Config{EnablePrefillerSampling: tt.sampling})
116+
s := NewProxy(Config{Port: "8000", EnablePrefillerSampling: tt.sampling})
117117
s.allowlistValidator = &AllowlistValidator{}
118118
// return a predictable sequence of values
119119
s.prefillSamplerFn = func(n int) int { return i % n }

pkg/sidecar/proxy/connector_epd_shared_storage_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -271,7 +271,7 @@ func TestFanoutEncoderPrimerDeduplication(t *testing.T) {
271271

272272
encoderURL, err := url.Parse(encoderBackend.URL)
273273
assert.NoError(t, err)
274-
srv := NewProxy("0", encoderURL, Config{})
274+
srv := NewProxy(Config{Port: "0", DecoderURL: encoderURL})
275275
srv.logger = log.Log
276276

277277
encoderHostPort := encoderURL.Host

pkg/sidecar/proxy/connector_nixlv2.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -229,8 +229,8 @@ func (s *Server) runNIXLProtocolV2(w http.ResponseWriter, r *http.Request, prefi
229229
decodeSpan.SetAttributes(attribute.Bool("llm_d.pd_proxy.decode.data_parallel", dataParallelUsed))
230230

231231
if !dataParallelUsed {
232-
s.logger.V(4).Info("sending request to decoder", "to", s.decoderURL.Host)
233-
decodeSpan.SetAttributes(attribute.String("llm_d.pd_proxy.decode.target", s.decoderURL.Host))
232+
s.logger.V(4).Info("sending request to decoder", "to", s.config.DecoderURL.Host)
233+
decodeSpan.SetAttributes(attribute.String("llm_d.pd_proxy.decode.target", s.config.DecoderURL.Host))
234234
s.decoderProxy.ServeHTTP(w, dreq)
235235
}
236236

pkg/sidecar/proxy/connector_nixlv2_test.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,8 +39,8 @@ var _ = Describe("NIXL Connector (v2)", func() {
3939
go func() {
4040
defer GinkgoRecover()
4141

42-
validator := &AllowlistValidator{enabled: false}
43-
err := testInfo.proxy.Start(testInfo.ctx, validator)
42+
testInfo.proxy.allowlistValidator = &AllowlistValidator{enabled: false}
43+
err := testInfo.proxy.Start(testInfo.ctx)
4444
Expect(err).ToNot(HaveOccurred())
4545

4646
testInfo.stoppedCh <- struct{}{}

pkg/sidecar/proxy/connector_sglang.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -152,7 +152,7 @@ func (s *Server) sendSGLangConcurrentRequests(w http.ResponseWriter, r *http.Req
152152
decodeDuration := time.Since(decodeStart)
153153
decodeSpan.SetAttributes(
154154
attribute.Float64("llm_d.pd_proxy.decode.duration_ms", float64(decodeDuration.Milliseconds())),
155-
attribute.String("llm_d.pd_proxy.decode.target", s.decoderURL.Host),
155+
attribute.String("llm_d.pd_proxy.decode.target", s.config.DecoderURL.Host),
156156
)
157157

158158
// Calculate end-to-end P/D timing metrics for concurrent P/D.

pkg/sidecar/proxy/connector_sglang_test.go

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -44,8 +44,8 @@ var _ = Describe("SGLang Connector", func() {
4444
go func() {
4545
defer GinkgoRecover()
4646

47-
validator := &AllowlistValidator{enabled: false}
48-
err := testInfo.proxy.Start(testInfo.ctx, validator)
47+
testInfo.proxy.allowlistValidator = &AllowlistValidator{enabled: false}
48+
err := testInfo.proxy.Start(testInfo.ctx)
4949
Expect(err).ToNot(HaveOccurred())
5050

5151
testInfo.stoppedCh <- struct{}{}
@@ -136,14 +136,16 @@ var _ = Describe("SGLang Connector", func() {
136136

137137
// Re-initialize proxy to fetch the new mock addresses
138138
cfg := Config{
139+
Port: "0",
140+
DecoderURL: testInfo.decodeURL,
139141
KVConnector: KVConnectorSGLang,
140142
}
141-
testInfo.proxy = NewProxy("0", testInfo.decodeURL, cfg)
143+
testInfo.proxy = NewProxy(cfg)
142144

143145
go func() {
144146
defer GinkgoRecover()
145-
validator := &AllowlistValidator{enabled: false}
146-
err := testInfo.proxy.Start(testInfo.ctx, validator)
147+
testInfo.proxy.allowlistValidator = &AllowlistValidator{enabled: false}
148+
err := testInfo.proxy.Start(testInfo.ctx)
147149
Expect(err).ToNot(HaveOccurred())
148150
testInfo.stoppedCh <- struct{}{}
149151
}()

pkg/sidecar/proxy/connector_test.go

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -57,8 +57,8 @@ var _ = Describe("Common Connector tests", func() {
5757
go func() {
5858
defer GinkgoRecover()
5959

60-
validator := &AllowlistValidator{enabled: false}
61-
err := testInfo.proxy.Start(testInfo.ctx, validator)
60+
testInfo.proxy.allowlistValidator = &AllowlistValidator{enabled: false}
61+
err := testInfo.proxy.Start(testInfo.ctx)
6262
Expect(err).ToNot(HaveOccurred())
6363

6464
testInfo.stoppedCh <- struct{}{}
@@ -118,8 +118,8 @@ var _ = Describe("Common Connector tests", func() {
118118
go func() {
119119
defer GinkgoRecover()
120120

121-
validator := &AllowlistValidator{enabled: false}
122-
err := testInfo.proxy.Start(testInfo.ctx, validator)
121+
testInfo.proxy.allowlistValidator = &AllowlistValidator{enabled: false}
122+
err := testInfo.proxy.Start(testInfo.ctx)
123123
Expect(err).ToNot(HaveOccurred())
124124

125125
testInfo.stoppedCh <- struct{}{}
@@ -200,8 +200,8 @@ func sidecarConnectionTestSetup(connector string) *sidecarTestInfo {
200200
url, err := url.Parse(testInfo.decodeBackend.URL)
201201
Expect(err).ToNot(HaveOccurred())
202202
testInfo.decodeURL = url
203-
cfg := Config{KVConnector: connector}
204-
testInfo.proxy = NewProxy("0", testInfo.decodeURL, cfg) // port 0 to automatically choose one that's available.
203+
cfg := Config{Port: "0", DecoderURL: testInfo.decodeURL, KVConnector: connector}
204+
testInfo.proxy = NewProxy(cfg)
205205

206206
return &testInfo
207207
}

pkg/sidecar/proxy/data_parallel.go

Lines changed: 17 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -37,17 +37,16 @@ func (s *Server) dataParallelHandler(w http.ResponseWriter, r *http.Request) boo
3737

3838
func (s *Server) startDataParallel(ctx context.Context, grp *errgroup.Group) error {
3939
podIP := os.Getenv("POD_IP")
40-
basePort, err := strconv.Atoi(s.port)
40+
basePort, err := strconv.Atoi(s.config.Port)
4141
if err != nil {
4242
return err
4343
}
44-
baseDecoderPort, err := strconv.Atoi(s.decoderURL.Port())
44+
baseDecoderPort, err := strconv.Atoi(s.config.DecoderURL.Port())
4545
if err != nil {
4646
return err
4747
}
48-
decoderScheme := s.decoderURL.Scheme // capture before goroutines launch
49-
50-
s.dataParallelProxies[net.JoinHostPort(podIP, s.port)] = s.decoderProxy
48+
decoderScheme := s.config.DecoderURL.Scheme // capture before goroutines launch
49+
s.dataParallelProxies[net.JoinHostPort(podIP, s.config.Port)] = s.decoderProxy
5150

5251
// Fill in map of proxies, thus avoiding locks
5352
for idx := range s.config.DataParallelSize - 1 {
@@ -58,24 +57,25 @@ func (s *Server) startDataParallel(ctx context.Context, grp *errgroup.Group) err
5857
if err != nil {
5958
return err
6059
}
61-
handler := s.createDecoderProxyHandler(decoderURL, s.config.DecoderInsecureSkipVerify)
60+
handler := s.createDecoderProxyHandler(decoderURL, s.config.InsecureSkipVerifyForDecoder)
6261
s.dataParallelProxies[hostPort] = handler
6362
}
6463

6564
for idx := range s.config.DataParallelSize - 1 {
66-
grp.Go(func() error {
67-
rankPort := strconv.Itoa(basePort + idx + 1)
68-
decoderPort := strconv.Itoa(baseDecoderPort + idx + 1)
69-
decoderURL, err := url.Parse(decoderScheme + "://localhost:" + decoderPort)
70-
if err != nil {
71-
return err
72-
}
65+
rankPort := strconv.Itoa(basePort + idx + 1)
66+
decoderPort := strconv.Itoa(baseDecoderPort + idx + 1)
67+
decoderURL, err := url.Parse(decoderScheme + "://localhost:" + decoderPort)
68+
if err != nil {
69+
return err
70+
}
7371

74-
clone := s.Clone()
72+
clone := s.Clone()
73+
clone.config.Port = rankPort
74+
clone.config.DecoderURL = decoderURL
75+
clone.forwardDataParallel = false
76+
77+
grp.Go(func() error {
7578
clone.logger = log.FromContext(ctx).WithName("proxy server on port " + rankPort)
76-
clone.port = rankPort
77-
clone.decoderURL = decoderURL
78-
clone.forwardDataParallel = false
7979
// Configure handlers
8080
clone.handler = clone.createRoutes()
8181
clone.setKVConnector()

0 commit comments

Comments
 (0)