This repository has no description
1//go:build linux
2
3package microvm
4
5import (
6 "context"
7 "errors"
8 "fmt"
9 "log/slog"
10 "net"
11 "net/http"
12 "net/url"
13 "strings"
14 "time"
15
16 "github.com/mdlayher/vsock"
17)
18
19type UploadCacheBackend interface {
20 http.Handler
21 Close() error
22}
23
24type UploadCacheProxy struct {
25 port uint32
26
27 ln *vsock.Listener
28 server *http.Server
29 backend UploadCacheBackend
30}
31
32func StartUploadCacheProxy(ctx context.Context, cid uint32, uploadURL string, readUpstreams []CacheUpstream, stagingDir string, logger *slog.Logger) (*UploadCacheProxy, error) {
33 if strings.TrimSpace(uploadURL) == "" {
34 return nil, nil
35 }
36
37 if logger == nil {
38 logger = slog.Default()
39 }
40 logger = logger.With("where", "upload_cache_proxy", "cid", cid, "uploadURL", uploadURL)
41
42 backend, err := newUploadCacheBackend(uploadURL, readUpstreams, stagingDir, logger)
43 if err != nil {
44 return nil, err
45 }
46
47 ln, port, err := listenRandomVsockUploadPort(ctx)
48 if err != nil {
49 return nil, fmt.Errorf("listen for cache upload proxy: %w", err)
50 }
51
52 proxy := &UploadCacheProxy{
53 port: port,
54 ln: ln,
55 backend: backend,
56 }
57 proxy.server = &http.Server{
58 Handler: backend,
59 Protocols: cacheProxyProtocols(),
60 ReadHeaderTimeout: 30 * time.Second,
61 }
62
63 filtered := &cidFilteredVsockListener{
64 Listener: ln,
65 cid: cid,
66 logger: logger,
67 }
68 go func() {
69 if err := proxy.server.Serve(filtered); err != nil && !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
70 logger.Warn("upload cache proxy stopped", "port", port, "error", err)
71 }
72 }()
73
74 logger.Info("started upload cache proxy", "port", port, "target", uploadURL, "readUpstreams", len(readUpstreams))
75 return proxy, nil
76}
77
78func newUploadCacheBackend(uploadURL string, readUpstreams []CacheUpstream, stagingDir string, logger *slog.Logger) (UploadCacheBackend, error) {
79 if strings.TrimSpace(uploadURL) == "" {
80 return nil, nil
81 }
82
83 target, err := url.Parse(uploadURL)
84 if err != nil {
85 return nil, fmt.Errorf("parse upload URL %q: %w", uploadURL, err)
86 }
87
88 switch target.Scheme {
89 case "http", "https":
90 if target.Host == "" {
91 return nil, fmt.Errorf("upload URL %q is missing host", uploadURL)
92 }
93 return newHTTPUploadProxyBackend(target, readUpstreams, logger), nil
94
95 case "ssh", "ssh-ng":
96 return newNixStoreUploadBackend(target.String(), stagingDir, readUpstreams, logger, nil)
97
98 case "":
99 switch uploadURL {
100 case "daemon", "local":
101 return newNixStoreUploadBackend(uploadURL, stagingDir, readUpstreams, logger, nil)
102 default:
103 return nil, fmt.Errorf("unsupported upload URL %q", uploadURL)
104 }
105
106 default:
107 return nil, fmt.Errorf("upload URL %q uses unsupported scheme %q", uploadURL, target.Scheme)
108 }
109}
110
111func (p *UploadCacheProxy) Port() uint32 {
112 if p == nil {
113 return 0
114 }
115 return p.port
116}
117
118func (p *UploadCacheProxy) Close() error {
119 if p == nil {
120 return nil
121 }
122
123 var closeErr error
124 if p.server != nil {
125 ctx, cancel := context.WithTimeout(context.Background(), time.Second)
126 closeErr = errors.Join(closeErr, p.server.Shutdown(ctx))
127 cancel()
128 p.server = nil
129 }
130 if p.ln != nil {
131 closeErr = errors.Join(closeErr, p.ln.Close())
132 p.ln = nil
133 }
134 if p.backend != nil {
135 closeErr = errors.Join(closeErr, p.backend.Close())
136 }
137 return closeErr
138}
139
140func listenRandomVsockUploadPort(ctx context.Context) (*vsock.Listener, uint32, error) {
141 var lastErr error
142 for range 32 {
143 port, err := randomVsockPort()
144 if err != nil {
145 return nil, 0, err
146 }
147 ln, err := vsock.ListenContextID(vsock.Host, port, nil)
148 if err == nil {
149 return ln, port, nil
150 }
151 lastErr = err
152
153 select {
154 case <-ctx.Done():
155 return nil, 0, ctx.Err()
156 default:
157 }
158 }
159 return nil, 0, fmt.Errorf("listen on random vsock upload port: %w", lastErr)
160}