main
go 513 lines 13.1 KB
Raw
1 package vultr
2
3 import (
4 "context"
5 "errors"
6 "fmt"
7 "strconv"
8 "strings"
9
10 "github.com/go-acme/lego/v4/challenge"
11 legovultr "github.com/go-acme/lego/v4/providers/dns/vultr"
12 "github.com/vultr/govultr/v3"
13 "golang.org/x/oauth2"
14
15 "github.com/gosuda/portal-tunnel/v2/utils"
16 )
17
18 const defaultRecordTTL = 60
19
20 type Provider struct {
21 apiKey string
22
23 zones *utils.Snapshot[map[string]string]
24 }
25
26 func New(apiKey string) *Provider {
27 return &Provider{
28 apiKey: strings.TrimSpace(apiKey),
29 zones: utils.NewSnapshot(map[string]string{}, utils.CloneMap[string, string]),
30 }
31 }
32
33 func (p *Provider) Name() string {
34 return "vultr"
35 }
36
37 func (p *Provider) ChallengeProvider(context.Context) (challenge.Provider, error) {
38 if p == nil {
39 return nil, errors.New("vultr provider is nil")
40 }
41 if p.apiKey == "" {
42 return nil, errors.New("vultr api key is required")
43 }
44
45 cfg := legovultr.NewDefaultConfig()
46 cfg.APIKey = p.apiKey
47
48 provider, err := legovultr.NewDNSProviderConfig(cfg)
49 if err != nil {
50 return nil, fmt.Errorf("create vultr lego provider: %w", err)
51 }
52 return provider, nil
53 }
54
55 func (p *Provider) EnsureARecords(ctx context.Context, baseDomain, publicIPv4 string) error {
56 if p == nil {
57 return errors.New("vultr provider is nil")
58 }
59 baseDomain = utils.NormalizeBaseDomain(baseDomain)
60 if baseDomain == "" {
61 return errors.New("base domain is required")
62 }
63 if err := utils.ValidateIPv4(publicIPv4); err != nil {
64 return err
65 }
66
67 client, err := p.newClient(ctx)
68 if err != nil {
69 return err
70 }
71 zone, err := p.findZone(ctx, client, baseDomain)
72 if err != nil {
73 return err
74 }
75
76 for _, recordName := range []string{baseDomain, "*." + baseDomain} {
77 if err := ensureRecord(ctx, client, zone, recordName, "A", strings.TrimSpace(publicIPv4)); err != nil {
78 return fmt.Errorf("upsert vultr A record %s: %w", recordName, err)
79 }
80 }
81 return nil
82 }
83
84 func (p *Provider) EnsureARecord(ctx context.Context, name, publicIPv4 string) error {
85 if p == nil {
86 return errors.New("vultr provider is nil")
87 }
88 name = utils.NormalizeHostname(name)
89 if name == "" {
90 return errors.New("record name is required")
91 }
92 if err := utils.ValidateIPv4(publicIPv4); err != nil {
93 return err
94 }
95
96 client, zone, err := p.clientAndZone(ctx, name)
97 if err != nil {
98 return err
99 }
100 if err := ensureRecord(ctx, client, zone, name, "A", strings.TrimSpace(publicIPv4)); err != nil {
101 return fmt.Errorf("upsert vultr A record %s: %w", name, err)
102 }
103 return nil
104 }
105
106 func (p *Provider) DeleteARecord(ctx context.Context, name string) error {
107 if p == nil {
108 return errors.New("vultr provider is nil")
109 }
110 name = utils.NormalizeHostname(name)
111 if name == "" {
112 return errors.New("record name is required")
113 }
114
115 client, zone, err := p.clientAndZone(ctx, name)
116 if err != nil {
117 return err
118 }
119 if err := deleteRecords(ctx, client, zone, name, "A", ""); err != nil {
120 return fmt.Errorf("delete vultr A record %s: %w", name, err)
121 }
122 return nil
123 }
124
125 func (p *Provider) EnsureTXTRecord(ctx context.Context, name, value string) error {
126 if p == nil {
127 return errors.New("vultr provider is nil")
128 }
129 name = utils.NormalizeHostname(name)
130 if name == "" {
131 return errors.New("record name is required")
132 }
133 value = strings.TrimSpace(value)
134 if value == "" {
135 return errors.New("txt record value is required")
136 }
137
138 client, zone, err := p.clientAndZone(ctx, name)
139 if err != nil {
140 return err
141 }
142 if err := ensureTXTRecord(ctx, client, zone, name, value); err != nil {
143 return fmt.Errorf("upsert vultr TXT record %s: %w", name, err)
144 }
145 return nil
146 }
147
148 func (p *Provider) DeleteTXTRecords(ctx context.Context, name, matchPrefix string) error {
149 if p == nil {
150 return errors.New("vultr provider is nil")
151 }
152 name = utils.NormalizeHostname(name)
153 if name == "" {
154 return errors.New("record name is required")
155 }
156 matchPrefix = strings.TrimSpace(matchPrefix)
157 if matchPrefix == "" {
158 return errors.New("txt record match prefix is required")
159 }
160
161 client, zone, err := p.clientAndZone(ctx, name)
162 if err != nil {
163 return err
164 }
165 if err := deleteRecords(ctx, client, zone, name, "TXT", matchPrefix); err != nil {
166 return fmt.Errorf("delete vultr TXT records %s: %w", name, err)
167 }
168 return nil
169 }
170
171 func (p *Provider) EnsureHTTPSRecord(ctx context.Context, name string, _ uint16, _, _, content string) error {
172 if p == nil {
173 return errors.New("vultr provider is nil")
174 }
175 name = utils.NormalizeHostname(name)
176 if name == "" {
177 return errors.New("record name is required")
178 }
179 content = strings.TrimSpace(content)
180 if content == "" {
181 return errors.New("https record content is required")
182 }
183
184 client, zone, err := p.clientAndZone(ctx, name)
185 if err != nil {
186 return err
187 }
188 if err := ensureRecord(ctx, client, zone, name, "HTTPS", content); err != nil {
189 return fmt.Errorf("upsert vultr HTTPS record %s: %w", name, err)
190 }
191 return nil
192 }
193
194 func (p *Provider) DeleteHTTPSRecord(ctx context.Context, name string) error {
195 if p == nil {
196 return errors.New("vultr provider is nil")
197 }
198 name = utils.NormalizeHostname(name)
199 if name == "" {
200 return errors.New("record name is required")
201 }
202
203 client, zone, err := p.clientAndZone(ctx, name)
204 if err != nil {
205 return err
206 }
207 if err := deleteRecords(ctx, client, zone, name, "HTTPS", ""); err != nil {
208 return fmt.Errorf("delete vultr HTTPS record %s: %w", name, err)
209 }
210 return nil
211 }
212
213 func (p *Provider) EnsureDNSSEC(ctx context.Context, baseDomain string) (state, dsRecord, message string, err error) {
214 if p == nil {
215 return "", "", "", errors.New("vultr provider is nil")
216 }
217 baseDomain = utils.NormalizeBaseDomain(baseDomain)
218 if baseDomain == "" {
219 return "", "", "", errors.New("base domain is required")
220 }
221
222 client, err := p.newClient(ctx)
223 if err != nil {
224 return "", "", "", err
225 }
226 zone, err := p.findZone(ctx, client, baseDomain)
227 if err != nil {
228 return "", "", "", err
229 }
230
231 domain, _, err := client.Domain.Get(ctx, zone)
232 if err != nil {
233 return "", "", "", fmt.Errorf("get vultr domain %s: %w", zone, err)
234 }
235 if domain != nil {
236 state = strings.TrimSpace(domain.DNSSec)
237 }
238 if !strings.EqualFold(state, "enabled") {
239 if err := client.Domain.Update(ctx, zone, "enabled"); err != nil {
240 return "", "", "", fmt.Errorf("enable vultr dnssec: %w", err)
241 }
242 domain, _, err = client.Domain.Get(ctx, zone)
243 if err != nil {
244 return "", "", "", fmt.Errorf("refresh vultr domain %s: %w", zone, err)
245 }
246 state = "enabled"
247 if domain != nil && strings.TrimSpace(domain.DNSSec) != "" {
248 state = strings.TrimSpace(domain.DNSSec)
249 }
250 }
251
252 records, _, err := client.Domain.GetDNSSec(ctx, zone)
253 if err != nil {
254 return "", "", "", fmt.Errorf("get vultr dnssec records: %w", err)
255 }
256 dsRecord = preferredDSRecord(records)
257 if dsRecord != "" {
258 message = "publish the DS record at the registrar after Vultr zone signing is enabled"
259 } else if strings.EqualFold(state, "enabled") {
260 message = "wait for the active Vultr DS record before updating the registrar"
261 }
262 return state, dsRecord, message, nil
263 }
264
265 func (p *Provider) clientAndZone(ctx context.Context, domain string) (*govultr.Client, string, error) {
266 client, err := p.newClient(ctx)
267 if err != nil {
268 return nil, "", err
269 }
270 zone, err := p.findZone(ctx, client, domain)
271 if err != nil {
272 return nil, "", err
273 }
274 return client, zone, nil
275 }
276
277 func (p *Provider) newClient(ctx context.Context) (*govultr.Client, error) {
278 if p == nil {
279 return nil, errors.New("vultr provider is nil")
280 }
281 if p.apiKey == "" {
282 return nil, errors.New("vultr api key is required")
283 }
284 return govultr.NewClient(oauth2.NewClient(ctx, oauth2.StaticTokenSource(&oauth2.Token{AccessToken: p.apiKey}))), nil
285 }
286
287 func (p *Provider) findZone(ctx context.Context, client *govultr.Client, domain string) (string, error) {
288 if client == nil {
289 return "", errors.New("vultr client is nil")
290 }
291 domain = utils.NormalizeHostname(domain)
292 candidates := utils.DomainCandidates(domain)
293
294 zones := p.zones.Load()
295 for _, candidate := range candidates {
296 if zone := zones[candidate]; zone != "" {
297 return zone, nil
298 }
299 }
300
301 listOptions := &govultr.ListOptions{PerPage: 100}
302 for {
303 domains, meta, _, err := client.Domain.List(ctx, listOptions)
304 if err != nil {
305 return "", fmt.Errorf("list vultr domains: %w", err)
306 }
307 for _, item := range domains {
308 zone := utils.NormalizeBaseDomain(item.Domain)
309 for _, candidate := range candidates {
310 if zone != candidate {
311 continue
312 }
313 p.zones.UpdateCopy(func(zones *map[string]string) {
314 if *zones == nil {
315 *zones = make(map[string]string)
316 }
317 (*zones)[candidate] = zone
318 })
319 return zone, nil
320 }
321 }
322 if meta == nil || meta.Links == nil || meta.Links.Next == "" {
323 break
324 }
325 listOptions.Cursor = meta.Links.Next
326 }
327
328 return "", fmt.Errorf("no vultr domain found for %s", domain)
329 }
330
331 func ensureRecord(ctx context.Context, client *govultr.Client, zone, fqdn, recordType, data string) error {
332 recordName, err := relativeRecordName(fqdn, zone)
333 if err != nil {
334 return err
335 }
336 existing, err := listRecords(ctx, client, zone, fqdn, recordType)
337 if err != nil {
338 return err
339 }
340 for _, record := range existing {
341 if strings.TrimSpace(record.Data) == data {
342 return nil
343 }
344 }
345 if len(existing) > 0 {
346 name := recordName
347 return client.DomainRecord.Update(ctx, zone, existing[0].ID, &govultr.DomainRecordUpdateReq{
348 Name: &name,
349 Type: recordType,
350 Data: data,
351 TTL: defaultRecordTTL,
352 })
353 }
354
355 _, _, err = client.DomainRecord.Create(ctx, zone, &govultr.DomainRecordCreateReq{
356 Name: recordName,
357 Type: recordType,
358 Data: data,
359 TTL: defaultRecordTTL,
360 })
361 return err
362 }
363
364 func ensureTXTRecord(ctx context.Context, client *govultr.Client, zone, fqdn, value string) error {
365 recordName, err := relativeRecordName(fqdn, zone)
366 if err != nil {
367 return err
368 }
369 existing, err := listRecords(ctx, client, zone, fqdn, "TXT")
370 if err != nil {
371 return err
372 }
373 for _, record := range existing {
374 if txtContent(record.Data) == value {
375 return nil
376 }
377 }
378
379 _, _, err = client.DomainRecord.Create(ctx, zone, &govultr.DomainRecordCreateReq{
380 Name: recordName,
381 Type: "TXT",
382 Data: value,
383 TTL: defaultRecordTTL,
384 })
385 return err
386 }
387
388 func deleteRecords(ctx context.Context, client *govultr.Client, zone, fqdn, recordType, matchPrefix string) error {
389 existing, err := listRecords(ctx, client, zone, fqdn, recordType)
390 if err != nil {
391 return err
392 }
393 for _, record := range existing {
394 if matchPrefix != "" && !strings.HasPrefix(txtContent(record.Data), matchPrefix) {
395 continue
396 }
397 if err := client.DomainRecord.Delete(ctx, zone, record.ID); err != nil {
398 return err
399 }
400 }
401 return nil
402 }
403
404 func listRecords(ctx context.Context, client *govultr.Client, zone, fqdn, recordType string) ([]govultr.DomainRecord, error) {
405 if client == nil {
406 return nil, errors.New("vultr client is nil")
407 }
408 recordName, err := relativeRecordName(fqdn, zone)
409 if err != nil {
410 return nil, err
411 }
412 recordType = strings.ToUpper(strings.TrimSpace(recordType))
413 listOptions := &govultr.ListOptions{PerPage: 100}
414
415 var filtered []govultr.DomainRecord
416 for {
417 records, meta, _, err := client.DomainRecord.List(ctx, zone, listOptions)
418 if err != nil {
419 return nil, err
420 }
421 for _, record := range records {
422 if !strings.EqualFold(strings.TrimSpace(record.Type), recordType) || !sameRecordName(record.Name, recordName, fqdn, zone) {
423 continue
424 }
425 filtered = append(filtered, record)
426 }
427 if meta == nil || meta.Links == nil || meta.Links.Next == "" {
428 break
429 }
430 listOptions.Cursor = meta.Links.Next
431 }
432 return filtered, nil
433 }
434
435 func relativeRecordName(fqdn, zone string) (string, error) {
436 fqdn = utils.NormalizeHostname(fqdn)
437 zone = utils.NormalizeBaseDomain(zone)
438 if fqdn == "" {
439 return "", errors.New("record name is required")
440 }
441 if zone == "" {
442 return "", errors.New("vultr zone is required")
443 }
444 if fqdn == zone {
445 return "@", nil
446 }
447 suffix := "." + zone
448 if !strings.HasSuffix(fqdn, suffix) {
449 return "", fmt.Errorf("hostname %q is outside vultr zone %q", fqdn, zone)
450 }
451 return strings.TrimSuffix(fqdn, suffix), nil
452 }
453
454 func sameRecordName(recordName, expected, fqdn, zone string) bool {
455 recordName = utils.NormalizeHostname(recordName)
456 expected = strings.TrimSpace(strings.ToLower(expected))
457 fqdn = utils.NormalizeHostname(fqdn)
458 zone = utils.NormalizeBaseDomain(zone)
459
460 if recordName == expected {
461 return true
462 }
463 if expected == "@" && (recordName == "" || recordName == zone || recordName == fqdn) {
464 return true
465 }
466 return recordName == fqdn
467 }
468
469 func txtContent(raw string) string {
470 unquoted, err := strconv.Unquote(strings.TrimSpace(raw))
471 if err == nil {
472 return unquoted
473 }
474 return strings.Trim(strings.TrimSpace(raw), "\"")
475 }
476
477 func preferredDSRecord(records []string) string {
478 candidates := make(map[string]string, len(records))
479 first := ""
480 for _, raw := range records {
481 ds := normalizeDSRecord(raw)
482 if ds == "" {
483 continue
484 }
485 if first == "" {
486 first = ds
487 }
488 fields := strings.Fields(ds)
489 if len(fields) < 4 {
490 continue
491 }
492 candidates[fields[2]] = ds
493 }
494 for _, digestType := range []string{"2", "4", "1"} {
495 if record := candidates[digestType]; record != "" {
496 return record
497 }
498 }
499 return first
500 }
501
502 func normalizeDSRecord(raw string) string {
503 fields := strings.Fields(strings.TrimSpace(raw))
504 for i, field := range fields {
505 if strings.EqualFold(field, "DS") && len(fields) >= i+5 {
506 return strings.Join(fields[i+1:i+5], " ")
507 }
508 }
509 if len(fields) == 4 {
510 return strings.Join(fields, " ")
511 }
512 return ""
513 }