238 lines
7.8 KiB
Go
238 lines
7.8 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
|
}
|