refact: tidy tls utils

rabbitprincess committed Apr 12, 2026 at 19:01 UTC 74294d0a7606f5c66e2ed68b47011e3127af7311
7 files changed +197 -196
portal/discovery/refresher.go
+43 -25
@@ -3,6 +3,7 @@ package discovery
3 import (
4 "context"
5 "crypto/tls"
6 + "crypto/x509"
7 "errors"
8 "fmt"
9 "net/http"
@@ -16,6 +17,7 @@ import (
17 )
18
19 const (
20 + defaultRequestTimeout = 15 * time.Second
21 DiscoveryPollInterval = 1 * time.Minute
22 defaultRecoveryFailures = 3
23 )
@@ -37,24 +39,26 @@ func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime)
39 if relaySet == nil {
40 return nil, errors.New("relay set is required")
41 }
40 - httpClient := http.DefaultClient
42 + var rootCAs *x509.CertPool
43 if len(rootCAPEM) > 0 {
42 - rootCAs, err := utils.CertPoolFromPEM(rootCAPEM)
43 - if err != nil {
44 - return nil, err
44 + rootCAs = x509.NewCertPool()
45 + if !rootCAs.AppendCertsFromPEM(rootCAPEM) {
46 + return nil, errors.New("failed to parse relay root ca")
47 }
46 - httpClient = &http.Client{
48 + }
49 + return &Refresher{
50 + relaySet: relaySet,
51 + httpClient: &http.Client{
52 Transport: &http.Transport{
53 TLSClientConfig: &tls.Config{
54 MinVersion: tls.VersionTLS12,
55 RootCAs: rootCAs,
56 + NextProtos: []string{"http/1.1"},
57 },
58 + ForceAttemptHTTP2: false,
59 },
53 - }
54 - }
55 - return &Refresher{
56 - relaySet: relaySet,
57 - httpClient: httpClient,
60 + Timeout: defaultRequestTimeout,
61 + },
62 overlay: overlay,
63 directRecoveryFailures: defaultRecoveryFailures,
64 overlayRecoveryFailures: defaultRecoveryFailures,
@@ -88,13 +92,26 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
92 continue
93 }
94 relay := state.Descriptor
91 - resp, err := r.discoverHTTPS(ctx, relay)
95 + baseURL, err := url.Parse(relay.APIHTTPSAddr)
96 if err != nil {
97 if ctx.Err() != nil {
98 return ctx.Err()
99 }
100 continue
101 }
102 + if utils.IsLocalRelayHost(baseURL.Hostname()) {
103 + log.Info().
104 + Str("relay", relay.APIHTTPSAddr).
105 + Msg("skip loopback relay as discovery source")
106 + continue
107 + }
108 + var resp types.DiscoveryResponse
109 + if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
110 + if ctx.Err() != nil {
111 + return ctx.Err()
112 + }
113 + continue
114 + }
115
116 now := time.Now().UTC()
117 _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
@@ -119,8 +136,22 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
136 continue
137 }
138 relay := state.Descriptor
122 - resp, err := r.discoverHTTPS(ctx, relay)
139 + baseURL, err := url.Parse(relay.APIHTTPSAddr)
140 if err != nil {
141 + if ctx.Err() != nil {
142 + return ctx.Err()
143 + }
144 + r.logDirectDiscoveryFailure(relay, fmt.Errorf("parse discovery base url: %w", err), r.directRecoveryFailures)
145 + continue
146 + }
147 + if utils.IsLocalRelayHost(baseURL.Hostname()) {
148 + log.Info().
149 + Str("relay", relay.APIHTTPSAddr).
150 + Msg("skip loopback relay as discovery source")
151 + continue
152 + }
153 + var resp types.DiscoveryResponse
154 + if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
155 if ctx.Err() != nil {
156 return ctx.Err()
157 }
@@ -138,19 +169,6 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
169 return ctx.Err()
170 }
171
141 -func (r *Refresher) discoverHTTPS(ctx context.Context, relay types.RelayDescriptor) (types.DiscoveryResponse, error) {
142 - baseURL, err := url.Parse(relay.APIHTTPSAddr)
143 - if err != nil {
144 - return types.DiscoveryResponse{}, fmt.Errorf("parse discovery base url: %w", err)
145 - }
146 -
147 - var resp types.DiscoveryResponse
148 - if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
149 - return types.DiscoveryResponse{}, err
150 - }
151 - return resp, nil
152 -}
153 -
172 func (r *Refresher) refreshOverlay(ctx context.Context) error {
173 r.relaySet.mu.RLock()
174 states := r.relaySet.relayStatesLocked()
portal/keyless/client.go
+1 -65
@@ -3,13 +3,10 @@ package keyless
3 import (
4 "context"
5 "crypto/tls"
6 - "encoding/pem"
6 "errors"
7 "fmt"
9 - "net"
8 "net/url"
9 "strings"
12 - "time"
10
11 keylesstls "github.com/gosuda/keyless_tls/keyless"
12
@@ -74,7 +71,7 @@ type ioCloser interface {
71 }
72
73 func ResolveMaterials(ctx context.Context, endpoint, serverName string) ([]byte, []byte, error) {
77 - chainPEM, err := FetchEndpointCertificateChain(ctx, endpoint, serverName)
74 + chainPEM, err := utils.FetchEndpointCertificateChain(ctx, endpoint, serverName)
75 if err != nil {
76 return nil, nil, fmt.Errorf("fetch signer certificate chain: %w", err)
77 }
@@ -91,64 +88,3 @@ func VerifyCertificateHostname(certPEM []byte, hostname string) error {
88 }
89 return leaf.VerifyHostname(hostname)
90 }
94 -
95 -func FetchEndpointCertificateChain(ctx context.Context, endpoint, serverName string) ([]byte, error) {
96 - raw := strings.TrimSpace(endpoint)
97 - if raw == "" {
98 - return nil, errors.New("endpoint is required")
99 - }
100 - if !strings.Contains(raw, "://") {
101 - raw = "https://" + raw
102 - }
103 -
104 - u, err := url.Parse(raw)
105 - if err != nil {
106 - return nil, fmt.Errorf("parse endpoint url: %w", err)
107 - }
108 - if !strings.EqualFold(u.Scheme, "https") {
109 - return nil, errors.New("keyless endpoint must use https")
110 - }
111 -
112 - host := u.Hostname()
113 - if host == "" {
114 - return nil, errors.New("endpoint hostname is empty")
115 - }
116 - port := u.Port()
117 - if port == "" {
118 - port = "443"
119 - }
120 - if serverName == "" {
121 - serverName = host
122 - }
123 -
124 - dialer := &net.Dialer{Timeout: 5 * time.Second}
125 - rawConn, err := dialer.DialContext(ctx, "tcp", net.JoinHostPort(host, port))
126 - if err != nil {
127 - return nil, fmt.Errorf("dial signer endpoint: %w", err)
128 - }
129 -
130 - tlsConn := tls.Client(rawConn, &tls.Config{
131 - MinVersion: tls.VersionTLS12,
132 - ServerName: serverName,
133 - InsecureSkipVerify: utils.IsLocalRelayHost(host),
134 - NextProtos: []string{"http/1.1"},
135 - })
136 - defer tlsConn.Close()
137 - if err := tlsConn.HandshakeContext(ctx); err != nil {
138 - return nil, fmt.Errorf("tls handshake with signer endpoint: %w", err)
139 - }
140 -
141 - peerCerts := tlsConn.ConnectionState().PeerCertificates
142 - if len(peerCerts) == 0 {
143 - return nil, errors.New("no peer certificates from signer endpoint")
144 - }
145 -
146 - var chainPEM []byte
147 - for _, cert := range peerCerts {
148 - chainPEM = append(chainPEM, pem.EncodeToMemory(&pem.Block{
149 - Type: "CERTIFICATE",
150 - Bytes: cert.Raw,
151 - })...)
152 - }
153 - return chainPEM, nil
154 -}
portal/keyless/http.go deleted
-41
@@ -1,41 +0,0 @@
1 -package keyless
2 -
3 -import (
4 - "context"
5 - "crypto/tls"
6 - "errors"
7 - "net/http"
8 - "net/url"
9 - "time"
10 -)
11 -
12 -func NewRelayHTTPClient(ctx context.Context, relayURL *url.URL, rootCAPEM []byte, timeout time.Duration) (*tls.Config, *http.Client, error) {
13 - if relayURL == nil {
14 - return nil, nil, errors.New("relay url is required")
15 - }
16 -
17 - serverName := relayURL.Hostname()
18 - if serverName == "" {
19 - return nil, nil, errors.New("relay hostname is required")
20 - }
21 -
22 - rootCAs, err := RelayRootCAs(ctx, relayURL.String(), serverName, rootCAPEM)
23 - if err != nil {
24 - return nil, nil, err
25 - }
26 -
27 - rawTLSConfig := &tls.Config{
28 - MinVersion: tls.VersionTLS12,
29 - ServerName: serverName,
30 - RootCAs: rootCAs,
31 - NextProtos: []string{"http/1.1"},
32 - }
33 - httpClient := &http.Client{
34 - Transport: &http.Transport{
35 - TLSClientConfig: rawTLSConfig.Clone(),
36 - ForceAttemptHTTP2: false,
37 - },
38 - Timeout: timeout,
39 - }
40 - return rawTLSConfig, httpClient, nil
41 -}
portal/keyless/tls.go
-17
@@ -1,32 +1,15 @@
1 package keyless
2
3 import (
4 - "context"
4 "crypto/tls"
6 - "crypto/x509"
5 "errors"
6 "fmt"
7 "io"
8 "net/http"
9
10 keylesstls "github.com/gosuda/keyless_tls/keyless"
13 -
14 - "github.com/gosuda/portal-tunnel/v2/utils"
11 )
12
17 -func RelayRootCAs(ctx context.Context, endpoint, serverName string, rootCAPEM []byte) (*x509.CertPool, error) {
18 - resolvedRootCAPEM := append([]byte(nil), rootCAPEM...)
19 - if len(resolvedRootCAPEM) == 0 && utils.IsLocalRelayHost(serverName) {
20 - _, fetchedRootCAPEM, err := ResolveMaterials(ctx, endpoint, serverName)
21 - if err != nil {
22 - return nil, fmt.Errorf("bootstrap localhost relay trust: %w", err)
23 - }
24 - resolvedRootCAPEM = fetchedRootCAPEM
25 - }
26 -
27 - return utils.CertPoolFromPEM(resolvedRootCAPEM)
28 -}
29 -
13 type TLSMaterialConfig struct {
14 Keyless *RemoteSignerConfig
15 CertPEM []byte
sdk/api_client.go
+1 -2
@@ -18,7 +18,6 @@ import (
18
19 "github.com/quic-go/quic-go"
20
21 - "github.com/gosuda/portal-tunnel/v2/portal/keyless"
21 "github.com/gosuda/portal-tunnel/v2/types"
22 "github.com/gosuda/portal-tunnel/v2/utils"
23 )
@@ -153,7 +152,7 @@ func (a *apiClient) ensureHTTPClient(ctx context.Context) error {
152 bootstrapCtx, cancel := context.WithTimeout(ctx, defaultDialTimeout+defaultHandshakeTimeout)
153 defer cancel()
154
156 - rawTLSConfig, httpClient, err := keyless.NewRelayHTTPClient(bootstrapCtx, a.baseURL, a.rootCAPEM, a.requestTimeout)
155 + rawTLSConfig, httpClient, err := utils.NewHTTPTLSClient(bootstrapCtx, a.baseURL, a.rootCAPEM, a.requestTimeout)
156 if err != nil {
157 return err
158 }
utils/tls.go new
+152
@@ -0,0 +1,152 @@
1 +package utils
2 +
3 +import (
4 + "context"
5 + "crypto"
6 + "crypto/ecdsa"
7 + "crypto/rsa"
8 + "crypto/tls"
9 + "crypto/x509"
10 + "encoding/pem"
11 + "errors"
12 + "fmt"
13 + "net"
14 + "net/http"
15 + "net/url"
16 + "strings"
17 + "time"
18 +)
19 +
20 +func NewHTTPTLSClient(ctx context.Context, relayURL *url.URL, rootCAPEM []byte, timeout time.Duration) (*tls.Config, *http.Client, error) {
21 + if relayURL == nil {
22 + return nil, nil, errors.New("relay url is required")
23 + }
24 +
25 + serverName := relayURL.Hostname()
26 + if serverName == "" {
27 + return nil, nil, errors.New("relay hostname is required")
28 + }
29 +
30 + resolvedRootCAPEM := append([]byte(nil), rootCAPEM...)
31 + if len(resolvedRootCAPEM) == 0 && IsLocalRelayHost(serverName) {
32 + fetchedRootCAPEM, err := FetchEndpointCertificateChain(ctx, relayURL.String(), serverName)
33 + if err != nil {
34 + return nil, nil, fmt.Errorf("bootstrap localhost relay trust: %w", err)
35 + }
36 + resolvedRootCAPEM = fetchedRootCAPEM
37 + }
38 +
39 + var rootCAs *x509.CertPool
40 + if len(resolvedRootCAPEM) > 0 {
41 + rootCAs = x509.NewCertPool()
42 + if !rootCAs.AppendCertsFromPEM(resolvedRootCAPEM) {
43 + return nil, nil, errors.New("failed to parse relay root ca")
44 + }
45 + }
46 +
47 + rawTLSConfig := &tls.Config{
48 + MinVersion: tls.VersionTLS12,
49 + ServerName: serverName,
50 + RootCAs: rootCAs,
51 + NextProtos: []string{"http/1.1"},
52 + }
53 + httpClient := &http.Client{
54 + Transport: &http.Transport{
55 + TLSClientConfig: rawTLSConfig.Clone(),
56 + ForceAttemptHTTP2: false,
57 + },
58 + Timeout: timeout,
59 + }
60 + return rawTLSConfig, httpClient, nil
61 +}
62 +
63 +func FetchEndpointCertificateChain(ctx context.Context, endpoint, serverName string) ([]byte, error) {
64 + raw := strings.TrimSpace(endpoint)
65 + if raw == "" {
66 + return nil, errors.New("endpoint is required")
67 + }
68 + if !strings.Contains(raw, "://") {
69 + raw = "https://" + raw
70 + }
71 +
72 + u, err := url.Parse(raw)
73 + if err != nil {
74 + return nil, fmt.Errorf("parse endpoint url: %w", err)
75 + }
76 + if !strings.EqualFold(u.Scheme, "https") {
77 + return nil, errors.New("relay endpoint must use https")
78 + }
79 +
80 + host := u.Hostname()
81 + if host == "" {
82 + return nil, errors.New("endpoint hostname is empty")
83 + }
84 + port := u.Port()
85 + if port == "" {
86 + port = "443"
87 + }
88 + if serverName == "" {
89 + serverName = host
90 + }
91 +
92 + dialer := &net.Dialer{Timeout: 5 * time.Second}
93 + rawConn, err := dialer.DialContext(ctx, "tcp", net.JoinHostPort(host, port))
94 + if err != nil {
95 + return nil, fmt.Errorf("dial relay endpoint: %w", err)
96 + }
97 +
98 + tlsConn := tls.Client(rawConn, &tls.Config{
99 + MinVersion: tls.VersionTLS12,
100 + ServerName: serverName,
101 + InsecureSkipVerify: IsLocalRelayHost(host),
102 + NextProtos: []string{"http/1.1"},
103 + })
104 + defer tlsConn.Close()
105 + if err := tlsConn.HandshakeContext(ctx); err != nil {
106 + return nil, fmt.Errorf("tls handshake with relay endpoint: %w", err)
107 + }
108 +
109 + peerCerts := tlsConn.ConnectionState().PeerCertificates
110 + if len(peerCerts) == 0 {
111 + return nil, errors.New("no peer certificates from relay endpoint")
112 + }
113 +
114 + var chainPEM []byte
115 + for _, cert := range peerCerts {
116 + chainPEM = append(chainPEM, pem.EncodeToMemory(&pem.Block{
117 + Type: "CERTIFICATE",
118 + Bytes: cert.Raw,
119 + })...)
120 + }
121 + return chainPEM, nil
122 +}
123 +
124 +func ParseCertificatePEM(pemData []byte) (*x509.Certificate, error) {
125 + block, _ := pem.Decode(pemData)
126 + if block == nil {
127 + return nil, errors.New("no pem block found")
128 + }
129 + return x509.ParseCertificate(block.Bytes)
130 +}
131 +
132 +func ParsePrivateKeyPEM(keyPEM []byte) (crypto.PrivateKey, error) {
133 + block, _ := pem.Decode(keyPEM)
134 + if block == nil {
135 + return nil, errors.New("invalid private key pem")
136 + }
137 + if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
138 + switch typed := key.(type) {
139 + case *ecdsa.PrivateKey:
140 + return typed, nil
141 + case *rsa.PrivateKey:
142 + return typed, nil
143 + }
144 + }
145 + if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
146 + return key, nil
147 + }
148 + if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
149 + return key, nil
150 + }
151 + return nil, errors.New("unsupported private key type")
152 +}
utils/utils.go
-46
@@ -2,14 +2,9 @@ package utils
2
3 import (
4 "context"
5 - "crypto"
6 - "crypto/ecdsa"
5 "crypto/rand"
8 - "crypto/rsa"
9 - "crypto/x509"
6 "encoding/base64"
7 "encoding/hex"
12 - "encoding/pem"
8 "errors"
9 "fmt"
10 "io"
@@ -482,47 +477,6 @@ func RandomHex(size int) (string, error) {
477 return hex.EncodeToString(buf), nil
478 }
479
485 -func CertPoolFromPEM(rootCAPEM []byte) (*x509.CertPool, error) {
486 - if len(rootCAPEM) == 0 {
487 - return nil, nil
488 - }
489 - pool := x509.NewCertPool()
490 - if !pool.AppendCertsFromPEM(rootCAPEM) {
491 - return nil, errors.New("failed to parse relay root ca")
492 - }
493 - return pool, nil
494 -}
495 -
496 -func ParseCertificatePEM(pemData []byte) (*x509.Certificate, error) {
497 - block, _ := pem.Decode(pemData)
498 - if block == nil {
499 - return nil, errors.New("no pem block found")
500 - }
501 - return x509.ParseCertificate(block.Bytes)
502 -}
503 -
504 -func ParsePrivateKeyPEM(keyPEM []byte) (crypto.PrivateKey, error) {
505 - block, _ := pem.Decode(keyPEM)
506 - if block == nil {
507 - return nil, errors.New("invalid private key pem")
508 - }
509 - if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
510 - switch typed := key.(type) {
511 - case *ecdsa.PrivateKey:
512 - return typed, nil
513 - case *rsa.PrivateKey:
514 - return typed, nil
515 - }
516 - }
517 - if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
518 - return key, nil
519 - }
520 - if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
521 - return key, nil
522 - }
523 - return nil, errors.New("unsupported private key type")
524 -}
525 -
480 func SleepOrDone(ctx context.Context, d time.Duration) bool {
481 timer := time.NewTimer(d)
482 defer timer.Stop()