Checksum simplification proposal (#29)

* Simplified checksum calculation and verification

* Fixed tests

* UDP: using NewBoundedFrame to create a Frame instance limited to the actual frame size.

* Minor: renamed some functions

* Commented and made more readable consecutive SetCRC calls.

* Fixed icmp checksum calculation after rebase

* Replaced NewBoundedFrame with explicit validation
This commit is contained in:
ddirect
2026-02-05 22:35:42 +02:00
committed by GitHub
parent cd51bd6c3f
commit fba97c990a
13 changed files with 139 additions and 228 deletions
+33 -53
View File
@@ -161,27 +161,26 @@ func (pc *PacketBreakdown) CaptureIPv6(dst []Frame, pkt []byte, bitOffset int) (
end := bitOffset + 40*octet
var protoErrs []error
var crc lneto.CRC791
if proto == lneto.IPProtoTCP {
ifrm6.CRCWritePseudo(&crc)
tfrm, err := tcp.NewFrame(ifrm6.Payload())
if err == nil {
tfrm.CRCWrite(&crc)
wantSum := crc.Sum16()
gotSum := tfrm.CRC()
if wantSum != gotSum {
protoErrs = append(protoErrs, &crcError16{protocol: "ipv6+tcp", want: wantSum, got: gotSum})
}
ifrm6.CRCWritePseudo(&crc)
switch proto {
case lneto.IPProtoTCP:
if crc.PayloadSum16(ifrm6.Payload()) != 0 {
protoErrs = append(protoErrs, lneto.ErrBadCRC)
}
} else if proto == lneto.IPProtoUDP || proto == lneto.IPProtoUDPLite {
ifrm6.CRCWritePseudo(&crc)
case lneto.IPProtoUDP, lneto.IPProtoUDPLite:
ufrm, err := udp.NewFrame(ifrm6.Payload())
if err == nil {
ufrm.CRCWriteIPv6(&crc)
wantSum := crc.Sum16()
gotSum := ufrm.CRC()
if wantSum != gotSum {
protoErrs = append(protoErrs, &crcError16{protocol: "ipv6+udp", want: wantSum, got: gotSum})
}
if err != nil {
protoErrs = append(protoErrs, err)
break
}
ufrm.ValidateSize(pc.validator())
if err = pc.validator().ErrPop(); err != nil {
protoErrs = append(protoErrs, err)
break
}
frameLen := ufrm.Length()
if crc.PayloadSum16(ufrm.RawData()[:frameLen]) != 0 {
protoErrs = append(protoErrs, lneto.ErrBadCRC)
}
}
return pc.captureIPProto(proto, dst, pkt, end, protoErrs...)
@@ -214,57 +213,48 @@ func (pc *PacketBreakdown) CaptureIPv4(dst []Frame, pkt []byte, bitOffset int) (
BitLength: octet * len(options),
})
}
gotSum := ifrm4.CRC()
wantSum := ifrm4.CalculateHeaderCRC()
if gotSum != wantSum {
finfo.Errors = append(finfo.Errors, &crcError16{protocol: "ipv4", want: wantSum, got: gotSum})
if ifrm4.CalculateHeaderCRC() != 0 {
finfo.Errors = append(finfo.Errors, lneto.ErrBadCRC)
}
dst = append(dst, finfo)
proto := ifrm4.Protocol()
end := bitOffset + octet*ifrm4.HeaderLength()
var protoErrs []error
var crc lneto.CRC791
payload := ifrm4.Payload()
switch proto {
case lneto.IPProtoTCP:
ifrm4.CRCWriteTCPPseudo(&crc)
tfrm, err := tcp.NewFrame(ifrm4.Payload())
tfrm, err := tcp.NewFrame(payload)
if err == nil {
tfrm.ValidateSize(pc.validator())
if pc.vld.HasError() {
println("BAD TCP")
return dst, pc.vld.ErrPop()
}
tfrm.CRCWrite(&crc)
wantSum := crc.Sum16()
gotSum := tfrm.CRC()
if wantSum != gotSum {
protoErrs = append(protoErrs, &crcError16{protocol: "ipv4+tcp", want: wantSum, got: gotSum})
ifrm4.CRCWriteTCPPseudo(&crc)
if crc.PayloadSum16(payload) != 0 {
protoErrs = append(protoErrs, lneto.ErrBadCRC)
}
}
case lneto.IPProtoUDP:
ifrm4.CRCWriteUDPPseudo(&crc)
ufrm, err := udp.NewFrame(ifrm4.Payload())
ufrm, err := udp.NewFrame(payload)
if err == nil {
ufrm.ValidateSize(pc.validator())
if pc.vld.HasError() {
println("BAD UDP")
return dst, pc.vld.ErrPop()
}
ufrm.CRCWriteIPv4(&crc)
wantSum := crc.Sum16()
gotSum := ufrm.CRC()
if wantSum != gotSum {
protoErrs = append(protoErrs, &crcError16{protocol: "ipv4+udp", want: wantSum, got: gotSum})
frameLen := ufrm.Length()
ifrm4.CRCWriteUDPPseudo(&crc, frameLen)
if crc.PayloadSum16(ufrm.RawData()[:frameLen]) != 0 {
protoErrs = append(protoErrs, lneto.ErrBadCRC)
}
}
case lneto.IPProtoICMP:
ifrm, err := icmpv4.NewFrame(ifrm4.Payload())
_, err := icmpv4.NewFrame(payload)
if err == nil {
ifrm.CRCWrite(&crc)
wantSum := crc.Sum16()
gotSum := ifrm.CRC()
if wantSum != gotSum {
protoErrs = append(protoErrs, &crcError16{protocol: "icmpv4", want: wantSum, got: gotSum})
if crc.PayloadSum16(payload) != 0 {
protoErrs = append(protoErrs, lneto.ErrBadCRC)
}
}
}
@@ -1309,13 +1299,3 @@ func remainingFrameInfo(proto any, class FieldClass, pktBitOffset, pktBitLen int
}},
}
}
type crcError16 struct {
protocol string
want uint16
got uint16
}
func (cerr *crcError16) Error() string {
return fmt.Sprintf("%s:incorrect checksum. want 0x%x, got 0x%x", cerr.protocol, cerr.want, cerr.got)
}
+1
View File
@@ -447,6 +447,7 @@ func ExampleFormatter_dhcp() {
totalLen := ipv4Size + udpSize + dhcpLen
ifrm.SetTotalLength(uint16(totalLen))
ufrm.SetLength(uint16(udpSize + dhcpLen))
// The CRC field is already zero here.
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
pkt = pkt[:ethSize+totalLen]
+30 -34
View File
@@ -10,7 +10,6 @@ import (
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
"github.com/soypat/lneto/ipv4"
"github.com/soypat/lneto/ipv4/icmpv4"
"github.com/soypat/lneto/tcp"
"github.com/soypat/lneto/udp"
)
@@ -91,11 +90,9 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
return err
}
// Incoming CRC Validation of common IP Protocols.
var crc lneto.CRC791
ifrm.CRCWriteHeader(&crc)
if !crc.VerifySum16(ifrm.CRC()) {
if ifrm.CalculateHeaderCRC() != 0 {
sb.handlers.error("ip:demux.crc")
return lneto.ErrBadCRC
}
off := ifrm.HeaderLength()
totalLen := ifrm.TotalLength()
@@ -111,28 +108,27 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
return lneto.ErrPacketDrop
}
// Incoming CRC Validation of common IP Protocols.
crc.Reset()
var crc lneto.CRC791
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
tfrm, err := tcp.NewFrame(ifrm.Payload())
if err != nil {
return err
}
tfrm.CRCWrite(&crc)
if !crc.VerifySum16(tfrm.CRC()) {
if crc.PayloadSum16(ifrm.Payload()) != 0 {
sb.handlers.error("ip:demux.tcpcrc")
return lneto.ErrBadCRC
}
case lneto.IPProtoUDP:
ifrm.CRCWriteUDPPseudo(&crc)
ufrm, err := udp.NewFrame(ifrm.Payload())
if err != nil {
return err
}
ufrm.CRCWriteIPv4(&crc)
// checksums are optional in UDP: the field is set to zero in this case
if gotCrc := ufrm.CRC(); gotCrc != 0 && !crc.VerifySum16(gotCrc) {
ufrm.ValidateSize(&sb.validator)
if err = sb.validator.ErrPop(); err != nil {
sb.handlers.error("ip:demux.udpvalidatesize")
return err
}
frameLen := ufrm.Length()
ifrm.CRCWriteUDPPseudo(&crc, frameLen)
if crc.PayloadSum16(ufrm.RawData()[:frameLen]) != 0 {
sb.handlers.error("ip:demux.udpcrc")
return lneto.ErrBadCRC
}
@@ -174,25 +170,29 @@ func (sb *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int
totalLen := n + headerlen
ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetProtocol(proto)
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
// Zero the CRC field so its value does not add to the final result.
ifrm.SetCRC(0)
crcValue := ifrm.CalculateHeaderCRC()
ifrm.SetCRC(crcValue)
// Calculate CRC for our newly generated packet.
var crc lneto.CRC791
payload := ifrm.Payload()
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
tfrm, _ := tcp.NewFrame(ifrm.Payload())
tfrm.CRCWrite(&crc)
tfrm.SetCRC(crc.Sum16())
tfrm, _ := tcp.NewFrame(payload)
// Zero the CRC field so its value does not add to the final result.
tfrm.SetCRC(0)
crcValue = crc.PayloadSum16(payload)
tfrm.SetCRC(crcValue)
case lneto.IPProtoUDP:
ifrm.CRCWriteUDPPseudo(&crc)
ufrm, _ := udp.NewFrame(ifrm.Payload())
ufrm, _ := udp.NewFrame(payload)
ifrm.CRCWriteUDPPseudo(&crc, uint16(n))
ufrm.SetLength(uint16(n))
ufrm.CRCWriteIPv4(&crc)
ufrm.SetCRC(crc.Sum16())
if n != int(ufrm.Length()) {
sb.handlers.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length())))
return 0, errors.New("invalid UDP length")
}
// Zero the CRC field so its value does not add to the final result.
ufrm.SetCRC(0)
crcValue = lneto.NeverZeroSum(crc.PayloadSum16(payload))
ufrm.SetCRC(crcValue)
}
return totalLen, err
}
@@ -206,13 +206,9 @@ func (sb *StackIP) Register(h StackNode) error {
}
func (sb *StackIP) recvicmp(carrierData []byte, offset int) error {
frameData := carrierData[offset:]
var crc lneto.CRC791
cfrm, err := icmpv4.NewFrame(carrierData[offset:])
if err != nil {
return err
}
cfrm.CRCWrite(&crc)
if !crc.VerifySum16(cfrm.CRC()) {
if crc.PayloadSum16(frameData) != 0 {
return errors.New("ICMP CRC mismatch")
}
return nil