mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 01:36:37 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -0,0 +1,63 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package config provides shared configuration types for edge applications.
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// MillisecondDuration is a time.Duration that marshals to/from JSON as an integer number of milliseconds
|
||||
// or as a Go duration string (e.g. "1s", "500ms").
|
||||
type MillisecondDuration time.Duration
|
||||
|
||||
// Duration returns the underlying time.Duration value.
|
||||
func (d MillisecondDuration) Duration() time.Duration {
|
||||
return time.Duration(d)
|
||||
}
|
||||
|
||||
func (d MillisecondDuration) String() string {
|
||||
return time.Duration(d).String()
|
||||
}
|
||||
|
||||
// UnmarshalJSON decodes either a numeric millisecond value or a quoted Go duration string.
|
||||
func (d *MillisecondDuration) UnmarshalJSON(data []byte) error {
|
||||
raw := strings.TrimSpace(string(data))
|
||||
if raw == "" || raw == "null" {
|
||||
*d = 0
|
||||
return nil
|
||||
}
|
||||
if strings.HasPrefix(raw, "\"") {
|
||||
var text string
|
||||
if err := json.Unmarshal(data, &text); err != nil {
|
||||
return err
|
||||
}
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
*d = 0
|
||||
return nil
|
||||
}
|
||||
parsed, err := time.ParseDuration(text)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid duration string %q: %w", text, err)
|
||||
}
|
||||
*d = MillisecondDuration(parsed)
|
||||
return nil
|
||||
}
|
||||
ms, err := strconv.ParseInt(raw, 10, 64)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid duration milliseconds %q: %w", raw, err)
|
||||
}
|
||||
*d = MillisecondDuration(time.Duration(ms) * time.Millisecond)
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalJSON encodes the duration as an integer number of milliseconds.
|
||||
func (d MillisecondDuration) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(time.Duration(d).Milliseconds())
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestMillisecondDurationUnmarshalJSON(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want time.Duration
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "null", input: "null", want: 0},
|
||||
{name: "empty", input: `""`, want: 0},
|
||||
{name: "integer milliseconds", input: "30000", want: 30 * time.Second},
|
||||
{name: "duration string", input: `"5s"`, want: 5 * time.Second},
|
||||
{name: "invalid number", input: "not-a-number", wantErr: true},
|
||||
{name: "invalid duration string", input: `"not-a-duration"`, wantErr: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var d MillisecondDuration
|
||||
err := json.Unmarshal([]byte(tt.input), &d)
|
||||
if tt.wantErr {
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if d.Duration() != tt.want {
|
||||
t.Fatalf("got %s, want %s", d, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMillisecondDurationMarshalJSON(t *testing.T) {
|
||||
d := MillisecondDuration(7 * time.Second)
|
||||
data, err := json.Marshal(d)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(data) != "7000" {
|
||||
t.Fatalf("unexpected marshaled value: %s", string(data))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package heartbeat handles periodic heartbeat and update checks.
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
edgeupdater "Wavelet/openflare/share/edge/updater"
|
||||
)
|
||||
|
||||
// AutoUpdateSettings defines the settings for automatic edge updates.
|
||||
type AutoUpdateSettings struct {
|
||||
AutoUpdate bool
|
||||
UpdateNow bool
|
||||
UpdateRepo string
|
||||
UpdateChannel string
|
||||
UpdateTag string
|
||||
}
|
||||
|
||||
// TryAutoUpdate attempts to check and apply auto updates for the edge service.
|
||||
func TryAutoUpdate(ctx context.Context, updater *edgeupdater.Service, settings *AutoUpdateSettings, logLabel string) {
|
||||
if settings == nil || updater == nil {
|
||||
return
|
||||
}
|
||||
force := settings.UpdateNow
|
||||
shouldCheck := settings.AutoUpdate || force
|
||||
if !shouldCheck || strings.TrimSpace(settings.UpdateRepo) == "" {
|
||||
return
|
||||
}
|
||||
channel := "stable"
|
||||
if force && strings.TrimSpace(settings.UpdateChannel) != "" {
|
||||
channel = settings.UpdateChannel
|
||||
}
|
||||
slog.Info("checking for "+logLabel+" updates", "repo", settings.UpdateRepo, "channel", channel, "force", force)
|
||||
err := updater.CheckAndUpdate(ctx, settings.UpdateRepo, edgeupdater.UpdateOptions{
|
||||
Channel: channel,
|
||||
TagName: settings.UpdateTag,
|
||||
Force: force,
|
||||
})
|
||||
if err != nil {
|
||||
slog.Error(logLabel+" update check failed", "error", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// RunLoop invokes fn immediately, then on each interval tick until ctx is cancelled.
|
||||
func RunLoop(ctx context.Context, interval time.Duration, fn func(context.Context)) {
|
||||
fn(ctx)
|
||||
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
fn(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestRunLoopImmediateAndTicker(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
var calls atomic.Int32
|
||||
interval := 20 * time.Millisecond
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
RunLoop(ctx, interval, func(context.Context) {
|
||||
calls.Add(1)
|
||||
})
|
||||
close(done)
|
||||
}()
|
||||
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
if got := calls.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 immediate call, got %d", got)
|
||||
}
|
||||
|
||||
time.Sleep(35 * time.Millisecond)
|
||||
if got := calls.Load(); got < 2 {
|
||||
t.Fatalf("expected at least 2 calls after ticker, got %d", got)
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("RunLoop did not exit after context cancellation")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package httpclient provides an authenticated HTTP client for edge services.
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Client is an HTTP client wrapper for communicating with remote HTTP services.
|
||||
type Client struct {
|
||||
baseURL string
|
||||
token string
|
||||
authHeader string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
// New creates a new Client instance.
|
||||
func New(baseURL, token string, timeout time.Duration, authHeader string) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
token: token,
|
||||
authHeader: authHeader,
|
||||
httpClient: &http.Client{Timeout: timeout},
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the client auth token.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.token = strings.TrimSpace(token)
|
||||
slog.Debug("http client token updated")
|
||||
}
|
||||
|
||||
// GetJSON sends a GET request and decodes the response body into target.
|
||||
func (c *Client) GetJSON(ctx context.Context, path string, target any) error {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.setAuthHeader(req)
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
// PostJSON sends a POST request with JSON body and decodes the response body into target.
|
||||
func (c *Client) PostJSON(ctx context.Context, path string, body any, target any) error {
|
||||
data, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+path, bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
c.setAuthHeader(req)
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
// DoRaw performs an HTTP request with custom headers and returns the raw response.
|
||||
func (c *Client) DoRaw(ctx context.Context, method, path string, headers map[string]string) (*http.Response, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.setAuthHeader(req)
|
||||
for key, value := range headers {
|
||||
req.Header.Set(key, value)
|
||||
}
|
||||
return c.httpClient.Do(req)
|
||||
}
|
||||
|
||||
func (c *Client) setAuthHeader(req *http.Request) {
|
||||
if c.authHeader != "" {
|
||||
req.Header.Set(c.authHeader, c.token)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) do(req *http.Request, target any) error {
|
||||
res, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
slog.Error("http request failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
if err := res.Body.Close(); err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
body, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
slog.Error("http response read failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
if res.StatusCode != http.StatusOK {
|
||||
slog.Warn("http request returned non-200", "method", req.Method, "path", req.URL.Path, "status", res.Status)
|
||||
return ReadBodyError(body, res.Status)
|
||||
}
|
||||
if target == nil {
|
||||
return nil
|
||||
}
|
||||
if err = json.Unmarshal(body, target); err != nil {
|
||||
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// APIError creates a new API error with the given message if it is not empty.
|
||||
func APIError(msg string) error {
|
||||
if strings.TrimSpace(msg) == "" {
|
||||
return nil
|
||||
}
|
||||
return errors.New(msg)
|
||||
}
|
||||
|
||||
// ReadBodyError parses the error message from the response body, or returns the fallback message.
|
||||
func ReadBodyError(body []byte, fallback string) error {
|
||||
var errBody struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &errBody); err == nil && strings.TrimSpace(errBody.ErrorMsg) != "" {
|
||||
return errors.New(errBody.ErrorMsg)
|
||||
}
|
||||
return errors.New(fallback)
|
||||
}
|
||||
|
||||
// ReadHTTPError reads the error message from the HTTP response.
|
||||
func ReadHTTPError(res *http.Response) error {
|
||||
body, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
return errors.New(res.Status)
|
||||
}
|
||||
return ReadBodyError(body, res.Status)
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logging configures structured logging for edge applications.
|
||||
package logging
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Options holds configuration options for the structured logger.
|
||||
type Options struct {
|
||||
AddSource bool
|
||||
}
|
||||
|
||||
// Setup initialises the default slog handler using the given options and the LOG_LEVEL environment variable.
|
||||
func Setup(opts Options) {
|
||||
handlerOpts := &slog.HandlerOptions{
|
||||
AddSource: opts.AddSource,
|
||||
Level: ParseLevel(os.Getenv("LOG_LEVEL")),
|
||||
}
|
||||
handler := slog.NewTextHandler(os.Stdout, handlerOpts)
|
||||
slog.SetDefault(slog.New(handler))
|
||||
}
|
||||
|
||||
// ParseLevel converts a log-level string (e.g. "debug", "warn") to the corresponding slog.Level.
|
||||
func ParseLevel(value string) slog.Level {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "debug":
|
||||
return slog.LevelDebug
|
||||
case "warn", "warning":
|
||||
return slog.LevelWarn
|
||||
case "error":
|
||||
return slog.LevelError
|
||||
default:
|
||||
return slog.LevelInfo
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package nodeip detects the preferred public IP address for edge nodes.
|
||||
package nodeip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/share/geoip"
|
||||
"Wavelet/openflare/share/geoip/iputil"
|
||||
)
|
||||
|
||||
const (
|
||||
outboundIPLookupTimeout = 5 * time.Second
|
||||
publicIPPriorityScore = 2 // matches iputil.Score for public IPv4 addresses
|
||||
ipCacheTTL = 10 * time.Minute
|
||||
)
|
||||
|
||||
// LookupOutboundIP and LookupLocalIP are the provider functions used to detect the node's outbound/local IP.
|
||||
// They are package-level variables so they can be overridden in tests.
|
||||
var (
|
||||
LookupOutboundIP = geoip.GetOutboundIP
|
||||
LookupLocalIP = DetectLocal
|
||||
|
||||
cacheMu sync.RWMutex
|
||||
cachedIP string
|
||||
lastDetected time.Time
|
||||
)
|
||||
|
||||
// Detect returns the best available outbound or local IPv4 address for this node.
|
||||
func Detect() string {
|
||||
return DetectWithContext(context.Background())
|
||||
}
|
||||
|
||||
// DetectWithContext returns the best available outbound or local IPv4 address, respecting ctx for cancellation.
|
||||
func DetectWithContext(ctx context.Context) string {
|
||||
cacheMu.RLock()
|
||||
if cachedIP != "" && time.Since(lastDetected) < ipCacheTTL {
|
||||
ip := cachedIP
|
||||
cacheMu.RUnlock()
|
||||
return ip
|
||||
}
|
||||
cacheMu.RUnlock()
|
||||
|
||||
var ip string
|
||||
if ip = detectOutbound(ctx); ip == "" {
|
||||
ip = LookupLocalIP()
|
||||
}
|
||||
|
||||
if ip != "" {
|
||||
cacheMu.Lock()
|
||||
cachedIP = ip
|
||||
lastDetected = time.Now()
|
||||
cacheMu.Unlock()
|
||||
}
|
||||
return ip
|
||||
}
|
||||
|
||||
func detectOutbound(ctx context.Context) string {
|
||||
ctx, cancel := context.WithTimeout(ctx, outboundIPLookupTimeout)
|
||||
defer cancel()
|
||||
ip, err := LookupOutboundIP(ctx)
|
||||
if err != nil || ip == nil {
|
||||
return ""
|
||||
}
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
// DetectLocal returns the highest-priority non-loopback local IPv4 address found on system interfaces.
|
||||
func DetectLocal() string {
|
||||
interfaces, err := net.Interfaces()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
bestIP := ""
|
||||
bestPriority := -1
|
||||
for _, iface := range interfaces {
|
||||
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
|
||||
continue
|
||||
}
|
||||
addrs, err := iface.Addrs()
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, addr := range addrs {
|
||||
ipNet, ok := addr.(*net.IPNet)
|
||||
if !ok || ipNet.IP == nil || ipNet.IP.IsLoopback() {
|
||||
continue
|
||||
}
|
||||
ipv4 := ipNet.IP.To4()
|
||||
if ipv4 == nil {
|
||||
continue
|
||||
}
|
||||
priority := iputil.Score(ipv4)
|
||||
if priority > bestPriority {
|
||||
bestIP = ipv4.String()
|
||||
bestPriority = priority
|
||||
}
|
||||
if bestPriority == publicIPPriorityScore {
|
||||
return bestIP
|
||||
}
|
||||
}
|
||||
}
|
||||
return bestIP
|
||||
}
|
||||
|
||||
// ResetCacheForTest clears the cached IP.
|
||||
func ResetCacheForTest() {
|
||||
cacheMu.Lock()
|
||||
cachedIP = ""
|
||||
lastDetected = time.Time{}
|
||||
cacheMu.Unlock()
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package observability provides helpers that read Linux /proc and /sys metrics for system monitoring.
|
||||
package observability
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
const (
|
||||
memInfoMinFieldCount = 2
|
||||
cpuStatMinFieldCount = 5
|
||||
netDevMinFieldCount = 16
|
||||
diskStatsMinFieldCount = 14
|
||||
)
|
||||
|
||||
// ReadLinuxOSRelease returns the OS name and version from /etc/os-release.
|
||||
func ReadLinuxOSRelease() (string, string) {
|
||||
file, err := os.Open("/etc/os-release")
|
||||
if err != nil {
|
||||
return runtime.GOOS, ""
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
values := make(map[string]string)
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
}
|
||||
key, value, ok := strings.Cut(line, "=")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
values[key] = strings.Trim(value, `"`)
|
||||
}
|
||||
if pretty := strings.TrimSpace(values["PRETTY_NAME"]); pretty != "" {
|
||||
return pretty, strings.TrimSpace(values["VERSION_ID"])
|
||||
}
|
||||
name := strings.TrimSpace(values["NAME"])
|
||||
if name == "" {
|
||||
name = runtime.GOOS
|
||||
}
|
||||
return name, strings.TrimSpace(values["VERSION_ID"])
|
||||
}
|
||||
|
||||
// ReadLinuxCPUModel returns the CPU model name from /proc/cpuinfo.
|
||||
func ReadLinuxCPUModel() string {
|
||||
file, err := os.Open("/proc/cpuinfo")
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
if strings.HasPrefix(strings.ToLower(line), "model name") {
|
||||
_, value, ok := strings.Cut(line, ":")
|
||||
if ok {
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// ReadMemInfo returns total and used memory bytes from /proc/meminfo.
|
||||
func ReadMemInfo() (int64, int64) {
|
||||
file, err := os.Open("/proc/meminfo")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var memTotalKB int64
|
||||
var memAvailableKB int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
switch {
|
||||
case strings.HasPrefix(line, "MemTotal:"):
|
||||
memTotalKB = parseMemInfoValue(line)
|
||||
case strings.HasPrefix(line, "MemAvailable:"):
|
||||
memAvailableKB = parseMemInfoValue(line)
|
||||
}
|
||||
}
|
||||
|
||||
total := memTotalKB * 1024
|
||||
if total == 0 {
|
||||
return 0, 0
|
||||
}
|
||||
used := max(total-(memAvailableKB*1024), 0)
|
||||
return total, used
|
||||
}
|
||||
|
||||
func parseMemInfoValue(line string) int64 {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < memInfoMinFieldCount {
|
||||
return 0
|
||||
}
|
||||
value, err := strconv.ParseInt(fields[1], 10, 64)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// ReadLinuxUptimeSeconds returns system uptime in seconds from /proc/uptime.
|
||||
func ReadLinuxUptimeSeconds() int64 {
|
||||
content, err := os.ReadFile("/proc/uptime")
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
fields := strings.Fields(string(content))
|
||||
if len(fields) == 0 {
|
||||
return 0
|
||||
}
|
||||
value, err := strconv.ParseFloat(fields[0], 64)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return int64(value)
|
||||
}
|
||||
|
||||
// ReadLinuxCPUStat returns aggregate CPU jiffies and idle jiffies from /proc/stat.
|
||||
func ReadLinuxCPUStat() (uint64, uint64) {
|
||||
content, err := os.ReadFile("/proc/stat")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
lines := strings.SplitSeq(string(content), "\n")
|
||||
for line := range lines {
|
||||
if !strings.HasPrefix(line, "cpu ") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < cpuStatMinFieldCount {
|
||||
return 0, 0
|
||||
}
|
||||
var total uint64
|
||||
for i := 1; i < len(fields); i++ {
|
||||
value, err := strconv.ParseUint(fields[i], 10, 64)
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
total += value
|
||||
}
|
||||
idle, err := strconv.ParseUint(fields[4], 10, 64)
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
return total, idle
|
||||
}
|
||||
return 0, 0
|
||||
}
|
||||
|
||||
// ReadLinuxNetworkTotals returns aggregate RX and TX byte totals from /proc/net/dev.
|
||||
func ReadLinuxNetworkTotals() (int64, int64) {
|
||||
file, err := os.Open("/proc/net/dev")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var rx int64
|
||||
var tx int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if !strings.Contains(line, ":") {
|
||||
continue
|
||||
}
|
||||
name, data, ok := strings.Cut(line, ":")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(name) == "lo" {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(data)
|
||||
if len(fields) < netDevMinFieldCount {
|
||||
continue
|
||||
}
|
||||
rxValue, err := strconv.ParseInt(fields[0], 10, 64)
|
||||
if err == nil {
|
||||
rx += rxValue
|
||||
}
|
||||
txValue, err := strconv.ParseInt(fields[8], 10, 64)
|
||||
if err == nil {
|
||||
tx += txValue
|
||||
}
|
||||
}
|
||||
return rx, tx
|
||||
}
|
||||
|
||||
// ReadLinuxDiskTotals returns aggregate disk read and write byte totals from /proc/diskstats.
|
||||
func ReadLinuxDiskTotals() (int64, int64) {
|
||||
file, err := os.Open("/proc/diskstats")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var readBytes int64
|
||||
var writeBytes int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
fields := strings.Fields(scanner.Text())
|
||||
if len(fields) < diskStatsMinFieldCount {
|
||||
continue
|
||||
}
|
||||
device := fields[2]
|
||||
if shouldSkipDiskDevice(device) {
|
||||
continue
|
||||
}
|
||||
readSectors, err := strconv.ParseInt(fields[5], 10, 64)
|
||||
if err == nil {
|
||||
readBytes += readSectors * 512
|
||||
}
|
||||
writeSectors, err := strconv.ParseInt(fields[9], 10, 64)
|
||||
if err == nil {
|
||||
writeBytes += writeSectors * 512
|
||||
}
|
||||
}
|
||||
return readBytes, writeBytes
|
||||
}
|
||||
|
||||
func shouldSkipDiskDevice(device string) bool {
|
||||
switch {
|
||||
case device == "":
|
||||
return true
|
||||
case strings.HasPrefix(device, "loop"),
|
||||
strings.HasPrefix(device, "ram"),
|
||||
strings.HasPrefix(device, "dm-"):
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// StatFilesystem returns total and used bytes for the filesystem containing path.
|
||||
func StatFilesystem(path string) (int64, int64) {
|
||||
if strings.TrimSpace(path) == "" {
|
||||
path = string(os.PathSeparator)
|
||||
}
|
||||
absPath := filepath.Clean(path)
|
||||
var stat syscall.Statfs_t
|
||||
if err := syscall.Statfs(absPath, &stat); err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
bsize := int64(stat.Bsize) //nolint:unconvert // Statfs_t.Bsize is int64 on Linux and uint32 on Darwin
|
||||
total := multiplyUint64Int64(stat.Blocks, bsize)
|
||||
free := multiplyUint64Int64(stat.Bavail, bsize)
|
||||
used := max(total-free, 0)
|
||||
return total, used
|
||||
}
|
||||
|
||||
// multiplyUint64Int64 multiplies a uint64 by a positive int64, saturating at
|
||||
// math.MaxInt64 to avoid int64 overflow.
|
||||
func multiplyUint64Int64(a uint64, b int64) int64 {
|
||||
if a == 0 || b <= 0 {
|
||||
return 0
|
||||
}
|
||||
v := a * uint64(b)
|
||||
if v > math.MaxInt64 {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return int64(v)
|
||||
}
|
||||
|
||||
// ReadFirstLine reads and returns the trimmed first line of a file.
|
||||
func ReadFirstLine(path string) string {
|
||||
content, err := os.ReadFile(path) //nolint:gosec // path is a fixed /proc or /sys path from internal callers, not user input
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(string(content))
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseMemInfoValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
line string
|
||||
want int64
|
||||
}{
|
||||
{line: "MemTotal: 16384000 kB", want: 16384000},
|
||||
{line: "MemAvailable: 8192000 kB", want: 8192000},
|
||||
{line: "invalid", want: 0},
|
||||
{line: "MemTotal: not-a-number kB", want: 0},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := parseMemInfoValue(tt.line); got != tt.want {
|
||||
t.Fatalf("parseMemInfoValue(%q) = %d, want %d", tt.line, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldSkipDiskDevice(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
device string
|
||||
want bool
|
||||
}{
|
||||
{device: "", want: true},
|
||||
{device: "loop0", want: true},
|
||||
{device: "ram0", want: true},
|
||||
{device: "dm-0", want: true},
|
||||
{device: "sda", want: false},
|
||||
{device: "nvme0n1", want: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := shouldSkipDiskDevice(tt.device); got != tt.want {
|
||||
t.Fatalf("shouldSkipDiskDevice(%q) = %v, want %v", tt.device, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFirstLine(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "sample.txt")
|
||||
if err := os.WriteFile(path, []byte(" first line\nsecond line\n"), 0o644); err != nil {
|
||||
t.Fatalf("write file: %v", err)
|
||||
}
|
||||
|
||||
if got := ReadFirstLine(path); got != "first line\nsecond line" {
|
||||
t.Fatalf("ReadFirstLine() = %q, want trimmed first line content", got)
|
||||
}
|
||||
if got := ReadFirstLine(filepath.Join(dir, "missing.txt")); got != "" {
|
||||
t.Fatalf("ReadFirstLine(missing) = %q, want empty string", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStatFilesystem(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
total, used := StatFilesystem(t.TempDir())
|
||||
if total <= 0 {
|
||||
t.Fatalf("StatFilesystem() total = %d, want > 0", total)
|
||||
}
|
||||
if used < 0 || used > total {
|
||||
t.Fatalf("StatFilesystem() used = %d, total = %d", used, total)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package runner provides shared WebSocket reconnect helpers for edge daemons.
|
||||
package runner
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
)
|
||||
|
||||
// WSConnection defines the minimum interface required for a WebSocket connection that can be closed.
|
||||
type WSConnection interface {
|
||||
Close() error
|
||||
}
|
||||
|
||||
// WSReconnectConfig specifies configuration parameters for the WebSocket reconnect loop.
|
||||
type WSReconnectConfig struct {
|
||||
ComponentName string
|
||||
ConnectBackoff time.Duration
|
||||
ReconnectDelay time.Duration
|
||||
OnShutdown func()
|
||||
}
|
||||
|
||||
// RunWSReconnectLoop runs a loop that attempts to keep a WebSocket connection active, automatically reconnecting when closed or failed.
|
||||
func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
|
||||
connect func(context.Context) (WSConnection, error),
|
||||
handle func(context.Context, WSConnection),
|
||||
) error {
|
||||
if cfg.ConnectBackoff <= 0 {
|
||||
cfg.ConnectBackoff = 5 * time.Second
|
||||
}
|
||||
if cfg.ReconnectDelay <= 0 {
|
||||
cfg.ReconnectDelay = 2 * time.Second
|
||||
}
|
||||
label := cfg.ComponentName
|
||||
if label == "" {
|
||||
label = "edge"
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if cfg.OnShutdown != nil {
|
||||
cfg.OnShutdown()
|
||||
}
|
||||
return ctx.Err()
|
||||
default:
|
||||
// Continue reconnect loop
|
||||
}
|
||||
|
||||
conn, err := connect(ctx)
|
||||
if err != nil {
|
||||
slog.Error(label+" ws connect failed, will retry", "error", err)
|
||||
SleepContext(ctx, cfg.ConnectBackoff)
|
||||
continue
|
||||
}
|
||||
|
||||
handle(ctx, conn)
|
||||
_ = conn.Close()
|
||||
slog.Info(label + " ws connection closed, reconnecting...")
|
||||
SleepContext(ctx, cfg.ReconnectDelay)
|
||||
}
|
||||
}
|
||||
|
||||
// SleepContext pauses execution for the given duration or until the context is canceled.
|
||||
func SleepContext(ctx context.Context, d time.Duration) {
|
||||
t := time.NewTimer(d)
|
||||
defer t.Stop()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-t.C:
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
//go:build !windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater provides capabilities to check for, download, and apply updates.
|
||||
package updater
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(execPath, backupPath); err != nil {
|
||||
renameErr := err
|
||||
if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) {
|
||||
slog.Error("remove tmp binary failed", "path", tmpPath, "error", err)
|
||||
return fmt.Errorf("backup current binary: %w; remove tmp binary: %w", renameErr, err)
|
||||
}
|
||||
return fmt.Errorf("backup current binary: %w", renameErr)
|
||||
}
|
||||
if err := os.Rename(tmpPath, execPath); err != nil {
|
||||
replaceErr := err
|
||||
if err := os.Rename(backupPath, execPath); err != nil {
|
||||
slog.Error("restore backup binary failed", "path", backupPath, "error", err)
|
||||
return fmt.Errorf("replace binary: %w; restore backup binary: %w", replaceErr, err)
|
||||
}
|
||||
return fmt.Errorf("replace binary: %w", replaceErr)
|
||||
}
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { //nolint:gosec // execPath is the validated edge updater binary path
|
||||
return fmt.Errorf("exec restart: %w", err)
|
||||
}
|
||||
return errors.New("unreachable after exec")
|
||||
}
|
||||
|
||||
func removeBackupBinary(path string) error {
|
||||
if err := os.Remove(path); err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
slog.Error("remove backup binary failed", "path", path, "error", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
//go:build !windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRemoveBackupBinaryIgnoresMissingFile(t *testing.T) {
|
||||
backupPath := filepath.Join(t.TempDir(), "openflare-agent.bak")
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
t.Fatalf("expected missing backup cleanup to be ignored: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
//go:build windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
scriptPath := execPath + ".update.cmd"
|
||||
script := fmt.Sprintf(`@echo off
|
||||
setlocal
|
||||
:waitloop
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 (
|
||||
ping 127.0.0.1 -n 2 >nul
|
||||
goto waitloop
|
||||
)
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 exit /b 1
|
||||
start "" %s
|
||||
del /Q "%s" >nul 2>nul
|
||||
del /Q "%%~f0" >nul 2>nul
|
||||
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
|
||||
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("write restart script: %w", err)
|
||||
}
|
||||
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
|
||||
if err := cmd.Start(); err != nil {
|
||||
os.Remove(scriptPath)
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("schedule restart: %w", err)
|
||||
}
|
||||
os.Exit(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildWindowsCommandLine(execPath string, args []string) string {
|
||||
parts := []string{quoteWindowsArg(execPath)}
|
||||
for _, arg := range args {
|
||||
parts = append(parts, quoteWindowsArg(arg))
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func quoteWindowsArg(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -0,0 +1,426 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater provides capabilities to check for, download, and apply updates.
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/share/ofutil"
|
||||
)
|
||||
|
||||
const (
|
||||
maxChecksumAssetSize = 64 * 1024
|
||||
goosWindows = "windows"
|
||||
updateTmpFilePerm = 0o600
|
||||
updateBinaryFilePerm = 0o755
|
||||
)
|
||||
|
||||
var replaceAndRestartFunc = replaceAndRestart
|
||||
|
||||
// Config defines the configuration for the update service.
|
||||
type Config struct {
|
||||
LocalVersion string
|
||||
AssetPrefix string
|
||||
LogLabel string
|
||||
}
|
||||
|
||||
// Service handles checking and applying application binary updates.
|
||||
type Service struct {
|
||||
httpClient *http.Client
|
||||
lastCheckKey string
|
||||
localVersion string
|
||||
assetPrefix string
|
||||
logLabel string
|
||||
}
|
||||
|
||||
// New creates a new updater Service with the provided configuration.
|
||||
func New(cfg Config) *Service {
|
||||
return &Service{
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
localVersion: cfg.LocalVersion,
|
||||
assetPrefix: cfg.AssetPrefix,
|
||||
logLabel: cfg.LogLabel,
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateOptions specifies parameters for checking and applying updates.
|
||||
type UpdateOptions struct {
|
||||
Channel string
|
||||
TagName string
|
||||
Force bool
|
||||
}
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
Assets []githubAsset `json:"assets"`
|
||||
}
|
||||
|
||||
type githubAsset struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
Digest string `json:"digest"`
|
||||
}
|
||||
|
||||
// CheckAndUpdate checks for a newer release on GitHub and performs an update if available.
|
||||
func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error {
|
||||
release, err := s.getRelease(ctx, repo, options)
|
||||
if err != nil {
|
||||
return fmt.Errorf("check latest release: %w", err)
|
||||
}
|
||||
if release == nil || release.TagName == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
remoteVersion := normalizeVersion(release.TagName)
|
||||
localVersion := normalizeVersion(s.localVersion)
|
||||
checkKey := buildReleaseCheckKey(options, remoteVersion)
|
||||
|
||||
if remoteVersion == localVersion {
|
||||
return nil
|
||||
}
|
||||
if !options.Force && checkKey != "" && checkKey == s.lastCheckKey {
|
||||
return nil
|
||||
}
|
||||
if !isNewer(localVersion, remoteVersion) {
|
||||
s.lastCheckKey = checkKey
|
||||
return nil
|
||||
}
|
||||
|
||||
slog.Info(s.logLabel+" update available", "from", localVersion, "to", remoteVersion)
|
||||
assetName := s.assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
|
||||
downloadURL, expectedChecksum, err := s.resolveReleaseAsset(ctx, release, assetName)
|
||||
if err != nil {
|
||||
if downloadURL == "" {
|
||||
s.lastCheckKey = checkKey
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get executable path: %w", err)
|
||||
}
|
||||
if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil {
|
||||
return fmt.Errorf("download and restart: %w", err)
|
||||
}
|
||||
s.lastCheckKey = checkKey
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) getRelease(ctx context.Context, repo string, options UpdateOptions) (*githubRelease, error) {
|
||||
tagName := strings.TrimSpace(options.TagName)
|
||||
if tagName != "" {
|
||||
return s.getReleaseByTag(ctx, repo, tagName)
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") {
|
||||
return s.getLatestPreviewRelease(ctx, repo)
|
||||
}
|
||||
return s.getLatestStableRelease(ctx, repo)
|
||||
}
|
||||
|
||||
func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) {
|
||||
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo)
|
||||
return s.fetchReleaseFromURL(ctx, url)
|
||||
}
|
||||
|
||||
func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) {
|
||||
url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("github api returned %s", resp.Status)
|
||||
}
|
||||
|
||||
var releases []githubRelease
|
||||
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, release := range releases {
|
||||
if release.Draft || !release.Prerelease {
|
||||
continue
|
||||
}
|
||||
releaseCopy := release
|
||||
return &releaseCopy, nil
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) {
|
||||
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag))
|
||||
return s.fetchReleaseFromURL(ctx, url)
|
||||
}
|
||||
|
||||
func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubRelease, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() {
|
||||
if err := resp.Body.Close(); err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, nil
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("github api returned %s", resp.Status)
|
||||
}
|
||||
|
||||
return decodeRelease(resp.Body)
|
||||
}
|
||||
|
||||
func decodeRelease(reader io.Reader) (*githubRelease, error) {
|
||||
var release githubRelease
|
||||
if err := json.NewDecoder(reader).Decode(&release); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func (s *Service) resolveReleaseAsset(ctx context.Context, release *githubRelease, assetName string) (downloadURL string, expectedChecksum string, err error) {
|
||||
checksumAssetName := assetName + ".sha256"
|
||||
var checksumURL string
|
||||
|
||||
for _, asset := range release.Assets {
|
||||
switch asset.Name {
|
||||
case assetName:
|
||||
downloadURL = asset.BrowserDownloadURL
|
||||
expectedChecksum = normalizeGitHubDigest(asset.Digest)
|
||||
case checksumAssetName:
|
||||
checksumURL = asset.BrowserDownloadURL
|
||||
}
|
||||
}
|
||||
if downloadURL == "" {
|
||||
return "", "", fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
|
||||
}
|
||||
if expectedChecksum != "" {
|
||||
return downloadURL, expectedChecksum, nil
|
||||
}
|
||||
if checksumURL == "" {
|
||||
return downloadURL, "", fmt.Errorf("no sha256 digest or checksum asset %q in release %s", checksumAssetName, release.TagName)
|
||||
}
|
||||
|
||||
expectedChecksum, err = s.downloadChecksum(ctx, checksumURL, assetName)
|
||||
if err != nil {
|
||||
return downloadURL, "", fmt.Errorf("download checksum: %w", err)
|
||||
}
|
||||
return downloadURL, expectedChecksum, nil
|
||||
}
|
||||
|
||||
func normalizeGitHubDigest(digest string) string {
|
||||
digest = strings.TrimSpace(digest)
|
||||
if digest == "" {
|
||||
return ""
|
||||
}
|
||||
const prefix = "sha256:"
|
||||
if strings.HasPrefix(strings.ToLower(digest), prefix) {
|
||||
digest = digest[len(prefix):]
|
||||
}
|
||||
digest = strings.ToLower(digest)
|
||||
if isSHA256Hex(digest) {
|
||||
return digest
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (s *Service) downloadChecksum(ctx context.Context, url string, assetName string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("checksum download returned %s", resp.Status)
|
||||
}
|
||||
|
||||
content, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumAssetSize+1))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if len(content) > maxChecksumAssetSize {
|
||||
return "", fmt.Errorf("checksum asset exceeds %d bytes", maxChecksumAssetSize)
|
||||
}
|
||||
checksum, err := parseSHA256Checksum(string(content), assetName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return checksum, nil
|
||||
}
|
||||
|
||||
func parseSHA256Checksum(content string, assetName string) (string, error) {
|
||||
assetName = strings.TrimSpace(assetName)
|
||||
for line := range strings.SplitSeq(content, "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
}
|
||||
if checksum, ok := parseSHA256Line(line, assetName); ok {
|
||||
return checksum, nil
|
||||
}
|
||||
}
|
||||
if assetName == "" {
|
||||
return "", errors.New("checksum asset does not contain a valid sha256 digest")
|
||||
}
|
||||
return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName)
|
||||
}
|
||||
|
||||
func parseSHA256Line(line string, assetName string) (string, bool) {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) == 1 && isSHA256Hex(fields[0]) {
|
||||
return strings.ToLower(fields[0]), true
|
||||
}
|
||||
if len(fields) >= 2 && isSHA256Hex(fields[0]) {
|
||||
fileName := strings.TrimPrefix(strings.TrimSpace(fields[1]), "*")
|
||||
if assetName == "" || fileName == assetName {
|
||||
return strings.ToLower(fields[0]), true
|
||||
}
|
||||
}
|
||||
|
||||
prefix := "SHA256("
|
||||
if strings.HasPrefix(line, prefix) {
|
||||
closing := strings.Index(line, ")")
|
||||
if closing > len(prefix) && closing+1 < len(line) {
|
||||
fileName := strings.TrimSpace(line[len(prefix):closing])
|
||||
rest := strings.TrimSpace(line[closing+1:])
|
||||
rest = strings.TrimPrefix(rest, "=")
|
||||
rest = strings.TrimSpace(rest)
|
||||
if isSHA256Hex(rest) && (assetName == "" || fileName == assetName) {
|
||||
return strings.ToLower(rest), true
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func isSHA256Hex(value string) bool {
|
||||
value = strings.TrimSpace(value)
|
||||
if len(value) != sha256.Size*2 {
|
||||
return false
|
||||
}
|
||||
_, err := hex.DecodeString(value)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error {
|
||||
expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum))
|
||||
if !isSHA256Hex(expectedChecksum) {
|
||||
return errors.New("invalid expected sha256 checksum")
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("download returned %s", resp.Status)
|
||||
}
|
||||
|
||||
tmpPath := targetPath + ".update"
|
||||
if runtime.GOOS == goosWindows && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") {
|
||||
tmpPath += ".exe"
|
||||
}
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, updateTmpFilePerm) //nolint:gosec // tmpPath is derived from the configured updater binary location
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
hasher := sha256.New()
|
||||
if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil {
|
||||
_ = tmpFile.Close()
|
||||
_ = os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
if err = tmpFile.Close(); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
actualChecksum := hex.EncodeToString(hasher.Sum(nil))
|
||||
if actualChecksum != expectedChecksum {
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("sha256 checksum mismatch: expected %s, got %s", expectedChecksum, actualChecksum)
|
||||
}
|
||||
if err = os.Chmod(tmpPath, updateBinaryFilePerm); err != nil && runtime.GOOS != goosWindows { //nolint:gosec // downloaded edge binary must remain executable
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("set executable permission: %w", err)
|
||||
}
|
||||
|
||||
slog.Info(s.logLabel + " binary updated, restarting")
|
||||
return replaceAndRestartFunc(targetPath, tmpPath)
|
||||
}
|
||||
|
||||
func (s *Service) assetNameForGOOSGOARCH(goos string, goarch string) string {
|
||||
name := fmt.Sprintf("%s-%s-%s", s.assetPrefix, goos, goarch)
|
||||
if goos == goosWindows {
|
||||
return name + ".exe"
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func normalizeVersion(v string) string {
|
||||
v = strings.TrimSpace(v)
|
||||
v = strings.TrimPrefix(v, "v")
|
||||
return v
|
||||
}
|
||||
|
||||
func isNewer(local, remote string) bool {
|
||||
return compareVersions(local, remote) < 0
|
||||
}
|
||||
|
||||
func buildReleaseCheckKey(options UpdateOptions, remoteVersion string) string {
|
||||
channel := strings.TrimSpace(options.Channel)
|
||||
if channel == "" {
|
||||
channel = "stable"
|
||||
}
|
||||
if tagName := strings.TrimSpace(options.TagName); tagName != "" {
|
||||
return channel + ":" + tagName
|
||||
}
|
||||
return channel + ":" + remoteVersion
|
||||
}
|
||||
|
||||
func compareVersions(local string, remote string) int {
|
||||
return ofutil.CompareVersions(local, remote)
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type roundTripFunc func(req *http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func testService(httpClient *http.Client) *Service {
|
||||
return &Service{
|
||||
httpClient: httpClient,
|
||||
localVersion: "v1.0.0",
|
||||
assetPrefix: "openflare-agent",
|
||||
logLabel: "agent",
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLatestPreviewRelease(t *testing.T) {
|
||||
service := testService(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases?per_page=20" {
|
||||
t.Fatalf("unexpected request url: %s", req.URL.String())
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`[
|
||||
{"tag_name":"v1.0.0","prerelease":false},
|
||||
{"tag_name":"v1.1.0-rc.1","prerelease":true}
|
||||
]`)),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
|
||||
release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{Channel: "preview"})
|
||||
if err != nil {
|
||||
t.Fatalf("expected preview release query to succeed: %v", err)
|
||||
}
|
||||
if release == nil || release.TagName != "v1.1.0-rc.1" {
|
||||
t.Fatalf("unexpected preview release: %#v", release)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetReleaseByTag(t *testing.T) {
|
||||
service := testService(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v1.1.0-rc.1" {
|
||||
t.Fatalf("unexpected request url: %s", req.URL.String())
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
|
||||
release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{Channel: "preview", TagName: "v1.1.0-rc.1", Force: true})
|
||||
if err != nil {
|
||||
t.Fatalf("expected tag release query to succeed: %v", err)
|
||||
}
|
||||
if release == nil || release.TagName != "v1.1.0-rc.1" {
|
||||
t.Fatalf("unexpected tag release: %#v", release)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckAndUpdateRequiresChecksumSource(t *testing.T) {
|
||||
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
service := testService(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" {
|
||||
t.Fatalf("unexpected request url: %s", req.URL.String())
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{
|
||||
"tag_name":"v1.0.1",
|
||||
"assets":[
|
||||
{"name":"` + assetName + `","browser_download_url":"https://example.test/agent"}
|
||||
]
|
||||
}`)),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
|
||||
err := service.CheckAndUpdate(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{})
|
||||
if err == nil || !strings.Contains(err.Error(), "no sha256 digest or checksum asset") {
|
||||
t.Fatalf("expected missing checksum source error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeGitHubDigest(t *testing.T) {
|
||||
checksum := strings.Repeat("a", sha256.Size*2)
|
||||
testCases := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{name: "prefixed digest", input: "sha256:" + checksum, want: checksum},
|
||||
{name: "bare hex", input: checksum, want: checksum},
|
||||
{name: "empty", input: "", want: ""},
|
||||
{name: "invalid", input: "sha256:not-a-digest", want: ""},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
if got := normalizeGitHubDigest(testCase.input); got != testCase.want {
|
||||
t.Fatalf("unexpected digest: got %q want %q", got, testCase.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveReleaseAssetPrefersDigest(t *testing.T) {
|
||||
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
checksum := strings.Repeat("b", sha256.Size*2)
|
||||
service := testService(nil)
|
||||
|
||||
downloadURL, expectedChecksum, err := service.resolveReleaseAsset(context.Background(), &githubRelease{
|
||||
TagName: "v1.0.1",
|
||||
Assets: []githubAsset{
|
||||
{
|
||||
Name: assetName,
|
||||
BrowserDownloadURL: "https://example.test/agent",
|
||||
Digest: "sha256:" + checksum,
|
||||
},
|
||||
{
|
||||
Name: assetName + ".sha256",
|
||||
BrowserDownloadURL: "https://example.test/agent.sha256",
|
||||
},
|
||||
},
|
||||
}, assetName)
|
||||
if err != nil {
|
||||
t.Fatalf("expected digest resolution to succeed: %v", err)
|
||||
}
|
||||
if downloadURL != "https://example.test/agent" {
|
||||
t.Fatalf("unexpected download url: %s", downloadURL)
|
||||
}
|
||||
if expectedChecksum != checksum {
|
||||
t.Fatalf("unexpected checksum: got %s want %s", expectedChecksum, checksum)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveReleaseAssetFallsBackToChecksumAsset(t *testing.T) {
|
||||
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
checksum := strings.Repeat("c", sha256.Size*2)
|
||||
service := testService(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.String() != "https://example.test/agent.sha256" {
|
||||
t.Fatalf("unexpected request url: %s", req.URL.String())
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(checksum + "\n")),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
|
||||
downloadURL, expectedChecksum, err := service.resolveReleaseAsset(context.Background(), &githubRelease{
|
||||
TagName: "v1.0.1",
|
||||
Assets: []githubAsset{
|
||||
{
|
||||
Name: assetName,
|
||||
BrowserDownloadURL: "https://example.test/agent",
|
||||
},
|
||||
{
|
||||
Name: assetName + ".sha256",
|
||||
BrowserDownloadURL: "https://example.test/agent.sha256",
|
||||
},
|
||||
},
|
||||
}, assetName)
|
||||
if err != nil {
|
||||
t.Fatalf("expected checksum fallback to succeed: %v", err)
|
||||
}
|
||||
if downloadURL != "https://example.test/agent" {
|
||||
t.Fatalf("unexpected download url: %s", downloadURL)
|
||||
}
|
||||
if expectedChecksum != checksum {
|
||||
t.Fatalf("unexpected checksum: got %s want %s", expectedChecksum, checksum)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseSHA256Checksum(t *testing.T) {
|
||||
checksum := strings.Repeat("a", sha256.Size*2)
|
||||
testCases := []struct {
|
||||
name string
|
||||
content string
|
||||
asset string
|
||||
want string
|
||||
}{
|
||||
{name: "single digest", content: checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
{name: "sha256sum format", content: checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
{name: "bsd format", content: "SHA256(openflare-agent-linux-amd64)= " + checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
{name: "selects matching file", content: strings.Repeat("b", sha256.Size*2) + " other\n" + checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
got, err := parseSHA256Checksum(testCase.content, testCase.asset)
|
||||
if err != nil {
|
||||
t.Fatalf("expected checksum parse to succeed: %v", err)
|
||||
}
|
||||
if got != testCase.want {
|
||||
t.Fatalf("unexpected checksum: got %s want %s", got, testCase.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadAndRestartVerifiesChecksum(t *testing.T) {
|
||||
payload := []byte("new-agent-binary")
|
||||
sum := sha256.Sum256(payload)
|
||||
expectedChecksum := hex.EncodeToString(sum[:])
|
||||
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
|
||||
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
|
||||
t.Fatalf("write target: %v", err)
|
||||
}
|
||||
|
||||
var replacedTarget string
|
||||
var replacedTemp string
|
||||
originalReplace := replaceAndRestartFunc
|
||||
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
|
||||
replacedTarget = execPath
|
||||
replacedTemp = tmpPath
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
replaceAndRestartFunc = originalReplace
|
||||
})
|
||||
|
||||
service := testService(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(string(payload))),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
|
||||
if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil {
|
||||
t.Fatalf("expected verified download to succeed: %v", err)
|
||||
}
|
||||
if replacedTarget != targetPath {
|
||||
t.Fatalf("unexpected replace target: %s", replacedTarget)
|
||||
}
|
||||
if replacedTemp == "" {
|
||||
t.Fatal("expected replacement temp path to be recorded")
|
||||
}
|
||||
if _, err := os.Stat(replacedTemp); err != nil {
|
||||
t.Fatalf("expected verified temp binary to remain for replacement: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadAndRestartRejectsChecksumMismatch(t *testing.T) {
|
||||
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
|
||||
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
|
||||
t.Fatalf("write target: %v", err)
|
||||
}
|
||||
|
||||
originalReplace := replaceAndRestartFunc
|
||||
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
|
||||
t.Fatal("replace should not run on checksum mismatch")
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
replaceAndRestartFunc = originalReplace
|
||||
})
|
||||
|
||||
service := testService(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader("tampered")),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
|
||||
err := service.downloadAndRestart(context.Background(), "https://example.test/agent", strings.Repeat("0", sha256.Size*2), targetPath)
|
||||
if err == nil || !strings.Contains(err.Error(), "sha256 checksum mismatch") {
|
||||
t.Fatalf("expected checksum mismatch error, got %v", err)
|
||||
}
|
||||
if _, err = os.Stat(targetPath + ".update"); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected temp update file to be removed, stat err=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNewerSupportsPrerelease(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
local string
|
||||
remote string
|
||||
expected bool
|
||||
}{
|
||||
{name: "stable newer than prerelease", local: "1.2.3-rc.1", remote: "1.2.3", expected: true},
|
||||
{name: "same stable not newer", local: "1.2.3", remote: "1.2.3-rc.1", expected: false},
|
||||
{name: "higher prerelease sequence", local: "1.2.3-rc.1", remote: "1.2.3-rc.2", expected: true},
|
||||
{name: "higher minor", local: "1.2.3", remote: "1.3.0-rc.1", expected: true},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
if actual := isNewer(testCase.local, testCase.remote); actual != testCase.expected {
|
||||
t.Fatalf("unexpected compare result: local=%s remote=%s actual=%v expected=%v", testCase.local, testCase.remote, actual, testCase.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package wsclient provides WebSocket client abstractions for edge node communication.
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
pkgprotocol "Wavelet/openflare/share/protocol"
|
||||
shared "Wavelet/openflare/share/wsclient"
|
||||
)
|
||||
|
||||
// WSMessage is an alias for the shared WebSocket message type.
|
||||
type WSMessage = shared.WSMessage
|
||||
|
||||
// MessageHandler is an alias for the shared WebSocket message handler type.
|
||||
type MessageHandler = shared.MessageHandler
|
||||
|
||||
// Preset represents a predefined connection configuration for a specific edge role.
|
||||
type Preset int
|
||||
|
||||
const (
|
||||
// PresetAgent is the configuration preset for agent connections.
|
||||
PresetAgent Preset = iota
|
||||
// PresetRelay is the configuration preset for relay connections.
|
||||
PresetRelay
|
||||
// PresetFlared is the configuration preset for flared (tunnel) connections.
|
||||
PresetFlared
|
||||
)
|
||||
|
||||
type presetConfig struct {
|
||||
HeaderKey string
|
||||
WSPath string
|
||||
}
|
||||
|
||||
var presets = map[Preset]presetConfig{
|
||||
PresetAgent: {HeaderKey: "X-Agent-Token", WSPath: "/api/v1/agent/ws"},
|
||||
PresetRelay: {HeaderKey: "X-Agent-Token", WSPath: "/api/v1/relay/ws"},
|
||||
PresetFlared: {HeaderKey: "X-Tunnel-Token", WSPath: "/api/v1/tunnel/ws"},
|
||||
}
|
||||
|
||||
// PresetHeaderKey returns the HTTP header key used for authentication with the given preset.
|
||||
func PresetHeaderKey(preset Preset) string {
|
||||
return presets[preset].HeaderKey
|
||||
}
|
||||
|
||||
// PresetWSPath returns the WebSocket path used for the given preset.
|
||||
func PresetWSPath(preset Preset) string {
|
||||
return presets[preset].WSPath
|
||||
}
|
||||
|
||||
// Client is a WebSocket client configured for a specific edge preset.
|
||||
type Client struct {
|
||||
sharedClient *shared.Client
|
||||
}
|
||||
|
||||
// New creates a new Client for the given preset, base URL, token, and timeout.
|
||||
func New(preset Preset, baseURL, token string, timeout time.Duration) *Client {
|
||||
cfg := presets[preset]
|
||||
return &Client{
|
||||
sharedClient: shared.New(shared.Config{
|
||||
BaseURL: baseURL,
|
||||
Token: token,
|
||||
Timeout: timeout,
|
||||
HeaderKey: cfg.HeaderKey,
|
||||
WSPath: cfg.WSPath,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the authentication token used by the client.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.sharedClient.SetToken(token)
|
||||
}
|
||||
|
||||
// URL returns the fully resolved WebSocket URL for this client.
|
||||
func (c *Client) URL() string {
|
||||
return c.sharedClient.URL()
|
||||
}
|
||||
|
||||
// Connection represents an established WebSocket connection to an edge node.
|
||||
type Connection struct {
|
||||
sharedConn *shared.Connection
|
||||
}
|
||||
|
||||
// Connect establishes a WebSocket connection using the client configuration.
|
||||
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||
conn, err := c.sharedClient.Connect(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Connection{sharedConn: conn}, nil
|
||||
}
|
||||
|
||||
// AgentConnection is a Connection specialized for agent node communication.
|
||||
type AgentConnection struct {
|
||||
Connection
|
||||
}
|
||||
|
||||
// ConnectAgent establishes a WebSocket connection and returns it as an AgentConnection.
|
||||
func (c *Client) ConnectAgent(ctx context.Context) (*AgentConnection, error) {
|
||||
conn, err := c.Connect(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &AgentConnection{Connection: *conn}, nil
|
||||
}
|
||||
|
||||
// URL returns the resolved WebSocket URL of this connection.
|
||||
func (conn *Connection) URL() string {
|
||||
if conn == nil || conn.sharedConn == nil {
|
||||
return ""
|
||||
}
|
||||
return conn.sharedConn.URL
|
||||
}
|
||||
|
||||
// SendPing sends a ping message over the connection.
|
||||
func (conn *Connection) SendPing() error {
|
||||
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypePing, nil)
|
||||
}
|
||||
|
||||
// SendPong sends a pong message over the connection.
|
||||
func (conn *Connection) SendPong() error {
|
||||
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypePong, nil)
|
||||
}
|
||||
|
||||
// SendMessage sends a typed message with an optional payload over the connection.
|
||||
func (conn *Connection) SendMessage(msgType string, payload any) error {
|
||||
return conn.sharedConn.SendMessage(msgType, payload)
|
||||
}
|
||||
|
||||
// Receive reads the next message from the connection.
|
||||
func (conn *Connection) Receive() (pkgprotocol.WSMessage, error) {
|
||||
var message pkgprotocol.WSMessage
|
||||
if err := conn.sharedConn.Receive(&message); err != nil {
|
||||
return message, err
|
||||
}
|
||||
return message, nil
|
||||
}
|
||||
|
||||
// RunReceiveLoop blocks and dispatches incoming messages to the handler until the context is canceled.
|
||||
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
|
||||
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||
}
|
||||
|
||||
// Close gracefully closes the WebSocket connection.
|
||||
func (conn *Connection) Close() error {
|
||||
if conn == nil || conn.sharedConn == nil {
|
||||
return nil
|
||||
}
|
||||
return conn.sharedConn.Close()
|
||||
}
|
||||
|
||||
// SendStatus sends a node status payload over the agent connection.
|
||||
func (conn *AgentConnection) SendStatus(payload pkgprotocol.NodePayload) error {
|
||||
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypeStatus, payload)
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestPresetConfig(t *testing.T) {
|
||||
tests := []struct {
|
||||
preset Preset
|
||||
headerKey string
|
||||
wsPath string
|
||||
}{
|
||||
{preset: PresetAgent, headerKey: "X-Agent-Token", wsPath: "/api/v1/agent/ws"},
|
||||
{preset: PresetRelay, headerKey: "X-Agent-Token", wsPath: "/api/v1/relay/ws"},
|
||||
{preset: PresetFlared, headerKey: "X-Tunnel-Token", wsPath: "/api/v1/tunnel/ws"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.wsPath, func(t *testing.T) {
|
||||
if got := PresetHeaderKey(tt.preset); got != tt.headerKey {
|
||||
t.Fatalf("PresetHeaderKey() = %q, want %q", got, tt.headerKey)
|
||||
}
|
||||
if got := PresetWSPath(tt.preset); got != tt.wsPath {
|
||||
t.Fatalf("PresetWSPath() = %q, want %q", got, tt.wsPath)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientURL(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
preset Preset
|
||||
baseURL string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "agent https",
|
||||
preset: PresetAgent,
|
||||
baseURL: "https://example.com",
|
||||
want: "wss://example.com/api/v1/agent/ws",
|
||||
},
|
||||
{
|
||||
name: "relay http with path prefix",
|
||||
preset: PresetRelay,
|
||||
baseURL: "http://example.com/api",
|
||||
want: "ws://example.com/api/api/v1/relay/ws",
|
||||
},
|
||||
{
|
||||
name: "flared wss",
|
||||
preset: PresetFlared,
|
||||
baseURL: "wss://edge.example.com",
|
||||
want: "wss://edge.example.com/api/v1/tunnel/ws",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
client := New(tt.preset, tt.baseURL, "token", time.Second)
|
||||
if got := client.URL(); got != tt.want {
|
||||
t.Fatalf("URL() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user