pia_test.go 11 KB

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