diff --git a/examples/xnet/main.go b/examples/xnet/main.go index 06ac571..29f7e40 100644 --- a/examples/xnet/main.go +++ b/examples/xnet/main.go @@ -19,6 +19,7 @@ import ( "github.com/soypat/lneto/internal" "github.com/soypat/lneto/internal/ltesto" "github.com/soypat/lneto/internet/pcap" + "github.com/soypat/lneto/tcp" "github.com/soypat/lneto/x/xnet" ) @@ -97,8 +98,6 @@ func run() (err error) { HardwareAddress: brHW, MTU: uint16(mtu), MaxTCPConns: 1, - TCPBufferSizeTx: 2048, - TCPBufferSizeRx: 2048, }) if err != nil { return err @@ -199,6 +198,12 @@ func run() (err error) { return fmt.Errorf("DNS of host %q failed: %w", flagHostToResolve, err) } fmt.Printf("DNS resolution of %q complete and resolved to %v\n", flagHostToResolve, addrs) + var conn tcp.Conn + conn.Configure(tcp.ConnConfig{ + RxBuf: make([]byte, mtu), + TxBuf: make([]byte, mtu), + TxPacketQueueSize: 3, + }) if flagHTTPGet { var hdr httpraw.Header hdr.SetMethod("GET") @@ -213,7 +218,7 @@ func run() (err error) { } const tcpDebugTimeout = 60 * time.Minute target := netip.AddrPortFrom(addrs[0], 80) - conn, err := rstack.DoDialTCP(uint16(softRand&0xefff)+1024, target, tcpDebugTimeout, internetRetries) + err = rstack.DoDialTCP(&conn, uint16(softRand&0xefff)+1024, target, tcpDebugTimeout, internetRetries) if err != nil { return fmt.Errorf("TCP failed: %w", err) } diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index f52d94e..197ef80 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -25,6 +25,10 @@ func (ls *StackEthernet) SetGateway6(gw [6]byte) { ls.gwmac = gw } +func (ls *StackEthernet) Gateway6() (gw [6]byte) { + return ls.gwmac +} + func (ls *StackEthernet) SetHardwareAddr6(mac [6]byte) { ls.mac = mac } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 946abd8..7adfee3 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -210,26 +210,42 @@ func (s *StackAsync) SetHardwareAddress(hw [6]byte) error { return s.resetARP() } +func (s *StackAsync) HardwareAddress() (hw [6]byte) { + s.mu.Lock() + defer s.mu.Unlock() + return s.link.HardwareAddr6() +} + func (s *StackAsync) SetGateway6(gwhw [6]byte) { s.mu.Lock() defer s.mu.Unlock() s.link.SetGateway6(gwhw) } -func (s *StackAsync) start() { - +func (s *StackAsync) Gateway6() [6]byte { + s.mu.Lock() + defer s.mu.Unlock() + return s.link.Gateway6() } func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrPort) (err error) { - if !conn.State().IsClosed() { - return errors.New("conn not closed") - } - conn.Abort() // Conn is closed, safe to abort. err = conn.OpenActive(localPort, addrp, tcp.Value(s.Prand32())) + if err != nil { + return err + } + err = s.tcps.Register(conn) if err != nil { conn.Abort() return err } + return nil +} + +func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) { + err = conn.OpenListen(localPort, tcp.Value(s.Prand32())) + if err != nil { + return err + } err = s.tcps.Register(conn) if err != nil { conn.Abort() diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index 84b9efa..1268649 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -4,9 +4,15 @@ import ( "net/netip" "testing" + "github.com/soypat/lneto" + "github.com/soypat/lneto/internet/pcap" "github.com/soypat/lneto/tcp" ) +const ( + synack = tcp.FlagSYN | tcp.FlagACK +) + func TestABC(t *testing.T) { const seed = 1234 const MTU = 1500 @@ -37,23 +43,8 @@ func TestABC(t *testing.T) { if err != nil { t.Fatal(err) } - // IMG_1084.MOV const svPort = 80 - var clconn tcp.Conn - err = clconn.Configure(tcp.ConnConfig{ - RxBuf: make([]byte, MTU), - TxBuf: make([]byte, MTU), - TxPacketQueueSize: 4, - }) - if err != nil { - t.Fatal(err) - } - err = client.DialTCP(&clconn, 1337, netip.AddrPortFrom(sv.Addr(), svPort)) - if err != nil { - t.Fatal(err) - } - var svconn tcp.Conn err = svconn.Configure(tcp.ConnConfig{ RxBuf: make([]byte, MTU), @@ -63,5 +54,92 @@ func TestABC(t *testing.T) { if err != nil { t.Fatal(err) } - // clconn.OpenListen() + err = sv.ListenTCP(&svconn, svPort) + if err != nil { + t.Fatal(err) + } + + var clconn tcp.Conn + err = clconn.Configure(tcp.ConnConfig{ + RxBuf: make([]byte, MTU), + TxBuf: make([]byte, MTU), + TxPacketQueueSize: 4, + }) + if err != nil { + t.Fatal(err) + } + client.SetGateway6(sv.HardwareAddress()) + sv.SetGateway6(client.HardwareAddress()) + + err = client.DialTCP(&clconn, 1337, netip.AddrPortFrom(sv.Addr(), svPort)) + if err != nil { + t.Fatal(err) + } + + var expected = []struct { + fromClient bool + flags tcp.Flags + }{ + { + fromClient: true, + flags: tcp.FlagSYN, + }, + { + fromClient: false, + flags: synack, + }, + { + fromClient: true, + flags: tcp.FlagACK, + }, + } + var cap pcap.PacketBreakdown + var frms []pcap.Frame + var buf [MTU]byte + for _, action := range expected { + var n int + switch action.fromClient { + case true: + n, err = client.Encapsulate(buf[:], 0) + case false: + n, err = sv.Encapsulate(buf[:], 0) + } + if err != nil { + t.Fatal(err) + } else if n == 0 { + t.Error("zero bits sent") + } + frms, err = cap.CaptureEthernet(frms[:0], buf[:n], 0) + if err != nil { + t.Fatal(err) + } + tfrm := getProtoFrame(frms, lneto.IPProtoTCP) + if tfrm == nil { + t.Fatal("where's the TCP?") + } + fidx, _ := tfrm.FieldByClass(pcap.FieldClassFlags) + flags, _ := tfrm.FieldAsUint(fidx, buf[:n]) + tflags := tcp.Flags(flags) + if tflags != action.flags { + t.Errorf("expected flags %s, got %s", action.flags.String(), tflags.String()) + } + switch action.fromClient { + case true: + err = sv.Demux(buf[:], 0) + case false: + err = client.Demux(buf[:], 0) + } + if err != nil { + t.Fatal(err) + } + } +} + +func getProtoFrame(frms []pcap.Frame, proto any) *pcap.Frame { + for i := range frms { + if frms[i].Protocol == proto { + return &frms[i] + } + } + return nil }