mirror of
https://github.com/soypat/lneto.git
synced 2026-09-01 04:19:05 +00:00
begin touching up a DNS Client implementation; internal.GetIP/SetIP refactor
This commit is contained in:
+124
@@ -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
@@ -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
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user