Skip to content
File

Blob: cmd/server/listener_group.go

go232 lines
1package server
2 
3import (
4 "context"
5 "crypto/tls"
6 "errors"
7 "net"
8 "sync"
9 
10 cmdlisten "go.miragespace.co/specter/cmd/internal/listen"
11 "go.miragespace.co/specter/spec/transport/q"
12 
13 "github.com/quic-go/quic-go"
14 "go.uber.org/atomic"
15)
16 
17type multiListener struct {
18 listeners []net.Listener
19 stop chan struct{}
20 connCh chan net.Conn
21 closeOnce sync.Once
22 wg sync.WaitGroup
23 err atomic.Value
24}
25 
26func newMultiListener(listeners []net.Listener) net.Listener {
27 if len(listeners) == 1 {
28 return listeners[0]
29 }
30 
31 m := &multiListener{
32 listeners: listeners,
33 stop: make(chan struct{}),
34 connCh: make(chan net.Conn, len(listeners)),
35 }
36 for _, l := range listeners {
37 m.wg.Add(1)
38 go m.serve(l)
39 }
40 go func() {
41 m.wg.Wait()
42 close(m.connCh)
43 }()
44 return m
45}
46 
47func (m *multiListener) serve(l net.Listener) {
48 defer m.wg.Done()
49 for {
50 conn, err := l.Accept()
51 if err != nil {
52 if !errors.Is(err, net.ErrClosed) && m.err.Load() == nil {
53 m.err.Store(err)
54 }
55 return
56 }
57 select {
58 case m.connCh <- conn:
59 case <-m.stop:
60 conn.Close()
61 return
62 }
63 }
64}
65 
66func (m *multiListener) Accept() (net.Conn, error) {
67 select {
68 case conn, ok := <-m.connCh:
69 if !ok {
70 if err, ok := m.err.Load().(error); ok && err != nil {
71 return nil, err
72 }
73 return nil, net.ErrClosed
74 }
75 return conn, nil
76 case <-m.stop:
77 return nil, net.ErrClosed
78 }
79}
80 
81func (m *multiListener) Close() error {
82 m.closeOnce.Do(func() {
83 close(m.stop)
84 for _, l := range m.listeners {
85 l.Close()
86 }
87 m.wg.Wait()
88 })
89 
90 if err, ok := m.err.Load().(error); ok && err != nil {
91 return err
92 }
93 return nil
94}
95 
96func (m *multiListener) Addr() net.Addr {
97 if len(m.listeners) == 0 {
98 return nil
99 }
100 return m.listeners[0].Addr()
101}
102 
103type multiQuicListener struct {
104 listeners []q.Listener
105 ctx context.Context
106 cancel context.CancelFunc
107 connCh chan *quic.Conn
108 wg sync.WaitGroup
109 err atomic.Value
110}
111 
112func newMultiQuicListener(parent context.Context, listeners []q.Listener) q.Listener {
113 if len(listeners) == 1 {
114 return listeners[0]
115 }
116 
117 ctx, cancel := context.WithCancel(parent)
118 m := &multiQuicListener{
119 listeners: listeners,
120 ctx: ctx,
121 cancel: cancel,
122 connCh: make(chan *quic.Conn, len(listeners)),
123 }
124 for _, l := range listeners {
125 m.wg.Add(1)
126 go m.serve(l)
127 }
128 go func() {
129 m.wg.Wait()
130 close(m.connCh)
131 }()
132 return m
133}
134 
135func (m *multiQuicListener) serve(l q.Listener) {
136 defer m.wg.Done()
137 for {
138 conn, err := l.Accept(m.ctx)
139 if err != nil {
140 if m.ctx.Err() == nil && !errors.Is(err, net.ErrClosed) && m.err.Load() == nil {
141 m.err.Store(err)
142 }
143 return
144 }
145 select {
146 case m.connCh <- conn:
147 case <-m.ctx.Done():
148 conn.CloseWithError(0, "listener closed")
149 return
150 }
151 }
152}
153 
154func (m *multiQuicListener) Accept(ctx context.Context) (*quic.Conn, error) {
155 select {
156 case <-ctx.Done():
157 return nil, ctx.Err()
158 case conn, ok := <-m.connCh:
159 if !ok {
160 if err, ok := m.err.Load().(error); ok && err != nil {
161 return nil, err
162 }
163 return nil, net.ErrClosed
164 }
165 return conn, nil
166 }
167}
168 
169func (m *multiQuicListener) Close() error {
170 m.cancel()
171 for _, l := range m.listeners {
172 l.Close()
173 }
174 m.wg.Wait()
175 if err, ok := m.err.Load().(error); ok && err != nil {
176 return err
177 }
178 return nil
179}
180 
181func (m *multiQuicListener) Addr() net.Addr {
182 if len(m.listeners) == 0 {
183 return nil
184 }
185 return m.listeners[0].Addr()
186}
187 
188type udpBinding struct {
189 listen cmdlisten.Address
190 packetConn net.PacketConn
191 transport *quic.Transport
192}
193 
194type multiDialer struct {
195 bindings []udpBinding
196}
197 
198func newMultiDialer(bindings []udpBinding) *multiDialer {
199 return &multiDialer{bindings: bindings}
200}
201 
202func (m *multiDialer) DialEarly(ctx context.Context, addr net.Addr, tlsConf *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
203 target := cmdlisten.IPAny
204 if udpAddr, ok := addr.(*net.UDPAddr); ok && udpAddr.IP != nil {
205 if udpAddr.IP.To4() != nil {
206 target = cmdlisten.IPV4
207 } else if udpAddr.IP.To16() != nil {
208 target = cmdlisten.IPV6
209 }
210 }
211 
212 if tr := m.pick(target); tr != nil {
213 return tr.DialEarly(ctx, addr, tlsConf, cfg)
214 }
215 
216 return nil, errors.New("no udp listeners configured for quic transport")
217}
218 
219func (m *multiDialer) pick(target cmdlisten.IPVersion) *quic.Transport {
220 if target != cmdlisten.IPAny {
221 for _, b := range m.bindings {
222 if b.listen.Version == target {
223 return b.transport
224 }
225 }
226 }
227 if len(m.bindings) > 0 {
228 return m.bindings[0].transport
229 }
230 return nil
231}