node_mtls_test.go 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138
  1. package service
  2. import (
  3. "crypto/tls"
  4. "crypto/x509"
  5. "encoding/pem"
  6. "testing"
  7. "github.com/go-playground/validator/v10"
  8. "github.com/mhsanaei/3x-ui/v3/internal/database"
  9. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  10. "github.com/mhsanaei/3x-ui/v3/internal/web/runtime"
  11. )
  12. func TestReloadMasterMtlsClientDoesNotMintMissingCredential(t *testing.T) {
  13. _ = setupSettingMtlsDB(t)
  14. runtime.SetMasterClientCertProvider(func() (tls.Certificate, error) {
  15. pair, err := (&SettingService{}).EnsureMasterClientCert()
  16. if err != nil {
  17. return tls.Certificate{}, err
  18. }
  19. return tls.X509KeyPair(pair.CertPEM, pair.KeyPEM)
  20. })
  21. t.Cleanup(func() { runtime.SetMasterClientCertProvider(nil) })
  22. if err := (&NodeService{}).ReloadMasterMtlsClient(); err == nil {
  23. t.Fatal("reload on a fresh database unexpectedly succeeded")
  24. }
  25. var count int64
  26. keys := []string{settingNodeMtlsCaCert, settingNodeMtlsCaKey, settingNodeMtlsClientCert, settingNodeMtlsClientKey}
  27. if err := database.GetDB().Model(&model.Setting{}).Where("key IN ?", keys).Count(&count).Error; err != nil {
  28. t.Fatalf("count mTLS settings: %v", err)
  29. }
  30. if count != 0 {
  31. t.Fatalf("reload created %d mTLS setting rows, want 0", count)
  32. }
  33. }
  34. func TestNormalizeKeepsMtls(t *testing.T) {
  35. s := &NodeService{}
  36. cases := []struct {
  37. name string
  38. in model.Node
  39. wantMode string
  40. wantErr bool
  41. }{
  42. {"mtls over https preserved", model.Node{Name: "n", Address: "node.example.com", Port: 2053, Scheme: "https", TlsVerifyMode: "mtls"}, "mtls", false},
  43. {"mtls over http rejected", model.Node{Name: "n", Address: "node.example.com", Port: 2053, Scheme: "http", TlsVerifyMode: "mtls"}, "", true},
  44. {"unknown mode clamped to verify", model.Node{Name: "n", Address: "node.example.com", Port: 2053, Scheme: "https", TlsVerifyMode: "bogus"}, "verify", false},
  45. }
  46. for _, c := range cases {
  47. t.Run(c.name, func(t *testing.T) {
  48. n := c.in
  49. err := s.normalize(&n)
  50. if c.wantErr {
  51. if err == nil {
  52. t.Fatal("expected an error")
  53. }
  54. return
  55. }
  56. if err != nil {
  57. t.Fatalf("normalize: %v", err)
  58. }
  59. if n.TlsVerifyMode != c.wantMode {
  60. t.Fatalf("TlsVerifyMode = %q, want %q", n.TlsVerifyMode, c.wantMode)
  61. }
  62. })
  63. }
  64. }
  65. func TestNodeTlsVerifyModeValidatorAcceptsMtls(t *testing.T) {
  66. v := validator.New(validator.WithRequiredStructEnabled())
  67. base := model.Node{Name: "n", Address: "node.example.com", Port: 2053, Scheme: "https", ApiToken: "t"}
  68. for _, m := range []string{"verify", "skip", "pin", "mtls"} {
  69. n := base
  70. n.TlsVerifyMode = m
  71. if err := v.Struct(n); err != nil {
  72. t.Fatalf("validator rejected valid TlsVerifyMode %q: %v", m, err)
  73. }
  74. }
  75. bad := base
  76. bad.TlsVerifyMode = "bogus"
  77. if err := v.Struct(bad); err == nil {
  78. t.Fatal("validator must reject an unknown TlsVerifyMode")
  79. }
  80. }
  81. func TestNodeMtlsCaCert(t *testing.T) {
  82. _ = setupSettingMtlsDB(t)
  83. got, err := (&NodeService{}).NodeMtlsCaCert()
  84. if err != nil {
  85. t.Fatalf("NodeMtlsCaCert: %v", err)
  86. }
  87. block, _ := pem.Decode([]byte(got))
  88. if block == nil || block.Type != "CERTIFICATE" {
  89. t.Fatalf("NodeMtlsCaCert must return a CERTIFICATE PEM, got %q", got)
  90. }
  91. cert, err := x509.ParseCertificate(block.Bytes)
  92. if err != nil {
  93. t.Fatalf("parse returned cert: %v", err)
  94. }
  95. if !cert.IsCA {
  96. t.Fatal("NodeMtlsCaCert must return the CA certificate (IsCA)")
  97. }
  98. }
  99. func TestSetNodeMtlsTrustCA(t *testing.T) {
  100. _ = setupSettingMtlsDB(t)
  101. ns := &NodeService{}
  102. settings := SettingService{}
  103. ca, err := settings.EnsureNodeMtlsCA()
  104. if err != nil {
  105. t.Fatalf("EnsureNodeMtlsCA: %v", err)
  106. }
  107. if err := ns.SetNodeMtlsTrustCA(string(ca.CertPEM)); err != nil {
  108. t.Fatalf("SetNodeMtlsTrustCA(valid): %v", err)
  109. }
  110. pool, err := settings.NodeMtlsClientCAPool()
  111. if err != nil || pool == nil {
  112. t.Fatalf("valid trust CA must persist + build a pool: pool=%v err=%v", pool, err)
  113. }
  114. if err := ns.SetNodeMtlsTrustCA("not a certificate"); err == nil {
  115. t.Fatal("invalid PEM must be rejected (fail closed)")
  116. }
  117. if err := ns.SetNodeMtlsTrustCA(""); err != nil {
  118. t.Fatalf("clearing the trust CA must be allowed: %v", err)
  119. }
  120. pool, _ = settings.NodeMtlsClientCAPool()
  121. if pool != nil {
  122. t.Fatal("cleared trust CA must yield a nil pool (mTLS off)")
  123. }
  124. }