dns.mx raw

   1  package ws
   2  
   3  import (
   4  	"git.smesh.lol/moxie/pkg/mxutil"
   5  	"fmt"
   6  	"os"
   7  	"sync"
   8  	"syscall"
   9  	"time"
  10  )
  11  
  12  const dnsTTL = 24 * time.Hour
  13  
  14  var dnsAddr [4]byte
  15  var dnsAddrInit bool
  16  
  17  func initDNSAddr() {
  18  	if dnsAddrInit {
  19  		return
  20  	}
  21  	dnsAddrInit = true
  22  	dnsAddr = [4]byte{8, 8, 8, 8}
  23  	data, err := os.ReadFile("/etc/resolv.conf")
  24  	if err != nil {
  25  		return
  26  	}
  27  	i := 0
  28  	for i < len(data) {
  29  		lineStart := i
  30  		for i < len(data) && data[i] != '\n' {
  31  			i++
  32  		}
  33  		line := data[lineStart:i]
  34  		if i < len(data) {
  35  			i++
  36  		}
  37  		if len(line) > 11 && string(line[:11]) == "nameserver " {
  38  			ns := line[11:]
  39  			for len(ns) > 0 && (ns[len(ns)-1] == ' ' || ns[len(ns)-1] == '\t' || ns[len(ns)-1] == '\r') {
  40  				ns = ns[:len(ns)-1]
  41  			}
  42  			if ip := parseIPv4(ns); ip != nil {
  43  				dnsAddr = [4]byte{ip[0], ip[1], ip[2], ip[3]}
  44  				return
  45  			}
  46  		}
  47  	}
  48  }
  49  
  50  func parseIPv4(s []byte) (buf []byte) {
  51  	var parts [4]byte
  52  	p := 0
  53  	v := 0
  54  	for i := 0; i <= len(s); i++ {
  55  		if i == len(s) || s[i] == '.' {
  56  			if v > 255 || p > 3 {
  57  				return nil
  58  			}
  59  			parts[p] = byte(v)
  60  			p++
  61  			v = 0
  62  		} else if s[i] >= '0' && s[i] <= '9' {
  63  			v = v*10 + int32(s[i]-'0')
  64  		} else {
  65  			return nil
  66  		}
  67  	}
  68  	if p != 4 {
  69  		return nil
  70  	}
  71  	return parts[:]
  72  }
  73  
  74  type dnsEntry struct {
  75  	ip  string
  76  	exp time.Time
  77  }
  78  
  79  // dnsCacheState owns the resolver cache. A cache written after init must live
  80  // behind a self-mutating type: the rule forbids stores to a package global
  81  // outside init, and a method on the owner is the sanctioned place for them.
  82  type dnsCacheState struct {
  83  	mu sync.Mutex
  84  	m  map[string]dnsEntry
  85  }
  86  
  87  func (c *dnsCacheState) Get(host string) (e dnsEntry, ok bool) {
  88  	c.mu.Lock()
  89  	e, ok = c.m[host]
  90  	c.mu.Unlock()
  91  	return
  92  }
  93  
  94  func (c *dnsCacheState) Set(host string, e dnsEntry) {
  95  	c.mu.Lock()
  96  	c.m[host] = e
  97  	c.mu.Unlock()
  98  }
  99  
 100  var dnsCache dnsCacheState
 101  
 102  func init() {
 103  	dnsCache.m = map[string]dnsEntry{}
 104  }
 105  
 106  // ResolveHost is the exported entrypoint for the cached DNS resolver.
 107  // Used by other packages that need to dial a hostname without going
 108  // through net/http (which is broken in Moxie).
 109  func ResolveHost(host string) (addr string, err error) { return resolveHost(host) }
 110  
 111  // resolveHost returns a cached IP for the hostname, or resolves via raw UDP DNS.
 112  func resolveHost(host string) (addr string, derr error) {
 113  	if e, ok := dnsCache.Get(host); ok && time.Now().Before(e.exp) {
 114  		return e.ip, nil
 115  	}
 116  
 117  	ip, err := dnsLookup(host)
 118  	if err != nil {
 119  		return "", err
 120  	}
 121  
 122  	dnsCache.Set(host, dnsEntry{ip: ip, exp: time.Now().Add(dnsTTL)})
 123  
 124  	return ip, nil
 125  }
 126  
 127  // dnsLookup sends a minimal DNS A-record query via raw syscalls.
 128  func dnsLookup(host string) (addr string, derr error) {
 129  	fd, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_DGRAM, 0)
 130  	if err != nil {
 131  		return "", fmt.Errorf("dns: socket: %w", err)
 132  	}
 133  	defer syscall.Close(fd)
 134  
 135  	tv := syscall.Timeval{Sec: 5}
 136  	syscall.SetsockoptTimeval(fd, syscall.SOL_SOCKET, syscall.SO_RCVTIMEO, &tv)
 137  
 138  	initDNSAddr()
 139  	sa := &syscall.SockaddrInet4{Port: 53, Addr: dnsAddr}
 140  	if err = syscall.Connect(fd, sa); err != nil {
 141  		return "", fmt.Errorf("dns: connect: %w", err)
 142  	}
 143  
 144  	// Build DNS A-record query.
 145  	var pkt []byte
 146  	pkt = push(pkt, 0xAB, 0xCD)                         // ID
 147  	pkt = push(pkt, 0x01, 0x00)                         // flags: recursion desired
 148  	pkt = push(pkt, 0x00, 0x01)                         // 1 question
 149  	pkt = push(pkt, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00) // 0 answer/auth/additional
 150  
 151  	// Encode QNAME.
 152  	j := 0
 153  	for j < len(host) {
 154  		dot := j
 155  		for dot < len(host) && host[dot] != '.' {
 156  			dot++
 157  		}
 158  		pkt = mxutil.Ensure(pkt, 1)
 159  		pkt = push(pkt, byte(dot-j))
 160  		pkt = pkt | host[j:dot]
 161  		j = dot + 1
 162  	}
 163  	pkt = push(pkt, 0x00)       // root
 164  	pkt = push(pkt, 0x00, 0x01) // QTYPE A
 165  	pkt = push(pkt, 0x00, 0x01) // QCLASS IN
 166  
 167  	if err = syscall.Sendto(fd, pkt, 0, sa); err != nil {
 168  		return "", fmt.Errorf("dns: send: %w", err)
 169  	}
 170  
 171  	buf := []byte{:512}
 172  	n, _, err := syscall.Recvfrom(fd, buf, 0)
 173  	if err != nil {
 174  		return "", fmt.Errorf("dns: recv: %w", err)
 175  	}
 176  	if n < 12 {
 177  		return "", fmt.Errorf("dns: response too short")
 178  	}
 179  
 180  	anCount := int32(buf[6])<<8 | int32(buf[7])
 181  	if anCount == 0 {
 182  		return "", fmt.Errorf("dns: no answers for %s", host)
 183  	}
 184  
 185  	// Skip question section.
 186  	pos := 12
 187  	for pos < n {
 188  		if buf[pos] == 0 {
 189  			pos++
 190  			break
 191  		}
 192  		if buf[pos]&0xC0 == 0xC0 {
 193  			pos += 2
 194  			break
 195  		}
 196  		pos += int32(buf[pos]) + 1
 197  	}
 198  	pos += 4 // QTYPE + QCLASS
 199  
 200  	// Find first A record.
 201  	for a := 0; a < anCount && pos+10 < n; a++ {
 202  		if buf[pos]&0xC0 == 0xC0 {
 203  			pos += 2
 204  		} else {
 205  			for pos < n && buf[pos] != 0 {
 206  				pos += int32(buf[pos]) + 1
 207  			}
 208  			pos++
 209  		}
 210  		if pos+10 > n {
 211  			break
 212  		}
 213  		rtype := int32(buf[pos])<<8 | int32(buf[pos+1])
 214  		rdlen := int32(buf[pos+8])<<8 | int32(buf[pos+9])
 215  		pos += 10
 216  		if rtype == 1 && rdlen == 4 && pos+4 <= n {
 217  			// Copy Sprintf result to avoid buffer reuse corruption.
 218  			// The bytes are widened to int32: Moxie's fmt type-switch
 219  			// matches uint8 but not the named type byte, so %d on a
 220  			// byte prints "?" and the resolver would return "?.?.?.?"
 221  			// as a successful address.
 222  			s := fmt.Sprintf("%d.%d.%d.%d", int32(buf[pos]), int32(buf[pos+1]), int32(buf[pos+2]), int32(buf[pos+3]))
 223  			ip := []byte{:len(s)}
 224  			copy(ip, s)
 225  			return string(ip), nil
 226  		}
 227  		pos += rdlen
 228  	}
 229  
 230  	return "", fmt.Errorf("dns: no A record for %s", host)
 231  }
 232