server_remote_cert_hash_test.go 1.1 KB

1234567891011121314151617181920212223242526272829303132333435363738
  1. package service
  2. import (
  3. "crypto/sha256"
  4. "encoding/hex"
  5. "errors"
  6. "net/http"
  7. "net/http/httptest"
  8. "strings"
  9. "testing"
  10. "github.com/mhsanaei/3x-ui/v3/internal/util/netsafe"
  11. )
  12. func TestGetRemoteCertHashGuardsPrivateTargets(t *testing.T) {
  13. srv := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
  14. defer srv.Close()
  15. target := strings.TrimPrefix(srv.URL, "https://")
  16. sum := sha256.Sum256(srv.Certificate().Raw)
  17. want := hex.EncodeToString(sum[:])
  18. t.Run("loopback refused without opt-in", func(t *testing.T) {
  19. hashes, err := (&ServerService{}).GetRemoteCertHash(target, false)
  20. if !errors.Is(err, netsafe.ErrPrivateAddressBlocked) {
  21. t.Fatalf("GetRemoteCertHash(%s) = %v, %v; want ErrPrivateAddressBlocked", target, hashes, err)
  22. }
  23. })
  24. t.Run("loopback read with opt-in", func(t *testing.T) {
  25. hashes, err := (&ServerService{}).GetRemoteCertHash(target, true)
  26. if err != nil {
  27. t.Fatalf("GetRemoteCertHash(%s, allowPrivate): %v", target, err)
  28. }
  29. if len(hashes) != 1 || hashes[0] != want {
  30. t.Fatalf("hashes = %v, want [%s]", hashes, want)
  31. }
  32. })
  33. }