diff --git a/examples/stack/main.go b/examples/stack/main.go index 81a23bf..fcf4d5e 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -3,34 +3,125 @@ package main import ( "errors" "io" + "log" + "log/slog" + "math/rand" "github.com/soypat/lneto" + "github.com/soypat/lneto/internal" + "github.com/soypat/lneto/internal/ltesto" + "github.com/soypat/lneto/tcp" ) func main() { + rng := rand.New(rand.NewSource(1)) + var gen ltesto.PacketGen + gen.RandomizeAddrs(rng) + slogger := logger{slog.Default()} + lStack := LinkStack{ + logger: slogger, + mac: gen.DstMAC, + mtu: 1500, + } + iStack := &IPv4Stack{ + ip: gen.DstIPv4, + logger: slogger, + } + tStack := &TCPStack{ + logger: slogger, + } + pStack := &TCPPort{ + lport: gen.DstTCP, + rport: gen.SrcTCP, + tcb: tcp.ControlBlock{}, + } + + err := iStack.Register(tStack, &gen.SrcIPv4) + if err != nil { + log.Fatal(err) + } + err = lStack.Register(iStack, gen.SrcMAC) + if err != nil { + log.Fatal(err) + } + err = tStack.Register(pStack, pStack.lport) + if err != nil { + log.Fatal(err) + } + + err = pStack.tcb.Open(tcp.Value(rng.Int()), 256, tcp.StateListen) + if err != nil { + log.Fatal(err) + } + + buf := make([]byte, lStack.mtu) + packet := gen.AppendRandomIPv4TCPPacket(buf[:0], rng) + err = lStack.RecvEth(packet) + if err != nil { + log.Fatal(err) + } +} + +type Handler interface { + Protocol() uint32 + Recv(frame []byte, off int) error + Handle(dstAndFrame []byte, dstOff int) (int, error) +} + +// handler is abstraction of a frame marshaller. +type handler struct { + raddr []byte + recv func([]byte, int) error + handle func([]byte, int) (int, error) + proto uint32 + lport uint16 } type LinkStack struct { + handlers []handler + logger mac [6]byte mtu uint16 } +func (ls *LinkStack) Register(h Handler, remoteHWAddr [6]byte) error { + proto := h.Protocol() + for i := range ls.handlers { + if proto == ls.handlers[i].proto { + return errors.New("protocol already registered") + } + } + // Pattern to add a handler and reuse underlying memory. + ls.handlers = append(ls.handlers, handler{}) + hh := &ls.handlers[len(ls.handlers)-1] + hh.handle = h.Handle + hh.recv = h.Recv + hh.proto = proto + hh.raddr = append(hh.raddr[:0], remoteHWAddr[:]...) + return nil +} + func (ls *LinkStack) RecvEth(ethFrame []byte) (err error) { - eframe, err := lneto.NewEthFrame(ethFrame) + efrm, err := lneto.NewEthFrame(ethFrame) if err != nil { return err } - if !eframe.IsBroadcast() && ls.mac != *eframe.DestinationHardwareAddr() { + if !efrm.IsBroadcast() && ls.mac != *efrm.DestinationHardwareAddr() { return errors.New("packet MAC mismatch") } - - // Convert to dynamic handling. - etype := eframe.EtherTypeOrSize() - if etype != lneto.EtherTypeARP && etype != lneto.EtherTypeIPv4 && etype != lneto.EtherTypeIPv6 { - return nil + var vld lneto.Validator + efrm.ValidateSize(&vld) + if err := vld.Err(); err != nil { + return err + } + etype := efrm.EtherTypeOrSize() + for i := range ls.handlers { + h := &ls.handlers[i] + if h.proto == uint32(etype) { + return h.recv(efrm.Payload(), 0) + } } - return nil } @@ -38,40 +129,328 @@ func (ls *LinkStack) HandleEth(dst []byte) (n int, err error) { if len(dst) < int(ls.mtu) { return 0, io.ErrShortBuffer } - n, addr, etype, err := ls.handleUpper(dst[14:]) - if err != nil || n == 0 { - return 0, err + for i := range ls.handlers { + h := &ls.handlers[i] + n, err = h.handle(dst[:ls.mtu], 14) + if err != nil { + ls.error("handling", slog.String("proto", lneto.EtherType(h.proto).String()), slog.String("err", err.Error())) + continue + } + if n > 0 { + // Found packet + efrm, _ := lneto.NewEthFrame(dst[:14]) + copy(efrm.DestinationHardwareAddr()[:], h.raddr) + *efrm.SourceHardwareAddr() = ls.mac + efrm.SetEtherType(lneto.EtherType(h.proto)) + return n + 14, nil + } } - eframe, _ := lneto.NewEthFrame(dst[:14]) - *eframe.DestinationHardwareAddr() = addr - *eframe.SourceHardwareAddr() = ls.mac - eframe.SetEtherType(etype) - return 14 + n, nil -} - -func (ls *LinkStack) handleUpper(dst []byte) (n int, dstAddr [6]byte, etype lneto.EtherType, err error) { - return + return 0, err } type IPv4Stack struct { ip [4]byte - mtu uint16 validator lneto.Validator + handlers []handler + logger } -func (is *IPv4Stack) Recv(ipframe []byte) error { - iframe, err := lneto.NewIPv4Frame(ipframe) +func (is *IPv4Stack) Protocol() uint32 { return uint32(lneto.EtherTypeIPv4) } + +func (is *IPv4Stack) Register(h Handler, remoteAddr *[4]byte) error { + proto := h.Protocol() + for i := range is.handlers { + if proto == is.handlers[i].proto { + return errors.New("protocol already registered") + } + } + // Pattern to add a handler and reuse underlying memory. + is.handlers = append(is.handlers, handler{}) + hh := &is.handlers[len(is.handlers)-1] + hh.handle = h.Handle + hh.recv = h.Recv + hh.proto = proto + if remoteAddr != nil { + // Remote IP address specified. + hh.raddr = append(hh.raddr, remoteAddr[:]...) + } + return nil +} + +func (is *IPv4Stack) Recv(ethFrame []byte, ipOff int) error { + ifrm, err := lneto.NewIPv4Frame(ethFrame[ipOff:]) if err != nil { return err } - if *iframe.DestinationAddr() != is.ip { + if *ifrm.DestinationAddr() != is.ip { return errors.New("packet not for us") } - iframe.Validate(&is.validator) - err = is.validator.Err() + is.validator.ResetErr() + ifrm.ValidateExceptCRC(&is.validator) + if err = is.validator.Err(); err != nil { + return err + } + gotCRC := ifrm.CRC() + wantCRC := ifrm.CalculateHeaderCRC() + if gotCRC != wantCRC { + is.error("IPv4Stack:Recv:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC))) + return errors.New("IPv4 CRC mismatch") + } + off := ifrm.HeaderLength() + totalLen := ifrm.TotalLength() + for i := range is.handlers { + h := &is.handlers[i] + if h.proto == uint32(ifrm.Protocol()) { + return h.recv(ethFrame[ipOff:totalLen], off) + } + } + return nil +} + +func (is *IPv4Stack) Handle(ethFrame []byte, ipOff int) (int, error) { + if len(ethFrame)-ipOff < 256 { + return 0, io.ErrShortBuffer + } + ifrm, _ := lneto.NewIPv4Frame(ethFrame[ipOff:]) + const ihl = 5 + const headerlen = ihl * 4 + ifrm.SetVersionAndIHL(4, 5) + *ifrm.SourceAddr() = is.ip + ifrm.SetToS(0) + for i := range is.handlers { + h := &is.handlers[i] + proto := lneto.IPProto(h.proto) + ifrm.SetProtocol(proto) + if len(h.raddr) == 4 { + copy(ifrm.DestinationAddr()[:], h.raddr) + } else { + copy(ifrm.DestinationAddr()[:], "\x00\x00\x00\x00") + } + + n, err := h.handle(ethFrame[ipOff:], headerlen) + if err != nil { + is.error("IPv4Stack:handle", slog.String("proto", proto.String()), slog.String("err", err.Error())) + continue + } + if n > 0 { + const dontFrag = 0x4000 + totalLen := n + headerlen + ifrm.SetTotalLength(uint16(totalLen)) + ifrm.SetID(0) + ifrm.SetFlags(dontFrag) + ifrm.SetTTL(64) + ifrm.SetCRC(ifrm.CalculateHeaderCRC()) + return totalLen, nil + } + } + return 0, nil +} + +type TCPStack struct { + validator lneto.Validator + handlers []handler + logger +} + +func (ts *TCPStack) Protocol() uint32 { return uint32(lneto.IPProtoTCP) } + +func (ts *TCPStack) Register(h Handler, lport uint16) error { + if lport == 0 { + return errors.New("got zero port") + } + ts.handlers = append(ts.handlers, handler{}) + hh := &ts.handlers[len(ts.handlers)-1] + hh.handle = h.Handle + hh.recv = h.Recv + hh.lport = lport + return nil +} + +func (ts *TCPStack) Recv(ipFrame []byte, tcpOff int) error { + ipVersion := ipFrame[0] >> 4 + if ipVersion != 4 && ipVersion != 6 { + return errors.New("invalid IP version") + } + tfrm, err := lneto.NewTCPFrame(ipFrame[tcpOff:]) + if err != nil { + return err + } + lport := tfrm.DestinationPort() + var h *handler + for i := range ts.handlers { + if lport == ts.handlers[i].lport { + h = &ts.handlers[i] + break + } + } + if h == nil { + return errors.New("port not found") + } + ts.validator.ResetErr() + tfrm.ValidateSize(&ts.validator) + if err = ts.validator.Err(); err != nil { + return err + } + var crc uint16 + switch ipVersion { + case 4: + ifrm, _ := lneto.NewIPv4Frame(ipFrame) + crc = tfrm.CalculateIPv4CRC(ifrm) + case 6: + ifrm, _ := lneto.NewIPv6Frame(ipFrame) + crc = tfrm.CalculateIPv6CRC(ifrm) + } + gotCRC := tfrm.CRC() + if crc != gotCRC { + ts.error("TCPStack:Recv:crc-mismatch", slog.Uint64("lport", uint64(lport)), slog.Uint64("want", uint64(crc)), slog.Uint64("got", uint64(gotCRC))) + return errors.New("TCP crc mismatch") + } + return h.recv(ipFrame[tcpOff:], 0) +} + +func (ts *TCPStack) Handle(ipFrame []byte, tcpOff int) (n int, err error) { + ipVersion := ipFrame[0] >> 4 + if ipVersion != 4 && ipVersion != 6 { + return 0, errors.New("invalid IP version") + } + var h *handler + for i := range ts.handlers { + h = &ts.handlers[i] + n, err = h.handle(ipFrame[tcpOff:], 0) + if err != nil { + if err == io.EOF { + ts.handlers = removeHandler(ts.handlers, i) + err = nil + } else { + ts.error("TCPStack:Handle", slog.Uint64("lport", uint64(h.lport))) + continue + } + } + if n > 0 { + break + } + } + if n == 0 { + return 0, err + } + // TCP packet written. + tfrm, _ := lneto.NewTCPFrame(ipFrame[tcpOff:]) + ts.validator.ResetErr() + tfrm.ValidateSize(&ts.validator) // Perform basic validation. + if err = ts.validator.Err(); err != nil { + return 0, err + } + var crc uint16 + switch ipVersion { + case 4: + ifrm, _ := lneto.NewIPv4Frame(ipFrame) + crc = tfrm.CalculateIPv4CRC(ifrm) + case 6: + ifrm, _ := lneto.NewIPv6Frame(ipFrame) + crc = tfrm.CalculateIPv6CRC(ifrm) + } + tfrm.SetCRC(crc) + return tcpOff + n, nil +} + +type TCPPort struct { + tcb tcp.ControlBlock + validator lneto.Validator + lport uint16 + rport uint16 +} + +func (tp *TCPPort) Protocol() uint32 { return uint32(lneto.IPProtoTCP) } + +func (tp *TCPPort) Recv(tcpFrame []byte, off int) error { + if off != 0 { + return errors.New("TCP API expected 0 offset") + } + tfrm, err := lneto.NewTCPFrame(tcpFrame) + if err != nil { + return err + } + tp.validator.ResetErr() + tfrm.ValidateExceptCRC(&tp.validator) + if err = tp.validator.Err(); err != nil { + return err + } + if tfrm.DestinationPort() != tp.lport { + return errors.New("port mismatch") + } + seg := tfrm.Segment(len(tfrm.Payload())) + err = tp.tcb.Recv(seg) if err != nil { return err } return nil - +} + +func (tp *TCPPort) Handle(tcpFrame []byte, off int) (n int, err error) { + if off != 0 { + return 0, errors.New("TCP API expected 0 offset") + } else if tp.tcb.State().IsClosed() { + return 0, io.EOF + } + tfrm, err := lneto.NewTCPFrame(tcpFrame) + if err != nil { + return 0, err + } + if !tp.tcb.HasPending() { + return 0, nil + } + + seg, ok := tp.tcb.PendingSegment(0) + if !ok { + return 0, nil + } + err = tp.tcb.Send(seg) + if err != nil { + return 0, err + } + tfrm.SetSourcePort(tp.lport) + tfrm.SetDestinationPort(tp.rport) + tfrm.SetOffsetAndFlags(5, seg.Flags) + tfrm.SetSeq(seg.SEQ) + tfrm.SetAck(seg.ACK) + tfrm.SetUrgentPtr(0) + tfrm.SetWindowSize(uint16(seg.WND)) + + return 20, nil +} + +type logger struct { + log *slog.Logger +} + +func (l logger) error(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, slog.LevelError, msg, attrs...) +} +func (l logger) info(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, slog.LevelInfo, msg, attrs...) +} +func (l logger) warn(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, slog.LevelWarn, msg, attrs...) +} +func (l logger) debug(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, slog.LevelDebug, msg, attrs...) +} + +func removeHandler(handlers []handler, idxRemoved int) []handler { + return append(handlers[:idxRemoved], handlers[idxRemoved+1:]...) +} + +func addHandler(handlers []handler, h Handler, remoteAddr []byte, lport uint16) []handler { + // Pattern to add a handler and reuse underlying memory. + handlers = append(handlers, handler{}) + hh := &handlers[len(handlers)-1] + hh.handle = h.Handle + hh.recv = h.Recv + hh.proto = h.Protocol() + if remoteAddr != nil { + // Remote IP address specified. + hh.raddr = append(hh.raddr, remoteAddr[:]...) + } + hh.lport = lport + return handlers } diff --git a/internal/ltesto/ltesto.go b/internal/ltesto/ltesto.go new file mode 100644 index 0000000..947e68b --- /dev/null +++ b/internal/ltesto/ltesto.go @@ -0,0 +1,152 @@ +package ltesto + +import ( + "bytes" + "math/rand" + + "github.com/soypat/lneto" + "github.com/soypat/lneto/tcp" +) + +const ( + sizeHeaderIPv4 = 20 + sizeHeaderTCP = 20 + sizeHeaderEthNoVLAN = 14 + sizeHeaderUDP = 8 + sizeHeaderARPv4 = 28 + sizeHeaderIPv6 = 40 +) + +type PacketGen struct { + SrcMAC, DstMAC [6]byte // hardware address + SrcIPv4, DstIPv4 [4]byte // address + SrcTCP, DstTCP uint16 // ports + EnableVLAN bool +} + +func (gen *PacketGen) RandomizeAddrs(rng *rand.Rand) { + rng.Read(gen.SrcMAC[:]) + rng.Read(gen.DstMAC[:]) + rng.Read(gen.SrcIPv4[:]) + rng.Read(gen.DstIPv4[:]) + ports := rng.Uint32() + gen.SrcTCP = uint16(ports) + gen.DstTCP = uint16(ports >> 16) +} + +func (gen *PacketGen) AppendRandomIPv4TCPPacket(dst []byte, rng *rand.Rand) []byte { + ri := rng.Int() + var ( + isVLAN = ri&(1<<0) != 0 + hasIPOpt = ri&(1<<1) != 0 + hasTCPOpt = ri&(1<<2) != 0 + hasPayload = ri&(1<<3) != 0 + ) + var etherType lneto.EtherType = lneto.EtherTypeIPv4 + var ipOpts []byte + if hasIPOpt { + ipOpts = []byte{1, 2, 3, 4} + } + ethsize := 14 + if gen.EnableVLAN && isVLAN { + etherType = lneto.EtherTypeVLAN + ethsize = 18 + } + var tcpOpts []byte + if hasTCPOpt { + tcpOpts = []byte{byte(tcp.OptSACKPermitted), 0, 1, 0} + } + var payloadLen int + if hasPayload { + payloadLen = (ri >> 16) % 1024 + } + ipOptWLen := sizeWord(len(ipOpts)) + tcpOptWlen := sizeWord(len(tcpOpts)) + off := len(dst) + dst = append(dst, make([]byte, ethsize+sizeHeaderIPv4+4*int(ipOptWLen)+sizeHeaderTCP+4*int(tcpOptWlen)+payloadLen)...) + efrm, err := lneto.NewEthFrame(dst[off:]) + if err != nil { + panic(err) + } + *efrm.DestinationHardwareAddr() = gen.DstMAC + *efrm.SourceHardwareAddr() = gen.SrcMAC + + efrm.SetEtherType(etherType) + if isVLAN { + efrm.SetVLANEtherType(lneto.EtherTypeIPv4) + efrm.SetVLANTag(1 << 4) + } + ethernetPayload := efrm.Payload() + ifrm, err := lneto.NewIPv4Frame(ethernetPayload) + if err != nil { + panic(err) + } + ifrm.SetVersionAndIHL(4, sizeWord(20+len(ipOpts))) + ifrm.SetToS(192) + ifrm.SetTotalLength(uint16(len(ethernetPayload))) + ifrm.SetID(uint16(rng.Uint32())) + ifrm.SetFlags(0x4001) // Don't fragment. + ifrm.SetTTL(64) + ifrm.SetProtocol(lneto.IPProtoTCP) + *ifrm.SourceAddr() = gen.SrcIPv4 + *ifrm.DestinationAddr() = gen.DstIPv4 + ifrm.SetCRC(ifrm.CalculateHeaderCRC()) + + ipPayload := ifrm.Payload() + tfrm, err := lneto.NewTCPFrame(ipPayload) + if err != nil { + panic(err) + } + tfrm.SetSourcePort(gen.SrcTCP) + tfrm.SetDestinationPort(gen.DstTCP) + tfrm.SetSeq(tcp.Value(rng.Uint32())) + tfrm.SetAck(tcp.Value(rng.Uint32())) + wlen := sizeWord(sizeHeaderTCP + len(tcpOpts)) + tfrm.SetOffsetAndFlags(wlen, tcp.Flags(rng.Uint32())) + tfrm.SetWindowSize(uint16(rng.Uint32())) + urgPtr := uint16(rng.Uint32()) + tfrm.SetUrgentPtr(urgPtr) + tcpPayload := tfrm.Payload() + var firstPayloadByte byte + if len(tcpPayload) > 0 { + rng.Read(tcpPayload) + firstPayloadByte = tcpPayload[0] + } + // Set Variable section of data. + copy(ifrm.Options(), ipOpts) + copy(tfrm.Options(), tcpOpts) + tcpCRC := tfrm.CalculateIPv4CRC(ifrm) + tfrm.SetCRC(tcpCRC) + switch { + case gen.SrcTCP != tfrm.SourcePort(): + panic("IP options overwrite TCP header") + case !bytes.Equal(ifrm.Options(), ipOpts): + panic("bad ip options written, ensure ip options length is multiple of 4") + case !bytes.Equal(tfrm.Options(), tcpOpts): + panic("bad tcp options written, ensure tcp options length is multiple of 4") + case *ifrm.DestinationAddr() != gen.DstIPv4: + panic("IP options overwrite own header") + case tfrm.UrgentPtr() != urgPtr: + panic("TCP options overwrite urgent pointer field?") + case len(tcpPayload) > 0 && firstPayloadByte != tcpPayload[0]: + panic("TCP options overwrite payload") + } + var vld lneto.Validator + efrm.ValidateSize(&vld) + if err = vld.Err(); err != nil { + panic(err) + } + ifrm.ValidateExceptCRC(&vld) + if err = vld.Err(); err != nil { + panic(err) + } + tfrm.ValidateSize(&vld) + if err = vld.Err(); err != nil { + panic(err) + } + return dst +} + +func sizeWord(l int) uint8 { + return uint8((l + 3) / 4) +} diff --git a/lneto_test.go b/lneto_test.go index b4fc5e6..91f218a 100644 --- a/lneto_test.go +++ b/lneto_test.go @@ -5,18 +5,18 @@ import ( "math/rand" "testing" - "github.com/soypat/lneto/tcp" + "github.com/soypat/lneto/internal/ltesto" ) func TestTCPMarshalUnmarshal(t *testing.T) { rng := rand.New(rand.NewSource(1)) - var gen packetGen - gen.randomizeAddrs(rng) + var gen ltesto.PacketGen + gen.RandomizeAddrs(rng) const maxSize = 4096 src := make([]byte, maxSize) dst := make([]byte, maxSize) for i := 0; i < 512; i++ { - src = gen.appendRandomIPv4TCPPacket(src[:0], rng) + src = gen.AppendRandomIPv4TCPPacket(src[:0], rng) dst = dst[:len(src)] testMoveTCPPacket(t, src, dst) if !bytes.Equal(src, dst) { @@ -105,134 +105,3 @@ func testMoveTCPPacket(t *testing.T, src, dst []byte) { t.Fatalf("payload mismatch %d %d", len(payload), len(tfrm2.Payload())) } } - -type packetGen struct { - srcMAC, dstMAC [6]byte // hardware address - srcIPv4, dstIPv4 [4]byte // address - srcTCP, dstTCP uint16 // ports -} - -func (gen *packetGen) randomizeAddrs(rng *rand.Rand) { - rng.Read(gen.srcMAC[:]) - rng.Read(gen.dstMAC[:]) - rng.Read(gen.srcIPv4[:]) - rng.Read(gen.dstIPv4[:]) - ports := rng.Uint32() - gen.srcTCP = uint16(ports) - gen.dstTCP = uint16(ports >> 16) -} - -func (gen *packetGen) appendRandomIPv4TCPPacket(dst []byte, rng *rand.Rand) []byte { - ri := rng.Int() - var ( - isVLAN = ri&(1<<0) != 0 - hasIPOpt = ri&(1<<1) != 0 - hasTCPOpt = ri&(1<<2) != 0 - hasPayload = ri&(1<<3) != 0 - ) - var etherType EtherType = EtherTypeIPv4 - var ipOpts []byte - if hasIPOpt { - ipOpts = []byte{1, 2, 3, 4} - } - ethsize := 14 - if isVLAN { - etherType = EtherTypeVLAN - ethsize = 18 - } - var tcpOpts []byte - if hasTCPOpt { - tcpOpts = []byte{byte(tcp.OptSACKPermitted), 0, 1, 0} - } - var payloadLen int - if hasPayload { - payloadLen = (ri >> 16) % 1024 - } - ipOptWLen := sizeWord(len(ipOpts)) - tcpOptWlen := sizeWord(len(tcpOpts)) - off := len(dst) - dst = append(dst, make([]byte, ethsize+sizeHeaderIPv4+4*int(ipOptWLen)+sizeHeaderTCP+4*int(tcpOptWlen)+payloadLen)...) - efrm, err := NewEthFrame(dst[off:]) - if err != nil { - panic(err) - } - *efrm.DestinationHardwareAddr() = gen.dstMAC - *efrm.SourceHardwareAddr() = gen.srcMAC - - efrm.SetEtherType(etherType) - if isVLAN { - efrm.SetVLANEtherType(EtherTypeIPv4) - efrm.SetVLANTag(1 << 4) - } - ethernetPayload := efrm.Payload() - ifrm, err := NewIPv4Frame(ethernetPayload) - if err != nil { - panic(err) - } - ifrm.SetVersionAndIHL(4, sizeWord(20+len(ipOpts))) - ifrm.SetToS(192) - ifrm.SetTotalLength(uint16(len(ethernetPayload))) - ifrm.SetID(uint16(rng.Uint32())) - ifrm.SetFlags(0x4001) // Don't fragment. - ifrm.SetTTL(64) - ifrm.SetProtocol(IPProtoTCP) - *ifrm.SourceAddr() = gen.srcIPv4 - *ifrm.DestinationAddr() = gen.dstIPv4 - ifrm.SetCRC(ifrm.CalculateHeaderCRC()) - - ipPayload := ifrm.Payload() - tfrm, err := NewTCPFrame(ipPayload) - if err != nil { - panic(err) - } - tfrm.SetSourcePort(gen.srcTCP) - tfrm.SetDestinationPort(gen.dstTCP) - tfrm.SetSeq(tcp.Value(rng.Uint32())) - tfrm.SetAck(tcp.Value(rng.Uint32())) - wlen := sizeWord(sizeHeaderTCP + len(tcpOpts)) - tfrm.SetOffsetAndFlags(wlen, tcp.Flags(rng.Uint32())) - tfrm.SetWindowSize(uint16(rng.Uint32())) - urgPtr := uint16(rng.Uint32()) - tfrm.SetUrgentPtr(urgPtr) - tcpPayload := tfrm.Payload() - var firstPayloadByte byte - if len(tcpPayload) > 0 { - rng.Read(tcpPayload) - firstPayloadByte = tcpPayload[0] - } - // Set Variable section of data. - copy(ifrm.Options(), ipOpts) - copy(tfrm.Options(), tcpOpts) - switch { - case gen.srcTCP != tfrm.SourcePort(): - panic("IP options overwrite TCP header") - case !bytes.Equal(ifrm.Options(), ipOpts): - panic("bad ip options written, ensure ip options length is multiple of 4") - case !bytes.Equal(tfrm.Options(), tcpOpts): - panic("bad tcp options written, ensure tcp options length is multiple of 4") - case *ifrm.DestinationAddr() != gen.dstIPv4: - panic("IP options overwrite own header") - case tfrm.UrgentPtr() != urgPtr: - panic("TCP options overwrite urgent pointer field?") - case len(tcpPayload) > 0 && firstPayloadByte != tcpPayload[0]: - panic("TCP options overwrite payload") - } - var vld Validator - efrm.ValidateSize(&vld) - if err = vld.Err(); err != nil { - panic(err) - } - ifrm.Validate(&vld) - if err = vld.Err(); err != nil { - panic(err) - } - tfrm.ValidateSize(&vld) - if err = vld.Err(); err != nil { - panic(err) - } - return dst -} - -func sizeWord(l int) uint8 { - return uint8((l + 3) / 4) -} diff --git a/validation.go b/validation.go index 6346163..dce7cba 100644 --- a/validation.go +++ b/validation.go @@ -17,11 +17,14 @@ var ( errBadIPVersion = errors.New("bad IP version field") errEvilPacket = errors.New("evil packet") + errZeroDstPort = errors.New("TCP zero destination port") + errZeroSrcPort = errors.New("TCP zero source port") ) type Validator struct { - checkEvil bool - accum []error + checkEvil bool + allowMultiErrs bool + accum []error } func (v *Validator) ResetErr() { @@ -38,6 +41,9 @@ func (v *Validator) Err() error { } func (v *Validator) gotErr(err error) { + if len(v.accum) != 0 && !v.allowMultiErrs { + return + } v.accum = append(v.accum, err) } @@ -80,7 +86,9 @@ func (ifrm IPv4Frame) ValidateSize(v *Validator) { } } -func (ifrm IPv4Frame) ValidateFields(v *Validator) { +// ValidateExceptCRC checks for invalid frame values but does not check CRC. +func (ifrm IPv4Frame) ValidateExceptCRC(v *Validator) { + ifrm.ValidateSize(v) flags := ifrm.Flags() if ifrm.version() != 4 { v.gotErr(errBadIPVersion) @@ -90,12 +98,6 @@ func (ifrm IPv4Frame) ValidateFields(v *Validator) { } } -// Validate checks for invalid frame values. -func (ifrm IPv4Frame) Validate(v *Validator) { - ifrm.ValidateSize(v) - ifrm.ValidateFields(v) -} - // ValidateSize checks the frame's size fields and compares with the actual buffer // the frame. It returns a non-nil error on finding an inconsistency. func (i6frm IPv6Frame) ValidateSize(v *Validator) { @@ -117,6 +119,16 @@ func (tfrm TCPFrame) ValidateSize(v *Validator) { } } +func (tfrm TCPFrame) ValidateExceptCRC(v *Validator) { + tfrm.ValidateSize(v) + if tfrm.DestinationPort() == 0 { + v.gotErr(errZeroDstPort) + } + if tfrm.SourcePort() == 0 { + v.gotErr(errZeroSrcPort) + } +} + // ValidateSize checks the frame's size fields and compares with the actual buffer // the frame. It returns a non-nil error on finding an inconsistency. func (ufrm UDPFrame) ValidateSize(v *Validator) {