Skip to content
File

Blob: overlay/alpn_mux.go

go113 lines
1package overlay
2 
3import (
4 "context"
5 "crypto/tls"
6 "errors"
7 "sync"
8 "time"
9 
10 "go.miragespace.co/specter/spec/transport/q"
11 "go.miragespace.co/specter/util/acceptor"
12 
13 "github.com/quic-go/quic-go"
14)
15 
16type ALPNMux struct {
17 listener *quic.Listener
18 mux sync.Map // map[string]*protoCfg
19}
20 
21type protoCfg struct {
22 acceptor *acceptor.HTTP3Acceptor
23 tls *tls.Config
24}
25 
26func NewMux(tr *quic.Transport) (*ALPNMux, error) {
27 a := &ALPNMux{}
28 q, err := tr.Listen(&tls.Config{
29 GetConfigForClient: a.getConfigForClient,
30 }, quicConfig)
31 if err != nil {
32 return nil, err
33 }
34 a.listener = q
35 return a, nil
36}
37 
38func (a *ALPNMux) getConfigForClient(hello *tls.ClientHelloInfo) (*tls.Config, error) {
39 var baseCfg *tls.Config
40 var xCfg *tls.Config
41 var selected string
42 var found bool
43 
44 for _, propose := range hello.SupportedProtos {
45 cfg, ok := a.mux.Load(propose)
46 if ok {
47 found = true
48 selected = propose
49 baseCfg = cfg.(*protoCfg).tls
50 break
51 }
52 }
53 
54 if found {
55 xCfg = baseCfg.Clone()
56 xCfg.NextProtos = []string{selected}
57 return xCfg, nil
58 }
59 
60 return nil, errors.New("cipher: no mutually supported protocols")
61}
62 
63func (a *ALPNMux) With(baseCfg *tls.Config, protos ...string) q.Listener {
64 cfg := &protoCfg{
65 acceptor: acceptor.NewH3Acceptor(a.listener),
66 tls: baseCfg,
67 }
68 for _, proto := range protos {
69 a.mux.Store(proto, cfg)
70 }
71 return cfg.acceptor
72}
73 
74func (a *ALPNMux) Accept(ctx context.Context) {
75 for {
76 conn, err := a.listener.Accept(ctx)
77 if err != nil {
78 return
79 }
80 go a.handleConnection(ctx, conn)
81 }
82}
83 
84func (a *ALPNMux) Close() {
85 a.mux.Range(func(_, value any) bool {
86 cfg := value.(*protoCfg)
87 cfg.acceptor.Close()
88 return true
89 })
90 a.listener.Close()
91}
92 
93func (a *ALPNMux) handleConnection(ctx context.Context, conn *quic.Conn) {
94 select {
95 case <-time.After(quicConfig.HandshakeIdleTimeout):
96 conn.CloseWithError(401, "Gone")
97 return
98 case <-ctx.Done():
99 conn.CloseWithError(401, "Gone")
100 return
101 case <-conn.HandshakeComplete():
102 }
103 
104 cs := conn.ConnectionState().TLS
105 
106 cfg, ok := a.mux.Load(cs.NegotiatedProtocol)
107 if !ok {
108 conn.CloseWithError(404, "Unsupported protocol")
109 return
110 }
111 cfg.(*protoCfg).acceptor.Handle(conn)
112}