resolving_bind_test.go 3.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102
  1. package amneziawgnet
  2. import (
  3. "context"
  4. "errors"
  5. "net/netip"
  6. "testing"
  7. awgconn "github.com/amnezia-vpn/amneziawg-go/v3/conn"
  8. )
  9. func mustResolvingBind(t *testing.T) *resolvingBind {
  10. t.Helper()
  11. return newResolvingBind("")
  12. }
  13. // endpointAddrPort reads any bind's endpoint; Windows' default bind has its own type.
  14. func endpointAddrPort(t *testing.T, ep awgconn.Endpoint) netip.AddrPort {
  15. t.Helper()
  16. ap, err := netip.ParseAddrPort(ep.DstToString())
  17. if err != nil {
  18. t.Fatalf("endpoint %q: %v", ep.DstToString(), err)
  19. }
  20. return ap
  21. }
  22. // ownEndpointBind accepts only endpoints it parsed itself, as WinRingBind does.
  23. type ownEndpointBind struct{ awgconn.Bind }
  24. type ownEndpoint struct{ awgconn.StdNetEndpoint }
  25. func (ownEndpointBind) ParseEndpoint(s string) (awgconn.Endpoint, error) {
  26. ap, err := netip.ParseAddrPort(s)
  27. if err != nil {
  28. return nil, err
  29. }
  30. return &ownEndpoint{awgconn.StdNetEndpoint{AddrPort: ap}}, nil
  31. }
  32. // WinRingBind, the default bind on Windows, refuses to send to an endpoint of any
  33. // other type, so a hand-built StdNetEndpoint killed every handshake there.
  34. func TestResolvingBind_ParseEndpointComesFromTheWrappedBind(t *testing.T) {
  35. ep, err := (&resolvingBind{Bind: ownEndpointBind{}}).ParseEndpoint("203.0.113.7:51820")
  36. if err != nil {
  37. t.Fatalf("ParseEndpoint: %v", err)
  38. }
  39. if _, ok := ep.(*ownEndpoint); !ok {
  40. t.Fatalf("endpoint is %T, not the wrapped bind's own type", ep)
  41. }
  42. }
  43. func TestResolvingBind_ParseEndpointIPLiteral(t *testing.T) {
  44. b := mustResolvingBind(t)
  45. ep, err := b.ParseEndpoint("203.0.113.7:51820")
  46. if err != nil {
  47. t.Fatalf("IP endpoint rejected: %v", err)
  48. }
  49. got := endpointAddrPort(t, ep)
  50. if got.Addr().String() != "203.0.113.7" || got.Port() != 51820 {
  51. t.Fatalf("endpoint = %v, want 203.0.113.7:51820", got)
  52. }
  53. }
  54. func TestResolvingBind_ParseEndpointHostnameResolves(t *testing.T) {
  55. orig := lookupEndpointHost
  56. lookupEndpointHost = func(ctx context.Context, host string) ([]netip.Addr, error) {
  57. if host != "peer.example.test" {
  58. t.Errorf("unexpected lookup host %q", host)
  59. }
  60. return []netip.Addr{netip.MustParseAddr("198.51.100.9")}, nil
  61. }
  62. defer func() { lookupEndpointHost = orig }()
  63. b := mustResolvingBind(t)
  64. ep, err := b.ParseEndpoint("peer.example.test:443")
  65. if err != nil {
  66. t.Fatalf("hostname endpoint rejected: %v", err)
  67. }
  68. if got := endpointAddrPort(t, ep); got.Addr().String() != "198.51.100.9" || got.Port() != 443 {
  69. t.Fatalf("endpoint = %v, want 198.51.100.9:443", got)
  70. }
  71. }
  72. func TestResolvingBind_ParseEndpointResolveFailureIsAnError(t *testing.T) {
  73. orig := lookupEndpointHost
  74. lookupEndpointHost = func(ctx context.Context, host string) ([]netip.Addr, error) {
  75. return nil, errors.New("no such host")
  76. }
  77. defer func() { lookupEndpointHost = orig }()
  78. b := mustResolvingBind(t)
  79. if _, err := b.ParseEndpoint("missing.example.test:80"); err == nil {
  80. t.Fatal("expected resolve failure to surface as an error")
  81. }
  82. }
  83. func TestResolvingBind_ParseEndpointBadPortRejected(t *testing.T) {
  84. b := mustResolvingBind(t)
  85. if _, err := b.ParseEndpoint("203.0.113.7:none"); err == nil {
  86. t.Fatal("expected bad port to be rejected")
  87. }
  88. }