diff --git a/dhcpv4/client.go b/dhcpv4/client.go index be520d1..2bdf3e3 100644 --- a/dhcpv4/client.go +++ b/dhcpv4/client.go @@ -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[:]) diff --git a/dhcpv4/definitions.go b/dhcpv4/definitions.go index 17bdd71..24f19ad 100644 --- a/dhcpv4/definitions.go +++ b/dhcpv4/definitions.go @@ -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)") diff --git a/dhcpv4/dhcp_test.go b/dhcpv4/dhcp_test.go new file mode 100644 index 0000000..a291f3c --- /dev/null +++ b/dhcpv4/dhcp_test.go @@ -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) + } +} diff --git a/dhcpv4/frame.go b/dhcpv4/frame.go index e2f4fd6..d5d3297 100644 --- a/dhcpv4/frame.go +++ b/dhcpv4/frame.go @@ -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) + } +} diff --git a/dhcpv4/server.go b/dhcpv4/server.go new file mode 100644 index 0000000..8c80202 --- /dev/null +++ b/dhcpv4/server.go @@ -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 +} diff --git a/examples/bridge/main.go b/examples/bridge/main.go index d2867e3..6389a78 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -3,16 +3,19 @@ package main import ( "crypto/rand" "encoding/binary" + "flag" "fmt" "net" "net/netip" "os" + "strings" "time" "github.com/soypat/lneto" "github.com/soypat/lneto/arp" "github.com/soypat/lneto/dhcpv4" "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" "github.com/soypat/lneto/internal/ltesto" "github.com/soypat/lneto/internet" "github.com/soypat/lneto/internet/pcap" @@ -28,16 +31,47 @@ func main() { } func run() (err error) { - br := ltesto.NewHTTPTapClient("http://127.0.0.1:7070") - defer br.Close() - - nicHW := br.HardwareAddr6() + var ( + flagInterface = "tap0" + flagUseHTTP = false + ) + 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[5]++ // We'll be using a similar HW address but with NIC specific identifier modified. - mtu := br.MTU() - nicAddr := br.IPPrefix() - + mtu, err := iface.MTU() + 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()) var stack Stack err = stack.Reset(brHW, nicAddr.Addr().Next(), uint16(mtu)) @@ -64,7 +98,7 @@ func run() (err error) { } else { fmt.Println("OU", iframes) } - n, err := br.Write(buf[:nwrite]) + n, err := iface.Write(buf[:nwrite]) if err != nil { return err } else if n != nwrite { @@ -73,7 +107,7 @@ func run() (err error) { } clear(buf) - nread, err := br.Read(buf) + nread, err := iface.Read(buf) if err != nil { return err } else if nread > 0 { diff --git a/examples/stack/main.go b/examples/stack/main.go index 5943fb7..8bec0dc 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -34,7 +34,7 @@ func main() { ip := netip.MustParseAddr(stackIP) tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070") - ippfx := tap.IPPrefix() + ippfx, _ := tap.IPMask() if !ippfx.Contains(ip) { log.Fatal("interface does not contain stack address") } @@ -43,8 +43,8 @@ func main() { Level: slog.LevelDebug, })) - gatewayMAC := tap.HardwareAddr6() - mtu := tap.MTU() + gatewayMAC, _ := tap.HardwareAddress6() + mtu, _ := tap.MTU() var stack Stack err := stack.Reset(stackHWAddr, gatewayMAC, addrPort.Addr(), mtu) diff --git a/examples/stackbasic/main.go b/examples/stackbasic/main.go index bd3eed9..a554215 100644 --- a/examples/stackbasic/main.go +++ b/examples/stackbasic/main.go @@ -36,7 +36,7 @@ func main() { ip := netip.MustParseAddr(stackIP) tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070") - ippfx := tap.IPPrefix() + ippfx, _ := tap.IPMask() if !ippfx.Contains(ip) { log.Fatal("interface does not contain stack address") } @@ -46,8 +46,8 @@ func main() { })) slogger := logger{lg} - gatewayMAC := tap.HardwareAddr6() - mtu := tap.MTU() + gatewayMAC, _ := tap.HardwareAddress6() + mtu, _ := tap.MTU() lStack, handler, err := NewEthernetTCPStack(stackHWAddr, gatewayMAC, addrPort, uint16(mtu), slogger) if err != nil { log.Fatal(err) diff --git a/examples/tap/main.go b/examples/tap/main.go index b5e1442..ca8d056 100644 --- a/examples/tap/main.go +++ b/examples/tap/main.go @@ -5,11 +5,9 @@ import ( "flag" "fmt" "log" - "log/slog" "net" "net/http" "net/netip" - "runtime" "strings" "time" @@ -87,26 +85,8 @@ func run() error { 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()) - go http.ListenAndServe(":7070", sv) - const standbyDuration = 5 * time.Second - 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() - } - } + http.ListenAndServe(":7070", sv) + return errors.New("finished") } func getTCPData(frames []pcap.Frame, pkt []byte) (flags tcp.Flags, src, dst uint16) { diff --git a/internal/ip.go b/internal/ip.go index e3ff6ee..2c3c13c 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -10,18 +10,22 @@ var ( errInvalidIPVersionToSetAddr = errors.New("invalid ip version to setDstAddr") ) -func GetIPSourceAddr(buf []byte) (addr []byte, id uint16, err error) { - version := buf[0] >> 4 - switch version { // +func GetIPSourceAddr(buf []byte) (addr []byte, id, ipEndOff uint16, err error) { + b0 := buf[0] + version := b0 >> 4 + switch version { case 4: - addr = buf[12:16] + ihl := b0 & 0xf + ipEndOff = 4 * uint16(ihl) id = binary.BigEndian.Uint16(buf[4:6]) + addr = buf[12:16] case 6: addr = buf[8:24] + ipEndOff = 40 default: err = errUnsupportedIP } - return addr, id, err + return addr, id, ipEndOff, err } 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 { case 4: dstaddr = buf[16:20] - binary.BigEndian.PutUint16(buf[4:6], id) + if id > 0 { + binary.BigEndian.PutUint16(buf[4:6], id) + } case 6: dstaddr = buf[24:40] default: diff --git a/internal/ltesto/httptap.go b/internal/ltesto/httptap.go index e73eea0..eaaf022 100644 --- a/internal/ltesto/httptap.go +++ b/internal/ltesto/httptap.go @@ -5,10 +5,14 @@ import ( "encoding/json" "errors" "fmt" + "io" + "log/slog" "net" "net/http" "net/netip" "net/url" + "sync" + "time" ) const minMTU = 256 @@ -22,6 +26,8 @@ type Interface interface { IPMask() (netip.Prefix, error) } +var _ Interface = (*HTTPTapClient)(nil) + // NewHTTPTapClient returns a HTTPTapClient ready for use. func NewHTTPTapClient(baseURL string) *HTTPTapClient { var h HTTPTapClient @@ -35,18 +41,19 @@ func NewHTTPTapClient(baseURL string) *HTTPTapClient { return &h } -func (h *HTTPTapClient) IPPrefix() netip.Prefix { - h.ensureMTU() - return h.ip +func (h *HTTPTapClient) IPMask() (netip.Prefix, error) { + err := h.ensureMTU() + return h.ip, err } -func (h *HTTPTapClient) MTU() int { - h.ensureMTU() - return len(h.buf) +func (h *HTTPTapClient) MTU() (int, error) { + err := h.ensureMTU() + return len(h.buf), err } -func (h *HTTPTapClient) HardwareAddr6() [6]byte { - return h.hwaddr +func (h *HTTPTapClient) HardwareAddress6() ([6]byte, error) { + err := h.ensureMTU() + return h.hwaddr, err } func (h *HTTPTapClient) ensureMTU() (err error) { @@ -91,14 +98,15 @@ type HTTPTapClient struct { buf []byte } -func (h *HTTPTapClient) ReadDiscard() error { +func (h *HTTPTapClient) ReadDiscard() (err error) { for { - d, _ := h.ReadBytes() // Empty remote data. + d, err2 := h.ReadBytes() // Empty remote data. if len(d) == 0 { + err = err2 break } } - return nil + return err } func (h *HTTPTapClient) ReadBytes() (data []byte, err error) { @@ -110,7 +118,8 @@ func (h *HTTPTapClient) ReadBytes() (data []byte, err error) { if err != nil { return nil, err } 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 err = json.NewDecoder(resp.Body).Decode(&buf) @@ -124,7 +133,7 @@ func (h *HTTPTapClient) Read(b []byte) (int, error) { err := h.ensureMTU() if err != nil { 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") } data, err := h.ReadBytes() @@ -139,7 +148,7 @@ func (h *HTTPTapClient) Write(b []byte) (int, error) { err := h.ensureMTU() if err != nil { return 0, err - } else if len(b) > h.MTU() { + } else if len(b) > len(h.buf) { return 0, errors.New("buffer larger than MTU") } data, _ := json.Marshal(b) @@ -147,7 +156,8 @@ func (h *HTTPTapClient) Write(b []byte) (int, error) { if err != nil { return 0, err } 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 } @@ -155,12 +165,12 @@ func (h *HTTPTapClient) Write(b []byte) (int, error) { func (h *HTTPTapClient) Close() error { return nil } type HTTPTapServer struct { - router *http.ServeMux - stack stack - tap Interface - buf []byte - onTx func(channel int, pkt []byte) - tapfailed bool + sendmu sync.Mutex + recvmu sync.Mutex + router *http.ServeMux + tap Interface + buf []byte + onTx func(channel int, pkt []byte) } type tapInfo struct { @@ -188,51 +198,68 @@ func NewHTTPTapServer(iface Interface, queueOut, queueIn int) (*HTTPTapServer, e return nil, err } - s := stack{ - out: make(chan []byte, queueOut), - in: make(chan []byte, queueIn), - } sv := http.NewServeMux() taps := &HTTPTapServer{ router: sv, - stack: s, tap: iface, buf: make([]byte, mtu), } 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 err := json.NewDecoder(r.Body).Decode(&data) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) - } else { - if taps.onTx != nil { - taps.onTx(1, data) - } - select { - case s.out <- data: - default: - http.Error(w, "outgoing packet queue full", http.StatusInternalServerError) - } + return + } + _, err = taps.tap.Write(data) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + if taps.onTx != nil { + taps.onTx(1, data) } }) sv.HandleFunc("/recv", func(w http.ResponseWriter, r *http.Request) { - select { - case data := <-s.in: - json.NewEncoder(w).Encode(data) - default: - json.NewEncoder(w).Encode("") // send empty string. + if !taps.recvmu.TryLock() { + 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) + return } + 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() sv.HandleFunc("/info", func(w http.ResponseWriter, r *http.Request) { info := tapInfo{ - MTU: mtu, - IPPrefix: ipstr, - } - hw, err := iface.HardwareAddress6() - if err == nil { - info.HardwareAddr = net.HardwareAddr(hw[:]).String() + MTU: mtu, + IPPrefix: ipstr, + HardwareAddr: hwstr, } json.NewEncoder(w).Encode(info) }) @@ -257,81 +284,3 @@ type HandleTapResult struct { SentSize 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 -} diff --git a/internal/prand.go b/internal/prand.go new file mode 100644 index 0000000..a5544c0 --- /dev/null +++ b/internal/prand.go @@ -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 +} diff --git a/internet/definitions.go b/internet/definitions.go index 2985ed5..e14e153 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -104,6 +104,28 @@ func getNode(nodes []node, port uint16, protocol uint16) (node *node) { 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. func (n *node) destroy() { *n = node{} @@ -116,5 +138,17 @@ func getNodeByProto(nodes []node, protocol uint16) int { return i } } + 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] +} diff --git a/internet/node-tcplistener.go b/internet/node-tcplistener.go index e43c37a..2621f9d 100644 --- a/internet/node-tcplistener.go +++ b/internet/node-tcplistener.go @@ -123,7 +123,7 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e if err != nil { return err } - addr, _, err := internal.GetIPSourceAddr(carrierData) + addr, _, _, err := internal.GetIPSourceAddr(carrierData) if err != nil { return err } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index d15a3fc..478391b 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -18,6 +18,7 @@ var _ StackNode = (*StackIP)(nil) type StackIP struct { connID uint64 + ipID uint16 ip [4]byte validator lneto.Validator handlers []node @@ -136,24 +137,31 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) ifrm, _ := ipv4.NewFrame(frame) const ihl = 5 const headerlen = ihl * 4 + const dontFrag = 0x4000 ifrm.SetVersionAndIHL(4, ihl) 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 + sb.ipID = id for i := range sb.handlers { h := &sb.handlers[i] proto := lneto.IPProto(h.proto) n, err := h.encapsulate(frame[:], headerlen) 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())) continue } else if n == 0 { continue } - const dontFrag = 0x4000 totalLen := n + headerlen ifrm.SetTotalLength(uint16(totalLen)) - ifrm.SetFlags(dontFrag) ifrm.SetTTL(64) ifrm.SetProtocol(proto) ifrm.SetCRC(ifrm.CalculateHeaderCRC()) @@ -168,8 +176,13 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) case lneto.IPProtoUDP: ifrm.CRCWriteUDPPseudo(&crc) ufrm, _ := udp.NewFrame(ifrm.Payload()) + ufrm.SetLength(uint16(n)) ufrm.CRCWriteIPv4(&crc) 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 } diff --git a/internet/stack-udpport.go b/internet/stack-udpport.go index c71a9f8..9d0bb55 100644 --- a/internet/stack-udpport.go +++ b/internet/stack-udpport.go @@ -47,7 +47,7 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error { if sudp.rmport != 0 && src != sudp.rmport { 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 checkNodeErr(&sudp.h, err) { sudp.h.destroy() @@ -68,11 +68,14 @@ func (sudp *StackUDPPort) Encapsulate(carrierData []byte, frameOffset int) (int, } ufrm.SetSourcePort(sudp.h.port) ufrm.SetDestinationPort(sudp.rmport) - n, err := sudp.h.encapsulate(carrierData[frameOffset:], 8) - if err != nil { - slog.Error("stackudp:demux", slog.String("err", err.Error())) + n, err := sudp.h.encapsulate(carrierData, frameOffset+8) + if n == 0 { + if err != nil { + slog.Error("stackudp:demux", slog.String("err", err.Error())) + } + return 0, err } - ufrm.SetLength(8 + uint16(n)) - // UDP CRC left to IP layer. - return n, err + // UDP CRC and length left to IP layer. + length := 8 + n + return length, err } diff --git a/tcp/conn.go b/tcp/conn.go index 4558072..0146bf9 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -194,7 +194,7 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) { if off >= len(buf) { 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 { return err } @@ -216,7 +216,7 @@ func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) { if len(conn.remoteAddr) == 0 { return 0, errors.New("unset IP address") } - raddr, _, err := internal.GetIPSourceAddr(buf[:off]) + raddr, _, _, err := internal.GetIPSourceAddr(buf[:off]) if err != nil { return 0, err } else if len(raddr) != len(conn.remoteAddr) {