mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 15:06:37 +08:00
refactor(repo): consolidate openflare-server to root and move subprojects to internal/apps
- Merge all files inside openflare-server to the repository root directory. - Relocate agent, relay, and flared subprojects from internal/ to internal/apps/. - Combine docker-compose files and update build context paths to root. - Update GitHub workflows and Dockerfiles to refer to new directories and package names. - Rewrite Go package imports across all files. - Resolve database renew test race condition and clean up docs.
This commit is contained in:
@@ -0,0 +1,158 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
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"`
|
||||
TunnelToken string `json:"tunnel_token"`
|
||||
FrpcPath string `json:"frpc_path"`
|
||||
DataDir string `json:"data_dir"`
|
||||
StatePath string `json:"state_path"`
|
||||
HeartbeatInterval MillisecondDuration `json:"heartbeat_interval"`
|
||||
SyncInterval MillisecondDuration `json:"sync_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_TUNNEL_TOKEN",
|
||||
"OPENFLARE_DATA_DIR",
|
||||
"OPENFLARE_FRPC_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_TUNNEL_TOKEN", &cfg.TunnelToken)
|
||||
overrideString("OPENFLARE_DATA_DIR", &cfg.DataDir)
|
||||
overrideString("OPENFLARE_FRPC_PATH", &cfg.FrpcPath)
|
||||
}
|
||||
|
||||
func applyDefaults(cfg *Config, baseDir string) {
|
||||
baseDir = filepath.Clean(baseDir)
|
||||
if cfg.FrpcPath == "" {
|
||||
cfg.FrpcPath = "frpc" // rely on PATH
|
||||
}
|
||||
if cfg.DataDir == "" {
|
||||
cfg.DataDir = filepath.Join(baseDir, "data")
|
||||
}
|
||||
if cfg.StatePath == "" {
|
||||
cfg.StatePath = filepath.Join(cfg.DataDir, "flared-state.json")
|
||||
}
|
||||
if cfg.HeartbeatInterval <= 0 {
|
||||
cfg.HeartbeatInterval = MillisecondDuration(10 * time.Second)
|
||||
}
|
||||
if cfg.SyncInterval <= 0 {
|
||||
cfg.SyncInterval = MillisecondDuration(30 * 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.TunnelToken) == "" {
|
||||
return errors.New("tunnel_token 不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cfg *Config) InitialAuthToken() string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(cfg.TunnelToken)
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
package config
|
||||
|
||||
var Version = "dev"
|
||||
@@ -0,0 +1,85 @@
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/frpc"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/heartbeat"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/httpclient"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/sync"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/wsclient"
|
||||
)
|
||||
|
||||
type Runner struct {
|
||||
Config *config.Config
|
||||
HeartbeatService *heartbeat.Service
|
||||
FrpcManager *frpc.Manager
|
||||
SyncService *sync.Service
|
||||
WebSocketService *wsclient.Client
|
||||
HttpClient *httpclient.Client
|
||||
}
|
||||
|
||||
func (r *Runner) Run(ctx context.Context) error {
|
||||
// Start background services
|
||||
go r.HeartbeatService.Run(ctx)
|
||||
go r.SyncService.Run(ctx)
|
||||
|
||||
// WebSocket reconnection loop
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
r.FrpcManager.Stop()
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
conn, err := r.WebSocketService.Connect(ctx)
|
||||
if err != nil {
|
||||
slog.Error("flared ws connect failed, will retry", "error", err)
|
||||
r.sleepContext(ctx, 5*time.Second)
|
||||
continue
|
||||
}
|
||||
|
||||
r.handleConnection(ctx, conn)
|
||||
_ = conn.Close()
|
||||
slog.Info("flared ws connection closed, reconnecting...")
|
||||
r.sleepContext(ctx, 2*time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
type flaredWSHandler struct {
|
||||
runner *Runner
|
||||
}
|
||||
|
||||
func (h *flaredWSHandler) OnConnect(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *flaredWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||
switch msg.Type {
|
||||
case "active_config":
|
||||
slog.Info("received config update notification from server")
|
||||
h.runner.SyncService.Trigger()
|
||||
default:
|
||||
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *flaredWSHandler) OnClose(err error) {
|
||||
slog.Error("flared ws receive failed", "error", err)
|
||||
}
|
||||
|
||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
||||
_ = conn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r})
|
||||
}
|
||||
|
||||
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(d):
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,350 @@
|
||||
package frpc
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
type Manager struct {
|
||||
cfg *config.Config
|
||||
processes map[string]*Process
|
||||
mu sync.RWMutex
|
||||
|
||||
currentVersion string
|
||||
currentChecksum string
|
||||
}
|
||||
|
||||
type Process struct {
|
||||
RelayID string
|
||||
Cmd *exec.Cmd
|
||||
Cancel context.CancelFunc
|
||||
Status string
|
||||
StartTime time.Time
|
||||
LastError string
|
||||
}
|
||||
|
||||
func NewManager(cfg *config.Config) *Manager {
|
||||
return &Manager{
|
||||
cfg: cfg,
|
||||
processes: make(map[string]*Process),
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) GetVersion() string {
|
||||
cmd := exec.Command(m.cfg.FrpcPath, "-v")
|
||||
out, err := cmd.Output()
|
||||
if err != nil {
|
||||
return "unknown"
|
||||
}
|
||||
return strings.TrimSpace(string(out))
|
||||
}
|
||||
|
||||
func (m *Manager) GetConnectedRelays() []service.FlaredConnectedRelay {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
result := make([]service.FlaredConnectedRelay, 0, len(m.processes))
|
||||
for relayID, proc := range m.processes {
|
||||
result = append(result, service.FlaredConnectedRelay{
|
||||
RelayNodeID: relayID,
|
||||
Status: proc.Status,
|
||||
})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (m *Manager) GetCurrentConfigVersion() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.currentVersion
|
||||
}
|
||||
|
||||
func (m *Manager) GetCurrentConfigChecksum() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.currentChecksum
|
||||
}
|
||||
|
||||
func (m *Manager) UpdateConfig(ctx context.Context, newConfig *service.FlaredTunnelConfigResponse) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if newConfig == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
versionChanged := newConfig.Version != m.currentVersion || newConfig.Checksum != m.currentChecksum
|
||||
if versionChanged {
|
||||
slog.Info("applying new tunnel config", "version", newConfig.Version)
|
||||
} else {
|
||||
slog.Debug("tunnel config version unchanged, ensuring processes are running", "version", newConfig.Version)
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(m.cfg.DataDir, 0o755); err != nil {
|
||||
return fmt.Errorf("create data dir failed: %w", err)
|
||||
}
|
||||
|
||||
activeRelays := make(map[string]struct{})
|
||||
|
||||
for _, relay := range newConfig.Relays {
|
||||
activeRelays[relay.RelayNodeID] = struct{}{}
|
||||
tomlContent := buildFrpcToml(relay, newConfig.Proxies)
|
||||
configPath := filepath.Join(m.cfg.DataDir, fmt.Sprintf("frpc_%s.toml", relay.RelayNodeID))
|
||||
|
||||
needsRestart := false
|
||||
existingData, err := os.ReadFile(configPath)
|
||||
if err != nil || string(existingData) != tomlContent {
|
||||
// 配置文件不存在或内容有变化,需要写入并重启
|
||||
needsRestart = true
|
||||
}
|
||||
|
||||
if needsRestart {
|
||||
if err := os.WriteFile(configPath, []byte(tomlContent), 0o644); err != nil {
|
||||
slog.Error("failed to write frpc config", "relay_id", relay.RelayNodeID, "error", err)
|
||||
continue
|
||||
}
|
||||
m.restartProcess(ctx, relay.RelayNodeID, configPath)
|
||||
} else if _, ok := m.processes[relay.RelayNodeID]; !ok {
|
||||
// 配置未变但进程不存在(如重启后),直接启动进程
|
||||
slog.Info("frpc process missing, starting", "relay_id", relay.RelayNodeID)
|
||||
m.restartProcess(ctx, relay.RelayNodeID, configPath)
|
||||
}
|
||||
}
|
||||
|
||||
// Stop obsolete processes
|
||||
for relayID, proc := range m.processes {
|
||||
if _, ok := activeRelays[relayID]; !ok {
|
||||
slog.Info("stopping obsolete frpc process", "relay_id", relayID)
|
||||
proc.Cancel()
|
||||
pidPath := filepath.Join(m.cfg.DataDir, fmt.Sprintf("frpc_%s.pid", relayID))
|
||||
_ = os.Remove(pidPath)
|
||||
delete(m.processes, relayID)
|
||||
}
|
||||
}
|
||||
|
||||
if versionChanged {
|
||||
m.currentVersion = newConfig.Version
|
||||
m.currentChecksum = newConfig.Checksum
|
||||
return m.saveState()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath string) {
|
||||
pidPath := filepath.Join(m.cfg.DataDir, fmt.Sprintf("frpc_%s.pid", relayID))
|
||||
if proc, ok := m.processes[relayID]; ok {
|
||||
proc.Cancel()
|
||||
_ = os.Remove(pidPath)
|
||||
}
|
||||
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
procCtx, cancel := context.WithCancel(ctx)
|
||||
proc := &Process{
|
||||
RelayID: relayID,
|
||||
Cancel: cancel,
|
||||
Status: "starting",
|
||||
StartTime: time.Now(),
|
||||
}
|
||||
m.processes[relayID] = proc
|
||||
|
||||
go func() {
|
||||
backoff := 1 * time.Second
|
||||
const maxBackoff = 60 * time.Second
|
||||
|
||||
for {
|
||||
m.mu.Lock()
|
||||
if procCtx.Err() != nil {
|
||||
m.mu.Unlock()
|
||||
return
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
ensureNoOrphanProcess(pidPath)
|
||||
|
||||
cmd := exec.CommandContext(procCtx, m.cfg.FrpcPath, "-c", configPath)
|
||||
|
||||
m.mu.Lock()
|
||||
proc.Cmd = cmd
|
||||
proc.Status = "running"
|
||||
m.mu.Unlock()
|
||||
|
||||
startedAt := time.Now()
|
||||
err := cmd.Start()
|
||||
if err == nil {
|
||||
_ = os.WriteFile(pidPath, []byte(fmt.Sprintf("%d", cmd.Process.Pid)), 0o644)
|
||||
err = cmd.Wait()
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
|
||||
m.mu.Lock()
|
||||
if procCtx.Err() != nil {
|
||||
proc.Status = "stopped"
|
||||
m.mu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
proc.LastError = err.Error()
|
||||
proc.Status = "error"
|
||||
slog.Error("frpc process exited unexpectedly", "relay_id", relayID, "error", err)
|
||||
} else {
|
||||
proc.Status = "stopped"
|
||||
proc.LastError = "exited unexpectedly with code 0"
|
||||
slog.Warn("frpc process exited unexpectedly with code 0", "relay_id", relayID)
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
if time.Since(startedAt) >= 10*time.Second {
|
||||
backoff = 1 * time.Second
|
||||
}
|
||||
|
||||
select {
|
||||
case <-procCtx.Done():
|
||||
return
|
||||
case <-time.After(backoff):
|
||||
backoff = backoff * 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (m *Manager) Stop() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for relayID, proc := range m.processes {
|
||||
if proc != nil && proc.Cancel != nil {
|
||||
proc.Cancel()
|
||||
}
|
||||
pidPath := filepath.Join(m.cfg.DataDir, fmt.Sprintf("frpc_%s.pid", relayID))
|
||||
_ = os.Remove(pidPath)
|
||||
delete(m.processes, relayID)
|
||||
}
|
||||
}
|
||||
|
||||
func buildFrpcToml(relay service.FlaredRelayInfo, proxies []service.FlaredProxyEntry) string {
|
||||
var buf bytes.Buffer
|
||||
|
||||
host, port := parseAddr(relay.Address)
|
||||
|
||||
buf.WriteString(fmt.Sprintf(`serverAddr = "%s"
|
||||
serverPort = %s
|
||||
`, host, port))
|
||||
|
||||
if relay.AuthToken != "" {
|
||||
buf.WriteString(fmt.Sprintf(`auth.method = "token"
|
||||
auth.token = "%s"
|
||||
`, relay.AuthToken))
|
||||
}
|
||||
|
||||
if relay.ProxyURL != "" {
|
||||
buf.WriteString(fmt.Sprintf(`transport.proxyURL = "%s"
|
||||
`, relay.ProxyURL))
|
||||
}
|
||||
|
||||
buf.WriteString("\n")
|
||||
|
||||
for _, proxy := range proxies {
|
||||
buf.WriteString(fmt.Sprintf("[[proxies]]\nname = \"%s\"\ntype = \"%s\"\nlocalIP = \"%s\"\nlocalPort = %d\n",
|
||||
proxy.Name, proxy.Type, proxy.LocalAddr, proxy.LocalPort))
|
||||
if len(proxy.CustomDomains) > 0 {
|
||||
buf.WriteString(fmt.Sprintf("customDomains = [\"%s\"]\n", strings.Join(proxy.CustomDomains, "\", \"")))
|
||||
}
|
||||
buf.WriteString("\n")
|
||||
}
|
||||
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func parseAddr(addr string) (string, string) {
|
||||
addr = strings.TrimSpace(addr)
|
||||
if addr == "" {
|
||||
return "127.0.0.1", "7000"
|
||||
}
|
||||
host, port, err := net.SplitHostPort(addr)
|
||||
if err == nil {
|
||||
return strings.Trim(host, "[]"), port
|
||||
}
|
||||
lastColon := strings.LastIndex(addr, ":")
|
||||
if lastColon > 0 && strings.Count(addr, ":") == 1 {
|
||||
return addr[:lastColon], addr[lastColon+1:]
|
||||
}
|
||||
return addr, "7000"
|
||||
}
|
||||
|
||||
// State persistence
|
||||
type ManagerState struct {
|
||||
Version string
|
||||
Checksum string
|
||||
}
|
||||
|
||||
func (m *Manager) saveState() error {
|
||||
state := ManagerState{
|
||||
Version: m.currentVersion,
|
||||
Checksum: m.currentChecksum,
|
||||
}
|
||||
data, err := json.MarshalIndent(state, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(m.cfg.StatePath, data, 0o644)
|
||||
}
|
||||
|
||||
func (m *Manager) LoadState() error {
|
||||
data, err := os.ReadFile(m.cfg.StatePath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
var state ManagerState
|
||||
if err := json.Unmarshal(data, &state); err != nil {
|
||||
return err
|
||||
}
|
||||
m.mu.Lock()
|
||||
m.currentVersion = state.Version
|
||||
m.currentChecksum = state.Checksum
|
||||
m.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
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,387 @@
|
||||
package frpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
// 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_frpc")
|
||||
|
||||
// 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, relayID string, expectedStatus string, timeout time.Duration) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
m.mu.RLock()
|
||||
proc, ok := m.processes[relayID]
|
||||
m.mu.RUnlock()
|
||||
if ok && proc.Status == expectedStatus {
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
m.mu.RLock()
|
||||
proc, ok := m.processes[relayID]
|
||||
var got string
|
||||
var errStr string
|
||||
if ok {
|
||||
got = proc.Status
|
||||
errStr = proc.LastError
|
||||
} else {
|
||||
got = "not_found"
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
t.Fatalf("expected status eventually %s, got %s (err: %s)", expectedStatus, got, errStr)
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
err := m.UpdateConfig(context.Background(), newConfig)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to UpdateConfig: %v", err)
|
||||
}
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
|
||||
proc.Cancel()
|
||||
assertStatusEventually(t, m, "relay-1", "stopped", 4*time.Second) // wait for clean stop
|
||||
}
|
||||
|
||||
func TestStartProcessFailureAndBackoff(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
invalidScriptPath := filepath.Join(dir, "non_existent_frpc")
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: invalidScriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
_ = m.UpdateConfig(context.Background(), newConfig)
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "error", 4*time.Second)
|
||||
|
||||
// Correct the path to dummy script
|
||||
scriptPath, _ := setupDummyScript(t)
|
||||
writeControl(t, filepath.Dir(scriptPath), 0, 5)
|
||||
|
||||
m.mu.Lock()
|
||||
m.cfg.FrpcPath = scriptPath
|
||||
m.mu.Unlock()
|
||||
|
||||
// Wait for backoff retry (1s backoff)
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
proc.Cancel()
|
||||
}
|
||||
|
||||
func TestUnexpectedExit0CPUProtection(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Start with immediate exit code 0
|
||||
writeControl(t, dir, 0, 0)
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
_ = m.UpdateConfig(context.Background(), newConfig)
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "stopped", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
if !strings.Contains(proc.LastError, "exited unexpectedly with code 0") {
|
||||
t.Errorf("expected LastError to record exit status 0 warning, got %s", proc.LastError)
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
proc.Cancel()
|
||||
}
|
||||
|
||||
func TestBackoffReset(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Rapid exit code 1 to increase backoff
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
_ = m.UpdateConfig(context.Background(), newConfig)
|
||||
|
||||
// Wait to crash
|
||||
assertStatusEventually(t, m, "relay-1", "error", 4*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, "relay-1", "running", 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, "relay-1", "error", 4*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 running now.
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
proc.Cancel()
|
||||
}
|
||||
|
||||
func TestUpdateConfigKillsOrphanProcessBeforeRestart(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
|
||||
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()
|
||||
}
|
||||
})
|
||||
|
||||
pidPath := filepath.Join(dir, "frpc_relay-1.pid")
|
||||
if err := os.WriteFile(pidPath, []byte(fmt.Sprintf("%d", orphan.Process.Pid)), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed orphan pid file: %v", err)
|
||||
}
|
||||
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
if err := m.UpdateConfig(context.Background(), newConfig); err != nil {
|
||||
t.Fatalf("failed to UpdateConfig: %v", err)
|
||||
}
|
||||
|
||||
assertCommandExitedEventually(t, orphan, 2*time.Second)
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
proc.Cancel()
|
||||
}
|
||||
|
||||
func TestStopCancelsRunningProcesses(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 30)
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
if err := m.UpdateConfig(context.Background(), newConfig); err != nil {
|
||||
t.Fatalf("failed to UpdateConfig: %v", err)
|
||||
}
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
if proc == nil || proc.Cmd == nil {
|
||||
m.mu.RUnlock()
|
||||
t.Fatal("expected running process to have a command handle")
|
||||
}
|
||||
cmd := proc.Cmd
|
||||
m.mu.RUnlock()
|
||||
|
||||
m.Stop()
|
||||
assertCommandExitedEventually(t, cmd, 2*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
if len(m.processes) != 0 {
|
||||
t.Fatalf("expected no managed processes after stop, got %d", len(m.processes))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/frpc"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/httpclient"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/updater"
|
||||
"github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
var (
|
||||
lookupOutboundIP = geoip.GetOutboundIP
|
||||
lookupLocalIP = detectLocalNodeIP
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
client *httpclient.Client
|
||||
frpcManager *frpc.Manager
|
||||
config *config.Config
|
||||
updater *updater.Service
|
||||
}
|
||||
|
||||
func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
frpcManager: manager,
|
||||
config: cfg,
|
||||
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 flared heartbeat")
|
||||
|
||||
ip := detectNodeIP()
|
||||
|
||||
payload := service.FlaredHeartbeatPayload{
|
||||
ClientVersion: config.Version,
|
||||
FrpVersion: s.frpcManager.GetVersion(),
|
||||
IP: ip,
|
||||
TunnelStatus: "running", // TODO implement proper status tracking
|
||||
ConnectedRelays: s.frpcManager.GetConnectedRelays(),
|
||||
CurrentVersion: s.frpcManager.GetCurrentConfigVersion(),
|
||||
CurrentChecksum: s.frpcManager.GetCurrentConfigChecksum(),
|
||||
}
|
||||
|
||||
resp, err := s.client.Heartbeat(ctx, payload)
|
||||
if err != nil {
|
||||
slog.Error("flared heartbeat failed", "error", err)
|
||||
return
|
||||
}
|
||||
slog.Debug("flared heartbeat succeeded")
|
||||
|
||||
if resp != nil && resp.TunnelSettings != nil {
|
||||
s.tryAutoUpdate(ctx, resp.TunnelSettings)
|
||||
}
|
||||
}
|
||||
|
||||
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 client 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("client update check failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func detectNodeIP() string {
|
||||
if ip := detectOutboundNodeIP(); ip != "" {
|
||||
return ip
|
||||
}
|
||||
return lookupLocalIP()
|
||||
}
|
||||
|
||||
func detectOutboundNodeIP() string {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ip, err := lookupOutboundIP(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,143 @@
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
type APIResponse[T any] struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
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.FlaredHeartbeatPayload) (*service.FlaredHeartbeatResponse, error) {
|
||||
resp := APIResponse[service.FlaredHeartbeatResponse]{}
|
||||
if err := c.postJSON(ctx, "/api/v1/tunnel/heartbeat", payload, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := apiError(resp.ErrorMsg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
func (c *Client) GetActiveConfig(ctx context.Context) (*service.FlaredTunnelConfigResponse, error) {
|
||||
resp := APIResponse[service.FlaredTunnelConfigResponse]{}
|
||||
if err := c.getJSON(ctx, "/api/v1/tunnel/config/active", &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := apiError(resp.ErrorMsg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
func (c *Client) ReportApplyLog(ctx context.Context, payload service.ApplyLogPayload) error {
|
||||
resp := APIResponse[any]{}
|
||||
if err := c.postJSON(ctx, "/api/v1/tunnel/apply-log", payload, &resp); err != nil {
|
||||
return err
|
||||
}
|
||||
return apiError(resp.ErrorMsg)
|
||||
}
|
||||
|
||||
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-Tunnel-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-Tunnel-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)
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func apiError(msg string) error {
|
||||
if strings.TrimSpace(msg) == "" {
|
||||
return nil
|
||||
}
|
||||
return errors.New(msg)
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package sync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/frpc"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/httpclient"
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
client *httpclient.Client
|
||||
frpcManager *frpc.Manager
|
||||
config *config.Config
|
||||
triggerCh chan struct{}
|
||||
}
|
||||
|
||||
func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
frpcManager: manager,
|
||||
config: cfg,
|
||||
triggerCh: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) Trigger() {
|
||||
select {
|
||||
case s.triggerCh <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) Run(ctx context.Context) {
|
||||
ticker := time.NewTicker(s.config.SyncInterval.Duration())
|
||||
defer ticker.Stop()
|
||||
|
||||
// initial sync
|
||||
s.doSync(ctx)
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.doSync(ctx)
|
||||
case <-s.triggerCh:
|
||||
s.doSync(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) doSync(ctx context.Context) {
|
||||
slog.Debug("fetching active tunnel config")
|
||||
configResp, err := s.client.GetActiveConfig(ctx)
|
||||
if err != nil {
|
||||
slog.Error("failed to fetch active tunnel config", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
// 不在 sync 层做版本早退,由 frpcManager.UpdateConfig 负责判断。
|
||||
// 原因:重启后进程全部消失,即使版本/checksum 未变,仍需重新拉起 frpc 进程。
|
||||
err = s.frpcManager.UpdateConfig(ctx, configResp)
|
||||
|
||||
result := "success"
|
||||
message := "apply success"
|
||||
if err != nil {
|
||||
result = "failed"
|
||||
message = err.Error()
|
||||
slog.Error("failed to apply tunnel config", "error", err)
|
||||
} else {
|
||||
slog.Info("tunnel config applied successfully", "version", configResp.Version)
|
||||
}
|
||||
|
||||
// Report apply log
|
||||
logPayload := service.ApplyLogPayload{
|
||||
Version: configResp.Version,
|
||||
Result: result,
|
||||
Message: message,
|
||||
Checksum: configResp.Checksum,
|
||||
}
|
||||
if reportErr := s.client.ReportApplyLog(ctx, logPayload); reportErr != nil {
|
||||
slog.Error("failed to report apply log", "error", reportErr)
|
||||
}
|
||||
}
|
||||
@@ -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/Wavelet/pkg/utils"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/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("flared 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("flared binary updated, restarting")
|
||||
return replaceAndRestartFunc(targetPath, tmpPath)
|
||||
}
|
||||
|
||||
func assetNameForGOOSGOARCH(goos string, goarch string) string {
|
||||
name := fmt.Sprintf("openflared-%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"
|
||||
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
shared "github.com/Rain-kl/Wavelet/pkg/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-Tunnel-Token",
|
||||
WSPath: "/api/v1/tunnel/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