main
go 174 lines 4.62 KB
Raw
1 package transport
2
3 import (
4 "context"
5 "crypto/tls"
6 "encoding/json"
7 "errors"
8 "fmt"
9 "io"
10 "strings"
11 "time"
12
13 "github.com/quic-go/quic-go"
14 )
15
16 const (
17 quicBackhaulALPN = "portal-tunnel"
18 quicBackhaulControlTimeout = 10 * time.Second
19 quicBackhaulControlBodyLimit = 4096
20 )
21
22 type quicBackhaulControlMessage struct {
23 AccessToken string `json:"access_token"`
24 }
25
26 type quicBackhaulControlResponse struct {
27 OK bool `json:"ok"`
28 Error string `json:"error,omitempty"`
29 }
30
31 type QUICBackhaulControl struct {
32 AccessToken string
33 conn *quic.Conn
34 stream *quic.Stream
35 }
36
37 func ListenQUICBackhaul(addr string, cert tls.Certificate) (*quic.Listener, error) {
38 listener, err := quic.ListenAddr(addr, quicBackhaulServerTLSConfig(cert), quicBackhaulConfig())
39 if err != nil {
40 return nil, fmt.Errorf("listen quic backhaul: %w", err)
41 }
42 return listener, nil
43 }
44
45 func DialQUICBackhaul(ctx context.Context, addr string, tlsConfig *tls.Config, accessToken string) (*quic.Conn, error) {
46 accessToken = strings.TrimSpace(accessToken)
47 if accessToken == "" {
48 return nil, errors.New("quic backhaul access token is required")
49 }
50
51 conn, err := quic.DialAddr(ctx, addr, quicBackhaulClientTLSConfig(tlsConfig), quicBackhaulConfig())
52 if err != nil {
53 return nil, fmt.Errorf("dial quic backhaul: %w", err)
54 }
55
56 stream, err := conn.OpenStreamSync(ctx)
57 if err != nil {
58 _ = conn.CloseWithError(1, "control stream open failed")
59 return nil, fmt.Errorf("open quic backhaul control stream: %w", err)
60 }
61
62 _ = stream.SetDeadline(time.Now().Add(quicBackhaulControlTimeout))
63 if err := json.NewEncoder(stream).Encode(quicBackhaulControlMessage{AccessToken: accessToken}); err != nil {
64 _ = conn.CloseWithError(1, "control write failed")
65 return nil, fmt.Errorf("write quic backhaul control message: %w", err)
66 }
67
68 var resp quicBackhaulControlResponse
69 if err := json.NewDecoder(io.LimitReader(stream, quicBackhaulControlBodyLimit)).Decode(&resp); err != nil {
70 _ = conn.CloseWithError(1, "control response read failed")
71 return nil, fmt.Errorf("read quic backhaul control response: %w", err)
72 }
73 _ = stream.SetDeadline(time.Time{})
74 _ = stream.Close()
75
76 if !resp.OK {
77 errText := strings.TrimSpace(resp.Error)
78 if errText == "" {
79 errText = "rejected"
80 }
81 _ = conn.CloseWithError(1, errText)
82 return nil, fmt.Errorf("quic backhaul rejected: %s", errText)
83 }
84 return conn, nil
85 }
86
87 func AcceptQUICBackhaulControl(ctx context.Context, conn *quic.Conn) (*QUICBackhaulControl, error) {
88 stream, err := conn.AcceptStream(ctx)
89 if err != nil {
90 return nil, fmt.Errorf("accept quic backhaul control stream: %w", err)
91 }
92
93 _ = stream.SetReadDeadline(time.Now().Add(quicBackhaulControlTimeout))
94 var msg quicBackhaulControlMessage
95 if err := json.NewDecoder(io.LimitReader(stream, quicBackhaulControlBodyLimit)).Decode(&msg); err != nil {
96 return nil, fmt.Errorf("read quic backhaul control message: %w", err)
97 }
98 _ = stream.SetReadDeadline(time.Time{})
99
100 accessToken := strings.TrimSpace(msg.AccessToken)
101 if accessToken == "" {
102 return nil, errors.New("quic backhaul access token is required")
103 }
104
105 return &QUICBackhaulControl{
106 AccessToken: accessToken,
107 conn: conn,
108 stream: stream,
109 }, nil
110 }
111
112 func (c *QUICBackhaulControl) Accept() error {
113 if c == nil || c.stream == nil {
114 return nil
115 }
116 err := json.NewEncoder(c.stream).Encode(quicBackhaulControlResponse{OK: true})
117 return errors.Join(err, c.stream.Close())
118 }
119
120 func (c *QUICBackhaulControl) Reject(code, reason string) error {
121 if c == nil || c.conn == nil {
122 return nil
123 }
124 code = strings.TrimSpace(code)
125 if code == "" {
126 code = "rejected"
127 }
128 reason = strings.TrimSpace(reason)
129 if reason == "" {
130 reason = code
131 }
132
133 var err error
134 if c.stream != nil {
135 err = errors.Join(
136 json.NewEncoder(c.stream).Encode(quicBackhaulControlResponse{OK: false, Error: code}),
137 c.stream.Close(),
138 )
139 }
140 return errors.Join(err, c.conn.CloseWithError(1, reason))
141 }
142
143 func quicBackhaulServerTLSConfig(cert tls.Certificate) *tls.Config {
144 return &tls.Config{
145 Certificates: []tls.Certificate{cert},
146 NextProtos: []string{quicBackhaulALPN},
147 MinVersion: tls.VersionTLS13,
148 }
149 }
150
151 func quicBackhaulClientTLSConfig(base *tls.Config) *tls.Config {
152 if base == nil {
153 return &tls.Config{
154 NextProtos: []string{quicBackhaulALPN},
155 MinVersion: tls.VersionTLS13,
156 }
157 }
158
159 cfg := base.Clone()
160 cfg.NextProtos = []string{quicBackhaulALPN}
161 if cfg.MinVersion == 0 || cfg.MinVersion < tls.VersionTLS13 {
162 cfg.MinVersion = tls.VersionTLS13
163 }
164 return cfg
165 }
166
167 func quicBackhaulConfig() *quic.Config {
168 return &quic.Config{
169 EnableDatagrams: true,
170 KeepAlivePeriod: 15 * time.Second,
171 MaxIdleTimeout: 60 * time.Second,
172 MaxIncomingStreams: 16,
173 }
174 }