This repository has no description
1package mill
2
3import (
4 "github.com/gorilla/websocket"
5 "io"
6 "net/http"
7 "slices"
8 "strings"
9 "sync"
10 "time"
11
12 millproto "tangled.org/core/spindle/mill/proto"
13 millv1 "tangled.org/core/spindle/mill/proto/gen"
14)
15
16var (
17 handshakeSem = make(chan struct{}, 16)
18 handshakeMu sync.Mutex
19 inFlightHandshakes = make(map[string]struct{})
20)
21
22type livenessReader struct {
23 r io.Reader
24 conn *websocket.Conn
25 readTimeout time.Duration
26}
27
28func (lr *livenessReader) Read(p []byte) (int, error) {
29 if err := lr.conn.SetReadDeadline(time.Now().Add(lr.readTimeout)); err != nil {
30 return 0, err
31 }
32 return lr.r.Read(p)
33}
34
35var upgrader = websocket.Upgrader{
36 ReadBufferSize: 1024,
37 WriteBufferSize: 1024,
38}
39
40// auth before upgrade, a bad token never opens a socket
41func (m *Mill) HandleExecutorConn(w http.ResponseWriter, r *http.Request) {
42 name, authorizedLabels, ok := m.authenticate(r)
43 if !ok {
44 http.Error(w, "unauthorized", http.StatusUnauthorized)
45 return
46 }
47
48 select {
49 case handshakeSem <- struct{}{}:
50 case <-r.Context().Done():
51 return
52 }
53 handshakeSlotHeld := true
54 defer func() {
55 if handshakeSlotHeld {
56 <-handshakeSem
57 }
58 }()
59
60 // enforces one in-flight handshake and one live session at a time per identity
61 handshakeMu.Lock()
62 if _, ok := inFlightHandshakes[name]; ok {
63 handshakeMu.Unlock()
64 http.Error(w, "handshake already in progress", http.StatusConflict)
65 return
66 }
67 m.mu.Lock()
68 old, exists := m.sessions[name]
69 isLive := exists && old.live(m.cfg.ReconnectGrace)
70 m.mu.Unlock()
71 if isLive {
72 handshakeMu.Unlock()
73 http.Error(w, "session already active", http.StatusConflict)
74 return
75 }
76 inFlightHandshakes[name] = struct{}{}
77 handshakeMu.Unlock()
78 identityHandshakeHeld := true
79
80 defer func() {
81 if identityHandshakeHeld {
82 handshakeMu.Lock()
83 delete(inFlightHandshakes, name)
84 handshakeMu.Unlock()
85 }
86 }()
87
88 conn, err := upgrader.Upgrade(w, r, nil)
89 if err != nil {
90 m.l.Error("fleet ws upgrade failed", "err", err)
91 return
92 }
93 defer conn.Close()
94
95 if err := conn.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil {
96 m.l.Error("failed to set pre-hello read deadline", "err", err)
97 return
98 }
99
100 stream := millproto.NewWSStream(conn)
101 enc := millproto.NewEncoder(stream)
102 dec := millproto.NewDecoder(stream)
103
104 hello, err := dec.Decode()
105 if err != nil {
106 m.l.Error("fleet read hello failed", "err", err)
107 return
108 }
109 h := hello.GetHello()
110 if h == nil {
111 m.l.Error("fleet first frame was not hello")
112 return
113 }
114 if h.GetProtocolVersion() != millproto.ProtocolVersion {
115 m.l.Error("fleet protocol version mismatch", "got", h.GetProtocolVersion(), "want", millproto.ProtocolVersion)
116 return
117 }
118 if h.GetEpoch() == "" {
119 m.l.Error("fleet hello missing epoch")
120 return
121 }
122
123 for _, l := range h.GetLabels() {
124 if !slices.Contains(authorizedLabels, l) {
125 m.l.Error("executor requested unauthorized label", "label", l, "authorized", authorizedLabels)
126 return
127 }
128 }
129
130 sess := newSession(name, h.GetEpoch(), authorizedLabels, enc, m.l)
131 sess.closeTransport = conn.Close
132 sess.labels = h.GetLabels()
133
134 resume, ok := m.attachSession(sess)
135 if !ok {
136 m.l.Warn("rejecting duplicate live executor session", "node", name)
137 return
138 }
139 m.l.Info("executor connected", "node", sess.nodeID, "arch", h.GetArch(), "labels", h.GetLabels(), "resume", resume)
140
141 if err := sess.send(&millproto.Message{Resume: &millv1.Resume{Epoch: h.GetEpoch(), AckSeqno: resume}}); err != nil {
142 m.l.Error("fleet send resume failed", "err", err)
143 m.detachSession(sess)
144 return
145 }
146 handshakeSlotHeld = false
147 <-handshakeSem
148 handshakeMu.Lock()
149 delete(inFlightHandshakes, name)
150 handshakeMu.Unlock()
151 identityHandshakeHeld = false
152 m.sessionReady(sess)
153
154 readTimeout := m.cfg.ReconnectGrace
155 if readTimeout <= 0 {
156 readTimeout = 45 * time.Second
157 }
158 liveDec := millproto.NewDecoder(&livenessReader{r: stream, conn: conn, readTimeout: readTimeout})
159
160 if err := sess.readLoop(m, liveDec); err != nil {
161 m.l.Debug("session read ended", "node", sess.nodeID, "err", err)
162 }
163 m.detachSession(sess)
164}
165
166// identity comes from the token hash. unknown or missing token fails closed
167func (m *Mill) authenticate(r *http.Request) (string, []string, bool) {
168 const prefix = "Bearer "
169 h := r.Header.Get("Authorization")
170 if !strings.HasPrefix(h, prefix) {
171 return "", nil, false
172 }
173 token := strings.TrimPrefix(h, prefix)
174 if token == "" || m.db == nil {
175 return "", nil, false
176 }
177 name, labels, ok, err := m.db.ResolveExecutorToken(HashToken(token))
178 if err != nil {
179 m.l.Error("executor token lookup failed", "err", err)
180 return "", nil, false
181 }
182 return name, labels, ok
183}