diff --git a/examples/stack/main.go b/examples/stack/main.go index 65ad8a4..e953c1b 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -1,17 +1,13 @@ package main import ( - "bytes" - "encoding/json" "errors" "fmt" "io" "log" "log/slog" "net" - "net/http" "net/netip" - "net/url" "os" "time" @@ -19,6 +15,7 @@ import ( "github.com/soypat/lneto/arp" "github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/internal" + "github.com/soypat/lneto/internal/ltesto" "github.com/soypat/lneto/ipv4" "github.com/soypat/lneto/ipv6" "github.com/soypat/lneto/tcp" @@ -57,7 +54,7 @@ func main() { log.Fatal(err) } - tap := NewHTTPTap("http://127.0.0.1:7070") + tap := ltesto.NewHTTPTapClient("http://127.0.0.1:7070") defer tap.Close() fmt.Println("hosting server at ", addrPort.String()) @@ -520,65 +517,6 @@ func addHandler(handlers []handler, h Handler, remoteAddr []byte, lport uint16) return handlers } -func NewHTTPTap(baseURL string) *HTTPTap { - var h HTTPTap - h.sendurl = baseURL + "/send" - h.recvurl = baseURL + "/recv" - _, err := url.Parse(h.sendurl) - if err != nil { - panic(err) - } - var data [2048]byte - var n int = -1 - for n != 0 { - n, _ = h.Read(data[:]) // Empty remote data. - } - return &h -} - -type TAPNop struct{} - -func (h *TAPNop) Read(b []byte) (int, error) { return 0, nil } -func (h *TAPNop) Write(b []byte) (int, error) { return 0, nil } -func (h *TAPNop) Close() error { return nil } - -type HTTPTap struct { - c http.Client - recvurl string - sendurl string -} - -func (h *HTTPTap) Read(b []byte) (int, error) { - resp, err := h.c.Get(h.recvurl) - if err != nil { - return 0, err - } else if resp.StatusCode != 200 { - return 0, errors.New(resp.Status + " for " + h.recvurl) - } - var data []byte - err = json.NewDecoder(resp.Body).Decode(&data) - if err != nil { - return 0, err - } else if len(b) < len(data) { - return 0, fmt.Errorf("got too large packet %d for buffer %d", len(data), len(b)) - } - copy(b, data) - return len(data), nil -} - -func (h *HTTPTap) Write(b []byte) (int, error) { - data, _ := json.Marshal(b) - resp, err := h.c.Post(h.sendurl, "application/json", bytes.NewReader(data)) - if err != nil { - return 0, err - } else if resp.StatusCode != 200 { - return 0, errors.New(resp.Status + " for " + h.sendurl) - } - return len(b), nil -} - -func (h *HTTPTap) Close() error { return nil } - func tcpChecksum(ipFrame []byte, tcpPayload int) uint16 { version := ipFrame[0] >> 4 var tfrm tcp.Frame diff --git a/examples/stackbasic/main.go b/examples/stackbasic/main.go new file mode 100644 index 0000000..7b79f4d --- /dev/null +++ b/examples/stackbasic/main.go @@ -0,0 +1,278 @@ +package main + +import ( + "errors" + "fmt" + "io" + "log" + "log/slog" + "net" + "net/netip" + "os" + "time" + + "github.com/soypat/lneto" + "github.com/soypat/lneto/arp" + "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" + "github.com/soypat/lneto/internal/ltesto" + "github.com/soypat/lneto/internet" +) + +const ( + mtu = 2048 + iface = "192.168.10.1/24" + stackIP = "192.168.10.2" + stackPort = 80 + iss = 100 +) + +var stackHWAddr = [6]byte{0xc0, 0xff, 0xee, 0x00, 0xde, 0xad} + +func main() { + ip := netip.MustParseAddr(stackIP) + iface := netip.MustParsePrefix(iface) + if !iface.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{ + Level: slog.LevelDebug, + })) + handler.SetLoggers(logger, logger) + + err = handler.OpenListen(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 + for { + nread, err := tap.Read(buf[:]) + if err != nil { + slogger.error("tap-err", slog.String("err", err.Error())) + log.Fatal(err) + } else if nread > 0 { + err = lStack.RecvEth(buf[:nread]) + if err != nil { + slogger.error("recv", slog.String("err", err.Error()), slog.Int("plen", nread)) + } else { + slogger.info("recv", slog.Int("plen", nread)) + } + } + nw, err := lStack.HandleEth(buf[:]) + if err != nil { + slogger.error("handle", slog.String("err", err.Error())) + } else if nw > 0 { + _, 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) + } + } +} + +func NewEthernetTCPStack(mac [6]byte, ip netip.AddrPort, slogger logger) (*LinkStack, *internet.TCPConn, error) { + var err error + lStack := LinkStack{ + logger: slogger, + mac: mac, + mtu: mtu, + } + + var ipStack internet.StackBasic + addr := ip.Addr() + addr4 := addr.As4() + _ = addr4 + ipStack.SetAddr(addr) + lStack.Register(handler{ + raddr: nil, //addr4[:], + recv: func(b []byte, i int) error { + return ipStack.Recv(b[i:]) + }, + handle: func(b []byte, i int) (int, error) { + return ipStack.Handle(b[i:]) + }, + proto: uint32(lneto.IPProtoIPv4), + lport: 0, + }) + var conn internet.TCPConn + err = conn.Configure(&internet.TCPConnConfig{ + RxBuf: make([]byte, mtu), + TxBuf: make([]byte, mtu), + TxPacketQueueSize: 3, + Logger: slog.Default(), + }) + if err != nil { + return nil, nil, err + } + err = conn.OpenListen(ip.Port(), 100) + if err != nil { + return nil, nil, err + } + err = ipStack.RegisterTCPConn(&conn) + if err != nil { + return nil, nil, err + } + proto := ethernet.TypeIPv4 + if ip.Addr().Is6() { + proto = ethernet.TypeIPv6 + } + arphandler, err := arp.NewHandler(arp.HandlerConfig{ + HardwareAddr: mac[:], + ProtocolAddr: ip.Addr().AsSlice(), + MaxQueries: 1, + MaxPending: 1, + HardwareType: 1, + ProtocolType: proto, + }) + if err != nil { + return nil, nil, err + } + arpStack := ARPStack{ + handler: *arphandler, + } + + err = lStack.Register(ipStack, mac) + if err != nil { + return nil, nil, err + } + return &lStack, &conn, nil +} + +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) error { + proto := h.proto + for i := range ls.handlers { + if proto == ls.handlers[i].proto { + return errors.New("protocol already registered") + } + } + ls.handlers = append(ls.handlers, h) + return nil +} + +func (ls *LinkStack) RecvEth(ethFrame []byte) (err error) { + + efrm, err := ethernet.NewFrame(ethFrame) + if err != nil { + return 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()) + } + var vld lneto.Validator + efrm.ValidateSize(&vld) + if err := vld.Err(); err != nil { + return err + } + + for i := range ls.handlers { + h := &ls.handlers[i] + if h.proto == uint32(etype) { + return h.recv(efrm.Payload(), 0) + } + } + + return nil +} + +func (ls *LinkStack) HandleEth(dst []byte) (n int, err error) { + if len(dst) < int(ls.mtu) { + return 0, io.ErrShortBuffer + } + 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())) + 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 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 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 (l logger) trace(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, internal.LevelTrace, msg, attrs...) +} diff --git a/examples/tap/main.go b/examples/tap/main.go index 2be079a..3b36372 100644 --- a/examples/tap/main.go +++ b/examples/tap/main.go @@ -1,15 +1,15 @@ package main import ( - "encoding/json" "errors" "fmt" "log" "log/slog" "net/http" "net/netip" + "time" - "github.com/soypat/lneto/internal" + "github.com/soypat/lneto/internal/ltesto" ) func main() { @@ -23,102 +23,32 @@ func main() { func run() error { var ( - flagNet = "192.168.10.1/24" - flagiface = "tap0" - flagMTU = 1500 + flagNet = "192.168.10.1/24" + flagiface = "tap0" + flagMTU = 1500 + flagPacketQueueSize = 2048 ) - slogger := slog.Default() ip, err := netip.ParsePrefix(flagNet) if err != nil { return err } - tap, err := internal.NewTap(flagiface, ip) + sv, err := ltesto.NewHTTPTapServer(flagiface, ip, flagMTU, flagPacketQueueSize, flagPacketQueueSize) if err != nil { return err } - defer tap.Close() - s := stack{ - out: make(chan []byte, 256), - in: make(chan []byte, 2048), - } - sv := http.NewServeMux() - sv.HandleFunc("/send", func(w http.ResponseWriter, r *http.Request) { - var data []byte - err := json.NewDecoder(r.Body).Decode(&data) - if err != nil { - http.Error(w, err.Error(), http.StatusBadRequest) - } else { - select { - case s.out <- data: - slog.Info("http-send", slog.Int("plen", len(data))) - default: - http.Error(w, "outgoing packet queue full", http.StatusInternalServerError) - } - } - }) - sv.HandleFunc("/recv", func(w http.ResponseWriter, r *http.Request) { - select { - case data := <-s.in: - json.NewEncoder(w).Encode(data) - slog.Info("http-recv", slog.Int("plen", len(data))) - default: - json.NewEncoder(w).Encode("") // send empty string. - } - }) + defer sv.Close() fmt.Println("listening on http://127.0.0.1:7070/recv and http://127.0.0.1:7070/send") go http.ListenAndServe(":7070", sv) - buf := make([]byte, flagMTU) for { - n, err := tap.Read(buf[:]) + result, err := sv.HandleTap() if err != nil { - log.Fatal(err) - } else if n > 0 { - err = s.recv(buf[:n]) - if err != nil { - slogger.Error("recv", slog.String("err", err.Error()), slog.Int("plen", n)) - } else { - slogger.Info("recv", slog.Int("plen", n)) - } + slog.Error("handletap:error", slog.String("err", err.Error()), slog.Any("result", result)) } - n, err = s.handle(buf[:]) - if err != nil { - slogger.Error("handle", slog.String("err", err.Error())) - } else if n > 0 { - _, err = tap.Write(buf[:n]) - if err != nil { - log.Fatal(err) - } else { - slogger.Info("write", slog.Int("plen", n)) - } + if result.Failed { + return errors.New("tap failed, exit program") + } else if result.ReceivedSize == 0 && result.SentSize == 0 { + time.Sleep(200 * time.Millisecond) // No data exchanged, sleep a bit to not hog CPU. } } } - -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/ltesto/httptap.go b/internal/ltesto/httptap.go new file mode 100644 index 0000000..90d627e --- /dev/null +++ b/internal/ltesto/httptap.go @@ -0,0 +1,211 @@ +package ltesto + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net/http" + "net/netip" + "net/url" + + "github.com/soypat/lneto/internal" +) + +// NewHTTPTapClient returns a HTTPTapClient ready for use. +func NewHTTPTapClient(baseURL string) *HTTPTapClient { + var h HTTPTapClient + h.sendurl = baseURL + "/send" + h.recvurl = baseURL + "/recv" + _, err := url.Parse(h.sendurl) + if err != nil { + panic(err) + } + return &h +} + +type HTTPTapClient struct { + c http.Client + recvurl string + sendurl string +} + +func (h *HTTPTapClient) ReadDiscard() { + var data [2048]byte + var n int = -1 + for n != 0 { + n, _ = h.Read(data[:]) // Empty remote data. + } +} + +func (h *HTTPTapClient) Read(b []byte) (int, error) { + resp, err := h.c.Get(h.recvurl) + if err != nil { + return 0, err + } else if resp.StatusCode != 200 { + return 0, errors.New(resp.Status + " for " + h.recvurl) + } + var data []byte + err = json.NewDecoder(resp.Body).Decode(&data) + if err != nil { + return 0, err + } else if len(b) < len(data) { + return 0, fmt.Errorf("got too large packet %d for buffer %d", len(data), len(b)) + } + copy(b, data) + return len(data), nil +} + +func (h *HTTPTapClient) Write(b []byte) (int, error) { + data, _ := json.Marshal(b) + resp, err := h.c.Post(h.sendurl, "application/json", bytes.NewReader(data)) + if err != nil { + return 0, err + } else if resp.StatusCode != 200 { + return 0, errors.New(resp.Status + " for " + h.sendurl) + } + return len(b), nil +} + +func (h *HTTPTapClient) Close() error { return nil } + +type HTTPTapServer struct { + router *http.ServeMux + stack stack + tap *internal.Tap + buf []byte + tapfailed bool +} + +func NewHTTPTapServer(iface string, ip netip.Prefix, mtu, queueOut, queueIn int) (*HTTPTapServer, error) { + tap, err := internal.NewTap(iface, ip) + if err != nil { + return nil, err + } + + s := stack{ + out: make(chan []byte, queueOut), + in: make(chan []byte, queueIn), + } + sv := http.NewServeMux() + sv.HandleFunc("/send", func(w http.ResponseWriter, r *http.Request) { + var data []byte + err := json.NewDecoder(r.Body).Decode(&data) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + } else { + select { + case s.out <- data: + slog.Info("http-send", slog.Int("plen", len(data))) + default: + http.Error(w, "outgoing packet queue full", http.StatusInternalServerError) + } + } + }) + sv.HandleFunc("/recv", func(w http.ResponseWriter, r *http.Request) { + select { + case data := <-s.in: + json.NewEncoder(w).Encode(data) + slog.Info("http-recv", slog.Int("plen", len(data))) + default: + json.NewEncoder(w).Encode("") // send empty string. + } + }) + taps := HTTPTapServer{ + router: sv, + stack: s, + tap: tap, + buf: make([]byte, mtu), + } + return &taps, nil +} + +func (sv *HTTPTapServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { + sv.router.ServeHTTP(w, r) +} + +func (sv *HTTPTapServer) Close() error { + return sv.tap.Close() +} + +type HandleTapResult struct { + Failed bool + 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 { + 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/internet/basicstack.go b/internet/basicstack.go index e7afe3e..25d3eb3 100644 --- a/internet/basicstack.go +++ b/internet/basicstack.go @@ -4,6 +4,7 @@ import ( "errors" "io" "log/slog" + "net/netip" "github.com/soypat/lneto" "github.com/soypat/lneto/internal" @@ -19,35 +20,42 @@ type StackBasic struct { } type handler struct { - proto lneto.IPProto - port uint16 recv func([]byte, int) error handle func([]byte, int) (int, error) + proto lneto.IPProto + port uint16 } -func (is *StackBasic) Recv(frame []byte) error { +func (sb *StackBasic) SetAddr(addr netip.Addr) { + if !addr.Is4() { + panic("only support IPv4") + } + sb.ip = addr.As4() +} + +func (sb *StackBasic) Recv(frame []byte) error { ifrm, err := ipv4.NewFrame(frame) if err != nil { return err } - if *ifrm.DestinationAddr() != is.ip { + if *ifrm.DestinationAddr() != sb.ip { return errors.New("packet not for us") } - is.validator.ResetErr() - ifrm.ValidateExceptCRC(&is.validator) - if err = is.validator.Err(); err != nil { + sb.validator.ResetErr() + ifrm.ValidateExceptCRC(&sb.validator) + if err = sb.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))) + sb.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] + for i := range sb.handlers { + h := &sb.handlers[i] if h.proto == ifrm.Protocol() { return h.recv(frame[:totalLen], off) } @@ -55,7 +63,7 @@ func (is *StackBasic) Recv(frame []byte) error { return nil } -func (is *StackBasic) Handle(frame []byte) (int, error) { +func (sb *StackBasic) Handle(frame []byte) (int, error) { if len(frame) < 256 { return 0, io.ErrShortBuffer } @@ -63,38 +71,44 @@ func (is *StackBasic) Handle(frame []byte) (int, error) { const ihl = 5 const headerlen = ihl * 4 ifrm.SetVersionAndIHL(4, 5) - *ifrm.SourceAddr() = is.ip + *ifrm.SourceAddr() = sb.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") - } - + ifrm.SetID(0) + for i := range sb.handlers { + h := &sb.handlers[i] n, err := h.handle(frame[:], headerlen) if err != nil { - is.error("IPv4Stack:handle", slog.String("proto", proto.String()), slog.String("err", err.Error())) + sb.error("IPv4Stack:handle", slog.String("proto", h.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()) if ifrm.Protocol() == lneto.IPProtoTCP { tfrm, _ := tcp.NewFrame(ifrm.Payload()) - is.info("IPv4Stack:send", slog.String("ip", ifrm.String()), slog.String("tcp", tfrm.String())) + sb.info("IPv4Stack:send", slog.String("ip", ifrm.String()), slog.String("tcp", tfrm.String())) } return totalLen, nil } } + return 0, nil +} + +func (sb *StackBasic) RegisterTCPConn(conn *TCPConn) error { + if conn.LocalPort() == 0 { + return errors.New("undefined local port") + } + sb.handlers = append(sb.handlers, handler{ + recv: conn.RecvIP, + handle: conn.HandleIP, + proto: lneto.IPProtoIPv4, + port: conn.LocalPort(), + }) + return nil } type logger struct { @@ -113,3 +127,6 @@ func (l logger) warn(msg string, attrs ...slog.Attr) { func (l logger) debug(msg string, attrs ...slog.Attr) { internal.LogAttrs(l.log, slog.LevelDebug, msg, attrs...) } +func (l logger) trace(msg string, attrs ...slog.Attr) { + internal.LogAttrs(l.log, internal.LevelTrace, msg, attrs...) +} diff --git a/internet/tcpconn.go b/internet/tcpconn.go new file mode 100644 index 0000000..4344098 --- /dev/null +++ b/internet/tcpconn.go @@ -0,0 +1,158 @@ +package internet + +import ( + "bytes" + "errors" + "log/slog" + "net/netip" + "time" + + "github.com/soypat/lneto/ipv4" + "github.com/soypat/lneto/ipv6" + "github.com/soypat/lneto/tcp" +) + +type TCPConn struct { + h tcp.Handler + remoteAddr []byte + logger + + rdead time.Time + wdead time.Time + lastTx time.Time + lastRx time.Time +} +type TCPConnConfig struct { + RxBuf []byte + TxBuf []byte + TxPacketQueueSize int + Logger *slog.Logger +} + +func (conn *TCPConn) Configure(config *TCPConnConfig) (err error) { + err = conn.h.SetBuffers(config.TxBuf, config.RxBuf, config.TxPacketQueueSize) + if err != nil { + return err + } + conn.logger.log = config.Logger + return nil +} + +// LocalPort returns the local port on which the socket is listening or connected to. +func (conn *TCPConn) LocalPort() uint16 { return conn.h.LocalPort() } + +// RemotePort returns the port of the incoming remote connection. Is non-zero if connection is established. +func (conn *TCPConn) RemotePort() uint16 { return conn.h.RemotePort() } + +// State returns the TCP state of the socket. +func (conn *TCPConn) State() tcp.State { return conn.h.State() } + +// BufferedInput returns the number of bytes in the socket's receive/input buffer. +func (conn *TCPConn) BufferedInput() int { return conn.h.BufferedInput() } + +// OpenActive opens a connection to a remote peer with a known IP address and port combination. +// iss is the initial send sequence number which is ideally a random number which is far away from the last sequence number used on a connection to the same host. +func (conn *TCPConn) OpenActive(remote netip.AddrPort, localPort uint16, iss tcp.Value) error { + err := conn.h.OpenActive(localPort, remote.Port(), iss) + if err != nil { + return err + } + conn.reset(conn.h) + raddr := remote.Addr() + if raddr.Is4() { + addr4 := raddr.As4() + conn.remoteAddr = append(conn.remoteAddr[:0], addr4[:]...) + } else if raddr.Is6() { + addr6 := raddr.As16() + conn.remoteAddr = append(conn.remoteAddr[:0], addr6[:]...) + } + return nil +} + +// OpenListen opens a passive connection which listens for the first SYN packet to be received on a local port. +// iss is the initial send sequence number which is usually a randomly chosen number. +func (conn *TCPConn) OpenListen(localPort uint16, iss tcp.Value) error { + err := conn.h.OpenListen(localPort, iss) + if err != nil { + return err + } + conn.reset(conn.h) + return nil +} + +func (conn *TCPConn) RecvIP(buf []byte, off int) (err error) { + conn.trace("tcpconn.Recv:start") + if off >= len(buf) { + return errors.New("bad offset in TCPConn.Recv") + } + raddr, err := getIPAddr(buf[:off]) + if err != nil { + return err + } + if conn.isRaddrSet() && !bytes.Equal(conn.remoteAddr, raddr) { + return errors.New("IP addr mismatch on TCPConn") + } + err = conn.h.Recv(buf[off:]) + if err != nil { + return err + } + if !conn.isRaddrSet() && conn.h.RemotePort() != 0 { + conn.remoteAddr = append(conn.remoteAddr[:0], raddr...) + } + return nil +} + +func (conn *TCPConn) HandleIP(buf []byte, off int) (n int, err error) { + if len(conn.remoteAddr) == 0 { + return 0, errors.New("unset IP address") + } + raddr, err := getIPAddr(buf[:off]) + if err != nil { + return 0, err + } else if len(raddr) != len(conn.remoteAddr) { + return 0, errors.New("mismatched IP version") + } + n, err = conn.h.Send(buf[off:]) + if err != nil { + return 0, err + } + copy(raddr, conn.remoteAddr) + return n, nil +} + +func (conn *TCPConn) Send(response []byte) (n int, err error) { + conn.trace("tcpconn.Send:start") + return conn.h.Send(response) +} + +func getIPAddr(buf []byte) (addr []byte, err error) { + switch buf[0] >> 4 { + case 4: + ifrm4, err := ipv4.NewFrame(buf) + if err != nil { + return addr, err + } + addr = ifrm4.SourceAddr()[:] + case 6: + ifrm6, err := ipv6.NewFrame(buf) + if err != nil { + return addr, err + } + addr = ifrm6.SourceAddr()[:] + default: + err = errors.New("unsupported IP version") + } + return addr, err +} + +func (conn *TCPConn) isRaddrSet() bool { + return len(conn.remoteAddr) != 0 +} + +func (conn *TCPConn) reset(h tcp.Handler) { + *conn = TCPConn{ + h: h, + remoteAddr: conn.remoteAddr[:0], + logger: conn.logger, + } +} diff --git a/tcp/definitions.go b/tcp/definitions.go index 86a75b1..c98bf66 100644 --- a/tcp/definitions.go +++ b/tcp/definitions.go @@ -16,7 +16,7 @@ var ( errDropSegment = errors.New("drop segment") errWindowTooLarge = errors.New("invalid window size > 2**16") - errBufferTooSmall = errors.New("buffer too small") + errBufferTooSmall = errors.New("tcp buffer too small") errNeedClosedTCBToOpen = errors.New("need closed TCB to call open") errInvalidState = errors.New("invalid state") errConnNotExist = errors.New("connection does not exist") diff --git a/tcp/handler.go b/tcp/handler.go index a00a3cd..1988b69 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -73,12 +73,14 @@ func (h *Handler) OpenActive(localPort, remotePort uint16, iss Value) error { } if h.scb.State() != StateClosed { return errNeedClosedTCBToOpen + } else if remotePort == 0 { + return errors.New("zero port on open call") } h.reset(localPort, remotePort, iss) return nil } -// OpenListen prepares a passive connection. Set remote port to non-zero to +// OpenListen prepares a passive connection. func (h *Handler) OpenListen(localPort uint16, iss Value) error { if h.bufRx.Size() < minBufferSize || h.bufTx.Size() < minBufferSize { return errBufferTooSmall @@ -230,7 +232,8 @@ func (h *Handler) Read(b []byte) (int, error) { return h.bufRx.Read(b) } -func (h *Handler) Buffered() int { +// BufferedInput returns amount of bytes buffered in receive buffer. +func (h *Handler) BufferedInput() int { if h.State().IsClosed() { return 0 } diff --git a/tcp/handler_test.go b/tcp/handler_test.go index c01e1f9..2dc6875 100644 --- a/tcp/handler_test.go +++ b/tcp/handler_test.go @@ -33,8 +33,8 @@ func sendDataFull(t *testing.T, client, server *Handler, data, packetBuf []byte) err = server.Recv(packetBuf[:n]) if err != nil { t.Fatal("server receiving:", err) - } else if server.Buffered() != len(data) { - t.Fatal("server did not receive full data packet", server.Buffered(), len(data)) + } else if server.BufferedInput() != len(data) { + t.Fatal("server did not receive full data packet", server.BufferedInput(), len(data)) } clear(packetBuf) n, err = server.Read(packetBuf)