| 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 | } |