mirror of
https://github.com/soypat/lneto.git
synced 2026-09-11 00:59:30 +00:00
add dhcp failing test
This commit is contained in:
+28
-3
@@ -10,6 +10,7 @@ import (
|
||||
"net"
|
||||
|
||||
"github.com/soypat/lneto"
|
||||
"github.com/soypat/lneto/ipv4"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
@@ -63,6 +64,24 @@ func (c *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) }
|
||||
func (c *Client) LocalPort() uint16 { return DefaultClientPort }
|
||||
func (c *Client) ConnectionID() *uint64 { return &c.connID }
|
||||
|
||||
func (c *Client) setIP(b []byte, frameOffset int) {
|
||||
if frameOffset < 28 {
|
||||
return // Not an IP/UDP frame.
|
||||
}
|
||||
ifrm, _ := ipv4.NewFrame(b)
|
||||
ifrm.SetID((uint16(c.currentXID) ^ uint16(c.currentXID>>16)) + uint16(c.state))
|
||||
if c.state > StateInit {
|
||||
// TODO(soypat): Document why disabling ToS used by DHCP server may cause Request to fail.
|
||||
// Apparently server sets ToS=192. Uncommenting this line causes DHCP to fail on my setup.
|
||||
// If left fixed at 192, DHCP does not work.
|
||||
// If left fixed at 0, DHCP does not work.
|
||||
// Apparently ToS is a function of which state of DHCP one is in. Not sure why code below works.
|
||||
// Note: Not exactly needed for all servers.
|
||||
const ecnmask = 0b1100_0000
|
||||
ifrm.SetToS(ecnmask)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) {
|
||||
if c.isClosed() {
|
||||
return 0, net.ErrClosed
|
||||
@@ -98,6 +117,9 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error)
|
||||
nextState = StateSelecting
|
||||
|
||||
case StateSelecting:
|
||||
if c.offer == ([4]byte{}) {
|
||||
return 0, nil // Offer not yet received.
|
||||
}
|
||||
// Send out request, we know we've received an offer by now.
|
||||
optBuf = AppendOption(optBuf, OptMessageType, byte(MsgRequest))
|
||||
optBuf = AppendOption(optBuf, OptRequestedIPaddress, c.offer[:]...)
|
||||
@@ -108,8 +130,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error)
|
||||
return 0, errors.New("unhandled state")
|
||||
}
|
||||
if len(c.reqHostname) > 0 {
|
||||
optBuf = append(optBuf, byte(OptHostName), byte(len(c.hostname)))
|
||||
optBuf = append(optBuf, c.hostname...)
|
||||
optBuf = AppendOptionString(optBuf, OptHostName, c.reqHostname)
|
||||
}
|
||||
optBuf = append(optBuf, 0xff) // End mark.
|
||||
options := frm.OptionsPayload()
|
||||
@@ -118,6 +139,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error)
|
||||
}
|
||||
c.setHeader(frm)
|
||||
n := copy(options, optBuf)
|
||||
c.setIP(carrierFrame, frameOffset)
|
||||
c.state = nextState
|
||||
return optionsOffset + n, nil
|
||||
}
|
||||
@@ -155,6 +177,7 @@ func (c *Client) Demux(carrierData []byte, frameOffset int) error {
|
||||
// Lock in on this offer.
|
||||
c.gateway = *frm.GIAddr()
|
||||
c.offer = *frm.YIAddr()
|
||||
c.svip = *frm.SIAddr()
|
||||
}
|
||||
|
||||
case StateRequesting:
|
||||
@@ -223,7 +246,9 @@ func (c *Client) setHeader(frm Frame) {
|
||||
frm.SetXID(c.currentXID)
|
||||
frm.SetHardware(1, 6, 0)
|
||||
frm.SetSecs(1)
|
||||
// copy(frm.CIAddr()[:], c.offer[:])
|
||||
if c.state == StateBound {
|
||||
// copy(frm.CIAddr()[:], c.offer[:])
|
||||
}
|
||||
copy(frm.SIAddr()[:], c.svip[:])
|
||||
copy(frm.YIAddr()[:], c.offer[:])
|
||||
copy(frm.CHAddrAs6()[:], c.clientMAC[:])
|
||||
|
||||
@@ -2,6 +2,7 @@ package dhcpv4
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
//go:generate stringer -type=OptNum,Op,MessageType,ClientState -linecomment -output stringers.go
|
||||
@@ -38,6 +39,11 @@ func AppendOption(dst []byte, opt OptNum, data ...byte) []byte {
|
||||
return dst
|
||||
}
|
||||
|
||||
func AppendOptionString(dst []byte, opt OptNum, data string) []byte {
|
||||
bdata := unsafe.Slice(unsafe.StringData(data), len(data))
|
||||
return AppendOption(dst, opt, bdata...)
|
||||
}
|
||||
|
||||
func EncodeOption(dst []byte, opt OptNum, data ...byte) (int, error) {
|
||||
if len(data) > 255 {
|
||||
return 0, errors.New("DHCPv4 option data too long (>255)")
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
package dhcpv4
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestClientServer(t *testing.T) {
|
||||
svAddr := [4]byte{192, 168, 1, 1}
|
||||
clAddr := svAddr
|
||||
clAddr[3]++
|
||||
var sv Server
|
||||
var cl Client
|
||||
err := cl.BeginRequest(123, RequestConfig{
|
||||
RequestedAddr: clAddr,
|
||||
ClientHardwareAddr: [6]byte{1, 2, 3, 4, 5, 6},
|
||||
Hostname: "lneto",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertClState := func(state ClientState) {
|
||||
t.Helper()
|
||||
if state != cl.State() {
|
||||
t.Errorf("want client state %s, got %s", state.String(), cl.State().String())
|
||||
}
|
||||
}
|
||||
sv.Reset(svAddr, DefaultServerPort)
|
||||
// CLIENT DISCOVER.
|
||||
assertClState(StateInit)
|
||||
var buf [1024]byte
|
||||
n, err := cl.Encapsulate(buf[:], 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
} else if n == 0 {
|
||||
t.Fatal("no data exchanged")
|
||||
}
|
||||
assertClState(StateSelecting)
|
||||
err = sv.Demux(buf[:n], 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// SERVER REPLY OFFER
|
||||
n, err = sv.Encapsulate(buf[:], 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
} else if n == 0 {
|
||||
t.Fatal("no data exchanged")
|
||||
}
|
||||
err = cl.Demux(buf[:n], 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertClState(StateRequesting)
|
||||
|
||||
// CLIENT SEND OUT ACK.
|
||||
n, err = cl.Encapsulate(buf[:], 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
} else if n == 0 {
|
||||
t.Fatal("no data exchanged")
|
||||
}
|
||||
err = sv.Demux(buf[:n], 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
+38
-11
@@ -3,6 +3,8 @@ package dhcpv4
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
|
||||
"github.com/soypat/lneto"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -25,7 +27,7 @@ const (
|
||||
// An error is returned if the buffer size is smaller than 240.
|
||||
func NewFrame(buf []byte) (Frame, error) {
|
||||
if len(buf) < optionsOffset {
|
||||
return Frame{}, errors.New("DHCPv4 short frame")
|
||||
return Frame{}, errSmallFrame
|
||||
}
|
||||
return Frame{buf: buf}, nil
|
||||
}
|
||||
@@ -41,7 +43,7 @@ type Frame struct {
|
||||
|
||||
// OptionsPayload returns the options portion of the DHCP frame. May be zero lengthed.
|
||||
func (frm Frame) OptionsPayload() []byte {
|
||||
return frm.buf[:optionsOffset]
|
||||
return frm.buf[optionsOffset:]
|
||||
}
|
||||
|
||||
func (frm Frame) Op() Op { return Op(frm.buf[0]) }
|
||||
@@ -97,7 +99,10 @@ func (frm Frame) CHAddr() *[16]byte {
|
||||
return (*[16]byte)(frm.buf[28:44])
|
||||
}
|
||||
|
||||
// MagicCookie returns the magic cookie of the header. Expect this to always be [MagicCookie].
|
||||
func (frm Frame) MagicCookie() uint32 { return binary.BigEndian.Uint32(frm.buf[magicCookieOffset:]) }
|
||||
|
||||
// SetMagicCookie sets the MagicCookie. Call this with [MagicCookie] to create a valid DHCP header.
|
||||
func (frm Frame) SetMagicCookie(cookie uint32) {
|
||||
binary.BigEndian.PutUint32(frm.buf[magicCookieOffset:], cookie)
|
||||
}
|
||||
@@ -109,18 +114,20 @@ func (frm Frame) ClearHeader() {
|
||||
}
|
||||
}
|
||||
|
||||
// ForEachOption iterates over all DHCPv4 options returning an error on a malformed option or when user provided callback returns an error.
|
||||
// If the user provided callback is nil then only option buffer validation is performed.
|
||||
func (frm Frame) ForEachOption(fn func(op OptNum, data []byte) error) error {
|
||||
if fn == nil {
|
||||
return errors.New("nil function to parse DHCP")
|
||||
}
|
||||
// Parse DHCP options.
|
||||
ptr := optionsOffset
|
||||
if ptr >= len(frm.buf) {
|
||||
return errors.New("short payload to parse DHCP options")
|
||||
if ptr > len(frm.buf) {
|
||||
return errSmallFrame
|
||||
} else if len(frm.buf[ptr:]) == 0 {
|
||||
return errNoOptions
|
||||
}
|
||||
callback := fn != nil
|
||||
for ptr+1 < len(frm.buf) {
|
||||
if int(frm.buf[ptr+1]) >= len(frm.buf) {
|
||||
return errors.New("DHCP option length exceeds payload")
|
||||
return errDHCPBadOption
|
||||
}
|
||||
optnum := OptNum(frm.buf[ptr])
|
||||
if optnum == 0xff {
|
||||
@@ -130,11 +137,31 @@ func (frm Frame) ForEachOption(fn func(op OptNum, data []byte) error) error {
|
||||
continue
|
||||
}
|
||||
optlen := frm.buf[ptr+1]
|
||||
optionData := frm.buf[ptr+2 : ptr+2+int(optlen)]
|
||||
if err := fn(optnum, optionData); err != nil {
|
||||
return err
|
||||
if callback {
|
||||
optionData := frm.buf[ptr+2 : ptr+2+int(optlen)]
|
||||
if err := fn(optnum, optionData); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
ptr += int(optlen) + 2
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
//
|
||||
// Validation API.
|
||||
//
|
||||
|
||||
var (
|
||||
errSmallFrame = errors.New("DHCPv4: frame size <240")
|
||||
errDHCPBadOption = errors.New("DHCPv4: opt length exceeds payload")
|
||||
errNoOptions = errors.New("DHCPv4: no options")
|
||||
errOptionNotFit = errors.New("DHCPv4: options dont fit")
|
||||
)
|
||||
|
||||
func (frm Frame) ValidateSize(vld *lneto.Validator) {
|
||||
err := frm.ForEachOption(nil) // Does all necessary validation.
|
||||
if err != nil {
|
||||
vld.AddError(errDHCPBadOption)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,239 @@
|
||||
package dhcpv4
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
|
||||
"github.com/soypat/lneto"
|
||||
"github.com/soypat/lneto/internal"
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
connID uint64
|
||||
nextAddr netip.Addr
|
||||
prefix netip.Prefix
|
||||
hosts map[[36]byte]serverEntry
|
||||
vld lneto.Validator
|
||||
pending int
|
||||
port uint16
|
||||
siaddr [4]byte
|
||||
gwaddr [4]byte
|
||||
}
|
||||
|
||||
type serverEntry struct {
|
||||
hostname string
|
||||
xid uint32
|
||||
port uint16
|
||||
addr [4]byte
|
||||
requestlist [10]byte
|
||||
hwaddr [6]byte
|
||||
clientIdlen uint8
|
||||
// Possible states:
|
||||
// - 0: No entry/uninitialized
|
||||
// - Init: Server received discover, pending Offer sent out.
|
||||
// - Selecting: Server sent out offer, request not received.
|
||||
// - Requesting: Request received, pending Ack sent out.
|
||||
// - Bound: Request sent out, no more pending data to be sent.
|
||||
state ClientState
|
||||
}
|
||||
|
||||
func (sv *Server) Reset(serverAddr [4]byte, port uint16) {
|
||||
*sv = Server{
|
||||
connID: sv.connID + 1,
|
||||
siaddr: serverAddr,
|
||||
port: port,
|
||||
hosts: sv.hosts,
|
||||
nextAddr: netip.AddrFrom4(serverAddr),
|
||||
}
|
||||
if sv.hosts == nil {
|
||||
sv.hosts = make(map[[36]byte]serverEntry)
|
||||
} else {
|
||||
for k := range sv.hosts {
|
||||
delete(sv.hosts, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (sv *Server) ConnectionID() *uint64 { return &sv.connID }
|
||||
func (sv *Server) Protocol() uint64 { return uint64(lneto.IPProtoUDP) }
|
||||
func (sv *Server) Port() uint16 { return sv.port }
|
||||
|
||||
func (sv *Server) Demux(carrierData []byte, frameOffset int) error {
|
||||
isIPLayer := frameOffset >= 28
|
||||
dhcpData := carrierData[frameOffset:]
|
||||
dfrm, err := NewFrame(dhcpData)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dfrm.ValidateSize(&sv.vld)
|
||||
if sv.vld.HasError() {
|
||||
return sv.vld.ErrPop()
|
||||
}
|
||||
|
||||
var msgType MessageType
|
||||
var clientID []byte
|
||||
var reqlist []byte
|
||||
var reqAddr []byte
|
||||
var hostname []byte
|
||||
err = dfrm.ForEachOption(func(op OptNum, data []byte) error {
|
||||
switch op {
|
||||
case OptMessageType:
|
||||
if len(data) == 1 {
|
||||
msgType = MessageType(data[0])
|
||||
}
|
||||
case OptHostName:
|
||||
if len(data) <= 36 {
|
||||
hostname = data
|
||||
}
|
||||
case OptClientIdentifier:
|
||||
if len(data) <= 36 {
|
||||
clientID = data
|
||||
}
|
||||
case OptParameterRequestList:
|
||||
if len(data) > 36 {
|
||||
return errors.New("too many request options")
|
||||
}
|
||||
reqlist = data
|
||||
case OptRequestedIPaddress:
|
||||
if len(data) == 4 {
|
||||
reqAddr = data
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
var clientIDRaw [36]byte
|
||||
var client serverEntry
|
||||
var clientExists bool
|
||||
if len(clientID) == 0 {
|
||||
client, clientIDRaw, clientExists = sv.getClientByIP(*dfrm.CIAddr())
|
||||
} else {
|
||||
copy(clientIDRaw[:], clientID)
|
||||
client, clientExists = sv.getClient(clientIDRaw)
|
||||
}
|
||||
|
||||
switch msgType {
|
||||
case MsgDiscover:
|
||||
if clientExists {
|
||||
err = errors.New("DHCP Discover on initialized client")
|
||||
break
|
||||
}
|
||||
if len(reqAddr) == 4 {
|
||||
println("requested", reqAddr[0], reqAddr[1], reqAddr[2], reqAddr[3])
|
||||
}
|
||||
sv.nextAddr = sv.nextAddr.Next()
|
||||
copy(client.requestlist[:], reqlist)
|
||||
client.addr = sv.nextAddr.As4()
|
||||
client.state = StateInit
|
||||
client.hostname = string(hostname)
|
||||
client.xid = dfrm.XID()
|
||||
client.hwaddr = *dfrm.CHAddrAs6()
|
||||
if isIPLayer {
|
||||
_, client.port, _ = getSrcIPPort(carrierData)
|
||||
}
|
||||
client.clientIdlen = uint8(len(clientID))
|
||||
sv.pending++
|
||||
|
||||
case MsgRequest:
|
||||
if client.state != StateSelecting && client.state != StateRequesting {
|
||||
err = errors.New("DHCP request unexpected state")
|
||||
break
|
||||
}
|
||||
client.state = StateBound
|
||||
sv.pending++
|
||||
|
||||
default:
|
||||
err = errors.New("unhandled message type")
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("msgtype=%s client=%+v: %w", msgType.String(), client, err)
|
||||
}
|
||||
sv.hosts[clientIDRaw] = client
|
||||
return nil
|
||||
// n := copy(dfrm.OptionsPayload(), optBuf)
|
||||
}
|
||||
|
||||
func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) {
|
||||
carrierIsIP := frameOffset >= 28
|
||||
dfrm, err := NewFrame(carrierData[frameOffset:])
|
||||
optBuf := dfrm.OptionsPayload()[:0]
|
||||
if err != nil {
|
||||
return 0, err
|
||||
} else if cap(optBuf) < 255 {
|
||||
return 0, errOptionNotFit
|
||||
}
|
||||
if sv.pending == 0 {
|
||||
return 0, nil // No pending outgoing frames.a
|
||||
}
|
||||
|
||||
var client serverEntry
|
||||
var clientID [36]byte
|
||||
for k, v := range sv.hosts {
|
||||
pending := v.state == StateInit || v.state == StateRequesting
|
||||
if pending {
|
||||
client = v
|
||||
clientID = k
|
||||
break
|
||||
}
|
||||
}
|
||||
if client.state == 0 {
|
||||
return 0, nil // Nothing to do.
|
||||
}
|
||||
futureState := ClientState(0)
|
||||
switch client.state {
|
||||
case StateInit:
|
||||
futureState = StateSelecting
|
||||
optBuf = AppendOption(optBuf, OptMessageType, byte(MsgOffer))
|
||||
case StateRequesting:
|
||||
futureState = StateBound
|
||||
optBuf = AppendOption(optBuf, OptMessageType, byte(MsgAck))
|
||||
*dfrm.CIAddr() = client.addr
|
||||
}
|
||||
|
||||
dfrm.ClearHeader()
|
||||
dfrm.SetOp(OpReply)
|
||||
dfrm.SetHardware(1, 6, 0)
|
||||
dfrm.SetXID(client.xid)
|
||||
dfrm.SetSecs(0)
|
||||
dfrm.SetFlags(0)
|
||||
*dfrm.YIAddr() = client.addr // Offer here.
|
||||
*dfrm.SIAddr() = sv.siaddr
|
||||
*dfrm.GIAddr() = sv.gwaddr
|
||||
copy(dfrm.CHAddrAs6()[:], client.hwaddr[:])
|
||||
dfrm.SetMagicCookie(MagicCookie)
|
||||
if carrierIsIP {
|
||||
internal.SetIPDestinationAddr(carrierData, 0, client.addr[:])
|
||||
}
|
||||
client.state = futureState
|
||||
|
||||
// Set server state.
|
||||
sv.hosts[clientID] = client
|
||||
sv.pending--
|
||||
return optionsOffset + len(optBuf), nil
|
||||
}
|
||||
|
||||
func (sv *Server) getClient(clientID [36]byte) (serverEntry, bool) {
|
||||
entry, ok := sv.hosts[clientID]
|
||||
return entry, ok
|
||||
}
|
||||
|
||||
func (sv *Server) getClientByIP(ip [4]byte) (serverEntry, [36]byte, bool) {
|
||||
for k, v := range sv.hosts {
|
||||
if v.addr == ip {
|
||||
return v, k, true
|
||||
}
|
||||
}
|
||||
return serverEntry{}, [36]byte{}, false
|
||||
}
|
||||
|
||||
func getSrcIPPort(ipCarrier []byte) (addr []byte, port uint16, err error) {
|
||||
addr, _, off, err := internal.GetIPSourceAddr(ipCarrier)
|
||||
if err != nil {
|
||||
return addr, port, err
|
||||
} else if len(ipCarrier[off:]) < 2 {
|
||||
return addr, port, errors.New("getSrcIPPort got only IP layer")
|
||||
}
|
||||
port = binary.BigEndian.Uint16(ipCarrier[off:]) // TCP and UDP share same port offsets.
|
||||
return addr, port, nil
|
||||
}
|
||||
Reference in New Issue
Block a user