Files
lneto/dns/client.go
T
Pat Whittingslow 75f1e02a20 dns: @hnw alternate proposal (#193)
* fix(dns): resolve CNAME chains and prevent invalid IP parsing

- Add automated in-band CNAME chain resolution to extract final A/AAAA IP addresses.
- Replace `MaxResponseAnswers` with `MaxIPs` and `MaxCNAMEs` to explicitly bound resource decoding and prevent memory exhaustion.
- Bound CNAME chain traversal to prevent infinite loops from cyclic records.

* refactor(dns): eagerly decode CNAME target into r.data, drop target field

Address review feedback on #189:

- Remove Resource.target: Resource.Decode expands CNAME RDATA in place
  into r.data (uncompressed wire format) while the full message is still
  available and updates header.Length, storing the name bytes exactly once.
- Add Resource.CNAMEView returning a length-bounded view of the expanded
  target; Message.WriteAnswers uses it.
- Rename matchesHost to ResourceHeader.pertainsTo.
- Make ResolveConfig.MaxCNAMEs explicit: drop the silent default in
  Client.StartResolve and let callers opt in (x/xnet).
- Reduce test suite to regression coverage only.

refs #189

* dns: simplify @hnw function and additional fixes

---------

Co-authored-by: Yoshio HANAWA <y@hnw.jp>
2026-08-28 21:39:12 -03:00

159 lines
3.9 KiB
Go

package dns
import (
"log/slog"
"math"
"net"
"net/netip"
"github.com/soypat/lneto"
"github.com/soypat/lneto/internal"
)
type Client struct {
connID uint64
txid uint16
lport uint16
msg Message
respFlags HeaderFlags
state StateClientQuery
enableRecursion bool
}
type ResolveConfig struct {
Questions []Question
Additional []Resource
EnableRecursion bool
// MaxResponseAnswers limits how many answer records are decoded from the
// DNS response. If zero it defaults to the number of Questions. Answers
// are decoded in wire order regardless of type, so a response resolved
// through CNAMEs needs room for the CNAME records as well as the addresses.
MaxResponseAnswers uint16
}
func (sudp *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) }
func (sudp *Client) LocalPort() uint16 { return sudp.lport }
func (sudp *Client) ConnectionID() *uint64 { return &sudp.connID }
func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error {
nd := len(cfg.Questions)
if nd > math.MaxUint16 || nd == 0 {
return lneto.ErrInvalidConfig
}
maxAns := cfg.MaxResponseAnswers
if maxAns == 0 {
maxAns = uint16(nd)
}
c.reset(localPort, txid, CQueryPending, cfg.EnableRecursion)
c.msg.LimitResourceDecoding(uint16(nd), maxAns, 0, 0)
c.msg.AddQuestions(cfg.Questions)
c.msg.AddAdditionals(cfg.Additional)
return nil
}
func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
if c.isClosed() {
return 0, net.ErrClosed
} else if c.state != CQueryPending {
return 0, nil
}
msg := &c.msg
frame := carrierData[offsetToFrame:]
msglen := msg.Len()
if msglen > uint16(len(frame)) {
return 0, errCalcLen
}
data, err := msg.AppendTo(frame[:0], c.txid, NewClientHeaderFlags(OpCodeQuery, c.enableRecursion))
if err != nil {
return 0, err
} else if len(data) > int(msglen) {
internal.LogAttrs(nil, slog.LevelError, "dns:unexpected-write", slog.Int("got", len(data)), slog.Int("want", int(msglen)))
return 0, lneto.ErrBug
}
c.state = CQueryOutstanding
// Unset don't frag since DNS requests go through LOTS of nodes.
// if frameOffset >= 28 {
// version := carrierData[0] >> 4
// if version == 4 {
// carrierData[6], carrierData[7] = 0, 0 // unset IP Flags.
// }
// }
return len(data), nil
}
func (c *Client) Demux(carrierData []byte, frameOffset int) error {
if c.isClosed() {
return net.ErrClosed
} else if c.state != CQueryOutstanding {
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 = CQueryDone
msg := &c.msg
_, incompleteButOK, err := msg.Decode(frame)
if err != nil && !incompleteButOK {
return err
}
return nil
}
func (c *Client) isClosed() bool {
return c.state == CQueryIdle || c.state == CQueryAborted
}
func (c *Client) ResponseCopyTo(dst *Message) (done bool, err error) {
if !c.respFlags.IsResponse() {
return false, nil
}
dst.CopyFrom(c.msg)
rcode := c.respFlags.ResponseCode()
if rcode != 0 {
return true, rcode
}
return true, nil
}
func (c *Client) ResponseAnswerLookup(dst []netip.Addr, host string) (uint16, error) {
if !c.respFlags.IsResponse() {
return 0, nil
}
rcode := c.respFlags.ResponseCode()
if rcode != 0 {
return 0, rcode
}
return c.msg.WriteAnswers(dst, host)
}
func (c *Client) ResponseFlags() (HeaderFlags, bool) {
return c.respFlags, c.respFlags.IsResponse()
}
func (c *Client) Abort() {
c.reset(0, 0, CQueryAborted, false)
}
func (c *Client) reset(lport, txid uint16, state StateClientQuery, enableRecursion bool) {
*c = Client{
connID: c.connID + 1,
lport: lport,
txid: txid,
msg: c.msg,
state: state,
enableRecursion: enableRecursion,
}
c.msg.Reset()
}