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
+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))
}