diff --git a/.codecov.yml b/.codecov.yml new file mode 100644 index 0000000..2e44464 --- /dev/null +++ b/.codecov.yml @@ -0,0 +1,2 @@ +ignore: + - "examples/**" diff --git a/examples/ntp-client/main.go b/examples/ntp-client/main.go new file mode 100644 index 0000000..bfd6a3c --- /dev/null +++ b/examples/ntp-client/main.go @@ -0,0 +1,90 @@ +// Command ntp-client performs a two-exchange NTP clock synchronization against +// a remote server and prints the corrected time, clock offset, and round-trip +// delay. +// +// Usage: +// +// go run ./examples/ntp-client/ -server pool.ntp.org:123 +// go run ./examples/ntp-client/ -server 127.0.0.1:10123 -debug +// +// This tool uses the standard library net package for UDP transport instead of +// lneto's own networking stack. These examples exercise one protocol layer at a +// time in isolation, keeping the transport concern separate so failures are +// clearly attributable to the NTP codec and state machine rather than the +// full-stack IP/UDP path. +package main + +import ( + "flag" + "fmt" + "log/slog" + "net" + "os" + "time" + + "github.com/soypat/lneto/ntp" +) + +func main() { + if err := run(); err != nil { + fmt.Println(err) + os.Exit(1) + } +} + +func run() error { + addr := flag.String("server", "pool.ntp.org:123", "NTP server address (host:port)") + debug := flag.Bool("debug", false, "enable debug logging") + flag.Parse() + + conn, err := net.DialTimeout("udp", *addr, 5*time.Second) + if err != nil { + return fmt.Errorf("dial: %w", err) + } + defer conn.Close() + + var precBuf [64]int64 + sysprec := ntp.CalculateSystemPrecision(nil, precBuf[:]) + + var client ntp.Client + client.Reset(sysprec, time.Now) + if *debug { + client.SetLogger(slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelDebug}))) + } + + const maxRetries = 10 + var buf [1500]byte + for attempt := 0; !client.IsDone() && attempt < maxRetries; attempt++ { + n, err := client.Encapsulate(buf[:ntp.SizeHeader], 0, 0) + if err != nil { + return fmt.Errorf("encapsulate: %w", err) + } + if n == 0 { + return fmt.Errorf("encapsulate returned 0 bytes unexpectedly") + } + + conn.SetDeadline(time.Now().Add(5 * time.Second)) + if _, err = conn.Write(buf[:n]); err != nil { + return fmt.Errorf("write: %w", err) + } + + rn, err := conn.Read(buf[:]) + if err != nil { + return fmt.Errorf("read: %w", err) + } + + if err = client.Demux(buf[:rn], 0); err != nil { + return fmt.Errorf("demux: %w", err) + } + } + + if !client.IsDone() { + return fmt.Errorf("NTP exchange did not complete within %d attempts", maxRetries) + } + + fmt.Printf("NTP time: %s\n", client.Now().Format(time.RFC3339Nano)) + fmt.Printf("Offset: %s\n", client.Offset()) + fmt.Printf("RTD: %s\n", client.RoundTripDelay()) + fmt.Printf("Stratum: %s\n", client.ServerStratum()) + return nil +} diff --git a/examples/ntp-server/main.go b/examples/ntp-server/main.go new file mode 100644 index 0000000..636b0b7 --- /dev/null +++ b/examples/ntp-server/main.go @@ -0,0 +1,96 @@ +// Command ntp-server is a minimal NTP server that listens for client requests +// on a UDP socket and responds with the current system time. It serves as an +// integration test target for the ntp-client example. +// +// Usage: +// +// go run ./examples/ntp-server/ -addr :10123 +// +// The listen address defaults to :123 (requires root). +// +// This tool uses the standard library net package for UDP transport instead of +// lneto's own networking stack. These examples exercise one protocol layer at a +// time in isolation, keeping the transport concern separate so failures are +// clearly attributable to the NTP codec and state machine rather than the +// full-stack IP/UDP path. +package main + +import ( + "flag" + "fmt" + "net" + "os" + "time" + + "github.com/soypat/lneto/ntp" +) + +func main() { + if err := run(); err != nil { + fmt.Println(err) + os.Exit(1) + } +} + +func run() error { + listenAddr := flag.String("addr", ":123", "UDP listen address (host:port)") + flag.Parse() + + pc, err := net.ListenPacket("udp", *listenAddr) + if err != nil { + return fmt.Errorf("listen: %w", err) + } + defer pc.Close() + fmt.Printf("NTP server listening on %s\n", pc.LocalAddr()) + + var precBuf [64]int64 + sysprec := ntp.CalculateSystemPrecision(nil, precBuf[:]) + + var handler ntp.Server + err = handler.Reset(ntp.ServerConfig{ + Now: time.Now, + Stratum: ntp.StratumPrimary, + Precision: sysprec, + RefID: [4]byte{'G', 'O', 'L', 'N'}, + MaxPending: 16, + }) + if err != nil { + return fmt.Errorf("server reset: %w", err) + } + + var buf [1500]byte + var backoff uint + for { + pc.SetReadDeadline(time.Now().Add(100 * time.Millisecond)) + n, raddr, err := pc.ReadFrom(buf[:]) + if err != nil { + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + backoff++ + if backoff > 10 { + time.Sleep(time.Duration(backoff) * time.Millisecond) + } + continue + } + fmt.Printf("read error: %v\n", err) + continue + } + backoff = 0 + + if err = handler.Demux(buf[:n], 0); err != nil { + fmt.Printf("demux error from %s: %v\n", raddr, err) + continue + } + + var resp [ntp.SizeHeader]byte + rn, err := handler.Encapsulate(resp[:], 0, 0) + if err != nil { + fmt.Printf("encapsulate error: %v\n", err) + continue + } + if rn > 0 { + if _, err = pc.WriteTo(resp[:rn], raddr); err != nil { + fmt.Printf("write error: %v\n", err) + } + } + } +} diff --git a/ntp/client.go b/ntp/client.go index aa3991e..ac37c30 100644 --- a/ntp/client.go +++ b/ntp/client.go @@ -1,9 +1,11 @@ package ntp import ( + "log/slog" "time" "github.com/soypat/lneto" + "github.com/soypat/lneto/internal" ) type state uint8 @@ -23,14 +25,15 @@ type Client struct { connID uint64 start time.Time _now func() time.Time - // t stores the time offsets needed to compute the time at client - // taking into consideration the round-trip delay. - // - t[0] (orig): Client timestamp of request packet transmission. - // - t[1] (rec): Server timestamp of request packet reception. - // - t[2] (xmt): Server timestamp of response packet transmission. - // - t[3]: Client timestamp of response packet reception. - t [4]Timestamp - // org Timestamp + logger logger + // t stores the four NTP timestamps per RFC 5905: + // - t[0] (T1): Client timestamp of request packet transmission. + // - t[1] (T2): Server timestamp of request packet reception. + // - t[2] (T3): Server timestamp of response packet transmission. + // - t[3] (T4): Client timestamp of response packet reception. + t [4]Timestamp + offset1 time.Duration // clock offset from first exchange, averaged with second in OffsetUnsynced. + rtt1 time.Duration // round-trip delay from first exchange, averaged with second in RoundTripDelay. state state serverStratum Stratum sysprec int8 @@ -40,11 +43,16 @@ func (c *Client) Reset(sysprec int8, now func() time.Time) { *c = Client{ connID: c.connID + 1, _now: now, + logger: c.logger, sysprec: sysprec, state: stateSend1, } } +// SetLogger configures a structured logger for debug output. +// Pass nil to disable logging (the default). +func (c *Client) SetLogger(l *slog.Logger) { c.logger.log = l } + func (c *Client) Protocol() uint64 { return 0 } func (c *Client) LocalPort() uint16 { return ClientPort } func (c *Client) ConnectionID() *uint64 { @@ -62,27 +70,32 @@ func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) } switch c.state { - case stateSend1: - c.start = c.now() - c.t[0] = TimestampFromUint64(0) - c.state = stateAwait1 - case stateSend2: - // c.xmt = c.unsyncTimestamp(c.now()) - c.state = stateDone + case stateSend1, stateSend2: + now := c.now() + c.start = now + var err error + if c.t[0], err = TimestampFromTime(now); err != nil { + return 0, err + } + if c.state == stateSend1 { + c.state = stateAwait1 + } else { + c.state = stateAwait2 + } default: return 0, nil // Nothing to handle. } - for i := range payload[:SizeHeader] { - payload[i] = 0 - } - frm.ClearHeader() frm.SetStratum(StratumUnsync) frm.SetPoll(6) frm.SetPrecision(c.sysprec) - frm.SetOriginTime(c.t[0]) + // RFC 5905 §8: client places T1 in TransmitTime of the request. + // The server will echo it back as OriginTime in its response. + frm.SetTransmitTime(c.t[0]) frm.SetFlags(ModeClient, Version4, LeapNoWarning) + c.logger.debug("ntp.Client:encapsulate", slog.Int("state", int(c.state)), + slog.Uint64("T1", c.t[0].Uint64())) return SizeHeader, nil } @@ -97,21 +110,45 @@ func (c *Client) Demux(carrierData []byte, frameOffset int) error { } switch c.state { - case stateAwait1: - xmt := frm.TransmitTime() - orig := frm.OriginTime() - if xmt == orig || orig != c.t[0] { - return lneto.ErrPacketDrop - } + case stateAwait1, stateAwait2: + default: + return nil // Not awaiting a response. + } - txelapsed := c.now().Sub(c.start) - c.t[1] = frm.ReceiveTime() - c.t[2] = xmt - c.t[3] = c.t[0].Add(txelapsed) + // RFC 5905 §8 validation: discard bogus packets. + // Bogus: origin timestamp does not echo our T1 (the transmit time we sent). + // Malformed: server's transmit time equals its own origin echo. + xmt := frm.TransmitTime() + orig := frm.OriginTime() + if xmt == orig || orig != c.t[0] { + c.logger.debug("ntp.Client:demux:drop", slog.String("reason", "origin mismatch"), + slog.Uint64("orig", orig.Uint64()), slog.Uint64("T1", c.t[0].Uint64())) + return lneto.ErrPacketDrop + } + + // Compute T4, then derive offset θ and round-trip delay δ per RFC 5905 §8. + txelapsed := c.now().Sub(c.start) + c.t[1] = frm.ReceiveTime() + c.t[2] = xmt + c.t[3] = c.t[0].Add(txelapsed) + + offset := (c.t[1].Sub(c.t[0]) + c.t[2].Sub(c.t[3])) / 2 + rtt := c.t[3].Sub(c.t[0]) - c.t[2].Sub(c.t[1]) + + if c.state == stateAwait1 { c.serverStratum = frm.Stratum() - c.state = stateDone // TODO: add second exchange part. - case stateAwait2: + c.offset1 = offset + c.rtt1 = rtt + c.state = stateSend2 + c.logger.debug("ntp.Client:demux:exchange1", + slog.Duration("offset", c.offset1), slog.Duration("rtt", c.rtt1), + slog.String("stratum", c.serverStratum.String())) + } else { c.state = stateDone + c.logger.debug("ntp.Client:demux:exchange2", + slog.Duration("offset", offset), slog.Duration("rtt", rtt), + slog.Duration("avg_offset", (c.offset1+offset)/2), + slog.Duration("avg_rtt", (c.rtt1+rtt)/2)) } return nil } @@ -147,28 +184,35 @@ func (c *Client) Offset() time.Duration { } func (c *Client) offsetAndNow() (clientNow time.Time, offset time.Duration) { - now := c.now() - serverToBase := c.OffsetUnsynced() - clientToBase := now.Sub(BaseTime()) - serverToClient := serverToBase - clientToBase - return now, serverToClient + return c.now(), c.OffsetUnsynced() } // OffsetUnsynced returns the absolute time offset difference between client and server clock -// as calculated by the clock synchonization algorithm. It is unsynchonized- the result of OffsetUnsynced will not change with time. +// as calculated by the clock synchronization algorithm. It is unsynced — the result will not +// change with time. When both exchanges are complete the result is the average of both exchanges. func (c *Client) OffsetUnsynced() time.Duration { if c.IsDone() { t := &c.t - return (t[1].Sub(t[0]) + t[2].Sub(t[3])) / 2 + offset2 := (t[1].Sub(t[0]) + t[2].Sub(t[3])) / 2 + return (c.offset1 + offset2) / 2 } return 0 } +// RoundTripDelay returns the average round-trip delay across both NTP exchanges. func (c *Client) RoundTripDelay() time.Duration { if c.IsDone() { - d0 := c.t[3].Sub(c.t[0]) - d1 := c.t[2].Sub(c.t[1]) - return d0 - d1 + rtt2 := c.t[3].Sub(c.t[0]) - c.t[2].Sub(c.t[1]) + return (c.rtt1 + rtt2) / 2 } return -1 } + +// logger provides non-allocating structured logging using [internal.LogAttrs]. +type logger struct { + log *slog.Logger +} + +func (l logger) debug(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, slog.LevelDebug, msg, attrs...) +} diff --git a/ntp/client_test.go b/ntp/client_test.go index 8460cbe..8a2f2ce 100644 --- a/ntp/client_test.go +++ b/ntp/client_test.go @@ -5,6 +5,32 @@ import ( "time" ) +// simulateServerResponse builds a server NTP response that echoes the client's +// TransmitTime as the response's OriginTime (RFC 5905 §8), then sets server +// receive and transmit timestamps. +func simulateServerResponse(t *testing.T, reqBuf []byte, serverRecv, serverXmt time.Time) []byte { + t.Helper() + reqFrm, _ := NewFrame(reqBuf) + respBuf := make([]byte, SizeHeader) + respFrm, _ := NewFrame(respBuf) + respFrm.SetFlags(ModeServer, Version4, LeapNoWarning) + respFrm.SetStratum(StratumPrimary) + respFrm.SetPrecision(-20) + // Server echoes client's TransmitTime as response OriginTime per RFC 5905 §8. + respFrm.SetOriginTime(reqFrm.TransmitTime()) + recvTS, err := TimestampFromTime(serverRecv) + if err != nil { + t.Fatal(err) + } + xmtTS, err := TimestampFromTime(serverXmt) + if err != nil { + t.Fatal(err) + } + respFrm.SetReceiveTime(recvTS) + respFrm.SetTransmitTime(xmtTS) + return respBuf +} + func TestClient_FullExchange(t *testing.T) { // Simulate a NTP client-server exchange without network. baseTime := BaseTime() @@ -19,7 +45,7 @@ func TestClient_FullExchange(t *testing.T) { t.Fatal("client should not be done before exchange") } - // Step 1: Client encapsulates request. + // Step 1: Client encapsulates first request. reqBuf := make([]byte, SizeHeader) n, err := client.Encapsulate(reqBuf, 0, 0) if err != nil { @@ -49,42 +75,53 @@ func TestClient_FullExchange(t *testing.T) { // Server receives at clientStart + serverOffset, sends response at clientStart + serverOffset + 10ms processing. serverRecvTime := clientStart.Add(serverOffset) serverXmtTime := serverRecvTime.Add(10 * time.Millisecond) - - respBuf := make([]byte, SizeHeader) - respFrm, _ := NewFrame(respBuf) - respFrm.SetFlags(ModeServer, Version4, LeapNoWarning) - respFrm.SetStratum(StratumPrimary) - respFrm.SetPrecision(-20) - - // Echo client's origin time. - respFrm.SetOriginTime(reqFrm.OriginTime()) - - // Set server timestamps. - recvTS, err := TimestampFromTime(serverRecvTime) - if err != nil { - t.Fatal(err) - } - xmtTS, err := TimestampFromTime(serverXmtTime) - if err != nil { - t.Fatal(err) - } - respFrm.SetReceiveTime(recvTS) - respFrm.SetTransmitTime(xmtTS) + respBuf := simulateServerResponse(t, reqBuf, serverRecvTime, serverXmtTime) // Advance client clock to simulate network delay. clockTime = clientStart.Add(100 * time.Millisecond) - // Step 3: Client demuxes response. + // Step 3: Client demuxes first response. err = client.Demux(respBuf, 0) if err != nil { t.Fatal(err) } - if !client.IsDone() { - t.Fatal("client should be done after exchange") + if client.IsDone() { + t.Fatal("client should not be done after first exchange only") + } + if client.ServerStratum() != StratumPrimary { + t.Errorf("server stratum = %s; want primary", client.ServerStratum()) } - // Step 4: Verify results. + // Step 4: Client encapsulates second request. + req2Buf := make([]byte, SizeHeader) + clockTime = clientStart.Add(200 * time.Millisecond) + n, err = client.Encapsulate(req2Buf, 0, 0) + if err != nil { + t.Fatal(err) + } + if n != SizeHeader { + t.Fatalf("second request: expected %d bytes, got %d", SizeHeader, n) + } + + // Step 5: Simulate second server response. + serverRecv2 := clientStart.Add(serverOffset + 200*time.Millisecond) + serverXmt2 := serverRecv2.Add(10 * time.Millisecond) + resp2Buf := simulateServerResponse(t, req2Buf, serverRecv2, serverXmt2) + + clockTime = clientStart.Add(300 * time.Millisecond) + + // Step 6: Client demuxes second response. + err = client.Demux(resp2Buf, 0) + if err != nil { + t.Fatal(err) + } + + if !client.IsDone() { + t.Fatal("client should be done after second exchange") + } + + // Step 7: Verify results. if client.ServerStratum() != StratumPrimary { t.Errorf("server stratum = %s; want primary", client.ServerStratum()) } @@ -174,6 +211,49 @@ func TestClient_OffsetBeforeDone(t *testing.T) { } } +func TestClient_SecondExchangeRejection(t *testing.T) { + baseTime := BaseTime() + clientStart := baseTime.Add(10 * time.Second) + serverOffset := 500 * time.Millisecond + clockTime := clientStart + + var client Client + client.Reset(-18, func() time.Time { return clockTime }) + + // Complete first exchange. + reqBuf := make([]byte, SizeHeader) + client.Encapsulate(reqBuf, 0, 0) + + serverRecv1 := clientStart.Add(serverOffset) + serverXmt1 := serverRecv1.Add(10 * time.Millisecond) + resp1Buf := simulateServerResponse(t, reqBuf, serverRecv1, serverXmt1) + clockTime = clientStart.Add(100 * time.Millisecond) + client.Demux(resp1Buf, 0) + + // Start second exchange. + req2Buf := make([]byte, SizeHeader) + clockTime = clientStart.Add(200 * time.Millisecond) + client.Encapsulate(req2Buf, 0, 0) + + // Build bogus response with wrong origin time. + bogus := make([]byte, SizeHeader) + frm, _ := NewFrame(bogus) + frm.SetFlags(ModeServer, Version4, LeapNoWarning) + frm.SetOriginTime(TimestampFromUint64(99999)) + xmt, _ := TimestampFromTime(clockTime.Add(time.Second)) + frm.SetTransmitTime(xmt) + frm.SetReceiveTime(xmt) + + clockTime = clientStart.Add(300 * time.Millisecond) + err := client.Demux(bogus, 0) + if err == nil { + t.Fatal("second exchange should reject mismatched origin") + } + if client.IsDone() { + t.Fatal("should not be done after rejected second response") + } +} + func TestClient_DemuxRejectsBogusResponse(t *testing.T) { var c Client clockTime := BaseTime().Add(time.Second) diff --git a/ntp/extensions.go b/ntp/extensions.go new file mode 100644 index 0000000..9d539a6 --- /dev/null +++ b/ntp/extensions.go @@ -0,0 +1,116 @@ +package ntp + +import ( + "encoding/binary" + + "github.com/soypat/lneto" +) + +// ExtType identifies the type of an NTP extension field. +// See RFC 7822 and RFC 8915. +type ExtType uint16 + +const ( + // NTS Unique Identifier extension field (RFC 8915 §5.3, critical). + // Contains a random nonce used to prevent replay attacks. + ExtNTSUniqueID ExtType = 0x0104 + // NTS Cookie extension field (RFC 8915 §5.4). + // Contains an encrypted cookie obtained during NTS-KE key exchange. + ExtNTSCookie ExtType = 0x0204 + // NTS Cookie Placeholder extension field (RFC 8915 §5.5). + // Requests additional cookies from the server in its response. + ExtNTSCookiePlaceholder ExtType = 0x0304 + // NTS Authenticator and Encrypted Extension Fields (RFC 8915 §5.6, critical). + // Contains the AEAD-authenticated and encrypted extension fields. + // + // Full NTS authentication using this field requires an AEAD cipher + // (AEAD_AES_SIV_CMAC_256 per RFC 8915 §5.7) and session keys obtained + // during NTS-KE (RFC 8915 §4). The crypto portion is not implemented + // here due to AES-SIV not being available in the Go standard library. + // Callers may supply their own cipher.AEAD to build/verify this field. + ExtNTSAuthAndEEF ExtType = 0x0404 +) + +// sizeExtHeader is the fixed 4-byte header size of every NTP extension field (RFC 7822). +const sizeExtHeader = 4 + +// ExtField provides zero-copy access to a single NTP extension field +// within an existing packet buffer. +type ExtField struct { + buf []byte +} + +// Type returns the extension field type. +func (ef ExtField) Type() ExtType { + return ExtType(binary.BigEndian.Uint16(ef.buf[0:2])) +} + +// TotalLen returns the total length of the extension field, including the +// 4-byte header. Always a multiple of 4. +func (ef ExtField) TotalLen() uint16 { + return binary.BigEndian.Uint16(ef.buf[2:4]) +} + +// Value returns the extension field value bytes (body only, without the 4-byte header). +// Returns nil if the length field is inconsistent with the buffer. +func (ef ExtField) Value() []byte { + n := int(ef.TotalLen()) + if n < sizeExtHeader || n > len(ef.buf) { + return nil + } + return ef.buf[sizeExtHeader:n] +} + +// RawData returns the complete extension field bytes including the 4-byte header. +func (ef ExtField) RawData() []byte { return ef.buf } + +// NextExtField parses the first NTP extension field from buf and returns it +// along with the number of bytes consumed. An empty buf returns a zero n +// with nil error. Use this in a loop: +// +// for off := 0; off < len(payload); { +// field, n, err := ntp.NextExtField(payload[off:]) +// if err != nil { break } +// // process field +// off += n +// } +func NextExtField(buf []byte) (field ExtField, n int, err error) { + if len(buf) == 0 { + return ExtField{}, 0, nil + } + if len(buf) < sizeExtHeader { + return ExtField{}, 0, lneto.ErrTruncatedFrame + } + totalLen := int(binary.BigEndian.Uint16(buf[2:4])) + if totalLen < sizeExtHeader { + return ExtField{}, 0, lneto.ErrInvalidLengthField + } + if totalLen%4 != 0 { + return ExtField{}, 0, lneto.ErrInvalidLengthField + } + if totalLen > len(buf) { + return ExtField{}, 0, lneto.ErrTruncatedFrame + } + return ExtField{buf: buf[:totalLen]}, totalLen, nil +} + +// AppendExtField appends a single NTP extension field with the given type and value +// to dst. The value is zero-padded to the nearest 4-byte boundary. Returns the +// extended dst slice. Panics if the padded total length exceeds 65535 (the uint16 maximum), +// which cannot occur with any valid NTP packet payload. +func AppendExtField(dst []byte, typ ExtType, value []byte) []byte { + padded := (len(value) + 3) &^ 3 + total := sizeExtHeader + padded + if total > 0xFFFF { + panic("ntp: AppendExtField: value too large to encode in uint16 length field") + } + var hdr [sizeExtHeader]byte + binary.BigEndian.PutUint16(hdr[0:2], uint16(typ)) + binary.BigEndian.PutUint16(hdr[2:4], uint16(total)) + dst = append(dst, hdr[:]...) + dst = append(dst, value...) + for i := len(value); i < padded; i++ { + dst = append(dst, 0) + } + return dst +} diff --git a/ntp/extensions_test.go b/ntp/extensions_test.go new file mode 100644 index 0000000..cb7e1cd --- /dev/null +++ b/ntp/extensions_test.go @@ -0,0 +1,176 @@ +package ntp + +import ( + "testing" + + "github.com/soypat/lneto" +) + +func TestNextExtField_Empty(t *testing.T) { + field, n, err := NextExtField(nil) + if err != nil { + t.Fatal(err) + } + if len(field.RawData()) != 0 { + t.Errorf("NextExtField(nil) RawData len = %d; want 0", len(field.RawData())) + } + if n != 0 { + t.Errorf("NextExtField(nil) n = %d; want 0", n) + } +} + +func TestAppendAndIterateExtFields(t *testing.T) { + uid := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16} + cookie := []byte{0xAA, 0xBB, 0xCC} + + var buf []byte + buf = AppendExtField(buf, ExtNTSUniqueID, uid) + buf = AppendExtField(buf, ExtNTSCookie, cookie) + + // Each field should be padded to 4-byte boundary. + // UID: 4 header + 16 value = 20 bytes (already aligned) + // Cookie: 4 header + 3 value + 1 padding = 8 bytes + const wantLen = 20 + 8 + if len(buf) != wantLen { + t.Fatalf("AppendExtField total len = %d; want %d", len(buf), wantLen) + } + + off := 0 + field, n, err := NextExtField(buf[off:]) + if err != nil { + t.Fatal(err) + } + off += n + if field.Type() != ExtNTSUniqueID { + t.Errorf("field 1 Type() = %#x; want ExtNTSUniqueID (%#x)", field.Type(), ExtNTSUniqueID) + } + if string(field.Value()) != string(uid) { + t.Errorf("field 1 Value() mismatch") + } + + field, n, err = NextExtField(buf[off:]) + if err != nil { + t.Fatal(err) + } + off += n + if field.Type() != ExtNTSCookie { + t.Errorf("field 2 Type() = %#x; want ExtNTSCookie (%#x)", field.Type(), ExtNTSCookie) + } + // Value() includes the 4-byte-aligned body (RFC 7822 §2.1 length includes padding). + wantCookiePadded := []byte{0xAA, 0xBB, 0xCC, 0x00} + if string(field.Value()) != string(wantCookiePadded) { + t.Errorf("field 2 Value() = %v; want %v", field.Value(), wantCookiePadded) + } + + field, n, err = NextExtField(buf[off:]) + if err != nil { + t.Fatal(err) + } + if len(field.RawData()) != 0 { + t.Errorf("NextExtField after last: RawData len = %d; want 0", len(field.RawData())) + } + _ = n +} + +func TestNextExtField_Errors(t *testing.T) { + tests := []struct { + name string + buf []byte + }{ + {name: "truncated/2bytes", buf: []byte{0x01, 0x04}}, + {name: "length_below_min", buf: []byte{0x01, 0x04, 0x00, 0x02}}, + {name: "length_unaligned", buf: []byte{0x01, 0x04, 0x00, 0x05, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF}}, + {name: "length_exceeds_buf", buf: []byte{0x01, 0x04, 0x00, 0x08, 0x00}}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + _, _, err := NextExtField(tc.buf) + if err == nil { + t.Errorf("NextExtField(%x) = nil error; want error", tc.buf) + } + }) + } +} + +func TestFrameExtensionFields(t *testing.T) { + t.Run("with_extensions", func(t *testing.T) { + buf := make([]byte, SizeHeader+8) + frm, err := NewFrame(buf) + if err != nil { + t.Fatal(err) + } + p := frm.ExtensionFields() + if len(p) != 8 { + t.Errorf("ExtensionFields() len = %d; want 8", len(p)) + } + }) + t.Run("header_only", func(t *testing.T) { + buf := make([]byte, SizeHeader) + frm, _ := NewFrame(buf) + if len(frm.ExtensionFields()) != 0 { + t.Errorf("ExtensionFields() len = %d; want 0 for header-only frame", len(frm.ExtensionFields())) + } + }) +} + +func TestFrameValidateSize(t *testing.T) { + t.Run("header_only", func(t *testing.T) { + var v lneto.Validator + buf := make([]byte, SizeHeader) + frm, _ := NewFrame(buf) + frm.ValidateSize(&v) + if v.HasError() { + t.Errorf("ValidateSize(header-only) = %v; want no error", v.ErrPop()) + } + }) + t.Run("valid_extension", func(t *testing.T) { + var v lneto.Validator + ext := AppendExtField(nil, ExtNTSUniqueID, make([]byte, 16)) + buf := make([]byte, SizeHeader+len(ext)) + copy(buf[SizeHeader:], ext) + frm, _ := NewFrame(buf) + frm.ValidateSize(&v) + if v.HasError() { + t.Errorf("ValidateSize(valid ext) = %v; want no error", v.ErrPop()) + } + }) + t.Run("malformed_extension", func(t *testing.T) { + var v lneto.Validator + buf := make([]byte, SizeHeader+4) + buf[SizeHeader+2] = 0x00 + buf[SizeHeader+3] = 0x05 // length = 5, not 4-byte aligned + frm, _ := NewFrame(buf) + frm.ValidateSize(&v) + if !v.HasError() { + t.Errorf("ValidateSize(malformed ext) = no error; want error") + } + v.ErrPop() + }) +} + +func FuzzNextExtField(f *testing.F) { + f.Add(AppendExtField(nil, ExtNTSUniqueID, make([]byte, 16))) + f.Add(AppendExtField(nil, ExtNTSCookie, make([]byte, 64))) + two := AppendExtField(nil, ExtNTSUniqueID, make([]byte, 32)) + two = AppendExtField(two, ExtNTSCookie, make([]byte, 8)) + f.Add(two) + f.Add([]byte{}) + f.Add([]byte{0x01}) + f.Add([]byte{0, 1, 0, 4}) + f.Fuzz(func(t *testing.T, data []byte) { + off := 0 + for off < len(data) { + field, n, err := NextExtField(data[off:]) + if err != nil { + return + } + if len(field.RawData()) == 0 { + return + } + _ = field.Type() + _ = field.TotalLen() + _ = field.Value() + off += n + } + }) +} diff --git a/ntp/ntp.go b/ntp/ntp.go index 70731b2..0ee8e89 100644 --- a/ntp/ntp.go +++ b/ntp/ntp.go @@ -131,6 +131,34 @@ func (frm Frame) SetTransmitTime(rt Timestamp) { rt.Put(frm.buf[40:48]) } +// RawData returns the underlying byte slice for the entire NTP packet. +func (frm Frame) RawData() []byte { return frm.buf } + +// ExtensionFields returns the extension fields area of the NTP packet (all +// bytes following the fixed 48-byte NTP header). The RFC calls these +// "extension fields" (RFC 7822 §2). +func (frm Frame) ExtensionFields() []byte { + return frm.buf[SizeHeader:] +} + +// ValidateSize checks that the NTP header is complete and that any extension +// fields are well-formed with valid lengths. +func (frm Frame) ValidateSize(v *lneto.Validator) { + if len(frm.buf) < SizeHeader { + v.AddError(lneto.ErrTruncatedFrame) + return + } + buf := frm.ExtensionFields() + for len(buf) > 0 { + _, n, err := NextExtField(buf) + if err != nil { + v.AddError(err) + return + } + buf = buf[n:] + } +} + // ClearHeader zeros out the header contents. func (frm Frame) ClearHeader() { for i := range frm.buf[:SizeHeader] { @@ -216,6 +244,10 @@ func (t Timestamp) Seconds() uint32 { return t.sec } func (t Timestamp) Fractions() uint32 { return t.fra } +// Uint64 returns the full 64-bit NTP timestamp with seconds in the upper 32 +// bits and fractions in the lower 32 bits. Suitable for logging and encoding. +func (t Timestamp) Uint64() uint64 { return uint64(t.sec)<<32 | uint64(t.fra) } + func (t Short) Seconds() uint16 { return uint16(t >> 16) } func (t Short) Fractions() uint16 { return uint16(t) } diff --git a/ntp/server.go b/ntp/server.go new file mode 100644 index 0000000..d217d8d --- /dev/null +++ b/ntp/server.go @@ -0,0 +1,133 @@ +package ntp + +import ( + "time" + + "github.com/soypat/lneto" +) + +// ServerConfig configures an NTP [Server]. +type ServerConfig struct { + Now func() time.Time + Stratum Stratum + Precision int8 + RefID [4]byte + MaxPending int +} + +// Server is a basic NTP server implementing [lneto.StackNode]. +// It receives client requests via [Server.Demux] and builds server +// responses via [Server.Encapsulate]. +// +// Server is not safe for concurrent use. +type Server struct { + connID uint64 + _now func() time.Time + stratum Stratum + prec int8 + refID [4]byte + pending []pendingRequest +} + +type pendingRequest struct { + origin Timestamp +} + +// Reset re-initialises the server with cfg. Increments connID. +func (h *Server) Reset(cfg ServerConfig) error { + if cfg.Now == nil { + return lneto.ErrInvalidConfig + } + if cfg.MaxPending <= 0 { + cfg.MaxPending = 4 + } + pending := h.pending[:0] + if cap(pending) < cfg.MaxPending { + pending = make([]pendingRequest, 0, cfg.MaxPending) + } + *h = Server{ + connID: h.connID + 1, + _now: cfg.Now, + stratum: cfg.Stratum, + prec: cfg.Precision, + refID: cfg.RefID, + pending: pending, + } + return nil +} + +// ConnectionID implements [lneto.StackNode]. +func (h *Server) ConnectionID() *uint64 { return &h.connID } + +// Protocol implements [lneto.StackNode]. +func (h *Server) Protocol() uint64 { return 0 } + +// LocalPort implements [lneto.StackNode]. +func (h *Server) LocalPort() uint16 { return ServerPort } + +// Encapsulate implements [lneto.StackNode]. It writes one pending NTP server +// response into carrierData. Returns 0 when no pending requests exist. +func (h *Server) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + if len(h.pending) == 0 { + return 0, nil + } + buf := carrierData[offsetToFrame:] + frm, err := NewFrame(buf) + if err != nil { + return 0, err + } + + req := h.pending[len(h.pending)-1] + h.pending = h.pending[:len(h.pending)-1] + + now := h.now() + xmt, err := TimestampFromTime(now) + if err != nil { + return 0, err + } + + frm.ClearHeader() + frm.SetFlags(ModeServer, Version4, LeapNoWarning) + frm.SetStratum(h.stratum) + frm.SetPrecision(h.prec) + frm.SetPoll(6) + *frm.ReferenceID() = h.refID + frm.SetOriginTime(req.origin) + frm.SetReceiveTime(xmt) + frm.SetTransmitTime(xmt) + return SizeHeader, nil +} + +// Demux implements [lneto.StackNode]. It validates an incoming NTP client +// request and queues it for response via [Server.Encapsulate]. +func (h *Server) Demux(carrierData []byte, frameOffset int) error { + buf := carrierData[frameOffset:] + frm, err := NewFrame(buf) + if err != nil { + return err + } + + mode, version, _ := frm.Flags() + if mode != ModeClient { + return lneto.ErrPacketDrop + } + if version != Version4 { + return lneto.ErrPacketDrop + } + + if len(h.pending) == cap(h.pending) { + return lneto.ErrExhausted + } + + h.pending = append(h.pending, pendingRequest{ + origin: frm.TransmitTime(), + }) + return nil +} + +func (h *Server) now() time.Time { + if h._now == nil { + return time.Now() + } + return h._now() +} diff --git a/ntp/server_test.go b/ntp/server_test.go new file mode 100644 index 0000000..da4adeb --- /dev/null +++ b/ntp/server_test.go @@ -0,0 +1,193 @@ +package ntp + +import ( + "testing" + "time" +) + +func TestServer_BasicExchange(t *testing.T) { + serverTime := BaseTime().Add(100 * time.Second) + var h Server + err := h.Reset(ServerConfig{ + Now: func() time.Time { return serverTime }, + Stratum: StratumPrimary, + Precision: -20, + RefID: [4]byte{'G', 'P', 'S', 0}, + }) + if err != nil { + t.Fatal(err) + } + + reqBuf := make([]byte, SizeHeader) + frm, _ := NewFrame(reqBuf) + frm.SetFlags(ModeClient, Version4, LeapNoWarning) + frm.SetStratum(StratumUnsync) + clientXmt := TimestampFromUint64(0x12345678_9abcdef0) + frm.SetTransmitTime(clientXmt) + + if err := h.Demux(reqBuf, 0); err != nil { + t.Fatal(err) + } + + respBuf := make([]byte, SizeHeader) + n, err := h.Encapsulate(respBuf, 0, 0) + if err != nil { + t.Fatal(err) + } + if n != SizeHeader { + t.Fatalf("expected %d bytes, got %d", SizeHeader, n) + } + + resp, _ := NewFrame(respBuf) + mode, version, _ := resp.Flags() + if mode != ModeServer { + t.Errorf("response mode = %d; want ModeServer", mode) + } + if version != Version4 { + t.Errorf("response version = %d; want 4", version) + } + if resp.Stratum() != StratumPrimary { + t.Errorf("response stratum = %s; want primary", resp.Stratum()) + } + if resp.OriginTime() != clientXmt { + t.Error("response origin time does not echo client transmit time (RFC 5905 §8)") + } + if resp.ReferenceID() == nil || *resp.ReferenceID() != [4]byte{'G', 'P', 'S', 0} { + t.Error("reference ID mismatch") + } +} + +func TestServer_RejectsNonClient(t *testing.T) { + modes := []struct { + name string + mode Mode + }{ + {name: "server", mode: ModeServer}, + {name: "broadcast", mode: ModeBroadcast}, + {name: "symmetric_active", mode: ModeSymmetricActive}, + {name: "symmetric_passive", mode: ModeSymmetricPassive}, + } + for _, tc := range modes { + t.Run(tc.name, func(t *testing.T) { + var h Server + h.Reset(ServerConfig{ + Now: time.Now, + Stratum: StratumPrimary, + }) + reqBuf := make([]byte, SizeHeader) + frm, _ := NewFrame(reqBuf) + frm.SetFlags(tc.mode, Version4, LeapNoWarning) + if err := h.Demux(reqBuf, 0); err == nil { + t.Errorf("Server.Demux(mode=%d) = nil; want error", tc.mode) + } + }) + } +} + +func TestServer_ExhaustedPending(t *testing.T) { + var h Server + h.Reset(ServerConfig{ + Now: time.Now, + Stratum: StratumPrimary, + MaxPending: 1, + }) + + reqBuf := make([]byte, SizeHeader) + frm, _ := NewFrame(reqBuf) + frm.SetFlags(ModeClient, Version4, LeapNoWarning) + + if err := h.Demux(reqBuf, 0); err != nil { + t.Fatal(err) + } + if err := h.Demux(reqBuf, 0); err == nil { + t.Fatal("expected exhausted error on second request") + } +} + +func TestServer_NoPendingReturnsZero(t *testing.T) { + var h Server + h.Reset(ServerConfig{ + Now: time.Now, + Stratum: StratumPrimary, + }) + + buf := make([]byte, SizeHeader) + n, err := h.Encapsulate(buf, 0, 0) + if err != nil { + t.Fatal(err) + } + if n != 0 { + t.Fatalf("expected 0 bytes when no pending, got %d", n) + } +} + +func TestServer_ClientServerRoundTrip(t *testing.T) { + baseTime := BaseTime() + clientStart := baseTime.Add(10 * time.Second) + serverOffset := 500 * time.Millisecond + clockTime := clientStart + + var client Client + client.Reset(-18, func() time.Time { return clockTime }) + + serverTime := clientStart.Add(serverOffset) + var server Server + server.Reset(ServerConfig{ + Now: func() time.Time { return serverTime }, + Stratum: StratumPrimary, + Precision: -20, + RefID: [4]byte{'G', 'P', 'S', 0}, + }) + + for exchange := range 2 { + reqBuf := make([]byte, SizeHeader) + n, err := client.Encapsulate(reqBuf, 0, 0) + if err != nil || n == 0 { + t.Fatalf("exchange %d: Encapsulate: n=%d err=%v", exchange, n, err) + } + + if err = server.Demux(reqBuf[:n], 0); err != nil { + t.Fatalf("exchange %d: server Demux: %v", exchange, err) + } + + respBuf := make([]byte, SizeHeader) + n, err = server.Encapsulate(respBuf, 0, 0) + if err != nil || n == 0 { + t.Fatalf("exchange %d: server Encapsulate: n=%d err=%v", exchange, n, err) + } + + clockTime = clockTime.Add(100 * time.Millisecond) + serverTime = serverTime.Add(100 * time.Millisecond) + + if err = client.Demux(respBuf[:n], 0); err != nil { + t.Fatalf("exchange %d: client Demux: %v", exchange, err) + } + } + + if !client.IsDone() { + t.Fatal("client should be done after two exchanges") + } + if client.RoundTripDelay() < 0 { + t.Errorf("RTD = %v; want >= 0", client.RoundTripDelay()) + } +} + +func FuzzServerDemux(f *testing.F) { + valid := make([]byte, SizeHeader) + frm, _ := NewFrame(valid) + frm.SetFlags(ModeClient, Version4, LeapNoWarning) + f.Add(valid) + f.Add(make([]byte, SizeHeader)) + f.Add([]byte{}) + f.Add(make([]byte, 10)) + f.Fuzz(func(t *testing.T, data []byte) { + var h Server + h.Reset(ServerConfig{ + Now: time.Now, + Stratum: StratumPrimary, + }) + _ = h.Demux(data, 0) + buf := make([]byte, SizeHeader) + h.Encapsulate(buf, 0, 0) + }) +}