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:
Masterain
2026-08-23 05:11:06 +08:00
committed by GitHub
parent a3e617215c
commit bd6a6aba43
73 changed files with 4095 additions and 31 deletions
+84
View File
@@ -0,0 +1,84 @@
package pia
import (
"bytes"
"context"
"fmt"
"mime/multipart"
"net/http"
"strings"
"time"
)
type AuthClient struct {
Endpoint string
HTTPClient *http.Client
MaxBody int64
UserAgent string
Now func() time.Time
}
func NewAuthClient(endpoint string) *AuthClient {
return &AuthClient{
Endpoint: endpoint,
MaxBody: DefaultMaxResponseBody,
UserAgent: DefaultUserAgent,
Now: time.Now,
HTTPClient: &http.Client{Timeout: DefaultRequestTimeout, CheckRedirect: noRedirect},
}
}
func (c *AuthClient) Authenticate(ctx context.Context, username string, password []byte) (Token, error) {
username = strings.TrimSpace(username)
if !validSecret([]byte(username), 1, 256) || !validSecret(password, 1, 1024) {
return Token{}, NewError(CodeInvalidCredentials, "Enter a valid PIA username and password.")
}
var body bytes.Buffer
writer := multipart.NewWriter(&body)
if err := writer.WriteField("username", username); err != nil {
return Token{}, WrapError(CodeAuthenticationUnavailable, "Could not prepare the authentication request.", err)
}
if err := writer.WriteField("password", string(password)); err != nil {
return Token{}, WrapError(CodeAuthenticationUnavailable, "Could not prepare the authentication request.", err)
}
if err := writer.Close(); err != nil {
return Token{}, WrapError(CodeAuthenticationUnavailable, "Could not prepare the authentication request.", err)
}
request, err := http.NewRequestWithContext(ctx, http.MethodPost, c.Endpoint, &body)
if err != nil {
return Token{}, WrapError(CodeAuthenticationUnavailable, "The authentication endpoint is invalid.", err)
}
request.Header.Set("Content-Type", writer.FormDataContentType())
request.Header.Set("Accept", "application/json")
request.Header.Set("User-Agent", c.UserAgent)
response, err := c.HTTPClient.Do(request)
if err != nil {
return Token{}, classifyNetworkError(ctx, CodeAuthenticationUnavailable, "PIA authentication could not be reached.", err)
}
defer response.Body.Close()
if response.StatusCode == http.StatusUnauthorized || response.StatusCode == http.StatusForbidden {
return Token{}, NewError(CodeInvalidCredentials, "The PIA username or password was rejected.")
}
if response.StatusCode != http.StatusOK {
return Token{}, NewError(CodeAuthenticationUnavailable, fmt.Sprintf("PIA authentication returned HTTP %d.", response.StatusCode))
}
if !expectedContentType(response.Header.Get("Content-Type"), "application/json") {
return Token{}, NewError(CodeAuthenticationUnavailable, "PIA authentication returned an unexpected content type.")
}
raw, err := readLimitedBody(response.Body, c.MaxBody)
if err != nil {
return Token{}, WrapError(CodeAuthenticationUnavailable, "PIA authentication returned an invalid response.", err)
}
var payload struct {
Token string `json:"token"`
}
if err := decodeSingleJSON(raw, &payload); err != nil || !validSecret([]byte(payload.Token), 16, 4096) {
return Token{}, NewError(CodeAuthenticationUnavailable, "PIA authentication returned an invalid token response.")
}
now := time.Now
if c.Now != nil {
now = c.Now
}
return Token{Value: []byte(payload.Token), ExpiresAt: now().Add(DefaultTokenTTL)}, nil
}
+143
View File
@@ -0,0 +1,143 @@
package pia
import (
"context"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"time"
)
func TestAuthClientSuccessAndReject(t *testing.T) {
successFixture, err := os.ReadFile(filepath.Join("testdata", "auth", "success.json"))
if err != nil {
t.Fatal(err)
}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if err := r.ParseMultipartForm(1 << 20); err != nil {
t.Errorf("parse form: %v", err)
}
if r.FormValue("username") != "p123" || r.FormValue("password") != "password" {
t.Errorf("unexpected credentials")
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write(successFixture)
}))
defer server.Close()
client := NewAuthClient(server.URL)
token, err := client.Authenticate(context.Background(), "p123", []byte("password"))
if err != nil || string(token.Value) != "test-token-value-that-is-long-enough" {
t.Fatalf("unexpected auth result: token=%q err=%v", token.Value, err)
}
rejected := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusUnauthorized) }))
defer rejected.Close()
client = NewAuthClient(rejected.URL)
_, err = client.Authenticate(context.Background(), "p123", []byte("wrong"))
if CodeOf(err) != CodeInvalidCredentials {
t.Fatalf("got %s, want %s", CodeOf(err), CodeInvalidCredentials)
}
}
func TestAuthClientRejectsInvalidResponsesAndTimeout(t *testing.T) {
htmlFixture, err := os.ReadFile(filepath.Join("testdata", "auth", "html.txt"))
if err != nil {
t.Fatal(err)
}
tests := []struct {
name, contentType, body string
status int
maxBody int64
wantCode string
}{
{name: "forbidden", status: http.StatusForbidden, contentType: "application/json", body: `{}`, wantCode: CodeInvalidCredentials},
{name: "html fixture", status: http.StatusOK, contentType: "text/html", body: string(htmlFixture), wantCode: CodeAuthenticationUnavailable},
{name: "malformed JSON", status: http.StatusOK, contentType: "application/json", body: `{`, wantCode: CodeAuthenticationUnavailable},
{name: "trailing JSON", status: http.StatusOK, contentType: "application/json", body: `{"token":"test-token-value-that-is-long-enough"}{}`, wantCode: CodeAuthenticationUnavailable},
{name: "short token", status: http.StatusOK, contentType: "application/json", body: `{"token":"short"}`, wantCode: CodeAuthenticationUnavailable},
{name: "oversized", status: http.StatusOK, contentType: "application/json", body: `{"token":"` + strings.Repeat("a", 100) + `"}`, maxBody: 32, wantCode: CodeAuthenticationUnavailable},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", test.contentType)
w.WriteHeader(test.status)
_, _ = w.Write([]byte(test.body))
}))
defer server.Close()
client := NewAuthClient(server.URL)
if test.maxBody > 0 {
client.MaxBody = test.maxBody
}
_, err := client.Authenticate(context.Background(), "p123", []byte("password"))
if CodeOf(err) != test.wantCode {
t.Fatalf("got %s, want %s: %v", CodeOf(err), test.wantCode, err)
}
})
}
timeoutServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
time.Sleep(100 * time.Millisecond)
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"token":"test-token-value-that-is-long-enough"}`))
}))
defer timeoutServer.Close()
client := NewAuthClient(timeoutServer.URL)
client.HTTPClient.Timeout = 25 * time.Millisecond
_, err = client.Authenticate(context.Background(), "p123", []byte("password"))
if CodeOf(err) != CodeTimeout {
t.Fatalf("timeout returned %s, want %s: %v", CodeOf(err), CodeTimeout, err)
}
}
func TestAuthClientDoesNotFollowRedirectWithSecrets(t *testing.T) {
var destinationHits atomic.Int32
destination := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
destinationHits.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer destination.Close()
origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, destination.URL, http.StatusTemporaryRedirect)
}))
defer origin.Close()
client := NewAuthClient(origin.URL)
_, _ = client.Authenticate(context.Background(), "p123", []byte("password"))
if destinationHits.Load() != 0 {
t.Fatal("authentication request followed a redirect and exposed credentials")
}
}
func TestAuthClientRejectsControlCharactersBeforeNetwork(t *testing.T) {
var hits atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
hits.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer server.Close()
client := NewAuthClient(server.URL)
_, err := client.Authenticate(context.Background(), "p123\r\nInjected", []byte("password"))
if CodeOf(err) != CodeInvalidCredentials || hits.Load() != 0 {
t.Fatalf("invalid credentials reached the network: code=%s hits=%d err=%v", CodeOf(err), hits.Load(), err)
}
}
func TestAuthErrorsOmitPassword(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
}))
defer server.Close()
client := NewAuthClient(server.URL)
password := "TEST-PIA-PASSWORD-MUST-NOT-LEAK"
_, err := client.Authenticate(context.Background(), "p123", []byte(password))
if err == nil {
t.Fatal("expected error")
}
if containsSecret(err.Error(), password) {
t.Fatalf("password leaked in error: %v", err)
}
}
+86
View File
@@ -0,0 +1,86 @@
package pia
import (
"context"
"sync"
"time"
)
type Catalog struct {
Source ServerListSource
CacheTTL time.Duration
Now func() time.Time
mu sync.Mutex
cached []Region
schema string
verified bool
fetchedAt time.Time
refreshing chan struct{}
}
func NewCatalog(source ServerListSource) *Catalog {
return &Catalog{Source: source, CacheTTL: DefaultCatalogFreshTTL, Now: time.Now}
}
func (c *Catalog) ListRegions(ctx context.Context) ([]Region, string, error) {
for {
c.mu.Lock()
age := c.Now().Sub(c.fetchedAt)
if len(c.cached) > 0 && c.verified && c.CacheTTL > 0 && age >= 0 && age < c.CacheTTL {
regions, schema := cloneRegions(c.cached), c.schema
c.mu.Unlock()
return regions, schema, nil
}
if wait := c.refreshing; wait != nil {
c.mu.Unlock()
select {
case <-ctx.Done():
return nil, "", ctx.Err()
case <-wait:
continue
}
}
done := make(chan struct{})
c.refreshing = done
c.mu.Unlock()
return c.fetchAndPublish(ctx, done)
}
}
func (c *Catalog) fetchAndPublish(ctx context.Context, done chan struct{}) ([]Region, string, error) {
defer func() {
c.mu.Lock()
c.refreshing = nil
close(done)
c.mu.Unlock()
}()
snapshot, err := c.Source.Fetch(ctx)
var regions []Region
var schema string
if err == nil && !snapshot.SignatureVerified {
err = NewError(CodeCatalogSignatureInvalid, "The PIA region list was not signature-verified.")
}
if err == nil {
regions, schema, err = ParseServerList(snapshot.Payload, snapshot.SchemaHint)
}
if err != nil {
return nil, "", err
}
c.mu.Lock()
c.cached = cloneRegions(regions)
c.schema = schema
c.verified = true
c.fetchedAt = c.Now()
c.mu.Unlock()
return cloneRegions(regions), schema, nil
}
func cloneRegions(regions []Region) []Region {
result := make([]Region, len(regions))
for i, region := range regions {
result[i] = region
result[i].WireGuard = append([]WireGuardServer(nil), region.WireGuard...)
}
return result
}
+125
View File
@@ -0,0 +1,125 @@
package pia
import (
"context"
"os"
"path/filepath"
"sync"
"sync/atomic"
"testing"
"time"
)
type fakeServerListSource struct {
snapshot ServerListSnapshot
err error
calls int
}
func (f *fakeServerListSource) Fetch(context.Context) (ServerListSnapshot, error) {
f.calls++
return f.snapshot, f.err
}
func TestCatalogCachesOnlyVerifiedParsedSnapshots(t *testing.T) {
raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", "v6_valid.json"))
if err != nil {
t.Fatal(err)
}
source := &fakeServerListSource{snapshot: ServerListSnapshot{Payload: raw, SchemaHint: "6", SignatureVerified: true}}
now := time.Unix(1_700_000_000, 0)
catalog := NewCatalog(source)
catalog.CacheTTL = 30 * time.Minute
catalog.Now = func() time.Time { return now }
first, schema, err := catalog.ListRegions(context.Background())
if err != nil || schema != "v6" || len(first) != 1 {
t.Fatalf("unexpected first result: schema=%q regions=%v err=%v", schema, first, err)
}
first[0].WireGuard[0].Hostname = "mutated-by-caller"
second, _, err := catalog.ListRegions(context.Background())
if err != nil || source.calls != 1 {
t.Fatalf("verified snapshot was not cached: calls=%d err=%v", source.calls, err)
}
if second[0].WireGuard[0].Hostname == "mutated-by-caller" {
t.Fatal("catalog returned mutable cached storage")
}
now = now.Add(-time.Second)
if _, _, err := catalog.ListRegions(context.Background()); err != nil || source.calls != 2 {
t.Fatalf("backward clock movement incorrectly extended the cache: calls=%d err=%v", source.calls, err)
}
now = now.Add(catalog.CacheTTL + time.Second)
if _, _, err := catalog.ListRegions(context.Background()); err != nil || source.calls != 3 {
t.Fatalf("expired snapshot was not refreshed: calls=%d err=%v", source.calls, err)
}
}
type gatedServerListSource struct {
snapshot ServerListSnapshot
started chan struct{}
release chan struct{}
startOnce sync.Once
calls atomic.Int32
}
func (g *gatedServerListSource) Fetch(context.Context) (ServerListSnapshot, error) {
g.calls.Add(1)
g.startOnce.Do(func() { close(g.started) })
<-g.release
return g.snapshot, nil
}
func TestCatalogCoalescesConcurrentRefresh(t *testing.T) {
raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", "v6_valid.json"))
if err != nil {
t.Fatal(err)
}
source := &gatedServerListSource{
snapshot: ServerListSnapshot{Payload: raw, SchemaHint: "6", SignatureVerified: true},
started: make(chan struct{}),
release: make(chan struct{}),
}
catalog := NewCatalog(source)
catalog.CacheTTL = time.Hour
errc := make(chan error, 2)
go func() {
_, _, err := catalog.ListRegions(context.Background())
errc <- err
}()
<-source.started
go func() {
_, _, err := catalog.ListRegions(context.Background())
errc <- err
}()
deadline := time.Now().Add(200 * time.Millisecond)
for time.Now().Before(deadline) {
if source.calls.Load() > 1 {
close(source.release)
t.Fatalf("concurrent refresh issued %d fetches, want 1", source.calls.Load())
}
time.Sleep(time.Millisecond)
}
close(source.release)
for i := 0; i < 2; i++ {
if err := <-errc; err != nil {
t.Fatal(err)
}
}
if source.calls.Load() != 1 {
t.Fatalf("concurrent refresh issued %d fetches, want 1", source.calls.Load())
}
}
func TestCatalogRejectsUnverifiedSnapshot(t *testing.T) {
source := &fakeServerListSource{snapshot: ServerListSnapshot{
Payload: []byte(`{"version":6,"groups":{"wg":[]},"regions":[]}`), SchemaHint: "6", SignatureVerified: false,
}}
catalog := NewCatalog(source)
_, _, err := catalog.ListRegions(context.Background())
if CodeOf(err) != CodeCatalogSignatureInvalid {
t.Fatalf("unverified snapshot returned %s, want %s: %v", CodeOf(err), CodeCatalogSignatureInvalid, err)
}
}
+24
View File
@@ -0,0 +1,24 @@
package pia
import (
_ "embed"
"time"
)
const (
DefaultTokenEndpoint = "https://www.privateinternetaccess.com/api/client/v2/token"
DefaultServerListEndpoint = "https://serverlist.piaservers.net/vpninfo/servers/v6"
DefaultAddKeyPort = uint16(1337)
DefaultUserAgent = "3x-ui-pia/1.0"
DefaultMaxServerListBody = int64(8 << 20)
DefaultMaxResponseBody = int64(64 << 10)
DefaultRequestTimeout = 20 * time.Second
DefaultCatalogFreshTTL = 6 * time.Hour
DefaultTokenTTL = 24 * time.Hour
)
//go:embed trust/ca.rsa.4096.crt
var EmbeddedPIACA []byte
//go:embed trust/serverlist_public_key.pem
var EmbeddedServerListPublicKey []byte
+73
View File
@@ -0,0 +1,73 @@
package pia
import (
"errors"
"fmt"
)
const (
CodeInvalidInput = "pia_invalid_input"
CodeInvalidCredentials = "pia_invalid_credentials"
CodeAuthenticationUnavailable = "pia_authentication_unavailable"
CodeTokenRejected = "pia_token_rejected"
CodeCatalogUnavailable = "pia_catalog_unavailable"
CodeCatalogSignatureInvalid = "pia_catalog_signature_invalid"
CodeCatalogSchemaUnsupported = "pia_catalog_schema_unsupported"
CodeServerNotFound = "pia_server_not_found"
CodeTLSValidation = "pia_tls_validation"
CodeRegistrationRejected = "pia_registration_rejected"
CodeRegistrationInvalid = "pia_registration_response_invalid"
CodeTimeout = "pia_timeout"
CodeCancelled = "pia_cancelled"
CodeNetworkUnavailable = "pia_network_unavailable"
)
type Error struct {
Code string
Message string
cause error
}
func NewError(code, message string) *Error {
return &Error{Code: code, Message: message}
}
func WrapError(code, message string, cause error) *Error {
return &Error{Code: code, Message: message, cause: cause}
}
func (e *Error) Error() string {
if e == nil {
return ""
}
return fmt.Sprintf("%s: %s", e.Code, e.Message)
}
func (e *Error) Unwrap() error {
if e == nil {
return nil
}
return e.cause
}
func CodeOf(err error) string {
if err == nil {
return ""
}
var pe *Error
if errors.As(err, &pe) && pe != nil {
return pe.Code
}
return CodeNetworkUnavailable
}
func MessageOf(err error) string {
var pe *Error
if errors.As(err, &pe) && pe != nil {
return pe.Message
}
if err == nil {
return ""
}
return "An unexpected error occurred."
}
+37
View File
@@ -0,0 +1,37 @@
package pia
import (
"encoding/base64"
"errors"
"testing"
)
func TestNilErrorHelpers(t *testing.T) {
if code := CodeOf(nil); code != "" {
t.Fatalf("CodeOf(nil)=%q, want empty", code)
}
var typed *Error
var err error = typed
if code := CodeOf(err); code != CodeNetworkUnavailable {
t.Fatalf("CodeOf(nil *Error)=%q, want %q", code, CodeNetworkUnavailable)
}
if unwrapped := errors.Unwrap(err); unwrapped != nil {
t.Fatalf("errors.Unwrap(nil *Error)=%v, want nil", unwrapped)
}
}
func TestValidWGKeyRequiresBase64Encoded32Bytes(t *testing.T) {
valid := base64.StdEncoding.EncodeToString(make([]byte, 32))
if !validWGKey(valid) {
t.Fatalf("valid WireGuard key rejected: %q", valid)
}
for _, invalid := range []string{
"!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!",
base64.StdEncoding.EncodeToString(make([]byte, 31)),
base64.StdEncoding.EncodeToString(make([]byte, 33)),
} {
if validWGKey(invalid) {
t.Fatalf("invalid WireGuard key accepted: %q", invalid)
}
}
}
+100
View File
@@ -0,0 +1,100 @@
package pia
import (
"bytes"
"context"
"crypto/x509"
"encoding/json"
"errors"
"fmt"
"io"
"mime"
"net"
"net/http"
"strings"
)
func readLimitedBody(body io.Reader, limit int64) ([]byte, error) {
raw, err := io.ReadAll(io.LimitReader(body, limit+1))
if err != nil {
return nil, err
}
if int64(len(raw)) > limit {
return nil, fmt.Errorf("response exceeds %d bytes", limit)
}
return raw, nil
}
func expectedContentType(header string, accepted ...string) bool {
mediaType, _, err := mime.ParseMediaType(header)
if err != nil {
return false
}
for _, candidate := range accepted {
if strings.EqualFold(mediaType, candidate) {
return true
}
}
return false
}
func noRedirect(_ *http.Request, _ []*http.Request) error {
return errors.New("redirects are disabled for this request")
}
func decodeSingleJSON(raw []byte, target any) error {
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.UseNumber()
if err := decoder.Decode(target); err != nil {
return err
}
var extra any
if err := decoder.Decode(&extra); !errors.Is(err, io.EOF) {
if err == nil {
return errors.New("multiple JSON values are not allowed")
}
return err
}
return nil
}
func classifyNetworkError(ctx context.Context, fallback, message string, err error) error {
cause := redactNetErr(err)
if errors.Is(ctx.Err(), context.Canceled) || errors.Is(err, context.Canceled) {
return WrapError(CodeCancelled, "The operation was cancelled.", cause)
}
if errors.Is(ctx.Err(), context.DeadlineExceeded) || errors.Is(err, context.DeadlineExceeded) {
return WrapError(CodeTimeout, "The network request timed out.", cause)
}
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
return WrapError(CodeTimeout, "The network request timed out.", cause)
}
var unknownAuthority x509.UnknownAuthorityError
var hostnameError x509.HostnameError
var invalidCertificate x509.CertificateInvalidError
if errors.As(err, &unknownAuthority) || errors.As(err, &hostnameError) || errors.As(err, &invalidCertificate) {
return WrapError(CodeTLSValidation, "PIA's server identity could not be verified.", cause)
}
return WrapError(fallback, message, cause)
}
type redactedCause struct{ kind string }
func (e redactedCause) Error() string { return e.kind }
func redactNetErr(err error) error {
if err == nil {
return nil
}
return redactedCause{kind: "network error"}
}
func containsSecret(s string, secrets ...string) bool {
for _, secret := range secrets {
if secret != "" && strings.Contains(s, secret) {
return true
}
}
return false
}
+147
View File
@@ -0,0 +1,147 @@
package pia
import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
"net"
"net/http"
"net/netip"
"net/url"
"strconv"
"time"
)
type RegistrationClient struct {
CAPEM []byte
Port uint16
MaxBody int64
Timeout time.Duration
UserAgent string
}
func NewRegistrationClient(caPEM []byte) *RegistrationClient {
return &RegistrationClient{
CAPEM: caPEM, Port: DefaultAddKeyPort, MaxBody: DefaultMaxResponseBody,
Timeout: DefaultRequestTimeout, UserAgent: DefaultUserAgent,
}
}
func (c *RegistrationClient) RegisterKey(ctx context.Context, server WireGuardServer, token string, publicKey string) (Registration, error) {
if !server.IP.IsValid() || !server.IP.Is4() || !validHostname(server.Hostname) {
return Registration{}, NewError(CodeInvalidInput, "The selected PIA WireGuard server is invalid.")
}
if !validSecret([]byte(token), 16, 4096) {
return Registration{}, NewError(CodeTokenRejected, "The PIA authentication token is invalid.")
}
if !validWGKey(publicKey) {
return Registration{}, NewError(CodeInvalidInput, "The WireGuard public key is invalid.")
}
roots := x509.NewCertPool()
if !roots.AppendCertsFromPEM(c.CAPEM) {
return Registration{}, NewError(CodeTLSValidation, "The built-in PIA certificate authority is invalid.")
}
port := c.Port
if port == 0 {
port = DefaultAddKeyPort
}
dialer := &net.Dialer{Timeout: 8 * time.Second, KeepAlive: 30 * time.Second}
transport := &http.Transport{
Proxy: nil,
DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) {
return dialer.DialContext(ctx, network, net.JoinHostPort(server.IP.String(), strconv.Itoa(int(port))))
},
TLSClientConfig: &tls.Config{ServerName: server.Hostname, RootCAs: roots, MinVersion: tls.VersionTLS12},
TLSHandshakeTimeout: 8 * time.Second, ResponseHeaderTimeout: 12 * time.Second, ForceAttemptHTTP2: true,
}
defer transport.CloseIdleConnections()
client := &http.Client{Transport: transport, Timeout: c.Timeout, CheckRedirect: noRedirect}
endpoint := url.URL{Scheme: "https", Host: net.JoinHostPort(server.Hostname, strconv.Itoa(int(port))), Path: "/addKey"}
query := endpoint.Query()
query.Set("pt", token)
query.Set("pubkey", publicKey)
endpoint.RawQuery = query.Encode()
request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil)
if err != nil {
return Registration{}, WrapError(CodeRegistrationRejected, "Could not prepare PIA key registration.", err)
}
request.Header.Set("Accept", "application/json")
request.Header.Set("User-Agent", c.UserAgent)
response, err := client.Do(request)
if err != nil {
return Registration{}, classifyNetworkError(ctx, CodeNetworkUnavailable, "The selected PIA WireGuard server could not be reached.", err)
}
defer response.Body.Close()
if response.StatusCode == http.StatusUnauthorized || response.StatusCode == http.StatusForbidden {
return Registration{}, NewError(CodeTokenRejected, "The PIA authentication token was rejected.")
}
if response.StatusCode != http.StatusOK {
return Registration{}, NewError(CodeRegistrationRejected, fmt.Sprintf("PIA key registration returned HTTP %d.", response.StatusCode))
}
if !expectedContentType(response.Header.Get("Content-Type"), "application/json") {
return Registration{}, NewError(CodeRegistrationInvalid, "PIA key registration returned an unexpected content type.")
}
raw, err := readLimitedBody(response.Body, c.MaxBody)
if err != nil {
return Registration{}, WrapError(CodeRegistrationInvalid, "PIA key registration returned an invalid response.", err)
}
return parseRegistration(raw)
}
func parseRegistration(raw []byte) (Registration, error) {
var payload struct {
Status string `json:"status"`
PeerIP string `json:"peer_ip"`
ServerKey string `json:"server_key"`
ServerIP string `json:"server_ip"`
ServerPort int `json:"server_port"`
DNSServers []string `json:"dns_servers"`
}
if err := decodeSingleJSON(raw, &payload); err != nil {
return Registration{}, NewError(CodeRegistrationInvalid, "PIA key registration returned malformed JSON.")
}
if payload.Status != "OK" {
return Registration{}, NewError(CodeRegistrationRejected, "The PIA server rejected WireGuard key registration.")
}
peerIP, err := parsePeerIP(payload.PeerIP)
if err != nil {
return Registration{}, NewError(CodeRegistrationInvalid, "PIA returned an invalid WireGuard peer address.")
}
if !validWGKey(payload.ServerKey) {
return Registration{}, NewError(CodeRegistrationInvalid, "PIA returned an invalid WireGuard server key.")
}
serverIP, err := netip.ParseAddr(payload.ServerIP)
if err != nil || !serverIP.Is4() || serverIP.IsUnspecified() {
return Registration{}, NewError(CodeRegistrationInvalid, "PIA returned an invalid WireGuard server address.")
}
if payload.ServerPort < 1 || payload.ServerPort > 65535 {
return Registration{}, NewError(CodeRegistrationInvalid, "PIA returned an invalid WireGuard server port.")
}
dns := make([]netip.Addr, 0, len(payload.DNSServers))
for _, value := range payload.DNSServers {
address, err := netip.ParseAddr(value)
if err != nil || !address.Is4() || address.IsUnspecified() {
continue
}
if len(dns) == 8 {
break
}
dns = append(dns, address)
}
return Registration{PeerIP: peerIP, ServerKey: payload.ServerKey, ServerIP: serverIP, ServerPort: uint16(payload.ServerPort), DNSServers: dns}, nil
}
func parsePeerIP(value string) (netip.Prefix, error) {
if address, err := netip.ParseAddr(value); err == nil {
if !address.Is4() || address.IsUnspecified() {
return netip.Prefix{}, fmt.Errorf("peer address is not a usable IPv4 address")
}
return netip.PrefixFrom(address, 32), nil
}
prefix, err := netip.ParsePrefix(value)
if err != nil || !prefix.Addr().Is4() || prefix.Addr().IsUnspecified() || prefix.Bits() != 32 {
return netip.Prefix{}, fmt.Errorf("peer address is not an IPv4 host prefix")
}
return prefix, nil
}
+260
View File
@@ -0,0 +1,260 @@
package pia
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/base64"
"encoding/pem"
"math/big"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"os"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"time"
)
func testPubKey() string {
raw := make([]byte, 32)
raw[0] = 1
return base64.StdEncoding.EncodeToString(raw)
}
func TestParseRegistrationFixture(t *testing.T) {
raw, err := os.ReadFile(filepath.Join("testdata", "addkey", "success.json"))
if err != nil {
t.Fatal(err)
}
result, err := parseRegistration(raw)
if err != nil {
t.Fatal(err)
}
if result.PeerIP.String() != "10.42.0.2/32" || result.ServerPort != 51820 || result.ServerIP.String() != "198.51.100.10" || len(result.DNSServers) != 2 {
t.Fatalf("unexpected registration result: %+v", result)
}
prefixed := []byte(`{"status":"OK","peer_ip":"10.42.0.3/32","server_key":"AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=","server_ip":"198.51.100.10","server_port":51820,"dns_servers":["10.0.0.242"]}`)
prefixedResult, err := parseRegistration(prefixed)
if err != nil || prefixedResult.PeerIP.String() != "10.42.0.3/32" {
t.Fatalf("expected an explicit /32 peer address to remain supported: result=%+v err=%v", prefixedResult, err)
}
missing, err := os.ReadFile(filepath.Join("testdata", "addkey", "missing_port.json"))
if err != nil {
t.Fatal(err)
}
if _, err := parseRegistration(missing); err == nil || CodeOf(err) != CodeRegistrationInvalid {
t.Fatalf("expected missing server port to be rejected: %v", err)
}
dnsRaw, err := os.ReadFile(filepath.Join("testdata", "addkey", "invalid_dns.json"))
if err != nil {
t.Fatal(err)
}
dnsResult, err := parseRegistration(dnsRaw)
if err != nil || len(dnsResult.DNSServers) != 0 {
t.Fatalf("invalid dns_servers must be ignored: result=%+v err=%v", dnsResult, err)
}
invalidFixtures := []struct{ file, wantCode string }{
{"status_error.json", CodeRegistrationRejected},
{"invalid_peer_ip.json", CodeRegistrationInvalid},
{"invalid_peer_prefix.json", CodeRegistrationInvalid},
{"invalid_server_key.json", CodeRegistrationInvalid},
{"invalid_server_ip.json", CodeRegistrationInvalid},
{"invalid_port.json", CodeRegistrationInvalid},
}
for _, test := range invalidFixtures {
t.Run(test.file, func(t *testing.T) {
raw, readErr := os.ReadFile(filepath.Join("testdata", "addkey", test.file))
if readErr != nil {
t.Fatal(readErr)
}
if _, parseErr := parseRegistration(raw); CodeOf(parseErr) != test.wantCode {
t.Fatalf("expected %s, got %s: %v", test.wantCode, CodeOf(parseErr), parseErr)
}
})
}
}
func TestRegistrationTLSHostnameAndCA(t *testing.T) {
fixture, err := os.ReadFile(filepath.Join("testdata", "addkey", "success.json"))
if err != nil {
t.Fatal(err)
}
const token = "test-token-value-that-is-long-enough"
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Query().Get("pt") != token || r.URL.Query().Get("pubkey") == "" {
t.Errorf("registration query is missing required values")
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write(fixture)
}))
server.StartTLS()
defer server.Close()
certificate := server.Certificate()
caPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate.Raw})
host, portText, err := net.SplitHostPort(server.Listener.Addr().String())
if err != nil {
t.Fatal(err)
}
port, err := net.LookupPort("tcp", portText)
if err != nil {
t.Fatal(err)
}
client := NewRegistrationClient(caPEM)
client.Port = uint16(port)
key := testPubKey()
t.Setenv("HTTPS_PROXY", "http://127.0.0.1:1")
_, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)}, token, key)
if err != nil {
t.Fatalf("expected TLS registration to succeed: %v", err)
}
_, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)}, token, "")
if CodeOf(err) != CodeInvalidInput {
t.Fatalf("zero public key returned %s, want %s", CodeOf(err), CodeInvalidInput)
}
_, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "wrong.example", IP: netip.MustParseAddr(host)}, token, key)
if CodeOf(err) != CodeTLSValidation {
t.Fatalf("wrong hostname returned %s, want %s: %v", CodeOf(err), CodeTLSValidation, err)
}
client = NewRegistrationClient([]byte("-----BEGIN CERTIFICATE-----\ninvalid\n-----END CERTIFICATE-----"))
client.Port = uint16(port)
_, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)}, token, key)
if CodeOf(err) != CodeTLSValidation {
t.Fatalf("wrong CA returned %s, want %s", CodeOf(err), CodeTLSValidation)
}
expiredServer, expiredCA := newExpiredTLSServer(t, fixture)
defer expiredServer.Close()
expiredHost, expiredPortText, err := net.SplitHostPort(expiredServer.Listener.Addr().String())
if err != nil {
t.Fatal(err)
}
expiredPort, err := net.LookupPort("tcp", expiredPortText)
if err != nil {
t.Fatal(err)
}
client = NewRegistrationClient(expiredCA)
client.Port = uint16(expiredPort)
_, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(expiredHost)}, token, key)
if CodeOf(err) != CodeTLSValidation {
t.Fatalf("expired certificate returned %s, want %s: %v", CodeOf(err), CodeTLSValidation, err)
}
}
func TestRegistrationResponseGuardsAndRedirect(t *testing.T) {
var destinationHits atomic.Int32
destination := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
destinationHits.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer destination.Close()
tests := []struct {
name, contentType, body, redirect string
maxBody int64
delay time.Duration
wantCode string
}{
{name: "HTML", contentType: "text/html", body: "<html>maintenance</html>", wantCode: CodeRegistrationInvalid},
{name: "oversized", contentType: "application/json", body: strings.Repeat("x", 65), maxBody: 64, wantCode: CodeRegistrationInvalid},
{name: "redirect", contentType: "application/json", redirect: destination.URL, wantCode: CodeNetworkUnavailable},
{name: "timeout", contentType: "application/json", body: `{"status":"OK"}`, delay: 100 * time.Millisecond, wantCode: CodeTimeout},
}
const token = "test-token-value-that-is-long-enough"
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if test.delay > 0 {
time.Sleep(test.delay)
}
if test.redirect != "" {
http.Redirect(w, r, test.redirect, http.StatusTemporaryRedirect)
return
}
w.Header().Set("Content-Type", test.contentType)
_, _ = w.Write([]byte(test.body))
}))
server.StartTLS()
defer server.Close()
caPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: server.Certificate().Raw})
host, portText, err := net.SplitHostPort(server.Listener.Addr().String())
if err != nil {
t.Fatal(err)
}
port, err := net.LookupPort("tcp", portText)
if err != nil {
t.Fatal(err)
}
client := NewRegistrationClient(caPEM)
client.Port = uint16(port)
if test.maxBody > 0 {
client.MaxBody = test.maxBody
}
if test.delay > 0 {
client.Timeout = 25 * time.Millisecond
}
_, err = client.RegisterKey(
context.Background(),
WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)},
token,
testPubKey(),
)
if err == nil {
t.Fatal("expected registration error")
}
if CodeOf(err) != test.wantCode {
t.Fatalf("got %s, want %s: %v", CodeOf(err), test.wantCode, err)
}
if containsSecret(err.Error(), token) {
t.Fatalf("token leaked in error: %v", err)
}
})
}
if destinationHits.Load() != 0 {
t.Fatal("registration request followed a redirect and exposed secrets")
}
}
func newExpiredTLSServer(t *testing.T, response []byte) (*httptest.Server, []byte) {
t.Helper()
key, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "example.com"},
DNSNames: []string{"example.com"},
NotBefore: time.Now().Add(-48 * time.Hour),
NotAfter: time.Now().Add(-24 * time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
IsCA: true,
BasicConstraintsValid: true,
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
if err != nil {
t.Fatal(err)
}
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
certificate, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
t.Fatal(err)
}
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write(response)
}))
server.TLS = &tls.Config{Certificates: []tls.Certificate{certificate}, MinVersion: tls.VersionTLS12}
server.StartTLS()
return server, certPEM
}
+79
View File
@@ -0,0 +1,79 @@
package pia
import (
"context"
"fmt"
"net/http"
"net/url"
"path"
"strings"
)
type ServerListSource interface {
Fetch(ctx context.Context) (ServerListSnapshot, error)
}
type ServerListSnapshot struct {
Payload []byte
SchemaHint string
SignatureVerified bool
}
type CatalogClient struct {
Endpoint string
PublicKeyPEM []byte
HTTPClient *http.Client
MaxBody int64
UserAgent string
}
func NewCatalogClient(endpoint string, publicKey []byte) *CatalogClient {
return &CatalogClient{
Endpoint: endpoint,
PublicKeyPEM: publicKey,
MaxBody: DefaultMaxServerListBody,
UserAgent: DefaultUserAgent,
HTTPClient: &http.Client{Timeout: DefaultRequestTimeout, CheckRedirect: noRedirect},
}
}
func (c *CatalogClient) Fetch(ctx context.Context) (ServerListSnapshot, error) {
request, err := http.NewRequestWithContext(ctx, http.MethodGet, c.Endpoint, nil)
if err != nil {
return ServerListSnapshot{}, WrapError(CodeCatalogUnavailable, "The PIA region-list endpoint is invalid.", err)
}
request.Header.Set("Accept", "application/json, text/plain;q=0.9")
request.Header.Set("User-Agent", c.UserAgent)
response, err := c.HTTPClient.Do(request)
if err != nil {
return ServerListSnapshot{}, classifyNetworkError(ctx, CodeCatalogUnavailable, "The PIA region list could not be downloaded.", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return ServerListSnapshot{}, NewError(CodeCatalogUnavailable, fmt.Sprintf("PIA returned HTTP %d for the region list.", response.StatusCode))
}
if !expectedContentType(response.Header.Get("Content-Type"), "application/json", "text/plain", "application/octet-stream") {
return ServerListSnapshot{}, NewError(CodeCatalogSchemaUnsupported, "PIA returned an unexpected region-list content type.")
}
raw, err := readLimitedBody(response.Body, c.MaxBody)
if err != nil {
return ServerListSnapshot{}, WrapError(CodeCatalogUnavailable, "The PIA region-list response is too large or incomplete.", err)
}
verified, err := VerifySignedServerList(raw, c.PublicKeyPEM)
if err != nil {
return ServerListSnapshot{}, err
}
return ServerListSnapshot{Payload: verified, SchemaHint: schemaHint(c.Endpoint), SignatureVerified: true}, nil
}
func schemaHint(endpoint string) string {
parsed, err := url.Parse(endpoint)
if err != nil {
return ""
}
base := strings.ToLower(path.Base(parsed.Path))
if base == "v6" || base == "v7" {
return strings.TrimPrefix(base, "v")
}
return ""
}
+108
View File
@@ -0,0 +1,108 @@
package pia
import (
"context"
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/pem"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestCatalogClientReturnsExplicitlyVerifiedSnapshot(t *testing.T) {
payload := []byte(`{"version":6,"groups":{"wg":[]},"regions":[]}`)
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
digest := sha256.Sum256(payload)
signature, err := rsa.SignPKCS1v15(rand.Reader, privateKey, crypto.SHA256, digest[:])
if err != nil {
t.Fatal(err)
}
publicDER, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey)
if err != nil {
t.Fatal(err)
}
signed := append(append(append([]byte{}, payload...), '\n'), []byte(base64.StdEncoding.EncodeToString(signature))...)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/octet-stream")
_, _ = w.Write(signed)
}))
defer server.Close()
client := NewCatalogClient(server.URL+"/v6", pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: publicDER}))
snapshot, err := client.Fetch(context.Background())
if err != nil {
t.Fatal(err)
}
if !snapshot.SignatureVerified || snapshot.SchemaHint != "6" || string(snapshot.Payload) != string(payload) {
t.Fatalf("unexpected verified snapshot: %+v", snapshot)
}
}
func TestCatalogClientRejectsUnsafeResponses(t *testing.T) {
tests := []struct {
name, contentType, body string
maxBody int64
wantCode string
}{
{name: "html", contentType: "text/html", body: "<html>maintenance</html>", wantCode: CodeCatalogSchemaUnsupported},
{name: "oversized", contentType: "application/octet-stream", body: strings.Repeat("x", 65), maxBody: 64, wantCode: CodeCatalogUnavailable},
{name: "unsigned", contentType: "application/json", body: `{"version":6,"groups":{},"regions":[]}`, wantCode: CodeCatalogSignatureInvalid},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", test.contentType)
_, _ = w.Write([]byte(test.body))
}))
defer server.Close()
client := NewCatalogClient(server.URL+"/v6", []byte("invalid public key"))
if test.maxBody > 0 {
client.MaxBody = test.maxBody
}
_, err := client.Fetch(context.Background())
if CodeOf(err) != test.wantCode {
t.Fatalf("got %s, want %s: %v", CodeOf(err), test.wantCode, err)
}
})
}
payload := []byte(`{"version":6,"groups":{"wg":[]},"regions":[]}`)
signingKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
digest := sha256.Sum256(payload)
signature, err := rsa.SignPKCS1v15(rand.Reader, signingKey, crypto.SHA256, digest[:])
if err != nil {
t.Fatal(err)
}
signed := append(append(append([]byte{}, payload...), '\n'), []byte(base64.StdEncoding.EncodeToString(signature))...)
wrongKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
wrongDER, err := x509.MarshalPKIXPublicKey(&wrongKey.PublicKey)
if err != nil {
t.Fatal(err)
}
t.Run("valid signature from unpinned key", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/octet-stream")
_, _ = w.Write(signed)
}))
defer server.Close()
client := NewCatalogClient(server.URL+"/v6", pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: wrongDER}))
_, err := client.Fetch(context.Background())
if CodeOf(err) != CodeCatalogSignatureInvalid {
t.Fatalf("got %s, want %s: %v", CodeOf(err), CodeCatalogSignatureInvalid, err)
}
})
}
+196
View File
@@ -0,0 +1,196 @@
package pia
import (
"encoding/json"
"fmt"
"net/netip"
"regexp"
"sort"
"strings"
)
var (
regionIDPattern = regexp.MustCompile(`^[A-Za-z0-9_-]{1,64}$`)
countryCodePattern = regexp.MustCompile(`^[A-Za-z]{2}$`)
)
type ServerListParser interface {
Schema() string
CanParse(raw []byte) bool
Parse(raw []byte) ([]Region, error)
}
type (
V6Parser struct{}
V7Parser struct{}
)
func (V6Parser) Schema() string { return "v6" }
func (V7Parser) Schema() string { return "v7" }
func (V6Parser) CanParse(raw []byte) bool { return schemaVersion(raw) == 0 || schemaVersion(raw) == 6 }
func (V7Parser) CanParse(raw []byte) bool { return schemaVersion(raw) == 7 }
func (V6Parser) Parse(raw []byte) ([]Region, error) { return parseCatalog(raw, false) }
func (V7Parser) Parse(raw []byte) ([]Region, error) { return parseCatalog(raw, true) }
type catalogEnvelope struct {
Version json.RawMessage `json:"version"`
Groups map[string]json.RawMessage `json:"groups"`
Regions []rawRegion `json:"regions"`
}
type rawRegion struct {
ID string `json:"id"`
Name string `json:"name"`
Country string `json:"country"`
Geo *bool `json:"geo"`
Offline *bool `json:"offline"`
PortForward *bool `json:"port_forward"`
PortForwarding *bool `json:"port_forwarding"`
Servers rawServers `json:"servers"`
}
type rawServers struct {
WireGuard []rawServer `json:"wg"`
}
type rawServer struct {
IP string `json:"ip"`
CN string `json:"cn"`
Hostname string `json:"hostname"`
}
func ParseServerList(raw []byte, schemaHint string) ([]Region, string, error) {
parsers := []ServerListParser{V7Parser{}, V6Parser{}}
version, present, err := detectSchemaVersion(raw)
if err != nil {
return nil, "", WrapError(CodeCatalogSchemaUnsupported, "PIA returned an invalid server-list version.", err)
}
if present {
for _, parser := range parsers {
if strings.TrimPrefix(parser.Schema(), "v") == fmt.Sprint(version) {
regions, parseErr := parser.Parse(raw)
return regions, parser.Schema(), parseErr
}
}
return nil, "", NewError(CodeCatalogSchemaUnsupported, "This PIA server-list schema is not supported.")
}
hint := strings.ToLower(strings.TrimPrefix(schemaHint, "v"))
if hint != "" {
for _, parser := range parsers {
if strings.TrimPrefix(parser.Schema(), "v") != hint {
continue
}
regions, err := parser.Parse(raw)
return regions, parser.Schema(), err
}
}
for _, parser := range parsers {
if parser.CanParse(raw) {
regions, err := parser.Parse(raw)
return regions, parser.Schema(), err
}
}
return nil, "", NewError(CodeCatalogSchemaUnsupported, "This PIA server-list schema is not supported.")
}
func schemaVersion(raw []byte) int {
version, present, err := detectSchemaVersion(raw)
if err != nil || !present {
return 0
}
return version
}
func detectSchemaVersion(raw []byte) (int, bool, error) {
var envelope struct {
Version json.RawMessage `json:"version"`
}
if err := json.Unmarshal(raw, &envelope); err != nil {
return 0, false, err
}
if len(envelope.Version) == 0 || string(envelope.Version) == "null" {
return 0, false, nil
}
var number int
if json.Unmarshal(envelope.Version, &number) == nil {
if number < 1 {
return 0, true, fmt.Errorf("version must be positive")
}
return number, true, nil
}
var text string
if json.Unmarshal(envelope.Version, &text) == nil {
text = strings.TrimPrefix(strings.ToLower(text), "v")
if _, err := fmt.Sscanf(text, "%d", &number); err == nil && fmt.Sprint(number) == text && number > 0 {
return number, true, nil
}
}
return 0, true, fmt.Errorf("version has an unsupported type or value")
}
func parseCatalog(raw []byte, allowV7Aliases bool) ([]Region, error) {
var envelope catalogEnvelope
if err := decodeSingleJSON(raw, &envelope); err != nil {
return nil, WrapError(CodeCatalogSchemaUnsupported, "PIA returned an invalid region list.", err)
}
if len(envelope.Groups) == 0 || len(envelope.Regions) == 0 {
return nil, NewError(CodeCatalogSchemaUnsupported, "The PIA region list is missing required fields.")
}
seen := make(map[string]struct{}, len(envelope.Regions))
regions := make([]Region, 0, len(envelope.Regions))
for _, rawRegion := range envelope.Regions {
if !regionIDPattern.MatchString(rawRegion.ID) || strings.TrimSpace(rawRegion.Name) == "" || len(rawRegion.Name) > 128 {
continue
}
idKey := strings.ToLower(rawRegion.ID)
if _, duplicate := seen[idKey]; duplicate {
continue
}
seen[idKey] = struct{}{}
if !countryCodePattern.MatchString(rawRegion.Country) || rawRegion.Geo == nil || rawRegion.Offline == nil {
continue
}
if *rawRegion.Offline {
continue
}
portForwarding := false
if rawRegion.PortForward != nil {
portForwarding = *rawRegion.PortForward
} else if allowV7Aliases && rawRegion.PortForwarding != nil {
portForwarding = *rawRegion.PortForwarding
}
servers := make([]WireGuardServer, 0, len(rawRegion.Servers.WireGuard))
for _, rawServer := range rawRegion.Servers.WireGuard {
hostname := rawServer.CN
if hostname == "" && allowV7Aliases {
hostname = rawServer.Hostname
}
ip, err := netip.ParseAddr(rawServer.IP)
if err != nil || !ip.Is4() || ip.IsUnspecified() || !validHostname(hostname) {
continue
}
servers = append(servers, WireGuardServer{Hostname: hostname, IP: ip})
}
if len(servers) == 0 {
continue
}
regions = append(regions, Region{
ID: rawRegion.ID, Name: rawRegion.Name, CountryCode: strings.ToUpper(rawRegion.Country), Geo: *rawRegion.Geo,
PortForwarding: portForwarding, WireGuard: servers,
})
}
if len(regions) == 0 {
return nil, NewError(CodeCatalogSchemaUnsupported, "The PIA region list contains no available WireGuard regions.")
}
sort.Slice(regions, func(i, j int) bool {
if regions[i].CountryCode == regions[j].CountryCode {
return regions[i].Name < regions[j].Name
}
return regions[i].CountryCode < regions[j].CountryCode
})
return regions, nil
}
+100
View File
@@ -0,0 +1,100 @@
package pia
import (
"os"
"path/filepath"
"testing"
)
func TestServerListAdapters(t *testing.T) {
tests := []struct {
file, hint, schema, id, hostname string
}{
{"v6_valid.json", "6", "v6", "us-east", "useast401"},
{"v7_valid.json", "7", "v7", "de-berlin", "berlin501"},
}
for _, test := range tests {
raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", test.file))
if err != nil {
t.Fatal(err)
}
regions, schema, err := ParseServerList(raw, test.hint)
if err != nil {
t.Fatalf("%s: %v", test.file, err)
}
if schema != test.schema || len(regions) != 1 || regions[0].ID != test.id || regions[0].WireGuard[0].Hostname != test.hostname {
t.Fatalf("unexpected parsed result for %s: schema=%s regions=%+v", test.file, schema, regions)
}
}
raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", "v7_valid.json"))
if err != nil {
t.Fatal(err)
}
regions, schema, err := ParseServerList(raw, "6")
if err != nil || schema != "v7" || regions[0].ID != "de-berlin" {
t.Fatalf("detected schema did not override a stale endpoint hint: schema=%q regions=%v err=%v", schema, regions, err)
}
legacy := []byte(`{"groups":{"wg":[]},"regions":[{"id":"legacy","name":"Legacy","country":"US","geo":false,"offline":false,"servers":{"wg":[{"ip":"198.51.100.9","cn":"legacy.example"}]}}]}`)
regions, schema, err = ParseServerList(legacy, "")
if err != nil || schema != "v6" || regions[0].ID != "legacy" {
t.Fatalf("versionless v6 fallback failed: schema=%q regions=%v err=%v", schema, regions, err)
}
}
func TestServerListRejectsMalformedFields(t *testing.T) {
raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", "malformed.json"))
if err != nil {
t.Fatal(err)
}
if _, _, err := ParseServerList(raw, "6"); err == nil || CodeOf(err) != CodeCatalogSchemaUnsupported {
t.Fatalf("expected %s for malformed server list, got %s: %v", CodeCatalogSchemaUnsupported, CodeOf(err), err)
}
}
func TestServerListRejectsUnsupportedDuplicateAndTrailingData(t *testing.T) {
tests := []struct {
name, raw, hint string
}{
{"unsupported schema", `{"version":99,"groups":{},"regions":[]}`, ""},
{"invalid version value", `{"version":"v7beta","groups":{"wg":[]},"regions":[]}`, ""},
{"trailing JSON", `{"version":6,"groups":{},"regions":[]} {}`, "6"},
{"wrong groups type", `{"version":6,"groups":[],"regions":[]}`, "6"},
{"wrong field type", `{"version":6,"groups":{"wg":[]},"regions":[{"id":7,"name":"One","country":"US","geo":false,"offline":false,"servers":{"wg":[]}}]}`, "6"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if _, _, err := ParseServerList([]byte(test.raw), test.hint); err == nil || CodeOf(err) != CodeCatalogSchemaUnsupported {
t.Fatalf("expected %s, got %s: %v", CodeCatalogSchemaUnsupported, CodeOf(err), err)
}
})
}
}
func TestServerListSkipsBadRows(t *testing.T) {
duplicate := []byte(`{"version":6,"groups":{"wg":[]},"regions":[{"id":"same","name":"One","country":"US","geo":false,"offline":false,"servers":{"wg":[{"ip":"198.51.100.1","cn":"one.example"}]}},{"id":"SAME","name":"Two","country":"US","geo":false,"offline":false,"servers":{"wg":[{"ip":"198.51.100.2","cn":"two.example"}]}}]}`)
regions, _, err := ParseServerList(duplicate, "6")
if err != nil || len(regions) != 1 || regions[0].ID != "same" || regions[0].WireGuard[0].Hostname != "one.example" {
t.Fatalf("duplicate region id should keep the first: regions=%+v err=%v", regions, err)
}
mixed := []byte(`{"version":6,"groups":{"wg":[]},"regions":[{"id":"us-east","name":"US East","country":"US","geo":false,"offline":false,"servers":{"wg":[{"ip":"2001:db8::1","cn":"bad6"},{"ip":"198.51.100.10","cn":"useast1"}]}}]}`)
regions, _, err = ParseServerList(mixed, "6")
if err != nil || len(regions) != 1 || len(regions[0].WireGuard) != 1 || regions[0].WireGuard[0].Hostname != "useast1" {
t.Fatalf("invalid WireGuard server should be skipped: regions=%+v err=%v", regions, err)
}
}
func FuzzParseServerList(f *testing.F) {
for _, name := range []string{"v6_valid.json", "v7_valid.json", "malformed.json"} {
raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", name))
if err != nil {
f.Fatal(err)
}
f.Add(raw)
}
f.Fuzz(func(t *testing.T, raw []byte) {
_, _, _ = ParseServerList(raw, "6")
})
}
+63
View File
@@ -0,0 +1,63 @@
package pia
import (
"bytes"
"crypto"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/pem"
"fmt"
"unicode"
)
func VerifySignedServerList(raw, publicKeyPEM []byte) ([]byte, error) {
jsonBody, signature, err := splitSignedServerList(raw)
if err != nil {
return nil, WrapError(CodeCatalogSignatureInvalid, "The PIA region list signature is missing or invalid.", err)
}
block, _ := pem.Decode(publicKeyPEM)
if block == nil {
return nil, NewError(CodeCatalogSignatureInvalid, "The built-in region-list public key is invalid.")
}
parsed, err := x509.ParsePKIXPublicKey(block.Bytes)
if err != nil {
return nil, WrapError(CodeCatalogSignatureInvalid, "The built-in region-list public key is invalid.", err)
}
publicKey, ok := parsed.(*rsa.PublicKey)
if !ok {
return nil, NewError(CodeCatalogSignatureInvalid, "The region-list public key is not RSA.")
}
digest := sha256.Sum256(jsonBody)
if err := rsa.VerifyPKCS1v15(publicKey, crypto.SHA256, digest[:], signature); err != nil {
return nil, WrapError(CodeCatalogSignatureInvalid, "The PIA region list signature does not match its content.", err)
}
return jsonBody, nil
}
func splitSignedServerList(raw []byte) ([]byte, []byte, error) {
if len(raw) == 0 || raw[0] != '{' {
return nil, nil, fmt.Errorf("response does not start with a JSON object")
}
end := bytes.LastIndexByte(raw, '}')
if end < 0 || end == len(raw)-1 {
return nil, nil, fmt.Errorf("appended signature is absent")
}
jsonBody := append([]byte(nil), raw[:end+1]...)
encoded := bytes.Map(func(r rune) rune {
if unicode.IsSpace(r) {
return -1
}
return r
}, raw[end+1:])
if len(encoded) == 0 {
return nil, nil, fmt.Errorf("appended signature is empty")
}
signature := make([]byte, base64.StdEncoding.DecodedLen(len(encoded)))
n, err := base64.StdEncoding.Decode(signature, encoded)
if err != nil {
return nil, nil, fmt.Errorf("decode signature: %w", err)
}
return jsonBody, signature[:n], nil
}
+63
View File
@@ -0,0 +1,63 @@
package pia
import (
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/pem"
"os"
"path/filepath"
"testing"
)
func TestVerifySignedServerList(t *testing.T) {
payload := []byte(`{"version":6,"groups":{},"regions":[]}`)
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
digest := sha256.Sum256(payload)
signature, err := rsa.SignPKCS1v15(rand.Reader, privateKey, crypto.SHA256, digest[:])
if err != nil {
t.Fatal(err)
}
publicDER, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey)
if err != nil {
t.Fatal(err)
}
publicPEM := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: publicDER})
signed := append(append(append([]byte{}, payload...), '\n', '\n'), []byte(base64.StdEncoding.EncodeToString(signature))...)
validSigned := append([]byte(nil), signed...)
verified, err := VerifySignedServerList(signed, publicPEM)
if err != nil {
t.Fatal(err)
}
if string(verified) != string(payload) {
t.Fatalf("verified payload changed: %s", verified)
}
signed[10] ^= 1
if _, err := VerifySignedServerList(signed, publicPEM); err == nil {
t.Fatal("expected tampered payload to fail signature verification")
}
for name, input := range map[string][]byte{
"missing signature": payload,
"invalid base64": append(append([]byte{}, payload...), []byte("\nnot-base64!")...),
"trailing garbage": append(append([]byte{}, validSigned...), []byte("\nextra")...),
} {
t.Run(name, func(t *testing.T) {
if _, err := VerifySignedServerList(input, publicPEM); err == nil {
t.Fatal("expected malformed signed response to be rejected")
}
})
}
fixture, err := os.ReadFile(filepath.Join("testdata", "serverlist", "invalid_signature.txt"))
if err != nil {
t.Fatal(err)
}
if _, err := VerifySignedServerList(fixture, publicPEM); err == nil {
t.Fatal("expected invalid-signature fixture to be rejected")
}
}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "OK",
"peer_ip": "10.0.0.2/32",
"server_key": "AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=",
"server_ip": "198.51.100.10",
"server_port": 51820,
"dns_servers": ["not-an-ip"]
}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "OK",
"peer_ip": "not-an-ip",
"server_key": "AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=",
"server_ip": "198.51.100.10",
"server_port": 51820,
"dns_servers": ["10.0.0.1"]
}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "OK",
"peer_ip": "10.0.0.0/24",
"server_key": "AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=",
"server_ip": "198.51.100.10",
"server_port": 51820,
"dns_servers": ["10.0.0.1"]
}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "OK",
"peer_ip": "10.0.0.2/32",
"server_key": "AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=",
"server_ip": "198.51.100.10",
"server_port": 65536,
"dns_servers": ["10.0.0.1"]
}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "OK",
"peer_ip": "10.0.0.2/32",
"server_key": "AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=",
"server_ip": "not-an-ip",
"server_port": 51820,
"dns_servers": ["10.0.0.1"]
}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "OK",
"peer_ip": "10.0.0.2/32",
"server_key": "not-a-wireguard-key",
"server_ip": "198.51.100.10",
"server_port": 51820,
"dns_servers": ["10.0.0.1"]
}
+1
View File
@@ -0,0 +1 @@
{"status":"OK","peer_ip":"10.42.0.2/32","server_key":"AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=","server_ip":"198.51.100.10","dns_servers":["10.0.0.242"]}
+8
View File
@@ -0,0 +1,8 @@
{
"status": "ERROR",
"peer_ip": "10.0.0.2/32",
"server_key": "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=",
"server_ip": "198.51.100.10",
"server_port": 51820,
"dns_servers": ["10.0.0.1"]
}
+1
View File
@@ -0,0 +1 @@
{"status":"OK","server_key":"AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=","server_port":51820,"server_ip":"198.51.100.10","server_vip":"10.42.0.1","peer_ip":"10.42.0.2","peer_pubkey":"AQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=","dns_servers":["10.0.0.242","10.0.0.243"]}
+1
View File
@@ -0,0 +1 @@
<!doctype html><title>upstream error</title>
+1
View File
@@ -0,0 +1 @@
{"token":"test-token-value-that-is-long-enough"}
@@ -0,0 +1,3 @@
{"version":6,"groups":{},"regions":[]}
QUFBQQ==
+1
View File
@@ -0,0 +1 @@
{"version":6,"groups":{},"regions":[{"id":12}]}
+1
View File
@@ -0,0 +1 @@
{"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":"useast401"}]}}]}
+1
View File
@@ -0,0 +1 @@
{"version":"v7","groups":{"wg":[{"name":"wireguard","ports":[1337]}]},"regions":[{"id":"de-berlin","name":"Germany Berlin","country":"DE","geo":true,"offline":false,"port_forwarding":false,"servers":{"wg":[{"ip":"203.0.113.20","hostname":"berlin501"}]}}]}
+43
View File
@@ -0,0 +1,43 @@
-----BEGIN CERTIFICATE-----
MIIHqzCCBZOgAwIBAgIJAJ0u+vODZJntMA0GCSqGSIb3DQEBDQUAMIHoMQswCQYD
VQQGEwJVUzELMAkGA1UECBMCQ0ExEzARBgNVBAcTCkxvc0FuZ2VsZXMxIDAeBgNV
BAoTF1ByaXZhdGUgSW50ZXJuZXQgQWNjZXNzMSAwHgYDVQQLExdQcml2YXRlIElu
dGVybmV0IEFjY2VzczEgMB4GA1UEAxMXUHJpdmF0ZSBJbnRlcm5ldCBBY2Nlc3Mx
IDAeBgNVBCkTF1ByaXZhdGUgSW50ZXJuZXQgQWNjZXNzMS8wLQYJKoZIhvcNAQkB
FiBzZWN1cmVAcHJpdmF0ZWludGVybmV0YWNjZXNzLmNvbTAeFw0xNDA0MTcxNzQw
MzNaFw0zNDA0MTIxNzQwMzNaMIHoMQswCQYDVQQGEwJVUzELMAkGA1UECBMCQ0Ex
EzARBgNVBAcTCkxvc0FuZ2VsZXMxIDAeBgNVBAoTF1ByaXZhdGUgSW50ZXJuZXQg
QWNjZXNzMSAwHgYDVQQLExdQcml2YXRlIEludGVybmV0IEFjY2VzczEgMB4GA1UE
AxMXUHJpdmF0ZSBJbnRlcm5ldCBBY2Nlc3MxIDAeBgNVBCkTF1ByaXZhdGUgSW50
ZXJuZXQgQWNjZXNzMS8wLQYJKoZIhvcNAQkBFiBzZWN1cmVAcHJpdmF0ZWludGVy
bmV0YWNjZXNzLmNvbTCCAiIwDQYJKoZIhvcNAQEBBQADggIPADCCAgoCggIBALVk
hjumaqBbL8aSgj6xbX1QPTfTd1qHsAZd2B97m8Vw31c/2yQgZNf5qZY0+jOIHULN
De4R9TIvyBEbvnAg/OkPw8n/+ScgYOeH876VUXzjLDBnDb8DLr/+w9oVsuDeFJ9K
V2UFM1OYX0SnkHnrYAN2QLF98ESK4NCSU01h5zkcgmQ+qKSfA9Ny0/UpsKPBFqsQ
25NvjDWFhCpeqCHKUJ4Be27CDbSl7lAkBuHMPHJs8f8xPgAbHRXZOxVCpayZ2SND
fCwsnGWpWFoMGvdMbygngCn6jA/W1VSFOlRlfLuuGe7QFfDwA0jaLCxuWt/BgZyl
p7tAzYKR8lnWmtUCPm4+BtjyVDYtDCiGBD9Z4P13RFWvJHw5aapx/5W/CuvVyI7p
Kwvc2IT+KPxCUhH1XI8ca5RN3C9NoPJJf6qpg4g0rJH3aaWkoMRrYvQ+5PXXYUzj
tRHImghRGd/ydERYoAZXuGSbPkm9Y/p2X8unLcW+F0xpJD98+ZI+tzSsI99Zs5wi
jSUGYr9/j18KHFTMQ8n+1jauc5bCCegN27dPeKXNSZ5riXFL2XX6BkY68y58UaNz
meGMiUL9BOV1iV+PMb7B7PYs7oFLjAhh0EdyvfHkrh/ZV9BEhtFa7yXp8XR0J6vz
1YV9R6DYJmLjOEbhU8N0gc3tZm4Qz39lIIG6w3FDAgMBAAGjggFUMIIBUDAdBgNV
HQ4EFgQUrsRtyWJftjpdRM0+925Y6Cl08SUwggEfBgNVHSMEggEWMIIBEoAUrsRt
yWJftjpdRM0+925Y6Cl08SWhge6kgeswgegxCzAJBgNVBAYTAlVTMQswCQYDVQQI
EwJDQTETMBEGA1UEBxMKTG9zQW5nZWxlczEgMB4GA1UEChMXUHJpdmF0ZSBJbnRl
cm5ldCBBY2Nlc3MxIDAeBgNVBAsTF1ByaXZhdGUgSW50ZXJuZXQgQWNjZXNzMSAw
HgYDVQQDExdQcml2YXRlIEludGVybmV0IEFjY2VzczEgMB4GA1UEKRMXUHJpdmF0
ZSBJbnRlcm5ldCBBY2Nlc3MxLzAtBgkqhkiG9w0BCQEWIHNlY3VyZUBwcml2YXRl
aW50ZXJuZXRhY2Nlc3MuY29tggkAnS7684Nkme0wDAYDVR0TBAUwAwEB/zANBgkq
hkiG9w0BAQ0FAAOCAgEAJsfhsPk3r8kLXLxY+v+vHzbr4ufNtqnL9/1Uuf8NrsCt
pXAoyZ0YqfbkWx3NHTZ7OE9ZRhdMP/RqHQE1p4N4Sa1nZKhTKasV6KhHDqSCt/dv
Em89xWm2MVA7nyzQxVlHa9AkcBaemcXEiyT19XdpiXOP4Vhs+J1R5m8zQOxZlV1G
tF9vsXmJqWZpOVPmZ8f35BCsYPvv4yMewnrtAC8PFEK/bOPeYcKN50bol22QYaZu
LfpkHfNiFTnfMh8sl/ablPyNY7DUNiP5DRcMdIwmfGQxR5WEQoHL3yPJ42LkB5zs
6jIm26DGNXfwura/mi105+ENH1CaROtRYwkiHb08U6qLXXJz80mWJkT90nr8Asj3
5xN2cUppg74nG3YVav/38P48T56hG1NHbYF5uOCske19F6wi9maUoto/3vEr0rnX
JUp2KODmKdvBI7co245lHBABWikk8VfejQSlCtDBXn644ZMtAdoxKNfR2WTFVEwJ
iyd1Fzx0yujuiXDROLhISLQDRjVVAvawrAtLZWYK31bY7KlezPlQnl/D9Asxe85l
8jO5+0LdJ6VyOs/Hd4w52alDW/MFySDZSfQHMTIc30hLBJ8OnCEIvluVQQ2UQvoW
+no177N9L2Y+M9TcTA62ZyMXShHQGeh20rb4kK8f+iFX8NxtdHVSkxMEFSfDDyQ=
-----END CERTIFICATE-----
@@ -0,0 +1,9 @@
-----BEGIN PUBLIC KEY-----
MIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEAzLYHwX5Ug/oUObZ5eH5P
rEwmfj4E/YEfSKLgFSsyRGGsVmmjiXBmSbX2s3xbj/ofuvYtkMkP/VPFHy9E/8ox
Y+cRjPzydxz46LPY7jpEw1NHZjOyTeUero5e1nkLhiQqO/cMVYmUnuVcuFfZyZvc
8Apx5fBrIp2oWpF/G9tpUZfUUJaaHiXDtuYP8o8VhYtyjuUu3h7rkQFoMxvuoOFH
6nkc0VQmBsHvCfq4T9v8gyiBtQRy543leapTBMT34mxVIQ4ReGLPVit/6sNLoGLb
gSnGe9Bk/a5V/5vlqeemWF0hgoRtUxMtU1hFbe7e8tSq1j+mu0SHMyKHiHd+OsmU
IQIDAQAB
-----END PUBLIC KEY-----
+52
View File
@@ -0,0 +1,52 @@
package pia
import (
"bytes"
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"testing"
"time"
)
func TestEmbeddedTrustAnchorsAreUsable(t *testing.T) {
t.Run("server list public key", func(t *testing.T) {
block, rest := pem.Decode(EmbeddedServerListPublicKey)
if block == nil {
t.Fatal("serverlist_public_key.pem does not decode as PEM")
}
if block.Type != "PUBLIC KEY" {
t.Fatalf("PEM block type = %q, want PUBLIC KEY", block.Type)
}
if len(bytes.TrimSpace(rest)) != 0 {
t.Fatalf("trailing data after the public key: %q", rest)
}
parsed, err := x509.ParsePKIXPublicKey(block.Bytes)
if err != nil {
t.Fatalf("ParsePKIXPublicKey: %v", err)
}
if _, ok := parsed.(*rsa.PublicKey); !ok {
t.Fatalf("public key type = %T, want *rsa.PublicKey", parsed)
}
})
t.Run("addKey certificate authority", func(t *testing.T) {
if !x509.NewCertPool().AppendCertsFromPEM(EmbeddedPIACA) {
t.Fatal("ca.rsa.4096.crt was rejected by AppendCertsFromPEM")
}
block, _ := pem.Decode(EmbeddedPIACA)
if block == nil {
t.Fatal("ca.rsa.4096.crt does not decode as PEM")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
t.Fatalf("ParseCertificate: %v", err)
}
if !cert.IsCA {
t.Fatal("the embedded PIA certificate is not a CA")
}
if !cert.NotAfter.After(time.Now()) {
t.Fatalf("the embedded PIA CA expired on %s", cert.NotAfter)
}
})
}
+43
View File
@@ -0,0 +1,43 @@
// Package pia is a standalone PIA WireGuard control-plane client.
package pia
import (
"context"
"net/netip"
"time"
)
type Region struct {
ID string
Name string
CountryCode string
Geo bool
PortForwarding bool
WireGuard []WireGuardServer
}
type WireGuardServer struct {
Hostname string
IP netip.Addr
}
type Token struct {
Value []byte
ExpiresAt time.Time
}
type Registration struct {
PeerIP netip.Prefix
ServerKey string
ServerIP netip.Addr
ServerPort uint16
DNSServers []netip.Addr
}
type Authenticator interface {
Authenticate(ctx context.Context, username string, password []byte) (Token, error)
}
type Registrar interface {
RegisterKey(ctx context.Context, server WireGuardServer, token string, publicKey string) (Registration, error)
}
+42
View File
@@ -0,0 +1,42 @@
package pia
import (
"encoding/base64"
"net"
"strings"
"unicode"
)
func validSecret(value []byte, min, max int) bool {
if len(value) < min || len(value) > max {
return false
}
for _, b := range value {
if b == 0 || b == '\r' || b == '\n' {
return false
}
}
return true
}
func validHostname(host string) bool {
if host == "" || len(host) > 253 || net.ParseIP(host) != nil || strings.HasSuffix(host, ".") {
return false
}
for _, label := range strings.Split(host, ".") {
if label == "" || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' {
return false
}
for _, r := range label {
if r > unicode.MaxASCII || (!unicode.IsLetter(r) && !unicode.IsDigit(r) && r != '-') {
return false
}
}
}
return true
}
func validWGKey(key string) bool {
decoded, err := base64.StdEncoding.DecodeString(key)
return err == nil && len(decoded) == 32
}