From 68c4d1ea719b6482af871273957bdf3b27f03e5e Mon Sep 17 00:00:00 2001 From: soypat Date: Sat, 12 Jul 2025 12:34:20 -0300 Subject: [PATCH] DNS working! --- dns/client.go | 21 ++++++++--- dns/dns.go | 78 ++++++++++++++++++++++++++++++----------- dns/dns_test.go | 6 ++-- examples/bridge/main.go | 31 +++++++++------- internet/stack-ip.go | 2 +- 5 files changed, 96 insertions(+), 42 deletions(-) diff --git a/dns/client.go b/dns/client.go index c8eda70..77cdbb0 100644 --- a/dns/client.go +++ b/dns/client.go @@ -12,6 +12,7 @@ import ( type Client struct { connID uint64 txid uint16 + lport uint16 msg Message respFlags HeaderFlags state clientState @@ -20,23 +21,25 @@ type Client struct { type ResolveConfig struct { Questions []Question + Additional []Resource EnableRecursion bool } func (sudp *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) } -func (sudp *Client) LocalPort() uint16 { return ClientPort } +func (sudp *Client) LocalPort() uint16 { return sudp.lport } func (sudp *Client) ConnectionID() *uint64 { return &sudp.connID } -func (c *Client) StartResolve(txid uint16, cfg ResolveConfig) error { +func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error { nd := len(cfg.Questions) if nd > math.MaxUint16 { return errors.New("overflow uint16 in DNS questions") } - c.reset(txid, dnsSendQuery, cfg.EnableRecursion) + c.reset(localPort, txid, dnsSendQuery, cfg.EnableRecursion) c.msg.LimitResourceDecoding(uint16(nd), uint16(nd), 0, 0) c.msg.AddQuestions(cfg.Questions) + c.msg.AddAdditionals(cfg.Additional) return nil } @@ -61,6 +64,13 @@ func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { return 0, fmt.Errorf("unexpected write %d v %d", len(data), msglen) } c.state = dnsAwaitResponse + // 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 } @@ -113,12 +123,13 @@ func (c *Client) Answers() []Resource { } func (c *Client) Abort() { - c.reset(0, 0, false) + c.reset(0, 0, 0, false) } -func (c *Client) reset(txid uint16, state clientState, enableRecursion bool) { +func (c *Client) reset(lport, txid uint16, state clientState, enableRecursion bool) { *c = Client{ connID: c.connID + 1, + lport: lport, txid: txid, msg: c.msg, state: state, diff --git a/dns/dns.go b/dns/dns.go index c5e99dc..105eb70 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -37,7 +37,7 @@ type Question struct { } type Resource struct { - Header ResourceHeader + header ResourceHeader data []byte } @@ -55,6 +55,35 @@ type Name struct { data []byte } +type ZFlags uint16 + +func NewResource(name Name, typ Type, class Class, ttl uint32, data []byte) Resource { + return Resource{ + header: ResourceHeader{ + Name: name, + Type: typ, + Class: class, + TTL: ttl, + Length: uint16(len(data)), + }, + data: data, + } +} + +func (r *Resource) SetEDNS0(UDPlength uint16, rcode RCode, zflags ZFlags, data []byte) { + if len(data) > math.MaxUint16-2 || len(data)+8+2*SizeHeader > int(UDPlength) { + panic("too large data") + } + r.header = ResourceHeader{ + Name: Name{data: rootDomain}, + Type: TypeOPT, + Class: Class(UDPlength), + TTL: uint32(rcode)<<24 | 0<<16 | uint32(zflags), + Length: uint16(len(data)), + } + r.data = append(r.data[:0], data...) +} + // Decode decodes the DNS message in b into m. It returns the number of bytes // consumed from b (0 if no bytes were consumed) and any error encountered. // If the message was not completely parsed due to LimitResourceDecoding, @@ -98,7 +127,8 @@ func (m *Message) Decode(msg []byte) (_ uint16, incompleteButOK bool, err error) } } // Skip undecoded questions. - for i := 0; i < int(hdr.QDCount())-nq; i++ { + qd := hdr.QDCount() + for i := 0; i < int(qd)-nq; i++ { off, err = skipQuestion(msg, off) if err != nil { return off, false, err @@ -245,9 +275,16 @@ func (m *Message) AddQuestions(questions []Question) { m.Questions = slices.Grow(m.Questions, len(questions)) m.Questions = m.Questions[:qoff+len(questions)] for i := range questions { - m.Questions[qoff+i].Name.CopyFrom(questions[i].Name) - m.Questions[qoff+i].Type = questions[i].Type - m.Questions[qoff+i].Class = questions[i].Class + m.Questions[qoff+i].CopyFrom(questions[i]) + } +} + +func (m *Message) AddAdditionals(rsc []Resource) { + aoff := len(m.Additionals) + m.Additionals = slices.Grow(m.Additionals, len(rsc)) + m.Additionals = m.Additionals[:aoff+len(rsc)] + for i := range rsc { + m.Additionals[aoff+i].CopyFrom(rsc[i]) } } @@ -272,16 +309,12 @@ func (h *ResourceHeader) String() string { } func (r *Resource) Reset() { - r.Header.Reset() + r.header.Reset() r.data = r.data[:0] } -func (r *Resource) Len() uint16 { - return r.Header.Name.Len() + 10 + uint16(len(r.data)) -} - func (r *Resource) RawData() []byte { - length := r.Header.Length + length := r.header.Length if int(length) > len(r.data) { length = uint16(len(r.data)) } @@ -330,20 +363,19 @@ func (q *Question) String() string { } func (r *Resource) Decode(b []byte, off uint16) (uint16, error) { - off, err := r.Header.Decode(b, off) + off, err := r.header.Decode(b, off) if err != nil { return off, err } - if r.Header.Length > uint16(len(b[off:])) { + if r.header.Length > uint16(len(b[off:])) { return off, errResourceLen } - r.data = append(r.data[:0], b[off:off+r.Header.Length]...) - return off + r.Header.Length, nil + r.data = append(r.data[:0], b[off:off+r.header.Length]...) + return off + r.header.Length, nil } func (r *Resource) appendTo(buf []byte) (_ []byte, err error) { - r.Header.Length = uint16(len(r.data)) - buf, err = r.Header.appendTo(buf) + buf, err = r.header.appendTo(buf) if err != nil { return buf, err } @@ -351,6 +383,10 @@ func (r *Resource) appendTo(buf []byte) (_ []byte, err error) { return buf, nil } +func (r *Resource) Len() uint16 { + return r.header.Name.Len() + 10 + uint16(len(r.data)) +} + func (rhdr *ResourceHeader) Decode(msg []byte, off uint16) (uint16, error) { off, err := rhdr.Name.Decode(msg, off) if err != nil { @@ -386,7 +422,7 @@ func MustNewName(s string) Name { return name } -var emptyDomain = []byte{0} +var rootDomain = []byte{0} // NewName parses a domain name and returns a new Name. func NewName(domain string) (Name, error) { @@ -394,7 +430,7 @@ func NewName(domain string) (Name, error) { return Name{}, errEmptyDomainName } if len(domain) == 1 && domain[0] == '.' { - return Name{data: emptyDomain}, nil + return Name{data: append([]byte{}, rootDomain...)}, nil } var name Name for len(domain) > 0 { @@ -462,7 +498,7 @@ func (n *Name) Decode(b []byte, off uint16) (uint16, error) { return off, nil } -// Reset resets the Name labels to be empty and reuses buffer. +// Reset resets the Name labels to be empty andatad reuses buffer. func (n *Name) Reset() { n.data = n.data[:0] } // CanAddLabel reports whether the label can be added to the name. @@ -605,7 +641,7 @@ func (dst *Question) CopyFrom(q Question) { } func (dst *Resource) CopyFrom(r Resource) { - dst.Header.CopyFrom(r.Header) + dst.header.CopyFrom(r.header) dst.data = append(dst.data[:0], r.data...) } diff --git a/dns/dns_test.go b/dns/dns_test.go index 3ed4a82..6ddaa97 100644 --- a/dns/dns_test.go +++ b/dns/dns_test.go @@ -97,7 +97,7 @@ func TestMessageAppendEncode(t *testing.T) { }, Answers: []Resource{ { - Header: ResourceHeader{ + header: ResourceHeader{ Name: MustNewName("."), Type: TypeA, Class: ClassINET, @@ -147,7 +147,7 @@ func TestMessageAppendEncodeIncompleteOK(t *testing.T) { }, Answers: []Resource{ { - Header: ResourceHeader{ + header: ResourceHeader{ Name: MustNewName("."), Type: TypeA, Class: ClassINET, @@ -157,7 +157,7 @@ func TestMessageAppendEncodeIncompleteOK(t *testing.T) { data: []byte{1, 2, 3}, }, { - Header: ResourceHeader{ + header: ResourceHeader{ Name: MustNewName("."), Type: TypeA, Class: ClassINET, diff --git a/examples/bridge/main.go b/examples/bridge/main.go index 97114cb..f8e1d43 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -179,13 +179,14 @@ func run() (err error) { } type Stack struct { - link internet.StackEthernet - ip internet.StackIP - arp internet.NodeARP - udps internet.StackPorts - dhcp dhcpv4.Client - dns dns.Client - lookup dns.Message + link internet.StackEthernet + ip internet.StackIP + arp internet.NodeARP + udps internet.StackPorts + dhcp dhcpv4.Client + dns dns.Client + ednsopt dns.Resource + lookup dns.Message // Packet capture and top level filtering. shark pcap.PacketBreakdown @@ -194,10 +195,10 @@ type Stack struct { func (s *Stack) Demux(b []byte, _ int) (err error) { s.aux, err = s.shark.CaptureEthernet(s.aux[:0], b, 0) - baseFrame := s.aux[len(s.aux)-1] - isOK := baseFrame.Protocol == "DHCPv4" || - (baseFrame.Protocol == lneto.IPProtoUDP && getField(baseFrame, b, pcap.FieldClassDst) == 53) || - baseFrame.Protocol == ethernet.TypeARP + topFrame := s.aux[len(s.aux)-1] + isOK := topFrame.Protocol == "DHCPv4" || // Allow DHCP responses. + (topFrame.Protocol == lneto.IPProtoUDP && getField(topFrame, b, pcap.FieldClassSrc) == 53) || // Allow DNS responses. + topFrame.Protocol == ethernet.TypeARP // Allow ARP responses. if !isOK { return nil } @@ -278,7 +279,10 @@ func (s *Stack) StartLookupIP(host string) error { if err != nil { return err } - err = s.dns.StartResolve(uint16(softRand), dns.ResolveConfig{ + s.link.SetHardwareAddr6([6]byte{0xd8, 0x5e, 0xd3, 0x43, 0x03, 0xeb}) + s.ip.SetAddr(netip.AddrFrom4([4]byte{192, 168, 1, 53})) + s.ednsopt.SetEDNS0(uint16(s.link.MTU())-100, 0, 0, nil) + err = s.dns.StartResolve(uint16(softRand>>1)+1024, uint16(softRand), dns.ResolveConfig{ Questions: []dns.Question{ { Name: name, @@ -286,6 +290,9 @@ func (s *Stack) StartLookupIP(host string) error { Class: dns.ClassINET, }, }, + Additional: []dns.Resource{ + s.ednsopt, + }, EnableRecursion: true, }) if err != nil { diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 771fbde..34315b8 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -148,6 +148,7 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) id := internal.Prand16(seed) ifrm.SetID(id) ifrm.SetFlags(dontFrag) + ifrm.SetTTL(64) *ifrm.SourceAddr() = sb.ip sb.ipID = id for i := range sb.handlers { @@ -166,7 +167,6 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) } totalLen := n + headerlen ifrm.SetTotalLength(uint16(totalLen)) - ifrm.SetTTL(64) ifrm.SetProtocol(proto) ifrm.SetCRC(ifrm.CalculateHeaderCRC()) // Calculate CRC for our newly generated packet.