diff --git a/dhcpv4/server.go b/dhcpv4/server.go index a0e0b43..c8609fe 100644 --- a/dhcpv4/server.go +++ b/dhcpv4/server.go @@ -203,7 +203,10 @@ func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) copy(dfrm.CHAddrAs6()[:], client.hwaddr[:]) dfrm.SetMagicCookie(MagicCookie) if carrierIsIP { - internal.SetIPDestinationAddr(carrierData, 0, client.addr[:]) + err = internal.SetIPAddrs(carrierData, 0, sv.siaddr[:], client.addr[:]) + if err != nil { + return 0, err + } } client.state = futureState @@ -227,13 +230,13 @@ func (sv *Server) getClientByIP(ip [4]byte) (serverEntry, [36]byte, bool) { return serverEntry{}, [36]byte{}, false } -func getSrcIPPort(ipCarrier []byte) (addr []byte, port uint16, err error) { - addr, _, off, err := internal.GetIPSourceAddr(ipCarrier) +func getSrcIPPort(ipCarrier []byte) (srcaddr []byte, port uint16, err error) { + srcaddr, _, _, off, err := internal.GetIPAddr(ipCarrier) if err != nil { - return addr, port, err + return srcaddr, port, err } else if len(ipCarrier[off:]) < 2 { - return addr, port, errors.New("getSrcIPPort got only IP layer") + return srcaddr, port, errors.New("getSrcIPPort got only IP layer") } port = binary.BigEndian.Uint16(ipCarrier[off:]) // TCP and UDP share same port offsets. - return addr, port, nil + return srcaddr, port, nil } diff --git a/dns/client.go b/dns/client.go new file mode 100644 index 0000000..aa96bfd --- /dev/null +++ b/dns/client.go @@ -0,0 +1,124 @@ +package dns + +import ( + "errors" + "math" + "net" + + "github.com/soypat/lneto" + "github.com/soypat/lneto/internal" +) + +type Client struct { + connID uint64 + txid uint16 + msg Message + respFlags HeaderFlags + state clientState + enableRecursion bool +} + +type ResolveConfig struct { + Questions []Question + EnableRecursion bool +} + +func (sudp *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) } + +func (sudp *Client) LocalPort() uint16 { return ClientPort } + +func (sudp *Client) ConnectionID() *uint64 { return &sudp.connID } + +func (c *Client) StartResolve(cfg ResolveConfig) error { + nd := len(cfg.Questions) + if nd > math.MaxUint16 { + return errors.New("overflow uint16 in DNS questions") + } + c.reset(internal.Prand16(c.txid^uint16(c.connID)), dnsSendQuery, cfg.EnableRecursion) + c.msg.LimitResourceDecoding(uint16(nd), uint16(nd), 0, 0) + c.msg.AddQuestions(cfg.Questions) + return nil +} + +func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { + if c.isClosed() { + return 0, net.ErrClosed + } else if c.state != dnsSendQuery { + return 0, nil + } + msg := &c.msg + frame := carrierData[frameOffset:] + msglen := msg.Len() + if msglen > uint16(len(frame)) { + return 0, errCalcLen + } + data, err := msg.AppendTo(frame, c.txid, NewClientHeaderFlags(OpCodeQuery, c.enableRecursion)) + if err != nil { + return 0, err + } else if len(data) > int(msglen) { + return 0, errors.New("unexpected write") + } + c.state = dnsAwaitResponse + return len(data), nil +} + +func (c *Client) Demux(carrierData []byte, frameOffset int) error { + if c.isClosed() { + return net.ErrClosed + } else if c.state != dnsAwaitResponse { + return nil + } + frame := carrierData[frameOffset:] + f, err := NewFrame(frame) + if err != nil { + return err + } + flags := f.Flags() + if f.TxID() != c.txid || !flags.IsResponse() { + return nil // Not meant for our client. + } + c.respFlags = flags + c.state = dnsDone + msg := &c.msg + _, incompleteButOK, err := msg.Decode(frame) + if err != nil && !incompleteButOK { + return err + } + return nil +} + +func (c *Client) isClosed() bool { + return c.state == dnsClosed || c.state == dnsAborted +} + +func (c *Client) Answers() []Resource { + if c.state != dnsDone { + return nil + } + return c.msg.Answers +} + +func (c *Client) Abort() { + c.reset(0, 0, false) +} + +func (c *Client) reset(txid uint16, state clientState, enableRecursion bool) { + *c = Client{ + connID: c.connID + 1, + txid: txid, + msg: c.msg, + state: state, + enableRecursion: enableRecursion, + } + c.msg.Reset() +} + +type clientState uint8 + +const ( + dnsClosed clientState = iota + dnsSendQuery + dnsAwaitResponse + dnsDone + dnsAborted +) diff --git a/dns/definitions.go b/dns/definitions.go index 2ca830a..ed8e6cf 100644 --- a/dns/definitions.go +++ b/dns/definitions.go @@ -13,7 +13,7 @@ var ( errNoNullTerm = errors.New("DNS name missing null terminator") errCalcLen = errors.New("DNS calculated name label length exceeds remaining buffer length") errCantAddLabel = errors.New("long/empty/zterm/escape DNS label or not enough space") - errBaseLen = errors.New("insufficient data for base length type") + errBaseLen = errors.New("DNS frame length too short") errReserved = errors.New("segment prefix is reserved") errTooManyPtr = errors.New("too many pointers (>10)") errInvalidPtr = errors.New("invalid pointer") @@ -41,8 +41,11 @@ type Frame struct { buf []byte } -func NewFrame(buf []byte) Frame { - return Frame{buf: buf} +func NewFrame(buf []byte) (Frame, error) { + if len(buf) < SizeHeader { + return Frame{}, errBaseLen + } + return Frame{buf: buf}, nil } func (frm Frame) TxID() uint16 { diff --git a/dns/dns.go b/dns/dns.go index 9226244..12a3069 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -66,7 +66,10 @@ func (m *Message) Decode(msg []byte) (_ uint16, incompleteButOK bool, err error) return 0, false, errResTooLong } m.Reset() - hdr := NewFrame(msg) + hdr, err := NewFrame(msg) + if err != nil { + return 0, false, err + } nq := int(hdr.QDCount()) off := uint16(SizeHeader) // Return tooManyErr if found to flag to the caller that the message was @@ -175,7 +178,10 @@ func (m *Message) AppendTo(buf []byte, txid uint16, flags HeaderFlags) (_ []byte nauth := uint16(len(m.Authorities)) nadd := uint16(len(m.Additionals)) var hdr [SizeHeader]byte - f := NewFrame(hdr[:]) + f, err := NewFrame(hdr[:]) + if err != nil { + return buf, err + } f.SetTxID(txid) f.SetFlags(flags) f.SetQDCount(nq) @@ -406,6 +412,9 @@ func NewName(domain string) (Name, error) { // Len returns the length over-the-wire of the encoded Name. func (n *Name) Len() uint16 { + if len(n.data) > math.MaxUint16 { + panic("size of DNS name data overflows 16bits") + } return uint16(len(n.data)) } diff --git a/examples/bridge/main.go b/examples/bridge/main.go index 6389a78..74c7129 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -14,6 +14,7 @@ import ( "github.com/soypat/lneto" "github.com/soypat/lneto/arp" "github.com/soypat/lneto/dhcpv4" + "github.com/soypat/lneto/dns" "github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/internal" "github.com/soypat/lneto/internal/ltesto" @@ -138,6 +139,7 @@ type Stack struct { arp internet.NodeARP udps internet.StackPorts dhcp dhcpv4.Client + dns dns.Client } func (s *Stack) Demux(b []byte, _ int) error { @@ -194,6 +196,34 @@ func (s *Stack) Reset(mac [6]byte, addr netip.Addr, mtu uint16) error { return nil } +func (s *Stack) StartLookupNetIP(host string) error { + name, err := dns.NewName(host) + if err != nil { + return err + } + err = s.dns.StartResolve(dns.ResolveConfig{ + Questions: []dns.Question{ + { + Name: name, + Type: dns.TypeA, + Class: dns.ClassINET, + }, + }, + EnableRecursion: true, + }) + if err != nil { + return err + } + var u internet.StackUDPPort + u.SetStackNode(&s.dns, nil, dns.ServerPort) + err = s.udps.Register(&u) + if err != nil { + return err + } + return err + return nil +} + func (s *Stack) BeginDHCPRequest() error { addr4 := s.ip.Addr().As4() var buf [4]byte @@ -208,7 +238,7 @@ func (s *Stack) BeginDHCPRequest() error { return err } var u internet.StackUDPPort - u.SetStackNode(&s.dhcp, dhcpv4.DefaultServerPort) + u.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) err = s.udps.Register(&u) if err != nil { return err diff --git a/internal/ip.go b/internal/ip.go index 2c3c13c..ac86109 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -10,7 +10,7 @@ var ( errInvalidIPVersionToSetAddr = errors.New("invalid ip version to setDstAddr") ) -func GetIPSourceAddr(buf []byte) (addr []byte, id, ipEndOff uint16, err error) { +func GetIPAddr(buf []byte) (src, dst []byte, id, ipEndOff uint16, err error) { b0 := buf[0] version := b0 >> 4 switch version { @@ -18,33 +18,41 @@ func GetIPSourceAddr(buf []byte) (addr []byte, id, ipEndOff uint16, err error) { ihl := b0 & 0xf ipEndOff = 4 * uint16(ihl) id = binary.BigEndian.Uint16(buf[4:6]) - addr = buf[12:16] + src = buf[12:16] + dst = buf[16:20] case 6: - addr = buf[8:24] + src = buf[8:24] + dst = buf[24:40] ipEndOff = 40 default: err = errUnsupportedIP } - return addr, id, ipEndOff, err + return src, dst, id, ipEndOff, err } -func SetIPDestinationAddr(buf []byte, id uint16, addr []byte) (err error) { - var dstaddr []byte +func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { + var dstaddr, srcaddr []byte version := buf[0] >> 4 switch version { case 4: + srcaddr = buf[12:16] dstaddr = buf[16:20] if id > 0 { binary.BigEndian.PutUint16(buf[4:6], id) } case 6: + srcaddr = buf[8:24] dstaddr = buf[24:40] default: - err = errUnsupportedIP + return errUnsupportedIP } - if err == nil && len(dstaddr) != len(addr) { - return errInvalidIPVersionToSetAddr + if src != nil && len(srcaddr) != len(src) { + return errors.New("mismatched length of ip src addr") } - copy(dstaddr, addr) + if dst != nil && len(dstaddr) != len(dst) { + return errors.New("mismatched length of ip dst addr") + } + copy(srcaddr, src) + copy(dstaddr, dst) return nil } diff --git a/internet/node-tcplistener.go b/internet/node-tcplistener.go index 2621f9d..8cac79e 100644 --- a/internet/node-tcplistener.go +++ b/internet/node-tcplistener.go @@ -123,7 +123,7 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e if err != nil { return err } - addr, _, _, err := internal.GetIPSourceAddr(carrierData) + srcaddr, _, _, _, err := internal.GetIPAddr(carrierData) if err != nil { return err } @@ -133,11 +133,11 @@ func (listener *NodeTCPListener) Demux(carrierData []byte, tcpFrameOffset int) e } src := tfrm.SourcePort() // Try to demux in accepted: - demuxed, err := listener.tryDemux(listener.accepted, src, addr, carrierData, tcpFrameOffset) + demuxed, err := listener.tryDemux(listener.accepted, src, srcaddr, carrierData, tcpFrameOffset) if demuxed { return err } - demuxed, err = listener.tryDemux(listener.ready, src, addr, carrierData, tcpFrameOffset) + demuxed, err = listener.tryDemux(listener.ready, src, srcaddr, carrierData, tcpFrameOffset) if demuxed { return err } diff --git a/internet/stack-udpport.go b/internet/stack-udpport.go index 9d0bb55..150c714 100644 --- a/internet/stack-udpport.go +++ b/internet/stack-udpport.go @@ -12,11 +12,13 @@ type StackUDPPort struct { h node vld lneto.Validator rmport uint16 + raddr []byte } -func (sudp *StackUDPPort) SetStackNode(node StackNode, rmport uint16) { +func (sudp *StackUDPPort) SetStackNode(node StackNode, raddr []byte, rmport uint16) { sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol()) sudp.rmport = rmport + sudp.raddr = append(sudp.raddr[:0], raddr...) } func (sudp *StackUDPPort) Protocol() uint64 { return uint64(lneto.IPProtoUDP) } @@ -42,6 +44,7 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error { if dst != sudp.h.port { return nil // Not meant for us. } + // TODO remote ip address handling. src := ufrm.SourcePort() if sudp.rmport != 0 && src != sudp.rmport { diff --git a/tcp/conn.go b/tcp/conn.go index 0146bf9..d2a4d4c 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -194,7 +194,7 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) { if off >= len(buf) { return errors.New("bad offset in TCPConn.Recv") } - raddr, id, _, err := internal.GetIPSourceAddr(buf[:off]) + raddr, _, id, _, err := internal.GetIPAddr(buf[:off]) if err != nil { return err } @@ -216,7 +216,7 @@ func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) { if len(conn.remoteAddr) == 0 { return 0, errors.New("unset IP address") } - raddr, _, _, err := internal.GetIPSourceAddr(buf[:off]) + raddr, _, _, _, err := internal.GetIPAddr(buf[:off]) if err != nil { return 0, err } else if len(raddr) != len(conn.remoteAddr) { @@ -226,7 +226,7 @@ func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) { if err != nil { return 0, err } - err = internal.SetIPDestinationAddr(buf[:off], conn.ipID, conn.remoteAddr) + err = internal.SetIPAddrs(buf[:off], conn.ipID, nil, conn.remoteAddr) if err != nil { return 0, err }