// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package spmd import ( "bufio" "bytes" "encoding/binary" "io" "net" "time" "sourcedock.dev/petrbalvin/tensor/internal/base" ) // A frame is one message on a link: a 24-byte header, then the payload. // // offset 0: tag u8 // offset 1: version u8 (frameVersion, a guard against a confused peer) // offset 2: reserved u16 // offset 4: from u32, little-endian, the sending rank // offset 8: dest u32, little-endian, the receiving rank // offset 12: length u64, little-endian, payload bytes that follow // // The header is the only place a length arrives from the wire, and no // read ever allocates a payload longer than the world's message // ceiling: a peer announcing more is an error before a byte of payload // is read. From and dest name the endpoints because a networked world // is a star over rank 0: a frame from rank 2 to rank 5 rides rank 2's // connection in, and rank 5's connection out. const ( frameHeaderLen = 24 frameVersion = 1 ) // The tags the collectives speak. The handshake speaks its own words on // the fresh connection, before any frame. const ( tagBarrier uint8 = 1 tagBarrierAck uint8 = 2 tagBroadcast uint8 = 3 tagScatterHead uint8 = 4 tagScatter uint8 = 5 tagGather uint8 = 6 tagShardsValues uint8 = 8 tagShardsWhole uint8 = 9 tagReduce uint8 = 10 tagHalo uint8 = 11 ) // message is one frame in the world's own terms. type message struct { tag uint8 from int dest int data []byte // pooled marks a payload whose buffer the hub's pump took from the // routed frame pool, so the drain, the frame's single consumer, // returns it there after the write. Every other frame leaves it // false and its buffer belongs to whoever holds the frame. pooled bool } // peer is one rank's link. An in-process world joins the two ranks' // channels directly: the sender writes into the receiver's inbox. A // networked hub pumps each of its links with one reader and one writer; // a networked rank that is not the hub drives its single connection // itself and lets the hub route by dest. type peer struct { rank int // inbox carries frames from this rank that this world reads. In // process it is the direct channel from the peer; at the hub a // reader feeds it; a non-hub rank leaves it nil and reads its one // connection itself. inbox chan message // outbox carries frames to this rank's connection. In process it // is the direct channel to the peer; at the hub a writer drains // it; a non-hub rank leaves it nil and writes its one connection // itself. outbox chan message // gone closes when the peer's world has left; only an in-process // link has one, and only a sender waits on it. gone chan struct{} // conn is this world's end of the peer's networked link. conn *tcpLink // pending holds frames that arrived before the collective asked // for their rank and tag, so no answer ever depends on the // arrival order. Owned by the one goroutine that drives the // world's collectives. pending []message } // tcpLink is a framed TCP connection to one peer rank. type tcpLink struct { conn net.Conn rd *bufio.Reader wr *bufio.Writer } func newTCPLink(conn net.Conn) *tcpLink { return &tcpLink{conn: conn, rd: bufio.NewReader(conn), wr: bufio.NewWriter(conn)} } func (l *tcpLink) readFrame(max int64, deadline time.Time) (message, error) { return l.readFrameInto(max, deadline, nil) } // readFrameInto reads one frame as readFrame does. Take, when not // nil, is asked with the parsed header and the payload length the // header announced for the buffer the payload reads into; a nil // answer allocates the payload as usual. A buffer take supplied // leaves the read with the message's pooled mark on, which commits it // to the single ownership chain the routed frame pool lives on: the // pump that reads the frame hands it through one outbox channel to // the one drain, which returns the buffer after the write. func (l *tcpLink) readFrameInto(max int64, deadline time.Time, take func(message, int64) []byte) (message, error) { if err := l.conn.SetReadDeadline(deadline); err != nil { return message{}, err } var head [frameHeaderLen]byte if _, err := io.ReadFull(l.rd, head[:]); err != nil { return message{}, err } if head[1] != frameVersion { return message{}, base.Errf("spmd: frame version %d from the link is not %d", head[1], frameVersion) } m := message{ tag: head[0], from: int(binary.LittleEndian.Uint32(head[4:])), dest: int(binary.LittleEndian.Uint32(head[8:])), } length := int64(binary.LittleEndian.Uint64(head[12:])) if length < 0 || length > max { return message{}, base.Errf("spmd: the link announces a %d byte payload beyond the %d byte ceiling", length, max) } if take != nil { if buf := take(m, length); buf != nil { m.data, m.pooled = buf, true } } if m.data == nil { m.data = make([]byte, length) } if _, err := io.ReadFull(l.rd, m.data); err != nil { return message{}, err } return m, nil } func (l *tcpLink) writeFrame(m message, deadline time.Time) error { if err := l.conn.SetWriteDeadline(deadline); err != nil { return err } var head [frameHeaderLen]byte head[0] = m.tag head[1] = frameVersion binary.LittleEndian.PutUint32(head[4:], uint32(m.from)) binary.LittleEndian.PutUint32(head[8:], uint32(m.dest)) binary.LittleEndian.PutUint64(head[12:], uint64(len(m.data))) if _, err := l.wr.Write(head[:]); err != nil { return err } if _, err := l.wr.Write(m.data); err != nil { return err } return l.wr.Flush() } // The handshake words on a fresh TCP connection: the joining rank // sends hello, the listening rank answers with the world's size and // the joining rank's place in it. const handshakeLen = 16 var handshakeMagic = [4]byte{'T', 'S', 'P', 'M'} // sendHello is the joining side's word. func sendHello(conn net.Conn) error { var hello [handshakeLen]byte copy(hello[0:4], handshakeMagic[:]) hello[4] = frameVersion _, err := conn.Write(hello[:]) return err } // readHello is the listening side's read of it. func readHello(conn net.Conn, deadline time.Time) error { if err := conn.SetReadDeadline(deadline); err != nil { return err } var hello [handshakeLen]byte if _, err := io.ReadFull(conn, hello[:]); err != nil { return err } if !bytes.Equal(hello[0:4], handshakeMagic[:]) { return base.Errf("spmd: the joining connection did not say the spmd magic") } if hello[4] != frameVersion { return base.Errf("spmd: joining protocol version %d is not %d", hello[4], frameVersion) } return nil } // sendWelcome is the listening side's answer: the world size and the // rank the joining connection carries. func sendWelcome(conn net.Conn, size, rank int) error { var welcome [handshakeLen]byte copy(welcome[0:4], handshakeMagic[:]) welcome[4] = frameVersion binary.LittleEndian.PutUint32(welcome[8:], uint32(size)) binary.LittleEndian.PutUint32(welcome[12:], uint32(rank)) _, err := conn.Write(welcome[:]) return err } // readWelcome is the joining side's read of it. func readWelcome(conn net.Conn, deadline time.Time) (size, rank int, err error) { if err := conn.SetReadDeadline(deadline); err != nil { return 0, 0, err } var welcome [handshakeLen]byte if _, err := io.ReadFull(conn, welcome[:]); err != nil { return 0, 0, err } if !bytes.Equal(welcome[0:4], handshakeMagic[:]) { return 0, 0, base.Errf("spmd: the listener did not answer with the spmd magic") } if welcome[4] != frameVersion { return 0, 0, base.Errf("spmd: listener protocol version %d is not %d", welcome[4], frameVersion) } size = int(binary.LittleEndian.Uint32(welcome[8:])) rank = int(binary.LittleEndian.Uint32(welcome[12:])) if size < 1 || rank < 0 || rank >= size { return 0, 0, base.Errf("spmd: the listener answered with an impossible size %d and rank %d", size, rank) } return size, rank, nil }