This commit is contained in:
jdl
2026-06-10 19:17:55 +02:00
parent db78964033
commit f89f5d5083
7 changed files with 260 additions and 65 deletions

View File

@@ -9,7 +9,6 @@ import (
"net/http"
"net/netip"
"os"
"path/filepath"
"golang.org/x/crypto/nacl/sign"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
@@ -32,21 +31,27 @@ type LocalState struct {
// localStateJSON is the on-disk representation.
type localStateJSON struct {
PrivKey string `json:"priv_key"` // standard base64
SignKey string `json:"sign_key"` // standard base64
VPNIP netip.Addr `json:"vpn_ip"`
VPNNet netip.Prefix `json:"vpn_net"`
WGPort uint16 `json:"wg_port"`
IsRelay bool `json:"is_relay"`
IsPublic bool `json:"is_public"`
LocalDomain string `json:"local_domain"`
PrivKey string
SignKey string
VPNIP netip.Addr
VPNNet netip.Prefix
WGPort uint16
IsRelay bool
IsPublic bool
LocalDomain string
}
// LoadOrInit loads LocalState from path, or registers with the hub and creates
// the file if it doesn't exist.
func LoadOrInit(statePath, hubURL, apiKey string) (LocalState, error) {
if data, err := os.ReadFile(statePath); err == nil {
return parseLocalState(data)
var state LocalState
switch err := loadJSON(statePath, &state); {
case err == nil:
return state, nil
case !os.IsNotExist(err):
// File exists but is unreadable/corrupt: surface it rather than
// silently regenerating a new identity and re-registering.
return LocalState{}, fmt.Errorf("load state: %w", err)
}
privKey, err := wgtypes.GeneratePrivateKey()
@@ -54,12 +59,12 @@ func LoadOrInit(statePath, hubURL, apiKey string) (LocalState, error) {
return LocalState{}, fmt.Errorf("generate key: %w", err)
}
state, err := initFromHub(hubURL, apiKey, privKey)
state, err = initFromHub(hubURL, apiKey, privKey)
if err != nil {
return LocalState{}, err
}
if err := saveLocalState(statePath, state); err != nil {
if err := storeJSON(statePath, state); err != nil {
return LocalState{}, fmt.Errorf("save state: %w", err)
}
return state, nil
@@ -130,42 +135,8 @@ func initFromHub(hubURL, apiKey string, privKey wgtypes.Key) (LocalState, error)
}, nil
}
func parseLocalState(data []byte) (LocalState, error) {
var j localStateJSON
if err := json.Unmarshal(data, &j); err != nil {
return LocalState{}, fmt.Errorf("parse state: %w", err)
}
keyBytes, err := base64.StdEncoding.DecodeString(j.PrivKey)
if err != nil {
return LocalState{}, fmt.Errorf("decode key: %w", err)
}
key, err := wgtypes.NewKey(keyBytes)
if err != nil {
return LocalState{}, fmt.Errorf("invalid key: %w", err)
}
signKeyBytes, err := base64.StdEncoding.DecodeString(j.SignKey)
if err != nil {
return LocalState{}, fmt.Errorf("decode sign key: %w", err)
}
if len(signKeyBytes) != 64 {
return LocalState{}, fmt.Errorf("invalid sign key length: %d", len(signKeyBytes))
}
var signKey [64]byte
copy(signKey[:], signKeyBytes)
return LocalState{
PrivKey: key,
SignKey: signKey,
VPNIP: j.VPNIP,
VPNNet: j.VPNNet,
WGPort: j.WGPort,
IsRelay: j.IsRelay,
IsPublic: j.IsPublic,
LocalDomain: j.LocalDomain,
}, nil
}
func saveLocalState(path string, s LocalState) error {
j := localStateJSON{
func (s LocalState) MarshalJSON() ([]byte, error) {
return json.Marshal(localStateJSON{
PrivKey: base64.StdEncoding.EncodeToString(s.PrivKey[:]),
SignKey: base64.StdEncoding.EncodeToString(s.SignKey[:]),
VPNIP: s.VPNIP,
@@ -174,13 +145,38 @@ func saveLocalState(path string, s LocalState) error {
IsRelay: s.IsRelay,
IsPublic: s.IsPublic,
LocalDomain: s.LocalDomain,
}
data, err := json.MarshalIndent(j, "", " ")
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(path), 0700); err != nil {
return err
}
return os.WriteFile(path, data, 0600)
})
}
func (s *LocalState) UnmarshalJSON(data []byte) error {
var j localStateJSON
if err := json.Unmarshal(data, &j); err != nil {
return err
}
keyBytes, err := base64.StdEncoding.DecodeString(j.PrivKey)
if err != nil {
return fmt.Errorf("decode key: %w", err)
}
key, err := wgtypes.NewKey(keyBytes)
if err != nil {
return fmt.Errorf("invalid key: %w", err)
}
signKeyBytes, err := base64.StdEncoding.DecodeString(j.SignKey)
if err != nil {
return fmt.Errorf("decode sign key: %w", err)
}
if len(signKeyBytes) != 64 {
return fmt.Errorf("invalid sign key length: %d", len(signKeyBytes))
}
*s = LocalState{
PrivKey: key,
SignKey: [64]byte(signKeyBytes),
VPNIP: j.VPNIP,
VPNNet: j.VPNNet,
WGPort: j.WGPort,
IsRelay: j.IsRelay,
IsPublic: j.IsPublic,
LocalDomain: j.LocalDomain,
}
return nil
}