feat: full NFSv4.2 server and client in pure Go
Test / test (push) Successful in 2m4s
Release / gates (push) Successful in 2m5s
Release / build (amd64, freebsd) (push) Successful in 1m27s
Release / build (amd64, linux) (push) Successful in 1m22s
Release / build (amd64, netbsd) (push) Successful in 1m19s
Release / build (amd64, openbsd) (push) Successful in 1m20s
Release / build (arm64, darwin) (push) Successful in 1m21s
Release / build (arm64, freebsd) (push) Successful in 1m26s
Release / build (arm64, linux) (push) Successful in 1m25s
Release / build (arm64, netbsd) (push) Successful in 1m31s
Release / build (arm64, openbsd) (push) Successful in 1m27s
Release / build (loong64, linux) (push) Successful in 1m37s
Release / build (riscv64, linux) (push) Successful in 1m21s
Release / release (push) Successful in 40s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-21 18:51:17 +02:00
commit a9b8039ef7
153 changed files with 34403 additions and 0 deletions
+55
View File
@@ -0,0 +1,55 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"bytes"
"testing"
)
// FuzzReadRecord feeds arbitrary fragment streams into the reassembler:
// no input may panic, allocate without bound or return anything but a
// record, io.EOF, ErrRecordTooLarge or a truncation error.
func FuzzReadRecord(f *testing.F) {
f.Add([]byte{0x80, 0, 0, 4, 'a', 'b', 'c', 'd'})
f.Add([]byte{0, 0, 0, 4, 'a', 'b', 'c', 'd', 0x80, 0, 0, 0})
f.Add([]byte{0xff, 0xff, 0xff, 0xff})
f.Add([]byte{0, 0, 0, 1})
f.Fuzz(func(t *testing.T, data []byte) {
rec, err := ReadRecord(bytes.NewReader(data), 1<<16)
if err == nil && len(rec) > 1<<16 {
t.Fatalf("a record of %d bytes against a limit of %d", len(rec), 1<<16)
}
})
}
// FuzzDecodeMessage feeds arbitrary records through the call and reply
// decoders: no input may panic, and every failure arrives as an error.
func FuzzDecodeMessage(f *testing.F) {
call, err := AppendCall(nil, Call{XID: 1, Program: 100003, Version: 4, Procedure: 1,
Cred: Auth{Flavor: FlavorSys, Body: []byte{1, 2}}})
if err != nil {
f.Fatal(err)
}
reply, err := AppendAcceptedReply(nil, 1, AuthNull, AcceptSuccess, Mismatch{})
if err != nil {
f.Fatal(err)
}
reply = append(reply, 0, 0, 0, 42)
f.Add(call)
f.Add(reply)
f.Add(AppendRejectedReply(nil, 2, AuthBadVerf))
f.Add([]byte{0, 0, 0, 1, 0, 0, 0, 1})
f.Add([]byte{0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 2})
f.Fuzz(func(t *testing.T, data []byte) {
// The property under test is that none of these panics; every
// malformed input must arrive as an ordinary error.
_, _, _ = DecodeCall(data)
_, _ = DecodeReply(data)
_, _, _ = PeekHeader(data)
_, _ = DecodeAuthSysBody(data)
_, _ = DecodeGSSCred(data)
_, _, _, _, _, _ = DecodeGSSInitRes(data)
})
}
+130
View File
@@ -0,0 +1,130 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// The RPCSEC_GSS credential of RFC 2203 as refined by RFC 5403: the
// credential body, the context establishment procedures and result, and
// the service levels none, integrity and privacy.
package rpc
import (
"sourcedock.dev/petrbalvin/nfs/internal/xdr"
)
// The RPCSEC_GSS authentication flavor.
const FlavorGSS = 6
// GSSVersion1 is the credential version of RFC 2203.
const GSSVersion1 = 1
// Credential procedures of the gss_proc union.
const (
GSSProcData = 0
GSSProcInit = 1
GSSProcContinue = 2
GSSProcDestroy = 3
)
// Service levels of rpc_gss_svc_t.
const (
SvcNone = 1
SvcIntegrity = 2
SvcPrivacy = 3
)
// A GSSCred is the decoded version one credential body: the version,
// the procedure, the sequence number, the service and the context
// handle, in the order RFC 2203 section 5.2.1 fixes for every
// procedure.
type GSSCred struct {
Version uint32
Proc uint32
Seq uint32
Service uint32
Handle []byte
}
// AppendGSSCred encodes the version one credential body. The context
// token of the control procedures travels in the procedure arguments,
// never in the credential.
func AppendGSSCred(b []byte, proc, seq, service uint32, handle []byte) []byte {
b = xdr.AppendUint32(b, GSSVersion1)
b = xdr.AppendUint32(b, proc)
b = xdr.AppendUint32(b, seq)
b = xdr.AppendUint32(b, service)
return xdr.AppendVarOpaque(b, handle)
}
// DecodeGSSCred decodes the version one credential body.
func DecodeGSSCred(body []byte) (GSSCred, error) {
d := xdr.NewDecoder(body)
var c GSSCred
var err error
if c.Version, err = d.Uint32(); err != nil {
return c, err
}
if c.Version != GSSVersion1 {
return c, ErrGSSCred
}
if c.Proc, err = d.Uint32(); err != nil {
return c, err
}
switch c.Proc {
case GSSProcData, GSSProcInit, GSSProcContinue, GSSProcDestroy:
default:
return c, ErrGSSCred
}
if c.Seq, err = d.Uint32(); err != nil {
return c, err
}
if c.Service, err = d.Uint32(); err != nil {
return c, err
}
c.Handle, err = d.VarOpaque()
return c, err
}
// ErrGSSCred marks a malformed version one credential: an unknown
// procedure or a version other than one.
var ErrGSSCred = &gssError{"malformed rpcsec gss version one credential"}
type gssError struct{ s string }
func (e *gssError) Error() string { return "rpc: " + e.s }
// AppendGSSInitRes encodes the RPCSEC_GSS_INIT result: the handle the
// server assigns, the major and minor status, the sequence window and
// the reply token.
func AppendGSSInitRes(b []byte, handle []byte, major, minor, window uint32, token []byte) []byte {
b = xdr.AppendVarOpaque(b, handle)
b = xdr.AppendUint32(b, major)
b = xdr.AppendUint32(b, minor)
b = xdr.AppendUint32(b, window)
return xdr.AppendVarOpaque(b, token)
}
// DecodeGSSInitRes decodes the RPCSEC_GSS_INIT result.
func DecodeGSSInitRes(payload []byte) (handle []byte, major, minor, window uint32, token []byte, err error) {
d := xdr.NewDecoder(payload)
if handle, err = d.VarOpaque(); err != nil {
return
}
if major, err = d.Uint32(); err != nil {
return
}
if minor, err = d.Uint32(); err != nil {
return
}
if window, err = d.Uint32(); err != nil {
return
}
token, err = d.VarOpaque()
return
}
// The AUTH_TLS authentication flavor of RFC 9289 and the STARTTLS
// token the server answers the probe with.
const (
FlavorTLS = 7
StarttlsToken = "STARTTLS"
)
+39
View File
@@ -0,0 +1,39 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"testing"
)
// The RPCSEC_GSS_INIT result round trips: handle, statuses, window and
// the reply token.
func TestGSSInitResRoundTrip(t *testing.T) {
res := AppendGSSInitRes(nil, []byte("handle-2"), 0, 1, 32, []byte("ap-rep"))
handle, major, minor, window, token, err := DecodeGSSInitRes(res)
if err != nil || string(handle) != "handle-2" || major != 0 || minor != 1 ||
window != 32 || string(token) != "ap-rep" {
t.Fatalf("res %q %d %d %d %q %v", handle, major, minor, window, token, err)
}
// The empty refusal answers an empty handle and no token.
res = AppendGSSInitRes(nil, nil, 16<<16, 1, 0, nil)
if handle, major, _, _, token, err := DecodeGSSInitRes(res); err != nil ||
len(handle) != 0 || major != 16<<16 || len(token) != 0 {
t.Fatalf("refusal %q %d %q %v", handle, major, token, err)
}
// The error text of the credential sentinel is stable.
if ErrGSSCred.Error() == "" {
t.Fatal("empty error text")
}
// A truncated init result is refused word by word.
for n := 0; n < len(res); n++ {
if _, _, _, _, _, err := DecodeGSSInitRes(res[:n]); err == nil {
t.Fatalf("a %d byte prefix decoded cleanly", n)
}
}
// The NULL procedure constant the control exchange rides on.
if ProcedureNull != 0 {
t.Fatalf("null procedure %d", ProcedureNull)
}
}
+106
View File
@@ -0,0 +1,106 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"bytes"
"encoding/hex"
"testing"
"sourcedock.dev/petrbalvin/nfs/internal/krb5"
)
// The version one credential pins the exact wire order RFC 2203
// section 5.2.1 fixes: version, procedure, sequence, service, handle.
func TestGSSCredWireOrder(t *testing.T) {
data := AppendGSSCred(nil, GSSProcData, 7, SvcIntegrity, []byte("handle-1"))
want, err := hex.DecodeString("00000001" + "00000000" + "00000007" + "00000002" +
"00000008" + "68616e646c652d31")
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(data, want) {
t.Fatalf("cred bytes %x, want %x", data, want)
}
cred, err := DecodeGSSCred(data)
if err != nil {
t.Fatal(err)
}
if cred.Proc != GSSProcData || cred.Version != GSSVersion1 || cred.Service != SvcIntegrity ||
string(cred.Handle) != "handle-1" || cred.Seq != 7 {
t.Fatalf("cred %+v", cred)
}
// The control credentials a conformant peer sends: the same order,
// an empty handle and the token in the procedure arguments.
init := AppendGSSCred(nil, GSSProcInit, 0, 0, nil)
if cred, err = DecodeGSSCred(init); err != nil || cred.Proc != GSSProcInit || len(cred.Handle) != 0 {
t.Fatalf("init cred %+v %v", cred, err)
}
cont := AppendGSSCred(nil, GSSProcContinue, 0, 0, nil)
if cred, err = DecodeGSSCred(cont); err != nil || cred.Proc != GSSProcContinue {
t.Fatalf("continue cred %+v %v", cred, err)
}
dest := AppendGSSCred(nil, GSSProcDestroy, 8, SvcIntegrity, []byte("handle-1"))
if cred, err = DecodeGSSCred(dest); err != nil || cred.Proc != GSSProcDestroy || string(cred.Handle) != "handle-1" {
t.Fatalf("destroy cred %+v %v", cred, err)
}
// A credential that opens with another version or names an unknown
// procedure is refused.
bad := append([]byte{}, dest...)
bad[0] = 3 // a version three body belongs to the GSSv3 decoder
if _, err = DecodeGSSCred(bad); err == nil {
t.Fatal("version three accepted by the version one decoder")
}
bad[0] = GSSVersion1
bad[4] = 9
if _, err = DecodeGSSCred(bad); err == nil {
t.Fatal("unknown procedure accepted")
}
}
// A full DATA call: the header with an empty verifier is checksummed,
// then the call is re-encoded with the MIC as the verifier.
func TestGSSCredAndVerf(t *testing.T) {
data := AppendGSSCred(nil, GSSProcData, 7, SvcIntegrity, []byte("handle-1"))
call := Call{XID: 99, Program: 100003, Version: 4, Procedure: 1,
Cred: Auth{Flavor: FlavorGSS, Body: data}}
prefix, err := AppendCall(nil, call)
if err != nil {
t.Fatal(err)
}
// Two halves of one established context share the session key.
clientCtx := &krb5.Context{Etype: krb5.EtypeAES256, Key: make([]byte, 32)}
serverCtx := &krb5.Context{Etype: krb5.EtypeAES256, Key: clientCtx.Key, Accepting: true}
mic, err := clientCtx.GetMIC(prefix)
if err != nil {
t.Fatal(err)
}
call.Verifier = Auth{Flavor: FlavorGSS, Body: mic}
full, err := AppendCall(nil, call)
if err != nil {
t.Fatal(err)
}
decoded, args, err := DecodeCall(full)
if err != nil {
t.Fatal(err)
}
if decoded.Verifier.Flavor != FlavorGSS {
t.Fatalf("verifier flavor %d", decoded.Verifier.Flavor)
}
// The receiver re-derives the prefix by re-encoding with an empty
// verifier and verifies the MIC over it.
again, err := AppendCall(nil, Call{XID: decoded.XID, Program: decoded.Program,
Version: decoded.Version, Procedure: decoded.Procedure, Cred: decoded.Cred})
if err != nil {
t.Fatal(err)
}
if err := serverCtx.VerifyMIC(again, decoded.Verifier.Body); err != nil {
t.Fatalf("verifier MIC: %v", err)
}
if len(args) != 0 {
t.Fatal("stray arguments after the header")
}
}
+298
View File
@@ -0,0 +1,298 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// The RPCSEC_GSSv3 structures of RFC 7861: the version three credential,
// the CREATE and LIST control procedures and the assertion payloads.
package rpc
import (
"errors"
"sourcedock.dev/petrbalvin/nfs/internal/xdr"
)
// Control procedure numbers of the rpc_gss_proc_t enumeration, RFC 7861
// section 5.1.
const (
GSSProcCreate = 5
GSSProcList = 6
)
// The credential version that carries the new control procedures.
const GSSVersion3 = 3
// Assertion types of the rgss3_assertion_type enumeration.
const (
AssertionLabel = 0
AssertionPrivs = 1
)
// ErrGSSv3 marks a malformed version three control message.
var ErrGSSv3 = errors.New("rpc: malformed rpcsec gssv3 message")
// A GSSv3Cred is the version three credential: the version field rides
// in front of the version one shape, RFC 7861 section 5.1.
type GSSv3Cred struct {
Proc uint32
Seq uint32
Service uint32
Handle []byte
}
// AppendGSSv3Cred encodes the version three credential body.
func AppendGSSv3Cred(b []byte, proc, seq, service uint32, handle []byte) []byte {
b = xdr.AppendUint32(b, GSSVersion3)
b = xdr.AppendUint32(b, proc)
b = xdr.AppendUint32(b, seq)
b = xdr.AppendUint32(b, service)
return xdr.AppendVarOpaque(b, handle)
}
// DecodeGSSv3Cred decodes the version three credential body: the
// leading version field is checked and skipped before the credential
// proper.
func DecodeGSSv3Cred(body []byte) (GSSv3Cred, error) {
d := xdr.NewDecoder(body)
var c GSSv3Cred
var err error
if c.Proc, err = d.Uint32(); err != nil {
return c, err
}
if c.Proc != GSSVersion3 {
return c, ErrGSSv3
}
if c.Proc, err = d.Uint32(); err != nil {
return c, err
}
if c.Seq, err = d.Uint32(); err != nil {
return c, err
}
if c.Service, err = d.Uint32(); err != nil {
return c, err
}
c.Handle, err = d.VarOpaque()
return c, err
}
// A Label is the rgss3_label assertion: the label format specifier and
// the opaque label payload.
type Label struct {
LfsId uint32
PiId uint32
Bytes []byte
}
// Privs is the rgss3_privs structured privilege: who grants what.
type Privs struct {
Who string
Grant string
Bytes []byte
}
// An Assertion is one rgss3_assertion_u union member.
type Assertion struct {
Type uint32
Label Label
Privs Privs
Ext []byte
}
// appendLabel and appendPrivs encode the assertion payloads.
func appendLabel(b []byte, l Label) []byte {
b = xdr.AppendUint32(b, l.LfsId)
b = xdr.AppendUint32(b, l.PiId)
return xdr.AppendVarOpaque(b, l.Bytes)
}
func appendPrivs(b []byte, p Privs) []byte {
b = xdr.AppendString(b, p.Who)
b = xdr.AppendString(b, p.Grant)
return xdr.AppendVarOpaque(b, p.Bytes)
}
// AppendAssertion encodes one rgss3_assertion_u union.
func AppendAssertion(b []byte, a Assertion) []byte {
b = xdr.AppendUint32(b, a.Type)
switch a.Type {
case AssertionLabel:
return appendLabel(b, a.Label)
case AssertionPrivs:
return appendPrivs(b, a.Privs)
default:
return xdr.AppendVarOpaque(b, a.Ext)
}
}
// DecodeAssertion decodes one rgss3_assertion_u union.
func DecodeAssertion(d *xdr.Decoder) (Assertion, error) {
var a Assertion
var err error
if a.Type, err = d.Uint32(); err != nil {
return a, err
}
switch a.Type {
case AssertionLabel:
if a.Label.LfsId, err = d.Uint32(); err != nil {
return a, err
}
if a.Label.PiId, err = d.Uint32(); err != nil {
return a, err
}
a.Label.Bytes, err = d.VarOpaque()
return a, err
case AssertionPrivs:
if a.Privs.Who, err = d.String(); err != nil {
return a, err
}
if a.Privs.Grant, err = d.String(); err != nil {
return a, err
}
a.Privs.Bytes, err = d.VarOpaque()
return a, err
default:
a.Ext, err = d.VarOpaque()
return a, err
}
}
// A MpAuth is the rgss3_gss_mp_auth multi-principal authentication
// payload: the inner context handle and a MIC of the RPC header made
// under the inner context.
type MpAuth struct {
InnerHandle []byte
HeaderMic []byte
}
// appendOptional encodes an XDR optional: the presence flag and the
// payload.
func appendOptional(b []byte, present bool, enc func([]byte) []byte) []byte {
b = xdr.AppendBool(b, present)
if present {
return enc(b)
}
return b
}
// AppendCreateArgs encodes the rgss3_create_args call data.
func AppendCreateArgs(b []byte, mpAuth *MpAuth, chanBinding []byte, assertions []Assertion) []byte {
b = appendOptional(b, mpAuth != nil, func(x []byte) []byte {
x = xdr.AppendVarOpaque(x, mpAuth.InnerHandle)
return xdr.AppendVarOpaque(x, mpAuth.HeaderMic)
})
b = appendOptional(b, chanBinding != nil, func(x []byte) []byte {
return xdr.AppendVarOpaque(x, chanBinding)
})
b = xdr.AppendUint32(b, uint32(len(assertions)))
for _, a := range assertions {
b = AppendAssertion(b, a)
}
return b
}
// DecodeCreateArgs decodes the rgss3_create_args call data.
func DecodeCreateArgs(payload []byte) (mpAuth *MpAuth, chanBinding []byte, assertions []Assertion, err error) {
d := xdr.NewDecoder(payload)
var present bool
if present, err = d.Bool(); err != nil {
return
}
if present {
mpAuth = &MpAuth{}
if mpAuth.InnerHandle, err = d.VarOpaque(); err != nil {
return
}
if mpAuth.HeaderMic, err = d.VarOpaque(); err != nil {
return
}
}
if present, err = d.Bool(); err != nil {
return
}
if present {
if chanBinding, err = d.VarOpaque(); err != nil {
return
}
}
var n uint32
if n, err = d.Uint32(); err != nil {
return
}
for i := uint32(0); i < n; i++ {
var a Assertion
if a, err = DecodeAssertion(d); err != nil {
return
}
assertions = append(assertions, a)
}
return
}
// AppendCreateRes encodes the rgss3_create_res reply: the child handle,
// the mirrored optional fields and the granted assertions in order.
func AppendCreateRes(b []byte, handle []byte, mpAuth *MpAuth, chanBinding []byte, assertions []Assertion) []byte {
b = xdr.AppendVarOpaque(b, handle)
b = appendOptional(b, mpAuth != nil, func(x []byte) []byte {
x = xdr.AppendVarOpaque(x, mpAuth.InnerHandle)
return xdr.AppendVarOpaque(x, mpAuth.HeaderMic)
})
b = appendOptional(b, chanBinding != nil, func(x []byte) []byte {
return xdr.AppendVarOpaque(x, chanBinding)
})
b = xdr.AppendUint32(b, uint32(len(assertions)))
for _, a := range assertions {
b = AppendAssertion(b, a)
}
return b
}
// DecodeCreateRes decodes the rgss3_create_res reply.
func DecodeCreateRes(payload []byte) (handle []byte, mpAuth *MpAuth, chanBinding []byte, assertions []Assertion, err error) {
d := xdr.NewDecoder(payload)
if handle, err = d.VarOpaque(); err != nil {
return
}
var present bool
if present, err = d.Bool(); err != nil {
return
}
if present {
mpAuth = &MpAuth{}
if mpAuth.InnerHandle, err = d.VarOpaque(); err != nil {
return
}
if mpAuth.HeaderMic, err = d.VarOpaque(); err != nil {
return
}
}
if present, err = d.Bool(); err != nil {
return
}
if present {
if chanBinding, err = d.VarOpaque(); err != nil {
return
}
}
var n uint32
if n, err = d.Uint32(); err != nil {
return
}
for i := uint32(0); i < n; i++ {
var a Assertion
if a, err = DecodeAssertion(d); err != nil {
return
}
assertions = append(assertions, a)
}
return
}
// AppendListRes encodes the RPCSEC_GSS_LIST reply: the supported
// assertion types, RFC 7861 section 5.3.
func AppendListRes(b []byte, types []uint32) []byte {
b = xdr.AppendUint32(b, uint32(len(types)))
for _, t := range types {
b = xdr.AppendUint32(b, t)
}
return b
}
+78
View File
@@ -0,0 +1,78 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"testing"
"sourcedock.dev/petrbalvin/nfs/internal/xdr"
)
// The RPCSEC_GSSv3 credential round trips with the version field in
// front, RFC 7861 section 5.1.
func TestGSSv3CredShape(t *testing.T) {
body := AppendGSSv3Cred(nil, GSSProcCreate, 7, SvcPrivacy, []byte("parent"))
cred, err := DecodeGSSv3Cred(body)
if err != nil {
t.Fatal(err)
}
if cred.Proc != GSSProcCreate || cred.Seq != 7 || cred.Service != SvcPrivacy ||
string(cred.Handle) != "parent" {
t.Fatalf("cred %+v", cred)
}
if _, err := DecodeGSSv3Cred(append([]byte{0, 0, 0, 1}, body[4:]...)); err != ErrGSSv3 {
t.Fatalf("version one body accepted: %v", err)
}
}
// The create arguments and reply round trip through their codecs,
// including the optional fields and the assertion union.
func TestGSSv3CreateShapes(t *testing.T) {
args := AppendCreateArgs(nil,
&MpAuth{InnerHandle: []byte("inner"), HeaderMic: []byte("mic")},
[]byte("binding"),
[]Assertion{
{Type: AssertionLabel, Label: Label{LfsId: 1, PiId: 2, Bytes: []byte("secret")}},
{Type: AssertionPrivs, Privs: Privs{Who: "petr", Grant: "admin", Bytes: []byte("x")}},
{Type: 9, Ext: []byte("ext")},
})
mp, bind, assertions, err := DecodeCreateArgs(args)
if err != nil {
t.Fatal(err)
}
if mp == nil || string(mp.InnerHandle) != "inner" || string(mp.HeaderMic) != "mic" {
t.Fatalf("mp auth %+v", mp)
}
if string(bind) != "binding" {
t.Fatalf("binding %q", bind)
}
if len(assertions) != 3 || assertions[0].Label.LfsId != 1 ||
assertions[1].Privs.Who != "petr" || assertions[2].Ext == nil {
t.Fatalf("assertions %+v", assertions)
}
// The reply mirrors the shape with the child handle.
res := AppendCreateRes(nil, []byte("child"), nil, nil,
[]Assertion{{Type: AssertionLabel, Label: Label{LfsId: 1, Bytes: []byte("secret")}}})
handle, _, _, granted, err := DecodeCreateRes(res)
if err != nil {
t.Fatal(err)
}
if string(handle) != "child" || len(granted) != 1 {
t.Fatalf("res handle %q granted %d", handle, len(granted))
}
// The list reply carries the type array.
list := AppendListRes(nil, []uint32{AssertionLabel, AssertionPrivs})
d := xdr.NewDecoder(list)
n, err := d.Uint32()
if err != nil || n != 2 {
t.Fatalf("list count %d: %v", n, err)
}
t1, _ := d.Uint32()
t2, _ := d.Uint32()
if t1 != AssertionLabel || t2 != AssertionPrivs {
t.Fatalf("list types %d %d", t1, t2)
}
}
+375
View File
@@ -0,0 +1,375 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"errors"
"fmt"
"sourcedock.dev/petrbalvin/nfs/internal/xdr"
)
// Version is the ONC RPC protocol version, fixed by RFC 5531.
const Version = 2
// Message types.
const (
MsgCall = 0
MsgReply = 1
)
// ProcedureNull is the NULL procedure every program reserves for
// control exchanges such as the RPCSEC_GSS context procedures of
// RFC 2203 section 5.1.3 and the STARTTLS probe of RFC 9289.
const ProcedureNull = 0
// Authentication flavors.
const (
FlavorNone = 0
FlavorSys = 1
FlavorShort = 2
)
// maxAuthBody is the largest opaque credential or verifier body the
// standard allows.
const maxAuthBody = 400
// ErrBadMessage is returned for a record that is not a well formed ONC RPC
// message.
var ErrBadMessage = errors.New("rpc: malformed message")
// An Auth is an opaque_auth: a flavor and its flavor defined body.
type Auth struct {
Flavor uint32
Body []byte
}
// AuthNull is the credential and verifier that carry nothing.
var AuthNull = Auth{}
func appendAuth(b []byte, a Auth) ([]byte, error) {
if len(a.Body) > maxAuthBody {
return nil, fmt.Errorf("rpc: auth body of %d bytes exceeds %d", len(a.Body), maxAuthBody)
}
b = xdr.AppendUint32(b, a.Flavor)
b = xdr.AppendVarOpaque(b, a.Body)
return b, nil
}
func decodeAuth(d *xdr.Decoder) (Auth, error) {
flavor, err := d.Uint32()
if err != nil {
return Auth{}, err
}
body, err := d.VarOpaque()
if err != nil {
return Auth{}, err
}
if len(body) > maxAuthBody {
return Auth{}, fmt.Errorf("rpc: auth body of %d bytes exceeds %d", len(body), maxAuthBody)
}
return Auth{Flavor: flavor, Body: body}, nil
}
// A Call is the header of an ONC RPC call. The procedure arguments follow
// the header in the same record.
type Call struct {
XID uint32
Program, Version, Procedure uint32
Cred, Verifier Auth
}
// AppendCall appends the call header to b. The caller appends the
// procedure arguments afterwards. Call.Version carries the program version
// the call targets; the ONC RPC protocol version is fixed at 2.
func AppendCall(b []byte, c Call) ([]byte, error) {
b = xdr.AppendUint32(b, c.XID)
b = xdr.AppendUint32(b, MsgCall)
b = xdr.AppendUint32(b, Version)
b = xdr.AppendUint32(b, c.Program)
b = xdr.AppendUint32(b, c.Version)
b = xdr.AppendUint32(b, c.Procedure)
var err error
if b, err = appendAuth(b, c.Cred); err != nil {
return nil, fmt.Errorf("rpc: credential: %w", err)
}
if b, err = appendAuth(b, c.Verifier); err != nil {
return nil, fmt.Errorf("rpc: verifier: %w", err)
}
return b, nil
}
// DecodeCall splits a record into its call header and the bytes that hold
// the procedure arguments.
func DecodeCall(record []byte) (Call, []byte, error) {
d := xdr.NewDecoder(record)
xid, err := d.Uint32()
if err != nil {
return Call{}, nil, ErrBadMessage
}
mtype, err := d.Uint32()
if err != nil {
return Call{}, nil, ErrBadMessage
}
if mtype != MsgCall {
return Call{}, nil, fmt.Errorf("%w: message type %d is not a call", ErrBadMessage, mtype)
}
var c Call
c.XID = xid
rpcvers, err := d.Uint32()
if err != nil {
return Call{}, nil, ErrBadMessage
}
if rpcvers != Version {
return Call{}, nil, fmt.Errorf("%w: rpc version %d, want %d", ErrBadMessage, rpcvers, Version)
}
if c.Program, err = d.Uint32(); err != nil {
return Call{}, nil, ErrBadMessage
}
if c.Version, err = d.Uint32(); err != nil {
return Call{}, nil, ErrBadMessage
}
if c.Procedure, err = d.Uint32(); err != nil {
return Call{}, nil, ErrBadMessage
}
if c.Cred, err = decodeAuth(d); err != nil {
return Call{}, nil, ErrBadMessage
}
if c.Verifier, err = decodeAuth(d); err != nil {
return Call{}, nil, ErrBadMessage
}
return c, record[len(record)-d.Remaining():], nil
}
// Accept statuses carried by an accepted reply.
const (
AcceptSuccess = 0
AcceptProgUnavail = 1
AcceptProgMismatch = 2
AcceptProcUnavail = 3
AcceptGarbageArgs = 4
AcceptSystemErr = 5
)
// Reject statuses carried by a rejected reply.
const (
RejectRPCMismatch = 0
RejectAuthError = 1
)
// Auth statistics of a rejected reply, RFC 5531 section 9 and
// RFC 2203 section 5.2.3.
const (
AuthBadCred = 1
AuthRejectedCred = 2
AuthBadVerf = 3
AuthRejectedVerf = 4
AuthTooWeak = 5
AuthInvalidResp = 6
AuthFailed = 7
AuthDenied = 8
AuthGSSCredProb = 14
AuthGSSCtxProb = 15
)
// A Mismatch carries the program version range a server accepts, sent when
// a call names a version the server does not.
type Mismatch struct {
Low, High uint32
}
// AppendAcceptedReply appends an accepted reply header. On success the
// caller appends the procedure results afterwards; on any other status the
// header carries the whole reply, and mismatch is read only when the status
// is AcceptProgMismatch.
func AppendAcceptedReply(b []byte, xid uint32, verifier Auth, status uint32, mismatch Mismatch) ([]byte, error) {
b = xdr.AppendUint32(b, xid)
b = xdr.AppendUint32(b, MsgReply)
b = xdr.AppendUint32(b, 0) // accepted
var err error
if b, err = appendAuth(b, verifier); err != nil {
return nil, fmt.Errorf("rpc: verifier: %w", err)
}
b = xdr.AppendUint32(b, status)
switch status {
case AcceptSuccess, AcceptProgUnavail, AcceptProcUnavail, AcceptGarbageArgs, AcceptSystemErr:
case AcceptProgMismatch:
b = xdr.AppendUint32(b, mismatch.Low)
b = xdr.AppendUint32(b, mismatch.High)
default:
return nil, fmt.Errorf("rpc: unknown accept status %d", status)
}
return b, nil
}
// A Reply is a decoded reply header. Body holds the procedure results
// when the status is AcceptSuccess. A reply the server rejected, with
// MSG_DENIED, carries Rejected set: Status stays zero and is
// meaningless there, AuthStat holds the auth error when the rejection
// is one, and Mismatch holds the version range on an RPC_MISMATCH.
type Reply struct {
XID uint32
Rejected bool
Status uint32
AuthStat uint32
Mismatch Mismatch
Verifier Auth
Body []byte
}
// DecodeReply splits a record into its reply header and the result bytes.
func DecodeReply(record []byte) (Reply, error) {
d := xdr.NewDecoder(record)
xid, err := d.Uint32()
if err != nil {
return Reply{}, ErrBadMessage
}
mtype, err := d.Uint32()
if err != nil {
return Reply{}, ErrBadMessage
}
if mtype != MsgReply {
return Reply{}, fmt.Errorf("%w: message type %d is not a reply", ErrBadMessage, mtype)
}
stat, err := d.Uint32()
if err != nil {
return Reply{}, ErrBadMessage
}
switch stat {
case 0: // accepted
r := Reply{XID: xid}
if r.Verifier, err = decodeAuth(d); err != nil {
return Reply{}, ErrBadMessage
}
if r.Status, err = d.Uint32(); err != nil {
return Reply{}, ErrBadMessage
}
switch r.Status {
case AcceptSuccess:
r.Body = record[len(record)-d.Remaining():]
return r, nil
case AcceptProgMismatch:
if r.Mismatch.Low, err = d.Uint32(); err != nil {
return Reply{}, ErrBadMessage
}
if r.Mismatch.High, err = d.Uint32(); err != nil {
return Reply{}, ErrBadMessage
}
return r, nil
default:
return r, nil
}
case 1: // rejected
r := Reply{XID: xid, Rejected: true}
kind, err := d.Uint32()
if err != nil {
return Reply{}, ErrBadMessage
}
switch kind {
case RejectRPCMismatch:
if r.Mismatch.Low, err = d.Uint32(); err != nil {
return Reply{}, ErrBadMessage
}
if r.Mismatch.High, err = d.Uint32(); err != nil {
return Reply{}, ErrBadMessage
}
case RejectAuthError:
if r.AuthStat, err = d.Uint32(); err != nil {
return Reply{}, ErrBadMessage
}
default:
return Reply{}, ErrBadMessage
}
return r, nil
default:
return Reply{}, fmt.Errorf("%w: unknown reply stat %d", ErrBadMessage, stat)
}
}
// An AuthSys is an AUTH_SYS credential body: the identity the client
// asserts for every request.
type AuthSys struct {
Stamp uint32
Machine string
UID uint32
GID uint32
GIDs []uint32
}
// Body encodes the credential in the AUTH_SYS layout.
func (a AuthSys) Body() ([]byte, error) {
if len(a.GIDs) > 16 {
return nil, fmt.Errorf("rpc: %d supplementary groups exceeds 16", len(a.GIDs))
}
b := xdr.AppendUint32(nil, a.Stamp)
b = xdr.AppendString(b, a.Machine)
b = xdr.AppendUint32(b, a.UID)
b = xdr.AppendUint32(b, a.GID)
b = xdr.AppendUint32(b, uint32(len(a.GIDs)))
for _, g := range a.GIDs {
b = xdr.AppendUint32(b, g)
}
return b, nil
}
// DecodeAuthSysBody decodes an AUTH_SYS credential body.
func DecodeAuthSysBody(body []byte) (AuthSys, error) {
d := xdr.NewDecoder(body)
var a AuthSys
var err error
if a.Stamp, err = d.Uint32(); err != nil {
return a, ErrBadMessage
}
if a.Machine, err = d.String(); err != nil {
return a, ErrBadMessage
}
if a.UID, err = d.Uint32(); err != nil {
return a, ErrBadMessage
}
if a.GID, err = d.Uint32(); err != nil {
return a, ErrBadMessage
}
n, err := d.Uint32()
if err != nil {
return a, ErrBadMessage
}
if n > 16 {
return a, fmt.Errorf("rpc: %d supplementary groups exceeds 16", n)
}
for range n {
g, err := d.Uint32()
if err != nil {
return a, ErrBadMessage
}
a.GIDs = append(a.GIDs, g)
}
return a, nil
}
// PeekHeader reads the XID and the message type of a record without fully
// decoding it. It is the demultiplexer's tool: a connection that carries
// both directions distinguishes a reply to its own call from a call the
// peer issued by the message type alone.
func PeekHeader(record []byte) (xid uint32, mtype uint32, err error) {
d := xdr.NewDecoder(record)
if xid, err = d.Uint32(); err != nil {
return 0, 0, ErrBadMessage
}
if mtype, err = d.Uint32(); err != nil {
return 0, 0, ErrBadMessage
}
return xid, mtype, nil
}
// MsgDenied is the reply_stat of a rejected reply, RFC 5531 section 8.
const MsgDenied = 1
// AppendRejectedReply encodes a MSG_DENIED reply with an auth error.
func AppendRejectedReply(b []byte, xid, authStat uint32) []byte {
b = xdr.AppendUint32(b, xid)
b = xdr.AppendUint32(b, MsgReply)
b = xdr.AppendUint32(b, MsgDenied)
b = xdr.AppendUint32(b, RejectAuthError)
return xdr.AppendUint32(b, authStat)
}
+199
View File
@@ -0,0 +1,199 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"bytes"
"errors"
"testing"
"sourcedock.dev/petrbalvin/nfs/internal/xdr"
)
func TestCallRoundTrip(t *testing.T) {
cred := Auth{Flavor: FlavorSys, Body: []byte{1, 2, 3}}
c := Call{XID: 0xfeedface, Program: 100003, Version: 4, Procedure: 1, Cred: cred}
rec, err := AppendCall(nil, c)
if err != nil {
t.Fatalf("encode: %v", err)
}
rec = append(rec, xdr.AppendString(nil, "tag")...)
got, args, err := DecodeCall(rec)
if err != nil {
t.Fatalf("decode: %v", err)
}
if got.XID != c.XID || got.Program != c.Program || got.Version != c.Version ||
got.Procedure != c.Procedure ||
got.Cred.Flavor != c.Cred.Flavor || !bytes.Equal(got.Cred.Body, c.Cred.Body) {
t.Fatalf("round trip: got %+v, want %+v", got, c)
}
tag, err := xdr.NewDecoder(args).String()
if err != nil || tag != "tag" {
t.Fatalf("arguments after the header: %q, %v", tag, err)
}
}
func TestDecodeCallRejectsReplies(t *testing.T) {
rec := xdr.AppendUint32(xdr.AppendUint32(nil, 1), MsgReply)
if _, _, err := DecodeCall(rec); err == nil {
t.Fatal("a reply record decoded as a call")
}
}
func TestOversizedAuthBody(t *testing.T) {
c := Call{XID: 1, Program: 100003, Version: Version, Procedure: 0,
Cred: Auth{Flavor: FlavorSys, Body: make([]byte, maxAuthBody+1)}}
if _, err := AppendCall(nil, c); err == nil {
t.Fatal("an oversized credential encoded without error")
}
}
func TestAcceptedReplyRoundTrip(t *testing.T) {
rec, err := AppendAcceptedReply(nil, 7, AuthNull, AcceptSuccess, Mismatch{})
if err != nil {
t.Fatalf("encode: %v", err)
}
rec = append(rec, xdr.AppendUint32(nil, 42)...)
r, err := DecodeReply(rec)
if err != nil {
t.Fatalf("decode: %v", err)
}
if r.XID != 7 || r.Status != AcceptSuccess {
t.Fatalf("got xid %d status %d", r.XID, r.Status)
}
body, err := xdr.NewDecoder(r.Body).Uint32()
if err != nil || body != 42 {
t.Fatalf("result body: %d, %v", body, err)
}
}
func TestProgMismatchReply(t *testing.T) {
rec, err := AppendAcceptedReply(nil, 9, AuthNull, AcceptProgMismatch, Mismatch{Low: 4, High: 4})
if err != nil {
t.Fatalf("encode: %v", err)
}
r, err := DecodeReply(rec)
if err != nil {
t.Fatalf("decode: %v", err)
}
if r.Status != AcceptProgMismatch || r.Mismatch != (Mismatch{4, 4}) {
t.Fatalf("got status %d mismatch %+v", r.Status, r.Mismatch)
}
}
func TestUnknownAcceptStatus(t *testing.T) {
if _, err := AppendAcceptedReply(nil, 1, AuthNull, 99, Mismatch{}); err == nil {
t.Fatal("an unknown accept status encoded without error")
}
}
func TestAuthSysRoundTrip(t *testing.T) {
a := AuthSys{Stamp: 12, Machine: "client", UID: 1000, GID: 100, GIDs: []uint32{100, 5, 27}}
body, err := a.Body()
if err != nil {
t.Fatalf("encode: %v", err)
}
got, err := DecodeAuthSysBody(body)
if err != nil {
t.Fatalf("decode: %v", err)
}
if got.Stamp != a.Stamp || got.Machine != a.Machine || got.UID != a.UID ||
got.GID != a.GID || !equalU32(got.GIDs, a.GIDs) {
t.Fatalf("round trip: got %+v, want %+v", got, a)
}
}
func equalU32(a, b []uint32) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
func TestAuthSysTooManyGroups(t *testing.T) {
a := AuthSys{GIDs: make([]uint32, 17)}
if _, err := a.Body(); err == nil {
t.Fatal("17 supplementary groups encoded without error")
}
if _, err := DecodeAuthSysBody(xdr.AppendUint32(xdr.AppendUint32(nil, 0), 17)); err == nil {
t.Fatal("a credential claiming 17 groups decoded without error")
}
}
func TestDecodeRejectedReply(t *testing.T) {
// RPC_MISMATCH carries the version range.
rec := xdr.AppendUint32(nil, 5)
rec = xdr.AppendUint32(rec, MsgReply)
rec = xdr.AppendUint32(rec, 1) // rejected
rec = xdr.AppendUint32(rec, RejectRPCMismatch)
rec = xdr.AppendUint32(rec, 2)
rec = xdr.AppendUint32(rec, 4)
r, err := DecodeReply(rec)
if err != nil || r.XID != 5 || !r.Rejected || r.Mismatch != (Mismatch{2, 4}) {
t.Fatalf("rpc mismatch: %+v, %v", r, err)
}
if r.Status != 0 || r.Body != nil {
t.Fatalf("a rejected reply carries no result: %+v", r)
}
// An auth error carries one status word.
rec = xdr.AppendUint32(nil, 6)
rec = xdr.AppendUint32(rec, MsgReply)
rec = xdr.AppendUint32(rec, 1)
rec = xdr.AppendUint32(rec, RejectAuthError)
rec = xdr.AppendUint32(rec, AuthBadVerf)
if r, err = DecodeReply(rec); err != nil || r.XID != 6 || !r.Rejected || r.AuthStat != AuthBadVerf {
t.Fatalf("auth error: %+v, %v", r, err)
}
}
func TestDecodeAuthOversizeBody(t *testing.T) {
// A credential body beyond 400 bytes is refused, not buffered.
huge := xdr.AppendUint32(nil, FlavorSys)
huge = xdr.AppendUint32(huge, maxAuthBody+1)
huge = append(huge, make([]byte, 8)...)
rec := xdr.AppendUint32(nil, 1)
rec = xdr.AppendUint32(rec, MsgCall)
rec = xdr.AppendUint32(rec, Version)
rec = xdr.AppendUint32(rec, 100003)
rec = xdr.AppendUint32(rec, 4)
rec = xdr.AppendUint32(rec, 0)
rec = append(rec, huge...)
if _, _, err := DecodeCall(rec); !errors.Is(err, ErrBadMessage) {
t.Fatalf("an oversized credential decoded as %v", err)
}
}
func TestDecodeReplyGarbage(t *testing.T) {
for _, rec := range [][]byte{
nil,
{0, 0, 0}, // short
xdr.AppendUint32(nil, 1), // no type
{0, 0, 0, 1, 0, 0, 0, 9}, // unknown reply stat 9
} {
if _, err := DecodeReply(rec); !errors.Is(err, ErrBadMessage) {
t.Fatalf("record %x decoded with %v, want ErrBadMessage", rec, err)
}
}
}
func TestPeekHeader(t *testing.T) {
// A minimal call record: xid and the call type.
b := xdr.AppendUint32(xdr.AppendUint32(nil, 0x55), MsgCall)
xid, mtype, err := PeekHeader(b)
if err != nil || xid != 0x55 || mtype != MsgCall {
t.Fatalf("call peek: %x %d, %v", xid, mtype, err)
}
if _, _, err := PeekHeader(nil); !errors.Is(err, ErrBadMessage) {
t.Fatalf("an empty record: %v", err)
}
}
+103
View File
@@ -0,0 +1,103 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package rpc implements the record marking layer of ONC RPC, RFC 5531.
//
// An ONC RPC message travels over a byte stream as one record: a sequence of
// one or more fragments, each headed by a 32 bit word whose high bit marks
// the last fragment of the record and whose low 31 bits carry the fragment
// length in bytes.
package rpc
import (
"errors"
"fmt"
"io"
"sourcedock.dev/petrbalvin/nfs/internal/xdr"
)
// LastFragment is the high bit of a fragment header, set on the final
// fragment of a record.
const LastFragment = 1 << 31
// maxFragment is the largest fragment length the low 31 bits can carry.
const maxFragment = 1<<31 - 1
// ErrRecordTooLarge is returned by ReadRecord when a record exceeds the
// caller's limit.
var ErrRecordTooLarge = errors.New("rpc: record exceeds the size limit")
// AppendFragmentHeader appends the record marking header of a fragment
// that carries n bytes. It panics when n is negative or above the
// largest fragment length; callers reach it through WriteRecord, which
// rejects such input with ErrRecordTooLarge instead.
func AppendFragmentHeader(b []byte, n int, last bool) []byte {
if n < 0 || n > maxFragment {
panic(fmt.Sprintf("rpc: fragment length %d out of range", n))
}
h := uint32(n)
if last {
h |= LastFragment
}
return xdr.AppendUint32(b, h)
}
// WriteRecord writes data to w as one record in a single final fragment.
// The caller keeps the record under maxFragment bytes; a call that carries a
// whole ONC RPC request or reply always fits.
func WriteRecord(w io.Writer, data []byte) error {
if len(data) > maxFragment {
return ErrRecordTooLarge
}
buf := AppendFragmentHeader(make([]byte, 0, 4+len(data)), len(data), true)
buf = append(buf, data...)
_, err := w.Write(buf)
return err
}
// ReadRecord reads one record from r and returns its reassembled bytes. The
// record may arrive in any number of fragments and may exceed the reader's
// own buffer only up to limit bytes; a longer record returns
// ErrRecordTooLarge before the limit is exceeded in memory.
func ReadRecord(r io.Reader, limit int) ([]byte, error) {
var header [4]byte
var record []byte
for {
if _, err := io.ReadFull(r, header[:]); err != nil {
if errors.Is(err, io.EOF) && len(record) == 0 {
return nil, io.EOF
}
return nil, fmt.Errorf("rpc: fragment header: %w", err)
}
h, err := xdr.NewDecoder(header[:]).Uint32()
if err != nil {
// Unreachable: a four byte input always holds a uint32.
return nil, err
}
n := int(h &^ LastFragment)
if n > limit-len(record) {
return nil, ErrRecordTooLarge
}
start := len(record)
record = append(record, make([]byte, n)...)
if _, err := io.ReadFull(r, record[start:]); err != nil {
return nil, fmt.Errorf("rpc: fragment body: %w", err)
}
if h&LastFragment != 0 {
return record, nil
}
}
}
// AppendRecord appends the record marking of one last fragment and the
// data behind it, for a caller that owns the destination buffer: a
// writer that recycles its wire buffers through this function saves the
// allocation WriteRecord makes per call.
func AppendRecord(dst, data []byte) []byte {
if len(data) > maxFragment {
return dst
}
dst = AppendFragmentHeader(dst, len(data), true)
return append(dst, data...)
}
+137
View File
@@ -0,0 +1,137 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package rpc
import (
"bytes"
"errors"
"io"
"strings"
"testing"
)
func TestFragmentHeaderLayout(t *testing.T) {
got := AppendFragmentHeader(nil, 5, true)
want := []byte{0x80, 0, 0, 5}
if !bytes.Equal(got, want) {
t.Fatalf("last fragment header: got %x, want %x", got, want)
}
got = AppendFragmentHeader(nil, 5, false)
want = []byte{0, 0, 0, 5}
if !bytes.Equal(got, want) {
t.Fatalf("continuation header: got %x, want %x", got, want)
}
}
func TestWriteReadRoundTrip(t *testing.T) {
records := [][]byte{
nil,
[]byte("a"),
[]byte("abc"),
[]byte("abcd"),
bytes.Repeat([]byte{0xc3}, 70000),
}
for _, want := range records {
pr, pw := io.Pipe()
go func() {
err := WriteRecord(pw, want)
if err != nil {
pw.CloseWithError(err)
return
}
pw.Close()
}()
got, err := ReadRecord(pr, 1<<20)
if err != nil {
t.Fatalf("record of %d bytes: %v", len(want), err)
}
if !bytes.Equal(got, want) {
t.Fatalf("record of %d bytes came back as %d bytes", len(want), len(got))
}
}
}
func TestReadReassemblesFragments(t *testing.T) {
// A record delivered as three fragments arrives as the same bytes.
var stream bytes.Buffer
stream.Write(AppendFragmentHeader(nil, 2, false))
stream.WriteString("ab")
stream.Write(AppendFragmentHeader(nil, 0, false))
stream.Write(AppendFragmentHeader(nil, 3, true))
stream.WriteString("cde")
got, err := ReadRecord(&stream, 64)
if err != nil {
t.Fatalf("read: %v", err)
}
if string(got) != "abcde" {
t.Fatalf("reassembled record: got %q, want %q", got, "abcde")
}
}
func TestReadRecordTooLarge(t *testing.T) {
var stream bytes.Buffer
stream.Write(AppendFragmentHeader(nil, 100, true))
stream.Write(bytes.Repeat([]byte{0}, 100))
if _, err := ReadRecord(&stream, 64); !errors.Is(err, ErrRecordTooLarge) {
t.Fatalf("a record of 100 bytes against a limit of 64 returned %v", err)
}
// A record split over fragments is bounded by the record total, not by
// the single fragment length.
stream.Reset()
stream.Write(AppendFragmentHeader(nil, 50, false))
stream.Write(bytes.Repeat([]byte{0}, 50))
stream.Write(AppendFragmentHeader(nil, 50, true))
stream.Write(bytes.Repeat([]byte{0}, 50))
if _, err := ReadRecord(&stream, 64); !errors.Is(err, ErrRecordTooLarge) {
t.Fatalf("a split record of 100 bytes against a limit of 64 returned %v", err)
}
}
func TestReadRecordAtLimit(t *testing.T) {
// A record of exactly the limit is legal; one byte more is not.
// The boundary is strict, so a record of limit bytes must arrive.
var stream bytes.Buffer
stream.Write(AppendFragmentHeader(nil, 64, true))
stream.Write(bytes.Repeat([]byte{0}, 64))
if got, err := ReadRecord(&stream, 64); err != nil || len(got) != 64 {
t.Fatalf("a record of exactly the limit: %d bytes, %v", len(got), err)
}
stream.Reset()
stream.Write(AppendFragmentHeader(nil, 65, true))
stream.Write(bytes.Repeat([]byte{0}, 65))
if _, err := ReadRecord(&stream, 64); !errors.Is(err, ErrRecordTooLarge) {
t.Fatalf("a record one byte over the limit returned %v", err)
}
}
func TestReadCleanEOF(t *testing.T) {
if _, err := ReadRecord(strings.NewReader(""), 64); !errors.Is(err, io.EOF) {
t.Fatalf("an empty stream returned %v, want io.EOF", err)
}
}
func TestReadTruncated(t *testing.T) {
// A header that promises ten bytes and a body of three is a truncated
// fragment, not a clean end of stream.
var stream bytes.Buffer
stream.Write(AppendFragmentHeader(nil, 10, true))
stream.WriteString("abc")
if _, err := ReadRecord(&stream, 64); !errors.Is(err, io.ErrUnexpectedEOF) {
t.Fatalf("a truncated body returned %v", err)
}
}
func TestReadTruncatedHeader(t *testing.T) {
// A record already reassembling that loses its next header is also a
// truncation, not a clean end of stream.
var stream bytes.Buffer
stream.Write(AppendFragmentHeader(nil, 1, false))
stream.WriteString("a")
stream.Write([]byte{0, 0})
if _, err := ReadRecord(&stream, 64); !errors.Is(err, io.ErrUnexpectedEOF) {
t.Fatalf("a truncated continuation header returned %v", err)
}
}