begin touching up a DNS Client implementation; internal.GetIP/SetIP refactor

This commit is contained in:
soypat
2025-06-21 01:29:38 -03:00
parent 48f1fa9103
commit 1fcde05284
9 changed files with 209 additions and 29 deletions
+9 -6
View File
@@ -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
}
+124
View File
@@ -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
)
+6 -3
View File
@@ -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 {
+11 -2
View File
@@ -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))
}
+31 -1
View File
@@ -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
+18 -10
View File
@@ -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
}
+3 -3
View File
@@ -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
}
+4 -1
View File
@@ -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 {
+3 -3
View File
@@ -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
}