mirror of
https://github.com/soypat/lneto.git
synced 2026-08-15 04:13:44 +00:00
begin touching up a DNS Client implementation; internal.GetIP/SetIP refactor
This commit is contained in:
+9
-6
@@ -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
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
+31
-1
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user