Files
tensor/spmd/transport.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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
}