pia_test.go 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375
  1. package integration
  2. import (
  3. "context"
  4. "encoding/base64"
  5. "encoding/json"
  6. "net/netip"
  7. "path/filepath"
  8. "strconv"
  9. "strings"
  10. "testing"
  11. "time"
  12. "github.com/mhsanaei/3x-ui/v3/internal/crypto/nodetoken"
  13. "github.com/mhsanaei/3x-ui/v3/internal/database"
  14. piaprotocol "github.com/mhsanaei/3x-ui/v3/internal/pia"
  15. )
  16. type fakePiaAuth struct{ token string }
  17. func (f fakePiaAuth) Authenticate(context.Context, string, []byte) (piaprotocol.Token, error) {
  18. return piaprotocol.Token{Value: []byte(f.token), ExpiresAt: time.Now().Add(24 * time.Hour)}, nil
  19. }
  20. type fakePiaCatalog struct{ payload []byte }
  21. func (f fakePiaCatalog) Fetch(context.Context) (piaprotocol.ServerListSnapshot, error) {
  22. return piaprotocol.ServerListSnapshot{Payload: f.payload, SchemaHint: "6", SignatureVerified: true}, nil
  23. }
  24. type fakePiaRegistrar struct {
  25. n int
  26. token string
  27. }
  28. func (f *fakePiaRegistrar) RegisterKey(_ context.Context, server piaprotocol.WireGuardServer, token string, _ string) (piaprotocol.Registration, error) {
  29. f.n++
  30. f.token = token
  31. key := make([]byte, 32)
  32. key[0] = byte(f.n)
  33. return piaprotocol.Registration{
  34. PeerIP: netip.MustParsePrefix("10.8.0." + strconv.Itoa(f.n) + "/32"),
  35. ServerKey: base64.StdEncoding.EncodeToString(key),
  36. ServerIP: server.IP,
  37. ServerPort: 1337,
  38. }, nil
  39. }
  40. func setupPiaService(t *testing.T) *PiaService {
  41. t.Helper()
  42. if err := database.InitDB(filepath.Join(t.TempDir(), "x-ui.db")); err != nil {
  43. t.Fatal(err)
  44. }
  45. t.Cleanup(func() { _ = database.CloseDB() })
  46. payload := []byte(`{"version":6,"groups":{"wg":[{"name":"wireguard","ports":[1337]}]},"regions":[{"id":"us-east","name":"US East","country":"US","geo":false,"offline":false,"port_forward":true,"servers":{"wg":[{"ip":"198.51.100.10","cn":"useast1"},{"ip":"198.51.100.20","cn":"useast2"}]}},{"id":"de-berlin","name":"Berlin","country":"DE","geo":false,"offline":false,"port_forward":false,"servers":{"wg":[{"ip":"203.0.113.10","cn":"berlin1"}]}}]}`)
  47. svc := NewPiaService()
  48. svc.Auth = fakePiaAuth{token: "tokentokentokentoken12"}
  49. svc.Catalog = piaprotocol.NewCatalog(fakePiaCatalog{payload: payload})
  50. svc.Registrar = &fakePiaRegistrar{}
  51. return svc
  52. }
  53. func TestPiaLoginStoresTokenAndHidesItFromData(t *testing.T) {
  54. svc := setupPiaService(t)
  55. view, err := svc.Login("p1234567", "TEST-PIA-PASSWORD-MUST-NOT-LEAK")
  56. if err != nil {
  57. t.Fatal(err)
  58. }
  59. if view.Username != "p1234567" || view.AccountHint != "p1****67" {
  60. t.Fatalf("account view: %+v", view)
  61. }
  62. data, err := svc.GetPiaData()
  63. if err != nil || data == nil || data.AccountHint != "p1****67" {
  64. t.Fatalf("data: %+v err=%v", data, err)
  65. }
  66. raw, _ := json.Marshal(data)
  67. if strings.Contains(string(raw), "TEST-PIA-PASSWORD-MUST-NOT-LEAK") || strings.Contains(string(raw), "tokentokentokentoken12") {
  68. t.Fatalf("secret leaked in data: %s", raw)
  69. }
  70. stored, err := svc.GetPia()
  71. if err != nil || !strings.Contains(stored, "tokentokentokentoken12") {
  72. t.Fatalf("token must be stored in settings: %q err=%v", stored, err)
  73. }
  74. }
  75. func TestPiaCountriesAndServers(t *testing.T) {
  76. svc := setupPiaService(t)
  77. countries, err := svc.GetCountries()
  78. if err != nil {
  79. t.Fatal(err)
  80. }
  81. if len(countries) != 2 || countries[0].Code != "DE" || countries[1].Code != "US" {
  82. t.Fatalf("countries: %+v", countries)
  83. }
  84. servers, err := svc.GetServers("US")
  85. if err != nil {
  86. t.Fatal(err)
  87. }
  88. if len(servers.Regions) != 1 || servers.Regions[0].ID != "us-east" || len(servers.Servers) != 2 {
  89. t.Fatalf("us servers: %+v", servers)
  90. }
  91. if servers.Servers[0].Hostname != "useast1" || servers.Servers[0].RegionID != "us-east" {
  92. t.Fatalf("first server: %+v", servers.Servers[0])
  93. }
  94. }
  95. func TestPiaAddKeyRegistersWireGuardPeer(t *testing.T) {
  96. svc := setupPiaService(t)
  97. if _, err := svc.AddKey("useast1"); err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeTokenRejected {
  98. t.Fatalf("addKey before login: %v", err)
  99. }
  100. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  101. t.Fatal(err)
  102. }
  103. key, err := svc.AddKey("useast1")
  104. if err != nil {
  105. t.Fatal(err)
  106. }
  107. if key.Tag != "pia-us-east-useast1" || key.Hostname != "useast1" || key.SecretKey == "" || key.PublicKey == "" {
  108. t.Fatalf("key: %+v", key)
  109. }
  110. if key.Address != "10.8.0.1/32" || key.Endpoint != "198.51.100.10:1337" {
  111. t.Fatalf("peer: %+v", key)
  112. }
  113. byTag, err := svc.AddKey("pia-us-east-useast1")
  114. if err != nil || byTag.Hostname != "useast1" || byTag.Tag != "pia-us-east-useast1" {
  115. t.Fatalf("addKey by tag: %+v err=%v", byTag, err)
  116. }
  117. if _, err := svc.AddKey("1a"); err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeServerNotFound {
  118. t.Fatalf("truncated hostname must not match: %v", err)
  119. }
  120. }
  121. func TestPiaExpiredTokenNeverReachesRegistrar(t *testing.T) {
  122. svc := setupPiaService(t)
  123. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  124. t.Fatal(err)
  125. }
  126. raw, err := svc.GetPia()
  127. if err != nil {
  128. t.Fatal(err)
  129. }
  130. var stored piaStored
  131. if err := json.Unmarshal([]byte(raw), &stored); err != nil {
  132. t.Fatal(err)
  133. }
  134. stored.TokenExpiresAt = time.Now().Add(-time.Minute).Unix()
  135. rewritten, err := json.Marshal(stored)
  136. if err != nil {
  137. t.Fatal(err)
  138. }
  139. if err := svc.SetPia(string(rewritten)); err != nil {
  140. t.Fatal(err)
  141. }
  142. reg := svc.Registrar.(*fakePiaRegistrar)
  143. before := reg.n
  144. _, err = svc.AddKey("useast1")
  145. if err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeTokenRejected {
  146. t.Fatalf("expired token: %v", err)
  147. }
  148. if piaprotocol.MessageOf(err) != "The PIA token has expired. Sign in again." {
  149. t.Fatalf("expired token message: %v", err)
  150. }
  151. if reg.n != before {
  152. t.Fatalf("expired token reached registrar: calls=%d", reg.n)
  153. }
  154. }
  155. func TestPiaDelClearsAccount(t *testing.T) {
  156. svc := setupPiaService(t)
  157. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  158. t.Fatal(err)
  159. }
  160. if err := svc.DelPiaData(); err != nil {
  161. t.Fatal(err)
  162. }
  163. data, err := svc.GetPiaData()
  164. if err != nil || data != nil {
  165. t.Fatalf("want nil data after logout, got %+v err=%v", data, err)
  166. }
  167. }
  168. func TestPiaOutboundTag(t *testing.T) {
  169. tests := []struct {
  170. region, host, want string
  171. }{
  172. {"us-east", "useast1", "pia-us-east-useast1"},
  173. {"US-East", "useast401.privacy.network", "pia-us-east-useast401"},
  174. {"us_california", "silicon_valley", "pia-us-california-silicon-valley"},
  175. {"", "berlin1", "pia-berlin1"},
  176. }
  177. for _, tt := range tests {
  178. t.Run(tt.want, func(t *testing.T) {
  179. if got := piaOutboundTag(tt.region, tt.host); got != tt.want {
  180. t.Fatalf("piaOutboundTag(%q, %q) = %q, want %q", tt.region, tt.host, got, tt.want)
  181. }
  182. })
  183. }
  184. }
  185. func TestPiaCorruptSettingIsNotTreatedAsLoggedOut(t *testing.T) {
  186. svc := setupPiaService(t)
  187. if err := svc.SetPia(`{"username":`); err != nil {
  188. t.Fatal(err)
  189. }
  190. data, err := svc.GetPiaData()
  191. if err == nil || data != nil {
  192. t.Fatalf("corrupt pia setting must not look logged-out: data=%+v err=%v", data, err)
  193. }
  194. }
  195. func enablePiaTokenEncryption(t *testing.T) {
  196. t.Helper()
  197. var k [32]byte
  198. for i := range k {
  199. k[i] = byte(i + 1)
  200. }
  201. ring := &nodetoken.Keyring{ActiveID: "t1", Keys: map[string][32]byte{"t1": k}}
  202. codec, err := nodetoken.NewCodec(nodetoken.ModeRequired, ring)
  203. if err != nil {
  204. t.Fatalf("new codec: %v", err)
  205. }
  206. nodetoken.Init(codec)
  207. t.Cleanup(func() {
  208. off, _ := nodetoken.NewCodec(nodetoken.ModeOff, nil)
  209. nodetoken.Init(off)
  210. })
  211. }
  212. func TestPiaLoginEncryptsTokenWhenRequired(t *testing.T) {
  213. svc := setupPiaService(t)
  214. enablePiaTokenEncryption(t)
  215. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  216. t.Fatal(err)
  217. }
  218. stored, err := svc.GetPia()
  219. if err != nil {
  220. t.Fatal(err)
  221. }
  222. if strings.Contains(stored, "tokentokentokentoken12") {
  223. t.Fatalf("plaintext token at rest: %s", stored)
  224. }
  225. var parsed piaStored
  226. if err := json.Unmarshal([]byte(stored), &parsed); err != nil {
  227. t.Fatal(err)
  228. }
  229. if !nodetoken.IsEncrypted(parsed.Token) {
  230. t.Fatalf("token at rest is not encrypted: %q", parsed.Token)
  231. }
  232. data, err := svc.GetPiaData()
  233. if err != nil || data == nil || data.AccountHint != "p1****67" {
  234. t.Fatalf("data: %+v err=%v", data, err)
  235. }
  236. raw, _ := json.Marshal(data)
  237. if strings.Contains(string(raw), "tokentokentokentoken12") {
  238. t.Fatalf("secret leaked in data: %s", raw)
  239. }
  240. }
  241. func TestPiaAddKeyDecryptsEncryptedToken(t *testing.T) {
  242. svc := setupPiaService(t)
  243. enablePiaTokenEncryption(t)
  244. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  245. t.Fatal(err)
  246. }
  247. key, err := svc.AddKey("useast1")
  248. if err != nil {
  249. t.Fatal(err)
  250. }
  251. if key.Tag != "pia-us-east-useast1" {
  252. t.Fatalf("key: %+v", key)
  253. }
  254. reg := svc.Registrar.(*fakePiaRegistrar)
  255. if reg.token != "tokentokentokentoken12" {
  256. t.Fatalf("addKey must decrypt the stored token, got %q", reg.token)
  257. }
  258. }
  259. func TestPiaEncryptedTokenRejectedWhenEncryptionOff(t *testing.T) {
  260. svc := setupPiaService(t)
  261. enablePiaTokenEncryption(t)
  262. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  263. t.Fatal(err)
  264. }
  265. off, _ := nodetoken.NewCodec(nodetoken.ModeOff, nil)
  266. nodetoken.Init(off)
  267. if _, err := svc.AddKey("useast1"); err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeTokenRejected {
  268. t.Fatalf("addKey with encrypted token and encryption off: %v", err)
  269. }
  270. }
  271. func TestPiaWrongAADCiphertextRejected(t *testing.T) {
  272. svc := setupPiaService(t)
  273. enablePiaTokenEncryption(t)
  274. enc, err := nodetoken.Encrypt(1, "tokentokentokentoken12")
  275. if err != nil {
  276. t.Fatal(err)
  277. }
  278. raw, _ := json.Marshal(piaStored{Username: "p1234567", Token: enc, TokenExpiresAt: time.Now().Add(time.Hour).Unix()})
  279. if err := svc.SetPia(string(raw)); err != nil {
  280. t.Fatal(err)
  281. }
  282. if _, err := svc.AddKey("useast1"); err == nil {
  283. t.Fatal("node-bound ciphertext must not decrypt as a PIA token")
  284. } else if !strings.Contains(err.Error(), "authentication failed") {
  285. t.Fatalf("wrong-AAD error: %v", err)
  286. }
  287. }
  288. func TestPiaPlaintextMigratesWhenEncryptionEnabled(t *testing.T) {
  289. svc := setupPiaService(t)
  290. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  291. t.Fatal(err)
  292. }
  293. before, err := svc.GetPia()
  294. if err != nil || !strings.Contains(before, "tokentokentokentoken12") {
  295. t.Fatalf("want plaintext before migrate: %q err=%v", before, err)
  296. }
  297. enablePiaTokenEncryption(t)
  298. if _, err := svc.GetPiaData(); err != nil {
  299. t.Fatal(err)
  300. }
  301. after, err := svc.GetPia()
  302. if err != nil || strings.Contains(after, "tokentokentokentoken12") {
  303. t.Fatalf("want ciphertext after migrate: %q err=%v", after, err)
  304. }
  305. var parsed piaStored
  306. if err := json.Unmarshal([]byte(after), &parsed); err != nil || !nodetoken.IsEncrypted(parsed.Token) {
  307. t.Fatalf("migrated token: %+v err=%v", parsed, err)
  308. }
  309. }
  310. func TestPiaReencryptsTokenToActiveKey(t *testing.T) {
  311. svc := setupPiaService(t)
  312. var k1, k2 [32]byte
  313. for i := range k1 {
  314. k1[i] = byte(i + 1)
  315. k2[i] = byte(i + 2)
  316. }
  317. c1, err := nodetoken.NewCodec(nodetoken.ModeRequired, &nodetoken.Keyring{
  318. ActiveID: "k1", Keys: map[string][32]byte{"k1": k1, "k2": k2},
  319. })
  320. if err != nil {
  321. t.Fatal(err)
  322. }
  323. nodetoken.Init(c1)
  324. t.Cleanup(func() {
  325. off, _ := nodetoken.NewCodec(nodetoken.ModeOff, nil)
  326. nodetoken.Init(off)
  327. })
  328. if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
  329. t.Fatal(err)
  330. }
  331. before, err := svc.GetPia()
  332. if err != nil || !strings.Contains(before, "enc:v1:k1:") {
  333. t.Fatalf("want k1 ciphertext: %q err=%v", before, err)
  334. }
  335. c2, err := nodetoken.NewCodec(nodetoken.ModeRequired, &nodetoken.Keyring{
  336. ActiveID: "k2", Keys: map[string][32]byte{"k1": k1, "k2": k2},
  337. })
  338. if err != nil {
  339. t.Fatal(err)
  340. }
  341. nodetoken.Init(c2)
  342. if _, err := svc.GetPiaData(); err != nil {
  343. t.Fatal(err)
  344. }
  345. after, err := svc.GetPia()
  346. if err != nil || !strings.Contains(after, "enc:v1:k2:") {
  347. t.Fatalf("want k2 ciphertext: %q err=%v", after, err)
  348. }
  349. if strings.Contains(after, "tokentokentokentoken12") {
  350. t.Fatalf("plaintext leaked after rotation: %s", after)
  351. }
  352. }