nodetoken_test.go 8.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275
  1. package nodetoken
  2. import (
  3. "encoding/base64"
  4. "encoding/json"
  5. "fmt"
  6. "os"
  7. "path/filepath"
  8. "runtime"
  9. "strings"
  10. "testing"
  11. )
  12. func testRing(t *testing.T, activeID string, ids ...string) *Keyring {
  13. t.Helper()
  14. kr := &Keyring{ActiveID: activeID, Keys: map[string][keyLen]byte{}}
  15. for _, id := range ids {
  16. var k [keyLen]byte
  17. for i := range k {
  18. k[i] = byte(i) + id[len(id)-1] // deterministic and distinct for k1/k2
  19. }
  20. kr.Keys[id] = k
  21. }
  22. return kr
  23. }
  24. func TestRoundTrip(t *testing.T) {
  25. c, err := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  26. if err != nil {
  27. t.Fatal(err)
  28. }
  29. enc, err := c.Encrypt(7, "s3cret-token")
  30. if err != nil {
  31. t.Fatal(err)
  32. }
  33. if !IsEncrypted(enc) || !strings.HasPrefix(enc, "enc:v1:k1:") {
  34. t.Fatalf("unexpected ciphertext form: %q", enc)
  35. }
  36. pt, err := c.Decrypt(7, enc)
  37. if err != nil {
  38. t.Fatal(err)
  39. }
  40. if pt != "s3cret-token" {
  41. t.Fatalf("round-trip mismatch: %q", pt)
  42. }
  43. }
  44. func TestAADBindsToNode(t *testing.T) {
  45. c, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  46. enc, _ := c.Encrypt(7, "tok")
  47. // Decrypting under a different node id must fail (ciphertext bound to row).
  48. if _, err := c.Decrypt(8, enc); err == nil {
  49. t.Fatal("expected AAD mismatch error decrypting under wrong node id")
  50. } else if !strings.Contains(err.Error(), "node 8") || !strings.Contains(err.Error(), "authentication failed") {
  51. t.Fatalf("wrong-node decrypt error: %v", err)
  52. }
  53. }
  54. func TestAADBindsSettingsApartFromNodes(t *testing.T) {
  55. c, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  56. enc, err := c.EncryptBound([]byte("settings/pia_token"), "tok")
  57. if err != nil {
  58. t.Fatal(err)
  59. }
  60. if _, err := c.Decrypt(1, enc); err == nil {
  61. t.Fatal("settings/pia_token ciphertext must not decrypt under nodes/api_token/1")
  62. }
  63. pt, err := c.DecryptBound([]byte("settings/pia_token"), enc)
  64. if err != nil || pt != "tok" {
  65. t.Fatalf("pia AAD round-trip: %q err=%v", pt, err)
  66. }
  67. }
  68. func TestNonceIsRandom(t *testing.T) {
  69. c, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  70. a, _ := c.Encrypt(1, "same")
  71. b, _ := c.Encrypt(1, "same")
  72. if a == b {
  73. t.Fatal("two encryptions of the same value produced identical ciphertext (nonce reuse)")
  74. }
  75. }
  76. func TestPlaintextPassThrough(t *testing.T) {
  77. // ModeOff: encrypt is a no-op, decrypt returns plaintext.
  78. c, _ := NewCodec(ModeOff, nil)
  79. enc, err := c.Encrypt(1, "plain")
  80. if err != nil || enc != "plain" {
  81. t.Fatalf("off-mode encrypt should be no-op, got %q err=%v", enc, err)
  82. }
  83. // A legacy plaintext row decrypts (passes through) in any mode.
  84. c2, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  85. if pt, err := c2.Decrypt(1, "legacy-plain"); err != nil || pt != "legacy-plain" {
  86. t.Fatalf("legacy plaintext should pass through, got %q err=%v", pt, err)
  87. }
  88. }
  89. func TestEncryptedNeverFallsBackToPlaintext(t *testing.T) {
  90. c, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  91. enc, _ := c.Encrypt(1, "tok")
  92. // Corrupt the ciphertext body — must error, never return raw bytes.
  93. bad := flipLastCiphertextBit(t, enc)
  94. if _, err := c.Decrypt(1, bad); err == nil {
  95. t.Fatal("corrupted ciphertext must fail, not fall back to plaintext")
  96. }
  97. // Unknown key id must error.
  98. c2, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  99. other := strings.Replace(enc, "enc:v1:k1:", "enc:v1:zz:", 1)
  100. if _, err := c2.Decrypt(1, other); err == nil {
  101. t.Fatal("unknown key id must fail")
  102. }
  103. }
  104. // flipLastCiphertextBit rewrites the body through its decoded bytes, because
  105. // editing the trailing base64 characters can leave those bytes untouched.
  106. func flipLastCiphertextBit(t *testing.T, stored string) string {
  107. t.Helper()
  108. cut := strings.LastIndex(stored, ":") + 1
  109. blob, err := base64.RawURLEncoding.DecodeString(stored[cut:])
  110. if err != nil {
  111. t.Fatalf("decode ciphertext body: %v", err)
  112. }
  113. blob[len(blob)-1] ^= 0x01
  114. return stored[:cut] + base64.RawURLEncoding.EncodeToString(blob)
  115. }
  116. func TestEncryptionMarkerPassesThroughWhenDisabled(t *testing.T) {
  117. c, _ := NewCodec(ModeOff, nil)
  118. stored := "enc:v1:not-ciphertext"
  119. if got, err := c.Decrypt(1, stored); err != nil || got != stored {
  120. t.Fatalf("off-mode changed a legacy token: got %q err=%v", got, err)
  121. }
  122. }
  123. func TestParseKeyringRejectsDelimiterInKeyID(t *testing.T) {
  124. key := base64.StdEncoding.EncodeToString(make([]byte, keyLen))
  125. for _, tc := range []struct {
  126. name, active string
  127. keys map[string]string
  128. }{
  129. {"active delimiter", "region:k1", map[string]string{"region:k1": key}},
  130. {"key delimiter", "k1", map[string]string{"k1": key, "old:k0": key}},
  131. {"empty key", "k1", map[string]string{"k1": key, "": key}},
  132. } {
  133. t.Run(tc.name, func(t *testing.T) {
  134. if _, err := parseKeyring(tc.active, tc.keys); err == nil {
  135. t.Fatal("invalid key id was accepted")
  136. }
  137. })
  138. }
  139. }
  140. func TestEncryptRoundTripSafe(t *testing.T) {
  141. // Re-submitting stored ciphertext (UI round-trip) must not double-encrypt.
  142. c, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  143. enc, _ := c.Encrypt(5, "tok")
  144. again, err := c.Encrypt(5, enc)
  145. if err != nil {
  146. t.Fatal(err)
  147. }
  148. if again != enc {
  149. t.Fatal("re-encrypting stored ciphertext changed it (double-encrypt)")
  150. }
  151. if pt, _ := c.Decrypt(5, again); pt != "tok" {
  152. t.Fatalf("round-trip-safe encrypt corrupted token: %q", pt)
  153. }
  154. }
  155. func TestEmptyTokenNeverEncrypted(t *testing.T) {
  156. c, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1"))
  157. if v, _ := c.Encrypt(1, ""); v != "" {
  158. t.Fatalf("empty token must stay empty, got %q", v)
  159. }
  160. }
  161. func TestRotation(t *testing.T) {
  162. // k2 active, k1 retained. Old-key value still decrypts; new writes use k2.
  163. ring := testRing(t, "k2", "k1", "k2")
  164. if ring.Keys["k1"] == ring.Keys["k2"] {
  165. t.Fatal("rotation fixture keys k1 and k2 are identical")
  166. }
  167. c, _ := NewCodec(ModeRequired, ring)
  168. // produce a k1 value via a codec whose active is k1
  169. c1, _ := NewCodec(ModeRequired, testRing(t, "k1", "k1", "k2"))
  170. old, _ := c1.Encrypt(3, "tok")
  171. if pt, err := c.Decrypt(3, old); err != nil || pt != "tok" {
  172. t.Fatalf("retained old key must decrypt, got %q err=%v", pt, err)
  173. }
  174. if c.EncryptedWithActive(old) {
  175. t.Fatal("k1 value should not count as encrypted-with-active(k2)")
  176. }
  177. neu, _ := c.Encrypt(3, "tok")
  178. if !c.EncryptedWithActive(neu) {
  179. t.Fatal("new write should be encrypted with active key")
  180. }
  181. }
  182. func TestNewCodecRequiresKey(t *testing.T) {
  183. if _, err := NewCodec(ModeRequired, nil); err == nil {
  184. t.Fatal("required mode without a key must fail (fail-closed)")
  185. }
  186. if _, err := NewCodec(ModeMigration, &Keyring{ActiveID: "x", Keys: nil}); err == nil {
  187. t.Fatal("migration mode with empty keyring must fail")
  188. }
  189. }
  190. func TestParseMode(t *testing.T) {
  191. for in, want := range map[string]Mode{"": ModeOff, "off": ModeOff, "Migration": ModeMigration, "REQUIRED": ModeRequired} {
  192. if m, err := ParseMode(in); err != nil || m != want {
  193. t.Fatalf("ParseMode(%q)=%v err=%v, want %v", in, m, err, want)
  194. }
  195. }
  196. if _, err := ParseMode("bogus"); err == nil {
  197. t.Fatal("unknown mode must error")
  198. }
  199. }
  200. // writeKeyFile writes a one-key keyring and chmods it, since WriteFile's mode
  201. // passes through the umask.
  202. func writeKeyFile(t *testing.T, mode os.FileMode) string {
  203. t.Helper()
  204. p := filepath.Join(t.TempDir(), "k.json")
  205. key := make([]byte, keyLen)
  206. body, _ := json.Marshal(keyFile{Active: "k1", Keys: map[string]string{"k1": base64.StdEncoding.EncodeToString(key)}})
  207. if err := os.WriteFile(p, body, mode); err != nil {
  208. t.Fatal(err)
  209. }
  210. if err := os.Chmod(p, mode); err != nil {
  211. t.Fatal(err)
  212. }
  213. return p
  214. }
  215. // Windows reports every writable file as 0666, so a mode check there refused
  216. // every key file, an owner-only one included.
  217. func TestFileKeySourceLoadsOwnerOnlyKeyFile(t *testing.T) {
  218. kr, err := (FileKeySource{Path: writeKeyFile(t, 0o600)}).Load()
  219. if err != nil {
  220. t.Fatalf("0600 key file should load: %v", err)
  221. }
  222. if kr.ActiveID != "k1" || len(kr.Keys) != 1 {
  223. t.Fatalf("unexpected keyring %+v", kr)
  224. }
  225. }
  226. func TestFileKeySourceRejectsLoosePerms(t *testing.T) {
  227. if runtime.GOOS == "windows" {
  228. t.Skip("POSIX permission bits are not meaningful on Windows")
  229. }
  230. p := writeKeyFile(t, 0o644)
  231. _, err := (FileKeySource{Path: p}).Load()
  232. want := fmt.Sprintf("nodetoken: key file %s has insecure mode 0644 (want 0600)", p)
  233. if err == nil || err.Error() != want {
  234. t.Fatalf("Load() error = %v, want %q", err, want)
  235. }
  236. }
  237. func TestEnvKeySource(t *testing.T) {
  238. key := make([]byte, keyLen)
  239. for i := range key {
  240. key[i] = byte(i)
  241. }
  242. t.Setenv("XUI_NODE_TOKEN_KEY_TEST", base64.StdEncoding.EncodeToString(key))
  243. kr, err := (EnvKeySource{Var: "XUI_NODE_TOKEN_KEY_TEST"}).Load()
  244. if err != nil {
  245. t.Fatal(err)
  246. }
  247. if kr.ActiveID != "env" {
  248. t.Fatalf("env key id should be 'env', got %q", kr.ActiveID)
  249. }
  250. c, _ := NewCodec(ModeRequired, kr)
  251. enc, _ := c.Encrypt(1, "x")
  252. if pt, _ := c.Decrypt(1, enc); pt != "x" {
  253. t.Fatal("env-sourced key failed round trip")
  254. }
  255. }