setting_mtls_test.go 7.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274
  1. package service
  2. import (
  3. "bytes"
  4. "crypto/x509"
  5. "encoding/pem"
  6. "path/filepath"
  7. "testing"
  8. "time"
  9. "github.com/mhsanaei/3x-ui/v3/internal/database"
  10. "github.com/mhsanaei/3x-ui/v3/internal/util/crypto"
  11. )
  12. func setupSettingMtlsDB(t *testing.T) *SettingService {
  13. t.Helper()
  14. dbDir := t.TempDir()
  15. t.Setenv("XUI_DB_FOLDER", dbDir)
  16. if err := database.InitDB(filepath.Join(dbDir, "x-ui.db")); err != nil {
  17. t.Fatalf("InitDB: %v", err)
  18. }
  19. t.Cleanup(func() { _ = database.CloseDB() })
  20. return &SettingService{}
  21. }
  22. func TestEnsureNodeMtlsCA_Idempotent(t *testing.T) {
  23. s := setupSettingMtlsDB(t)
  24. first, err := s.EnsureNodeMtlsCA()
  25. if err != nil {
  26. t.Fatalf("EnsureNodeMtlsCA (first): %v", err)
  27. }
  28. block, _ := pem.Decode(first.CertPEM)
  29. if block == nil {
  30. t.Fatal("CA cert is not valid PEM")
  31. }
  32. caCert, err := x509.ParseCertificate(block.Bytes)
  33. if err != nil {
  34. t.Fatalf("parse CA cert: %v", err)
  35. }
  36. if !caCert.IsCA {
  37. t.Fatal("stored CA must have IsCA=true")
  38. }
  39. second, err := s.EnsureNodeMtlsCA()
  40. if err != nil {
  41. t.Fatalf("EnsureNodeMtlsCA (second): %v", err)
  42. }
  43. if !bytes.Equal(first.CertPEM, second.CertPEM) || !bytes.Equal(first.KeyPEM, second.KeyPEM) {
  44. t.Fatal("EnsureNodeMtlsCA must be idempotent: second call returned different PEMs")
  45. }
  46. }
  47. func TestEnsureMasterClientCert_VerifiesAndIdempotent(t *testing.T) {
  48. s := setupSettingMtlsDB(t)
  49. ca, err := s.EnsureNodeMtlsCA()
  50. if err != nil {
  51. t.Fatalf("EnsureNodeMtlsCA: %v", err)
  52. }
  53. client, err := s.EnsureMasterClientCert()
  54. if err != nil {
  55. t.Fatalf("EnsureMasterClientCert: %v", err)
  56. }
  57. cblock, _ := pem.Decode(client.CertPEM)
  58. if cblock == nil {
  59. t.Fatal("client cert is not valid PEM")
  60. }
  61. leaf, err := x509.ParseCertificate(cblock.Bytes)
  62. if err != nil {
  63. t.Fatalf("parse client cert: %v", err)
  64. }
  65. caBlock, _ := pem.Decode(ca.CertPEM)
  66. roots := x509.NewCertPool()
  67. roots.AddCert(mustParse(t, caBlock.Bytes))
  68. if _, err := leaf.Verify(x509.VerifyOptions{
  69. Roots: roots,
  70. KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth},
  71. }); err != nil {
  72. t.Fatalf("master client cert must verify against the node CA for client auth: %v", err)
  73. }
  74. again, err := s.EnsureMasterClientCert()
  75. if err != nil {
  76. t.Fatalf("EnsureMasterClientCert (second): %v", err)
  77. }
  78. if !bytes.Equal(client.CertPEM, again.CertPEM) || !bytes.Equal(client.KeyPEM, again.KeyPEM) {
  79. t.Fatal("EnsureMasterClientCert must be idempotent")
  80. }
  81. }
  82. func TestEnsureMasterClientCertRejectsMismatchedStoredKey(t *testing.T) {
  83. s := setupSettingMtlsDB(t)
  84. first, err := s.EnsureMasterClientCert()
  85. if err != nil {
  86. t.Fatal(err)
  87. }
  88. ca, err := s.EnsureNodeMtlsCA()
  89. if err != nil {
  90. t.Fatal(err)
  91. }
  92. other, err := crypto.IssueClientCert(ca, "other master")
  93. if err != nil {
  94. t.Fatal(err)
  95. }
  96. if err := s.setString(settingNodeMtlsClientCert, string(first.CertPEM)); err != nil {
  97. t.Fatal(err)
  98. }
  99. if err := s.setString(settingNodeMtlsClientKey, string(other.KeyPEM)); err != nil {
  100. t.Fatal(err)
  101. }
  102. if _, err := s.EnsureMasterClientCert(); err == nil {
  103. t.Fatal("mismatched stored certificate and key were accepted")
  104. }
  105. }
  106. func TestEnsureMasterClientCert_ReissuesLeafWhenCAStillExists(t *testing.T) {
  107. s := setupSettingMtlsDB(t)
  108. client, err := s.EnsureMasterClientCert()
  109. if err != nil {
  110. t.Fatalf("EnsureMasterClientCert: %v", err)
  111. }
  112. pin, err := clientCertSHA256FromPEM(client.CertPEM)
  113. if err != nil {
  114. t.Fatalf("clientCertSHA256FromPEM: %v", err)
  115. }
  116. if err := s.setString(settingNodeMtlsClientPin, pin); err != nil {
  117. t.Fatalf("persist client pin: %v", err)
  118. }
  119. if err := s.setString(settingNodeMtlsClientCert, ""); err != nil {
  120. t.Fatalf("clear client cert: %v", err)
  121. }
  122. if err := s.setString(settingNodeMtlsClientKey, ""); err != nil {
  123. t.Fatalf("clear client key: %v", err)
  124. }
  125. reissued, err := s.EnsureMasterClientCert()
  126. if err != nil {
  127. t.Fatalf("reissue with surviving CA: %v", err)
  128. }
  129. newPin, err := clientCertSHA256FromPEM(reissued.CertPEM)
  130. if err != nil {
  131. t.Fatalf("new pin: %v", err)
  132. }
  133. if newPin == pin {
  134. t.Fatal("reissued credential kept the lost leaf identity")
  135. }
  136. stored, err := s.getString(settingNodeMtlsClientPin)
  137. if err != nil || stored != newPin {
  138. t.Fatalf("stored pin = %q, error = %v, want %q", stored, err, newPin)
  139. }
  140. }
  141. func TestEnsureMasterClientCertConcurrentFirstUseMintsOneCredential(t *testing.T) {
  142. s := setupSettingMtlsDB(t)
  143. blocked := make(chan struct{})
  144. masterClientCredentialMu.Lock()
  145. go func() {
  146. _, _ = s.EnsureMasterClientCert()
  147. close(blocked)
  148. }()
  149. select {
  150. case <-blocked:
  151. masterClientCredentialMu.Unlock()
  152. t.Fatal("EnsureMasterClientCert returned while its serialization lock was held")
  153. case <-time.After(50 * time.Millisecond):
  154. }
  155. masterClientCredentialMu.Unlock()
  156. select {
  157. case <-blocked:
  158. case <-time.After(time.Second):
  159. t.Fatal("EnsureMasterClientCert remained blocked after serialization lock release")
  160. }
  161. const callers = 8
  162. start := make(chan struct{})
  163. results := make(chan crypto.CertKeyPEM, callers)
  164. errs := make(chan error, callers)
  165. for range callers {
  166. go func() {
  167. <-start
  168. credential, err := s.EnsureMasterClientCert()
  169. results <- credential
  170. errs <- err
  171. }()
  172. }
  173. close(start)
  174. var first crypto.CertKeyPEM
  175. for i := 0; i < callers; i++ {
  176. credential := <-results
  177. if err := <-errs; err != nil {
  178. t.Fatalf("caller %d: %v", i, err)
  179. }
  180. if i == 0 {
  181. first = credential
  182. } else if !bytes.Equal(first.CertPEM, credential.CertPEM) || !bytes.Equal(first.KeyPEM, credential.KeyPEM) {
  183. t.Fatalf("caller %d received a different credential", i)
  184. }
  185. }
  186. storedCert, _ := s.getString(settingNodeMtlsClientCert)
  187. storedKey, _ := s.getString(settingNodeMtlsClientKey)
  188. if storedCert != string(first.CertPEM) || storedKey != string(first.KeyPEM) {
  189. t.Fatal("persisted credential differs from concurrent callers")
  190. }
  191. }
  192. func TestEnsureMasterClientCert_PersistsCredentialAtomically(t *testing.T) {
  193. s := setupSettingMtlsDB(t)
  194. db := database.GetDB()
  195. trigger := `CREATE TRIGGER fail_master_pin_insert
  196. BEFORE INSERT ON settings
  197. WHEN NEW.key = 'nodeMtlsClientCertSha256'
  198. BEGIN SELECT RAISE(ABORT, 'injected pin failure'); END`
  199. if err := db.Exec(trigger).Error; err != nil {
  200. t.Fatalf("create failure trigger: %v", err)
  201. }
  202. if _, err := s.EnsureMasterClientCert(); err == nil {
  203. t.Fatal("injected persistence failure unexpectedly succeeded")
  204. }
  205. for _, key := range []string{settingNodeMtlsClientCert, settingNodeMtlsClientKey, settingNodeMtlsClientPin} {
  206. got, err := s.getString(key)
  207. if err != nil {
  208. t.Fatalf("get %s after rollback: %v", key, err)
  209. }
  210. if got != "" {
  211. t.Fatalf("%s persisted despite transaction rollback", key)
  212. }
  213. }
  214. if err := db.Exec("DROP TRIGGER fail_master_pin_insert").Error; err != nil {
  215. t.Fatalf("drop failure trigger: %v", err)
  216. }
  217. if _, err := s.EnsureMasterClientCert(); err != nil {
  218. t.Fatalf("retry after rollback: %v", err)
  219. }
  220. }
  221. func TestNodeMtlsClientCAPool(t *testing.T) {
  222. s := setupSettingMtlsDB(t)
  223. pool, err := s.NodeMtlsClientCAPool()
  224. if err != nil {
  225. t.Fatalf("NodeMtlsClientCAPool (unset): %v", err)
  226. }
  227. if pool != nil {
  228. t.Fatal("with no trust CA configured, the pool must be nil (mTLS off; listener unchanged)")
  229. }
  230. ca, err := s.EnsureNodeMtlsCA()
  231. if err != nil {
  232. t.Fatalf("EnsureNodeMtlsCA: %v", err)
  233. }
  234. if err := s.setString("nodeMtlsClientCAPem", string(ca.CertPEM)); err != nil {
  235. t.Fatalf("set trust CA: %v", err)
  236. }
  237. pool, err = s.NodeMtlsClientCAPool()
  238. if err != nil {
  239. t.Fatalf("NodeMtlsClientCAPool (set): %v", err)
  240. }
  241. if pool == nil {
  242. t.Fatal("with a trust CA configured, the pool must be non-nil")
  243. }
  244. }
  245. func mustParse(t *testing.T, der []byte) *x509.Certificate {
  246. t.Helper()
  247. c, err := x509.ParseCertificate(der)
  248. if err != nil {
  249. t.Fatalf("ParseCertificate: %v", err)
  250. }
  251. return c
  252. }