mirror of
https://github.com/fosrl/newt.git
synced 2026-10-01 02:09:07 +02:00
+20
-6
@@ -14,6 +14,7 @@ import (
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"software.sslmate.com/src/go-pkcs12"
|
||||
@@ -51,10 +52,11 @@ type Client struct {
|
||||
serverVersion string
|
||||
configVersion int64 // Latest config version received from server
|
||||
configVersionMux sync.RWMutex
|
||||
processingMessage bool // Flag to track if a message is currently being processed
|
||||
processingMux sync.RWMutex // Protects processingMessage
|
||||
processingWg sync.WaitGroup // WaitGroup to wait for message processing to complete
|
||||
justProvisioned bool // Set to true when provisionIfNeeded exchanges a key for permanent credentials
|
||||
processingMessage bool // Flag to track if a message is currently being processed
|
||||
processingMux sync.RWMutex // Protects processingMessage
|
||||
processingWg sync.WaitGroup // WaitGroup to wait for message processing to complete
|
||||
justProvisioned bool // Set to true when provisionIfNeeded exchanges a key for permanent credentials
|
||||
consecutiveFailures atomic.Int32 // Counts consecutive connection failures for log suppression
|
||||
}
|
||||
|
||||
type ClientOption func(*Client)
|
||||
@@ -419,7 +421,7 @@ func (c *Client) getToken() (string, error) {
|
||||
logger.Debug("Token response body: %s", string(body))
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
logger.Error("Failed to get token with status code: %d", resp.StatusCode)
|
||||
logger.Debug("Failed to get token with status code: %d", resp.StatusCode)
|
||||
telemetry.IncConnAttempt(ctx, "auth", "failure")
|
||||
etype := "io_error"
|
||||
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
|
||||
@@ -502,6 +504,12 @@ func classifyWSDisconnect(err error) (result, reason string) {
|
||||
}
|
||||
}
|
||||
|
||||
// consecutiveFailureThreshold is the number of consecutive failures before
|
||||
// connection errors are promoted from Debug to Error. This suppresses the noisy
|
||||
// transient errors that occur during routine server updates while still surfacing
|
||||
// genuine connectivity problems.
|
||||
const consecutiveFailureThreshold = 3
|
||||
|
||||
func (c *Client) connectWithRetry() {
|
||||
for {
|
||||
select {
|
||||
@@ -510,10 +518,16 @@ func (c *Client) connectWithRetry() {
|
||||
default:
|
||||
err := c.establishConnection()
|
||||
if err != nil {
|
||||
logger.Error("Failed to connect: %v. Retrying in %v...", err, c.reconnectInterval)
|
||||
n := c.consecutiveFailures.Add(1)
|
||||
if n >= consecutiveFailureThreshold {
|
||||
logger.Error("Failed to connect: %v. Retrying in %v...", err, c.reconnectInterval)
|
||||
} else {
|
||||
logger.Debug("Failed to connect (attempt %d): %v. Retrying in %v...", n, err, c.reconnectInterval)
|
||||
}
|
||||
time.Sleep(c.reconnectInterval)
|
||||
continue
|
||||
}
|
||||
c.consecutiveFailures.Store(0)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
+48
-59
@@ -19,6 +19,10 @@ import (
|
||||
"github.com/fosrl/newt/logger"
|
||||
)
|
||||
|
||||
// getConfigPath returns the resolved config/credentials file path. Path
|
||||
// resolution (CLI flag / env var / OS default) is now owned by the calling
|
||||
// binary's root config loader; this is a fallback used only when a caller
|
||||
// (e.g. a test) doesn't supply an explicit path via WithConfigFile.
|
||||
func getConfigPath(clientType string, overridePath string) string {
|
||||
if overridePath != "" {
|
||||
return overridePath
|
||||
@@ -46,21 +50,14 @@ func getConfigPath(clientType string, overridePath string) string {
|
||||
return configFile
|
||||
}
|
||||
|
||||
// loadConfig no longer merges credentials in from the file: the caller
|
||||
// resolves ID/Secret/Endpoint/TlsClientCert/ProvisioningKey/Name (from its
|
||||
// own settings file, env, and CLI flags) before constructing the Client. This
|
||||
// only figures out whether the file needs to be (re)written so that the
|
||||
// caller-provided values get cached to disk on first successful connect.
|
||||
func (c *Client) loadConfig() error {
|
||||
originalConfig := *c.config // Store original config to detect changes
|
||||
configPath := getConfigPath(c.clientType, c.configFilePath)
|
||||
|
||||
if c.config.ID != "" && c.config.Secret != "" && c.config.Endpoint != "" {
|
||||
logger.Debug("Config already provided, skipping loading from file")
|
||||
// Check if config file exists, if not, we should save it
|
||||
if _, err := os.Stat(configPath); os.IsNotExist(err) {
|
||||
logger.Info("Config file does not exist at %s, will create it", configPath)
|
||||
c.configNeedsSave = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
logger.Info("Loading config from: %s", configPath)
|
||||
data, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
@@ -73,57 +70,16 @@ func (c *Client) loadConfig() error {
|
||||
if len(bytes.TrimSpace(data)) == 0 {
|
||||
logger.Info("Config file at %s is empty, will initialize it with provided values", configPath)
|
||||
c.configNeedsSave = true
|
||||
return nil
|
||||
}
|
||||
|
||||
var config Config
|
||||
if err := json.Unmarshal(data, &config); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Track what was loaded from file vs provided by CLI
|
||||
fileHadID := c.config.ID == ""
|
||||
fileHadSecret := c.config.Secret == ""
|
||||
fileHadCert := c.config.TlsClientCert == ""
|
||||
fileHadEndpoint := c.config.Endpoint == ""
|
||||
|
||||
if c.config.ID == "" {
|
||||
c.config.ID = config.ID
|
||||
}
|
||||
if c.config.Secret == "" {
|
||||
c.config.Secret = config.Secret
|
||||
}
|
||||
if c.config.TlsClientCert == "" {
|
||||
c.config.TlsClientCert = config.TlsClientCert
|
||||
}
|
||||
if c.config.Endpoint == "" {
|
||||
c.config.Endpoint = config.Endpoint
|
||||
c.baseURL = config.Endpoint
|
||||
}
|
||||
// Always load the provisioning key from the file if not already set
|
||||
if c.config.ProvisioningKey == "" {
|
||||
c.config.ProvisioningKey = config.ProvisioningKey
|
||||
}
|
||||
// Always load the name from the file if not already set
|
||||
if c.config.Name == "" {
|
||||
c.config.Name = config.Name
|
||||
}
|
||||
|
||||
// Check if CLI args provided values that override file values
|
||||
if (!fileHadID && originalConfig.ID != "") ||
|
||||
(!fileHadSecret && originalConfig.Secret != "") ||
|
||||
(!fileHadCert && originalConfig.TlsClientCert != "") ||
|
||||
(!fileHadEndpoint && originalConfig.Endpoint != "") {
|
||||
logger.Info("CLI arguments provided, config will be updated")
|
||||
c.configNeedsSave = true
|
||||
}
|
||||
|
||||
logger.Debug("Loaded config from %s", configPath)
|
||||
logger.Debug("Config: %+v", c.config)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// saveConfig persists the credential fields (id, secret, endpoint,
|
||||
// tlsClientCert, provisioningKey, name) into the config file. It's a
|
||||
// read-modify-write against the raw JSON so that any other settings a user
|
||||
// has placed in the same file (mtu, dns, etc.) are preserved rather than
|
||||
// clobbered.
|
||||
func (c *Client) saveConfig() error {
|
||||
if !c.configNeedsSave {
|
||||
logger.Debug("Config has not changed, skipping save")
|
||||
@@ -131,7 +87,31 @@ func (c *Client) saveConfig() error {
|
||||
}
|
||||
|
||||
configPath := getConfigPath(c.clientType, c.configFilePath)
|
||||
data, err := json.MarshalIndent(c.config, "", " ")
|
||||
|
||||
existing := map[string]interface{}{}
|
||||
if data, err := os.ReadFile(configPath); err == nil && len(bytes.TrimSpace(data)) > 0 {
|
||||
if err := json.Unmarshal(data, &existing); err != nil {
|
||||
logger.Warn("Existing config file at %s is not valid JSON, other settings in it will be lost: %v", configPath, err)
|
||||
existing = map[string]interface{}{}
|
||||
}
|
||||
}
|
||||
|
||||
existing["id"] = c.config.ID
|
||||
existing["secret"] = c.config.Secret
|
||||
existing["endpoint"] = c.config.Endpoint
|
||||
|
||||
setOrDelete := func(key, value string) {
|
||||
if value != "" {
|
||||
existing[key] = value
|
||||
} else {
|
||||
delete(existing, key)
|
||||
}
|
||||
}
|
||||
setOrDelete("tlsClientCert", c.config.TlsClientCert)
|
||||
setOrDelete("provisioningKey", c.config.ProvisioningKey)
|
||||
setOrDelete("name", c.config.Name)
|
||||
|
||||
data, err := json.MarshalIndent(existing, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -159,6 +139,15 @@ func interpolateString(s string) string {
|
||||
})
|
||||
}
|
||||
|
||||
// EnsureProvisioned exchanges a provisioning key for permanent credentials
|
||||
// synchronously, if one is configured and credentials aren't already present.
|
||||
// Callers that need the resolved ID/Secret before Connect() runs (which only
|
||||
// provisions lazily, in the background, on first connection attempt) should
|
||||
// call this first.
|
||||
func (c *Client) EnsureProvisioned() error {
|
||||
return c.provisionIfNeeded()
|
||||
}
|
||||
|
||||
// provisionIfNeeded checks whether a provisioning key is present and, if so,
|
||||
// exchanges it for a newt ID and secret by calling the registration endpoint.
|
||||
// On success the config is updated in-place and flagged for saving so that
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
@@ -33,3 +34,63 @@ func TestLoadConfig_EmptyFileMarksConfigForSave(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveConfig_PreservesUnrelatedSettings(t *testing.T) {
|
||||
t.Setenv("CONFIG_FILE", "")
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
configPath := filepath.Join(tmpDir, "config.json")
|
||||
initial := `{
|
||||
"mtu": 1300,
|
||||
"dns": "1.1.1.1",
|
||||
"provisioningKey": "spk-test"
|
||||
}`
|
||||
if err := os.WriteFile(configPath, []byte(initial), 0o644); err != nil {
|
||||
t.Fatalf("failed to create config file: %v", err)
|
||||
}
|
||||
|
||||
client := &Client{
|
||||
config: &Config{
|
||||
ID: "newt-id",
|
||||
Secret: "newt-secret",
|
||||
Endpoint: "https://example.com",
|
||||
// ProvisioningKey cleared, simulating a completed provisioning exchange.
|
||||
},
|
||||
clientType: "newt",
|
||||
configFilePath: configPath,
|
||||
configNeedsSave: true,
|
||||
}
|
||||
|
||||
if err := client.saveConfig(); err != nil {
|
||||
t.Fatalf("saveConfig returned error: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read saved config: %v", err)
|
||||
}
|
||||
|
||||
var saved map[string]interface{}
|
||||
if err := json.Unmarshal(data, &saved); err != nil {
|
||||
t.Fatalf("saved config is not valid JSON: %v", err)
|
||||
}
|
||||
|
||||
if saved["mtu"] != float64(1300) {
|
||||
t.Errorf("expected mtu to be preserved, got %v", saved["mtu"])
|
||||
}
|
||||
if saved["dns"] != "1.1.1.1" {
|
||||
t.Errorf("expected dns to be preserved, got %v", saved["dns"])
|
||||
}
|
||||
if saved["id"] != "newt-id" {
|
||||
t.Errorf("expected id to be updated, got %v", saved["id"])
|
||||
}
|
||||
if saved["secret"] != "newt-secret" {
|
||||
t.Errorf("expected secret to be updated, got %v", saved["secret"])
|
||||
}
|
||||
if _, ok := saved["provisioningKey"]; ok {
|
||||
t.Errorf("expected provisioningKey to be cleared after provisioning, got %v", saved["provisioningKey"])
|
||||
}
|
||||
if client.configNeedsSave {
|
||||
t.Error("expected configNeedsSave to be reset after a successful save")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user