Skip to content
File

Blob: tun/server/server.go

go306 lines
1package server
2 
3import (
4 "context"
5 "errors"
6 "net"
7 "time"
8 
9 "go.miragespace.co/specter/spec/chord"
10 "go.miragespace.co/specter/spec/cipher"
11 "go.miragespace.co/specter/spec/protocol"
12 "go.miragespace.co/specter/spec/rpc"
13 "go.miragespace.co/specter/spec/transport"
14 "go.miragespace.co/specter/spec/tun"
15 "go.miragespace.co/specter/util/acceptor"
16 
17 "github.com/Yiling-J/theine-go"
18 "go.uber.org/zap"
19 "golang.org/x/sync/singleflight"
20)
21 
22// Gateway procedure:
23// use the hostname to lookup chordHash(hostname-[1..3])
24// lookup those keys sequentially on DHT
25// if all failed, that means hostname does not point to a valid tunnel
26// if any of the one is good, fetch (ClientIdentity, ServerIdentity) from DHT
27// if ServerIdentity is ourself, we have a direct connection to client
28// otherwise, DialDirect to that server to get a tunnel Stream
29// in either case, pipe the gateway connection to (direct|indirect) Stream
30// tunnel is now established
31 
32type Config struct {
33 Logger *zap.Logger
34 ParentContext context.Context
35 Chord chord.VNode
36 TunnelTransport transport.Transport
37 ChordTransport transport.Transport
38 Resolver tun.DNSResolver
39 CertProvider cipher.CertProvider
40 Apex string
41 Acme string
42}
43 
44type Server struct {
45 sessions *sessionRegistry
46 ephemeralLoads chan struct{}
47 rpcAcceptor *acceptor.HTTP2Acceptor
48 routeCache *theine.Cache[string, *routesResult]
49 routeLoads singleflight.Group
50 keylessCache *theine.LoadingCache[string, keylessCertResult]
51 Config
52}
53 
54var _ tun.Server = (*Server)(nil)
55 
56var _ protocol.TunnelService = (*Server)(nil)
57 
58func New(cfg Config) *Server {
59 s := &Server{
60 Config: cfg,
61 sessions: newSessionRegistry(),
62 ephemeralLoads: make(chan struct{}, 16),
63 rpcAcceptor: acceptor.NewH2Acceptor(nil),
64 }
65 s.initRouteCache()
66 s.initKeylessCache()
67 return s
68}
69 
70func (s *Server) Identity() *protocol.Node {
71 return s.TunnelTransport.Identity()
72}
73 
74func (s *Server) AttachRouter(ctx context.Context, router *transport.StreamRouter) {
75 router.HandleChord(protocol.Stream_PROXY, nil, func(delegate *transport.StreamDelegate) {
76 l := s.Logger.With(
77 zap.Object("peer", delegate.Identity),
78 zap.String("remote", delegate.RemoteAddr().String()),
79 zap.String("local", delegate.LocalAddr().String()),
80 )
81 defer func() {
82 if err := recover(); err != nil {
83 l.Warn("Panic recovered while handling proxy connection", zap.Any("error", err))
84 }
85 }()
86 s.handleProxyConn(ctx, delegate)
87 })
88 router.HandleTunnel(protocol.Stream_DIRECT, func(delegate *transport.StreamDelegate) {
89 // client uses this to register connection
90 // but it is a no-op on the server side
91 delegate.Close()
92 })
93 s.attachRPC(ctx, router)
94}
95 
96func (s *Server) MustRegister(ctx context.Context) {
97 s.Logger.Info("Publishing destinations to chord")
98 
99 publishCtx, cancel := context.WithTimeout(context.Background(), time.Second*30)
100 defer cancel()
101 
102 if err := s.publishDestinations(publishCtx); err != nil {
103 s.Logger.Panic("Error publishing destinations", zap.Error(err))
104 }
105 
106 s.Logger.Info("specter server started")
107}
108 
109func (s *Server) Stop() {
110 defer s.routeCache.Close()
111 defer s.keylessCache.Close()
112 
113 s.rpcAcceptor.Close()
114 
115 ctx, cancel := context.WithTimeout(context.Background(), time.Second*30)
116 defer cancel()
117 
118 s.unpublishDestinations(ctx)
119}
120 
121func (s *Server) handleProxyConn(ctx context.Context, delegation *transport.StreamDelegate) {
122 var err error
123 var clientConn net.Conn
124 
125 defer func() {
126 tun.SendStatusProto(delegation, err)
127 if err != nil {
128 delegation.Close()
129 return
130 }
131 tun.Pipe(delegation, clientConn)
132 }()
133 
134 route := &protocol.TunnelRoute{}
135 err = rpc.BoundedReceive(delegation, route, 2048)
136 if err != nil {
137 s.Logger.Error("Error receiving remote tunnel negotiation", zap.Error(err))
138 return
139 }
140 
141 l := s.Logger.With(
142 zap.String("hostname", route.GetHostname()),
143 zap.Uint64("client", route.GetClientDestination().GetId()),
144 )
145 l.Debug("Received proxy stream from remote node",
146 zap.Object("remote_chord", delegation.Identity),
147 zap.Object("chord", route.GetChordDestination()),
148 zap.Object("tunnel", route.GetTunnelDestination()))
149 
150 if route.GetTunnelDestination().GetAddress() != s.TunnelTransport.Identity().GetAddress() {
151 l.Warn("Received remote connection for the wrong server",
152 zap.String("expected", s.TunnelTransport.Identity().GetAddress()),
153 zap.String("got", route.GetTunnelDestination().GetAddress()),
154 )
155 err = tun.ErrDestinationNotFound
156 return
157 }
158 
159 clientConn, err = s.dialDirect(ctx, route)
160 if err != nil && !tun.IsNoDirect(err) {
161 l.Error("Error dialing connection to connected client", zap.Error(err))
162 }
163}
164 
165func (s *Server) getConn(ctx context.Context, route *protocol.TunnelRoute, link *protocol.Link) (net.Conn, error) {
166 l := s.Logger.With(
167 zap.String("hostname", route.GetHostname()),
168 zap.Uint64("client", route.GetClientDestination().GetId()),
169 )
170 
171 var conn net.Conn
172 var err error
173 direct := route.GetTunnelDestination().GetAddress() == s.TunnelTransport.Identity().GetAddress()
174 if direct {
175 l.Debug("client is connected to us, opening direct stream")
176 conn, err = s.dialDirect(ctx, route)
177 } else {
178 l.Debug("client is connected to remote node, opening proxy stream",
179 zap.Object("chord", route.GetChordDestination()),
180 zap.Object("tunnel", route.GetTunnelDestination()))
181 
182 conn, err = s.ChordTransport.DialStream(ctx, route.GetChordDestination(), protocol.Stream_PROXY)
183 }
184 if err != nil {
185 return nil, err
186 }
187 
188 // Transports may retain the dialing context for a physical connection. Keep
189 // this stream's negotiation timeout separate from that connection lifecycle.
190 ctx, cancel := context.WithTimeout(ctx, lookupTimeout)
191 defer cancel()
192 
193 // Bound stream negotiation, including writes, and close failed streams before
194 // trying another route. Cancelled callers also interrupt blocked I/O.
195 stop := context.AfterFunc(ctx, func() { conn.Close() })
196 defer stop()
197 ready := false
198 defer func() {
199 if !ready {
200 conn.Close()
201 }
202 }()
203 deadline, _ := ctx.Deadline()
204 if err := conn.SetDeadline(deadline); err != nil {
205 return nil, err
206 }
207 if !direct {
208 if err := rpc.Send(conn, route); err != nil {
209 l.Error("sending remote tunnel negotiation", zap.Error(err))
210 return nil, err
211 }
212 status := &protocol.TunnelStatus{}
213 if err := rpc.BoundedReceive(conn, status, 1024); err != nil {
214 l.Error("Error receiving remote tunnel status", zap.Error(err))
215 return nil, err
216 }
217 switch status.GetStatus() {
218 case protocol.TunnelStatusCode_STATUS_OK:
219 case protocol.TunnelStatusCode_NO_DIRECT:
220 return nil, tun.ErrTunnelClientNotConnected
221 default:
222 return nil, errors.New(status.GetError())
223 }
224 }
225 if err := rpc.Send(conn, link); err != nil {
226 return nil, err
227 }
228 if !stop() {
229 return nil, ctx.Err()
230 }
231 if err := conn.SetDeadline(time.Time{}); err != nil {
232 return nil, err
233 }
234 ready = true
235 return conn, nil
236}
237 
238func (s *Server) DialInternal(ctx context.Context, node *protocol.Node) (net.Conn, error) {
239 if node.GetAddress() == "" || node.GetUnknown() {
240 return nil, transport.ErrNoDirect
241 }
242 return s.ChordTransport.DialStream(ctx, node, protocol.Stream_INTERNAL)
243}
244 
245func (s *Server) DialClient(ctx context.Context, link *protocol.Link) (net.Conn, error) {
246 ret, err := s.lookupRoutes(ctx, link.GetHostname(), nil)
247 if err != nil {
248 return nil, err
249 }
250 
251 var isNoRoute bool
252 for attempt := 0; attempt < 2; attempt++ {
253 if ret.err != nil {
254 return nil, ret.err
255 }
256 for _, route := range ret.routes {
257 if err := ctx.Err(); err != nil {
258 return nil, err
259 }
260 conn, err := s.getConn(ctx, route, link)
261 if err == nil {
262 return conn, nil
263 }
264 if tun.IsNoDirect(err) {
265 isNoRoute = true
266 } else {
267 s.Logger.Error("Failed to establish connection to client",
268 zap.String("hostname", link.GetHostname()),
269 zap.Object("chord", route.GetChordDestination()),
270 zap.Object("tunnel", route.GetTunnelDestination()),
271 zap.Object("client", route.GetClientDestination()),
272 zap.Error(err),
273 )
274 }
275 }
276 if err := ctx.Err(); err != nil {
277 return nil, err
278 }
279 if attempt == 1 {
280 break
281 }
282 fresh, err := s.lookupRoutes(ctx, link.GetHostname(), ret)
283 if err != nil {
284 return nil, err
285 }
286 if fresh == ret {
287 break
288 }
289 ret = fresh
290 }
291 
292 if isNoRoute {
293 return nil, tun.ErrTunnelClientNotConnected
294 }
295 
296 return nil, tun.ErrDestinationNotFound // fallback to not found
297}
298 
299func (s *Server) dialDirect(ctx context.Context, route *protocol.TunnelRoute) (net.Conn, error) {
300 client := route.GetClientDestination()
301 if tun.IsSessionAlias(client.GetAddress()) {
302 return s.sessions.dial(client.GetAddress(), route.GetHostname())
303 }
304 return s.TunnelTransport.DialStream(ctx, client, protocol.Stream_DIRECT)
305}