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