Skip to content

Commit 3ee3b4d

Browse files
xenOs76cursoragent
andcommitted
fix: propagate context through MCP requests and certinfo TLS paths
Thread cancellation from MCP tool timeouts into HandleRequests, http.NewRequestWithContext, and certinfo DialContext handshakes. Remove parallel os.Args mutation in TestIsMCPCommand subtests. Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent aa22215 commit 3ee3b4d

13 files changed

Lines changed: 146 additions & 71 deletions

‎internal/certinfo/certinfo.go‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package certinfo
22

33
import (
4+
"context"
45
"crypto"
56
"crypto/x509"
67
"fmt"
@@ -180,7 +181,7 @@ func (c *Config) SetPrivateKeyFromFile(
180181
}
181182

182183
// SetTLSEndpoint parses a host:port string and fetches the remote certificates from that endpoint.
183-
func (c *Config) SetTLSEndpoint(hostport string) error {
184+
func (c *Config) SetTLSEndpoint(ctx context.Context, hostport string) error {
184185
if hostport != emptyString {
185186
eHost, ePort, err := net.SplitHostPort(hostport)
186187
if err != nil {
@@ -191,7 +192,7 @@ func (c *Config) SetTLSEndpoint(hostport string) error {
191192
c.TLSEndpointHost = eHost
192193
c.TLSEndpointPort = ePort
193194

194-
err = c.GetRemoteCerts()
195+
err = c.GetRemoteCerts(ctx)
195196
if err != nil {
196197
return fmt.Errorf("unable to get endpoint certificates: %w", err)
197198
}

‎internal/certinfo/certinfo_handlers.go‎

Lines changed: 56 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ package certinfo
66

77
import (
88
"cmp"
9+
"context"
910
"crypto/sha256"
1011
"crypto/tls"
1112
"crypto/x509"
@@ -29,7 +30,7 @@ import (
2930
// to the provided writer in a human-readable format.
3031
//
3132
//nolint:revive
32-
func (c *Config) PrintData(w io.Writer) error {
33+
func (c *Config) PrintData(ctx context.Context, w io.Writer) error {
3334
ks := style.ItemKey.PaddingBottom(0).PaddingTop(1).PaddingLeft(1)
3435
sl := style.CertKeyP4.Bold(true)
3536
sv := style.CertValue.Bold(false)
@@ -49,7 +50,7 @@ func (c *Config) PrintData(w io.Writer) error {
4950
}
5051

5152
if c.TLSInfoRequested {
52-
_ = c.ProbeTLSInfo()
53+
_ = c.ProbeTLSInfo(ctx)
5354
c.printTLSInfo(w, ks, sl, sv)
5455
}
5556

@@ -174,30 +175,47 @@ func (c *Config) printCACerts(w io.Writer, ks, sl, sv lipgloss.Style) error {
174175
return nil
175176
}
176177

178+
// dialTLS connects to serverAddr and completes a TLS handshake using ctx for cancellation.
179+
func dialTLS(ctx context.Context, serverAddr string, tlsConfig *tls.Config) (*tls.Conn, error) {
180+
dialer := &net.Dialer{Timeout: TLSTimeout}
181+
182+
rawConn, err := dialer.DialContext(ctx, "tcp", serverAddr)
183+
if err != nil {
184+
return nil, err
185+
}
186+
187+
conn := tls.Client(rawConn, tlsConfig)
188+
189+
if err = conn.HandshakeContext(ctx); err != nil {
190+
_ = rawConn.Close()
191+
192+
return nil, err
193+
}
194+
195+
return conn, nil
196+
}
197+
177198
// GetRemoteCerts establishes a TLS connection to the configured endpoint and retrieves
178199
// the peer certificate chain. It also performs certificate verification unless TLSInsecure is true.
179-
func (c *Config) GetRemoteCerts() error {
200+
func (c *Config) GetRemoteCerts(ctx context.Context) error {
180201
tlsConfig := &tls.Config{
181202
RootCAs: c.CACertsPool,
182203
InsecureSkipVerify: c.TLSInsecure,
183204
}
184205

185-
if c.TLSServerName != emptyString {
206+
verifyName := c.TLSServerName
207+
switch {
208+
case c.TLSServerName != emptyString:
186209
tlsConfig.ServerName = c.TLSServerName
210+
case c.TLSEndpointHost != emptyString:
211+
tlsConfig.ServerName = c.TLSEndpointHost
212+
verifyName = c.TLSEndpointHost
213+
default:
187214
}
188215

189216
serverAddr := net.JoinHostPort(c.TLSEndpointHost, c.TLSEndpointPort)
190217

191-
dialer := &net.Dialer{
192-
Timeout: TLSTimeout,
193-
}
194-
195-
conn, err := tls.DialWithDialer(
196-
dialer,
197-
"tcp",
198-
serverAddr,
199-
tlsConfig,
200-
)
218+
conn, err := dialTLS(ctx, serverAddr, tlsConfig)
201219
if err != nil {
202220
return fmt.Errorf("TLS handshake failed: %w", err)
203221
}
@@ -214,7 +232,7 @@ func (c *Config) GetRemoteCerts() error {
214232
}
215233

216234
opts := x509.VerifyOptions{
217-
DNSName: c.TLSServerName,
235+
DNSName: verifyName,
218236
Roots: c.CACertsPool,
219237
Intermediates: x509.NewCertPool(),
220238
}
@@ -425,7 +443,7 @@ func tlsVersionToString(version uint16) string {
425443
}
426444

427445
// probeProtocol tests whether the TLS endpoint supports a specific TLS protocol version.
428-
func (c *Config) probeProtocol(version uint16) bool {
446+
func (c *Config) probeProtocol(ctx context.Context, version uint16) bool {
429447
tlsConfig := &tls.Config{
430448
MinVersion: version,
431449
MaxVersion: version,
@@ -434,17 +452,15 @@ func (c *Config) probeProtocol(version uint16) bool {
434452

435453
if c.TLSServerName != emptyString {
436454
tlsConfig.ServerName = c.TLSServerName
455+
} else if c.TLSEndpointHost != emptyString {
456+
tlsConfig.ServerName = c.TLSEndpointHost
437457
}
438458

439459
serverAddr := net.JoinHostPort(c.TLSEndpointHost, c.TLSEndpointPort)
440460

441-
dialer := &net.Dialer{
442-
Timeout: TLSTimeout,
443-
}
444-
445-
conn, err := tls.DialWithDialer(dialer, "tcp", serverAddr, tlsConfig)
461+
conn, err := dialTLS(ctx, serverAddr, tlsConfig)
446462
if err == nil {
447-
conn.Close()
463+
_ = conn.Close()
448464

449465
return true
450466
}
@@ -453,7 +469,7 @@ func (c *Config) probeProtocol(version uint16) bool {
453469
}
454470

455471
// probeCipher tests whether a specific TLS 1.0-1.2 cipher suite is supported.
456-
func (c *Config) probeCipher(suite *tls.CipherSuite) (bool, string) {
472+
func (c *Config) probeCipher(ctx context.Context, suite *tls.CipherSuite) (bool, string) {
457473
tlsConfig := &tls.Config{
458474
MinVersion: tls.VersionTLS10,
459475
MaxVersion: tls.VersionTLS12,
@@ -463,19 +479,17 @@ func (c *Config) probeCipher(suite *tls.CipherSuite) (bool, string) {
463479

464480
if c.TLSServerName != emptyString {
465481
tlsConfig.ServerName = c.TLSServerName
482+
} else if c.TLSEndpointHost != emptyString {
483+
tlsConfig.ServerName = c.TLSEndpointHost
466484
}
467485

468486
serverAddr := net.JoinHostPort(c.TLSEndpointHost, c.TLSEndpointPort)
469487

470-
dialer := &net.Dialer{
471-
Timeout: TLSTimeout,
472-
}
473-
474-
conn, err := tls.DialWithDialer(dialer, "tcp", serverAddr, tlsConfig)
488+
conn, err := dialTLS(ctx, serverAddr, tlsConfig)
475489
if err == nil {
476490
state := conn.ConnectionState()
477491

478-
conn.Close()
492+
_ = conn.Close()
479493

480494
return true, tlsVersionToString(state.Version)
481495
}
@@ -484,7 +498,7 @@ func (c *Config) probeCipher(suite *tls.CipherSuite) (bool, string) {
484498
}
485499

486500
// ProbeTLSInfo concurrently scans the endpoint for supported TLS versions and cipher suites.
487-
func (c *Config) ProbeTLSInfo() error {
501+
func (c *Config) ProbeTLSInfo(ctx context.Context) error {
488502
if c.TLSEndpoint == emptyString {
489503
return nil
490504
}
@@ -495,23 +509,27 @@ func (c *Config) ProbeTLSInfo() error {
495509
versions := []uint16{tls.VersionTLS10, tls.VersionTLS11, tls.VersionTLS12, tls.VersionTLS13}
496510

497511
for _, v := range versions {
498-
supported := c.probeProtocol(v)
512+
if err := ctx.Err(); err != nil {
513+
return err
514+
}
515+
516+
supported := c.probeProtocol(ctx, v)
499517

500518
c.ProbedProtocols[tlsVersionToString(v)] = supported
501519
}
502520

503521
// 2. Probe ciphers concurrently
504522
suites := append(tls.CipherSuites(), tls.InsecureCipherSuites()...)
505523

506-
c.ProbedCiphers = c.probeCiphersConcurrently(suites)
524+
c.ProbedCiphers = c.probeCiphersConcurrently(ctx, suites)
507525

508526
return nil
509527
}
510528

511529
// probeCiphersConcurrently manages the worker pool to concurrently scan cipher suites.
512530
//
513531
//nolint:gocognit,revive,wsl
514-
func (c *Config) probeCiphersConcurrently(suites []*tls.CipherSuite) []ProbedCipher {
532+
func (c *Config) probeCiphersConcurrently(ctx context.Context, suites []*tls.CipherSuite) []ProbedCipher {
515533
type job struct {
516534
suite *tls.CipherSuite
517535
}
@@ -539,6 +557,10 @@ func (c *Config) probeCiphersConcurrently(suites []*tls.CipherSuite) []ProbedCip
539557
defer wg.Done()
540558

541559
for j := range jobs {
560+
if err := ctx.Err(); err != nil {
561+
return
562+
}
563+
542564
suite := j.suite
543565
isTLS13 := false
544566

@@ -559,7 +581,7 @@ func (c *Config) probeCiphersConcurrently(suites []*tls.CipherSuite) []ProbedCip
559581
supported = c.ProbedProtocols["TLS 1.3"]
560582
protoName = "TLS 1.3"
561583
} else {
562-
ok, name := c.probeCipher(suite)
584+
ok, name := c.probeCipher(ctx, suite)
563585

564586
supported = ok
565587
protoName = name

‎internal/certinfo/certinfo_handlers_test.go‎

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package certinfo
22

33
import (
44
"bytes"
5+
"context"
56
"crypto/x509"
67
"crypto/x509/pkix"
78
"math/big"
@@ -145,10 +146,10 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) {
145146

146147
cc.SetTLSServerName(tt.srvCfg.serverName)
147148
cc.SetCaPoolFromFile(tt.caCertFile, inputReader)
148-
cc.SetTLSEndpoint(tt.srvCfg.serverAddr)
149+
cc.SetTLSEndpoint(context.Background(), tt.srvCfg.serverAddr)
149150
cc.SetTLSInsecure(tt.insecure)
150151

151-
err = cc.GetRemoteCerts()
152+
err = cc.GetRemoteCerts(context.Background())
152153
if !tt.expectError {
153154
require.NoError(t, err, "check error not expected")
154155
require.Equal(t, tt.srvCfg.serverName, cc.TLSServerName, "check TLSServerName")
@@ -491,7 +492,7 @@ func TestCertinfo_PrintData(t *testing.T) {
491492
})
492493
cc.CertsBundleFilePath = "dummy"
493494

494-
errPrint := cc.PrintData(&buffer)
495+
errPrint := cc.PrintData(context.Background(), &buffer)
495496
require.Error(t, errPrint)
496497
require.ErrorContains(t, errPrint, "unable to check if private key matches local certificate")
497498
})
@@ -508,7 +509,7 @@ func TestCertinfo_PrintData(t *testing.T) {
508509
cc.TLSEndpointHost = "localhost"
509510
cc.TLSEndpointPort = "443"
510511

511-
errPrint := cc.PrintData(&buffer)
512+
errPrint := cc.PrintData(context.Background(), &buffer)
512513
require.Error(t, errPrint)
513514
require.ErrorContains(t, errPrint, "unable to check if private key matches remote TLS Endpoint certificate")
514515
})
@@ -520,7 +521,7 @@ func TestCertinfo_PrintData(t *testing.T) {
520521

521522
cc.CACertsFilePath = "non_existent_file.pem"
522523

523-
errPrint := cc.PrintData(&buffer)
524+
errPrint := cc.PrintData(context.Background(), &buffer)
524525
require.Error(t, errPrint)
525526
require.ErrorContains(t, errPrint, "unable for read Root certificates")
526527
})
@@ -559,15 +560,15 @@ func runPrintDataSubtest(t *testing.T, tt printDataTestCase) {
559560
cc.SetTLSServerName(tt.tlsServerName)
560561
cc.SetTLSInsecure(tt.tlsInsecure)
561562

562-
err = cc.SetTLSEndpoint(tt.tlsEndpoint)
563+
err = cc.SetTLSEndpoint(context.Background(), tt.tlsEndpoint)
563564
if tt.expectCertsFetchErr {
564565
require.EqualError(t, err, tt.expectCertsFetcMsg)
565566
} else {
566567
require.NoError(t, err, "SetTLSEndpoint require NoError")
567568
}
568569
}
569570

570-
errPrint := cc.PrintData(&buffer)
571+
errPrint := cc.PrintData(context.Background(), &buffer)
571572
require.NoError(t, errPrint)
572573

573574
got := buffer.String()

‎internal/certinfo/certinfo_test.go‎

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package certinfo
22

33
import (
44
"bytes"
5+
"context"
56
"crypto/tls"
67
"fmt"
78
"net/http"
@@ -400,7 +401,7 @@ func TestCertinfo_SetTLSEndpoint(t *testing.T) {
400401
cc, errNew := New()
401402
require.NoError(t, errNew)
402403

403-
err := cc.SetTLSEndpoint(tt.endpoint)
404+
err := cc.SetTLSEndpoint(context.Background(), tt.endpoint)
404405

405406
if !tt.processErr {
406407
// skip requiring NoError since SetTLSEndpoint will always return network errors
@@ -440,10 +441,10 @@ func TestCertinfo_ProbeTLSInfo(t *testing.T) {
440441
cc.SetTLSInsecure(true)
441442
cc.SetTLSServerName("example.com")
442443

443-
err = cc.SetTLSEndpoint(u.Host)
444+
err = cc.SetTLSEndpoint(context.Background(), u.Host)
444445
require.NoError(t, err)
445446

446-
err = cc.ProbeTLSInfo()
447+
err = cc.ProbeTLSInfo(context.Background())
447448
require.NoError(t, err)
448449

449450
// Since it's a local TLS server run by Go's httptest, it supports TLS 1.3 or TLS 1.2
@@ -470,7 +471,7 @@ func TestCertinfo_ProbeTLSInfo_NotRequested(t *testing.T) {
470471
cc.SetTLSInfoRequested(false)
471472
require.False(t, cc.TLSInfoRequested)
472473

473-
err = cc.ProbeTLSInfo()
474+
err = cc.ProbeTLSInfo(context.Background())
474475
require.NoError(t, err)
475476
require.Empty(t, cc.NegotiatedProtocol)
476477
}
@@ -483,7 +484,7 @@ func TestCertinfo_ProbeTLSInfo_NoEndpoint(t *testing.T) {
483484

484485
cc.SetTLSInfoRequested(true)
485486

486-
err = cc.ProbeTLSInfo()
487+
err = cc.ProbeTLSInfo(context.Background())
487488
require.NoError(t, err)
488489
require.Empty(t, cc.ProbedProtocols)
489490
}
@@ -502,7 +503,7 @@ func TestCertinfo_ProbeTLSInfo_Unreachable(t *testing.T) {
502503
cc.TLSEndpointHost = "127.0.0.1"
503504
cc.TLSEndpointPort = "54321"
504505

505-
err = cc.ProbeTLSInfo()
506+
err = cc.ProbeTLSInfo(context.Background())
506507
require.NoError(t, err)
507508

508509
// When unreachable, all scanned protocols should be unsupported
@@ -655,7 +656,7 @@ func TestCertinfo_ProbeTLSInfo_SingleCipher(t *testing.T) {
655656
},
656657
}
657658

658-
res := cc.probeCiphersConcurrently(ciphers)
659+
res := cc.probeCiphersConcurrently(context.Background(), ciphers)
659660
require.Len(t, res, 1)
660661
require.Equal(t, "TLS_AES_128_GCM_SHA256", res[0].Name)
661662
require.False(t, res[0].Supported)
@@ -673,7 +674,7 @@ func TestCertinfo_PrintData_WithTLSInfo(t *testing.T) {
673674

674675
var buf bytes.Buffer
675676

676-
err = cc.PrintData(&buf)
677+
err = cc.PrintData(context.Background(), &buf)
677678
require.NoError(t, err)
678679
require.Contains(t, buf.String(), "Negotiated TLS Connection")
679680
}

0 commit comments

Comments
 (0)