Something went wrong. Try again.
Monorepo for Tangled tangled.org
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199package mill
import ( "github.com/gorilla/websocket" "io" "net/http" "slices" "strings" "sync" "time"
millproto "tangled.org/core/spindle/mill/proto" millv1 "tangled.org/core/spindle/mill/proto/gen")
var ( handshakeSem = make(chan struct{}, 16) handshakeMu sync.Mutex inFlightHandshakes = make(map[string]struct{}))
type livenessReader struct { r io.Reader conn *websocket.Conn readTimeout time.Duration}
func (lr *livenessReader) Read(p []byte) (int, error) { if err := lr.conn.SetReadDeadline(time.Now().Add(lr.readTimeout)); err != nil { return 0, err } return lr.r.Read(p)}
var upgrader = websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024,}
// auth before upgrade, a bad token never opens a socketfunc (m *Mill) HandleExecutorConn(w http.ResponseWriter, r *http.Request) { name, authorizedLabels, ok := m.authenticate(r) if !ok { http.Error(w, "unauthorized", http.StatusUnauthorized) return }
select { case handshakeSem <- struct{}{}: case <-r.Context().Done(): return } handshakeSlotHeld := true defer func() { if handshakeSlotHeld { <-handshakeSem } }()
// enforces one in-flight handshake and one live session at a time per identity handshakeMu.Lock() if _, ok := inFlightHandshakes[name]; ok { handshakeMu.Unlock() http.Error(w, "handshake already in progress", http.StatusConflict) return } m.mu.Lock() old, exists := m.sessions[name] isLive := exists && old.live(m.cfg.ReconnectGrace) m.mu.Unlock() if isLive { handshakeMu.Unlock() http.Error(w, "session already active", http.StatusConflict) return } inFlightHandshakes[name] = struct{}{} handshakeMu.Unlock() identityHandshakeHeld := true
defer func() { if identityHandshakeHeld { handshakeMu.Lock() delete(inFlightHandshakes, name) handshakeMu.Unlock() } }()
conn, err := upgrader.Upgrade(w, r, nil) if err != nil { m.l.Error("fleet ws upgrade failed", "err", err) return } defer conn.Close()
if err := conn.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil { m.l.Error("failed to set pre-hello read deadline", "err", err) return }
stream := millproto.NewWSStream(conn) enc := millproto.NewEncoder(stream) dec := millproto.NewDecoder(stream)
hello, err := dec.Decode() if err != nil { m.l.Error("fleet read hello failed", "err", err) return } h := hello.GetHello() if h == nil { m.l.Error("fleet first frame was not hello") return } protocolVersion := h.GetProtocolVersion() minProtocol, maxProtocol := h.GetMinProtocolVersion(), h.GetMaxProtocolVersion() legacyProtocol := minProtocol == 0 && maxProtocol == 0 if legacyProtocol { // v4 executors predate capability negotiation. keep them available for // ordinary jobs, but never admit them for cache placement. minProtocol, maxProtocol = protocolVersion, protocolVersion } if minProtocol == 0 || maxProtocol == 0 || minProtocol > maxProtocol || protocolVersion < minProtocol || protocolVersion > maxProtocol || protocolVersion < millproto.ProtocolMinVersion || protocolVersion > millproto.ProtocolMaxVersion { m.l.Error("fleet protocol version range mismatch", "version", protocolVersion, "min", minProtocol, "max", maxProtocol, "supportedMin", millproto.ProtocolMinVersion, "supportedMax", millproto.ProtocolMaxVersion) return } if h.GetEpoch() == "" { m.l.Error("fleet hello missing epoch") return }
for _, l := range h.GetLabels() { if !slices.Contains(authorizedLabels, l) { m.l.Error("executor requested unauthorized label", "label", l, "authorized", authorizedLabels) return } }
sess := newSession(name, h.GetEpoch(), authorizedLabels, enc, m.l) sess.closeTransport = conn.Close sess.labels = h.GetLabels() sess.arch = h.GetArch() sess.cacheStoreID = h.GetCacheStoreId() sess.cacheNamespace = h.GetCacheNamespace() sess.legacyProtocol = legacyProtocol
resume, ok := m.attachSession(sess) if !ok { m.l.Warn("rejecting duplicate live executor session", "node", name) return } m.l.Info("executor connected", "node", sess.nodeID, "arch", h.GetArch(), "labels", h.GetLabels(), "resume", resume)
if err := sess.send(&millproto.Message{Resume: &millv1.Resume{Epoch: h.GetEpoch(), AckSeqno: resume}}); err != nil { m.l.Error("fleet send resume failed", "err", err) m.detachSession(sess) return } handshakeSlotHeld = false <-handshakeSem handshakeMu.Lock() delete(inFlightHandshakes, name) handshakeMu.Unlock() identityHandshakeHeld = false m.sessionReady(sess)
readTimeout := m.cfg.ReconnectGrace if readTimeout <= 0 { readTimeout = 45 * time.Second } liveDec := millproto.NewDecoder(&livenessReader{r: stream, conn: conn, readTimeout: readTimeout})
if err := sess.readLoop(m, liveDec); err != nil { m.l.Debug("session read ended", "node", sess.nodeID, "err", err) m.noteSessionError(sess, err) } m.detachSession(sess)}
// identity comes from the token hash. unknown or missing token fails closedfunc (m *Mill) authenticate(r *http.Request) (string, []string, bool) { const prefix = "Bearer " h := r.Header.Get("Authorization") if !strings.HasPrefix(h, prefix) { return "", nil, false } token := strings.TrimPrefix(h, prefix) if token == "" || m.db == nil { return "", nil, false } name, labels, ok, err := m.db.ResolveExecutorToken(HashToken(token)) if err != nil { m.l.Error("executor token lookup failed", "err", err) return "", nil, false } return name, labels, ok}