@@ -6,6 +6,7 @@ package certinfo
66
77import (
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
0 commit comments