File
Blob: acme/dns.go
| 1 | // implementation is adopted from https://github.com/joohoi/acme-dns/blob/master/dns.go |
| 2 | package acme |
| 3 | |
| 4 | import ( |
| 5 | "context" |
| 6 | "net" |
| 7 | "sort" |
| 8 | "strconv" |
| 9 | "strings" |
| 10 | "time" |
| 11 | |
| 12 | "go.miragespace.co/specter/spec" |
| 13 | "go.miragespace.co/specter/spec/chord" |
| 14 | "go.miragespace.co/specter/timing" |
| 15 | "go.miragespace.co/specter/util" |
| 16 | |
| 17 | "github.com/miekg/dns" |
| 18 | "go.uber.org/zap" |
| 19 | ) |
| 20 | |
| 21 | type DNS struct { |
| 22 | parentCtx context.Context |
| 23 | logger *zap.Logger |
| 24 | storage chord.KV |
| 25 | soa *dns.SOA |
| 26 | records map[string][]dns.RR |
| 27 | domain string |
| 28 | } |
| 29 | |
| 30 | var _ dns.Handler = (*DNS)(nil) |
| 31 | |
| 32 | func NewDNS(ctx context.Context, logger *zap.Logger, kv chord.KV, email, domain string, ns map[string][]string) *DNS { |
| 33 | if !strings.HasSuffix(domain, ".") { |
| 34 | domain = domain + "." |
| 35 | } |
| 36 | server := &DNS{ |
| 37 | parentCtx: ctx, |
| 38 | logger: logger, |
| 39 | storage: kv, |
| 40 | domain: strings.ToLower(domain), |
| 41 | records: map[string][]dns.RR{ |
| 42 | domain: make([]dns.RR, 0), |
| 43 | }, |
| 44 | } |
| 45 | server.initStatic(email, ns) |
| 46 | |
| 47 | logger.Info("ACME DNS configured", zap.String("domain", server.domain), zap.String("mbox", server.soa.Mbox), zap.Any("ns", ns)) |
| 48 | |
| 49 | return server |
| 50 | } |
| 51 | |
| 52 | func (d *DNS) initStatic(email string, ns map[string][]string) { |
| 53 | nameservers := make([]string, 0) |
| 54 | for ns, ips := range ns { |
| 55 | ns = dns.Fqdn(ns) |
| 56 | |
| 57 | nsRR := new(dns.NS) |
| 58 | nsRR.Hdr = dns.RR_Header{Name: d.domain, Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 3600} |
| 59 | nsRR.Ns = ns |
| 60 | d.records[d.domain] = append(d.records[d.domain], nsRR) |
| 61 | |
| 62 | if _, ok := d.records[ns]; !ok { |
| 63 | d.records[ns] = make([]dns.RR, 0) |
| 64 | } |
| 65 | |
| 66 | for _, ip := range ips { |
| 67 | parsed := net.ParseIP(ip) |
| 68 | if parsed == nil { |
| 69 | continue |
| 70 | } |
| 71 | |
| 72 | if strings.Contains(ip, ":") { |
| 73 | aaaaRR := new(dns.AAAA) |
| 74 | aaaaRR.Hdr = dns.RR_Header{Name: dns.Fqdn(ns), Rrtype: dns.TypeAAAA, Class: dns.ClassINET, Ttl: 3600} |
| 75 | aaaaRR.AAAA = parsed |
| 76 | d.records[ns] = append(d.records[ns], aaaaRR) |
| 77 | } else { |
| 78 | aRR := new(dns.A) |
| 79 | aRR.Hdr = dns.RR_Header{Name: dns.Fqdn(ns), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 3600} |
| 80 | aRR.A = parsed |
| 81 | d.records[ns] = append(d.records[ns], aRR) |
| 82 | } |
| 83 | } |
| 84 | nameservers = append(nameservers, ns) |
| 85 | } |
| 86 | sort.Strings(nameservers) |
| 87 | |
| 88 | mbox := dns.CanonicalName(strings.ReplaceAll(email, "@", ".")) |
| 89 | |
| 90 | var ( |
| 91 | unix int64 |
| 92 | serial uint32 |
| 93 | ) |
| 94 | unix, err := strconv.ParseInt(spec.BuildTime, 10, 64) |
| 95 | if err != nil { |
| 96 | unix = time.Now().Unix() |
| 97 | } |
| 98 | serial = uint32(util.Must(strconv.ParseUint(time.Unix(unix, 0).UTC().Format("2006010215"), 10, 32))) |
| 99 | |
| 100 | soaRR := &dns.SOA{ |
| 101 | Hdr: dns.RR_Header{Name: d.domain, Rrtype: dns.TypeSOA, Class: dns.ClassINET, Ttl: 3600}, |
| 102 | Ns: nameservers[0], |
| 103 | Mbox: mbox, |
| 104 | Serial: serial, |
| 105 | Refresh: 28800, |
| 106 | Retry: 7200, |
| 107 | Expire: 604800, |
| 108 | Minttl: 86400, |
| 109 | } |
| 110 | d.records[d.domain] = append(d.records[d.domain], soaRR) |
| 111 | d.soa = soaRR |
| 112 | } |
| 113 | |
| 114 | func (d *DNS) ServeDNS(w dns.ResponseWriter, r *dns.Msg) { |
| 115 | m := new(dns.Msg) |
| 116 | m.SetReply(r) |
| 117 | defer w.WriteMsg(m) |
| 118 | |
| 119 | opt := r.IsEdns0() |
| 120 | if opt != nil { |
| 121 | m.SetEdns0(512, false) |
| 122 | if opt.Version() != 0 { |
| 123 | m.MsgHdr.Rcode = dns.RcodeBadVers |
| 124 | return |
| 125 | } |
| 126 | } |
| 127 | if r.Opcode == dns.OpcodeQuery { |
| 128 | d.readQuery(m) |
| 129 | } |
| 130 | } |
| 131 | |
| 132 | func (d *DNS) readQuery(m *dns.Msg) { |
| 133 | var authoritative = false |
| 134 | |
| 135 | for _, que := range m.Question { |
| 136 | rr, rc, auth := d.answer(que) |
| 137 | if auth { |
| 138 | authoritative = auth |
| 139 | } |
| 140 | m.MsgHdr.Rcode = rc |
| 141 | m.Answer = append(m.Answer, rr...) |
| 142 | } |
| 143 | |
| 144 | m.MsgHdr.Authoritative = authoritative |
| 145 | if authoritative && m.MsgHdr.Rcode == dns.RcodeNameError { |
| 146 | m.Ns = append(m.Ns, d.soa) |
| 147 | } |
| 148 | } |
| 149 | |
| 150 | func (d *DNS) isImmediate(q dns.Question) bool { |
| 151 | qname := strings.ToLower(q.Name) |
| 152 | query := strings.Split(qname, ".") |
| 153 | self := strings.Split(d.domain, ".") |
| 154 | return strings.HasSuffix(qname, d.domain) && |
| 155 | len(query) >= len(self) && |
| 156 | len(query)-len(self) <= 1 |
| 157 | } |
| 158 | |
| 159 | func (d *DNS) answer(q dns.Question) (rr []dns.RR, rcode int, auth bool) { |
| 160 | auth = d.isImmediate(q) |
| 161 | if !auth { |
| 162 | return nil, dns.RcodeNameError, true |
| 163 | } |
| 164 | |
| 165 | if q.Qtype == dns.TypeANY { |
| 166 | rcode = dns.RcodeNotImplemented |
| 167 | return |
| 168 | } |
| 169 | |
| 170 | qname := strings.ToLower(q.Name) |
| 171 | defined := d.records[qname] |
| 172 | for _, ri := range defined { |
| 173 | if ri.Header().Rrtype == q.Qtype { |
| 174 | rr = append(rr, ri) |
| 175 | } |
| 176 | } |
| 177 | |
| 178 | if q.Qtype == dns.TypeTXT { |
| 179 | txtRRs, err := d.answerTXT(q) |
| 180 | if err != nil { |
| 181 | rcode = dns.RcodeServerFailure |
| 182 | } else { |
| 183 | rr = append(rr, txtRRs...) |
| 184 | } |
| 185 | } |
| 186 | |
| 187 | if len(rr) == 0 && rcode != dns.RcodeServerFailure { |
| 188 | rcode = dns.RcodeNameError |
| 189 | } |
| 190 | |
| 191 | d.logger.Debug("Answering DNS query", zap.String("qtype", dns.TypeToString[q.Qtype]), zap.String("rcode", dns.RcodeToString[rcode]), zap.String("domain", qname)) |
| 192 | |
| 193 | return rr, rcode, auth |
| 194 | } |
| 195 | |
| 196 | func (d *DNS) answerTXT(q dns.Question) ([]dns.RR, error) { |
| 197 | var ra []dns.RR |
| 198 | |
| 199 | qname := strings.ToLower(q.Name) |
| 200 | idx := strings.Index(qname, d.domain) |
| 201 | if idx <= 0 { |
| 202 | return ra, nil |
| 203 | } |
| 204 | subdomain := qname[0 : idx-1] |
| 205 | |
| 206 | callCtx, cancel := context.WithTimeout(d.parentCtx, timing.DNSLookupTimeout) |
| 207 | defer cancel() |
| 208 | |
| 209 | vals, err := d.storage.PrefixList(callCtx, []byte(dnsKeyName(subdomain))) |
| 210 | if err != nil { |
| 211 | d.logger.Error("Failed to lookup TXT", zap.String("subdomain", subdomain), zap.Error(err)) |
| 212 | return nil, err |
| 213 | } |
| 214 | |
| 215 | for _, v := range vals { |
| 216 | if len(v) > 0 { |
| 217 | r := new(dns.TXT) |
| 218 | r.Hdr = dns.RR_Header{Name: q.Name, Rrtype: dns.TypeTXT, Class: dns.ClassINET, Ttl: 1} |
| 219 | r.Txt = append(r.Txt, string(v)) |
| 220 | ra = append(ra, r) |
| 221 | } |
| 222 | } |
| 223 | |
| 224 | return ra, nil |
| 225 | } |