diff --git a/arp/handler.go b/arp/handler.go index 4ac7a6a..076a335 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -9,6 +9,7 @@ import ( ) type Handler struct { + connID uint64 ourHWAddr []byte ourProtoAddr []byte htype uint16 @@ -26,22 +27,31 @@ type HandlerConfig struct { ProtocolType ethernet.Type } -func NewHandler(cfg HandlerConfig) (*Handler, error) { +func (c *Handler) Reset(cfg HandlerConfig) error { if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 || len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 { - return nil, errors.New("invalid Handler address config") + return errors.New("invalid Handler address config") } else if cfg.MaxQueries <= 0 || cfg.MaxPending <= 0 { - return nil, errors.New("invalid Handler query or pending config") + return errors.New("invalid Handler query or pending config") } - h := &Handler{ - pending: make([][sizeHeaderv6]byte, 0, cfg.MaxPending), + *c = Handler{ + connID: c.connID + 1, + ourHWAddr: c.ourHWAddr[:0], + ourProtoAddr: c.ourProtoAddr[:0], htype: cfg.HardwareType, protoType: cfg.ProtocolType, - ourHWAddr: cfg.HardwareAddr, - ourProtoAddr: cfg.ProtocolAddr, - queries: make([]queryResult, 0, cfg.MaxQueries), + pending: c.pending[:0], + queries: c.queries[:0], } - return h, nil + c.ourHWAddr = append(c.ourHWAddr, cfg.HardwareAddr...) + c.ourProtoAddr = append(c.ourProtoAddr, cfg.ProtocolAddr...) + if cap(c.pending) < cfg.MaxPending { + c.pending = make([][52]byte, cfg.MaxPending)[:0] + } + if cap(c.queries) < cfg.MaxQueries { + c.queries = make([]queryResult, cfg.MaxQueries)[:0] + } + return nil } type queryResult struct { @@ -50,8 +60,8 @@ type queryResult struct { querysent bool } -// ResetState drops pending queries and incoming requests. -func (c *Handler) ResetState() { +// AbortPending drops pending queries and incoming requests. +func (c *Handler) AbortPending() { c.pending = c.pending[:0] c.queries = c.queries[:0] } @@ -60,6 +70,10 @@ func (c *Handler) expectSize() int { return sizeHeader + 2*len(c.ourHWAddr) + 2*len(c.ourProtoAddr) } +func (c *Handler) ConnectionID() *uint64 { + return &c.connID +} + func (c *Handler) QueryResult(protoAddr []byte) (hwAddr []byte, err error) { for i := range c.queries { if bytes.Equal(protoAddr, c.queries[i].protoaddr) { diff --git a/arp/handler_test.go b/arp/handler_test.go index 85bf8f6..009e293 100644 --- a/arp/handler_test.go +++ b/arp/handler_test.go @@ -5,11 +5,13 @@ import ( "log" "testing" + "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" ) func TestHandler(t *testing.T) { - c1, err := NewHandler(HandlerConfig{ + var c1, c2 Handler + err := c1.Reset(HandlerConfig{ HardwareAddr: []byte{0xde, 0xad, 0xbe, 0xef, 0x00, 0x00}, ProtocolAddr: []byte{192, 168, 1, 1}, MaxQueries: 1, @@ -20,7 +22,7 @@ func TestHandler(t *testing.T) { if err != nil { t.Fatal(err) } - c2, err := NewHandler(HandlerConfig{ + err = c2.Reset(HandlerConfig{ HardwareAddr: []byte{0xc0, 0xff, 0xee, 0xc0, 0xff, 0xee}, ProtocolAddr: []byte{192, 168, 1, 2}, MaxQueries: 1, @@ -58,6 +60,7 @@ func TestHandler(t *testing.T) { } else if n == 0 { t.Fatal("expected send of data after first query") } + validateARP(t, buf[:]) err = c2.Recv(buf[:n]) // Receive request. if err != nil { t.Fatal(err) @@ -69,6 +72,7 @@ func TestHandler(t *testing.T) { } else if n == 0 { t.Fatal("got no response to request") } + validateARP(t, buf[:]) n, err = c2.Send(discard[:]) // Double tap check, should send nothing. if err != nil { t.Fatal("double tap send error:", err) @@ -99,3 +103,19 @@ func TestHandler(t *testing.T) { t.Fatal("expected no data") } } + +func validateARP(t *testing.T, buf []byte) { + t.Helper() + afrm, err := NewFrame(buf) + if err != nil { + t.Error(err) + return + } + var vld lneto.Validator + afrm.ValidateSize(&vld) + if vld.HasError() { + t.Errorf("invalid arp: %s", vld.Err()) + } else if err := vld.Err(); err != nil { + panic("unreachable: " + err.Error()) + } +} diff --git a/examples/stack/main.go b/examples/stack/main.go index e953c1b..edbf158 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -1,29 +1,26 @@ package main import ( - "errors" "fmt" - "io" "log" "log/slog" "net" "net/netip" "os" + "runtime" "time" "github.com/soypat/lneto" "github.com/soypat/lneto/arp" "github.com/soypat/lneto/ethernet" - "github.com/soypat/lneto/internal" + "github.com/soypat/lneto/http/httpraw" "github.com/soypat/lneto/internal/ltesto" - "github.com/soypat/lneto/ipv4" - "github.com/soypat/lneto/ipv6" + "github.com/soypat/lneto/internet" + "github.com/soypat/lneto/internet/pcap" "github.com/soypat/lneto/tcp" ) const ( - mtu = 2048 - iface = "192.168.10.1/24" stackIP = "192.168.10.2" stackPort = 80 iss = 100 @@ -33,513 +30,228 @@ var stackHWAddr = [6]byte{0xc0, 0xff, 0xee, 0x00, 0xde, 0xad} func main() { ip := netip.MustParseAddr(stackIP) - iface := netip.MustParsePrefix(iface) - if !iface.Contains(ip) { + tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070") + + ippfx := tap.IPPrefix() + if !ippfx.Contains(ip) { log.Fatal("interface does not contain stack address") } addrPort := netip.AddrPortFrom(ip, stackPort) - slogger := logger{slog.Default()} - lStack, handler, err := NewEthernetTCPStack(stackHWAddr, addrPort, slogger) - if err != nil { - log.Fatal(err) - } - - logger := slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + lg := slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ Level: slog.LevelDebug, })) - handler.SetLoggers(logger, logger) - err = handler.OpenListen(addrPort.Port(), iss) + gatewayMAC := tap.HardwareAddr6() + mtu := tap.MTU() + stack, err := NewEthernetTCPStack(stackHWAddr, gatewayMAC, addrPort, uint16(mtu)) + if err != nil { + log.Fatal(err) + } + handler, err := stack.OpenPassiveTCP(addrPort.Port(), iss) if err != nil { log.Fatal(err) } - tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070") defer tap.Close() - - fmt.Println("hosting server at ", addrPort.String()) - var buf [mtu]byte + tap.ReadDiscard() // Discard all unread content. + fmt.Println("hosting server at ", addrPort.String(), "over tap interface of mtu:", mtu, "prefix:", ippfx, "gateway:", net.HardwareAddr(gatewayMAC[:]).String()) + buf := make([]byte, mtu) + var hdr httpraw.Header + hdr.Reset(make([]byte, 0, 1024)) + const standbyDuration = 5 * time.Second + lastHit := time.Now().Add(-standbyDuration) + var cap pcap.PacketBreakdown for { nread, err := tap.Read(buf[:]) if err != nil { - slogger.error("tap-err", slog.String("err", err.Error())) + lg.Error("tap-err", slog.String("err", err.Error())) log.Fatal(err) } else if nread > 0 { - err = lStack.RecvEth(buf[:nread]) + frames, err := cap.CaptureEthernet(nil, buf[:nread], 0) + if err == nil { + flags := getTCPFlags(frames, buf[:nread]) + if flags == 0 { + fmt.Println("IN", time.Now().Format("15:04:05.000"), frames) + } else { + fmt.Println("IN", time.Now().Format("15:04:05.000"), frames, flags.String()) + } + } + err = stack.ethernet.Demux(buf[:nread], 0) if err != nil { - slogger.error("recv", slog.String("err", err.Error()), slog.Int("plen", nread)) - } else { - slogger.info("recv", slog.Int("plen", nread)) + lg.Error("recv", slog.String("err", err.Error()), slog.Int("plen", nread)) } } - nw, err := lStack.HandleEth(buf[:]) + doHTTP(handler, &hdr) + nw, err := stack.ethernet.Encapsulate(buf[:], 0) if err != nil { - slogger.error("handle", slog.String("err", err.Error())) + lg.Error("handle", slog.String("err", err.Error())) } else if nw > 0 { + frames, err := cap.CaptureEthernet(nil, buf[:nread], 0) + if err == nil { + flags := getTCPFlags(frames, buf[:nread]) + if flags == 0 { + fmt.Println("OU", time.Now().Format("15:04:05.000"), frames) + } else { + fmt.Println("OU", time.Now().Format("15:04:05.000"), frames, flags.String()) + } + } _, err = tap.Write(buf[:nw]) if err != nil { log.Fatal(err) - } else { - slogger.info("write", slog.Int("plen", nw)) } } - if nread == 0 && nw == 0 { - time.Sleep(5 * time.Millisecond) + hit := nread > 0 || nw > 0 + if hit { + // slogger.info("exchange", slog.Int("read", nread), slog.Int("nwrite", nw)) + lastHit = time.Now() + } else { + if time.Since(lastHit) > standbyDuration { + time.Sleep(5 * time.Millisecond) + } else { + runtime.Gosched() + } } } } -func NewEthernetTCPStack(mac [6]byte, ip netip.AddrPort, slogger logger) (*LinkStack, *tcp.Handler, error) { - lStack := LinkStack{ - logger: slogger, - mac: mac, - mtu: mtu, +func doHTTP(conn *internet.TCPConn, hdr *httpraw.Header) error { + const asRequest = false + if conn.State() != tcp.StateEstablished || conn.BufferedInput() == 0 { + return nil // No data yet. } - ipStack := &IPv4Stack{ - ip: ip.Addr().As4(), - logger: slogger, + fmt.Println("state is established; check request and send response") + _, err := hdr.ReadFromLimited(conn, hdr.BufferFree()) + if err != nil { + return err } - tcpStack := &TCPStack{ - logger: slogger, + needMore, err := hdr.TryParse(asRequest) + if err != nil { + if !needMore { + fmt.Println("IT's SO GOVER") + conn.Close() + } + return err } - tcpPortStack := &TCPPort{ - handler: tcp.Handler{}, + // HTTP parsed succesfully! + fmt.Println("GOT HTTP:\n", hdr.String()) + fmt.Println("sending response...") + hdr.Reset(nil) + hdr.SetStatus("200", "OK") + data := `{"ok":true}` + response, err := hdr.AppendResponse(nil) + if err != nil { + return err } - proto := ethernet.TypeIPv4 - if ip.Addr().Is6() { - proto = ethernet.TypeIPv6 + response = append(response, data...) + _, err = conn.Write(response) + if err != nil { + return err } - arphandler, err := arp.NewHandler(arp.HandlerConfig{ - HardwareAddr: mac[:], - ProtocolAddr: ip.Addr().AsSlice(), - MaxQueries: 1, - MaxPending: 1, - HardwareType: 1, - ProtocolType: proto, + err = conn.Close() + if err != nil { + return err + } + return nil +} + +type Stack struct { + ethernet internet.StackLinkLayer + ip internet.StackIP + tcpports internet.StackPorts + arp internet.NodeARP + + onlyConn internet.TCPConn +} + +func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*internet.TCPConn, error) { + mtu := stack.ethernet.MTU() + conn := new(internet.TCPConn) + err := conn.Configure(&internet.TCPConnConfig{ + RxBuf: make([]byte, mtu), + TxBuf: make([]byte, mtu), + TxPacketQueueSize: 3, }) if err != nil { - return nil, nil, err + return nil, err } - arpStack := ARPStack{ - handler: *arphandler, - } - - port := ip.Port() - txbuf := make([]byte, mtu) - rxbuf := make([]byte, mtu) - err = tcpPortStack.handler.SetBuffers(txbuf, rxbuf, 3) + err = conn.OpenListen(port, iss) if err != nil { - return nil, nil, err + return nil, err } - err = ipStack.Register(tcpStack, nil) + err = stack.tcpports.Register(conn) if err != nil { - return nil, nil, err + return nil, err } - err = lStack.Register(ipStack, mac) - if err != nil { - return nil, nil, err - } - err = lStack.Register(&arpStack, [6]byte{0xff, 0xff, 0xff, 0xff, 0xff, 0xff}) - if err != nil { - return nil, nil, err - } - err = tcpStack.Register(tcpPortStack, port) - if err != nil { - return nil, nil, err - } - return &lStack, &tcpPortStack.handler, nil + return conn, nil } -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) { - - efrm, err := ethernet.NewFrame(ethFrame) +func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16) (*Stack, error) { + var stack Stack + var err error + err = stack.ethernet.Reset6(ourMAC, gwMAC, int(mtu)) if err != nil { - return err + return nil, err } - etype := efrm.EtherTypeOrSize() - dstaddr := efrm.DestinationHardwareAddr() - if !efrm.IsBroadcast() && ls.mac != *dstaddr { - return fmt.Errorf("incoming %s mismatch hwaddr %s", etype.String(), net.HardwareAddr(dstaddr[:]).String()) + err = stack.ip.Reset(ip.Addr()) + if err != nil { + return nil, err } - var vld lneto.Validator - efrm.ValidateSize(&vld) - if err := vld.Err(); err != nil { - return err + stack.tcpports.Reset(uint64(lneto.IPProtoTCP), 2) + ipaddr := ip.Addr().As4() + err = stack.arp.Reset(arp.HandlerConfig{ + HardwareAddr: ourMAC[:], + ProtocolAddr: ipaddr[:], + MaxQueries: 2, + MaxPending: 2, + HardwareType: 1, + ProtocolType: ethernet.TypeIPv4, + }) + if err != nil { + return nil, err } - for i := range ls.handlers { - h := &ls.handlers[i] - if h.proto == uint32(etype) { - return h.recv(efrm.Payload(), 0) - } + // Register stacks and nodes. + err = stack.ethernet.Register(&stack.arp) + if err != nil { + return nil, err } - - return nil + err = stack.ethernet.Register(&stack.ip) + if err != nil { + return nil, err + } + err = stack.ip.Register(&stack.tcpports) + if err != nil { + return nil, err + } + return &stack, nil } -func (ls *LinkStack) HandleEth(dst []byte) (n int, err error) { - if len(dst) < int(ls.mtu) { - return 0, io.ErrShortBuffer +func debugHex(b []byte) string { + var d []byte + for i := 0; i < len(b); i++ { + c1 := tblhex[b[i]&0xf] + c2 := tblhex[b[i]>>4] + d = append(d, c2, c1, ' ') } - 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", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) + return string(d) +} + +const tblhex = "0123456789abcdef" + +func getTCPFlags(frames []pcap.Frame, pkt []byte) (flags tcp.Flags) { + for i := range frames { + if frames[i].Protocol != lneto.IPProtoTCP { continue } - if n > 0 { - // Found packet - efrm, _ := ethernet.NewFrame(dst[:14]) - copy(efrm.DestinationHardwareAddr()[:], h.raddr) - *efrm.SourceHardwareAddr() = ls.mac - efrm.SetEtherType(ethernet.Type(h.proto)) - - return n + 14, nil - } - } - return 0, err -} - -type IPv4Stack struct { - ip [4]byte - validator lneto.Validator - handlers []handler - logger -} - -func (is *IPv4Stack) Protocol() uint32 { return uint32(ethernet.TypeIPv4) } - -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 := ipv4.NewFrame(ethFrame[ipOff:]) - if err != nil { - return err - } - if *ifrm.DestinationAddr() != is.ip { - return errors.New("packet not for us") - } - 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, _ := ipv4.NewFrame(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) + iflags, err := frames[i].FieldByClass(pcap.FieldClassFlags) if err != nil { - is.error("IPv4Stack:handle", slog.String("proto", proto.String()), slog.String("err", err.Error())) - continue + return 0 } - 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()) - if ifrm.Protocol() == lneto.IPProtoTCP { - tfrm, _ := tcp.NewFrame(ifrm.Payload()) - is.info("IPv4Stack:send", slog.String("ip", ifrm.String()), slog.String("tcp", tfrm.String())) - } - 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 := tcp.NewFrame(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 - } - crc := tcpChecksum(ipFrame, len(tfrm.RawData())) - 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) + v, err := frames[i].FieldAsUint(iflags, pkt) 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 { - ipFrame = ipFrame[:tcpOff+n] - break + return 0 } + return tcp.Flags(v) } - if n == 0 { - return 0, err - } - // TCP packet written. - tfrm, _ := tcp.NewFrame(ipFrame[tcpOff:]) - - ts.validator.ResetErr() - tfrm.ValidateSize(&ts.validator) // Perform basic validation. - if err = ts.validator.Err(); err != nil { - return 0, err - } - crc := tcpChecksum(ipFrame, n) - tfrm.SetCRC(crc) - return n, nil -} - -type ARPStack struct { - handler arp.Handler -} - -func (as *ARPStack) Protocol() uint32 { return uint32(ethernet.TypeARP) } - -func (as *ARPStack) Recv(EtherFrame []byte, arpOff int) error { - afrm, _ := arp.NewFrame(EtherFrame[arpOff:]) - slog.Info("recv", slog.String("in", afrm.String())) - return as.handler.Recv(EtherFrame[arpOff:]) -} - -func (as *ARPStack) Handle(EtherFrame []byte, arpOff int) (int, error) { - n, err := as.handler.Send(EtherFrame[arpOff:]) - if err != nil || n == 0 { - return 0, err - } - afrm, _ := arp.NewFrame(EtherFrame[arpOff:]) - hwaddr, _ := afrm.Target() - efrm, _ := ethernet.NewFrame(EtherFrame) - copy(efrm.DestinationHardwareAddr()[:], hwaddr) - slog.Info("handle", slog.String("out", afrm.String())) - return n, err -} - -type TCPPort struct { - handler tcp.Handler -} - -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") - } - return tp.handler.Recv(tcpFrame) -} - -func (tp *TCPPort) Handle(tcpFrame []byte, off int) (n int, err error) { - if off != 0 { - return 0, errors.New("TCP API expected 0 offset") - } - return tp.handler.Send(tcpFrame) -} - -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 -} - -func tcpChecksum(ipFrame []byte, tcpPayload int) uint16 { - version := ipFrame[0] >> 4 - var tfrm tcp.Frame - var crc lneto.CRC791 - switch version { - case 4: - ifrm, _ := ipv4.NewFrame(ipFrame) - crc.Write(ifrm.SourceAddr()[:]) - crc.Write(ifrm.DestinationAddr()[:]) - crc.AddUint16(uint16(tcpPayload)) - crc.AddUint16(6) - tfrm, _ = tcp.NewFrame(ifrm.Payload()) - case 6: - i6frm, _ := ipv6.NewFrame(ipFrame) - crc.Write(i6frm.SourceAddr()[:]) - crc.Write(i6frm.DestinationAddr()[:]) - crc.AddUint32(uint32(tcpPayload)) - crc.AddUint32(6) - i6frm.CRCWritePseudo(&crc) - tfrm, _ = tcp.NewFrame(i6frm.Payload()) - default: - panic("invalid IP version") - } - tfrm.CRCWrite(&crc) - return crc.Sum16() + return 0 } diff --git a/examples/stack/main_test.go b/examples/stack/main_test.go deleted file mode 100644 index caf07e3..0000000 --- a/examples/stack/main_test.go +++ /dev/null @@ -1,41 +0,0 @@ -package main - -import ( - "testing" - - "github.com/soypat/lneto/ethernet" - "github.com/soypat/lneto/ipv4" - "github.com/soypat/lneto/tcp" -) - -func TestChecksum(t *testing.T) { - for _, epacket := range ethpackets { - efrm, _ := ethernet.NewFrame(epacket) - ifrm, _ := ipv4.NewFrame(efrm.Payload()) - ipPayload := ifrm.Payload() - tfrm, _ := tcp.NewFrame(ipPayload) - - crc := tcpChecksum(ifrm.RawData(), len(ipPayload)) - wantCRC := tfrm.CRC() - if crc != wantCRC { - t.Fatalf("crc mismatch, got %x, want %x", crc, wantCRC) - } - } -} - -var ethpackets = [][]byte{ - - { - 0xc0, 0xff, 0xee, 0x00, 0xde, 0xad, 0x3a, 0xd1, 0x6d, 0x82, 0x6b, 0x1a, 0x08, 0x00, 0x45, 0x00, - 0x00, 0x3c, 0xe3, 0xc6, 0x40, 0x00, 0x40, 0x06, 0xc1, 0xa1, 0xc0, 0xa8, 0x0a, 0x01, 0xc0, 0xa8, - 0x0a, 0x02, 0xd5, 0x70, 0x00, 0x50, 0xc4, 0x10, 0x30, 0x49, 0x00, 0x00, 0x00, 0x00, 0xa0, 0x02, - 0xfa, 0xf0, 0x2c, 0x80, 0x00, 0x00, 0x02, 0x04, 0x05, 0xb4, 0x04, 0x02, 0x08, 0x0a, 0xe8, 0x22, - 0xd8, 0xfd, 0x00, 0x00, 0x00, 0x00, 0x01, 0x03, 0x03, 0x07, - }, - { - 0xc0, 0xff, 0xee, 0x00, 0xde, 0xad, 0xc0, 0xff, 0xee, 0x00, 0xde, 0xad, 0x08, 0x00, 0x45, 0x00, - 0x00, 0x28, 0x00, 0x00, 0x40, 0x00, 0x40, 0x06, 0x70, 0x26, 0xc0, 0xa8, 0x0a, 0x02, 0x00, 0x00, - 0x00, 0x00, 0x00, 0x50, 0xa5, 0xb8, 0x00, 0x00, 0x00, 0x64, 0x2f, 0x46, 0x4c, 0xe5, 0x50, 0x12, - 0x08, 0x00, 0xBA, 0x90, 0x00, 0x00, - }, -} diff --git a/examples/stackbasic/main.go b/examples/stackbasic/main.go index 6e35dc9..f99595c 100644 --- a/examples/stackbasic/main.go +++ b/examples/stackbasic/main.go @@ -203,7 +203,8 @@ func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16, s if ip.Addr().Is6() { proto = ethernet.TypeIPv6 } - arphandler, err := arp.NewHandler(arp.HandlerConfig{ + var narp arp.Handler + err = narp.Reset(arp.HandlerConfig{ HardwareAddr: ourMAC[:], ProtocolAddr: ip.Addr().AsSlice(), MaxQueries: 4, @@ -215,7 +216,7 @@ func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16, s return nil, nil, err } arpStack := ARPStack{ - handler: *arphandler, + handler: narp, } err = lStack.Register(handler{ recv: arpStack.Recv, diff --git a/internet/definitions.go b/internet/definitions.go index 71565d4..23aa4f8 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -52,10 +52,10 @@ var ( ) func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) { - if nodeIdx >= len(*nodesPtr) { - panic("unreachable") - } if err != nil { + if nodeIdx >= len(*nodesPtr) { + panic("unreachable") + } nodes := *nodesPtr badConnID := nodes[nodeIdx].connID != nil && *nodes[nodeIdx].connID != nodes[nodeIdx].currConnID if err == net.ErrClosed || nodes[nodeIdx].lastErrs[0] == err || nodes[nodeIdx].lastErrs[1] == err || badConnID { diff --git a/internet/endpoint-arp.go b/internet/endpoint-arp.go deleted file mode 100644 index a948bcc..0000000 --- a/internet/endpoint-arp.go +++ /dev/null @@ -1,4 +0,0 @@ -package internet - -type ARPEndpoint struct { -} diff --git a/internet/node-arp.go b/internet/node-arp.go new file mode 100644 index 0000000..1a3a184 --- /dev/null +++ b/internet/node-arp.go @@ -0,0 +1,51 @@ +package internet + +import ( + "log/slog" + + "github.com/soypat/lneto" + "github.com/soypat/lneto/arp" + "github.com/soypat/lneto/ethernet" +) + +type NodeARP struct { + handler arp.Handler + vld lneto.Validator +} + +func (narp *NodeARP) Reset(cfg arp.HandlerConfig) error { + return narp.handler.Reset(cfg) +} + +func (narp *NodeARP) LocalPort() uint16 { return 0 } + +func (narp *NodeARP) Protocol() uint64 { return uint64(ethernet.TypeARP) } + +func (narp *NodeARP) ConnectionID() *uint64 { return narp.handler.ConnectionID() } + +func (narp *NodeARP) Demux(EtherFrame []byte, arpOff int) error { + afrm, err := arp.NewFrame(EtherFrame[arpOff:]) + if err != nil { + slog.Error("bad-ARP", slog.String("err", err.Error())) + return nil + } + afrm.ValidateSize(&narp.vld) + if narp.vld.HasError() { + slog.Error("invalid-ARP", slog.String("err", narp.vld.Err().Error())) + return nil + } + return narp.handler.Recv(EtherFrame[arpOff:]) +} + +func (narp *NodeARP) Encapsulate(EtherFrame []byte, arpOff int) (int, error) { + n, err := narp.handler.Send(EtherFrame[arpOff:]) + if err != nil || n == 0 { + return 0, err // end with error. + } + afrm, _ := arp.NewFrame(EtherFrame[arpOff:]) + hwaddr, _ := afrm.Target() + efrm, _ := ethernet.NewFrame(EtherFrame) + copy(efrm.DestinationHardwareAddr()[:], hwaddr) + slog.Info("handle", slog.String("out", afrm.String())) + return n, err +} diff --git a/internet/stack-ip.go b/internet/stack-ip.go index c0adb3f..39fbcfc 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -152,18 +152,14 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) } func (sb *StackIP) Register(h StackNode) error { - port := h.LocalPort() proto := h.Protocol() - if port <= 0 { - return errZeroPort - } else if proto > 255 { + if proto > 255 { return errInvalidProto } sb.handlers = append(sb.handlers, node{ demux: h.Demux, encapsulate: h.Encapsulate, proto: uint16(proto), - port: port, }) return nil } @@ -173,8 +169,8 @@ func (sb *StackIP) RegisterTCPConn(conn *TCPConn) error { return errZeroPort } sb.handlers = append(sb.handlers, node{ - demux: conn.RecvIP, - encapsulate: conn.HandleIP, + demux: conn.Demux, + encapsulate: conn.Encapsulate, proto: uint16(lneto.IPProtoTCP), port: conn.LocalPort(), }) diff --git a/internet/stack-linklayer.go b/internet/stack-linklayer.go index 489e2eb..d72be6f 100644 --- a/internet/stack-linklayer.go +++ b/internet/stack-linklayer.go @@ -35,6 +35,8 @@ func (ls *StackLinkLayer) Reset6(mac, gateway [6]byte, mtu int) error { return nil } +func (ls *StackLinkLayer) MTU() int { return int(ls.mtu) } + func (ls *StackLinkLayer) ConnectionID() *uint64 { return &ls.connID } func (ls *StackLinkLayer) LocalPort() uint16 { return 0 } diff --git a/internet/stack-port.go b/internet/stack-port.go index 0a7ed92..cff9811 100644 --- a/internet/stack-port.go +++ b/internet/stack-port.go @@ -5,28 +5,28 @@ import ( "io" ) -type StackPort struct { +type StackPorts struct { connID uint64 protocol uint64 handlers []node dstPortOff int } -func (ps *StackPort) Reset(protocol uint64, dstPortOffset int) { - *ps = StackPort{ +func (ps *StackPorts) Reset(protocol uint64, dstPortOffset int) { + *ps = StackPorts{ connID: ps.connID + 1, handlers: ps.handlers[:0], dstPortOff: dstPortOffset, protocol: protocol, } } -func (ps *StackPort) LocalPort() uint16 { return 0 } +func (ps *StackPorts) LocalPort() uint16 { return 0 } -func (ps *StackPort) Protocol() uint64 { return ps.protocol } +func (ps *StackPorts) Protocol() uint64 { return ps.protocol } -func (ps *StackPort) ConnectionID() *uint64 { return &ps.connID } +func (ps *StackPorts) ConnectionID() *uint64 { return &ps.connID } -func (ps *StackPort) Encapsulate(b []byte, offset int) (n int, err error) { +func (ps *StackPorts) Encapsulate(b []byte, offset int) (n int, err error) { if ps.dstPortOff+offset+2 > len(b) { return 0, io.ErrShortBuffer } @@ -41,7 +41,7 @@ func (ps *StackPort) Encapsulate(b []byte, offset int) (n int, err error) { return n, err } -func (ps *StackPort) Demux(b []byte, offset int) (err error) { +func (ps *StackPorts) Demux(b []byte, offset int) (err error) { if ps.dstPortOff+offset+2 > len(b) { return io.ErrShortBuffer } @@ -60,7 +60,7 @@ func (ps *StackPort) Demux(b []byte, offset int) (err error) { return err } -func (ps *StackPort) Register(h StackNode) error { +func (ps *StackPorts) Register(h StackNode) error { port := h.LocalPort() proto := h.Protocol() if port <= 0 { @@ -76,6 +76,6 @@ func (ps *StackPort) Register(h StackNode) error { return nil } -func (ps *StackPort) handleResult(handlerIdx, n int, err error) { +func (ps *StackPorts) handleResult(handlerIdx, n int, err error) { handleNodeError(&ps.handlers, handlerIdx, err) } diff --git a/internet/tcpconn.go b/internet/tcpconn.go index 27499bd..fd93b16 100644 --- a/internet/tcpconn.go +++ b/internet/tcpconn.go @@ -98,7 +98,7 @@ func (conn *TCPConn) Close() error { return conn.h.Close() } -func (conn *TCPConn) RecvIP(buf []byte, off int) (err error) { +func (conn *TCPConn) Demux(buf []byte, off int) (err error) { conn.trace("tcpconn.Recv:start") if off >= len(buf) { return errors.New("bad offset in TCPConn.Recv") @@ -198,7 +198,7 @@ func (conn *TCPConn) checkPipeOpen() error { return nil } -func (conn *TCPConn) HandleIP(buf []byte, off int) (n int, err error) { +func (conn *TCPConn) Encapsulate(buf []byte, off int) (n int, err error) { if len(conn.remoteAddr) == 0 { return 0, errors.New("unset IP address") }