diff --git a/core/http/auth/public_routes.go b/core/http/auth/public_routes.go index 658205a78..e1db88a02 100644 --- a/core/http/auth/public_routes.go +++ b/core/http/auth/public_routes.go @@ -5,6 +5,8 @@ package auth import ( "net/http" "strings" + + clusterep "github.com/mudler/LocalAI/core/http/endpoints/cluster" ) type publicRouteRule struct { @@ -77,5 +79,10 @@ func isPublicRoute(method, path string) bool { // usesAlternativeAuthentication identifies requests whose credentials are // validated by route-group middleware instead of the global auth middleware. func usesAlternativeAuthentication(path string) bool { - return strings.HasPrefix(path, "/api/node/") + // The peer link carries the cluster token in an Authorization header that + // no browser session ever sets, and its handler checks that token itself. + // The prefix comes from the endpoints package so the route and the + // exemption cannot drift apart. + return strings.HasPrefix(path, "/api/node/") || + strings.HasPrefix(path, clusterep.AlternativeAuthPrefix) } diff --git a/core/http/endpoints/cluster/cluster_suite_test.go b/core/http/endpoints/cluster/cluster_suite_test.go new file mode 100644 index 000000000..00f45e37a --- /dev/null +++ b/core/http/endpoints/cluster/cluster_suite_test.go @@ -0,0 +1,13 @@ +package cluster_test + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestClusterEndpoints(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Cluster Endpoints Suite") +} diff --git a/core/http/endpoints/cluster/peer.go b/core/http/endpoints/cluster/peer.go new file mode 100644 index 000000000..4b6ef542c --- /dev/null +++ b/core/http/endpoints/cluster/peer.go @@ -0,0 +1,104 @@ +// SPDX-License-Identifier: MIT + +// Package cluster serves the replica-to-replica link that a LocalAI frontend +// uses to reach a worker tunnel it does not own. A peer dials +// GET /api/cluster/peer, the connection becomes one multiplexed yamux session, +// and the relay opens a stream on it per request. +package cluster + +import ( + "crypto/subtle" + "net/http" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + "github.com/libp2p/go-yamux/v5" + "github.com/mudler/xlog" +) + +// AlternativeAuthPrefix is the path prefix whose credentials are checked by +// this package rather than by the global session middleware. The auth layer +// consults this same constant, so the two cannot drift apart and leave every +// peer dial answering 401. +const AlternativeAuthPrefix = "/api/cluster/" + +// PeerPath is the route a peer replica dials. +const PeerPath = AlternativeAuthPrefix + "peer" + +// RegisterClusterRoutes registers the peer link. onPeer receives every +// authenticated session; see PeerHandler for what it is expected to do with it. +func RegisterClusterRoutes(e *echo.Echo, token string, onPeer func(string, *yamux.Session)) { + e.GET(PeerPath, PeerHandler(token, onPeer)) +} + +// PeerHandler upgrades an authenticated peer dial to a WebSocket, wraps it as +// a yamux server session and hands it to onSession. +// +// onSession runs on the request goroutine, so it must return promptly; the +// session outlives the handler because the upgrade hijacks the connection, and +// closing it is the caller's job. +func PeerHandler(token string, onSession func(peerID string, sess *yamux.Session)) echo.HandlerFunc { + // gorilla's default CheckOrigin already restricts a browser to same-origin + // and lets a header-less client (which every peer is) through, so the + // zero value is what this link wants. + upgrader := websocket.Upgrader{} + + return func(c echo.Context) error { + // Reject before upgrading. Upgrading and then closing would give the + // dialer a WebSocket error in place of an HTTP status, and both the + // route-coverage test and a peer's own retry logic read the status. + if !authorizedPeer(c.Request(), token) { + return echo.NewHTTPError(http.StatusUnauthorized, "unauthorized") + } + + peerID := c.QueryParam("id") + if peerID == "" { + return echo.NewHTTPError(http.StatusBadRequest, "missing peer id") + } + + ws, err := upgrader.Upgrade(c.Response(), c.Request(), nil) + if err != nil { + // Upgrade has already written its own failure to the client. + xlog.Debug("cluster peer link upgrade failed", "peer", peerID, "error", err) + return nil + } + + // Server side of the mux: the dialing peer is the client, so it owns + // the odd stream IDs and this side the even ones. + sess, err := yamux.Server(WebsocketConn(ws), nil, nil) + if err != nil { + xlog.Error("cluster peer link session setup failed", "peer", peerID, "error", err) + _ = ws.Close() + return nil + } + + if onSession == nil { + // Nothing will ever read from this session, so do not leave the + // peer believing it has a live link. + _ = sess.Close() + return nil + } + + xlog.Debug("cluster peer link established", "peer", peerID, "remote", ws.RemoteAddr().String()) + onSession(peerID, sess) + return nil + } +} + +// authorizedPeer compares the request's bearer token with the cluster token in +// constant time, matching the check the worker file-transfer server makes. +// +// Unlike that one, an empty configured token authorizes nobody: this route is +// registered in every deployment, so failing open would publish an +// unauthenticated mux to any caller that can reach the port. +func authorizedPeer(r *http.Request, expected string) bool { + if expected == "" { + return false + } + const prefix = "Bearer " + header := r.Header.Get("Authorization") + if len(header) < len(prefix) || header[:len(prefix)] != prefix { + return false + } + return subtle.ConstantTimeCompare([]byte(header[len(prefix):]), []byte(expected)) == 1 +} diff --git a/core/http/endpoints/cluster/peer_test.go b/core/http/endpoints/cluster/peer_test.go new file mode 100644 index 000000000..b4aa98356 --- /dev/null +++ b/core/http/endpoints/cluster/peer_test.go @@ -0,0 +1,117 @@ +package cluster_test + +import ( + "net/http" + "net/http/httptest" + "strings" + + clusterep "github.com/mudler/LocalAI/core/http/endpoints/cluster" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + "github.com/libp2p/go-yamux/v5" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Peer link handler", func() { + var ( + srv *httptest.Server + sessions chan *yamux.Session + ) + + BeforeEach(func() { + sessions = make(chan *yamux.Session, 1) + e := echo.New() + clusterep.RegisterClusterRoutes(e, "peer-token", func(_ string, s *yamux.Session) { + sessions <- s + }) + srv = httptest.NewServer(e) + DeferCleanup(srv.Close) + }) + + wsURL := func(s *httptest.Server) string { + return "ws" + strings.TrimPrefix(s.URL, "http") + "/api/cluster/peer?id=peer-1" + } + + It("rejects a connection with no token", func() { + _, resp, err := websocket.DefaultDialer.Dial(wsURL(srv), nil) + Expect(err).To(HaveOccurred()) + Expect(resp).ToNot(BeNil()) + Expect(resp.StatusCode).To(Equal(http.StatusUnauthorized)) + }) + + It("rejects a connection with the wrong token", func() { + h := http.Header{} + h.Set("Authorization", "Bearer wrong") + _, resp, err := websocket.DefaultDialer.Dial(wsURL(srv), h) + Expect(err).To(HaveOccurred()) + Expect(resp.StatusCode).To(Equal(http.StatusUnauthorized)) + }) + + It("accepts an authenticated peer and yields a usable yamux session", func() { + h := http.Header{} + h.Set("Authorization", "Bearer peer-token") + conn, _, err := websocket.DefaultDialer.Dial(wsURL(srv), h) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = conn.Close() }) + + var serverSess *yamux.Session + Eventually(sessions, "5s").Should(Receive(&serverSess)) + Expect(serverSess).ToNot(BeNil()) + + // The client wraps its side as a yamux CLIENT and opens a stream; the + // server must accept it. This proves the WebSocket was adapted into a + // stream-oriented conn correctly, which is the part most likely to be + // subtly wrong. + clientSess, err := yamux.Client(clusterep.WebsocketConn(conn), nil, nil) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = clientSess.Close() }) + + go func() { + defer GinkgoRecover() + st, e := clientSess.OpenStream(GinkgoT().Context()) + if e == nil { + _, _ = st.Write([]byte("hello")) + } + }() + + accepted := make(chan []byte, 1) + go func() { + defer GinkgoRecover() + st, e := serverSess.AcceptStream() + if e != nil { + return + } + buf := make([]byte, 5) + if _, e := st.Read(buf); e == nil { + accepted <- buf + } + }() + Eventually(accepted, "10s").Should(Receive(Equal([]byte("hello")))) + }) + + It("reports the peer id it was given", func() { + h := http.Header{} + h.Set("Authorization", "Bearer peer-token") + ids := make(chan string, 1) + e := echo.New() + clusterep.RegisterClusterRoutes(e, "peer-token", func(id string, _ *yamux.Session) { ids <- id }) + s2 := httptest.NewServer(e) + DeferCleanup(s2.Close) + + conn, _, err := websocket.DefaultDialer.Dial(wsURL(s2), h) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = conn.Close() }) + + Eventually(ids, "5s").Should(Receive(Equal("peer-1"))) + }) +}) + +var _ = Describe("Peer link auth prefix", func() { + It("is covered by the alternative-authentication prefix list", func() { + // /api/cluster/ authenticates with the cluster token, not the global + // session middleware, so it must be listed or every peer dial 401s. + Expect(clusterep.AlternativeAuthPrefix).To(Equal("/api/cluster/")) + }) +}) diff --git a/core/http/endpoints/cluster/wsconn.go b/core/http/endpoints/cluster/wsconn.go new file mode 100644 index 000000000..1f1f5306c --- /dev/null +++ b/core/http/endpoints/cluster/wsconn.go @@ -0,0 +1,129 @@ +// SPDX-License-Identifier: MIT + +package cluster + +import ( + "fmt" + "io" + "net" + "sync" + "time" + + "github.com/gorilla/websocket" +) + +// WebsocketConn adapts a gorilla WebSocket into the net.Conn that a yamux +// session drives. +// +// The two disagree about framing: WebSocket delivers whole messages, yamux +// wants an undelimited byte stream. The adapter therefore keeps the reader of +// the message it is part-way through between calls, so a Read whose buffer is +// smaller than the message hands back a prefix now and the rest next time +// instead of dropping the tail. That case is not hypothetical: yamux reads +// through a 4 KiB bufio.Reader while a single stream write can put a much +// larger data frame on the wire in one Write, so any message above the buffer +// size is read in pieces. +// +// The returned conn is safe for one reader and one writer concurrently, which +// is all yamux uses: its recvLoop reads and its sendLoop writes. It is not a +// general-purpose net.Conn. +func WebsocketConn(ws *websocket.Conn) net.Conn { + return &wsConn{ws: ws} +} + +type wsConn struct { + ws *websocket.Conn + + // readMu guards frame, which carries a partially consumed message across + // Read calls. gorilla allows a single concurrent reader, and this keeps + // the adapter to that contract even if a caller reads from two goroutines. + readMu sync.Mutex + frame io.Reader + + // writeMu keeps to gorilla's one-concurrent-writer contract. + writeMu sync.Mutex +} + +func (c *wsConn) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + + c.readMu.Lock() + defer c.readMu.Unlock() + + for { + if c.frame == nil { + messageType, r, err := c.ws.NextReader() + if err != nil { + return 0, translateReadErr(err) + } + // Binary is the only type this link speaks. Skipping an unexpected + // text message would silently desynchronise the yamux framing, so + // it is reported instead. + if messageType != websocket.BinaryMessage { + return 0, fmt.Errorf("cluster: peer link received websocket message type %d, want binary", messageType) + } + c.frame = r + } + + n, err := c.frame.Read(p) + if err == io.EOF { + // End of one message, not end of the stream: drop the reader so + // the next call pulls the next message. Passing io.EOF up would + // end the yamux session at an arbitrary message boundary. + c.frame = nil + err = nil + } + if n > 0 || err != nil { + return n, err + } + // A zero-length message yields nothing to return, and (0, nil) reads + // look like a stalled stream to some callers, so wait for the next one. + } +} + +func (c *wsConn) Write(p []byte) (int, error) { + c.writeMu.Lock() + defer c.writeMu.Unlock() + + if err := c.ws.WriteMessage(websocket.BinaryMessage, p); err != nil { + return 0, err + } + return len(p), nil +} + +// Close drops the underlying network connection without negotiating a +// WebSocket close handshake. yamux has already sent its own go-away by this +// point, and a close frame would need the write lock that a blocked sendLoop +// may still hold. +func (c *wsConn) Close() error { + return c.ws.Close() +} + +func (c *wsConn) LocalAddr() net.Addr { return c.ws.LocalAddr() } +func (c *wsConn) RemoteAddr() net.Addr { return c.ws.RemoteAddr() } + +func (c *wsConn) SetDeadline(t time.Time) error { + if err := c.ws.SetReadDeadline(t); err != nil { + return err + } + return c.ws.SetWriteDeadline(t) +} + +func (c *wsConn) SetReadDeadline(t time.Time) error { return c.ws.SetReadDeadline(t) } +func (c *wsConn) SetWriteDeadline(t time.Time) error { return c.ws.SetWriteDeadline(t) } + +// translateReadErr maps a peer hanging up cleanly onto io.EOF, which is how a +// yamux session recognises a normal ending. Any other close code, and any +// transport error, is passed through so the session reports a real failure. +func translateReadErr(err error) error { + if websocket.IsCloseError(err, + websocket.CloseNormalClosure, + websocket.CloseGoingAway, + websocket.CloseNoStatusReceived, + ) { + return io.EOF + } + return err +} diff --git a/core/http/endpoints/cluster/wsconn_test.go b/core/http/endpoints/cluster/wsconn_test.go new file mode 100644 index 000000000..f3ce4ddf0 --- /dev/null +++ b/core/http/endpoints/cluster/wsconn_test.go @@ -0,0 +1,216 @@ +package cluster_test + +import ( + "bytes" + "crypto/rand" + "io" + "net" + "net/http" + "net/http/httptest" + "strings" + "time" + + clusterep "github.com/mudler/LocalAI/core/http/endpoints/cluster" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v4" + "github.com/libp2p/go-yamux/v5" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// wsPair returns the two ends of one live WebSocket connection. +func wsPair() (clientSide, serverSide *websocket.Conn) { + GinkgoHelper() + + upgrader := websocket.Upgrader{} + accepted := make(chan *websocket.Conn, 1) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ws, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + accepted <- ws + })) + DeferCleanup(srv.Close) + + c, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(srv.URL, "http"), nil) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = c.Close() }) + + var s *websocket.Conn + Eventually(accepted, "5s").Should(Receive(&s)) + DeferCleanup(func() { _ = s.Close() }) + + return c, s +} + +var _ = Describe("WebsocketConn framing", func() { + // The specs below exist because the brief's end-to-end yamux spec cannot + // catch a lost message tail: yamux reads through a 4 KiB bufio.Reader, so + // every small message arrives whole no matter how the adapter behaves. + // These drive the adapter directly with buffers smaller than the message. + + It("returns the rest of a message on the following Read", func() { + clientWS, serverWS := wsPair() + writer := clusterep.WebsocketConn(clientWS) + reader := clusterep.WebsocketConn(serverWS) + + // A lost tail would otherwise park the reassembly below forever; with a + // deadline it fails as a timeout on the read that has nothing left. + Expect(reader.SetReadDeadline(time.Now().Add(10 * time.Second))).To(Succeed()) + + payload := []byte("0123456789abcdefghijklmnopqrstuvwxyz") + n, err := writer.Write(payload) + Expect(err).ToNot(HaveOccurred()) + Expect(n).To(Equal(len(payload))) + + // Deliberately smaller than the message: a naive adapter that starts a + // fresh NextReader on every call drops everything past the first 7 + // bytes, and this reassembly fails. + got := make([]byte, 0, len(payload)) + buf := make([]byte, 7) + for len(got) < len(payload) { + read, err := reader.Read(buf) + Expect(err).ToNot(HaveOccurred()) + Expect(read).To(BeNumerically(">", 0)) + Expect(read).To(BeNumerically("<=", len(buf))) + got = append(got, buf[:read]...) + } + Expect(got).To(Equal(payload)) + }) + + It("streams a message larger than the yamux read buffer without loss or reordering", func() { + clientWS, serverWS := wsPair() + writer := clusterep.WebsocketConn(clientWS) + reader := clusterep.WebsocketConn(serverWS) + + Expect(reader.SetReadDeadline(time.Now().Add(20 * time.Second))).To(Succeed()) + + payload := make([]byte, 256*1024) + _, err := rand.Read(payload) + Expect(err).ToNot(HaveOccurred()) + + go func() { + defer GinkgoRecover() + _, _ = writer.Write(payload) + }() + + // 4096 is the buffer yamux's bufio.Reader actually hands down. + got := make([]byte, len(payload)) + _, err = io.ReadFull(reader, got) + Expect(err).ToNot(HaveOccurred()) + Expect(bytes.Equal(got, payload)).To(BeTrue()) + }) + + It("presents consecutive messages as one continuous byte stream", func() { + clientWS, serverWS := wsPair() + writer := clusterep.WebsocketConn(clientWS) + reader := clusterep.WebsocketConn(serverWS) + + Expect(reader.SetReadDeadline(time.Now().Add(10 * time.Second))).To(Succeed()) + + for _, chunk := range []string{"abc", "", "de", "fghij"} { + _, err := writer.Write([]byte(chunk)) + Expect(err).ToNot(HaveOccurred()) + } + + // A read spanning three messages must be satisfied, and the empty + // message must not surface as a premature (0, nil) or an EOF. + got := make([]byte, 10) + _, err := io.ReadFull(reader, got) + Expect(err).ToNot(HaveOccurred()) + Expect(string(got)).To(Equal("abcdefghij")) + }) + + It("reports a clean peer close as io.EOF", func() { + clientWS, serverWS := wsPair() + reader := clusterep.WebsocketConn(serverWS) + + Expect(clientWS.WriteMessage(websocket.CloseMessage, + websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""))).To(Succeed()) + + _, err := reader.Read(make([]byte, 8)) + Expect(err).To(MatchError(io.EOF)) + }) + + It("refuses a text message rather than desynchronising the stream", func() { + clientWS, serverWS := wsPair() + reader := clusterep.WebsocketConn(serverWS) + + Expect(clientWS.WriteMessage(websocket.TextMessage, []byte("not a frame"))).To(Succeed()) + + _, err := reader.Read(make([]byte, 32)) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("want binary")) + }) + + It("satisfies net.Conn, including the deadlines yamux sets on every write", func() { + clientWS, _ := wsPair() + var conn net.Conn = clusterep.WebsocketConn(clientWS) + + Expect(conn.LocalAddr()).ToNot(BeNil()) + Expect(conn.RemoteAddr()).ToNot(BeNil()) + // yamux's sendLoop calls SetWriteDeadline before every flush, so an + // adapter that dropped the call would let a stalled peer block the + // session forever instead of failing it. + Expect(conn.SetWriteDeadline(time.Now().Add(time.Minute))).To(Succeed()) + Expect(conn.SetReadDeadline(time.Now().Add(time.Minute))).To(Succeed()) + Expect(conn.SetDeadline(time.Time{})).To(Succeed()) + }) +}) + +var _ = Describe("Peer link payloads", func() { + It("carries a payload far larger than one yamux frame end to end", func() { + sessions := make(chan *yamux.Session, 1) + e := echo.New() + clusterep.RegisterClusterRoutes(e, "peer-token", func(_ string, s *yamux.Session) { sessions <- s }) + srv := httptest.NewServer(e) + DeferCleanup(srv.Close) + + h := http.Header{} + h.Set("Authorization", "Bearer peer-token") + conn, _, err := websocket.DefaultDialer.Dial( + "ws"+strings.TrimPrefix(srv.URL, "http")+"/api/cluster/peer?id=peer-1", h) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = conn.Close() }) + + var serverSess *yamux.Session + Eventually(sessions, "5s").Should(Receive(&serverSess)) + + clientSess, err := yamux.Client(clusterep.WebsocketConn(conn), nil, nil) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = clientSess.Close() }) + + payload := make([]byte, 1<<20) + _, err = rand.Read(payload) + Expect(err).ToNot(HaveOccurred()) + + go func() { + defer GinkgoRecover() + st, e := clientSess.OpenStream(GinkgoT().Context()) + if e != nil { + return + } + defer func() { _ = st.Close() }() + _, _ = io.Copy(st, bytes.NewReader(payload)) + }() + + received := make(chan []byte, 1) + go func() { + defer GinkgoRecover() + st, e := serverSess.AcceptStream() + if e != nil { + return + } + buf := make([]byte, len(payload)) + if _, e := io.ReadFull(st, buf); e == nil { + received <- buf + } + }() + + var got []byte + Eventually(received, "30s").Should(Receive(&got)) + Expect(bytes.Equal(got, payload)).To(BeTrue()) + }) +})