mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 15:06:36 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -0,0 +1 @@
|
||||
/data/
|
||||
@@ -0,0 +1,33 @@
|
||||
# syntax=docker/dockerfile:1.7
|
||||
ARG VERSION=dev
|
||||
|
||||
FROM golang:1.25-alpine AS builder
|
||||
|
||||
ARG VERSION
|
||||
|
||||
WORKDIR /build
|
||||
COPY go.mod go.sum ./
|
||||
RUN --mount=type=cache,target=/go/pkg/mod \
|
||||
go mod download
|
||||
|
||||
COPY openflare-server ./openflare-server
|
||||
COPY openflare-relay ./openflare-relay
|
||||
RUN --mount=type=cache,target=/go/pkg/mod \
|
||||
--mount=type=cache,target=/root/.cache/go-build \
|
||||
CGO_ENABLED=0 GOOS=linux go build -trimpath -ldflags "-s -w -X 'github.com/rain-kl/openflare/openflare-relay/internal/config.Version=$VERSION'" -o /build/bin/openflare-relay ./openflare-relay/cmd/relay
|
||||
|
||||
# Final runtime image
|
||||
FROM fatedier/frps:v0.69.0
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Copy openflare-relay binary
|
||||
COPY --from=builder /build/bin/openflare-relay ./openflare-relay
|
||||
|
||||
VOLUME ["/app/data"]
|
||||
|
||||
ENV OPENFLARE_FRPS_PATH=/usr/bin/frps
|
||||
ENV OPENFLARE_DATA_DIR=/app/data
|
||||
|
||||
ENTRYPOINT ["/app/openflare-relay"]
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strings"
|
||||
"syscall"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/config"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/frps"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/heartbeat"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/httpclient"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/relay"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/state"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/wsclient"
|
||||
)
|
||||
|
||||
func main() {
|
||||
// Setup simple structured logging
|
||||
slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{
|
||||
Level: parseLevel(os.Getenv("LOG_LEVEL")),
|
||||
})))
|
||||
|
||||
configPath := flag.String("config", "./relay.json", "relay config path")
|
||||
flag.Parse()
|
||||
|
||||
cfg, err := config.Load(*configPath)
|
||||
if err != nil {
|
||||
slog.Error("load relay config failed", "error", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
slog.Info("relay config loaded",
|
||||
"server", cfg.ServerURL,
|
||||
"node", cfg.NodeName,
|
||||
"ip", cfg.NodeIP,
|
||||
"frps_path", cfg.FrpsPath,
|
||||
"data_dir", cfg.DataDir,
|
||||
"heartbeat_interval", cfg.HeartbeatInterval,
|
||||
)
|
||||
|
||||
stateStore := state.NewStore(cfg.StatePath)
|
||||
_ = stateStore // In the future we may use stateStore for auth caching
|
||||
|
||||
frpsManager := frps.NewManager(cfg.FrpsPath, cfg.DataDir, cfg.InitialAuthToken())
|
||||
|
||||
slog.Info("detected frps version", "version", frpsManager.GetVersion())
|
||||
|
||||
httpClient := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
|
||||
wsClient := wsclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
|
||||
|
||||
runner := &relay.Runner{
|
||||
Config: cfg,
|
||||
StateStore: stateStore,
|
||||
FrpsManager: frpsManager,
|
||||
HttpClient: httpClient,
|
||||
WebSocketService: wsClient,
|
||||
HeartbeatService: heartbeat.New(httpClient, frpsManager, cfg, stateStore),
|
||||
}
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
||||
defer stop()
|
||||
|
||||
slog.Info("relay process started")
|
||||
|
||||
if err := runner.Run(ctx); err != nil && err != context.Canceled {
|
||||
slog.Error("relay process exited with error", "error", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
slog.Info("relay process stopped")
|
||||
}
|
||||
|
||||
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,235 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
|
||||
"github.com/rain-kl/openflare/openflare-server/utils/geoip/iputil"
|
||||
)
|
||||
|
||||
type MillisecondDuration time.Duration
|
||||
|
||||
func (d *MillisecondDuration) UnmarshalJSON(b []byte) error {
|
||||
var v interface{}
|
||||
if err := json.Unmarshal(b, &v); err != nil {
|
||||
return err
|
||||
}
|
||||
switch value := v.(type) {
|
||||
case float64:
|
||||
*d = MillisecondDuration(time.Duration(value) * time.Millisecond)
|
||||
return nil
|
||||
case string:
|
||||
duration, err := time.ParseDuration(value)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*d = MillisecondDuration(duration)
|
||||
return nil
|
||||
default:
|
||||
return errors.New("invalid duration format")
|
||||
}
|
||||
}
|
||||
|
||||
func (d MillisecondDuration) Duration() time.Duration {
|
||||
return time.Duration(d)
|
||||
}
|
||||
|
||||
func (d MillisecondDuration) String() string {
|
||||
return time.Duration(d).String()
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
AgentToken string `json:"agent_token"`
|
||||
DiscoveryToken string `json:"discovery_token"`
|
||||
NodeName string `json:"node_name"`
|
||||
NodeIP string `json:"node_ip"`
|
||||
FrpsPath string `json:"frps_path"`
|
||||
DataDir string `json:"data_dir"`
|
||||
StatePath string `json:"state_path"`
|
||||
HeartbeatInterval MillisecondDuration `json:"heartbeat_interval"`
|
||||
RequestTimeout MillisecondDuration `json:"request_timeout"`
|
||||
configPath string
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
cfg := &Config{}
|
||||
if err == nil {
|
||||
if err = json.Unmarshal(data, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if err != nil && !hasEnvConfig() {
|
||||
return nil, err
|
||||
}
|
||||
cfg.configPath = path
|
||||
applyEnvOverrides(cfg)
|
||||
applyDefaults(cfg, filepath.Dir(path))
|
||||
if err = validate(cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func hasEnvConfig() bool {
|
||||
for _, key := range []string{
|
||||
"OPENFLARE_SERVER_URL",
|
||||
"OPENFLARE_AGENT_TOKEN",
|
||||
"OPENFLARE_DISCOVERY_TOKEN",
|
||||
"OPENFLARE_NODE_NAME",
|
||||
"OPENFLARE_NODE_IP",
|
||||
"OPENFLARE_DATA_DIR",
|
||||
"OPENFLARE_FRPS_PATH",
|
||||
} {
|
||||
if strings.TrimSpace(os.Getenv(key)) != "" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func applyEnvOverrides(cfg *Config) {
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
overrideString := func(key string, target *string) {
|
||||
if value := strings.TrimSpace(os.Getenv(key)); value != "" {
|
||||
*target = value
|
||||
}
|
||||
}
|
||||
overrideString("OPENFLARE_SERVER_URL", &cfg.ServerURL)
|
||||
overrideString("OPENFLARE_AGENT_TOKEN", &cfg.AgentToken)
|
||||
overrideString("OPENFLARE_DISCOVERY_TOKEN", &cfg.DiscoveryToken)
|
||||
overrideString("OPENFLARE_NODE_NAME", &cfg.NodeName)
|
||||
overrideString("OPENFLARE_NODE_IP", &cfg.NodeIP)
|
||||
overrideString("OPENFLARE_DATA_DIR", &cfg.DataDir)
|
||||
overrideString("OPENFLARE_FRPS_PATH", &cfg.FrpsPath)
|
||||
}
|
||||
|
||||
func applyDefaults(cfg *Config, baseDir string) {
|
||||
baseDir = filepath.Clean(baseDir)
|
||||
if cfg.FrpsPath == "" {
|
||||
cfg.FrpsPath = "frps" // rely on PATH
|
||||
}
|
||||
if cfg.DataDir == "" {
|
||||
cfg.DataDir = filepath.Join(baseDir, "data")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
host, _ := os.Hostname()
|
||||
cfg.NodeName = strings.TrimSpace(host)
|
||||
}
|
||||
if cfg.NodeIP == "" {
|
||||
cfg.NodeIP = detectNodeIP()
|
||||
}
|
||||
if cfg.StatePath == "" {
|
||||
cfg.StatePath = filepath.Join(cfg.DataDir, "relay-state.json")
|
||||
}
|
||||
if cfg.HeartbeatInterval <= 0 {
|
||||
cfg.HeartbeatInterval = MillisecondDuration(10 * time.Second)
|
||||
}
|
||||
if cfg.RequestTimeout <= 0 {
|
||||
cfg.RequestTimeout = MillisecondDuration(10 * time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
func validate(cfg *Config) error {
|
||||
if cfg.ServerURL == "" {
|
||||
return errors.New("server_url 不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.AgentToken) == "" && strings.TrimSpace(cfg.DiscoveryToken) == "" {
|
||||
return errors.New("agent_token 和 discovery_token 不能同时为空")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
return errors.New("node_name 不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cfg *Config) InitialAuthToken() string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
}
|
||||
if token := strings.TrimSpace(cfg.AgentToken); token != "" {
|
||||
return token
|
||||
}
|
||||
return strings.TrimSpace(cfg.DiscoveryToken)
|
||||
}
|
||||
|
||||
func (cfg *Config) Save() error {
|
||||
if cfg == nil {
|
||||
return errors.New("config 不能为空")
|
||||
}
|
||||
if cfg.configPath == "" {
|
||||
return errors.New("config path 未初始化")
|
||||
}
|
||||
data, err := json.MarshalIndent(cfg, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(cfg.configPath, data, 0o644)
|
||||
}
|
||||
|
||||
func detectNodeIP() string {
|
||||
if ip := detectOutboundNodeIP(); ip != "" {
|
||||
return ip
|
||||
}
|
||||
return detectLocalNodeIP()
|
||||
}
|
||||
|
||||
func detectOutboundNodeIP() string {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ip, err := geoip.GetOutboundIP(ctx)
|
||||
if err != nil || ip == nil {
|
||||
return ""
|
||||
}
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
func detectLocalNodeIP() 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 == 2 {
|
||||
return bestIP
|
||||
}
|
||||
}
|
||||
}
|
||||
return bestIP
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
package config
|
||||
|
||||
var Version = "dev"
|
||||
@@ -0,0 +1,306 @@
|
||||
package frps
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
)
|
||||
|
||||
type Manager struct {
|
||||
frpsPath string
|
||||
dataDir string
|
||||
configPath string
|
||||
pidPath string
|
||||
agentToken string
|
||||
|
||||
mu sync.RWMutex
|
||||
activeConfig *service.RelayConfig
|
||||
cmd *exec.Cmd
|
||||
status string
|
||||
lastError string
|
||||
generation uint64
|
||||
stopping bool
|
||||
}
|
||||
|
||||
type RuntimeStatus struct {
|
||||
Status string
|
||||
LastError string
|
||||
Connections int
|
||||
ProxyCount int
|
||||
ClientCount int
|
||||
Proxies []service.RelayProxyStat
|
||||
ProcessAlive bool
|
||||
}
|
||||
|
||||
func NewManager(frpsPath string, dataDir string, agentToken string) *Manager {
|
||||
return &Manager{
|
||||
frpsPath: frpsPath,
|
||||
dataDir: dataDir,
|
||||
configPath: filepath.Join(dataDir, "frps.toml"),
|
||||
pidPath: filepath.Join(dataDir, "frps.pid"),
|
||||
status: "unknown", // 启动阶段尚未获取配置,状态未知;避免首次 heartbeat 误报 frps_unhealthy
|
||||
agentToken: agentToken,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) GetVersion() string {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cmd := exec.CommandContext(ctx, m.frpsPath, "-v")
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
slog.Error("failed to get frps version", "error", err)
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(string(out))
|
||||
}
|
||||
|
||||
func (m *Manager) GetStatus() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.status
|
||||
}
|
||||
|
||||
func (m *Manager) GetRuntimeStatus() RuntimeStatus {
|
||||
m.mu.RLock()
|
||||
status := m.status
|
||||
lastError := m.lastError
|
||||
cmd := m.cmd
|
||||
m.mu.RUnlock()
|
||||
|
||||
return RuntimeStatus{
|
||||
Status: status,
|
||||
LastError: lastError,
|
||||
Connections: 0,
|
||||
ProxyCount: 0,
|
||||
ClientCount: 0,
|
||||
Proxies: nil,
|
||||
ProcessAlive: cmd != nil && cmd.Process != nil,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) UpdateConfig(cfg *service.RelayConfig) {
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// Check if config changed
|
||||
if m.activeConfig != nil &&
|
||||
m.activeConfig.BindPort == cfg.BindPort &&
|
||||
m.activeConfig.VhostHTTPPort == cfg.VhostHTTPPort &&
|
||||
m.activeConfig.AuthToken == cfg.AuthToken &&
|
||||
m.activeConfig.WebServerEnabled == cfg.WebServerEnabled {
|
||||
if m.cmd == nil && !m.stopping {
|
||||
slog.Warn("frps config unchanged but process is not running, restarting")
|
||||
m.stopping = false
|
||||
m.generation++
|
||||
generation := m.generation
|
||||
if err := m.renderConfig(cfg); err != nil {
|
||||
slog.Error("failed to render frps config", "error", err)
|
||||
m.status = "unhealthy"
|
||||
m.lastError = err.Error()
|
||||
return
|
||||
}
|
||||
go m.supervise(generation)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
m.activeConfig = cfg
|
||||
m.stopping = false
|
||||
m.generation++
|
||||
generation := m.generation
|
||||
slog.Info("relay config updated, reloading frps")
|
||||
|
||||
if m.cmd != nil && m.cmd.Process != nil {
|
||||
slog.Debug("stopping existing frps process")
|
||||
_ = m.cmd.Process.Kill()
|
||||
m.cmd = nil
|
||||
}
|
||||
|
||||
if err := m.renderConfig(cfg); err != nil {
|
||||
slog.Error("failed to render frps config", "error", err)
|
||||
m.status = "unhealthy"
|
||||
m.lastError = err.Error()
|
||||
return
|
||||
}
|
||||
|
||||
go m.supervise(generation)
|
||||
}
|
||||
|
||||
func (m *Manager) renderConfig(cfg *service.RelayConfig) error {
|
||||
if err := os.MkdirAll(m.dataDir, 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
buf.WriteString(fmt.Sprintf("bindPort = %d\n", cfg.BindPort))
|
||||
if cfg.VhostHTTPPort > 0 {
|
||||
buf.WriteString(fmt.Sprintf("vhostHTTPPort = %d\n", cfg.VhostHTTPPort))
|
||||
}
|
||||
if cfg.AuthToken != "" {
|
||||
buf.WriteString("[auth]\n")
|
||||
buf.WriteString("method = \"token\"\n")
|
||||
buf.WriteString(fmt.Sprintf("token = \"%s\"\n", cfg.AuthToken))
|
||||
}
|
||||
|
||||
// WebServer configuration
|
||||
buf.WriteString("\n[webServer]\n")
|
||||
if cfg.WebServerEnabled {
|
||||
buf.WriteString("addr = \"0.0.0.0\"\n")
|
||||
} else {
|
||||
buf.WriteString("addr = \"127.0.0.1\"\n")
|
||||
}
|
||||
buf.WriteString(fmt.Sprintf("port = %d\n", 17500))
|
||||
buf.WriteString("user = \"admin\"\n")
|
||||
|
||||
password := m.agentToken
|
||||
if password == "" {
|
||||
password = "admin"
|
||||
}
|
||||
buf.WriteString(fmt.Sprintf("password = \"%s\"\n", password))
|
||||
|
||||
return os.WriteFile(m.configPath, buf.Bytes(), 0644)
|
||||
}
|
||||
|
||||
func (m *Manager) supervise(generation uint64) {
|
||||
backoff := 1 * time.Second
|
||||
const maxBackoff = 60 * time.Second
|
||||
|
||||
for {
|
||||
m.mu.Lock()
|
||||
if m.stopping || m.generation != generation {
|
||||
m.mu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
ensureNoOrphanProcess(m.pidPath)
|
||||
|
||||
cmd := exec.Command(m.frpsPath, "-c", m.configPath)
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
err := cmd.Start()
|
||||
if err != nil {
|
||||
m.status = "unhealthy"
|
||||
m.lastError = fmt.Sprintf("failed to start: %v", err)
|
||||
slog.Error("failed to start frps", "error", err, "generation", generation)
|
||||
m.mu.Unlock()
|
||||
|
||||
if !m.sleepOrInterrupt(generation, backoff) {
|
||||
return
|
||||
}
|
||||
backoff = backoff * 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
_ = os.WriteFile(m.pidPath, []byte(fmt.Sprintf("%d", cmd.Process.Pid)), 0644)
|
||||
|
||||
m.cmd = cmd
|
||||
m.status = "healthy"
|
||||
m.lastError = ""
|
||||
m.mu.Unlock()
|
||||
|
||||
startedAt := time.Now()
|
||||
waitErr := cmd.Wait()
|
||||
_ = os.Remove(m.pidPath)
|
||||
|
||||
m.mu.Lock()
|
||||
if m.cmd == cmd {
|
||||
m.cmd = nil
|
||||
m.status = "unhealthy"
|
||||
if waitErr != nil {
|
||||
m.lastError = fmt.Sprintf("exited with error: %v", waitErr)
|
||||
} else {
|
||||
m.lastError = "exited unexpectedly"
|
||||
}
|
||||
slog.Warn("frps process exited unexpectedly", "error", waitErr, "generation", generation)
|
||||
}
|
||||
shouldContinue := !m.stopping && m.generation == generation
|
||||
m.mu.Unlock()
|
||||
|
||||
if !shouldContinue {
|
||||
return
|
||||
}
|
||||
|
||||
if time.Since(startedAt) >= 10*time.Second {
|
||||
backoff = 1 * time.Second
|
||||
}
|
||||
|
||||
if !m.sleepOrInterrupt(generation, backoff) {
|
||||
return
|
||||
}
|
||||
backoff = backoff * 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) sleepOrInterrupt(generation uint64, d time.Duration) bool {
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
deadline := time.Now().Add(d)
|
||||
for time.Now().Before(deadline) {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
m.mu.RLock()
|
||||
interrupted := m.stopping || m.generation != generation
|
||||
m.mu.RUnlock()
|
||||
if interrupted {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (m *Manager) Stop() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.stopping = true
|
||||
m.generation++
|
||||
if m.cmd != nil && m.cmd.Process != nil {
|
||||
_ = m.cmd.Process.Kill()
|
||||
m.cmd = nil
|
||||
}
|
||||
_ = os.Remove(m.pidPath)
|
||||
m.status = "unhealthy"
|
||||
}
|
||||
|
||||
func ensureNoOrphanProcess(pidPath string) {
|
||||
data, err := os.ReadFile(pidPath)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
var pid int
|
||||
if _, err := fmt.Sscanf(string(data), "%d", &pid); err != nil {
|
||||
return
|
||||
}
|
||||
if pid <= 0 {
|
||||
return
|
||||
}
|
||||
process, err := os.FindProcess(pid)
|
||||
if err == nil && process != nil {
|
||||
slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath)
|
||||
_ = process.Kill()
|
||||
// Wait a little bit to ensure the OS has reclaimed ports
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package frps
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
)
|
||||
|
||||
// Helper to write control file for the dummy script
|
||||
func writeControl(t *testing.T, dir string, exitCode int, delaySeconds int) {
|
||||
controlPath := filepath.Join(dir, "control.txt")
|
||||
content := fmt.Sprintf("%d %d\n", exitCode, delaySeconds)
|
||||
err := os.WriteFile(controlPath, []byte(content), 0644)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write control file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Setup a dummy executable script that reads control.txt to decide exit code and sleep duration
|
||||
func setupDummyScript(t *testing.T) (string, string) {
|
||||
dir := t.TempDir()
|
||||
scriptPath := filepath.Join(dir, "dummy_frps")
|
||||
|
||||
// On macOS/Linux, we write a shell script
|
||||
scriptContent := fmt.Sprintf(`#!/bin/sh
|
||||
control_file="%s/control.txt"
|
||||
EXIT_CODE=0
|
||||
DELAY=0
|
||||
if [ -f "$control_file" ]; then
|
||||
read -r EXIT_CODE DELAY < "$control_file"
|
||||
fi
|
||||
if [ -n "$DELAY" ] && [ "$DELAY" -gt 0 ] 2>/dev/null; then
|
||||
sleep "$DELAY"
|
||||
fi
|
||||
exit "${EXIT_CODE:-0}"
|
||||
`, dir)
|
||||
|
||||
err := os.WriteFile(scriptPath, []byte(scriptContent), 0755)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write dummy script: %v", err)
|
||||
}
|
||||
|
||||
return scriptPath, dir
|
||||
}
|
||||
|
||||
// Helper to poll for status to eliminate timing flakiness in tests
|
||||
func assertStatusEventually(t *testing.T, m *Manager, expectedStatus string, timeout time.Duration) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
rt := m.GetRuntimeStatus()
|
||||
if rt.Status == expectedStatus {
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
rt := m.GetRuntimeStatus()
|
||||
t.Fatalf("expected status eventually %s, got %s (err: %s)", expectedStatus, rt.Status, rt.LastError)
|
||||
}
|
||||
|
||||
func assertCommandExitedEventually(t *testing.T, cmd *exec.Cmd, timeout time.Duration) {
|
||||
t.Helper()
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- cmd.Wait()
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-time.After(timeout):
|
||||
t.Fatalf("expected process pid=%d to exit within %s", cmd.Process.Pid, timeout)
|
||||
case <-done:
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartProcessSuccess(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 5) // exit code 0, sleep 5s
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
|
||||
rt := m.GetRuntimeStatus()
|
||||
if !rt.ProcessAlive {
|
||||
t.Error("expected process to be alive")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartProcessFailureAndBackoff(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
invalidScriptPath := filepath.Join(dir, "non_existent_frps")
|
||||
|
||||
m := NewManager(invalidScriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
rt := m.GetRuntimeStatus()
|
||||
if !strings.Contains(rt.LastError, "failed to start") {
|
||||
t.Errorf("expected error message containing 'failed to start', got %s", rt.LastError)
|
||||
}
|
||||
|
||||
// Correct the path to dummy script
|
||||
scriptPath, _ := setupDummyScript(t)
|
||||
writeControl(t, filepath.Dir(scriptPath), 0, 5)
|
||||
|
||||
m.mu.Lock()
|
||||
m.frpsPath = scriptPath
|
||||
m.mu.Unlock()
|
||||
|
||||
// Wait for backoff retry (1s backoff)
|
||||
assertStatusEventually(t, m, "healthy", 3*time.Second)
|
||||
|
||||
rt = m.GetRuntimeStatus()
|
||||
if !rt.ProcessAlive {
|
||||
t.Error("expected process to be alive now")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnexpectedExitAndAutorestart(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Start with immediate exit code 1
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
rt := m.GetRuntimeStatus()
|
||||
if !strings.Contains(rt.LastError, "exited with error") {
|
||||
t.Errorf("expected exit error, got %s", rt.LastError)
|
||||
}
|
||||
|
||||
// Change control to be healthy (runs for 5s, exit 0)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Wait for the retry to fire (backoff was 1s)
|
||||
assertStatusEventually(t, m, "healthy", 3*time.Second)
|
||||
}
|
||||
|
||||
func TestBackoffReset(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Rapid exit to increase backoff
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
// Crashed once, backoff is 2s
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
// Now make it run successfully for 11 seconds (exit code 0, sleep 11s)
|
||||
writeControl(t, dir, 0, 11)
|
||||
|
||||
// Wait for next retry to start running
|
||||
assertStatusEventually(t, m, "healthy", 4*time.Second)
|
||||
|
||||
// Wait for process to run for 10.5 seconds to trigger backoff reset
|
||||
time.Sleep(10500 * time.Millisecond)
|
||||
|
||||
// Now make it crash again (exit code 1, sleep 0s)
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
// Wait for it to finish and crash
|
||||
assertStatusEventually(t, m, "unhealthy", 3*time.Second)
|
||||
|
||||
// It crashed. Since it ran for > 10s, backoff should have been reset to 1s.
|
||||
// We make it healthy again (exit code 0, sleep 5)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Wait 1.5 seconds. If backoff was reset to 1s, it should be healthy now.
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
}
|
||||
|
||||
func TestImmediateRestartOnSameConfigDeadProcess(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Crashes immediately
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
// Let it crash
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
// Make it start successfully
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Send same config block to trigger immediate restart bypass of backoff sleep
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
// Check if it started immediately
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
}
|
||||
|
||||
func TestSupervisorGenerationInterrupt(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 10)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
|
||||
m.mu.Lock()
|
||||
gen1 := m.generation
|
||||
cmd1 := m.cmd
|
||||
m.mu.Unlock()
|
||||
|
||||
if cmd1 == nil {
|
||||
t.Fatal("expected active process")
|
||||
}
|
||||
|
||||
// Update configuration with new bind port to trigger new generation
|
||||
cfg2 := &service.RelayConfig{
|
||||
BindPort: 7001,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
m.UpdateConfig(cfg2)
|
||||
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
|
||||
m.mu.Lock()
|
||||
gen2 := m.generation
|
||||
cmd2 := m.cmd
|
||||
m.mu.Unlock()
|
||||
|
||||
if gen2 <= gen1 {
|
||||
t.Errorf("expected generation incremented, got gen1=%d gen2=%d", gen1, gen2)
|
||||
}
|
||||
if cmd2 == cmd1 {
|
||||
t.Error("expected old process killed and new command started")
|
||||
}
|
||||
|
||||
// Verify old process is actually killed
|
||||
var cmd1Finished int32
|
||||
go func() {
|
||||
_ = cmd1.Wait()
|
||||
atomic.StoreInt32(&cmd1Finished, 1)
|
||||
}()
|
||||
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
if atomic.LoadInt32(&cmd1Finished) != 1 {
|
||||
t.Error("expected first process to be killed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateConfigKillsOrphanProcessBeforeRestart(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
orphan := exec.Command("sh", "-c", "sleep 30")
|
||||
if err := orphan.Start(); err != nil {
|
||||
t.Fatalf("failed to start orphan process: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if orphan.Process != nil {
|
||||
_ = orphan.Process.Kill()
|
||||
}
|
||||
})
|
||||
|
||||
if err := os.WriteFile(m.pidPath, []byte(fmt.Sprintf("%d", orphan.Process.Pid)), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed orphan pid file: %v", err)
|
||||
}
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertCommandExitedEventually(t, orphan, 2*time.Second)
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/config"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/frps"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/httpclient"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/observability"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/state"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/updater"
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
client *httpclient.Client
|
||||
frpsManager *frps.Manager
|
||||
config *config.Config
|
||||
stateStore *state.Store
|
||||
updater *updater.Service
|
||||
}
|
||||
|
||||
func New(client *httpclient.Client, manager *frps.Manager, cfg *config.Config, stateStore *state.Store) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
frpsManager: manager,
|
||||
config: cfg,
|
||||
stateStore: stateStore,
|
||||
updater: updater.New(),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) Run(ctx context.Context) {
|
||||
ticker := time.NewTicker(s.config.HeartbeatInterval.Duration())
|
||||
defer ticker.Stop()
|
||||
|
||||
// initial heartbeat
|
||||
s.doHeartbeat(ctx)
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.doHeartbeat(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) doHeartbeat(ctx context.Context) {
|
||||
slog.Debug("sending heartbeat")
|
||||
|
||||
runtimeStatus := s.frpsManager.GetRuntimeStatus()
|
||||
payload := service.RelayHeartbeatPayload{
|
||||
Version: config.Version,
|
||||
ExtVersion: s.frpsManager.GetVersion(),
|
||||
RelayStatus: runtimeStatus.Status,
|
||||
FrpsConnCount: runtimeStatus.Connections,
|
||||
FrpsProxyCount: runtimeStatus.ProxyCount,
|
||||
FrpsClientCount: runtimeStatus.ClientCount,
|
||||
FrpsProxies: runtimeStatus.Proxies,
|
||||
Name: s.config.NodeName,
|
||||
IP: s.config.NodeIP,
|
||||
Profile: observability.BuildProfile(s.config, s.stateStore),
|
||||
Snapshot: observability.BuildSnapshot(s.config, s.stateStore),
|
||||
HealthEvents: observability.BuildHealthEvents(runtimeStatus),
|
||||
}
|
||||
|
||||
resp, err := s.client.Heartbeat(ctx, payload)
|
||||
if err != nil {
|
||||
slog.Error("heartbeat failed", "error", err)
|
||||
return
|
||||
}
|
||||
slog.Debug("heartbeat succeeded")
|
||||
|
||||
// Update configs if changed
|
||||
s.frpsManager.UpdateConfig(resp.RelayConfig)
|
||||
|
||||
if resp != nil && resp.RelaySettings != nil {
|
||||
s.tryAutoUpdate(ctx, resp.RelaySettings)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) tryAutoUpdate(ctx context.Context, settings *service.RelaySettings) {
|
||||
if settings == nil || s.updater == nil {
|
||||
return
|
||||
}
|
||||
force := settings.UpdateNow
|
||||
shouldCheck := settings.AutoUpdate || force
|
||||
if !shouldCheck || settings.UpdateRepo == "" {
|
||||
return
|
||||
}
|
||||
channel := "stable"
|
||||
if force && settings.UpdateChannel != "" {
|
||||
channel = settings.UpdateChannel
|
||||
}
|
||||
slog.Info("checking for relay updates", "repo", settings.UpdateRepo, "channel", channel, "force", force)
|
||||
err := s.updater.CheckAndUpdate(ctx, settings.UpdateRepo, updater.UpdateOptions{
|
||||
Channel: channel,
|
||||
TagName: settings.UpdateTag,
|
||||
Force: force,
|
||||
})
|
||||
if err != nil {
|
||||
slog.Error("relay update check failed", "error", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
)
|
||||
|
||||
type APIResponse[T any] struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
Data T `json:"data"`
|
||||
}
|
||||
|
||||
type Client struct {
|
||||
baseURL string
|
||||
token string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
token: token,
|
||||
httpClient: &http.Client{
|
||||
Timeout: timeout,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Heartbeat(ctx context.Context, payload service.RelayHeartbeatPayload) (*service.RelayHeartbeatResponse, error) {
|
||||
resp := APIResponse[service.RelayHeartbeatResponse]{}
|
||||
if err := c.postJSON(ctx, "/api/relay/heartbeat", payload, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Success {
|
||||
return nil, errors.New(resp.Message)
|
||||
}
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.token = strings.TrimSpace(token)
|
||||
slog.Debug("http client token updated")
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
req.Header.Set("X-Agent-Token", c.token)
|
||||
return c.do(req, 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")
|
||||
req.Header.Set("X-Agent-Token", c.token)
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
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(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}(res.Body)
|
||||
if res.StatusCode != http.StatusOK {
|
||||
slog.Warn("http request returned non-200", "method", req.Method, "path", req.URL.Path, "status", res.Status)
|
||||
return errors.New(res.Status)
|
||||
}
|
||||
if target == nil {
|
||||
var wrapper APIResponse[json.RawMessage]
|
||||
if err = json.NewDecoder(res.Body).Decode(&wrapper); err != nil {
|
||||
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
if !wrapper.Success {
|
||||
slog.Warn("http api response failed", "method", req.Method, "path", req.URL.Path, "message", wrapper.Message)
|
||||
return errors.New(wrapper.Message)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err = json.NewDecoder(res.Body).Decode(target); err != nil {
|
||||
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,325 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/config"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/frps"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/state"
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
)
|
||||
|
||||
func BuildProfile(cfg *config.Config, stateStore *state.Store) *service.AgentNodeSystemProfile {
|
||||
profile := collectProfile(cfg)
|
||||
if profile == nil || stateStore == nil {
|
||||
return profile
|
||||
}
|
||||
fingerprint := fingerprintProfile(profile)
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
return profile
|
||||
}
|
||||
if snapshot.LastProfileFingerprint == fingerprint {
|
||||
return nil
|
||||
}
|
||||
snapshot.LastProfileFingerprint = fingerprint
|
||||
if err = stateStore.Save(snapshot); err != nil {
|
||||
return profile
|
||||
}
|
||||
return profile
|
||||
}
|
||||
|
||||
func BuildSnapshot(cfg *config.Config, stateStore *state.Store) *service.AgentNodeMetricSnapshot {
|
||||
now := time.Now().UTC()
|
||||
metric := &service.AgentNodeMetricSnapshot{CapturedAtUnix: now.Unix()}
|
||||
|
||||
metric.MemoryTotalBytes, metric.MemoryUsedBytes = readMemInfo()
|
||||
metric.StorageTotalBytes, metric.StorageUsedBytes = statFilesystem(cfg.DataDir)
|
||||
metric.NetworkRxBytes, metric.NetworkTxBytes = readLinuxNetworkTotals()
|
||||
metric.DiskReadBytes, metric.DiskWriteBytes = readLinuxDiskTotals()
|
||||
|
||||
if stateStore == nil {
|
||||
return metric
|
||||
}
|
||||
totalCPU, idleCPU := readLinuxCPUStat()
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
return metric
|
||||
}
|
||||
if snapshot.LastCPUStatTotal > 0 && totalCPU > snapshot.LastCPUStatTotal && idleCPU >= snapshot.LastCPUStatIdle {
|
||||
deltaTotal := totalCPU - snapshot.LastCPUStatTotal
|
||||
deltaIdle := idleCPU - snapshot.LastCPUStatIdle
|
||||
if deltaTotal > 0 && deltaIdle <= deltaTotal {
|
||||
metric.CPUUsagePercent = float64(deltaTotal-deltaIdle) / float64(deltaTotal) * 100
|
||||
}
|
||||
}
|
||||
snapshot.LastCPUStatTotal = totalCPU
|
||||
snapshot.LastCPUStatIdle = idleCPU
|
||||
snapshot.LastMetricAtUnix = now.Unix()
|
||||
_ = stateStore.Save(snapshot)
|
||||
return metric
|
||||
}
|
||||
|
||||
func BuildHealthEvents(status frps.RuntimeStatus) []service.AgentNodeHealthEvent {
|
||||
if strings.TrimSpace(status.Status) == "healthy" {
|
||||
return []service.AgentNodeHealthEvent{}
|
||||
}
|
||||
message := strings.TrimSpace(status.LastError)
|
||||
if message == "" {
|
||||
message = "frps runtime is not healthy"
|
||||
}
|
||||
return []service.AgentNodeHealthEvent{{
|
||||
EventType: "frps_unhealthy",
|
||||
Severity: "critical",
|
||||
Message: message,
|
||||
TriggeredAtUnix: time.Now().UTC().Unix(),
|
||||
}}
|
||||
}
|
||||
|
||||
func collectProfile(cfg *config.Config) *service.AgentNodeSystemProfile {
|
||||
hostname, _ := os.Hostname()
|
||||
osName, osVersion := readLinuxOSRelease()
|
||||
totalMemory, _ := readMemInfo()
|
||||
totalDisk, _ := statFilesystem(cfg.DataDir)
|
||||
return &service.AgentNodeSystemProfile{
|
||||
Hostname: strings.TrimSpace(hostname),
|
||||
OSName: osName,
|
||||
OSVersion: osVersion,
|
||||
KernelVersion: readFirstLine("/proc/sys/kernel/osrelease"),
|
||||
Architecture: runtime.GOARCH,
|
||||
CPUModel: readLinuxCPUModel(),
|
||||
CPUCores: runtime.NumCPU(),
|
||||
TotalMemoryBytes: totalMemory,
|
||||
TotalDiskBytes: totalDisk,
|
||||
UptimeSeconds: readLinuxUptimeSeconds(),
|
||||
ReportedAtUnix: time.Now().UTC().Unix(),
|
||||
}
|
||||
}
|
||||
|
||||
func fingerprintProfile(profile *service.AgentNodeSystemProfile) string {
|
||||
raw, err := json.Marshal(profile)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
sum := sha256.Sum256(raw)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func readLinuxOSRelease() (string, string) {
|
||||
file, err := os.Open("/etc/os-release")
|
||||
if err != nil {
|
||||
return runtime.GOOS, ""
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
values := make(map[string]string)
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
key, value, ok := strings.Cut(strings.TrimSpace(scanner.Text()), "=")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
values[key] = strings.Trim(value, `"`)
|
||||
}
|
||||
if pretty := strings.TrimSpace(values["PRETTY_NAME"]); pretty != "" {
|
||||
return pretty, strings.TrimSpace(values["VERSION_ID"])
|
||||
}
|
||||
if name := strings.TrimSpace(values["NAME"]); name != "" {
|
||||
return name, strings.TrimSpace(values["VERSION_ID"])
|
||||
}
|
||||
return runtime.GOOS, ""
|
||||
}
|
||||
|
||||
func readLinuxCPUModel() string {
|
||||
file, err := os.Open("/proc/cpuinfo")
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
defer 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 ""
|
||||
}
|
||||
|
||||
func readMemInfo() (int64, int64) {
|
||||
file, err := os.Open("/proc/meminfo")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
var totalKB, availableKB int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
if strings.HasPrefix(line, "MemTotal:") {
|
||||
totalKB = parseMemInfoValue(line)
|
||||
}
|
||||
if strings.HasPrefix(line, "MemAvailable:") {
|
||||
availableKB = parseMemInfoValue(line)
|
||||
}
|
||||
}
|
||||
total := totalKB * 1024
|
||||
used := total - availableKB*1024
|
||||
if used < 0 {
|
||||
used = 0
|
||||
}
|
||||
return total, used
|
||||
}
|
||||
|
||||
func parseMemInfoValue(line string) int64 {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < 2 {
|
||||
return 0
|
||||
}
|
||||
value, err := strconv.ParseInt(fields[1], 10, 64)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
func readLinuxCPUStat() (uint64, uint64) {
|
||||
content, err := os.ReadFile("/proc/stat")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
for _, line := range strings.Split(string(content), "\n") {
|
||||
if !strings.HasPrefix(line, "cpu ") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < 5 {
|
||||
return 0, 0
|
||||
}
|
||||
var total uint64
|
||||
for index := 1; index < len(fields); index++ {
|
||||
value, err := strconv.ParseUint(fields[index], 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
|
||||
}
|
||||
|
||||
func readLinuxNetworkTotals() (int64, int64) {
|
||||
file, err := os.Open("/proc/net/dev")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
var rx, tx int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
name, data, ok := strings.Cut(strings.TrimSpace(scanner.Text()), ":")
|
||||
if !ok || strings.TrimSpace(name) == "lo" {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(data)
|
||||
if len(fields) < 16 {
|
||||
continue
|
||||
}
|
||||
if value, err := strconv.ParseInt(fields[0], 10, 64); err == nil {
|
||||
rx += value
|
||||
}
|
||||
if value, err := strconv.ParseInt(fields[8], 10, 64); err == nil {
|
||||
tx += value
|
||||
}
|
||||
}
|
||||
return rx, tx
|
||||
}
|
||||
|
||||
func readLinuxDiskTotals() (int64, int64) {
|
||||
file, err := os.Open("/proc/diskstats")
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
var readBytes, writeBytes int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
fields := strings.Fields(scanner.Text())
|
||||
if len(fields) < 14 || shouldSkipDiskDevice(fields[2]) {
|
||||
continue
|
||||
}
|
||||
if value, err := strconv.ParseInt(fields[5], 10, 64); err == nil {
|
||||
readBytes += value * 512
|
||||
}
|
||||
if value, err := strconv.ParseInt(fields[9], 10, 64); err == nil {
|
||||
writeBytes += value * 512
|
||||
}
|
||||
}
|
||||
return readBytes, writeBytes
|
||||
}
|
||||
|
||||
func shouldSkipDiskDevice(device string) bool {
|
||||
return device == "" || strings.HasPrefix(device, "loop") || strings.HasPrefix(device, "ram") || strings.HasPrefix(device, "dm-")
|
||||
}
|
||||
|
||||
func statFilesystem(path string) (int64, int64) {
|
||||
if strings.TrimSpace(path) == "" {
|
||||
path = string(os.PathSeparator)
|
||||
}
|
||||
var stat syscall.Statfs_t
|
||||
if err := syscall.Statfs(filepath.Clean(path), &stat); err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
total := int64(stat.Blocks) * int64(stat.Bsize)
|
||||
used := total - int64(stat.Bavail)*int64(stat.Bsize)
|
||||
if used < 0 {
|
||||
used = 0
|
||||
}
|
||||
return total, used
|
||||
}
|
||||
|
||||
func readFirstLine(path string) string {
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(string(content))
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/config"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/frps"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/heartbeat"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/httpclient"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/state"
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/wsclient"
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
)
|
||||
|
||||
type Runner struct {
|
||||
Config *config.Config
|
||||
StateStore *state.Store
|
||||
HeartbeatService *heartbeat.Service
|
||||
FrpsManager *frps.Manager
|
||||
WebSocketService *wsclient.Client
|
||||
HttpClient *httpclient.Client
|
||||
}
|
||||
|
||||
func (r *Runner) Run(ctx context.Context) error {
|
||||
// Start heartbeat loop in background
|
||||
go r.HeartbeatService.Run(ctx)
|
||||
|
||||
// WebSocket reconnection loop
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
r.FrpsManager.Stop()
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
conn, err := r.WebSocketService.Connect(ctx)
|
||||
if err != nil {
|
||||
slog.Error("relay ws connect failed, will retry", "error", err)
|
||||
r.sleepContext(ctx, 5*time.Second)
|
||||
continue
|
||||
}
|
||||
|
||||
r.handleConnection(ctx, conn)
|
||||
_ = conn.Close()
|
||||
slog.Info("relay ws connection closed, reconnecting...")
|
||||
r.sleepContext(ctx, 2*time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
type relayWSHandler struct {
|
||||
runner *Runner
|
||||
}
|
||||
|
||||
func (h *relayWSHandler) OnConnect(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *relayWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||
switch msg.Type {
|
||||
case "relay_config":
|
||||
var cfg service.RelayConfig
|
||||
if err := json.Unmarshal(msg.Payload, &cfg); err != nil {
|
||||
slog.Error("failed to unmarshal relay_config", "error", err)
|
||||
return nil
|
||||
}
|
||||
h.runner.FrpsManager.UpdateConfig(&cfg)
|
||||
default:
|
||||
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *relayWSHandler) OnClose(err error) {
|
||||
slog.Error("relay ws receive failed", "error", err)
|
||||
}
|
||||
|
||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
||||
_ = conn.RunReceiveLoop(ctx, &relayWSHandler{runner: r})
|
||||
}
|
||||
|
||||
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(d):
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"os"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type Store struct {
|
||||
path string
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
type State struct {
|
||||
LastAuthToken string `json:"last_auth_token"`
|
||||
LastProfileFingerprint string `json:"last_profile_fingerprint"`
|
||||
LastCPUStatTotal uint64 `json:"last_cpu_stat_total"`
|
||||
LastCPUStatIdle uint64 `json:"last_cpu_stat_idle"`
|
||||
LastMetricAtUnix int64 `json:"last_metric_at_unix"`
|
||||
}
|
||||
|
||||
func NewStore(path string) *Store {
|
||||
return &Store{
|
||||
path: path,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Store) Load() (*State, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
data, err := os.ReadFile(s.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return &State{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var state State
|
||||
if err := json.Unmarshal(data, &state); err != nil {
|
||||
return &State{}, nil // Return empty state on corrupted file
|
||||
}
|
||||
return &state, nil
|
||||
}
|
||||
|
||||
func (s *Store) Save(state *State) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
data, err := json.MarshalIndent(state, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
slog.Debug("saving relay state")
|
||||
return os.WriteFile(s.path, data, 0644)
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
//go:build !windows
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"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: %v", 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: %v", 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 {
|
||||
return fmt.Errorf("exec restart: %w", err)
|
||||
}
|
||||
return fmt.Errorf("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,53 @@
|
||||
//go:build windows
|
||||
|
||||
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,371 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/utils"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-relay/internal/config"
|
||||
)
|
||||
|
||||
const maxChecksumAssetSize = 64 * 1024
|
||||
|
||||
var replaceAndRestartFunc = replaceAndRestart
|
||||
|
||||
type Service struct {
|
||||
httpClient *http.Client
|
||||
lastCheckKey string
|
||||
}
|
||||
|
||||
func New() *Service {
|
||||
return &Service{
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
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"`
|
||||
}
|
||||
|
||||
type UpdateOptions struct {
|
||||
Channel string
|
||||
TagName string
|
||||
Force bool
|
||||
}
|
||||
|
||||
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(config.Version)
|
||||
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("relay update available", "from", localVersion, "to", remoteVersion)
|
||||
assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
checksumAssetName := assetName + ".sha256"
|
||||
|
||||
var downloadURL string
|
||||
var checksumURL string
|
||||
for _, asset := range release.Assets {
|
||||
switch asset.Name {
|
||||
case assetName:
|
||||
downloadURL = asset.BrowserDownloadURL
|
||||
case checksumAssetName:
|
||||
checksumURL = asset.BrowserDownloadURL
|
||||
}
|
||||
}
|
||||
if downloadURL == "" {
|
||||
s.lastCheckKey = checkKey
|
||||
return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
|
||||
}
|
||||
if checksumURL == "" {
|
||||
return fmt.Errorf("no matching checksum asset %q in release %s", checksumAssetName, release.TagName)
|
||||
}
|
||||
|
||||
expectedChecksum, err := s.downloadChecksum(ctx, checksumURL, assetName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("download checksum: %w", 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 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(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}(resp.Body)
|
||||
|
||||
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) 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 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.Split(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 "", fmt.Errorf("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 fmt.Errorf("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 resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("download returned %s", resp.Status)
|
||||
}
|
||||
|
||||
tmpPath := targetPath + ".update"
|
||||
if runtime.GOOS == "windows" && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") {
|
||||
tmpPath += ".exe"
|
||||
}
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
|
||||
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, 0o755); err != nil && runtime.GOOS != "windows" {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("set executable permission: %w", err)
|
||||
}
|
||||
|
||||
slog.Info("relay binary updated, restarting")
|
||||
return replaceAndRestartFunc(targetPath, tmpPath)
|
||||
}
|
||||
|
||||
func assetNameForGOOSGOARCH(goos string, goarch string) string {
|
||||
name := fmt.Sprintf("openflare-relay-%s-%s", goos, goarch)
|
||||
if goos == "windows" {
|
||||
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 utils.CompareVersions(local, remote)
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/service"
|
||||
shared "github.com/rain-kl/openflare/openflare-server/utils/wsclient"
|
||||
)
|
||||
|
||||
type WSMessage = shared.WSMessage
|
||||
type MessageHandler = shared.MessageHandler
|
||||
|
||||
type Client struct {
|
||||
sharedClient *shared.Client
|
||||
}
|
||||
|
||||
type Connection struct {
|
||||
sharedConn *shared.Connection
|
||||
}
|
||||
|
||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
sharedClient: shared.New(shared.Config{
|
||||
BaseURL: baseURL,
|
||||
Token: token,
|
||||
Timeout: timeout,
|
||||
HeaderKey: "X-Agent-Token",
|
||||
WSPath: "/api/relay/ws",
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.sharedClient.SetToken(token)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (conn *Connection) SendPing() error {
|
||||
return conn.sharedConn.SendMessage("ping", nil)
|
||||
}
|
||||
|
||||
func (conn *Connection) SendPong() error {
|
||||
return conn.sharedConn.SendMessage("pong", nil)
|
||||
}
|
||||
|
||||
func (conn *Connection) Receive() (service.WSMessage, error) {
|
||||
var raw struct {
|
||||
Type string `json:"type"`
|
||||
Payload json.RawMessage `json:"payload,omitempty"`
|
||||
}
|
||||
if err := conn.sharedConn.Receive(&raw); err != nil {
|
||||
return service.WSMessage{}, err
|
||||
}
|
||||
return service.WSMessage{
|
||||
Type: raw.Type,
|
||||
Payload: raw.Payload,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
|
||||
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||
}
|
||||
|
||||
func (conn *Connection) Close() error {
|
||||
return conn.sharedConn.Close()
|
||||
}
|
||||
Reference in New Issue
Block a user