mirror of
https://github.com/MHSanaei/3x-ui.git
synced 2026-08-24 11:57:15 +00:00
feat(pia): add PIA login-and-add WireGuard outbounds (#6272)
* feat(pia): add login-and-add WireGuard outbounds (#2) * fix(pia): keep PIA outbounds identifiable after the editor strips hostname The outbound editor drops piaHostname, so last-segment matching failed for hyphenated servers. Identify rows by the computed tag, re-encrypt stored tokens onto the active key, skip unusable catalog rows, and always release the catalog refresh latch.
This commit is contained in:
@@ -0,0 +1,295 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/crypto/nodetoken"
|
||||
piaprotocol "github.com/mhsanaei/3x-ui/v3/internal/pia"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/wireguard"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/web/service"
|
||||
)
|
||||
|
||||
var piaTokenAAD = []byte("settings/pia_token")
|
||||
|
||||
type PiaService struct {
|
||||
service.SettingService
|
||||
Auth piaprotocol.Authenticator
|
||||
Catalog *piaprotocol.Catalog
|
||||
Registrar piaprotocol.Registrar
|
||||
}
|
||||
|
||||
type piaStored struct {
|
||||
Username string `json:"username"`
|
||||
Token string `json:"token"`
|
||||
TokenExpiresAt int64 `json:"tokenExpiresAt"`
|
||||
}
|
||||
|
||||
type PiaAccountView struct {
|
||||
Username string `json:"username"`
|
||||
AccountHint string `json:"accountHint"`
|
||||
}
|
||||
|
||||
type PiaCountryView struct {
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
type PiaRegionView struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
type PiaServerView struct {
|
||||
Hostname string `json:"hostname"`
|
||||
IP string `json:"ip"`
|
||||
RegionID string `json:"regionId"`
|
||||
RegionName string `json:"regionName"`
|
||||
}
|
||||
|
||||
type PiaServersView struct {
|
||||
Regions []PiaRegionView `json:"regions"`
|
||||
Servers []PiaServerView `json:"servers"`
|
||||
}
|
||||
|
||||
type PiaKeyView struct {
|
||||
Tag string `json:"tag"`
|
||||
Hostname string `json:"hostname"`
|
||||
SecretKey string `json:"secretKey"`
|
||||
Address string `json:"address"`
|
||||
PublicKey string `json:"publicKey"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
}
|
||||
|
||||
func NewPiaService() *PiaService {
|
||||
return &PiaService{
|
||||
Auth: piaprotocol.NewAuthClient(piaprotocol.DefaultTokenEndpoint),
|
||||
Catalog: piaprotocol.NewCatalog(piaprotocol.NewCatalogClient(piaprotocol.DefaultServerListEndpoint, piaprotocol.EmbeddedServerListPublicKey)),
|
||||
Registrar: piaprotocol.NewRegistrationClient(piaprotocol.EmbeddedPIACA),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *PiaService) Login(username, password string) (*PiaAccountView, error) {
|
||||
tok, err := s.Auth.Authenticate(context.Background(), username, []byte(password))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stored := piaStored{
|
||||
Username: strings.TrimSpace(username),
|
||||
Token: string(tok.Value),
|
||||
TokenExpiresAt: tok.ExpiresAt.Unix(),
|
||||
}
|
||||
if err := s.saveStored(stored); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return accountView(stored.Username), nil
|
||||
}
|
||||
|
||||
func (s *PiaService) GetPiaData() (*PiaAccountView, error) {
|
||||
stored, err := s.loadStored()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if stored == nil || stored.Token == "" {
|
||||
return nil, nil
|
||||
}
|
||||
return accountView(stored.Username), nil
|
||||
}
|
||||
|
||||
func (s *PiaService) DelPiaData() error {
|
||||
return s.SetPia("")
|
||||
}
|
||||
|
||||
func (s *PiaService) GetCountries() ([]PiaCountryView, error) {
|
||||
regions, err := s.regions()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
seen := map[string]struct{}{}
|
||||
out := make([]PiaCountryView, 0)
|
||||
for _, region := range regions {
|
||||
code := strings.ToUpper(strings.TrimSpace(region.CountryCode))
|
||||
if !validCountryCode(code) {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[code]; ok {
|
||||
continue
|
||||
}
|
||||
seen[code] = struct{}{}
|
||||
out = append(out, PiaCountryView{Code: code})
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].Code < out[j].Code })
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *PiaService) GetServers(countryCode string) (*PiaServersView, error) {
|
||||
code := strings.ToUpper(strings.TrimSpace(countryCode))
|
||||
if !validCountryCode(code) {
|
||||
return nil, piaprotocol.NewError(piaprotocol.CodeInvalidInput, "Select a country.")
|
||||
}
|
||||
regions, err := s.regions()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
view := &PiaServersView{Regions: []PiaRegionView{}, Servers: []PiaServerView{}}
|
||||
for _, region := range regions {
|
||||
if strings.ToUpper(strings.TrimSpace(region.CountryCode)) != code {
|
||||
continue
|
||||
}
|
||||
view.Regions = append(view.Regions, PiaRegionView{ID: region.ID, Name: region.Name})
|
||||
for _, server := range region.WireGuard {
|
||||
view.Servers = append(view.Servers, PiaServerView{
|
||||
Hostname: server.Hostname,
|
||||
IP: server.IP.String(),
|
||||
RegionID: region.ID,
|
||||
RegionName: region.Name,
|
||||
})
|
||||
}
|
||||
}
|
||||
sort.Slice(view.Regions, func(i, j int) bool { return view.Regions[i].Name < view.Regions[j].Name })
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func (s *PiaService) AddKey(hostname string) (*PiaKeyView, error) {
|
||||
hostname = strings.TrimSpace(hostname)
|
||||
if hostname == "" {
|
||||
return nil, piaprotocol.NewError(piaprotocol.CodeInvalidInput, "Select a PIA server.")
|
||||
}
|
||||
stored, err := s.loadStored()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if stored == nil || stored.Token == "" {
|
||||
return nil, piaprotocol.NewError(piaprotocol.CodeTokenRejected, "Sign in with a PIA account first.")
|
||||
}
|
||||
if stored.TokenExpiresAt > 0 && time.Now().Unix() >= stored.TokenExpiresAt {
|
||||
return nil, piaprotocol.NewError(piaprotocol.CodeTokenRejected, "The PIA token has expired. Sign in again.")
|
||||
}
|
||||
region, server, err := s.findServer(hostname)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
priv, pub, err := wireguard.GenerateWireguardKeypair()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reg, err := s.Registrar.RegisterKey(context.Background(), server, stored.Token, pub)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PiaKeyView{
|
||||
Tag: piaOutboundTag(region.ID, server.Hostname),
|
||||
Hostname: server.Hostname,
|
||||
SecretKey: priv,
|
||||
Address: reg.PeerIP.String(),
|
||||
PublicKey: reg.ServerKey,
|
||||
Endpoint: net.JoinHostPort(reg.ServerIP.String(), strconv.Itoa(int(reg.ServerPort))),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *PiaService) regions() ([]piaprotocol.Region, error) {
|
||||
if s.Catalog == nil {
|
||||
return nil, piaprotocol.NewError(piaprotocol.CodeCatalogUnavailable, "The PIA server list is not available.")
|
||||
}
|
||||
regions, _, err := s.Catalog.ListRegions(context.Background())
|
||||
return regions, err
|
||||
}
|
||||
|
||||
func (s *PiaService) findServer(hostname string) (piaprotocol.Region, piaprotocol.WireGuardServer, error) {
|
||||
regions, err := s.regions()
|
||||
if err != nil {
|
||||
return piaprotocol.Region{}, piaprotocol.WireGuardServer{}, err
|
||||
}
|
||||
for _, region := range regions {
|
||||
for _, server := range region.WireGuard {
|
||||
if server.Hostname == hostname || piaOutboundTag(region.ID, server.Hostname) == hostname {
|
||||
return region, server, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return piaprotocol.Region{}, piaprotocol.WireGuardServer{}, piaprotocol.NewError(piaprotocol.CodeServerNotFound, "The selected PIA server was not found.")
|
||||
}
|
||||
|
||||
func piaOutboundTag(regionID, hostname string) string {
|
||||
region := piaTagPart(regionID, false)
|
||||
server := piaTagPart(hostname, true)
|
||||
if region == "" {
|
||||
return "pia-" + server
|
||||
}
|
||||
return "pia-" + region + "-" + server
|
||||
}
|
||||
|
||||
func piaTagPart(s string, stripDomain bool) string {
|
||||
s = strings.ToLower(strings.TrimSpace(s))
|
||||
if stripDomain {
|
||||
if i := strings.IndexByte(s, '.'); i > 0 {
|
||||
s = s[:i]
|
||||
}
|
||||
}
|
||||
return strings.ReplaceAll(s, "_", "-")
|
||||
}
|
||||
|
||||
func (s *PiaService) saveStored(stored piaStored) error {
|
||||
enc, err := nodetoken.EncryptBound(piaTokenAAD, stored.Token)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
stored.Token = enc
|
||||
raw, err := json.Marshal(stored)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.SetPia(string(raw))
|
||||
}
|
||||
|
||||
func (s *PiaService) loadStored() (*piaStored, error) {
|
||||
raw, err := s.GetPia()
|
||||
if err != nil || strings.TrimSpace(raw) == "" {
|
||||
return nil, err
|
||||
}
|
||||
var stored piaStored
|
||||
if err := json.Unmarshal([]byte(raw), &stored); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
atRest := stored.Token
|
||||
if atRest == "" {
|
||||
return &stored, nil
|
||||
}
|
||||
if nodetoken.IsEncrypted(atRest) && !nodetoken.Enabled() {
|
||||
return nil, piaprotocol.NewError(piaprotocol.CodeTokenRejected, "The PIA token is encrypted but NODE_TOKEN_ENCRYPTION is off. Sign in again.")
|
||||
}
|
||||
plain, err := nodetoken.DecryptBound(piaTokenAAD, atRest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stored.Token = plain
|
||||
if nodetoken.Enabled() && (!nodetoken.IsEncrypted(atRest) || !nodetoken.Active().EncryptedWithActive(atRest)) {
|
||||
if err := s.saveStored(stored); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return &stored, nil
|
||||
}
|
||||
|
||||
func accountView(username string) *PiaAccountView {
|
||||
return &PiaAccountView{Username: username, AccountHint: piaAccountHint(username)}
|
||||
}
|
||||
|
||||
func piaAccountHint(username string) string {
|
||||
u := strings.TrimSpace(username)
|
||||
if len(u) <= 4 {
|
||||
return strings.Repeat("*", len(u))
|
||||
}
|
||||
return u[:2] + strings.Repeat("*", len(u)-4) + u[len(u)-2:]
|
||||
}
|
||||
|
||||
func validCountryCode(code string) bool {
|
||||
if len(code) != 2 {
|
||||
return false
|
||||
}
|
||||
return code[0] >= 'A' && code[0] <= 'Z' && code[1] >= 'A' && code[1] <= 'Z'
|
||||
}
|
||||
@@ -0,0 +1,375 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"net/netip"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/crypto/nodetoken"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/database"
|
||||
piaprotocol "github.com/mhsanaei/3x-ui/v3/internal/pia"
|
||||
)
|
||||
|
||||
type fakePiaAuth struct{ token string }
|
||||
|
||||
func (f fakePiaAuth) Authenticate(context.Context, string, []byte) (piaprotocol.Token, error) {
|
||||
return piaprotocol.Token{Value: []byte(f.token), ExpiresAt: time.Now().Add(24 * time.Hour)}, nil
|
||||
}
|
||||
|
||||
type fakePiaCatalog struct{ payload []byte }
|
||||
|
||||
func (f fakePiaCatalog) Fetch(context.Context) (piaprotocol.ServerListSnapshot, error) {
|
||||
return piaprotocol.ServerListSnapshot{Payload: f.payload, SchemaHint: "6", SignatureVerified: true}, nil
|
||||
}
|
||||
|
||||
type fakePiaRegistrar struct {
|
||||
n int
|
||||
token string
|
||||
}
|
||||
|
||||
func (f *fakePiaRegistrar) RegisterKey(_ context.Context, server piaprotocol.WireGuardServer, token string, _ string) (piaprotocol.Registration, error) {
|
||||
f.n++
|
||||
f.token = token
|
||||
key := make([]byte, 32)
|
||||
key[0] = byte(f.n)
|
||||
return piaprotocol.Registration{
|
||||
PeerIP: netip.MustParsePrefix("10.8.0." + strconv.Itoa(f.n) + "/32"),
|
||||
ServerKey: base64.StdEncoding.EncodeToString(key),
|
||||
ServerIP: server.IP,
|
||||
ServerPort: 1337,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func setupPiaService(t *testing.T) *PiaService {
|
||||
t.Helper()
|
||||
if err := database.InitDB(filepath.Join(t.TempDir(), "x-ui.db")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = database.CloseDB() })
|
||||
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"}]}}]}`)
|
||||
svc := NewPiaService()
|
||||
svc.Auth = fakePiaAuth{token: "tokentokentokentoken12"}
|
||||
svc.Catalog = piaprotocol.NewCatalog(fakePiaCatalog{payload: payload})
|
||||
svc.Registrar = &fakePiaRegistrar{}
|
||||
return svc
|
||||
}
|
||||
|
||||
func TestPiaLoginStoresTokenAndHidesItFromData(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
view, err := svc.Login("p1234567", "TEST-PIA-PASSWORD-MUST-NOT-LEAK")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if view.Username != "p1234567" || view.AccountHint != "p1****67" {
|
||||
t.Fatalf("account view: %+v", view)
|
||||
}
|
||||
data, err := svc.GetPiaData()
|
||||
if err != nil || data == nil || data.AccountHint != "p1****67" {
|
||||
t.Fatalf("data: %+v err=%v", data, err)
|
||||
}
|
||||
raw, _ := json.Marshal(data)
|
||||
if strings.Contains(string(raw), "TEST-PIA-PASSWORD-MUST-NOT-LEAK") || strings.Contains(string(raw), "tokentokentokentoken12") {
|
||||
t.Fatalf("secret leaked in data: %s", raw)
|
||||
}
|
||||
stored, err := svc.GetPia()
|
||||
if err != nil || !strings.Contains(stored, "tokentokentokentoken12") {
|
||||
t.Fatalf("token must be stored in settings: %q err=%v", stored, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaCountriesAndServers(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
countries, err := svc.GetCountries()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(countries) != 2 || countries[0].Code != "DE" || countries[1].Code != "US" {
|
||||
t.Fatalf("countries: %+v", countries)
|
||||
}
|
||||
servers, err := svc.GetServers("US")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(servers.Regions) != 1 || servers.Regions[0].ID != "us-east" || len(servers.Servers) != 2 {
|
||||
t.Fatalf("us servers: %+v", servers)
|
||||
}
|
||||
if servers.Servers[0].Hostname != "useast1" || servers.Servers[0].RegionID != "us-east" {
|
||||
t.Fatalf("first server: %+v", servers.Servers[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaAddKeyRegistersWireGuardPeer(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
if _, err := svc.AddKey("useast1"); err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeTokenRejected {
|
||||
t.Fatalf("addKey before login: %v", err)
|
||||
}
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
key, err := svc.AddKey("useast1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if key.Tag != "pia-us-east-useast1" || key.Hostname != "useast1" || key.SecretKey == "" || key.PublicKey == "" {
|
||||
t.Fatalf("key: %+v", key)
|
||||
}
|
||||
if key.Address != "10.8.0.1/32" || key.Endpoint != "198.51.100.10:1337" {
|
||||
t.Fatalf("peer: %+v", key)
|
||||
}
|
||||
byTag, err := svc.AddKey("pia-us-east-useast1")
|
||||
if err != nil || byTag.Hostname != "useast1" || byTag.Tag != "pia-us-east-useast1" {
|
||||
t.Fatalf("addKey by tag: %+v err=%v", byTag, err)
|
||||
}
|
||||
if _, err := svc.AddKey("1a"); err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeServerNotFound {
|
||||
t.Fatalf("truncated hostname must not match: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaExpiredTokenNeverReachesRegistrar(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, err := svc.GetPia()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var stored piaStored
|
||||
if err := json.Unmarshal([]byte(raw), &stored); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stored.TokenExpiresAt = time.Now().Add(-time.Minute).Unix()
|
||||
rewritten, err := json.Marshal(stored)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.SetPia(string(rewritten)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reg := svc.Registrar.(*fakePiaRegistrar)
|
||||
before := reg.n
|
||||
_, err = svc.AddKey("useast1")
|
||||
if err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeTokenRejected {
|
||||
t.Fatalf("expired token: %v", err)
|
||||
}
|
||||
if piaprotocol.MessageOf(err) != "The PIA token has expired. Sign in again." {
|
||||
t.Fatalf("expired token message: %v", err)
|
||||
}
|
||||
if reg.n != before {
|
||||
t.Fatalf("expired token reached registrar: calls=%d", reg.n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaDelClearsAccount(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.DelPiaData(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
data, err := svc.GetPiaData()
|
||||
if err != nil || data != nil {
|
||||
t.Fatalf("want nil data after logout, got %+v err=%v", data, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaOutboundTag(t *testing.T) {
|
||||
tests := []struct {
|
||||
region, host, want string
|
||||
}{
|
||||
{"us-east", "useast1", "pia-us-east-useast1"},
|
||||
{"US-East", "useast401.privacy.network", "pia-us-east-useast401"},
|
||||
{"us_california", "silicon_valley", "pia-us-california-silicon-valley"},
|
||||
{"", "berlin1", "pia-berlin1"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.want, func(t *testing.T) {
|
||||
if got := piaOutboundTag(tt.region, tt.host); got != tt.want {
|
||||
t.Fatalf("piaOutboundTag(%q, %q) = %q, want %q", tt.region, tt.host, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaCorruptSettingIsNotTreatedAsLoggedOut(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
if err := svc.SetPia(`{"username":`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
data, err := svc.GetPiaData()
|
||||
if err == nil || data != nil {
|
||||
t.Fatalf("corrupt pia setting must not look logged-out: data=%+v err=%v", data, err)
|
||||
}
|
||||
}
|
||||
|
||||
func enablePiaTokenEncryption(t *testing.T) {
|
||||
t.Helper()
|
||||
var k [32]byte
|
||||
for i := range k {
|
||||
k[i] = byte(i + 1)
|
||||
}
|
||||
ring := &nodetoken.Keyring{ActiveID: "t1", Keys: map[string][32]byte{"t1": k}}
|
||||
codec, err := nodetoken.NewCodec(nodetoken.ModeRequired, ring)
|
||||
if err != nil {
|
||||
t.Fatalf("new codec: %v", err)
|
||||
}
|
||||
nodetoken.Init(codec)
|
||||
t.Cleanup(func() {
|
||||
off, _ := nodetoken.NewCodec(nodetoken.ModeOff, nil)
|
||||
nodetoken.Init(off)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPiaLoginEncryptsTokenWhenRequired(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
enablePiaTokenEncryption(t)
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stored, err := svc.GetPia()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(stored, "tokentokentokentoken12") {
|
||||
t.Fatalf("plaintext token at rest: %s", stored)
|
||||
}
|
||||
var parsed piaStored
|
||||
if err := json.Unmarshal([]byte(stored), &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !nodetoken.IsEncrypted(parsed.Token) {
|
||||
t.Fatalf("token at rest is not encrypted: %q", parsed.Token)
|
||||
}
|
||||
data, err := svc.GetPiaData()
|
||||
if err != nil || data == nil || data.AccountHint != "p1****67" {
|
||||
t.Fatalf("data: %+v err=%v", data, err)
|
||||
}
|
||||
raw, _ := json.Marshal(data)
|
||||
if strings.Contains(string(raw), "tokentokentokentoken12") {
|
||||
t.Fatalf("secret leaked in data: %s", raw)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaAddKeyDecryptsEncryptedToken(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
enablePiaTokenEncryption(t)
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
key, err := svc.AddKey("useast1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if key.Tag != "pia-us-east-useast1" {
|
||||
t.Fatalf("key: %+v", key)
|
||||
}
|
||||
reg := svc.Registrar.(*fakePiaRegistrar)
|
||||
if reg.token != "tokentokentokentoken12" {
|
||||
t.Fatalf("addKey must decrypt the stored token, got %q", reg.token)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaEncryptedTokenRejectedWhenEncryptionOff(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
enablePiaTokenEncryption(t)
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
off, _ := nodetoken.NewCodec(nodetoken.ModeOff, nil)
|
||||
nodetoken.Init(off)
|
||||
if _, err := svc.AddKey("useast1"); err == nil || piaprotocol.CodeOf(err) != piaprotocol.CodeTokenRejected {
|
||||
t.Fatalf("addKey with encrypted token and encryption off: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaWrongAADCiphertextRejected(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
enablePiaTokenEncryption(t)
|
||||
enc, err := nodetoken.Encrypt(1, "tokentokentokentoken12")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _ := json.Marshal(piaStored{Username: "p1234567", Token: enc, TokenExpiresAt: time.Now().Add(time.Hour).Unix()})
|
||||
if err := svc.SetPia(string(raw)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := svc.AddKey("useast1"); err == nil {
|
||||
t.Fatal("node-bound ciphertext must not decrypt as a PIA token")
|
||||
} else if !strings.Contains(err.Error(), "authentication failed") {
|
||||
t.Fatalf("wrong-AAD error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaPlaintextMigratesWhenEncryptionEnabled(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
before, err := svc.GetPia()
|
||||
if err != nil || !strings.Contains(before, "tokentokentokentoken12") {
|
||||
t.Fatalf("want plaintext before migrate: %q err=%v", before, err)
|
||||
}
|
||||
enablePiaTokenEncryption(t)
|
||||
if _, err := svc.GetPiaData(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
after, err := svc.GetPia()
|
||||
if err != nil || strings.Contains(after, "tokentokentokentoken12") {
|
||||
t.Fatalf("want ciphertext after migrate: %q err=%v", after, err)
|
||||
}
|
||||
var parsed piaStored
|
||||
if err := json.Unmarshal([]byte(after), &parsed); err != nil || !nodetoken.IsEncrypted(parsed.Token) {
|
||||
t.Fatalf("migrated token: %+v err=%v", parsed, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPiaReencryptsTokenToActiveKey(t *testing.T) {
|
||||
svc := setupPiaService(t)
|
||||
var k1, k2 [32]byte
|
||||
for i := range k1 {
|
||||
k1[i] = byte(i + 1)
|
||||
k2[i] = byte(i + 2)
|
||||
}
|
||||
c1, err := nodetoken.NewCodec(nodetoken.ModeRequired, &nodetoken.Keyring{
|
||||
ActiveID: "k1", Keys: map[string][32]byte{"k1": k1, "k2": k2},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
nodetoken.Init(c1)
|
||||
t.Cleanup(func() {
|
||||
off, _ := nodetoken.NewCodec(nodetoken.ModeOff, nil)
|
||||
nodetoken.Init(off)
|
||||
})
|
||||
if _, err := svc.Login("p1234567", "password-long-enough"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
before, err := svc.GetPia()
|
||||
if err != nil || !strings.Contains(before, "enc:v1:k1:") {
|
||||
t.Fatalf("want k1 ciphertext: %q err=%v", before, err)
|
||||
}
|
||||
c2, err := nodetoken.NewCodec(nodetoken.ModeRequired, &nodetoken.Keyring{
|
||||
ActiveID: "k2", Keys: map[string][32]byte{"k1": k1, "k2": k2},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
nodetoken.Init(c2)
|
||||
if _, err := svc.GetPiaData(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
after, err := svc.GetPia()
|
||||
if err != nil || !strings.Contains(after, "enc:v1:k2:") {
|
||||
t.Fatalf("want k2 ciphertext: %q err=%v", after, err)
|
||||
}
|
||||
if strings.Contains(after, "tokentokentokentoken12") {
|
||||
t.Fatalf("plaintext leaked after rotation: %s", after)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user