mirror of
https://github.com/soypat/lneto.git
synced 2026-07-26 10:38:47 +00:00
feat(ntp): add two-exchange client, server handler, and extension fields (#101)
Implement full two-exchange NTP client state machine (RFC 5905 §8). First exchange stores offset/RTT, second averages both for improved accuracy. Client places T1 in TransmitTime per RFC 5905 §8; server echoes it as OriginTime in the response. Add NTP extension field codec (RFC 7822) with NextExtField iterator and AppendExtField builder. Add NTS extension type constants from RFC 8915. Add NTP Server (StackNode) that receives client requests via Demux and builds server responses via Encapsulate with configurable stratum, precision, reference ID, and pending request queue. Add Frame accessor methods: RawData(), ExtensionFields(), ValidateSize(), and Timestamp.Uint64(). Add ntp-client and ntp-server example programs with CalculateSystemPrecision, retry limits, and backoff. Use internal.LogAttrs pattern for non-allocating structured logging in the client, matching the tcp/debug.go convention. Generated with Claude assistance. Signed-off-by: Marvin Drees <marvin.drees@9elements.com>
This commit is contained in:
@@ -0,0 +1,2 @@
|
|||||||
|
ignore:
|
||||||
|
- "examples/**"
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+75
-31
@@ -1,9 +1,11 @@
|
|||||||
package ntp
|
package ntp
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"log/slog"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/soypat/lneto"
|
"github.com/soypat/lneto"
|
||||||
|
"github.com/soypat/lneto/internal"
|
||||||
)
|
)
|
||||||
|
|
||||||
type state uint8
|
type state uint8
|
||||||
@@ -23,14 +25,15 @@ type Client struct {
|
|||||||
connID uint64
|
connID uint64
|
||||||
start time.Time
|
start time.Time
|
||||||
_now func() time.Time
|
_now func() time.Time
|
||||||
// t stores the time offsets needed to compute the time at client
|
logger logger
|
||||||
// taking into consideration the round-trip delay.
|
// t stores the four NTP timestamps per RFC 5905:
|
||||||
// - t[0] (orig): Client timestamp of request packet transmission.
|
// - t[0] (T1): Client timestamp of request packet transmission.
|
||||||
// - t[1] (rec): Server timestamp of request packet reception.
|
// - t[1] (T2): Server timestamp of request packet reception.
|
||||||
// - t[2] (xmt): Server timestamp of response packet transmission.
|
// - t[2] (T3): Server timestamp of response packet transmission.
|
||||||
// - t[3]: Client timestamp of response packet reception.
|
// - t[3] (T4): Client timestamp of response packet reception.
|
||||||
t [4]Timestamp
|
t [4]Timestamp
|
||||||
// org 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
|
state state
|
||||||
serverStratum Stratum
|
serverStratum Stratum
|
||||||
sysprec int8
|
sysprec int8
|
||||||
@@ -40,11 +43,16 @@ func (c *Client) Reset(sysprec int8, now func() time.Time) {
|
|||||||
*c = Client{
|
*c = Client{
|
||||||
connID: c.connID + 1,
|
connID: c.connID + 1,
|
||||||
_now: now,
|
_now: now,
|
||||||
|
logger: c.logger,
|
||||||
sysprec: sysprec,
|
sysprec: sysprec,
|
||||||
state: stateSend1,
|
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) Protocol() uint64 { return 0 }
|
||||||
func (c *Client) LocalPort() uint16 { return ClientPort }
|
func (c *Client) LocalPort() uint16 { return ClientPort }
|
||||||
func (c *Client) ConnectionID() *uint64 {
|
func (c *Client) ConnectionID() *uint64 {
|
||||||
@@ -62,27 +70,32 @@ func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int)
|
|||||||
}
|
}
|
||||||
|
|
||||||
switch c.state {
|
switch c.state {
|
||||||
case stateSend1:
|
case stateSend1, stateSend2:
|
||||||
c.start = c.now()
|
now := c.now()
|
||||||
c.t[0] = TimestampFromUint64(0)
|
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
|
c.state = stateAwait1
|
||||||
case stateSend2:
|
} else {
|
||||||
// c.xmt = c.unsyncTimestamp(c.now())
|
c.state = stateAwait2
|
||||||
c.state = stateDone
|
}
|
||||||
default:
|
default:
|
||||||
return 0, nil // Nothing to handle.
|
return 0, nil // Nothing to handle.
|
||||||
}
|
}
|
||||||
|
|
||||||
for i := range payload[:SizeHeader] {
|
|
||||||
payload[i] = 0
|
|
||||||
}
|
|
||||||
|
|
||||||
frm.ClearHeader()
|
frm.ClearHeader()
|
||||||
frm.SetStratum(StratumUnsync)
|
frm.SetStratum(StratumUnsync)
|
||||||
frm.SetPoll(6)
|
frm.SetPoll(6)
|
||||||
frm.SetPrecision(c.sysprec)
|
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)
|
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
|
return SizeHeader, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -97,21 +110,45 @@ func (c *Client) Demux(carrierData []byte, frameOffset int) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
switch c.state {
|
switch c.state {
|
||||||
case stateAwait1:
|
case stateAwait1, stateAwait2:
|
||||||
|
default:
|
||||||
|
return nil // Not awaiting a response.
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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()
|
xmt := frm.TransmitTime()
|
||||||
orig := frm.OriginTime()
|
orig := frm.OriginTime()
|
||||||
if xmt == orig || orig != c.t[0] {
|
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
|
return lneto.ErrPacketDrop
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Compute T4, then derive offset θ and round-trip delay δ per RFC 5905 §8.
|
||||||
txelapsed := c.now().Sub(c.start)
|
txelapsed := c.now().Sub(c.start)
|
||||||
c.t[1] = frm.ReceiveTime()
|
c.t[1] = frm.ReceiveTime()
|
||||||
c.t[2] = xmt
|
c.t[2] = xmt
|
||||||
c.t[3] = c.t[0].Add(txelapsed)
|
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.serverStratum = frm.Stratum()
|
||||||
c.state = stateDone // TODO: add second exchange part.
|
c.offset1 = offset
|
||||||
case stateAwait2:
|
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.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
|
return nil
|
||||||
}
|
}
|
||||||
@@ -147,28 +184,35 @@ func (c *Client) Offset() time.Duration {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) offsetAndNow() (clientNow time.Time, offset time.Duration) {
|
func (c *Client) offsetAndNow() (clientNow time.Time, offset time.Duration) {
|
||||||
now := c.now()
|
return c.now(), c.OffsetUnsynced()
|
||||||
serverToBase := c.OffsetUnsynced()
|
|
||||||
clientToBase := now.Sub(BaseTime())
|
|
||||||
serverToClient := serverToBase - clientToBase
|
|
||||||
return now, serverToClient
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// OffsetUnsynced returns the absolute time offset difference between client and server clock
|
// 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 {
|
func (c *Client) OffsetUnsynced() time.Duration {
|
||||||
if c.IsDone() {
|
if c.IsDone() {
|
||||||
t := &c.t
|
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
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RoundTripDelay returns the average round-trip delay across both NTP exchanges.
|
||||||
func (c *Client) RoundTripDelay() time.Duration {
|
func (c *Client) RoundTripDelay() time.Duration {
|
||||||
if c.IsDone() {
|
if c.IsDone() {
|
||||||
d0 := c.t[3].Sub(c.t[0])
|
rtt2 := c.t[3].Sub(c.t[0]) - c.t[2].Sub(c.t[1])
|
||||||
d1 := c.t[2].Sub(c.t[1])
|
return (c.rtt1 + rtt2) / 2
|
||||||
return d0 - d1
|
|
||||||
}
|
}
|
||||||
return -1
|
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...)
|
||||||
|
}
|
||||||
|
|||||||
+106
-26
@@ -5,6 +5,32 @@ import (
|
|||||||
"time"
|
"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) {
|
func TestClient_FullExchange(t *testing.T) {
|
||||||
// Simulate a NTP client-server exchange without network.
|
// Simulate a NTP client-server exchange without network.
|
||||||
baseTime := BaseTime()
|
baseTime := BaseTime()
|
||||||
@@ -19,7 +45,7 @@ func TestClient_FullExchange(t *testing.T) {
|
|||||||
t.Fatal("client should not be done before exchange")
|
t.Fatal("client should not be done before exchange")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Step 1: Client encapsulates request.
|
// Step 1: Client encapsulates first request.
|
||||||
reqBuf := make([]byte, SizeHeader)
|
reqBuf := make([]byte, SizeHeader)
|
||||||
n, err := client.Encapsulate(reqBuf, 0, 0)
|
n, err := client.Encapsulate(reqBuf, 0, 0)
|
||||||
if err != nil {
|
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.
|
// Server receives at clientStart + serverOffset, sends response at clientStart + serverOffset + 10ms processing.
|
||||||
serverRecvTime := clientStart.Add(serverOffset)
|
serverRecvTime := clientStart.Add(serverOffset)
|
||||||
serverXmtTime := serverRecvTime.Add(10 * time.Millisecond)
|
serverXmtTime := serverRecvTime.Add(10 * time.Millisecond)
|
||||||
|
respBuf := simulateServerResponse(t, reqBuf, serverRecvTime, serverXmtTime)
|
||||||
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)
|
|
||||||
|
|
||||||
// Advance client clock to simulate network delay.
|
// Advance client clock to simulate network delay.
|
||||||
clockTime = clientStart.Add(100 * time.Millisecond)
|
clockTime = clientStart.Add(100 * time.Millisecond)
|
||||||
|
|
||||||
// Step 3: Client demuxes response.
|
// Step 3: Client demuxes first response.
|
||||||
err = client.Demux(respBuf, 0)
|
err = client.Demux(respBuf, 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !client.IsDone() {
|
if client.IsDone() {
|
||||||
t.Fatal("client should be done after exchange")
|
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 {
|
if client.ServerStratum() != StratumPrimary {
|
||||||
t.Errorf("server stratum = %s; want primary", client.ServerStratum())
|
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) {
|
func TestClient_DemuxRejectsBogusResponse(t *testing.T) {
|
||||||
var c Client
|
var c Client
|
||||||
clockTime := BaseTime().Add(time.Second)
|
clockTime := BaseTime().Add(time.Second)
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+32
@@ -131,6 +131,34 @@ func (frm Frame) SetTransmitTime(rt Timestamp) {
|
|||||||
rt.Put(frm.buf[40:48])
|
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.
|
// ClearHeader zeros out the header contents.
|
||||||
func (frm Frame) ClearHeader() {
|
func (frm Frame) ClearHeader() {
|
||||||
for i := range frm.buf[:SizeHeader] {
|
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 }
|
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) Seconds() uint16 { return uint16(t >> 16) }
|
||||||
func (t Short) Fractions() uint16 { return uint16(t) }
|
func (t Short) Fractions() uint16 { return uint16(t) }
|
||||||
|
|
||||||
|
|||||||
+133
@@ -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()
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
})
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user