Skip to content
File

Blob: tun/client/server.go

go369 lines
1package client
2 
3import (
4 "context"
5 "encoding/json"
6 "errors"
7 "fmt"
8 "io"
9 "net"
10 "net/http"
11 "net/url"
12 "os"
13 "time"
14 
15 "go.miragespace.co/specter/spec/acme"
16 "go.miragespace.co/specter/spec/protocol"
17 "go.miragespace.co/specter/spec/rpc"
18 "go.miragespace.co/specter/spec/transport"
19 "go.miragespace.co/specter/ui"
20 "go.miragespace.co/specter/util"
21 
22 "github.com/go-chi/chi/v5"
23 "github.com/go-chi/chi/v5/middleware"
24 "github.com/twitchtv/twirp"
25 "go.uber.org/zap"
26)
27 
28var _ protocol.ClientQueryService = (*Client)(nil)
29 
30func (c *Client) attachRPC(ctx context.Context, router *transport.StreamRouter) {
31 queryTwirp := protocol.NewClientQueryServiceServer(c)
32 
33 rpcHandler := chi.NewRouter()
34 rpcHandler.Use(middleware.Recoverer)
35 rpcHandler.Use(util.LimitBody(1 << 10)) // 1KB
36 rpcHandler.Mount(queryTwirp.PathPrefix(), queryTwirp)
37 
38 srv := &http.Server{
39 BaseContext: func(l net.Listener) context.Context {
40 return ctx
41 },
42 ConnContext: func(ctx context.Context, c net.Conn) context.Context {
43 return rpc.WithDelegation(ctx, c.(*transport.StreamDelegate))
44 },
45 MaxHeaderBytes: 1 << 10, // 1KB
46 ReadHeaderTimeout: time.Second * 3,
47 Handler: rpcHandler,
48 ErrorLog: util.GetStdLogger(c.Logger, "queryServer"),
49 }
50 
51 go srv.Serve(c.rpcAcceptor)
52 
53 router.HandleTunnel(protocol.Stream_RPC, func(delegate *transport.StreamDelegate) {
54 c.rpcAcceptor.Handle(delegate)
55 })
56}
57 
58func (c *Client) ListTunnels(ctx context.Context, _ *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error) {
59 c.configMu.RLock()
60 cfg := c.Configuration.clone()
61 c.configMu.RUnlock()
62 
63 tunnels := make([]*protocol.ClientTunnel, 0)
64 for _, tunnel := range cfg.Tunnels {
65 tunnels = append(tunnels, &protocol.ClientTunnel{
66 Hostname: tunnel.Hostname,
67 Target: tunnel.Target,
68 })
69 }
70 
71 return &protocol.ListTunnelsResponse{
72 Tunnels: tunnels,
73 }, nil
74}
75 
76type ClientStatus struct {
77 Apex string `json:"apex"`
78 ConnectedNodes []*protocol.Node `json:"connectedNodes"`
79 Synchronization SyncResult `json:"synchronization"`
80 Pending bool `json:"pending"`
81 RetryAt *time.Time `json:"retryAt,omitempty"`
82}
83 
84func (c *Client) getStatus() ClientStatus {
85 c.configMu.RLock()
86 defer c.configMu.RUnlock()
87 c.syncStateMu.RLock()
88 defer c.syncStateMu.RUnlock()
89 
90 result := c.lastSync
91 result.Tunnels = make([]TunnelSyncResult, 0, len(c.Configuration.Tunnels))
92 for _, tunnel := range c.Configuration.Tunnels {
93 outcome := TunnelSyncResult{Hostname: tunnel.Hostname, Target: tunnel.Target}
94 for _, previous := range c.lastSync.Tunnels {
95 if previous.Hostname == tunnel.Hostname && previous.Target == tunnel.Target {
96 outcome = previous
97 break
98 }
99 }
100 result.Tunnels = append(result.Tunnels, outcome)
101 }
102 status := ClientStatus{
103 Apex: c.Configuration.Apex, ConnectedNodes: c.getConnectedNodes(),
104 Synchronization: result, Pending: result.pendingPublication(),
105 }
106 if status.ConnectedNodes == nil {
107 status.ConnectedNodes = []*protocol.Node{}
108 }
109 if status.Pending {
110 next := c.nextSync
111 status.RetryAt = &next
112 }
113 return status
114}
115 
116func writeJSONResult(w http.ResponseWriter, status int, result any) {
117 w.Header().Set("Content-Type", "application/json")
118 w.WriteHeader(status)
119 json.NewEncoder(w).Encode(result)
120}
121 
122func writeActionError(w http.ResponseWriter, err error) {
123 var saveError *ConfigSaveError
124 writeJSONResult(w, http.StatusInternalServerError, SyncResult{
125 Applied: errors.As(err, &saveError), Error: err.Error(), Tunnels: []TunnelSyncResult{},
126 })
127}
128 
129func (c *Client) localHandler() http.Handler {
130 r := chi.NewRouter()
131 
132 r.Use(middleware.Heartbeat("/healthz"))
133 
134 api := chi.NewRouter()
135 
136 api.Route("/tokens", func(tokens chi.Router) {
137 tokens.Use(func(next http.Handler) http.Handler {
138 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
139 w.Header().Set("Cache-Control", "no-store")
140 w.Header().Set("Content-Type", "application/json")
141 next.ServeHTTP(w, r)
142 })
143 })
144 tokens.Get("/", func(w http.ResponseWriter, r *http.Request) {
145 resp, err := c.ListDelegations(r.Context())
146 if err != nil {
147 writeDelegationError(w, err)
148 return
149 }
150 c.FormatDelegations(resp, w)
151 })
152 tokens.With(util.LimitBody(1<<10)).Post("/", func(w http.ResponseWriter, r *http.Request) {
153 var body struct {
154 Hostname string `json:"hostname"`
155 ExpiresAt string `json:"expiresAt"`
156 }
157 decoder := json.NewDecoder(r.Body)
158 if err := decoder.Decode(&body); err != nil {
159 writeDelegationError(w, twirp.InvalidArgument.Error("invalid token request body"))
160 return
161 }
162 if err := decoder.Decode(new(any)); err != io.EOF {
163 writeDelegationError(w, twirp.InvalidArgument.Error("invalid token request body"))
164 return
165 }
166 var expiry time.Time
167 if body.ExpiresAt != "" {
168 parsed, err := time.Parse(time.RFC3339, body.ExpiresAt)
169 if err != nil {
170 writeDelegationError(w, twirp.InvalidArgument.Error("expiresAt must be RFC3339"))
171 return
172 }
173 expiry = parsed
174 }
175 resp, err := c.MintDelegation(r.Context(), body.Hostname, expiry)
176 if err != nil {
177 writeDelegationError(w, err)
178 return
179 }
180 w.WriteHeader(http.StatusCreated)
181 c.FormatMintedDelegation(resp, w)
182 })
183 tokens.Post("/{id}/revoke", func(w http.ResponseWriter, r *http.Request) {
184 resp, err := c.RevokeDelegation(r.Context(), chi.URLParam(r, "id"))
185 if err != nil {
186 writeDelegationError(w, err)
187 return
188 }
189 c.FormatRevokedDelegation(resp, w)
190 })
191 })
192 
193 api.Post("/reload", func(w http.ResponseWriter, r *http.Request) {
194 c.Logger.Info("Received request from API, reloading config")
195 result := c.doReload(r.Context())
196 if result.Error == "" {
197 w.WriteHeader(http.StatusNoContent)
198 return
199 }
200 status := http.StatusInternalServerError
201 if !result.Applied {
202 status = http.StatusBadRequest
203 }
204 writeJSONResult(w, status, result)
205 })
206 
207 api.Get("/status", func(w http.ResponseWriter, r *http.Request) {
208 w.Header().Set("Cache-Control", "no-store")
209 writeJSONResult(w, http.StatusOK, c.getStatus())
210 })
211 
212 api.Get("/config", func(w http.ResponseWriter, r *http.Request) {
213 cfg := c.GetCurrentConfig()
214 f, err := os.Open(cfg.path)
215 if err != nil {
216 http.Error(w, err.Error(), 500)
217 return
218 }
219 defer f.Close()
220 io.Copy(w, f)
221 })
222 
223 api.Post("/unpublish/{hostname}", func(w http.ResponseWriter, r *http.Request) {
224 hostname := chi.URLParam(r, "hostname")
225 hostname, err := url.PathUnescape(hostname)
226 if err != nil {
227 http.Error(w, err.Error(), http.StatusBadRequest)
228 return
229 }
230 hostname, err = acme.Normalize(hostname)
231 if err != nil {
232 http.Error(w, err.Error(), http.StatusBadRequest)
233 return
234 }
235 
236 err = c.UnpublishTunnel(r.Context(), Tunnel{
237 Hostname: hostname,
238 })
239 if err != nil {
240 writeActionError(w, err)
241 return
242 }
243 
244 fmt.Fprintf(w, "Tunnel %s unpublished from network\n", hostname)
245 })
246 
247 api.Post("/release/{hostname}", func(w http.ResponseWriter, r *http.Request) {
248 hostname := chi.URLParam(r, "hostname")
249 hostname, err := url.PathUnescape(hostname)
250 if err != nil {
251 http.Error(w, err.Error(), http.StatusBadRequest)
252 return
253 }
254 hostname, err = acme.Normalize(hostname)
255 if err != nil {
256 http.Error(w, err.Error(), http.StatusBadRequest)
257 return
258 }
259 
260 err = c.ReleaseTunnel(r.Context(), Tunnel{
261 Hostname: hostname,
262 })
263 if err != nil {
264 writeActionError(w, err)
265 return
266 }
267 
268 fmt.Fprintf(w, "Tunnel %s released from network\n", hostname)
269 })
270 
271 api.Get("/acme/{hostname}", func(w http.ResponseWriter, r *http.Request) {
272 hostname := chi.URLParam(r, "hostname")
273 hostname, err := url.PathUnescape(hostname)
274 if err != nil {
275 http.Error(w, err.Error(), http.StatusBadRequest)
276 return
277 }
278 hostname, err = acme.Normalize(hostname)
279 if err != nil {
280 http.Error(w, err.Error(), http.StatusBadRequest)
281 return
282 }
283 resp, err := c.GetAcmeInstruction(r.Context(), hostname)
284 if err != nil {
285 http.Error(w, err.Error(), http.StatusInternalServerError)
286 return
287 }
288 c.FormatAcme(resp, w)
289 })
290 
291 api.Get("/validate/{hostname}", func(w http.ResponseWriter, r *http.Request) {
292 hostname := chi.URLParam(r, "hostname")
293 hostname, err := url.PathUnescape(hostname)
294 if err != nil {
295 http.Error(w, err.Error(), http.StatusBadRequest)
296 return
297 }
298 hostname, err = acme.Normalize(hostname)
299 if err != nil {
300 http.Error(w, err.Error(), http.StatusBadRequest)
301 return
302 }
303 resp, err := c.RequestAcmeValidation(r.Context(), hostname)
304 if err != nil {
305 http.Error(w, err.Error(), http.StatusInternalServerError)
306 return
307 }
308 c.FormatValidate(hostname, resp, w)
309 })
310 
311 api.Get("/ls", func(w http.ResponseWriter, r *http.Request) {
312 hostnames, err := c.GetRegisteredHostnames(r.Context())
313 if err != nil {
314 http.Error(w, err.Error(), http.StatusInternalServerError)
315 return
316 }
317 c.FormatList(hostnames, w)
318 })
319 
320 r.Mount("/api", api)
321 r.Handle("/ui/*", http.StripPrefix("/ui", ui.Assets()))
322 r.Handle("/", ui.ClientPage())
323 
324 return r
325}
326 
327func (c *Client) startLocalServer(ctx context.Context) {
328 if c.ServerListener == nil {
329 return
330 }
331 
332 srv := &http.Server{
333 Handler: c.localHandler(),
334 ReadHeaderTimeout: connectTimeout,
335 ErrorLog: util.GetStdLogger(c.Logger, "localServer"),
336 BaseContext: func(l net.Listener) context.Context {
337 return ctx
338 },
339 }
340 
341 c.Logger.Info("Local server started", zap.String("listen", c.ServerListener.Addr().String()))
342 
343 go srv.Serve(c.ServerListener)
344}
345 
346func writeDelegationError(w http.ResponseWriter, err error) {
347 status := http.StatusInternalServerError
348 var te twirp.Error
349 if errors.Is(err, ErrDelegationUnsupported) {
350 status = http.StatusNotImplemented
351 } else if errors.As(err, &te) {
352 switch te.Code() {
353 case twirp.InvalidArgument:
354 status = http.StatusBadRequest
355 case twirp.PermissionDenied:
356 status = http.StatusForbidden
357 case twirp.Unauthenticated:
358 status = http.StatusUnauthorized
359 case twirp.NotFound:
360 status = http.StatusNotFound
361 case twirp.ResourceExhausted:
362 status = http.StatusTooManyRequests
363 }
364 }
365 writeJSONResult(w, status, struct {
366 Error string `json:"error"`
367 }{err.Error()})
368}