1
0

forwarded_trust_test.go 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205
  1. package sub
  2. import (
  3. "net/http"
  4. "net/http/httptest"
  5. "testing"
  6. "github.com/gin-gonic/gin"
  7. "github.com/mhsanaei/3x-ui/v3/internal/database"
  8. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  9. "github.com/mhsanaei/3x-ui/v3/internal/web/service"
  10. )
  11. func requestFrom(t *testing.T, remoteAddr string, headers map[string]string) *gin.Context {
  12. t.Helper()
  13. req := httptest.NewRequest(http.MethodGet, "/sub/abc", nil)
  14. req.Host = "panel.example.com:2096"
  15. req.RemoteAddr = remoteAddr
  16. for k, v := range headers {
  17. req.Header.Set(k, v)
  18. }
  19. c, _ := gin.CreateTestContext(httptest.NewRecorder())
  20. c.Request = req
  21. return c
  22. }
  23. func setTrustedProxyCIDRs(t *testing.T, value string) {
  24. t.Helper()
  25. if err := database.GetDB().Create(&model.Setting{Key: "trustedProxyCIDRs", Value: value}).Error; err != nil {
  26. t.Fatalf("set trustedProxyCIDRs: %v", err)
  27. }
  28. settingService := service.SettingService{}
  29. stored, err := settingService.GetTrustedProxyCIDRs()
  30. if err != nil {
  31. t.Fatalf("read trustedProxyCIDRs back through SettingService: %v", err)
  32. }
  33. if stored != value {
  34. t.Fatalf("SettingService reads trustedProxyCIDRs as %q, want %q — the key this helper writes has drifted", stored, value)
  35. }
  36. }
  37. func storedAs(value string) *string {
  38. return &value
  39. }
  40. func TestResolveRequest_ForwardedHeaderTrust(t *testing.T) {
  41. tests := []struct {
  42. name string
  43. stored *string
  44. remoteAddr string
  45. wantScheme string
  46. wantHost string
  47. wantHostWithPort string
  48. wantHostHeader string
  49. }{
  50. {
  51. name: "no stored row keeps trusting forwarded headers",
  52. stored: nil,
  53. remoteAddr: "203.0.113.9:51000",
  54. wantScheme: "https",
  55. wantHost: "sub.example.net",
  56. wantHostWithPort: "sub.example.net",
  57. wantHostHeader: "sub.example.net",
  58. },
  59. {
  60. name: "empty stored value keeps trusting forwarded headers",
  61. stored: storedAs(""),
  62. remoteAddr: "203.0.113.9:51000",
  63. wantScheme: "https",
  64. wantHost: "sub.example.net",
  65. wantHostWithPort: "sub.example.net",
  66. wantHostHeader: "sub.example.net",
  67. },
  68. {
  69. name: "stored shipped default keeps trusting forwarded headers",
  70. stored: storedAs(service.DefaultTrustedProxyCIDRs),
  71. remoteAddr: "203.0.113.9:51000",
  72. wantScheme: "https",
  73. wantHost: "sub.example.net",
  74. wantHostWithPort: "sub.example.net",
  75. wantHostHeader: "sub.example.net",
  76. },
  77. {
  78. name: "declared boundary ignores an origin outside it",
  79. stored: storedAs("10.0.0.0/8"),
  80. remoteAddr: "203.0.113.9:51000",
  81. wantScheme: "http",
  82. wantHost: "panel.example.com",
  83. wantHostWithPort: "panel.example.com:2096",
  84. wantHostHeader: "panel.example.com",
  85. },
  86. {
  87. name: "declared boundary trusts an origin inside it",
  88. stored: storedAs("10.0.0.0/8"),
  89. remoteAddr: "10.1.2.3:44000",
  90. wantScheme: "https",
  91. wantHost: "sub.example.net",
  92. wantHostWithPort: "sub.example.net",
  93. wantHostHeader: "sub.example.net",
  94. },
  95. {
  96. name: "declared boundary ignores an unparsable origin",
  97. stored: storedAs("10.0.0.0/8"),
  98. remoteAddr: "not-an-ip",
  99. wantScheme: "http",
  100. wantHost: "panel.example.com",
  101. wantHostWithPort: "panel.example.com:2096",
  102. wantHostHeader: "panel.example.com",
  103. },
  104. }
  105. for _, tc := range tests {
  106. t.Run(tc.name, func(t *testing.T) {
  107. initSubDB(t)
  108. if tc.stored != nil {
  109. setTrustedProxyCIDRs(t, *tc.stored)
  110. }
  111. s := &SubService{}
  112. c := requestFrom(t, tc.remoteAddr, map[string]string{
  113. "X-Forwarded-Host": "sub.example.net",
  114. "X-Forwarded-Proto": "https",
  115. })
  116. scheme, host, hostWithPort, hostHeader := s.ResolveRequest(c)
  117. if scheme != tc.wantScheme {
  118. t.Errorf("scheme = %q, want %q", scheme, tc.wantScheme)
  119. }
  120. if host != tc.wantHost {
  121. t.Errorf("host = %q, want %q", host, tc.wantHost)
  122. }
  123. if hostWithPort != tc.wantHostWithPort {
  124. t.Errorf("hostWithPort = %q, want %q", hostWithPort, tc.wantHostWithPort)
  125. }
  126. if hostHeader != tc.wantHostHeader {
  127. t.Errorf("hostHeader = %q, want %q", hostHeader, tc.wantHostHeader)
  128. }
  129. })
  130. }
  131. }
  132. func TestResolveRequest_GatesRealIPFallback(t *testing.T) {
  133. initSubDB(t)
  134. setTrustedProxyCIDRs(t, "10.0.0.0/8")
  135. s := &SubService{}
  136. c := requestFrom(t, "203.0.113.9:51000", map[string]string{
  137. "X-Real-IP": "198.51.100.7",
  138. })
  139. _, host, _, hostHeader := s.ResolveRequest(c)
  140. if host != "panel.example.com" {
  141. t.Errorf("host = %q, want the request host — X-Real-IP from an untrusted origin must be ignored", host)
  142. }
  143. if hostHeader != "panel.example.com" {
  144. t.Errorf("hostHeader = %q, want the request host", hostHeader)
  145. }
  146. }
  147. func TestHasForwardedHeaders(t *testing.T) {
  148. tests := []struct {
  149. name string
  150. headers map[string]string
  151. want bool
  152. }{
  153. {name: "no forwarded headers", want: false},
  154. {name: "forwarded host", headers: map[string]string{"X-Forwarded-Host": "sub.example.net"}, want: true},
  155. {name: "forwarded proto", headers: map[string]string{"X-Forwarded-Proto": "https"}, want: true},
  156. {name: "real ip", headers: map[string]string{"X-Real-IP": "10.1.2.3"}, want: true},
  157. }
  158. for _, tc := range tests {
  159. t.Run(tc.name, func(t *testing.T) {
  160. if got := hasForwardedHeaders(requestFrom(t, "10.1.2.3:1234", tc.headers)); got != tc.want {
  161. t.Errorf("hasForwardedHeaders() = %v, want %v", got, tc.want)
  162. }
  163. })
  164. }
  165. }
  166. func TestRemoteAddrInCIDRs(t *testing.T) {
  167. tests := []struct {
  168. name string
  169. remoteAddr string
  170. cidrs string
  171. want bool
  172. }{
  173. {name: "inside cidr", remoteAddr: "10.1.2.3:1234", cidrs: "10.0.0.0/8", want: true},
  174. {name: "ipv4 mapped address inside cidr", remoteAddr: "[::ffff:10.1.2.3]:1234", cidrs: "10.0.0.0/8", want: true},
  175. {name: "outside cidr", remoteAddr: "203.0.113.9:1234", cidrs: "10.0.0.0/8", want: false},
  176. {name: "bare address entry", remoteAddr: "192.168.1.5:80", cidrs: "192.168.1.5", want: true},
  177. {name: "ipv6 loopback", remoteAddr: "[::1]:8080", cidrs: "::1/128", want: true},
  178. {name: "no port", remoteAddr: "10.1.2.3", cidrs: "10.0.0.0/8", want: true},
  179. {name: "unparsable origin", remoteAddr: "not-an-ip", cidrs: "10.0.0.0/8", want: false},
  180. {name: "empty entries skipped", remoteAddr: "10.1.2.3:1", cidrs: " , 10.0.0.0/8 , ", want: true},
  181. }
  182. for _, tc := range tests {
  183. t.Run(tc.name, func(t *testing.T) {
  184. if got := remoteAddrInCIDRs(tc.remoteAddr, tc.cidrs); got != tc.want {
  185. t.Errorf("remoteAddrInCIDRs(%q, %q) = %v, want %v", tc.remoteAddr, tc.cidrs, got, tc.want)
  186. }
  187. })
  188. }
  189. }