resolving_bind.go 1.9 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667
  1. package amneziawgnet
  2. import (
  3. "context"
  4. "fmt"
  5. "net"
  6. "net/netip"
  7. "strconv"
  8. "strings"
  9. "time"
  10. awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
  11. )
  12. // endpointResolveTimeout bounds the one-time DNS lookup in ParseEndpoint.
  13. const endpointResolveTimeout = 5 * time.Second
  14. // resolvingBind lets peer endpoints be hostnames: StdNetBind has no DNS and
  15. // an unresolved name kills the whole IpcSet. Resolved once at configure.
  16. type resolvingBind struct {
  17. awgconn.Bind
  18. }
  19. var lookupEndpointHost = defaultLookupEndpointHost
  20. func defaultLookupEndpointHost(ctx context.Context, host string) ([]netip.Addr, error) {
  21. addrs, err := net.DefaultResolver.LookupNetIP(ctx, "ip", host)
  22. if err != nil {
  23. return nil, err
  24. }
  25. out := make([]netip.Addr, 0, len(addrs))
  26. for _, a := range addrs {
  27. out = append(out, a.Unmap())
  28. }
  29. return out, nil
  30. }
  31. func newResolvingBind() *resolvingBind {
  32. return &resolvingBind{Bind: awgconn.NewDefaultBind()}
  33. }
  34. // ParseEndpoint resolves hostnames before handing the address to amneziawg-go
  35. // (whose own implementation accepts literal IPs only).
  36. func (b *resolvingBind) ParseEndpoint(s string) (awgconn.Endpoint, error) {
  37. host, portStr, err := net.SplitHostPort(strings.TrimSpace(s))
  38. if err != nil {
  39. return nil, fmt.Errorf("endpoint %q: %w", s, err)
  40. }
  41. port64, err := strconv.ParseUint(portStr, 10, 16)
  42. if err != nil || port64 == 0 {
  43. return nil, fmt.Errorf("endpoint %q: bad port", s)
  44. }
  45. addr, err := netip.ParseAddr(host)
  46. if err != nil {
  47. ctx, cancel := context.WithTimeout(context.Background(), endpointResolveTimeout)
  48. defer cancel()
  49. addrs, rerr := lookupEndpointHost(ctx, host)
  50. if rerr != nil {
  51. return nil, fmt.Errorf("endpoint %q: resolve host: %w", s, rerr)
  52. }
  53. if len(addrs) == 0 {
  54. return nil, fmt.Errorf("endpoint %q: host resolved to no addresses", s)
  55. }
  56. addr = addrs[0]
  57. }
  58. return &awgconn.StdNetEndpoint{AddrPort: netip.AddrPortFrom(addr.Unmap(), uint16(port64))}, nil
  59. }