mirror of
https://github.com/soypat/lneto.git
synced 2026-09-11 09:09:30 +00:00
add dhcp failing test
This commit is contained in:
+28
-3
@@ -10,6 +10,7 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
|
|
||||||
"github.com/soypat/lneto"
|
"github.com/soypat/lneto"
|
||||||
|
"github.com/soypat/lneto/ipv4"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Client struct {
|
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) LocalPort() uint16 { return DefaultClientPort }
|
||||||
func (c *Client) ConnectionID() *uint64 { return &c.connID }
|
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) {
|
func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) {
|
||||||
if c.isClosed() {
|
if c.isClosed() {
|
||||||
return 0, net.ErrClosed
|
return 0, net.ErrClosed
|
||||||
@@ -98,6 +117,9 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error)
|
|||||||
nextState = StateSelecting
|
nextState = StateSelecting
|
||||||
|
|
||||||
case 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.
|
// Send out request, we know we've received an offer by now.
|
||||||
optBuf = AppendOption(optBuf, OptMessageType, byte(MsgRequest))
|
optBuf = AppendOption(optBuf, OptMessageType, byte(MsgRequest))
|
||||||
optBuf = AppendOption(optBuf, OptRequestedIPaddress, c.offer[:]...)
|
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")
|
return 0, errors.New("unhandled state")
|
||||||
}
|
}
|
||||||
if len(c.reqHostname) > 0 {
|
if len(c.reqHostname) > 0 {
|
||||||
optBuf = append(optBuf, byte(OptHostName), byte(len(c.hostname)))
|
optBuf = AppendOptionString(optBuf, OptHostName, c.reqHostname)
|
||||||
optBuf = append(optBuf, c.hostname...)
|
|
||||||
}
|
}
|
||||||
optBuf = append(optBuf, 0xff) // End mark.
|
optBuf = append(optBuf, 0xff) // End mark.
|
||||||
options := frm.OptionsPayload()
|
options := frm.OptionsPayload()
|
||||||
@@ -118,6 +139,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error)
|
|||||||
}
|
}
|
||||||
c.setHeader(frm)
|
c.setHeader(frm)
|
||||||
n := copy(options, optBuf)
|
n := copy(options, optBuf)
|
||||||
|
c.setIP(carrierFrame, frameOffset)
|
||||||
c.state = nextState
|
c.state = nextState
|
||||||
return optionsOffset + n, nil
|
return optionsOffset + n, nil
|
||||||
}
|
}
|
||||||
@@ -155,6 +177,7 @@ func (c *Client) Demux(carrierData []byte, frameOffset int) error {
|
|||||||
// Lock in on this offer.
|
// Lock in on this offer.
|
||||||
c.gateway = *frm.GIAddr()
|
c.gateway = *frm.GIAddr()
|
||||||
c.offer = *frm.YIAddr()
|
c.offer = *frm.YIAddr()
|
||||||
|
c.svip = *frm.SIAddr()
|
||||||
}
|
}
|
||||||
|
|
||||||
case StateRequesting:
|
case StateRequesting:
|
||||||
@@ -223,7 +246,9 @@ func (c *Client) setHeader(frm Frame) {
|
|||||||
frm.SetXID(c.currentXID)
|
frm.SetXID(c.currentXID)
|
||||||
frm.SetHardware(1, 6, 0)
|
frm.SetHardware(1, 6, 0)
|
||||||
frm.SetSecs(1)
|
frm.SetSecs(1)
|
||||||
// copy(frm.CIAddr()[:], c.offer[:])
|
if c.state == StateBound {
|
||||||
|
// copy(frm.CIAddr()[:], c.offer[:])
|
||||||
|
}
|
||||||
copy(frm.SIAddr()[:], c.svip[:])
|
copy(frm.SIAddr()[:], c.svip[:])
|
||||||
copy(frm.YIAddr()[:], c.offer[:])
|
copy(frm.YIAddr()[:], c.offer[:])
|
||||||
copy(frm.CHAddrAs6()[:], c.clientMAC[:])
|
copy(frm.CHAddrAs6()[:], c.clientMAC[:])
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package dhcpv4
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"unsafe"
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:generate stringer -type=OptNum,Op,MessageType,ClientState -linecomment -output stringers.go
|
//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
|
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) {
|
func EncodeOption(dst []byte, opt OptNum, data ...byte) (int, error) {
|
||||||
if len(data) > 255 {
|
if len(data) > 255 {
|
||||||
return 0, errors.New("DHCPv4 option data too long (>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 (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
|
|
||||||
|
"github.com/soypat/lneto"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -25,7 +27,7 @@ const (
|
|||||||
// An error is returned if the buffer size is smaller than 240.
|
// An error is returned if the buffer size is smaller than 240.
|
||||||
func NewFrame(buf []byte) (Frame, error) {
|
func NewFrame(buf []byte) (Frame, error) {
|
||||||
if len(buf) < optionsOffset {
|
if len(buf) < optionsOffset {
|
||||||
return Frame{}, errors.New("DHCPv4 short frame")
|
return Frame{}, errSmallFrame
|
||||||
}
|
}
|
||||||
return Frame{buf: buf}, nil
|
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.
|
// OptionsPayload returns the options portion of the DHCP frame. May be zero lengthed.
|
||||||
func (frm Frame) OptionsPayload() []byte {
|
func (frm Frame) OptionsPayload() []byte {
|
||||||
return frm.buf[:optionsOffset]
|
return frm.buf[optionsOffset:]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (frm Frame) Op() Op { return Op(frm.buf[0]) }
|
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])
|
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:]) }
|
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) {
|
func (frm Frame) SetMagicCookie(cookie uint32) {
|
||||||
binary.BigEndian.PutUint32(frm.buf[magicCookieOffset:], cookie)
|
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 {
|
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.
|
// Parse DHCP options.
|
||||||
ptr := optionsOffset
|
ptr := optionsOffset
|
||||||
if ptr >= len(frm.buf) {
|
if ptr > len(frm.buf) {
|
||||||
return errors.New("short payload to parse DHCP options")
|
return errSmallFrame
|
||||||
|
} else if len(frm.buf[ptr:]) == 0 {
|
||||||
|
return errNoOptions
|
||||||
}
|
}
|
||||||
|
callback := fn != nil
|
||||||
for ptr+1 < len(frm.buf) {
|
for ptr+1 < len(frm.buf) {
|
||||||
if int(frm.buf[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])
|
optnum := OptNum(frm.buf[ptr])
|
||||||
if optnum == 0xff {
|
if optnum == 0xff {
|
||||||
@@ -130,11 +137,31 @@ func (frm Frame) ForEachOption(fn func(op OptNum, data []byte) error) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
optlen := frm.buf[ptr+1]
|
optlen := frm.buf[ptr+1]
|
||||||
optionData := frm.buf[ptr+2 : ptr+2+int(optlen)]
|
if callback {
|
||||||
if err := fn(optnum, optionData); err != nil {
|
optionData := frm.buf[ptr+2 : ptr+2+int(optlen)]
|
||||||
return err
|
if err := fn(optnum, optionData); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
ptr += int(optlen) + 2
|
ptr += int(optlen) + 2
|
||||||
}
|
}
|
||||||
return nil
|
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
|
||||||
|
}
|
||||||
+43
-9
@@ -3,16 +3,19 @@ package main
|
|||||||
import (
|
import (
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/soypat/lneto"
|
"github.com/soypat/lneto"
|
||||||
"github.com/soypat/lneto/arp"
|
"github.com/soypat/lneto/arp"
|
||||||
"github.com/soypat/lneto/dhcpv4"
|
"github.com/soypat/lneto/dhcpv4"
|
||||||
"github.com/soypat/lneto/ethernet"
|
"github.com/soypat/lneto/ethernet"
|
||||||
|
"github.com/soypat/lneto/internal"
|
||||||
"github.com/soypat/lneto/internal/ltesto"
|
"github.com/soypat/lneto/internal/ltesto"
|
||||||
"github.com/soypat/lneto/internet"
|
"github.com/soypat/lneto/internet"
|
||||||
"github.com/soypat/lneto/internet/pcap"
|
"github.com/soypat/lneto/internet/pcap"
|
||||||
@@ -28,16 +31,47 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func run() (err error) {
|
func run() (err error) {
|
||||||
br := ltesto.NewHTTPTapClient("http://127.0.0.1:7070")
|
var (
|
||||||
defer br.Close()
|
flagInterface = "tap0"
|
||||||
|
flagUseHTTP = false
|
||||||
nicHW := br.HardwareAddr6()
|
)
|
||||||
|
flag.StringVar(&flagInterface, "i", flagInterface, "Interface to use. Either tap* or the name of an existing interface to bridge to.")
|
||||||
|
flag.BoolVar(&flagUseHTTP, "http", flagUseHTTP, "Use HTTP tap interface.")
|
||||||
|
flag.Parse()
|
||||||
|
var iface ltesto.Interface
|
||||||
|
if flagUseHTTP {
|
||||||
|
iface = ltesto.NewHTTPTapClient("http://127.0.0.1:7070")
|
||||||
|
} else {
|
||||||
|
if strings.HasPrefix(flagInterface, "tap") {
|
||||||
|
tap, err := internal.NewTap(flagInterface, netip.MustParsePrefix("192.168.1.1/24"))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
iface = tap
|
||||||
|
} else {
|
||||||
|
bridge, err := internal.NewBridge(flagInterface)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
iface = bridge
|
||||||
|
}
|
||||||
|
}
|
||||||
|
defer iface.Close()
|
||||||
|
|
||||||
|
nicHW, err := iface.HardwareAddress6()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
brHW := nicHW
|
brHW := nicHW
|
||||||
brHW[5]++ // We'll be using a similar HW address but with NIC specific identifier modified.
|
brHW[5]++ // We'll be using a similar HW address but with NIC specific identifier modified.
|
||||||
mtu := br.MTU()
|
mtu, err := iface.MTU()
|
||||||
nicAddr := br.IPPrefix()
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
nicAddr, err := iface.IPMask()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
fmt.Println("NIC hardware address:", net.HardwareAddr(nicHW[:]).String(), "bridgeHW:", net.HardwareAddr(brHW[:]).String(), "mtu:", mtu, "addr:", nicAddr.String())
|
fmt.Println("NIC hardware address:", net.HardwareAddr(nicHW[:]).String(), "bridgeHW:", net.HardwareAddr(brHW[:]).String(), "mtu:", mtu, "addr:", nicAddr.String())
|
||||||
var stack Stack
|
var stack Stack
|
||||||
err = stack.Reset(brHW, nicAddr.Addr().Next(), uint16(mtu))
|
err = stack.Reset(brHW, nicAddr.Addr().Next(), uint16(mtu))
|
||||||
@@ -64,7 +98,7 @@ func run() (err error) {
|
|||||||
} else {
|
} else {
|
||||||
fmt.Println("OU", iframes)
|
fmt.Println("OU", iframes)
|
||||||
}
|
}
|
||||||
n, err := br.Write(buf[:nwrite])
|
n, err := iface.Write(buf[:nwrite])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
} else if n != nwrite {
|
} else if n != nwrite {
|
||||||
@@ -73,7 +107,7 @@ func run() (err error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
clear(buf)
|
clear(buf)
|
||||||
nread, err := br.Read(buf)
|
nread, err := iface.Read(buf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
} else if nread > 0 {
|
} else if nread > 0 {
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ func main() {
|
|||||||
ip := netip.MustParseAddr(stackIP)
|
ip := netip.MustParseAddr(stackIP)
|
||||||
tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070")
|
tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070")
|
||||||
|
|
||||||
ippfx := tap.IPPrefix()
|
ippfx, _ := tap.IPMask()
|
||||||
if !ippfx.Contains(ip) {
|
if !ippfx.Contains(ip) {
|
||||||
log.Fatal("interface does not contain stack address")
|
log.Fatal("interface does not contain stack address")
|
||||||
}
|
}
|
||||||
@@ -43,8 +43,8 @@ func main() {
|
|||||||
Level: slog.LevelDebug,
|
Level: slog.LevelDebug,
|
||||||
}))
|
}))
|
||||||
|
|
||||||
gatewayMAC := tap.HardwareAddr6()
|
gatewayMAC, _ := tap.HardwareAddress6()
|
||||||
mtu := tap.MTU()
|
mtu, _ := tap.MTU()
|
||||||
|
|
||||||
var stack Stack
|
var stack Stack
|
||||||
err := stack.Reset(stackHWAddr, gatewayMAC, addrPort.Addr(), mtu)
|
err := stack.Reset(stackHWAddr, gatewayMAC, addrPort.Addr(), mtu)
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ func main() {
|
|||||||
ip := netip.MustParseAddr(stackIP)
|
ip := netip.MustParseAddr(stackIP)
|
||||||
tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070")
|
tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070")
|
||||||
|
|
||||||
ippfx := tap.IPPrefix()
|
ippfx, _ := tap.IPMask()
|
||||||
if !ippfx.Contains(ip) {
|
if !ippfx.Contains(ip) {
|
||||||
log.Fatal("interface does not contain stack address")
|
log.Fatal("interface does not contain stack address")
|
||||||
}
|
}
|
||||||
@@ -46,8 +46,8 @@ func main() {
|
|||||||
}))
|
}))
|
||||||
|
|
||||||
slogger := logger{lg}
|
slogger := logger{lg}
|
||||||
gatewayMAC := tap.HardwareAddr6()
|
gatewayMAC, _ := tap.HardwareAddress6()
|
||||||
mtu := tap.MTU()
|
mtu, _ := tap.MTU()
|
||||||
lStack, handler, err := NewEthernetTCPStack(stackHWAddr, gatewayMAC, addrPort, uint16(mtu), slogger)
|
lStack, handler, err := NewEthernetTCPStack(stackHWAddr, gatewayMAC, addrPort, uint16(mtu), slogger)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Fatal(err)
|
log.Fatal(err)
|
||||||
|
|||||||
+2
-22
@@ -5,11 +5,9 @@ import (
|
|||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"log/slog"
|
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -87,26 +85,8 @@ func run() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
fmt.Println("listening on http://127.0.0.1:7070/recv and http://127.0.0.1:7070/send on hwaddr:", net.HardwareAddr(hwaddr[:]).String())
|
fmt.Println("listening on http://127.0.0.1:7070/recv and http://127.0.0.1:7070/send on hwaddr:", net.HardwareAddr(hwaddr[:]).String())
|
||||||
go http.ListenAndServe(":7070", sv)
|
http.ListenAndServe(":7070", sv)
|
||||||
const standbyDuration = 5 * time.Second
|
return errors.New("finished")
|
||||||
lastHit := time.Now().Add(-standbyDuration)
|
|
||||||
for {
|
|
||||||
result, err := sv.HandleTap()
|
|
||||||
if err != nil {
|
|
||||||
slog.Error("handletap:error", slog.String("err", err.Error()), slog.Any("result", result))
|
|
||||||
}
|
|
||||||
if result.Failed {
|
|
||||||
return errors.New("tap failed, exit program")
|
|
||||||
} else if result.ReceivedSize == 0 && result.SentSize == 0 {
|
|
||||||
if time.Since(lastHit) > standbyDuration {
|
|
||||||
time.Sleep(5 * time.Millisecond) // Enter standby.
|
|
||||||
} else {
|
|
||||||
runtime.Gosched()
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
lastHit = time.Now()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func getTCPData(frames []pcap.Frame, pkt []byte) (flags tcp.Flags, src, dst uint16) {
|
func getTCPData(frames []pcap.Frame, pkt []byte) (flags tcp.Flags, src, dst uint16) {
|
||||||
|
|||||||
+12
-6
@@ -10,18 +10,22 @@ var (
|
|||||||
errInvalidIPVersionToSetAddr = errors.New("invalid ip version to setDstAddr")
|
errInvalidIPVersionToSetAddr = errors.New("invalid ip version to setDstAddr")
|
||||||
)
|
)
|
||||||
|
|
||||||
func GetIPSourceAddr(buf []byte) (addr []byte, id uint16, err error) {
|
func GetIPSourceAddr(buf []byte) (addr []byte, id, ipEndOff uint16, err error) {
|
||||||
version := buf[0] >> 4
|
b0 := buf[0]
|
||||||
switch version { //
|
version := b0 >> 4
|
||||||
|
switch version {
|
||||||
case 4:
|
case 4:
|
||||||
addr = buf[12:16]
|
ihl := b0 & 0xf
|
||||||
|
ipEndOff = 4 * uint16(ihl)
|
||||||
id = binary.BigEndian.Uint16(buf[4:6])
|
id = binary.BigEndian.Uint16(buf[4:6])
|
||||||
|
addr = buf[12:16]
|
||||||
case 6:
|
case 6:
|
||||||
addr = buf[8:24]
|
addr = buf[8:24]
|
||||||
|
ipEndOff = 40
|
||||||
default:
|
default:
|
||||||
err = errUnsupportedIP
|
err = errUnsupportedIP
|
||||||
}
|
}
|
||||||
return addr, id, err
|
return addr, id, ipEndOff, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func SetIPDestinationAddr(buf []byte, id uint16, addr []byte) (err error) {
|
func SetIPDestinationAddr(buf []byte, id uint16, addr []byte) (err error) {
|
||||||
@@ -30,7 +34,9 @@ func SetIPDestinationAddr(buf []byte, id uint16, addr []byte) (err error) {
|
|||||||
switch version {
|
switch version {
|
||||||
case 4:
|
case 4:
|
||||||
dstaddr = buf[16:20]
|
dstaddr = buf[16:20]
|
||||||
binary.BigEndian.PutUint16(buf[4:6], id)
|
if id > 0 {
|
||||||
|
binary.BigEndian.PutUint16(buf[4:6], id)
|
||||||
|
}
|
||||||
case 6:
|
case 6:
|
||||||
dstaddr = buf[24:40]
|
dstaddr = buf[24:40]
|
||||||
default:
|
default:
|
||||||
|
|||||||
+74
-125
@@ -5,10 +5,14 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
const minMTU = 256
|
const minMTU = 256
|
||||||
@@ -22,6 +26,8 @@ type Interface interface {
|
|||||||
IPMask() (netip.Prefix, error)
|
IPMask() (netip.Prefix, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var _ Interface = (*HTTPTapClient)(nil)
|
||||||
|
|
||||||
// NewHTTPTapClient returns a HTTPTapClient ready for use.
|
// NewHTTPTapClient returns a HTTPTapClient ready for use.
|
||||||
func NewHTTPTapClient(baseURL string) *HTTPTapClient {
|
func NewHTTPTapClient(baseURL string) *HTTPTapClient {
|
||||||
var h HTTPTapClient
|
var h HTTPTapClient
|
||||||
@@ -35,18 +41,19 @@ func NewHTTPTapClient(baseURL string) *HTTPTapClient {
|
|||||||
return &h
|
return &h
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *HTTPTapClient) IPPrefix() netip.Prefix {
|
func (h *HTTPTapClient) IPMask() (netip.Prefix, error) {
|
||||||
h.ensureMTU()
|
err := h.ensureMTU()
|
||||||
return h.ip
|
return h.ip, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *HTTPTapClient) MTU() int {
|
func (h *HTTPTapClient) MTU() (int, error) {
|
||||||
h.ensureMTU()
|
err := h.ensureMTU()
|
||||||
return len(h.buf)
|
return len(h.buf), err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *HTTPTapClient) HardwareAddr6() [6]byte {
|
func (h *HTTPTapClient) HardwareAddress6() ([6]byte, error) {
|
||||||
return h.hwaddr
|
err := h.ensureMTU()
|
||||||
|
return h.hwaddr, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *HTTPTapClient) ensureMTU() (err error) {
|
func (h *HTTPTapClient) ensureMTU() (err error) {
|
||||||
@@ -91,14 +98,15 @@ type HTTPTapClient struct {
|
|||||||
buf []byte
|
buf []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *HTTPTapClient) ReadDiscard() error {
|
func (h *HTTPTapClient) ReadDiscard() (err error) {
|
||||||
for {
|
for {
|
||||||
d, _ := h.ReadBytes() // Empty remote data.
|
d, err2 := h.ReadBytes() // Empty remote data.
|
||||||
if len(d) == 0 {
|
if len(d) == 0 {
|
||||||
|
err = err2
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *HTTPTapClient) ReadBytes() (data []byte, err error) {
|
func (h *HTTPTapClient) ReadBytes() (data []byte, err error) {
|
||||||
@@ -110,7 +118,8 @@ func (h *HTTPTapClient) ReadBytes() (data []byte, err error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
} else if resp.StatusCode != 200 {
|
} else if resp.StatusCode != 200 {
|
||||||
return nil, errors.New(resp.Status + " for " + h.recvurl)
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
return nil, fmt.Errorf("bad server response %s %s: %s", h.sendurl, resp.Status, b)
|
||||||
}
|
}
|
||||||
buf := h.buf
|
buf := h.buf
|
||||||
err = json.NewDecoder(resp.Body).Decode(&buf)
|
err = json.NewDecoder(resp.Body).Decode(&buf)
|
||||||
@@ -124,7 +133,7 @@ func (h *HTTPTapClient) Read(b []byte) (int, error) {
|
|||||||
err := h.ensureMTU()
|
err := h.ensureMTU()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
} else if len(b) < h.MTU() {
|
} else if len(b) < len(h.buf) {
|
||||||
return 0, errors.New("buffer must have at least MTU size")
|
return 0, errors.New("buffer must have at least MTU size")
|
||||||
}
|
}
|
||||||
data, err := h.ReadBytes()
|
data, err := h.ReadBytes()
|
||||||
@@ -139,7 +148,7 @@ func (h *HTTPTapClient) Write(b []byte) (int, error) {
|
|||||||
err := h.ensureMTU()
|
err := h.ensureMTU()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
} else if len(b) > h.MTU() {
|
} else if len(b) > len(h.buf) {
|
||||||
return 0, errors.New("buffer larger than MTU")
|
return 0, errors.New("buffer larger than MTU")
|
||||||
}
|
}
|
||||||
data, _ := json.Marshal(b)
|
data, _ := json.Marshal(b)
|
||||||
@@ -147,7 +156,8 @@ func (h *HTTPTapClient) Write(b []byte) (int, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
} else if resp.StatusCode != 200 {
|
} else if resp.StatusCode != 200 {
|
||||||
return 0, errors.New(resp.Status + " for " + h.sendurl)
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
return 0, fmt.Errorf("bad server response for plen %d @ %s %s: %s", len(b), h.sendurl, resp.Status, b)
|
||||||
}
|
}
|
||||||
return len(b), nil
|
return len(b), nil
|
||||||
}
|
}
|
||||||
@@ -155,12 +165,12 @@ func (h *HTTPTapClient) Write(b []byte) (int, error) {
|
|||||||
func (h *HTTPTapClient) Close() error { return nil }
|
func (h *HTTPTapClient) Close() error { return nil }
|
||||||
|
|
||||||
type HTTPTapServer struct {
|
type HTTPTapServer struct {
|
||||||
router *http.ServeMux
|
sendmu sync.Mutex
|
||||||
stack stack
|
recvmu sync.Mutex
|
||||||
tap Interface
|
router *http.ServeMux
|
||||||
buf []byte
|
tap Interface
|
||||||
onTx func(channel int, pkt []byte)
|
buf []byte
|
||||||
tapfailed bool
|
onTx func(channel int, pkt []byte)
|
||||||
}
|
}
|
||||||
|
|
||||||
type tapInfo struct {
|
type tapInfo struct {
|
||||||
@@ -188,51 +198,68 @@ func NewHTTPTapServer(iface Interface, queueOut, queueIn int) (*HTTPTapServer, e
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
s := stack{
|
|
||||||
out: make(chan []byte, queueOut),
|
|
||||||
in: make(chan []byte, queueIn),
|
|
||||||
}
|
|
||||||
sv := http.NewServeMux()
|
sv := http.NewServeMux()
|
||||||
taps := &HTTPTapServer{
|
taps := &HTTPTapServer{
|
||||||
router: sv,
|
router: sv,
|
||||||
stack: s,
|
|
||||||
tap: iface,
|
tap: iface,
|
||||||
buf: make([]byte, mtu),
|
buf: make([]byte, mtu),
|
||||||
}
|
}
|
||||||
sv.HandleFunc("/send", func(w http.ResponseWriter, r *http.Request) {
|
sv.HandleFunc("/send", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
retries := 10
|
||||||
|
for {
|
||||||
|
if taps.sendmu.TryLock() {
|
||||||
|
defer taps.sendmu.Unlock()
|
||||||
|
break
|
||||||
|
} else if retries == 0 {
|
||||||
|
slog.Error("send-overload")
|
||||||
|
http.Error(w, "resource in use", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
retries--
|
||||||
|
time.Sleep(100 * time.Microsecond) // approx duration of what one request processing takes on my machine.
|
||||||
|
}
|
||||||
var data []byte
|
var data []byte
|
||||||
err := json.NewDecoder(r.Body).Decode(&data)
|
err := json.NewDecoder(r.Body).Decode(&data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
} else {
|
return
|
||||||
if taps.onTx != nil {
|
}
|
||||||
taps.onTx(1, data)
|
_, err = taps.tap.Write(data)
|
||||||
}
|
if err != nil {
|
||||||
select {
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
case s.out <- data:
|
return
|
||||||
default:
|
}
|
||||||
http.Error(w, "outgoing packet queue full", http.StatusInternalServerError)
|
if taps.onTx != nil {
|
||||||
}
|
taps.onTx(1, data)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
sv.HandleFunc("/recv", func(w http.ResponseWriter, r *http.Request) {
|
sv.HandleFunc("/recv", func(w http.ResponseWriter, r *http.Request) {
|
||||||
select {
|
if !taps.recvmu.TryLock() {
|
||||||
case data := <-s.in:
|
http.Error(w, "resource in use: recv may take a while, are you using concurrent access or have you restarted your client? please wait!", http.StatusInternalServerError)
|
||||||
json.NewEncoder(w).Encode(data)
|
return
|
||||||
default:
|
|
||||||
json.NewEncoder(w).Encode("") // send empty string.
|
|
||||||
}
|
}
|
||||||
|
defer taps.recvmu.Unlock()
|
||||||
|
n, err := taps.tap.Read(taps.buf)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if taps.onTx != nil {
|
||||||
|
taps.onTx(0, taps.buf[:n])
|
||||||
|
}
|
||||||
|
json.NewEncoder(w).Encode(taps.buf[:n])
|
||||||
})
|
})
|
||||||
|
hw6, err := iface.HardwareAddress6()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("acquiring hardware address: %w", err)
|
||||||
|
}
|
||||||
|
hwstr := net.HardwareAddr(hw6[:]).String()
|
||||||
ipstr := netmask.String()
|
ipstr := netmask.String()
|
||||||
sv.HandleFunc("/info", func(w http.ResponseWriter, r *http.Request) {
|
sv.HandleFunc("/info", func(w http.ResponseWriter, r *http.Request) {
|
||||||
info := tapInfo{
|
info := tapInfo{
|
||||||
MTU: mtu,
|
MTU: mtu,
|
||||||
IPPrefix: ipstr,
|
IPPrefix: ipstr,
|
||||||
}
|
HardwareAddr: hwstr,
|
||||||
hw, err := iface.HardwareAddress6()
|
|
||||||
if err == nil {
|
|
||||||
info.HardwareAddr = net.HardwareAddr(hw[:]).String()
|
|
||||||
}
|
}
|
||||||
json.NewEncoder(w).Encode(info)
|
json.NewEncoder(w).Encode(info)
|
||||||
})
|
})
|
||||||
@@ -257,81 +284,3 @@ type HandleTapResult struct {
|
|||||||
SentSize int
|
SentSize int
|
||||||
ReceivedSize int
|
ReceivedSize int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sv *HTTPTapServer) HandleTap() (result HandleTapResult, err error) {
|
|
||||||
result.ReceivedSize, err = sv.readTap()
|
|
||||||
result.Failed = sv.tapfailed
|
|
||||||
if result.Failed && err != nil {
|
|
||||||
return result, err
|
|
||||||
}
|
|
||||||
var err2 error
|
|
||||||
result.ReceivedSize, err2 = sv.writeTap()
|
|
||||||
result.Failed = result.Failed || sv.tapfailed
|
|
||||||
if err2 != nil && err == nil {
|
|
||||||
err = err2
|
|
||||||
} else if err2 != nil {
|
|
||||||
err = errors.Join(err, err2)
|
|
||||||
}
|
|
||||||
return result, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (sv *HTTPTapServer) readTap() (int, error) {
|
|
||||||
buf := sv.buf
|
|
||||||
n, err := sv.tap.Read(buf[:])
|
|
||||||
if err != nil {
|
|
||||||
sv.tapfailed = true
|
|
||||||
return n, err
|
|
||||||
} else if n > 0 {
|
|
||||||
if sv.onTx != nil {
|
|
||||||
sv.onTx(0, buf[:n])
|
|
||||||
}
|
|
||||||
err = sv.stack.recv(buf[:n])
|
|
||||||
if err != nil {
|
|
||||||
return n, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return n, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (sv *HTTPTapServer) writeTap() (int, error) {
|
|
||||||
buf := sv.buf
|
|
||||||
n, err := sv.stack.handle(buf[:])
|
|
||||||
if err != nil {
|
|
||||||
return n, err
|
|
||||||
} else if n > 0 {
|
|
||||||
n, err = sv.tap.Write(buf[:n])
|
|
||||||
if err != nil {
|
|
||||||
sv.tapfailed = true
|
|
||||||
return n, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return n, err
|
|
||||||
}
|
|
||||||
|
|
||||||
type stack struct {
|
|
||||||
out chan []byte
|
|
||||||
in chan []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *stack) recv(b []byte) (err error) {
|
|
||||||
bcopy := append([]byte{}, b...)
|
|
||||||
RETRY:
|
|
||||||
select {
|
|
||||||
case s.in <- bcopy:
|
|
||||||
default:
|
|
||||||
err = errors.New("receive queue packet full, dropping packet")
|
|
||||||
<-s.in
|
|
||||||
goto RETRY
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *stack) handle(b []byte) (n int, _ error) {
|
|
||||||
select {
|
|
||||||
case incoming := <-s.out:
|
|
||||||
n = copy(b, incoming)
|
|
||||||
default:
|
|
||||||
// pass if no data available.
|
|
||||||
}
|
|
||||||
return n, nil
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,19 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
// Prand16 generates a pseudo random number from a seed.
|
||||||
|
func Prand16(seed uint16) uint16 {
|
||||||
|
// 16bit Xorshift https://en.wikipedia.org/wiki/Xorshift
|
||||||
|
seed ^= seed << 7
|
||||||
|
seed ^= seed >> 9
|
||||||
|
seed ^= seed << 8
|
||||||
|
return seed
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prand32 generates a pseudo random number from a seed.
|
||||||
|
func Prand32[T ~uint32](seed T) T {
|
||||||
|
/* Algorithm "xor" from p. 4 of Marsaglia, "Xorshift RNGs" */
|
||||||
|
seed ^= seed << 13
|
||||||
|
seed ^= seed >> 17
|
||||||
|
seed ^= seed << 5
|
||||||
|
return seed
|
||||||
|
}
|
||||||
@@ -104,6 +104,28 @@ func getNode(nodes []node, port uint16, protocol uint16) (node *node) {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func getEncapsulateNode(nodes *[]node, carrierData []byte, frameOffset int) (nodeIdx int, written int, err error) {
|
||||||
|
destroyed := false
|
||||||
|
for i := range *nodes {
|
||||||
|
node := &(*nodes)[i]
|
||||||
|
if checkNode(node) {
|
||||||
|
destroyed = true
|
||||||
|
node.destroy()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
written, err = node.encapsulate(carrierData, frameOffset)
|
||||||
|
if written > 0 {
|
||||||
|
return i, written, err
|
||||||
|
} else if err != nil {
|
||||||
|
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if destroyed {
|
||||||
|
*nodes = nodesCompact(*nodes)
|
||||||
|
}
|
||||||
|
return -1, 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
// destroy removes all references to underlying StackNode. Allows garbage collection of node if possible.
|
// destroy removes all references to underlying StackNode. Allows garbage collection of node if possible.
|
||||||
func (n *node) destroy() {
|
func (n *node) destroy() {
|
||||||
*n = node{}
|
*n = node{}
|
||||||
@@ -116,5 +138,17 @@ func getNodeByProto(nodes []node, protocol uint16) int {
|
|||||||
return i
|
return i
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return -1
|
return -1
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func nodesCompact(nodes []node) []node {
|
||||||
|
nilOff := 0
|
||||||
|
for i := 0; i < len(nodes); i++ {
|
||||||
|
if !checkNode(&nodes[i]) {
|
||||||
|
nodes[nilOff] = nodes[i]
|
||||||
|
nilOff++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nodes[:nilOff]
|
||||||
|
}
|
||||||
|
|||||||
@@ -123,7 +123,7 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
addr, _, err := internal.GetIPSourceAddr(carrierData)
|
addr, _, _, err := internal.GetIPSourceAddr(carrierData)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
+16
-3
@@ -18,6 +18,7 @@ var _ StackNode = (*StackIP)(nil)
|
|||||||
|
|
||||||
type StackIP struct {
|
type StackIP struct {
|
||||||
connID uint64
|
connID uint64
|
||||||
|
ipID uint16
|
||||||
ip [4]byte
|
ip [4]byte
|
||||||
validator lneto.Validator
|
validator lneto.Validator
|
||||||
handlers []node
|
handlers []node
|
||||||
@@ -136,24 +137,31 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error)
|
|||||||
ifrm, _ := ipv4.NewFrame(frame)
|
ifrm, _ := ipv4.NewFrame(frame)
|
||||||
const ihl = 5
|
const ihl = 5
|
||||||
const headerlen = ihl * 4
|
const headerlen = ihl * 4
|
||||||
|
const dontFrag = 0x4000
|
||||||
ifrm.SetVersionAndIHL(4, ihl)
|
ifrm.SetVersionAndIHL(4, ihl)
|
||||||
ifrm.SetToS(0)
|
ifrm.SetToS(0)
|
||||||
ifrm.SetID(0)
|
seed := sb.ipID + uint16(sb.connID)
|
||||||
|
id := internal.Prand16(seed)
|
||||||
|
ifrm.SetID(id)
|
||||||
|
ifrm.SetFlags(dontFrag)
|
||||||
*ifrm.SourceAddr() = sb.ip
|
*ifrm.SourceAddr() = sb.ip
|
||||||
|
sb.ipID = id
|
||||||
for i := range sb.handlers {
|
for i := range sb.handlers {
|
||||||
h := &sb.handlers[i]
|
h := &sb.handlers[i]
|
||||||
proto := lneto.IPProto(h.proto)
|
proto := lneto.IPProto(h.proto)
|
||||||
n, err := h.encapsulate(frame[:], headerlen)
|
n, err := h.encapsulate(frame[:], headerlen)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if handleNodeError(&sb.handlers, i, err) {
|
||||||
|
println("NODE REMOVED", proto.String(), h.port)
|
||||||
|
h.destroy()
|
||||||
|
}
|
||||||
sb.error("StackIP:handle", slog.String("proto", proto.String()), slog.String("err", err.Error()))
|
sb.error("StackIP:handle", slog.String("proto", proto.String()), slog.String("err", err.Error()))
|
||||||
continue
|
continue
|
||||||
} else if n == 0 {
|
} else if n == 0 {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
const dontFrag = 0x4000
|
|
||||||
totalLen := n + headerlen
|
totalLen := n + headerlen
|
||||||
ifrm.SetTotalLength(uint16(totalLen))
|
ifrm.SetTotalLength(uint16(totalLen))
|
||||||
ifrm.SetFlags(dontFrag)
|
|
||||||
ifrm.SetTTL(64)
|
ifrm.SetTTL(64)
|
||||||
ifrm.SetProtocol(proto)
|
ifrm.SetProtocol(proto)
|
||||||
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
|
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
|
||||||
@@ -168,8 +176,13 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error)
|
|||||||
case lneto.IPProtoUDP:
|
case lneto.IPProtoUDP:
|
||||||
ifrm.CRCWriteUDPPseudo(&crc)
|
ifrm.CRCWriteUDPPseudo(&crc)
|
||||||
ufrm, _ := udp.NewFrame(ifrm.Payload())
|
ufrm, _ := udp.NewFrame(ifrm.Payload())
|
||||||
|
ufrm.SetLength(uint16(n))
|
||||||
ufrm.CRCWriteIPv4(&crc)
|
ufrm.CRCWriteIPv4(&crc)
|
||||||
ufrm.SetCRC(crc.Sum16())
|
ufrm.SetCRC(crc.Sum16())
|
||||||
|
if n != int(ufrm.Length()) {
|
||||||
|
sb.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length())))
|
||||||
|
return 0, errors.New("invalid UDP length")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return totalLen, nil
|
return totalLen, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error {
|
|||||||
if sudp.rmport != 0 && src != sudp.rmport {
|
if sudp.rmport != 0 && src != sudp.rmport {
|
||||||
return nil // Not from our target remote port.
|
return nil // Not from our target remote port.
|
||||||
}
|
}
|
||||||
err = sudp.h.demux(ufrm.Payload(), 8)
|
err = sudp.h.demux(carrierData, frameOffset+8)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if checkNodeErr(&sudp.h, err) {
|
if checkNodeErr(&sudp.h, err) {
|
||||||
sudp.h.destroy()
|
sudp.h.destroy()
|
||||||
@@ -68,11 +68,14 @@ func (sudp *StackUDPPort) Encapsulate(carrierData []byte, frameOffset int) (int,
|
|||||||
}
|
}
|
||||||
ufrm.SetSourcePort(sudp.h.port)
|
ufrm.SetSourcePort(sudp.h.port)
|
||||||
ufrm.SetDestinationPort(sudp.rmport)
|
ufrm.SetDestinationPort(sudp.rmport)
|
||||||
n, err := sudp.h.encapsulate(carrierData[frameOffset:], 8)
|
n, err := sudp.h.encapsulate(carrierData, frameOffset+8)
|
||||||
if err != nil {
|
if n == 0 {
|
||||||
slog.Error("stackudp:demux", slog.String("err", err.Error()))
|
if err != nil {
|
||||||
|
slog.Error("stackudp:demux", slog.String("err", err.Error()))
|
||||||
|
}
|
||||||
|
return 0, err
|
||||||
}
|
}
|
||||||
ufrm.SetLength(8 + uint16(n))
|
// UDP CRC and length left to IP layer.
|
||||||
// UDP CRC left to IP layer.
|
length := 8 + n
|
||||||
return n, err
|
return length, err
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-2
@@ -194,7 +194,7 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) {
|
|||||||
if off >= len(buf) {
|
if off >= len(buf) {
|
||||||
return errors.New("bad offset in TCPConn.Recv")
|
return errors.New("bad offset in TCPConn.Recv")
|
||||||
}
|
}
|
||||||
raddr, id, err := internal.GetIPSourceAddr(buf[:off])
|
raddr, id, _, err := internal.GetIPSourceAddr(buf[:off])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -216,7 +216,7 @@ func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) {
|
|||||||
if len(conn.remoteAddr) == 0 {
|
if len(conn.remoteAddr) == 0 {
|
||||||
return 0, errors.New("unset IP address")
|
return 0, errors.New("unset IP address")
|
||||||
}
|
}
|
||||||
raddr, _, err := internal.GetIPSourceAddr(buf[:off])
|
raddr, _, _, err := internal.GetIPSourceAddr(buf[:off])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
} else if len(raddr) != len(conn.remoteAddr) {
|
} else if len(raddr) != len(conn.remoteAddr) {
|
||||||
|
|||||||
Reference in New Issue
Block a user