From db9a9f98fdf73f24bd3e000167b379a426299ae5 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 19 Jun 2026 14:56:00 +0800 Subject: [PATCH] =?UTF-8?q?refactor(edge):=20=E6=8A=BD=E5=8F=96=E8=BE=B9?= =?UTF-8?q?=E7=BC=98=E8=BF=90=E8=A1=8C=E6=97=B6=E5=85=B1=E4=BA=AB=E5=8C=85?= =?UTF-8?q?=E5=B9=B6=E5=AE=8C=E6=88=90=20Phase=203=20=E9=87=8D=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 internal/apps/edge/,三组件改为薄包装,删除 3000+ 行重复代码 - Agent 心跳周期下沉至 heartbeat/cycle.go - 协议类型迁入 pkg/protocol/agent.go - 补充设计文档与 changelog --- cmd/agent/main.go | 24 +- cmd/flared/main.go | 20 +- cmd/relay/main.go | 20 +- docs/changelog/index.md | 3 + docs/design/edge-runtime-refactor.md | 121 ++++++ docs/design/index.md | 1 + docs/plan/20260619-edge-phase3-tasks.md | 53 +++ docs/plan/index.md | 1 + internal/apps/agent/agent/runner.go | 239 +---------- internal/apps/agent/agent/runner_test.go | 55 ++- internal/apps/agent/config/config.go | 72 +--- internal/apps/agent/config/config_test.go | 22 +- internal/apps/agent/config/duration.go | 53 +-- internal/apps/agent/heartbeat/cycle.go | 217 ++++++++++ internal/apps/agent/httpclient/client.go | 133 +----- internal/apps/agent/logging/logger.go | 28 +- .../apps/agent/observability/collector.go | 280 +------------ internal/apps/agent/protocol/alias.go | 43 ++ internal/apps/agent/updater/restart_unix.go | 51 --- .../apps/agent/updater/restart_windows.go | 53 --- internal/apps/agent/updater/updater.go | 367 +---------------- internal/apps/edge/config/duration.go | 54 +++ internal/apps/edge/config/duration_test.go | 53 +++ internal/apps/edge/heartbeat/autoupdate.go | 41 ++ internal/apps/edge/heartbeat/loop.go | 23 ++ internal/apps/edge/heartbeat/loop_test.go | 41 ++ internal/apps/edge/httpclient/client.go | 132 ++++++ internal/apps/edge/logging/logging.go | 33 ++ internal/apps/edge/nodeip/nodeip.go | 69 ++++ internal/apps/edge/observability/linux.go | 268 ++++++++++++ .../apps/edge/observability/linux_test.go | 78 ++++ internal/apps/edge/runner/wsreconnect.go | 64 +++ .../{flared => edge}/updater/restart_unix.go | 2 +- .../updater/restart_unix_test.go | 2 +- .../updater/restart_windows.go | 2 +- internal/apps/edge/updater/service.go | 381 ++++++++++++++++++ .../updater/service_test.go} | 146 ++++--- internal/apps/flared/config/config.go | 33 +- internal/apps/flared/flared/runner.go | 46 +-- internal/apps/flared/heartbeat/service.go | 119 +----- internal/apps/flared/httpclient/client.go | 107 +---- .../apps/flared/updater/restart_windows.go | 53 --- internal/apps/flared/updater/updater.go | 372 +---------------- internal/apps/relay/config/config.go | 91 +---- internal/apps/relay/heartbeat/service.go | 48 +-- internal/apps/relay/httpclient/client.go | 99 +---- .../apps/relay/observability/collector.go | 238 +---------- internal/apps/relay/relay/runner.go | 46 +-- internal/apps/relay/updater/restart_unix.go | 51 --- internal/apps/relay/updater/updater.go | 372 +---------------- .../agent_api.go => pkg/protocol/agent.go | 18 +- pkg/protocol/agent_test.go | 95 +++++ pkg/protocol/server.go | 5 - 53 files changed, 2091 insertions(+), 2947 deletions(-) create mode 100644 docs/design/edge-runtime-refactor.md create mode 100644 docs/plan/20260619-edge-phase3-tasks.md create mode 100644 internal/apps/agent/heartbeat/cycle.go create mode 100644 internal/apps/agent/protocol/alias.go delete mode 100644 internal/apps/agent/updater/restart_unix.go delete mode 100644 internal/apps/agent/updater/restart_windows.go create mode 100644 internal/apps/edge/config/duration.go create mode 100644 internal/apps/edge/config/duration_test.go create mode 100644 internal/apps/edge/heartbeat/autoupdate.go create mode 100644 internal/apps/edge/heartbeat/loop.go create mode 100644 internal/apps/edge/heartbeat/loop_test.go create mode 100644 internal/apps/edge/httpclient/client.go create mode 100644 internal/apps/edge/logging/logging.go create mode 100644 internal/apps/edge/nodeip/nodeip.go create mode 100644 internal/apps/edge/observability/linux.go create mode 100644 internal/apps/edge/observability/linux_test.go create mode 100644 internal/apps/edge/runner/wsreconnect.go rename internal/apps/{flared => edge}/updater/restart_unix.go (99%) rename internal/apps/{agent => edge}/updater/restart_unix_test.go (99%) rename internal/apps/{relay => edge}/updater/restart_windows.go (99%) create mode 100644 internal/apps/edge/updater/service.go rename internal/apps/{agent/updater/updater_test.go => edge/updater/service_test.go} (68%) delete mode 100644 internal/apps/flared/updater/restart_windows.go delete mode 100644 internal/apps/relay/updater/restart_unix.go rename internal/apps/agent/protocol/agent_api.go => pkg/protocol/agent.go (92%) create mode 100644 pkg/protocol/agent_test.go diff --git a/cmd/agent/main.go b/cmd/agent/main.go index 8324d154..dec1cb3a 100644 --- a/cmd/agent/main.go +++ b/cmd/agent/main.go @@ -92,15 +92,23 @@ func main() { } syncService := syncservice.New(client, runtimeManager, stateStore) syncService.SetPagesDir(cfg.PagesDir) + heartbeatService := heartbeat.New(client) + updateService := updater.New() runner := &agent.Runner{ - Config: cfg, - StateStore: stateStore, - ObservabilityBuffer: observabilityBuffer, - HeartbeatService: heartbeat.New(client), - SyncService: syncService, - Updater: updater.New(), - RuntimeManager: runtimeManager, - WebSocketService: wsClient, + Config: cfg, + StateStore: stateStore, + HeartbeatCycle: &heartbeat.Cycle{ + Config: cfg, + StateStore: stateStore, + ObservabilityBuffer: observabilityBuffer, + Heartbeat: heartbeatService, + Sync: syncService, + Updater: updateService, + }, + HeartbeatService: heartbeatService, + SyncService: syncService, + RuntimeManager: runtimeManager, + WebSocketService: wsClient, } ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) diff --git a/cmd/flared/main.go b/cmd/flared/main.go index c9af5341..bfc4d9df 100644 --- a/cmd/flared/main.go +++ b/cmd/flared/main.go @@ -6,9 +6,9 @@ import ( "log/slog" "os" "os/signal" - "strings" "syscall" + edgelogging "github.com/Rain-kl/Wavelet/internal/apps/edge/logging" "github.com/Rain-kl/Wavelet/internal/apps/flared/config" "github.com/Rain-kl/Wavelet/internal/apps/flared/flared" "github.com/Rain-kl/Wavelet/internal/apps/flared/frpc" @@ -19,10 +19,7 @@ import ( ) func main() { - // Setup simple structured logging - slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ - Level: parseLevel(os.Getenv("LOG_LEVEL")), - }))) + edgelogging.Setup(edgelogging.Options{}) configPath := flag.String("config", "./flared.json", "flared config path") flag.Parse() @@ -72,16 +69,3 @@ func main() { } slog.Info("flared 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 - } -} diff --git a/cmd/relay/main.go b/cmd/relay/main.go index 1ff7c419..1bc24321 100644 --- a/cmd/relay/main.go +++ b/cmd/relay/main.go @@ -6,9 +6,9 @@ import ( "log/slog" "os" "os/signal" - "strings" "syscall" + edgelogging "github.com/Rain-kl/Wavelet/internal/apps/edge/logging" "github.com/Rain-kl/Wavelet/internal/apps/relay/config" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/heartbeat" @@ -19,10 +19,7 @@ import ( ) func main() { - // Setup simple structured logging - slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ - Level: parseLevel(os.Getenv("LOG_LEVEL")), - }))) + edgelogging.Setup(edgelogging.Options{}) configPath := flag.String("config", "./relay.json", "relay config path") flag.Parse() @@ -72,16 +69,3 @@ func main() { } 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 - } -} diff --git a/docs/changelog/index.md b/docs/changelog/index.md index c6fdc8d5..234ba25d 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -18,6 +18,9 @@ sidebar: false ### 变更 +- 边缘组件运行时去重:新增 `internal/apps/edge/` 共享包(`updater`、`httpclient`、`nodeip`、`logging`、`heartbeat/autoupdate`、`runner`),Agent/Relay/Flared 三组件改为薄包装委托,删除约 1100 行重复自更新与 HTTP 传输层代码。 +- 边缘运行时 Phase 3 Batch 1:抽取 `edge/config/duration`、`edge/observability/linux`、`edge/heartbeat/loop`,统一 MillisecondDuration、Linux 指标采集与 relay/flared 心跳循环。 +- 边缘运行时 Phase 3 Batch 2:Agent 心跳周期下沉至 `heartbeat/cycle.go`;`pkg/protocol/agent.go` 统一 Agent 客户端协议类型。 - 合并并简化仓库结构:将 `openflare-server` 单体目录下的所有文件/目录提升至仓库根目录(去除了 `openflare-server` 嵌套层级),保留 `.github` 目录不变;统一配置 `docker-compose.yaml` 及所有 Dockerfile 的构建上下文为根目录。 - 调整子项目结构与包路径:将 `agent`、`relay` 和 `flared` 子项目从 `internal/` 移动至 `internal/apps/`(分别为 `internal/apps/agent`、`internal/apps/relay` 和 `internal/apps/flared`),并递归更新了所有涉及的 Go 导入路径(如 `github.com/Rain-kl/Wavelet/internal/apps/agent` 等)。 - 调整编译产物输出名称与 Makefile: diff --git a/docs/design/edge-runtime-refactor.md b/docs/design/edge-runtime-refactor.md new file mode 100644 index 00000000..0e00a693 --- /dev/null +++ b/docs/design/edge-runtime-refactor.md @@ -0,0 +1,121 @@ +# 边缘运行时重构设计 + +你会学到:Agent、Relay、OpenFlared 三组件的重复代码如何收敛到 `internal/apps/edge/`,以及后续演进路线。 + +--- + +## 背景 + +三类边缘守护进程共享同一运行时骨架: + +```text +配置加载 → HTTP/WS 客户端 → 定时心跳 →(可选)配置同步 → 自更新 → 信号优雅退出 +``` + +重构前,以下模块在三个组件间近乎复制粘贴: + +| 模块 | 重复度 | +| --- | --- | +| `updater/` + `restart_{unix,windows}.go` | ~98% | +| `httpclient` 传输层 (`do/postJSON/getJSON`) | ~90% | +| `tryAutoUpdate` | ~98% | +| `detectNodeIP` | ~95% | +| `relay/flared runner` WS 重连环 | ~85% | +| `parseLevel`(main 内联) | 100% | + +Agent 额外包含 nginx 栈、geoip、观测缓冲等**领域特有**逻辑,不宜强行合并。 + +--- + +## 共享包结构 + +``` +internal/apps/edge/ +├── logging/ # Setup、ParseLevel +├── nodeip/ # Detect、DetectLocal(可注入 LookupOutboundIP) +├── httpclient/ # 基础 HTTP 客户端(鉴权头可配置) +├── updater/ # GitHub Release 自更新 + 二进制替换重启 +├── heartbeat/ # TryAutoUpdate 统一入口 +└── runner/ # WS 重连循环、SleepContext +``` + +### 组件层保留 + +各组件仅保留**薄包装**与**领域逻辑**: + +| 组件 | 保留模块 | +| --- | --- | +| Agent | `nginx/`、`sync/`(OpenResty)、`geoipupdate/`、`agent/runner`(discovery/WS 混合) | +| Relay | `frps/`、`observability/` | +| Flared | `frpc/`、`sync/`(tunnel) | + +各组件 `updater/`、`httpclient/` 变为类型别名 + `New()` 工厂函数。 + +--- + +## API 约定 + +### 自更新 + +```go +edgeupdater.New(edgeupdater.Config{ + LocalVersion: config.Version, + AssetPrefix: "openflare-agent", // relay: openflare-relay, flared: openflared + LogLabel: "agent", +}) +``` + +### HTTP 客户端 + +```go +edgehttp.New(baseURL, token, timeout, "X-Agent-Token") // Agent/Relay +edgehttp.New(baseURL, token, timeout, "X-Tunnel-Token") // Flared +``` + +### 节点 IP 探测 + +```go +nodeip.Detect() // outbound → local 回退 +``` + +测试可通过替换 `nodeip.LookupOutboundIP` / `nodeip.LookupLocalIP` 注入桩。 + +--- + +## 已完成(Phase 0–2) + +- [x] `edge/updater` — 三组件 updater 收敛(删除 ~1100 行重复) +- [x] `edge/logging` — relay/flared main 统一日志初始化 +- [x] `edge/nodeip` — 删除三处 detectNodeIP 重复 +- [x] `edge/httpclient` — 三组件 HTTP 传输层收敛 +- [x] `edge/heartbeat/autoupdate` — tryAutoUpdate 统一 +- [x] `edge/runner` — relay/flared WS 重连环收敛 + +--- + +## 已完成(Phase 3 Batch 1) + +- [x] `edge/config/duration.go` — MillisecondDuration 三处合并(含 MarshalJSON) +- [x] `edge/observability/linux.go` — agent/relay collector 底层 Linux 指标采集收敛 +- [x] `edge/heartbeat/loop.go` — relay/flared 心跳 ticker 循环统一 + +## 已完成(Phase 3 Batch 2) + +- [x] `heartbeat/cycle.go` — Agent HTTP 心跳周期从 runner 下沉(payload 构建、同步、自动更新) +- [x] `pkg/protocol/agent.go` — Agent 客户端协议类型迁入公共包,`internal/apps/agent/protocol` 保留别名 re-export + +## 可选后续 + +| 项 | 说明 | +| --- | --- | +| Server 侧协议统一 | 评估 `internal/apps/openflare/agent` 与 `pkg/protocol` 类型去重 | +| `wsclient` 薄包装收敛 | relay/flared/agent wsclient 配置表化 | + +--- + +## 迁移原则 + +1. **领域逻辑不下沉**:nginx/frps/frpc/sync 核心业务保留在各自组件。 +2. **鉴权头显式传入**:禁止 httpclient 默认 Token Header,避免 Agent/Tunnel 混用。 +3. **小步 PR**:每阶段独立可测,自更新路径需集成验证。 +4. **测试随包迁移**:updater 测试已迁至 `edge/updater/`。 \ No newline at end of file diff --git a/docs/design/index.md b/docs/design/index.md index 9a4958a1..df2b6b51 100644 --- a/docs/design/index.md +++ b/docs/design/index.md @@ -189,3 +189,4 @@ OpenFlare 已收敛为**单 monorepo**(Go 模块 `github.com/Rain-kl/Wavelet` * 发布、同步、回滚与 Agent 模型变化:更新 [Agent 与发布模型](./agent-design.md)。 * 部署方式变化:更新 [部署说明](../deployment/deployment.md) 与 README。 * 配置项变化:更新 [配置项参考](../reference/configuration.md)。 +* 边缘组件共享运行时重构:更新 [边缘运行时重构](./edge-runtime-refactor.md)。 diff --git a/docs/plan/20260619-edge-phase3-tasks.md b/docs/plan/20260619-edge-phase3-tasks.md new file mode 100644 index 00000000..29554905 --- /dev/null +++ b/docs/plan/20260619-edge-phase3-tasks.md @@ -0,0 +1,53 @@ +# 边缘运行时 Phase 3 — 任务拆解 + +> **状态**:Batch 1 + Batch 2 已完成(2026-06-19) +> **前置**:[边缘运行时重构设计](../design/edge-runtime-refactor.md) Phase 0–2 已完成 + +--- + +## 任务依赖图 + +```text +Batch 1(并行,互不影响) +├── T1 edge/config/duration.go +├── T2 edge/observability/linux.go +└── T3 edge/heartbeat/loop.go(仅 relay/flared) + +Batch 2(串行,依赖 Batch 1 或需独立评审) +├── T4 Agent heartbeat 架构对齐(runner ↔ heartbeat service) +└── T5 agent/protocol → pkg/protocol(影响 Server 侧) +``` + +--- + +## Batch 1 — 并行任务 + +| ID | 任务 | 修改范围 | 风险 | 委派 | +| --- | --- | --- | --- | --- | +| **T1** | 抽取 `MillisecondDuration` | `edge/config/` + `agent/relay/flared/config` | 低 | ✅ 子代理 A | +| **T2** | 抽取 Linux 指标采集 | `edge/observability/` + `agent/relay/observability/collector.go` | 中 | ✅ 子代理 B | +| **T3** | 统一心跳 ticker 循环 | `edge/heartbeat/loop.go` + `relay/flared/heartbeat` | 低 | ✅ 子代理 C | + +### 隔离规则 + +- **T1** 禁止修改 `heartbeat/`、`observability/`、`agent/runner.go` +- **T2** 禁止修改 `config/`、`heartbeat/` +- **T3** 禁止修改 `agent/` 任何文件(Agent 心跳留在 runner,Batch 2 处理) + +--- + +## Batch 2 — 并行任务(已完成) + +| ID | 任务 | 修改范围 | 状态 | +| --- | --- | --- | --- | +| **T4** | Agent heartbeat 架构对齐 | `heartbeat/cycle.go` + 精简 `agent/runner.go` | ✅ 子代理 D | +| **T5** | `agent/protocol` → `pkg/protocol` | `pkg/protocol/agent.go` + `protocol/alias.go` | ✅ 子代理 E | + +--- + +## 验收标准(Batch 1) + +```bash +go build ./cmd/agent ./cmd/relay ./cmd/flared +go test ./internal/apps/edge/... ./internal/apps/agent/... ./internal/apps/relay/... ./internal/apps/flared/... -count=1 +``` \ No newline at end of file diff --git a/docs/plan/index.md b/docs/plan/index.md index 808cbd6b..74e60cc1 100644 --- a/docs/plan/index.md +++ b/docs/plan/index.md @@ -19,6 +19,7 @@ | [OpenFlare 前端迁移 — AI 委派](./handover-openflare-frontend-migration.md) | 前端迁移任务队列与验收状态 | | [前端路由验证](./verify-frontend-routes.md) · [Service 验证](./verify-frontend-services.md) · [UI 验证](./verify-frontend-ui.md) · [构建验证](./verify-frontend-build.md) | 多角度迁移验收报告 | | [文档结构更新 — AI 接手](./handover-docs-restructure-update.md) | 重构后多智能体分析结论与文档批量更新记录 | +| [边缘运行时 Phase 3 任务拆解](./20260619-edge-phase3-tasks.md) | Batch 1 并行任务(duration / observability / heartbeat loop) | ## 使用建议 diff --git a/internal/apps/agent/agent/runner.go b/internal/apps/agent/agent/runner.go index d27c5a2c..b4a44de8 100644 --- a/internal/apps/agent/agent/runner.go +++ b/internal/apps/agent/agent/runner.go @@ -9,10 +9,11 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/agent/config" - "github.com/Rain-kl/Wavelet/internal/apps/agent/observability" + agentheartbeat "github.com/Rain-kl/Wavelet/internal/apps/agent/heartbeat" "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" "github.com/Rain-kl/Wavelet/internal/apps/agent/state" "github.com/Rain-kl/Wavelet/internal/apps/agent/wsclient" + edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat" ) type HeartbeatService interface { @@ -29,10 +30,6 @@ type SyncService interface { ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error } -type Updater interface { - CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error -} - type RuntimeManager interface { CheckHealth(ctx context.Context) error Restart(ctx context.Context) error @@ -44,32 +41,23 @@ type WebSocketService interface { URL() string } -type UpdateOptions struct { - Channel string - TagName string - Force bool -} - type Runner struct { Config *config.Config StateStore *state.Store - ObservabilityBuffer *state.ObservabilityBufferStore + HeartbeatCycle *agentheartbeat.Cycle HeartbeatService HeartbeatService SyncService SyncService - Updater Updater RuntimeManager RuntimeManager WebSocketService WebSocketService - autoUpdate bool - updateNow bool - updateRepo string - updateChan string - updateTag string restartOpenrestyNow bool websocketUpgradeEnabled bool } func (r *Runner) Run(ctx context.Context) error { + if r.HeartbeatCycle != nil { + r.HeartbeatCycle.RecordSyncError = r.recordSyncError + } nodeID, err := r.StateStore.EnsureNodeID() if err != nil { return err @@ -149,36 +137,15 @@ func (r *Runner) Run(ctx context.Context) error { func (r *Runner) performHeartbeatCycle(ctx context.Context, nodeID string, startup bool) (bool, error) { r.refreshOpenrestyHealth(ctx) - payload, ackWindows := r.prepareHeartbeatPayload(nodeID) - heartbeatResult, err := r.HeartbeatService.Heartbeat(ctx, payload) - if err != nil { - return false, err - } - r.ackObservabilityWindows(ackWindows) - if heartbeatResult == nil { - heartbeatResult = &protocol.HeartbeatResult{} - } - mode := "periodic" - if startup { - mode = "startup" - } - slog.Debug("agent heartbeat succeeded", "mode", mode, "node_id", nodeID) - changed := r.applySettings(heartbeatResult.AgentSettings) - r.applyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups) - if startup { - if err = r.SyncService.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil { - r.recordSyncError(err) - slog.Error("agent startup sync failed", "error", err) - } else { - slog.Debug("agent startup sync completed") - } - } else if err = r.SyncService.SyncOnce(ctx, heartbeatResult.ActiveConfig); err != nil { - r.recordSyncError(err) - slog.Error("agent sync failed", "error", err) - } + return r.HeartbeatCycle.Perform(ctx, nodeID, startup, r) +} + +func (r *Runner) Apply(settings *protocol.AgentSettings) bool { + return r.applySettings(settings) +} + +func (r *Runner) RestartOpenrestyIfNeeded(ctx context.Context) { r.tryRestartOpenresty(ctx) - r.tryAutoUpdate(ctx) - return changed, nil } func (r *Runner) shouldUseWebSocket() bool { @@ -254,7 +221,6 @@ func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol. childCtx, cancel := context.WithCancel(ctx) defer cancel() - // Start status ticker sender in background go func() { for { select { @@ -285,11 +251,11 @@ func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol. func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error { r.refreshOpenrestyHealth(ctx) - payload, ackWindows := r.prepareHeartbeatPayload(nodeID) + payload, ackWindows := r.HeartbeatCycle.PrepareHeartbeatPayload(nodeID) if err := conn.SendStatus(payload); err != nil { return err } - r.ackObservabilityWindows(ackWindows) + r.HeartbeatCycle.AckObservabilityWindows(ackWindows) return nil } @@ -303,7 +269,7 @@ func (r *Runner) handleWebSocketMessage(ctx context.Context, message protocol.WS } changed := r.applySettings(&settings) r.tryRestartOpenresty(ctx) - r.tryAutoUpdate(ctx) + edgeheartbeat.TryAutoUpdate(ctx, r.HeartbeatCycle.Updater, agentheartbeat.AgentSettingsToAutoUpdate(&settings), "agent") if !r.websocketUpgradeEnabled { slog.Debug("agent ws disabled by server settings; falling back to http heartbeat") return changed, errors.New("websocket upgrade disabled by server") @@ -339,7 +305,7 @@ func (r *Runner) handleWebSocketMessage(ctx context.Context, message protocol.WS slog.Debug("agent ws waf ip groups decode failed", "error", err) return false, nil } - r.applyWAFIPGroups(ctx, groups) + r.HeartbeatCycle.ApplyWAFIPGroups(ctx, groups) return false, nil case protocol.WSMessageTypePing: slog.Debug("agent ws ping received") @@ -409,11 +375,6 @@ func (r *Runner) applySettings(settings *protocol.AgentSettings) bool { slog.Debug("agent websocket upgrade setting updated", "from", r.websocketUpgradeEnabled, "to", settings.WebsocketUpgradeEnabled) } r.websocketUpgradeEnabled = settings.WebsocketUpgradeEnabled - r.autoUpdate = settings.AutoUpdate - r.updateNow = settings.UpdateNow - r.updateRepo = strings.TrimSpace(settings.UpdateRepo) - r.updateChan = strings.TrimSpace(settings.UpdateChannel) - r.updateTag = strings.TrimSpace(settings.UpdateTag) r.restartOpenrestyNow = settings.RestartOpenrestyNow return changed } @@ -436,37 +397,12 @@ func (r *Runner) tryRestartOpenresty(ctx context.Context) { r.recordOpenrestyHealthy() } -func (r *Runner) tryAutoUpdate(ctx context.Context) { - force := r.updateNow - shouldCheck := r.autoUpdate || force - r.updateNow = false - r.updateTag = strings.TrimSpace(r.updateTag) - if !shouldCheck || r.Updater == nil || r.updateRepo == "" { - return - } - channel := "stable" - if force && r.updateChan != "" { - channel = r.updateChan - } - if err := r.Updater.CheckAndUpdate(ctx, r.updateRepo, UpdateOptions{ - Channel: channel, - TagName: r.updateTag, - Force: force, - }); err != nil { - slog.Error("agent update check failed", "error", err) - } - if force { - r.updateTag = "" - r.updateChan = "" - } -} - func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error { if strings.TrimSpace(r.Config.DiscoveryToken) == "" { return errors.New("agent_token 为空且未配置 discovery_token") } slog.Info("agent discovery registration started") - response, err := r.HeartbeatService.Register(ctx, r.nodePayload(*nodeID)) + response, err := r.HeartbeatService.Register(ctx, r.HeartbeatCycle.NodePayload(*nodeID)) if err != nil { return err } @@ -493,26 +429,9 @@ func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error { *nodeID = response.NodeID slog.Info("agent discovery registration succeeded", "node_id", response.NodeID) r.refreshOpenrestyHealth(ctx) - payload, ackWindows := r.prepareHeartbeatPayload(*nodeID) - heartbeatResult, heartbeatErr := r.HeartbeatService.Heartbeat(ctx, payload) - if heartbeatErr != nil { - slog.Error("agent post-register heartbeat failed", "error", heartbeatErr) - return nil + if _, err = r.HeartbeatCycle.Perform(ctx, *nodeID, true, r); err != nil { + slog.Error("agent post-register heartbeat failed", "error", err) } - r.ackObservabilityWindows(ackWindows) - if heartbeatResult == nil { - heartbeatResult = &protocol.HeartbeatResult{} - } - r.applySettings(heartbeatResult.AgentSettings) - r.applyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups) - if err = r.SyncService.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil { - r.recordSyncError(err) - slog.Error("agent post-register startup sync failed", "error", err) - } else { - slog.Debug("agent post-register startup sync completed") - } - r.tryRestartOpenresty(ctx) - r.tryAutoUpdate(ctx) return nil } @@ -582,118 +501,4 @@ func (r *Runner) recordOpenrestyUnhealthy(err error, fallbackOnly bool) { if saveErr := r.StateStore.Save(snapshot); saveErr != nil { slog.Error("save state after recording openresty error failed", "error", saveErr) } -} - -func (r *Runner) nodePayload(nodeID string) protocol.NodePayload { - snapshot, _ := r.StateStore.Load() - openrestyStatus := strings.TrimSpace(snapshot.OpenrestyStatus) - if openrestyStatus == "" { - openrestyStatus = protocol.OpenrestyStatusUnknown - } - profile := observability.BuildProfile(r.Config, r.StateStore) - managedOpenRestyMetrics := observability.CollectManagedOpenRestyMetrics(r.Config) - trafficReport, accessLogs, fallbackMetrics := observability.BuildTrafficObservability(r.Config, r.StateStore, managedOpenRestyMetrics) - if managedOpenRestyMetrics == nil { - managedOpenRestyMetrics = fallbackMetrics - } - metricSnapshot := observability.BuildSnapshot(r.Config, r.StateStore) - openrestyObservation := observability.BuildOpenrestyObservation(managedOpenRestyMetrics) - healthEvents := observability.BuildHealthEvents(snapshot) - payload := protocol.NodePayload{ - NodeID: nodeID, - Name: r.Config.NodeName, - IP: r.Config.NodeIP, - Version: r.Config.Version, - ExtVersion: r.Config.ExtVersion, - CurrentVersion: snapshot.CurrentVersion, - LastError: snapshot.LastError, - OpenrestyStatus: openrestyStatus, - OpenrestyMessage: snapshot.OpenrestyMessage, - Profile: profile, - Snapshot: metricSnapshot, - OpenrestyObservation: openrestyObservation, - TrafficReport: trafficReport, - AccessLogs: accessLogs, - HealthEvents: healthEvents, - } - if r.SyncService != nil { - checksums, err := r.SyncService.WAFIPGroupChecksums() - if err != nil { - slog.Debug("load local waf ip group checksums failed", "error", err) - } else if len(checksums) > 0 { - payload.WAFIPGroupChecksums = checksums - } - } - return payload -} - -func (r *Runner) applyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) { - if len(groups) == 0 || r.SyncService == nil { - return - } - if err := r.SyncService.ApplyWAFIPGroups(ctx, groups); err != nil { - r.recordSyncError(err) - slog.Error("agent apply waf ip groups failed", "error", err) - } -} - -func (r *Runner) prepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []int64) { - payload := r.nodePayload(nodeID) - if r.ObservabilityBuffer == nil || (payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0) { - return payload, nil - } - now := time.Now().UTC() - retainAfterUnix := now.Add(-time.Duration(r.Config.ObservabilityReplayMinutes) * time.Minute).Unix() - windowStartedAtUnix := state.ObservabilityWindowStartedAt(payload.Snapshot, payload.OpenrestyObservation, payload.TrafficReport) - if windowStartedAtUnix <= 0 { - return payload, nil - } - - record := state.ObservabilityBufferRecord{ - WindowStartedAtUnix: windowStartedAtUnix, - Snapshot: payload.Snapshot, - OpenrestyObservation: payload.OpenrestyObservation, - TrafficReport: payload.TrafficReport, - AccessLogs: payload.AccessLogs, - QueuedAtUnix: now.Unix(), - } - if err := r.ObservabilityBuffer.Upsert(record, retainAfterUnix); err != nil { - slog.Error("upsert observability buffer failed", "error", err) - return payload, nil - } - - records, err := r.ObservabilityBuffer.Replayable(windowStartedAtUnix, retainAfterUnix) - if err != nil { - slog.Error("load replayable observability buffer failed", "error", err) - return payload, []int64{windowStartedAtUnix} - } - - ackWindows := make([]int64, 0, len(records)+1) - buffered := make([]protocol.BufferedObservabilityRecord, 0, len(records)) - for _, item := range records { - if item.WindowStartedAtUnix <= 0 { - continue - } - buffered = append(buffered, protocol.BufferedObservabilityRecord{ - WindowStartedAtUnix: item.WindowStartedAtUnix, - Snapshot: item.Snapshot, - OpenrestyObservation: item.OpenrestyObservation, - TrafficReport: item.TrafficReport, - AccessLogs: item.AccessLogs, - }) - ackWindows = append(ackWindows, item.WindowStartedAtUnix) - } - payload.BufferedObservability = buffered - ackWindows = append(ackWindows, windowStartedAtUnix) - return payload, ackWindows -} - -func (r *Runner) ackObservabilityWindows(windowStartedAtUnix []int64) { - if r.ObservabilityBuffer == nil || len(windowStartedAtUnix) == 0 { - return - } - retainAfterUnix := time.Now().UTC().Add(-time.Duration(r.Config.ObservabilityReplayMinutes) * time.Minute).Unix() - if err := r.ObservabilityBuffer.Ack(windowStartedAtUnix, retainAfterUnix); err != nil { - slog.Error("ack observability buffer failed", "error", err) - } -} +} \ No newline at end of file diff --git a/internal/apps/agent/agent/runner_test.go b/internal/apps/agent/agent/runner_test.go index 7e4b9267..d81b9110 100644 --- a/internal/apps/agent/agent/runner_test.go +++ b/internal/apps/agent/agent/runner_test.go @@ -11,10 +11,24 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/agent/config" + agentheartbeat "github.com/Rain-kl/Wavelet/internal/apps/agent/heartbeat" "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" "github.com/Rain-kl/Wavelet/internal/apps/agent/state" + "github.com/Rain-kl/Wavelet/internal/apps/agent/updater" ) +func withHeartbeatCycle(runner *Runner, observabilityBuffer *state.ObservabilityBufferStore) *Runner { + runner.HeartbeatCycle = &agentheartbeat.Cycle{ + Config: runner.Config, + StateStore: runner.StateStore, + ObservabilityBuffer: observabilityBuffer, + Heartbeat: runner.HeartbeatService, + Sync: runner.SyncService, + Updater: updater.New(), + } + return runner +} + type fakeHeartbeatService struct { mu sync.Mutex registerCalls int @@ -192,7 +206,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) { syncService := &fakeSyncService{ startupErr: errors.New("当前没有激活版本,保持当前 OpenResty 配置"), } - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ AccessToken: "agent-token", NodeName: "edge-01", @@ -204,7 +218,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) { StateStore: stateStore, HeartbeatService: heartbeatService, SyncService: syncService, - } + }, nil) err := runner.Run(ctx) if !errors.Is(err, context.Canceled) { @@ -245,7 +259,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) { } }, } - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ AccessToken: "agent-token", NodeName: "edge-01", @@ -257,7 +271,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) { StateStore: stateStore, HeartbeatService: heartbeatService, SyncService: syncService, - } + }, nil) err := runner.Run(ctx) if !errors.Is(err, context.Canceled) { @@ -303,7 +317,7 @@ func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) { healthErr: errors.New("docker openresty container is not running"), clearHealthOnRestart: true, } - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ AccessToken: "agent-token", NodeName: "edge-01", @@ -316,7 +330,7 @@ func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) { HeartbeatService: heartbeatService, SyncService: &fakeSyncService{}, RuntimeManager: runtimeManager, - } + }, nil) err := runner.Run(ctx) if !errors.Is(err, context.Canceled) { @@ -357,7 +371,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) { t.Fatalf("failed to seed state: %v", err) } - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ NodeName: "edge-observe-1", NodeIP: "10.0.0.51", @@ -369,7 +383,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) { HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond), }, StateStore: stateStore, - } + }, nil) if err := os.MkdirAll(filepath.Dir(runner.Config.AccessLogPath), 0o755); err != nil { t.Fatalf("failed to prepare access log dir: %v", err) } @@ -381,7 +395,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) { t.Fatalf("failed to prepare access log: %v", err) } - firstPayload := runner.nodePayload("node-observe") + firstPayload := runner.HeartbeatCycle.NodePayload("node-observe") if firstPayload.Profile == nil { t.Fatal("expected first heartbeat payload to include system profile") } @@ -398,7 +412,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) { t.Fatalf("expected health events for openresty and sync error, got %+v", firstPayload.HealthEvents) } - secondPayload := runner.nodePayload("node-observe") + secondPayload := runner.HeartbeatCycle.NodePayload("node-observe") if secondPayload.Profile != nil { t.Fatal("expected unchanged profile to be omitted on subsequent heartbeat") } @@ -439,7 +453,7 @@ func TestRunnerReplaysBufferedObservabilityAfterHeartbeatRecovery(t *testing.T) } }, } - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ AccessToken: "agent-token", NodeName: "edge-buffer-01", @@ -451,11 +465,10 @@ func TestRunnerReplaysBufferedObservabilityAfterHeartbeatRecovery(t *testing.T) HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond), ObservabilityReplayMinutes: 15, }, - StateStore: stateStore, - ObservabilityBuffer: bufferStore, - HeartbeatService: heartbeatService, - SyncService: &fakeSyncService{}, - } + StateStore: stateStore, + HeartbeatService: heartbeatService, + SyncService: &fakeSyncService{}, + }, bufferStore) if err := os.MkdirAll(filepath.Dir(runner.Config.RouteConfigPath), 0o755); err != nil { t.Fatalf("failed to prepare route config dir: %v", err) } @@ -518,7 +531,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) { if err != nil { t.Fatalf("failed to load config: %v", err) } - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ ServerURL: cfg.ServerURL, DiscoveryToken: cfg.DiscoveryToken, @@ -531,7 +544,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) { StateStore: stateStore, HeartbeatService: heartbeatService, SyncService: syncService, - } + }, nil) runner.Config = cfg runner.Config.Version = config.Version runner.Config.ExtVersion = "1.27.1.2" @@ -561,7 +574,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) { func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) { syncService := &fakeSyncService{} - runner := &Runner{SyncService: syncService} + runner := withHeartbeatCycle(&Runner{SyncService: syncService}, nil) payload, err := json.Marshal(protocol.ActiveConfigMeta{ Version: "20260529-001", Checksum: "checksum-ws", @@ -589,12 +602,12 @@ func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) { } func TestRunnerHandlesWebSocketSettingsDisabled(t *testing.T) { - runner := &Runner{ + runner := withHeartbeatCycle(&Runner{ Config: &config.Config{ HeartbeatInterval: config.MillisecondDuration(10 * time.Second), }, websocketUpgradeEnabled: true, - } + }, nil) payload, err := json.Marshal(protocol.AgentSettings{ HeartbeatInterval: 15000, WebsocketUpgradeEnabled: false, diff --git a/internal/apps/agent/config/config.go b/internal/apps/agent/config/config.go index ec093c39..0e2054e5 100644 --- a/internal/apps/agent/config/config.go +++ b/internal/apps/agent/config/config.go @@ -1,11 +1,9 @@ package config import ( - "context" "encoding/json" "errors" "fmt" - "net" "os" pathpkg "path" "path/filepath" @@ -13,8 +11,7 @@ import ( "strings" "time" - "github.com/Rain-kl/Wavelet/pkg/geoip" - "github.com/Rain-kl/Wavelet/pkg/geoip/iputil" + "github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip" "github.com/Rain-kl/Wavelet/pkg/utils" ) @@ -35,11 +32,6 @@ const ( defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb" ) -var ( - lookupOutboundIP = geoip.GetOutboundIP - lookupLocalIP = detectLocalNodeIP -) - type Config struct { ServerURL string `json:"server_url"` AccessToken string `json:"agent_token"` @@ -166,7 +158,7 @@ func applyDefaults(cfg *Config, baseDir string) { cfg.NodeName = detectHostname() } if cfg.NodeIP == "" { - cfg.NodeIP = detectNodeIP() + cfg.NodeIP = nodeip.Detect() } if cfg.MainConfigPath == "" { cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultMainConfigRelativePath) @@ -400,64 +392,4 @@ func detectHostname() string { return strings.TrimSpace(host) } -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 := normalizeIPv4(ipNet.IP) - priority := nodeIPPriority(ipv4) - if priority > bestPriority { - bestIP = ipv4.String() - bestPriority = priority - } - if bestPriority == 2 { - return bestIP - } - } - } - return bestIP -} - -func normalizeIPv4(ip net.IP) net.IP { - if ip == nil { - return nil - } - return ip.To4() -} - -func nodeIPPriority(ip net.IP) int { - return iputil.Score(ip) -} diff --git a/internal/apps/agent/config/config_test.go b/internal/apps/agent/config/config_test.go index 3979c9e4..11649a1a 100644 --- a/internal/apps/agent/config/config_test.go +++ b/internal/apps/agent/config/config_test.go @@ -10,7 +10,9 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip" "github.com/Rain-kl/Wavelet/pkg/geoip" + "github.com/Rain-kl/Wavelet/pkg/geoip/iputil" ) func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) { @@ -248,12 +250,12 @@ func TestLoadUsesEnvConfigWhenFileIsMissing(t *testing.T) { } func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) { - previousLookup := lookupOutboundIP - lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) { + previousLookup := nodeip.LookupOutboundIP + nodeip.LookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) { return net.ParseIP("8.8.8.8"), nil } defer func() { - lookupOutboundIP = previousLookup + nodeip.LookupOutboundIP = previousLookup }() dir := t.TempDir() @@ -281,17 +283,17 @@ func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) { } func TestLoadFallsBackToLocalIPWhenOutboundLookupFails(t *testing.T) { - previousOutboundLookup := lookupOutboundIP - previousLocalLookup := lookupLocalIP - lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) { + previousOutboundLookup := nodeip.LookupOutboundIP + previousLocalLookup := nodeip.LookupLocalIP + nodeip.LookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) { return nil, errors.New("realip.cc unavailable") } - lookupLocalIP = func() string { + nodeip.LookupLocalIP = func() string { return "9.9.9.9" } defer func() { - lookupOutboundIP = previousOutboundLookup - lookupLocalIP = previousLocalLookup + nodeip.LookupOutboundIP = previousOutboundLookup + nodeip.LookupLocalIP = previousLocalLookup }() dir := t.TempDir() @@ -512,7 +514,7 @@ func TestNodeIPPriority(t *testing.T) { if tt.ip != "" { parsed = net.ParseIP(tt.ip) } - if got := nodeIPPriority(parsed); got != tt.expected { + if got := iputil.Score(parsed); got != tt.expected { t.Fatalf("unexpected priority for %q: got %d want %d", tt.ip, got, tt.expected) } }) diff --git a/internal/apps/agent/config/duration.go b/internal/apps/agent/config/duration.go index 17ca3f87..254dcbd8 100644 --- a/internal/apps/agent/config/duration.go +++ b/internal/apps/agent/config/duration.go @@ -1,54 +1,5 @@ package config -import ( - "encoding/json" - "fmt" - "strconv" - "strings" - "time" -) +import edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config" -type MillisecondDuration time.Duration - -func (d MillisecondDuration) Duration() time.Duration { - return time.Duration(d) -} - -func (d MillisecondDuration) String() string { - return time.Duration(d).String() -} - -func (d *MillisecondDuration) UnmarshalJSON(data []byte) error { - raw := strings.TrimSpace(string(data)) - if raw == "" || raw == "null" { - *d = 0 - return nil - } - if strings.HasPrefix(raw, "\"") { - var text string - if err := json.Unmarshal(data, &text); err != nil { - return err - } - text = strings.TrimSpace(text) - if text == "" { - *d = 0 - return nil - } - parsed, err := time.ParseDuration(text) - if err != nil { - return fmt.Errorf("invalid duration string %q: %w", text, err) - } - *d = MillisecondDuration(parsed) - return nil - } - ms, err := strconv.ParseInt(raw, 10, 64) - if err != nil { - return fmt.Errorf("invalid duration milliseconds %q: %w", raw, err) - } - *d = MillisecondDuration(time.Duration(ms) * time.Millisecond) - return nil -} - -func (d MillisecondDuration) MarshalJSON() ([]byte, error) { - return json.Marshal(time.Duration(d).Milliseconds()) -} +type MillisecondDuration = edgeconfig.MillisecondDuration \ No newline at end of file diff --git a/internal/apps/agent/heartbeat/cycle.go b/internal/apps/agent/heartbeat/cycle.go new file mode 100644 index 00000000..39e2a381 --- /dev/null +++ b/internal/apps/agent/heartbeat/cycle.go @@ -0,0 +1,217 @@ +package heartbeat + +import ( + "context" + "log/slog" + "strings" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/agent/config" + "github.com/Rain-kl/Wavelet/internal/apps/agent/observability" + "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" + "github.com/Rain-kl/Wavelet/internal/apps/agent/state" + "github.com/Rain-kl/Wavelet/internal/apps/agent/updater" + edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat" +) + +type HeartbeatClient interface { + Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) +} + +type SyncService interface { + SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error + SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error + WAFIPGroupChecksums() (map[string]string, error) + ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error +} + +type SettingsApplier interface { + Apply(settings *protocol.AgentSettings) (intervalChanged bool) + RestartOpenrestyIfNeeded(ctx context.Context) +} + +type Cycle struct { + Config *config.Config + StateStore *state.Store + ObservabilityBuffer *state.ObservabilityBufferStore + Heartbeat HeartbeatClient + Sync SyncService + Updater *updater.Service + RecordSyncError func(err error) +} + +func (c *Cycle) Perform(ctx context.Context, nodeID string, startup bool, settings SettingsApplier) (bool, error) { + payload, ackWindows := c.PrepareHeartbeatPayload(nodeID) + heartbeatResult, err := c.Heartbeat.Heartbeat(ctx, payload) + if err != nil { + return false, err + } + c.AckObservabilityWindows(ackWindows) + if heartbeatResult == nil { + heartbeatResult = &protocol.HeartbeatResult{} + } + mode := "periodic" + if startup { + mode = "startup" + } + slog.Debug("agent heartbeat succeeded", "mode", mode, "node_id", nodeID) + + var changed bool + if settings != nil { + changed = settings.Apply(heartbeatResult.AgentSettings) + } + c.ApplyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups) + if startup { + if err = c.Sync.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil { + c.recordSyncError(err) + slog.Error("agent startup sync failed", "error", err) + } else { + slog.Debug("agent startup sync completed") + } + } else if err = c.Sync.SyncOnce(ctx, heartbeatResult.ActiveConfig); err != nil { + c.recordSyncError(err) + slog.Error("agent sync failed", "error", err) + } + if settings != nil { + settings.RestartOpenrestyIfNeeded(ctx) + } + edgeheartbeat.TryAutoUpdate(ctx, c.Updater, agentSettingsToAutoUpdate(heartbeatResult.AgentSettings), "agent") + return changed, nil +} + +func (c *Cycle) NodePayload(nodeID string) protocol.NodePayload { + snapshot, _ := c.StateStore.Load() + openrestyStatus := strings.TrimSpace(snapshot.OpenrestyStatus) + if openrestyStatus == "" { + openrestyStatus = protocol.OpenrestyStatusUnknown + } + profile := observability.BuildProfile(c.Config, c.StateStore) + managedOpenRestyMetrics := observability.CollectManagedOpenRestyMetrics(c.Config) + trafficReport, accessLogs, fallbackMetrics := observability.BuildTrafficObservability(c.Config, c.StateStore, managedOpenRestyMetrics) + if managedOpenRestyMetrics == nil { + managedOpenRestyMetrics = fallbackMetrics + } + metricSnapshot := observability.BuildSnapshot(c.Config, c.StateStore) + openrestyObservation := observability.BuildOpenrestyObservation(managedOpenRestyMetrics) + healthEvents := observability.BuildHealthEvents(snapshot) + payload := protocol.NodePayload{ + NodeID: nodeID, + Name: c.Config.NodeName, + IP: c.Config.NodeIP, + Version: c.Config.Version, + ExtVersion: c.Config.ExtVersion, + CurrentVersion: snapshot.CurrentVersion, + LastError: snapshot.LastError, + OpenrestyStatus: openrestyStatus, + OpenrestyMessage: snapshot.OpenrestyMessage, + Profile: profile, + Snapshot: metricSnapshot, + OpenrestyObservation: openrestyObservation, + TrafficReport: trafficReport, + AccessLogs: accessLogs, + HealthEvents: healthEvents, + } + if c.Sync != nil { + checksums, err := c.Sync.WAFIPGroupChecksums() + if err != nil { + slog.Debug("load local waf ip group checksums failed", "error", err) + } else if len(checksums) > 0 { + payload.WAFIPGroupChecksums = checksums + } + } + return payload +} + +func (c *Cycle) PrepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []int64) { + payload := c.NodePayload(nodeID) + if c.ObservabilityBuffer == nil || (payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0) { + return payload, nil + } + now := time.Now().UTC() + retainAfterUnix := now.Add(-time.Duration(c.Config.ObservabilityReplayMinutes) * time.Minute).Unix() + windowStartedAtUnix := state.ObservabilityWindowStartedAt(payload.Snapshot, payload.OpenrestyObservation, payload.TrafficReport) + if windowStartedAtUnix <= 0 { + return payload, nil + } + + record := state.ObservabilityBufferRecord{ + WindowStartedAtUnix: windowStartedAtUnix, + Snapshot: payload.Snapshot, + OpenrestyObservation: payload.OpenrestyObservation, + TrafficReport: payload.TrafficReport, + AccessLogs: payload.AccessLogs, + QueuedAtUnix: now.Unix(), + } + if err := c.ObservabilityBuffer.Upsert(record, retainAfterUnix); err != nil { + slog.Error("upsert observability buffer failed", "error", err) + return payload, nil + } + + records, err := c.ObservabilityBuffer.Replayable(windowStartedAtUnix, retainAfterUnix) + if err != nil { + slog.Error("load replayable observability buffer failed", "error", err) + return payload, []int64{windowStartedAtUnix} + } + + ackWindows := make([]int64, 0, len(records)+1) + buffered := make([]protocol.BufferedObservabilityRecord, 0, len(records)) + for _, item := range records { + if item.WindowStartedAtUnix <= 0 { + continue + } + buffered = append(buffered, protocol.BufferedObservabilityRecord{ + WindowStartedAtUnix: item.WindowStartedAtUnix, + Snapshot: item.Snapshot, + OpenrestyObservation: item.OpenrestyObservation, + TrafficReport: item.TrafficReport, + AccessLogs: item.AccessLogs, + }) + ackWindows = append(ackWindows, item.WindowStartedAtUnix) + } + payload.BufferedObservability = buffered + ackWindows = append(ackWindows, windowStartedAtUnix) + return payload, ackWindows +} + +func (c *Cycle) AckObservabilityWindows(windowStartedAtUnix []int64) { + if c.ObservabilityBuffer == nil || len(windowStartedAtUnix) == 0 { + return + } + retainAfterUnix := time.Now().UTC().Add(-time.Duration(c.Config.ObservabilityReplayMinutes) * time.Minute).Unix() + if err := c.ObservabilityBuffer.Ack(windowStartedAtUnix, retainAfterUnix); err != nil { + slog.Error("ack observability buffer failed", "error", err) + } +} + +func (c *Cycle) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) { + if len(groups) == 0 || c.Sync == nil { + return + } + if err := c.Sync.ApplyWAFIPGroups(ctx, groups); err != nil { + c.recordSyncError(err) + slog.Error("agent apply waf ip groups failed", "error", err) + } +} + +func (c *Cycle) recordSyncError(err error) { + if c.RecordSyncError != nil { + c.RecordSyncError(err) + } +} + +func AgentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.AutoUpdateSettings { + if settings == nil { + return nil + } + return &edgeheartbeat.AutoUpdateSettings{ + AutoUpdate: settings.AutoUpdate, + UpdateNow: settings.UpdateNow, + UpdateRepo: settings.UpdateRepo, + UpdateChannel: settings.UpdateChannel, + UpdateTag: settings.UpdateTag, + } +} + +func agentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.AutoUpdateSettings { + return AgentSettingsToAutoUpdate(settings) +} \ No newline at end of file diff --git a/internal/apps/agent/httpclient/client.go b/internal/apps/agent/httpclient/client.go index 32438d90..f29d7e25 100644 --- a/internal/apps/agent/httpclient/client.go +++ b/internal/apps/agent/httpclient/client.go @@ -1,55 +1,44 @@ package httpclient import ( - "bytes" "context" "encoding/json" - "errors" "fmt" "io" - "log/slog" "net/http" - "strings" "time" + edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient" "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" ) type Client struct { - baseURL string - token string - httpClient *http.Client + base *edgehttp.Client } func New(baseURL string, token string, timeout time.Duration) *Client { return &Client{ - baseURL: strings.TrimRight(baseURL, "/"), - token: token, - httpClient: &http.Client{ - Timeout: timeout, - }, + base: edgehttp.New(baseURL, token, timeout, "X-Agent-Token"), } } func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) { - slog.Debug("http register node request", "node_id", payload.NodeID, "current_version", payload.CurrentVersion) resp := protocol.APIResponse[protocol.RegisterNodeResponse]{} - if err := c.postJSON(ctx, "/api/v1/agent/nodes/register", payload, &resp); err != nil { + if err := c.base.PostJSON(ctx, "/api/v1/agent/nodes/register", payload, &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } - slog.Debug("http register node response", "node_id", resp.Data.NodeID) return &resp.Data, nil } func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) { resp := protocol.APIResponse[protocol.HeartbeatData]{} - if err := c.postJSON(ctx, "/api/v1/agent/nodes/heartbeat", payload, &resp); err != nil { + if err := c.base.PostJSON(ctx, "/api/v1/agent/nodes/heartbeat", payload, &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } return &protocol.HeartbeatResult{ @@ -61,134 +50,46 @@ func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (* func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) { resp := protocol.APIResponse[protocol.ActiveConfigResponse]{} - if err := c.getJSON(ctx, "/api/v1/agent/config-versions/active", &resp); err != nil { + if err := c.base.GetJSON(ctx, "/api/v1/agent/config-versions/active", &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } - slog.Debug("http get active config response", "version", resp.Data.Version, "checksum", resp.Data.Checksum, "support_files", len(resp.Data.SupportFiles)) return &resp.Data, nil } func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error { - slog.Debug("http report apply log request", "node_id", payload.NodeID, "version", payload.Version, "result", payload.Result) resp := protocol.APIResponse[json.RawMessage]{} - if err := c.postJSON(ctx, "/api/v1/agent/apply-logs", payload, &resp); err != nil { + if err := c.base.PostJSON(ctx, "/api/v1/agent/apply-logs", payload, &resp); err != nil { return err } - return apiError(resp.ErrorMsg) + return edgehttp.APIError(resp.ErrorMsg) } func (c *Client) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) { resp := protocol.APIResponse[protocol.WAFIPGroupSyncResponse]{} - if err := c.postJSON(ctx, "/api/v1/agent/waf/ip-groups/sync", payload, &resp); err != nil { + if err := c.base.PostJSON(ctx, "/api/v1/agent/waf/ip-groups/sync", payload, &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } return &resp.Data, nil } func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil) - if err != nil { - return nil, err - } - req.Header.Set("X-Agent-Token", c.token) - res, err := c.httpClient.Do(req) + res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil) if err != nil { return nil, err } defer res.Body.Close() if res.StatusCode != http.StatusOK { - return nil, readHTTPError(res) + return nil, edgehttp.ReadHTTPError(res) } return io.ReadAll(res.Body) } 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) - - 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 readHTTPError(res *http.Response) error { - body, err := io.ReadAll(res.Body) - if err != nil { - return errors.New(res.Status) - } - return readBodyError(body, res.Status) -} - -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) -} + c.base.SetToken(token) +} \ No newline at end of file diff --git a/internal/apps/agent/logging/logger.go b/internal/apps/agent/logging/logger.go index f8a2e762..1e630a51 100644 --- a/internal/apps/agent/logging/logger.go +++ b/internal/apps/agent/logging/logger.go @@ -1,29 +1,7 @@ package logging -import ( - "log/slog" - "os" - "strings" -) +import edgelogging "github.com/Rain-kl/Wavelet/internal/apps/edge/logging" func Setup() { - opts := &slog.HandlerOptions{ - AddSource: true, - Level: parseLevel(os.Getenv("LOG_LEVEL")), - } - handler := slog.NewTextHandler(os.Stdout, opts) - slog.SetDefault(slog.New(handler)) -} - -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 - } -} + edgelogging.Setup(edgelogging.Options{AddSource: true}) +} \ No newline at end of file diff --git a/internal/apps/agent/observability/collector.go b/internal/apps/agent/observability/collector.go index 2f20818b..bf9953e9 100644 --- a/internal/apps/agent/observability/collector.go +++ b/internal/apps/agent/observability/collector.go @@ -1,21 +1,18 @@ package observability import ( - "bufio" "crypto/sha256" "encoding/hex" "encoding/json" "os" - "path/filepath" "runtime" - "strconv" "strings" - "syscall" "time" "github.com/Rain-kl/Wavelet/internal/apps/agent/config" "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" "github.com/Rain-kl/Wavelet/internal/apps/agent/state" + edgeobs "github.com/Rain-kl/Wavelet/internal/apps/edge/observability" ) func BuildProfile(cfg *config.Config, stateStore *state.Store) *protocol.NodeSystemProfile { @@ -47,22 +44,22 @@ func BuildSnapshot(cfg *config.Config, stateStore *state.Store) *protocol.NodeMe CapturedAtUnix: now.Unix(), } - memTotal, memUsed := readMemInfo() + memTotal, memUsed := edgeobs.ReadMemInfo() metric.MemoryTotalBytes = memTotal metric.MemoryUsedBytes = memUsed - storageTotal, storageUsed := statFilesystem(cfg.DataDir) + storageTotal, storageUsed := edgeobs.StatFilesystem(cfg.DataDir) metric.StorageTotalBytes = storageTotal metric.StorageUsedBytes = storageUsed - metric.NetworkRxBytes, metric.NetworkTxBytes = readLinuxNetworkTotals() - metric.DiskReadBytes, metric.DiskWriteBytes = readLinuxDiskTotals() + metric.NetworkRxBytes, metric.NetworkTxBytes = edgeobs.ReadLinuxNetworkTotals() + metric.DiskReadBytes, metric.DiskWriteBytes = edgeobs.ReadLinuxDiskTotals() if stateStore == nil { return metric } - totalCPU, idleCPU := readLinuxCPUStat() + totalCPU, idleCPU := edgeobs.ReadLinuxCPUStat() snapshot, err := stateStore.Load() if err != nil { return metric @@ -121,12 +118,12 @@ func BuildHealthEvents(snapshot *state.Snapshot) []protocol.NodeHealthEvent { func collectProfile(cfg *config.Config) *protocol.NodeSystemProfile { hostname, _ := os.Hostname() - osName, osVersion := readLinuxOSRelease() - kernelVersion := readFirstLine("/proc/sys/kernel/osrelease") - cpuModel := readLinuxCPUModel() - totalMemory, _ := readMemInfo() - totalDisk, _ := statFilesystem(cfg.DataDir) - uptimeSeconds := readLinuxUptimeSeconds() + osName, osVersion := edgeobs.ReadLinuxOSRelease() + kernelVersion := edgeobs.ReadFirstLine("/proc/sys/kernel/osrelease") + cpuModel := edgeobs.ReadLinuxCPUModel() + totalMemory, _ := edgeobs.ReadMemInfo() + totalDisk, _ := edgeobs.StatFilesystem(cfg.DataDir) + uptimeSeconds := edgeobs.ReadLinuxUptimeSeconds() return &protocol.NodeSystemProfile{ Hostname: strings.TrimSpace(hostname), @@ -150,255 +147,4 @@ func fingerprintProfile(profile *protocol.NodeSystemProfile) string { } 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() { - line := strings.TrimSpace(scanner.Text()) - if line == "" || strings.HasPrefix(line, "#") { - continue - } - key, value, ok := strings.Cut(line, "=") - if !ok { - continue - } - values[key] = strings.Trim(value, `"`) - } - if pretty := strings.TrimSpace(values["PRETTY_NAME"]); pretty != "" { - return pretty, strings.TrimSpace(values["VERSION_ID"]) - } - name := strings.TrimSpace(values["NAME"]) - if name == "" { - name = runtime.GOOS - } - return name, strings.TrimSpace(values["VERSION_ID"]) -} - -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 memTotalKB int64 - var memAvailableKB int64 - scanner := bufio.NewScanner(file) - for scanner.Scan() { - line := scanner.Text() - switch { - case strings.HasPrefix(line, "MemTotal:"): - memTotalKB = parseMemInfoValue(line) - case strings.HasPrefix(line, "MemAvailable:"): - memAvailableKB = parseMemInfoValue(line) - } - } - - total := memTotalKB * 1024 - if total == 0 { - return 0, 0 - } - used := total - (memAvailableKB * 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 - } - lines := strings.Split(string(content), "\n") - for _, line := range lines { - if !strings.HasPrefix(line, "cpu ") { - continue - } - fields := strings.Fields(line) - if len(fields) < 5 { - return 0, 0 - } - var total uint64 - for i := 1; i < len(fields); i++ { - value, err := strconv.ParseUint(fields[i], 10, 64) - if err != nil { - return 0, 0 - } - total += value - if i == 4 { - // idle - } - } - 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 int64 - var tx int64 - scanner := bufio.NewScanner(file) - for scanner.Scan() { - line := strings.TrimSpace(scanner.Text()) - if !strings.Contains(line, ":") { - continue - } - name, data, ok := strings.Cut(line, ":") - if !ok { - continue - } - if strings.TrimSpace(name) == "lo" { - continue - } - fields := strings.Fields(data) - if len(fields) < 16 { - continue - } - rxValue, err := strconv.ParseInt(fields[0], 10, 64) - if err == nil { - rx += rxValue - } - txValue, err := strconv.ParseInt(fields[8], 10, 64) - if err == nil { - tx += txValue - } - } - return rx, tx -} - -func readLinuxDiskTotals() (int64, int64) { - file, err := os.Open("/proc/diskstats") - if err != nil { - return 0, 0 - } - defer file.Close() - - var readBytes int64 - var writeBytes int64 - scanner := bufio.NewScanner(file) - for scanner.Scan() { - fields := strings.Fields(scanner.Text()) - if len(fields) < 14 { - continue - } - device := fields[2] - if shouldSkipDiskDevice(device) { - continue - } - readSectors, err := strconv.ParseInt(fields[5], 10, 64) - if err == nil { - readBytes += readSectors * 512 - } - writeSectors, err := strconv.ParseInt(fields[9], 10, 64) - if err == nil { - writeBytes += writeSectors * 512 - } - } - return readBytes, writeBytes -} - -func shouldSkipDiskDevice(device string) bool { - switch { - case device == "": - return true - case strings.HasPrefix(device, "loop"), - strings.HasPrefix(device, "ram"), - strings.HasPrefix(device, "dm-"): - return true - default: - return false - } -} - -func statFilesystem(path string) (int64, int64) { - if strings.TrimSpace(path) == "" { - path = string(os.PathSeparator) - } - absPath := filepath.Clean(path) - var stat syscall.Statfs_t - if err := syscall.Statfs(absPath, &stat); err != nil { - return 0, 0 - } - total := int64(stat.Blocks) * int64(stat.Bsize) - free := int64(stat.Bavail) * int64(stat.Bsize) - used := total - free - 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)) -} +} \ No newline at end of file diff --git a/internal/apps/agent/protocol/alias.go b/internal/apps/agent/protocol/alias.go new file mode 100644 index 00000000..9d21e26c --- /dev/null +++ b/internal/apps/agent/protocol/alias.go @@ -0,0 +1,43 @@ +package protocol + +import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol" + +type APIResponse[T any] = pkgprotocol.APIResponse[T] +type HeartbeatData = pkgprotocol.HeartbeatData +type HeartbeatResult = pkgprotocol.HeartbeatResult +type AgentSettings = pkgprotocol.AgentSettings +type WSMessage = pkgprotocol.WSMessage +type WSOutboundMessage = pkgprotocol.WSOutboundMessage +type WebSocketConnection = pkgprotocol.WebSocketConnection +type NodePayload = pkgprotocol.NodePayload +type NodeSystemProfile = pkgprotocol.NodeSystemProfile +type NodeMetricSnapshot = pkgprotocol.NodeMetricSnapshot +type NodeOpenrestyObservation = pkgprotocol.NodeOpenrestyObservation +type NodeTrafficReport = pkgprotocol.NodeTrafficReport +type NodeAccessLog = pkgprotocol.NodeAccessLog +type BufferedObservabilityRecord = pkgprotocol.BufferedObservabilityRecord +type NodeHealthEvent = pkgprotocol.NodeHealthEvent +type RegisterNodeResponse = pkgprotocol.RegisterNodeResponse +type ApplyLogPayload = pkgprotocol.ApplyLogPayload +type ActiveConfigResponse = pkgprotocol.ActiveConfigResponse +type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta +type WAFIPGroup = pkgprotocol.WAFIPGroup +type WAFIPGroupSyncRequest = pkgprotocol.WAFIPGroupSyncRequest +type WAFIPGroupSyncResponse = pkgprotocol.WAFIPGroupSyncResponse +type SupportFile = pkgprotocol.SupportFile + +const ( + WSMessageTypeStatus = pkgprotocol.WSMessageTypeStatus + WSMessageTypeSettings = pkgprotocol.WSMessageTypeSettings + WSMessageTypeActiveConfig = pkgprotocol.WSMessageTypeActiveConfig + WSMessageTypeForceSyncConfig = pkgprotocol.WSMessageTypeForceSyncConfig + WSMessageTypeWAFIPGroups = pkgprotocol.WSMessageTypeWAFIPGroups + WSMessageTypePing = pkgprotocol.WSMessageTypePing + WSMessageTypePong = pkgprotocol.WSMessageTypePong +) + +const ( + OpenrestyStatusHealthy = pkgprotocol.OpenrestyStatusHealthy + OpenrestyStatusUnhealthy = pkgprotocol.OpenrestyStatusUnhealthy + OpenrestyStatusUnknown = pkgprotocol.OpenrestyStatusUnknown +) \ No newline at end of file diff --git a/internal/apps/agent/updater/restart_unix.go b/internal/apps/agent/updater/restart_unix.go deleted file mode 100644 index 173af72b..00000000 --- a/internal/apps/agent/updater/restart_unix.go +++ /dev/null @@ -1,51 +0,0 @@ -//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 -} diff --git a/internal/apps/agent/updater/restart_windows.go b/internal/apps/agent/updater/restart_windows.go deleted file mode 100644 index 66cb5e61..00000000 --- a/internal/apps/agent/updater/restart_windows.go +++ /dev/null @@ -1,53 +0,0 @@ -//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, `"`, `""`) + `"` -} diff --git a/internal/apps/agent/updater/updater.go b/internal/apps/agent/updater/updater.go index 0a1de102..30cb8796 100644 --- a/internal/apps/agent/updater/updater.go +++ b/internal/apps/agent/updater/updater.go @@ -1,366 +1,17 @@ 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/agent/agent" + edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater" "github.com/Rain-kl/Wavelet/internal/apps/agent/config" ) -const maxChecksumAssetSize = 64 * 1024 - -var replaceAndRestartFunc = replaceAndRestart - -type Service struct { - httpClient *http.Client - lastCheckKey string -} +type Service = edgeupdater.Service +type UpdateOptions = edgeupdater.UpdateOptions 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"` -} - -func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options agent.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("agent 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 agent.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("agent binary updated, restarting") - return replaceAndRestartFunc(targetPath, tmpPath) -} - -func assetNameForGOOSGOARCH(goos string, goarch string) string { - name := fmt.Sprintf("openflare-agent-%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 agent.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) -} + return edgeupdater.New(edgeupdater.Config{ + LocalVersion: config.Version, + AssetPrefix: "openflare-agent", + LogLabel: "agent", + }) +} \ No newline at end of file diff --git a/internal/apps/edge/config/duration.go b/internal/apps/edge/config/duration.go new file mode 100644 index 00000000..abed860b --- /dev/null +++ b/internal/apps/edge/config/duration.go @@ -0,0 +1,54 @@ +package config + +import ( + "encoding/json" + "fmt" + "strconv" + "strings" + "time" +) + +type MillisecondDuration time.Duration + +func (d MillisecondDuration) Duration() time.Duration { + return time.Duration(d) +} + +func (d MillisecondDuration) String() string { + return time.Duration(d).String() +} + +func (d *MillisecondDuration) UnmarshalJSON(data []byte) error { + raw := strings.TrimSpace(string(data)) + if raw == "" || raw == "null" { + *d = 0 + return nil + } + if strings.HasPrefix(raw, "\"") { + var text string + if err := json.Unmarshal(data, &text); err != nil { + return err + } + text = strings.TrimSpace(text) + if text == "" { + *d = 0 + return nil + } + parsed, err := time.ParseDuration(text) + if err != nil { + return fmt.Errorf("invalid duration string %q: %w", text, err) + } + *d = MillisecondDuration(parsed) + return nil + } + ms, err := strconv.ParseInt(raw, 10, 64) + if err != nil { + return fmt.Errorf("invalid duration milliseconds %q: %w", raw, err) + } + *d = MillisecondDuration(time.Duration(ms) * time.Millisecond) + return nil +} + +func (d MillisecondDuration) MarshalJSON() ([]byte, error) { + return json.Marshal(time.Duration(d).Milliseconds()) +} \ No newline at end of file diff --git a/internal/apps/edge/config/duration_test.go b/internal/apps/edge/config/duration_test.go new file mode 100644 index 00000000..52d59a70 --- /dev/null +++ b/internal/apps/edge/config/duration_test.go @@ -0,0 +1,53 @@ +package config + +import ( + "encoding/json" + "testing" + "time" +) + +func TestMillisecondDurationUnmarshalJSON(t *testing.T) { + tests := []struct { + name string + input string + want time.Duration + wantErr bool + }{ + {name: "null", input: "null", want: 0}, + {name: "empty", input: `""`, want: 0}, + {name: "integer milliseconds", input: "30000", want: 30 * time.Second}, + {name: "duration string", input: `"5s"`, want: 5 * time.Second}, + {name: "invalid number", input: "not-a-number", wantErr: true}, + {name: "invalid duration string", input: `"not-a-duration"`, wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var d MillisecondDuration + err := json.Unmarshal([]byte(tt.input), &d) + if tt.wantErr { + if err == nil { + t.Fatal("expected error") + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if d.Duration() != tt.want { + t.Fatalf("got %s, want %s", d, tt.want) + } + }) + } +} + +func TestMillisecondDurationMarshalJSON(t *testing.T) { + d := MillisecondDuration(7 * time.Second) + data, err := json.Marshal(d) + if err != nil { + t.Fatalf("Marshal failed: %v", err) + } + if string(data) != "7000" { + t.Fatalf("unexpected marshaled value: %s", string(data)) + } +} \ No newline at end of file diff --git a/internal/apps/edge/heartbeat/autoupdate.go b/internal/apps/edge/heartbeat/autoupdate.go new file mode 100644 index 00000000..ec6f8e19 --- /dev/null +++ b/internal/apps/edge/heartbeat/autoupdate.go @@ -0,0 +1,41 @@ +package heartbeat + +import ( + "context" + "log/slog" + "strings" + + edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater" +) + +type AutoUpdateSettings struct { + AutoUpdate bool + UpdateNow bool + UpdateRepo string + UpdateChannel string + UpdateTag string +} + +func TryAutoUpdate(ctx context.Context, updater *edgeupdater.Service, settings *AutoUpdateSettings, logLabel string) { + if settings == nil || updater == nil { + return + } + force := settings.UpdateNow + shouldCheck := settings.AutoUpdate || force + if !shouldCheck || strings.TrimSpace(settings.UpdateRepo) == "" { + return + } + channel := "stable" + if force && strings.TrimSpace(settings.UpdateChannel) != "" { + channel = settings.UpdateChannel + } + slog.Info("checking for "+logLabel+" updates", "repo", settings.UpdateRepo, "channel", channel, "force", force) + err := updater.CheckAndUpdate(ctx, settings.UpdateRepo, edgeupdater.UpdateOptions{ + Channel: channel, + TagName: settings.UpdateTag, + Force: force, + }) + if err != nil { + slog.Error(logLabel+" update check failed", "error", err) + } +} \ No newline at end of file diff --git a/internal/apps/edge/heartbeat/loop.go b/internal/apps/edge/heartbeat/loop.go new file mode 100644 index 00000000..ee9ce506 --- /dev/null +++ b/internal/apps/edge/heartbeat/loop.go @@ -0,0 +1,23 @@ +package heartbeat + +import ( + "context" + "time" +) + +// RunLoop invokes fn immediately, then on each interval tick until ctx is cancelled. +func RunLoop(ctx context.Context, interval time.Duration, fn func(context.Context)) { + fn(ctx) + + ticker := time.NewTicker(interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + fn(ctx) + } + } +} \ No newline at end of file diff --git a/internal/apps/edge/heartbeat/loop_test.go b/internal/apps/edge/heartbeat/loop_test.go new file mode 100644 index 00000000..c347c432 --- /dev/null +++ b/internal/apps/edge/heartbeat/loop_test.go @@ -0,0 +1,41 @@ +package heartbeat + +import ( + "context" + "sync/atomic" + "testing" + "time" +) + +func TestRunLoopImmediateAndTicker(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var calls atomic.Int32 + interval := 20 * time.Millisecond + + done := make(chan struct{}) + go func() { + RunLoop(ctx, interval, func(context.Context) { + calls.Add(1) + }) + close(done) + }() + + time.Sleep(5 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Fatalf("expected 1 immediate call, got %d", got) + } + + time.Sleep(35 * time.Millisecond) + if got := calls.Load(); got < 2 { + t.Fatalf("expected at least 2 calls after ticker, got %d", got) + } + + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("RunLoop did not exit after context cancellation") + } +} \ No newline at end of file diff --git a/internal/apps/edge/httpclient/client.go b/internal/apps/edge/httpclient/client.go new file mode 100644 index 00000000..11d17c5b --- /dev/null +++ b/internal/apps/edge/httpclient/client.go @@ -0,0 +1,132 @@ +package httpclient + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "log/slog" + "net/http" + "strings" + "time" +) + +type Client struct { + baseURL string + token string + authHeader string + httpClient *http.Client +} + +func New(baseURL, token string, timeout time.Duration, authHeader string) *Client { + return &Client{ + baseURL: strings.TrimRight(baseURL, "/"), + token: token, + authHeader: authHeader, + httpClient: &http.Client{Timeout: timeout}, + } +} + +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 + } + c.setAuthHeader(req) + 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") + c.setAuthHeader(req) + return c.do(req, target) +} + +func (c *Client) DoRaw(ctx context.Context, method, path string, headers map[string]string) (*http.Response, error) { + req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, nil) + if err != nil { + return nil, err + } + c.setAuthHeader(req) + for key, value := range headers { + req.Header.Set(key, value) + } + return c.httpClient.Do(req) +} + +func (c *Client) setAuthHeader(req *http.Request) { + if c.authHeader != "" { + req.Header.Set(c.authHeader, c.token) + } +} + +func (c *Client) do(req *http.Request, target any) error { + res, err := c.httpClient.Do(req) + if err != nil { + slog.Error("http request failed", "method", req.Method, "path", req.URL.Path, "error", err) + return err + } + defer func(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) +} + +func ReadHTTPError(res *http.Response) error { + body, err := io.ReadAll(res.Body) + if err != nil { + return errors.New(res.Status) + } + return ReadBodyError(body, res.Status) +} \ No newline at end of file diff --git a/internal/apps/edge/logging/logging.go b/internal/apps/edge/logging/logging.go new file mode 100644 index 00000000..0663a86a --- /dev/null +++ b/internal/apps/edge/logging/logging.go @@ -0,0 +1,33 @@ +package logging + +import ( + "log/slog" + "os" + "strings" +) + +type Options struct { + AddSource bool +} + +func Setup(opts Options) { + handlerOpts := &slog.HandlerOptions{ + AddSource: opts.AddSource, + Level: ParseLevel(os.Getenv("LOG_LEVEL")), + } + handler := slog.NewTextHandler(os.Stdout, handlerOpts) + slog.SetDefault(slog.New(handler)) +} + +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 + } +} \ No newline at end of file diff --git a/internal/apps/edge/nodeip/nodeip.go b/internal/apps/edge/nodeip/nodeip.go new file mode 100644 index 00000000..f1d57338 --- /dev/null +++ b/internal/apps/edge/nodeip/nodeip.go @@ -0,0 +1,69 @@ +package nodeip + +import ( + "context" + "net" + "time" + + "github.com/Rain-kl/Wavelet/pkg/geoip" + "github.com/Rain-kl/Wavelet/pkg/geoip/iputil" +) + +var ( + LookupOutboundIP = geoip.GetOutboundIP + LookupLocalIP = DetectLocal +) + +func Detect() string { + if ip := detectOutbound(); ip != "" { + return ip + } + return LookupLocalIP() +} + +func detectOutbound() 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 DetectLocal() string { + interfaces, err := net.Interfaces() + if err != nil { + return "" + } + bestIP := "" + bestPriority := -1 + for _, iface := range interfaces { + if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 { + continue + } + addrs, err := iface.Addrs() + if err != nil { + continue + } + for _, addr := range addrs { + ipNet, ok := addr.(*net.IPNet) + if !ok || ipNet.IP == nil || ipNet.IP.IsLoopback() { + continue + } + ipv4 := ipNet.IP.To4() + if ipv4 == nil { + continue + } + priority := iputil.Score(ipv4) + if priority > bestPriority { + bestIP = ipv4.String() + bestPriority = priority + } + if bestPriority == 2 { + return bestIP + } + } + } + return bestIP +} \ No newline at end of file diff --git a/internal/apps/edge/observability/linux.go b/internal/apps/edge/observability/linux.go new file mode 100644 index 00000000..6d2b1bc3 --- /dev/null +++ b/internal/apps/edge/observability/linux.go @@ -0,0 +1,268 @@ +package observability + +import ( + "bufio" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "syscall" +) + +// ReadLinuxOSRelease returns the OS name and version from /etc/os-release. +func ReadLinuxOSRelease() (string, string) { + file, err := os.Open("/etc/os-release") + if err != nil { + return runtime.GOOS, "" + } + defer file.Close() + + values := make(map[string]string) + scanner := bufio.NewScanner(file) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + key, value, ok := strings.Cut(line, "=") + if !ok { + continue + } + values[key] = strings.Trim(value, `"`) + } + if pretty := strings.TrimSpace(values["PRETTY_NAME"]); pretty != "" { + return pretty, strings.TrimSpace(values["VERSION_ID"]) + } + name := strings.TrimSpace(values["NAME"]) + if name == "" { + name = runtime.GOOS + } + return name, strings.TrimSpace(values["VERSION_ID"]) +} + +// ReadLinuxCPUModel returns the CPU model name from /proc/cpuinfo. +func ReadLinuxCPUModel() string { + file, err := os.Open("/proc/cpuinfo") + if err != nil { + return "" + } + defer file.Close() + + scanner := bufio.NewScanner(file) + for scanner.Scan() { + line := scanner.Text() + if strings.HasPrefix(strings.ToLower(line), "model name") { + _, value, ok := strings.Cut(line, ":") + if ok { + return strings.TrimSpace(value) + } + } + } + return "" +} + +// ReadMemInfo returns total and used memory bytes from /proc/meminfo. +func ReadMemInfo() (int64, int64) { + file, err := os.Open("/proc/meminfo") + if err != nil { + return 0, 0 + } + defer file.Close() + + var memTotalKB int64 + var memAvailableKB int64 + scanner := bufio.NewScanner(file) + for scanner.Scan() { + line := scanner.Text() + switch { + case strings.HasPrefix(line, "MemTotal:"): + memTotalKB = parseMemInfoValue(line) + case strings.HasPrefix(line, "MemAvailable:"): + memAvailableKB = parseMemInfoValue(line) + } + } + + total := memTotalKB * 1024 + if total == 0 { + return 0, 0 + } + used := total - (memAvailableKB * 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 +} + +// ReadLinuxUptimeSeconds returns system uptime in seconds from /proc/uptime. +func ReadLinuxUptimeSeconds() int64 { + content, err := os.ReadFile("/proc/uptime") + if err != nil { + return 0 + } + fields := strings.Fields(string(content)) + if len(fields) == 0 { + return 0 + } + value, err := strconv.ParseFloat(fields[0], 64) + if err != nil { + return 0 + } + return int64(value) +} + +// ReadLinuxCPUStat returns aggregate CPU jiffies and idle jiffies from /proc/stat. +func ReadLinuxCPUStat() (uint64, uint64) { + content, err := os.ReadFile("/proc/stat") + if err != nil { + return 0, 0 + } + lines := strings.Split(string(content), "\n") + for _, line := range lines { + if !strings.HasPrefix(line, "cpu ") { + continue + } + fields := strings.Fields(line) + if len(fields) < 5 { + return 0, 0 + } + var total uint64 + for i := 1; i < len(fields); i++ { + value, err := strconv.ParseUint(fields[i], 10, 64) + if err != nil { + return 0, 0 + } + total += value + } + idle, err := strconv.ParseUint(fields[4], 10, 64) + if err != nil { + return 0, 0 + } + return total, idle + } + return 0, 0 +} + +// ReadLinuxNetworkTotals returns aggregate RX and TX byte totals from /proc/net/dev. +func ReadLinuxNetworkTotals() (int64, int64) { + file, err := os.Open("/proc/net/dev") + if err != nil { + return 0, 0 + } + defer file.Close() + + var rx int64 + var tx int64 + scanner := bufio.NewScanner(file) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if !strings.Contains(line, ":") { + continue + } + name, data, ok := strings.Cut(line, ":") + if !ok { + continue + } + if strings.TrimSpace(name) == "lo" { + continue + } + fields := strings.Fields(data) + if len(fields) < 16 { + continue + } + rxValue, err := strconv.ParseInt(fields[0], 10, 64) + if err == nil { + rx += rxValue + } + txValue, err := strconv.ParseInt(fields[8], 10, 64) + if err == nil { + tx += txValue + } + } + return rx, tx +} + +// ReadLinuxDiskTotals returns aggregate disk read and write byte totals from /proc/diskstats. +func ReadLinuxDiskTotals() (int64, int64) { + file, err := os.Open("/proc/diskstats") + if err != nil { + return 0, 0 + } + defer file.Close() + + var readBytes int64 + var writeBytes int64 + scanner := bufio.NewScanner(file) + for scanner.Scan() { + fields := strings.Fields(scanner.Text()) + if len(fields) < 14 { + continue + } + device := fields[2] + if shouldSkipDiskDevice(device) { + continue + } + readSectors, err := strconv.ParseInt(fields[5], 10, 64) + if err == nil { + readBytes += readSectors * 512 + } + writeSectors, err := strconv.ParseInt(fields[9], 10, 64) + if err == nil { + writeBytes += writeSectors * 512 + } + } + return readBytes, writeBytes +} + +func shouldSkipDiskDevice(device string) bool { + switch { + case device == "": + return true + case strings.HasPrefix(device, "loop"), + strings.HasPrefix(device, "ram"), + strings.HasPrefix(device, "dm-"): + return true + default: + return false + } +} + +// StatFilesystem returns total and used bytes for the filesystem containing path. +func StatFilesystem(path string) (int64, int64) { + if strings.TrimSpace(path) == "" { + path = string(os.PathSeparator) + } + absPath := filepath.Clean(path) + var stat syscall.Statfs_t + if err := syscall.Statfs(absPath, &stat); err != nil { + return 0, 0 + } + total := int64(stat.Blocks) * int64(stat.Bsize) + free := int64(stat.Bavail) * int64(stat.Bsize) + used := total - free + if used < 0 { + used = 0 + } + return total, used +} + +// ReadFirstLine reads and returns the trimmed first line of a file. +func ReadFirstLine(path string) string { + content, err := os.ReadFile(path) + if err != nil { + return "" + } + return strings.TrimSpace(string(content)) +} \ No newline at end of file diff --git a/internal/apps/edge/observability/linux_test.go b/internal/apps/edge/observability/linux_test.go new file mode 100644 index 00000000..12a11562 --- /dev/null +++ b/internal/apps/edge/observability/linux_test.go @@ -0,0 +1,78 @@ +package observability + +import ( + "os" + "path/filepath" + "testing" +) + +func TestParseMemInfoValue(t *testing.T) { + t.Parallel() + + tests := []struct { + line string + want int64 + }{ + {line: "MemTotal: 16384000 kB", want: 16384000}, + {line: "MemAvailable: 8192000 kB", want: 8192000}, + {line: "invalid", want: 0}, + {line: "MemTotal: not-a-number kB", want: 0}, + } + + for _, tt := range tests { + if got := parseMemInfoValue(tt.line); got != tt.want { + t.Fatalf("parseMemInfoValue(%q) = %d, want %d", tt.line, got, tt.want) + } + } +} + +func TestShouldSkipDiskDevice(t *testing.T) { + t.Parallel() + + tests := []struct { + device string + want bool + }{ + {device: "", want: true}, + {device: "loop0", want: true}, + {device: "ram0", want: true}, + {device: "dm-0", want: true}, + {device: "sda", want: false}, + {device: "nvme0n1", want: false}, + } + + for _, tt := range tests { + if got := shouldSkipDiskDevice(tt.device); got != tt.want { + t.Fatalf("shouldSkipDiskDevice(%q) = %v, want %v", tt.device, got, tt.want) + } + } +} + +func TestReadFirstLine(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + path := filepath.Join(dir, "sample.txt") + if err := os.WriteFile(path, []byte(" first line\nsecond line\n"), 0o644); err != nil { + t.Fatalf("write file: %v", err) + } + + if got := ReadFirstLine(path); got != "first line\nsecond line" { + t.Fatalf("ReadFirstLine() = %q, want trimmed first line content", got) + } + if got := ReadFirstLine(filepath.Join(dir, "missing.txt")); got != "" { + t.Fatalf("ReadFirstLine(missing) = %q, want empty string", got) + } +} + +func TestStatFilesystem(t *testing.T) { + t.Parallel() + + total, used := StatFilesystem(t.TempDir()) + if total <= 0 { + t.Fatalf("StatFilesystem() total = %d, want > 0", total) + } + if used < 0 || used > total { + t.Fatalf("StatFilesystem() used = %d, total = %d", used, total) + } +} \ No newline at end of file diff --git a/internal/apps/edge/runner/wsreconnect.go b/internal/apps/edge/runner/wsreconnect.go new file mode 100644 index 00000000..579bcefe --- /dev/null +++ b/internal/apps/edge/runner/wsreconnect.go @@ -0,0 +1,64 @@ +package runner + +import ( + "context" + "log/slog" + "time" +) + +type WSConnection interface { + Close() error +} + +type WSReconnectConfig struct { + ComponentName string + ConnectBackoff time.Duration + ReconnectDelay time.Duration + OnShutdown func() +} + +func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig, + connect func(context.Context) (WSConnection, error), + handle func(context.Context, WSConnection), +) error { + if cfg.ConnectBackoff <= 0 { + cfg.ConnectBackoff = 5 * time.Second + } + if cfg.ReconnectDelay <= 0 { + cfg.ReconnectDelay = 2 * time.Second + } + label := cfg.ComponentName + if label == "" { + label = "edge" + } + + for { + select { + case <-ctx.Done(): + if cfg.OnShutdown != nil { + cfg.OnShutdown() + } + return ctx.Err() + default: + } + + conn, err := connect(ctx) + if err != nil { + slog.Error(label+" ws connect failed, will retry", "error", err) + SleepContext(ctx, cfg.ConnectBackoff) + continue + } + + handle(ctx, conn) + _ = conn.Close() + slog.Info(label + " ws connection closed, reconnecting...") + SleepContext(ctx, cfg.ReconnectDelay) + } +} + +func SleepContext(ctx context.Context, d time.Duration) { + select { + case <-ctx.Done(): + case <-time.After(d): + } +} \ No newline at end of file diff --git a/internal/apps/flared/updater/restart_unix.go b/internal/apps/edge/updater/restart_unix.go similarity index 99% rename from internal/apps/flared/updater/restart_unix.go rename to internal/apps/edge/updater/restart_unix.go index 173af72b..2f6dddd9 100644 --- a/internal/apps/flared/updater/restart_unix.go +++ b/internal/apps/edge/updater/restart_unix.go @@ -48,4 +48,4 @@ func removeBackupBinary(path string) error { return err } return nil -} +} \ No newline at end of file diff --git a/internal/apps/agent/updater/restart_unix_test.go b/internal/apps/edge/updater/restart_unix_test.go similarity index 99% rename from internal/apps/agent/updater/restart_unix_test.go rename to internal/apps/edge/updater/restart_unix_test.go index a984346a..55dd002a 100644 --- a/internal/apps/agent/updater/restart_unix_test.go +++ b/internal/apps/edge/updater/restart_unix_test.go @@ -12,4 +12,4 @@ func TestRemoveBackupBinaryIgnoresMissingFile(t *testing.T) { if err := removeBackupBinary(backupPath); err != nil { t.Fatalf("expected missing backup cleanup to be ignored: %v", err) } -} +} \ No newline at end of file diff --git a/internal/apps/relay/updater/restart_windows.go b/internal/apps/edge/updater/restart_windows.go similarity index 99% rename from internal/apps/relay/updater/restart_windows.go rename to internal/apps/edge/updater/restart_windows.go index 66cb5e61..3ecc4e6f 100644 --- a/internal/apps/relay/updater/restart_windows.go +++ b/internal/apps/edge/updater/restart_windows.go @@ -50,4 +50,4 @@ func buildWindowsCommandLine(execPath string, args []string) string { func quoteWindowsArg(value string) string { return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` -} +} \ No newline at end of file diff --git a/internal/apps/edge/updater/service.go b/internal/apps/edge/updater/service.go new file mode 100644 index 00000000..0d0f5021 --- /dev/null +++ b/internal/apps/edge/updater/service.go @@ -0,0 +1,381 @@ +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" +) + +const maxChecksumAssetSize = 64 * 1024 + +var replaceAndRestartFunc = replaceAndRestart + +type Config struct { + LocalVersion string + AssetPrefix string + LogLabel string +} + +type Service struct { + httpClient *http.Client + lastCheckKey string + localVersion string + assetPrefix string + logLabel string +} + +func New(cfg Config) *Service { + return &Service{ + httpClient: &http.Client{Timeout: 30 * time.Second}, + localVersion: cfg.LocalVersion, + assetPrefix: cfg.AssetPrefix, + logLabel: cfg.LogLabel, + } +} + +type UpdateOptions struct { + Channel string + TagName string + Force bool +} + +type githubRelease struct { + TagName string `json:"tag_name"` + Prerelease bool `json:"prerelease"` + Draft bool `json:"draft"` + Assets []githubAsset `json:"assets"` +} + +type githubAsset struct { + Name string `json:"name"` + BrowserDownloadURL string `json:"browser_download_url"` +} + +func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error { + release, err := s.getRelease(ctx, repo, options) + if err != nil { + return fmt.Errorf("check latest release: %w", err) + } + if release == nil || release.TagName == "" { + return nil + } + + remoteVersion := normalizeVersion(release.TagName) + localVersion := normalizeVersion(s.localVersion) + checkKey := buildReleaseCheckKey(options, remoteVersion) + + if remoteVersion == localVersion { + return nil + } + if !options.Force && checkKey != "" && checkKey == s.lastCheckKey { + return nil + } + if !isNewer(localVersion, remoteVersion) { + s.lastCheckKey = checkKey + return nil + } + + slog.Info(s.logLabel+" update available", "from", localVersion, "to", remoteVersion) + assetName := s.assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH) + 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(s.logLabel + " binary updated, restarting") + return replaceAndRestartFunc(targetPath, tmpPath) +} + +func (s *Service) assetNameForGOOSGOARCH(goos string, goarch string) string { + name := fmt.Sprintf("%s-%s-%s", s.assetPrefix, goos, goarch) + if goos == "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) +} \ No newline at end of file diff --git a/internal/apps/agent/updater/updater_test.go b/internal/apps/edge/updater/service_test.go similarity index 68% rename from internal/apps/agent/updater/updater_test.go rename to internal/apps/edge/updater/service_test.go index f7856412..35a11a42 100644 --- a/internal/apps/agent/updater/updater_test.go +++ b/internal/apps/edge/updater/service_test.go @@ -11,9 +11,6 @@ import ( "runtime" "strings" "testing" - - "github.com/Rain-kl/Wavelet/internal/apps/agent/agent" - "github.com/Rain-kl/Wavelet/internal/apps/agent/config" ) type roundTripFunc func(req *http.Request) (*http.Response, error) @@ -22,26 +19,33 @@ func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { return f(req) } +func testService(httpClient *http.Client) *Service { + return &Service{ + httpClient: httpClient, + localVersion: "v1.0.0", + assetPrefix: "openflare-agent", + logLabel: "agent", + } +} + func TestGetLatestPreviewRelease(t *testing.T) { - service := &Service{ - httpClient: &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases?per_page=20" { - t.Fatalf("unexpected request url: %s", req.URL.String()) - } - return &http.Response{ - StatusCode: http.StatusOK, - Header: make(http.Header), - Body: io.NopCloser(strings.NewReader(`[ + service := testService(&http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases?per_page=20" { + t.Fatalf("unexpected request url: %s", req.URL.String()) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`[ {"tag_name":"v1.0.0","prerelease":false}, {"tag_name":"v1.1.0-rc.1","prerelease":true} ]`)), - }, nil - }), - }, - } + }, nil + }), + }) - release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{Channel: "preview"}) + release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{Channel: "preview"}) if err != nil { t.Fatalf("expected preview release query to succeed: %v", err) } @@ -51,22 +55,20 @@ func TestGetLatestPreviewRelease(t *testing.T) { } func TestGetReleaseByTag(t *testing.T) { - service := &Service{ - httpClient: &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v1.1.0-rc.1" { - t.Fatalf("unexpected request url: %s", req.URL.String()) - } - return &http.Response{ - StatusCode: http.StatusOK, - Header: make(http.Header), - Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)), - }, nil - }), - }, - } + service := testService(&http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v1.1.0-rc.1" { + t.Fatalf("unexpected request url: %s", req.URL.String()) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)), + }, nil + }), + }) - release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{Channel: "preview", TagName: "v1.1.0-rc.1", Force: true}) + release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{Channel: "preview", TagName: "v1.1.0-rc.1", Force: true}) if err != nil { t.Fatalf("expected tag release query to succeed: %v", err) } @@ -76,34 +78,26 @@ func TestGetReleaseByTag(t *testing.T) { } func TestCheckAndUpdateRequiresChecksumAsset(t *testing.T) { - originalVersion := config.Version - config.Version = "v1.0.0" - t.Cleanup(func() { - config.Version = originalVersion - }) - - assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH) - service := &Service{ - httpClient: &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" { - t.Fatalf("unexpected request url: %s", req.URL.String()) - } - return &http.Response{ - StatusCode: http.StatusOK, - Header: make(http.Header), - Body: io.NopCloser(strings.NewReader(`{ + assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH) + service := testService(&http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" { + t.Fatalf("unexpected request url: %s", req.URL.String()) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{ "tag_name":"v1.0.1", "assets":[ {"name":"` + assetName + `","browser_download_url":"https://example.test/agent"} ] }`)), - }, nil - }), - }, - } + }, nil + }), + }) - err := service.CheckAndUpdate(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{}) + err := service.CheckAndUpdate(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{}) if err == nil || !strings.Contains(err.Error(), "no matching checksum asset") { t.Fatalf("expected missing checksum asset error, got %v", err) } @@ -157,17 +151,15 @@ func TestDownloadAndRestartVerifiesChecksum(t *testing.T) { replaceAndRestartFunc = originalReplace }) - service := &Service{ - httpClient: &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - return &http.Response{ - StatusCode: http.StatusOK, - Header: make(http.Header), - Body: io.NopCloser(strings.NewReader(string(payload))), - }, nil - }), - }, - } + service := testService(&http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(string(payload))), + }, nil + }), + }) if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil { t.Fatalf("expected verified download to succeed: %v", err) @@ -198,17 +190,15 @@ func TestDownloadAndRestartRejectsChecksumMismatch(t *testing.T) { replaceAndRestartFunc = originalReplace }) - service := &Service{ - httpClient: &http.Client{ - Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { - return &http.Response{ - StatusCode: http.StatusOK, - Header: make(http.Header), - Body: io.NopCloser(strings.NewReader("tampered")), - }, nil - }), - }, - } + service := testService(&http.Client{ + Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader("tampered")), + }, nil + }), + }) err := service.downloadAndRestart(context.Background(), "https://example.test/agent", strings.Repeat("0", sha256.Size*2), targetPath) if err == nil || !strings.Contains(err.Error(), "sha256 checksum mismatch") { @@ -239,4 +229,4 @@ func TestIsNewerSupportsPrerelease(t *testing.T) { } }) } -} +} \ No newline at end of file diff --git a/internal/apps/flared/config/config.go b/internal/apps/flared/config/config.go index 09e8abfa..017bdb34 100644 --- a/internal/apps/flared/config/config.go +++ b/internal/apps/flared/config/config.go @@ -7,38 +7,11 @@ import ( "path/filepath" "strings" "time" + + edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config" ) -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 MillisecondDuration = edgeconfig.MillisecondDuration type Config struct { ServerURL string `json:"server_url"` diff --git a/internal/apps/flared/flared/runner.go b/internal/apps/flared/flared/runner.go index d9f444bc..3c62abad 100644 --- a/internal/apps/flared/flared/runner.go +++ b/internal/apps/flared/flared/runner.go @@ -3,8 +3,8 @@ package flared import ( "context" "log/slog" - "time" + edgerunner "github.com/Rain-kl/Wavelet/internal/apps/edge/runner" "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" @@ -23,31 +23,17 @@ type Runner struct { } 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 - } - + return edgerunner.RunWSReconnectLoop(ctx, edgerunner.WSReconnectConfig{ + ComponentName: "flared", + OnShutdown: r.FrpcManager.Stop, + }, func(ctx context.Context) (edgerunner.WSConnection, error) { + return r.WebSocketService.Connect(ctx) + }, func(ctx context.Context, conn edgerunner.WSConnection) { r.handleConnection(ctx, conn) - _ = conn.Close() - slog.Info("flared ws connection closed, reconnecting...") - r.sleepContext(ctx, 2*time.Second) - } + }) } type flaredWSHandler struct { @@ -73,13 +59,11 @@ 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): +func (r *Runner) handleConnection(ctx context.Context, conn edgerunner.WSConnection) { + wsConn, ok := conn.(*wsclient.Connection) + if !ok { + slog.Error("flared ws connection has unexpected type") + return } -} + _ = wsConn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r}) +} \ No newline at end of file diff --git a/internal/apps/flared/heartbeat/service.go b/internal/apps/flared/heartbeat/service.go index fcdb87c5..214c8ad1 100644 --- a/internal/apps/flared/heartbeat/service.go +++ b/internal/apps/flared/heartbeat/service.go @@ -3,23 +3,16 @@ package heartbeat import ( "context" "log/slog" - "net" - "time" + edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat" + "github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip" "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 @@ -37,32 +30,17 @@ func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) * } 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) - } - } + edgeheartbeat.RunLoop(ctx, s.config.HeartbeatInterval.Duration(), s.doHeartbeat) } 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 + IP: nodeip.Detect(), + TunnelStatus: "running", ConnectedRelays: s.frpcManager.GetConnectedRelays(), CurrentVersion: s.frpcManager.GetCurrentConfigVersion(), CurrentChecksum: s.frpcManager.GetCurrentConfigChecksum(), @@ -76,84 +54,19 @@ func (s *Service) doHeartbeat(ctx context.Context) { slog.Debug("flared heartbeat succeeded") if resp != nil && resp.TunnelSettings != nil { - s.tryAutoUpdate(ctx, resp.TunnelSettings) + edgeheartbeat.TryAutoUpdate(ctx, s.updater, tunnelSettingsToAutoUpdate(resp.TunnelSettings), "flared") } } -func (s *Service) tryAutoUpdate(ctx context.Context, settings *service.RelaySettings) { - if settings == nil || s.updater == nil { - return +func tunnelSettingsToAutoUpdate(settings *service.RelaySettings) *edgeheartbeat.AutoUpdateSettings { + if settings == nil { + return nil } - force := settings.UpdateNow - shouldCheck := settings.AutoUpdate || force - if !shouldCheck || settings.UpdateRepo == "" { - return + return &edgeheartbeat.AutoUpdateSettings{ + AutoUpdate: settings.AutoUpdate, + UpdateNow: settings.UpdateNow, + UpdateRepo: settings.UpdateRepo, + UpdateChannel: settings.UpdateChannel, + UpdateTag: settings.UpdateTag, } - 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 -} +} \ No newline at end of file diff --git a/internal/apps/flared/httpclient/client.go b/internal/apps/flared/httpclient/client.go index 821bdb35..17160296 100644 --- a/internal/apps/flared/httpclient/client.go +++ b/internal/apps/flared/httpclient/client.go @@ -1,16 +1,10 @@ package httpclient import ( - "bytes" "context" - "encoding/json" - "errors" - "io" - "log/slog" - "net/http" - "strings" "time" + edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient" service "github.com/Rain-kl/Wavelet/pkg/protocol" ) @@ -20,27 +14,21 @@ type APIResponse[T any] struct { } type Client struct { - baseURL string - token string - httpClient *http.Client + base *edgehttp.Client } func New(baseURL string, token string, timeout time.Duration) *Client { return &Client{ - baseURL: strings.TrimRight(baseURL, "/"), - token: token, - httpClient: &http.Client{ - Timeout: timeout, - }, + base: edgehttp.New(baseURL, token, timeout, "X-Tunnel-Token"), } } 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 { + if err := c.base.PostJSON(ctx, "/api/v1/tunnel/heartbeat", payload, &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } return &resp.Data, nil @@ -48,10 +36,10 @@ func (c *Client) Heartbeat(ctx context.Context, payload service.FlaredHeartbeatP 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 { + if err := c.base.GetJSON(ctx, "/api/v1/tunnel/config/active", &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } return &resp.Data, nil @@ -59,85 +47,12 @@ func (c *Client) GetActiveConfig(ctx context.Context) (*service.FlaredTunnelConf 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 { + if err := c.base.PostJSON(ctx, "/api/v1/tunnel/apply-log", payload, &resp); err != nil { return err } - return apiError(resp.ErrorMsg) + return edgehttp.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) -} + c.base.SetToken(token) +} \ No newline at end of file diff --git a/internal/apps/flared/updater/restart_windows.go b/internal/apps/flared/updater/restart_windows.go deleted file mode 100644 index 66cb5e61..00000000 --- a/internal/apps/flared/updater/restart_windows.go +++ /dev/null @@ -1,53 +0,0 @@ -//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, `"`, `""`) + `"` -} diff --git a/internal/apps/flared/updater/updater.go b/internal/apps/flared/updater/updater.go index 5d4faf46..9579cfd9 100644 --- a/internal/apps/flared/updater/updater.go +++ b/internal/apps/flared/updater/updater.go @@ -1,371 +1,17 @@ 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" - + edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater" "github.com/Rain-kl/Wavelet/internal/apps/flared/config" ) -const maxChecksumAssetSize = 64 * 1024 - -var replaceAndRestartFunc = replaceAndRestart - -type Service struct { - httpClient *http.Client - lastCheckKey string -} +type Service = edgeupdater.Service +type UpdateOptions = edgeupdater.UpdateOptions 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) -} + return edgeupdater.New(edgeupdater.Config{ + LocalVersion: config.Version, + AssetPrefix: "openflared", + LogLabel: "flared", + }) +} \ No newline at end of file diff --git a/internal/apps/relay/config/config.go b/internal/apps/relay/config/config.go index 1e00ef6f..41ae3dac 100644 --- a/internal/apps/relay/config/config.go +++ b/internal/apps/relay/config/config.go @@ -1,49 +1,18 @@ package config import ( - "context" "encoding/json" "errors" - "net" "os" "path/filepath" "strings" "time" - "github.com/Rain-kl/Wavelet/pkg/geoip" - "github.com/Rain-kl/Wavelet/pkg/geoip/iputil" + edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config" + "github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip" ) -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 MillisecondDuration = edgeconfig.MillisecondDuration type Config struct { ServerURL string `json:"server_url"` @@ -130,7 +99,7 @@ func applyDefaults(cfg *Config, baseDir string) { cfg.NodeName = strings.TrimSpace(host) } if cfg.NodeIP == "" { - cfg.NodeIP = detectNodeIP() + cfg.NodeIP = nodeip.Detect() } if cfg.StatePath == "" { cfg.StatePath = filepath.Join(cfg.DataDir, "relay-state.json") @@ -180,56 +149,4 @@ func (cfg *Config) Save() error { 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 -} diff --git a/internal/apps/relay/heartbeat/service.go b/internal/apps/relay/heartbeat/service.go index aa30fc3b..8dd91444 100644 --- a/internal/apps/relay/heartbeat/service.go +++ b/internal/apps/relay/heartbeat/service.go @@ -3,13 +3,13 @@ package heartbeat import ( "context" "log/slog" - "time" "github.com/Rain-kl/Wavelet/internal/apps/relay/config" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/httpclient" "github.com/Rain-kl/Wavelet/internal/apps/relay/observability" "github.com/Rain-kl/Wavelet/internal/apps/relay/state" + edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat" "github.com/Rain-kl/Wavelet/internal/apps/relay/updater" service "github.com/Rain-kl/Wavelet/pkg/protocol" ) @@ -33,20 +33,7 @@ func New(client *httpclient.Client, manager *frps.Manager, cfg *config.Config, s } 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) - } - } + edgeheartbeat.RunLoop(ctx, s.config.HeartbeatInterval.Duration(), s.doHeartbeat) } func (s *Service) doHeartbeat(ctx context.Context) { @@ -79,30 +66,19 @@ func (s *Service) doHeartbeat(ctx context.Context) { s.frpsManager.UpdateConfig(resp.RelayConfig) if resp != nil && resp.RelaySettings != nil { - s.tryAutoUpdate(ctx, resp.RelaySettings) + edgeheartbeat.TryAutoUpdate(ctx, s.updater, relaySettingsToAutoUpdate(resp.RelaySettings), "relay") } } -func (s *Service) tryAutoUpdate(ctx context.Context, settings *service.RelaySettings) { - if settings == nil || s.updater == nil { - return +func relaySettingsToAutoUpdate(settings *service.RelaySettings) *edgeheartbeat.AutoUpdateSettings { + if settings == nil { + return nil } - 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) + return &edgeheartbeat.AutoUpdateSettings{ + AutoUpdate: settings.AutoUpdate, + UpdateNow: settings.UpdateNow, + UpdateRepo: settings.UpdateRepo, + UpdateChannel: settings.UpdateChannel, + UpdateTag: settings.UpdateTag, } } diff --git a/internal/apps/relay/httpclient/client.go b/internal/apps/relay/httpclient/client.go index d580326a..641e4c7f 100644 --- a/internal/apps/relay/httpclient/client.go +++ b/internal/apps/relay/httpclient/client.go @@ -1,16 +1,10 @@ package httpclient import ( - "bytes" "context" - "encoding/json" - "errors" - "io" - "log/slog" - "net/http" - "strings" "time" + edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient" service "github.com/Rain-kl/Wavelet/pkg/protocol" ) @@ -20,105 +14,26 @@ type APIResponse[T any] struct { } type Client struct { - baseURL string - token string - httpClient *http.Client + base *edgehttp.Client } func New(baseURL string, token string, timeout time.Duration) *Client { return &Client{ - baseURL: strings.TrimRight(baseURL, "/"), - token: token, - httpClient: &http.Client{ - Timeout: timeout, - }, + base: edgehttp.New(baseURL, token, timeout, "X-Agent-Token"), } } func (c *Client) Heartbeat(ctx context.Context, payload service.RelayHeartbeatPayload) (*service.RelayHeartbeatResponse, error) { resp := APIResponse[service.RelayHeartbeatResponse]{} - if err := c.postJSON(ctx, "/api/v1/relay/heartbeat", payload, &resp); err != nil { + if err := c.base.PostJSON(ctx, "/api/v1/relay/heartbeat", payload, &resp); err != nil { return nil, err } - if err := apiError(resp.ErrorMsg); err != nil { + if err := edgehttp.APIError(resp.ErrorMsg); err != nil { return nil, err } 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) - - 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) -} + c.base.SetToken(token) +} \ No newline at end of file diff --git a/internal/apps/relay/observability/collector.go b/internal/apps/relay/observability/collector.go index dc103f6d..b2d4117c 100644 --- a/internal/apps/relay/observability/collector.go +++ b/internal/apps/relay/observability/collector.go @@ -1,18 +1,15 @@ package observability import ( - "bufio" "crypto/sha256" "encoding/hex" "encoding/json" "os" - "path/filepath" "runtime" - "strconv" "strings" - "syscall" "time" + edgeobs "github.com/Rain-kl/Wavelet/internal/apps/edge/observability" "github.com/Rain-kl/Wavelet/internal/apps/relay/config" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/state" @@ -43,15 +40,15 @@ func BuildSnapshot(cfg *config.Config, stateStore *state.Store) *service.AgentNo 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() + metric.MemoryTotalBytes, metric.MemoryUsedBytes = edgeobs.ReadMemInfo() + metric.StorageTotalBytes, metric.StorageUsedBytes = edgeobs.StatFilesystem(cfg.DataDir) + metric.NetworkRxBytes, metric.NetworkTxBytes = edgeobs.ReadLinuxNetworkTotals() + metric.DiskReadBytes, metric.DiskWriteBytes = edgeobs.ReadLinuxDiskTotals() if stateStore == nil { return metric } - totalCPU, idleCPU := readLinuxCPUStat() + totalCPU, idleCPU := edgeobs.ReadLinuxCPUStat() snapshot, err := stateStore.Load() if err != nil { return metric @@ -88,20 +85,20 @@ func BuildHealthEvents(status frps.RuntimeStatus) []service.AgentNodeHealthEvent func collectProfile(cfg *config.Config) *service.AgentNodeSystemProfile { hostname, _ := os.Hostname() - osName, osVersion := readLinuxOSRelease() - totalMemory, _ := readMemInfo() - totalDisk, _ := statFilesystem(cfg.DataDir) + osName, osVersion := edgeobs.ReadLinuxOSRelease() + totalMemory, _ := edgeobs.ReadMemInfo() + totalDisk, _ := edgeobs.StatFilesystem(cfg.DataDir) return &service.AgentNodeSystemProfile{ Hostname: strings.TrimSpace(hostname), OSName: osName, OSVersion: osVersion, - KernelVersion: readFirstLine("/proc/sys/kernel/osrelease"), + KernelVersion: edgeobs.ReadFirstLine("/proc/sys/kernel/osrelease"), Architecture: runtime.GOARCH, - CPUModel: readLinuxCPUModel(), + CPUModel: edgeobs.ReadLinuxCPUModel(), CPUCores: runtime.NumCPU(), TotalMemoryBytes: totalMemory, TotalDiskBytes: totalDisk, - UptimeSeconds: readLinuxUptimeSeconds(), + UptimeSeconds: edgeobs.ReadLinuxUptimeSeconds(), ReportedAtUnix: time.Now().UTC().Unix(), } } @@ -113,213 +110,4 @@ func fingerprintProfile(profile *service.AgentNodeSystemProfile) string { } 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)) -} +} \ No newline at end of file diff --git a/internal/apps/relay/relay/runner.go b/internal/apps/relay/relay/runner.go index 57ddd024..f6a0c64e 100644 --- a/internal/apps/relay/relay/runner.go +++ b/internal/apps/relay/relay/runner.go @@ -4,8 +4,8 @@ import ( "context" "encoding/json" "log/slog" - "time" + edgerunner "github.com/Rain-kl/Wavelet/internal/apps/edge/runner" "github.com/Rain-kl/Wavelet/internal/apps/relay/config" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/heartbeat" @@ -25,30 +25,16 @@ type Runner struct { } 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 - } - + return edgerunner.RunWSReconnectLoop(ctx, edgerunner.WSReconnectConfig{ + ComponentName: "relay", + OnShutdown: r.FrpsManager.Stop, + }, func(ctx context.Context) (edgerunner.WSConnection, error) { + return r.WebSocketService.Connect(ctx) + }, func(ctx context.Context, conn edgerunner.WSConnection) { r.handleConnection(ctx, conn) - _ = conn.Close() - slog.Info("relay ws connection closed, reconnecting...") - r.sleepContext(ctx, 2*time.Second) - } + }) } type relayWSHandler struct { @@ -78,13 +64,11 @@ 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): +func (r *Runner) handleConnection(ctx context.Context, conn edgerunner.WSConnection) { + wsConn, ok := conn.(*wsclient.Connection) + if !ok { + slog.Error("relay ws connection has unexpected type") + return } -} + _ = wsConn.RunReceiveLoop(ctx, &relayWSHandler{runner: r}) +} \ No newline at end of file diff --git a/internal/apps/relay/updater/restart_unix.go b/internal/apps/relay/updater/restart_unix.go deleted file mode 100644 index 173af72b..00000000 --- a/internal/apps/relay/updater/restart_unix.go +++ /dev/null @@ -1,51 +0,0 @@ -//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 -} diff --git a/internal/apps/relay/updater/updater.go b/internal/apps/relay/updater/updater.go index 524be4a3..3536cda2 100644 --- a/internal/apps/relay/updater/updater.go +++ b/internal/apps/relay/updater/updater.go @@ -1,371 +1,17 @@ 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" - + edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater" "github.com/Rain-kl/Wavelet/internal/apps/relay/config" ) -const maxChecksumAssetSize = 64 * 1024 - -var replaceAndRestartFunc = replaceAndRestart - -type Service struct { - httpClient *http.Client - lastCheckKey string -} +type Service = edgeupdater.Service +type UpdateOptions = edgeupdater.UpdateOptions 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) -} + return edgeupdater.New(edgeupdater.Config{ + LocalVersion: config.Version, + AssetPrefix: "openflare-relay", + LogLabel: "relay", + }) +} \ No newline at end of file diff --git a/internal/apps/agent/protocol/agent_api.go b/pkg/protocol/agent.go similarity index 92% rename from internal/apps/agent/protocol/agent_api.go rename to pkg/protocol/agent.go index 66c8ef3c..b81c9e0e 100644 --- a/internal/apps/agent/protocol/agent_api.go +++ b/pkg/protocol/agent.go @@ -159,17 +159,6 @@ type RegisterNodeResponse struct { Name string `json:"name"` } -type ApplyLogPayload struct { - NodeID string `json:"node_id"` - Version string `json:"version"` - Result string `json:"result"` - Message string `json:"message"` - Checksum string `json:"checksum"` - MainConfigChecksum string `json:"main_config_checksum"` - RouteConfigChecksum string `json:"route_config_checksum"` - SupportFileCount int `json:"support_file_count"` -} - type ActiveConfigResponse struct { Version string `json:"version"` Checksum string `json:"checksum"` @@ -178,11 +167,6 @@ type ActiveConfigResponse struct { CreatedAt string `json:"created_at"` } -type ActiveConfigMeta struct { - Version string `json:"version"` - Checksum string `json:"checksum"` -} - type WAFIPGroup struct { ID uint `json:"id"` Name string `json:"name"` @@ -204,4 +188,4 @@ type WAFIPGroupSyncResponse struct { type SupportFile struct { Path string `json:"path"` Content string `json:"content"` -} +} \ No newline at end of file diff --git a/pkg/protocol/agent_test.go b/pkg/protocol/agent_test.go new file mode 100644 index 00000000..d89b5303 --- /dev/null +++ b/pkg/protocol/agent_test.go @@ -0,0 +1,95 @@ +package protocol + +import ( + "encoding/json" + "reflect" + "testing" +) + +func TestAgentProtocolJSONTags(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + value any + expected map[string]string + }{ + { + name: "NodePayload", + value: NodePayload{}, + expected: map[string]string{ + "NodeID": "node_id", + "Name": "name", + }, + }, + { + name: "AgentSettings", + value: AgentSettings{}, + expected: map[string]string{ + "HeartbeatInterval": "heartbeat_interval", + "WebsocketUpgradeEnabled": "websocket_upgrade_enabled", + "RestartOpenrestyNow": "restart_openresty_now", + }, + }, + { + name: "WSMessage", + value: WSMessage{}, + expected: map[string]string{ + "Type": "type", + "Payload": "payload,omitempty", + }, + }, + { + name: "RegisterNodeResponse", + value: RegisterNodeResponse{}, + expected: map[string]string{ + "NodeID": "node_id", + "AccessToken": "agent_token", + }, + }, + } + + for _, tc := range cases { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + typ := reflect.TypeOf(tc.value) + for field, wantTag := range tc.expected { + structField, ok := typ.FieldByName(field) + if !ok { + t.Fatalf("field %q not found on %s", field, tc.name) + } + gotTag := structField.Tag.Get("json") + if gotTag != wantTag { + t.Fatalf("field %q json tag = %q, want %q", field, gotTag, wantTag) + } + } + }) + } +} + +func TestNodePayloadJSONRoundTrip(t *testing.T) { + t.Parallel() + + payload := NodePayload{ + NodeID: "node-1", + Name: "edge-a", + OpenrestyStatus: OpenrestyStatusHealthy, + HealthEvents: []NodeHealthEvent{}, + } + + encoded, err := json.Marshal(payload) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var decoded NodePayload + if err := json.Unmarshal(encoded, &decoded); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if decoded.NodeID != payload.NodeID || decoded.Name != payload.Name { + t.Fatalf("round trip mismatch: %+v", decoded) + } +} \ No newline at end of file diff --git a/pkg/protocol/server.go b/pkg/protocol/server.go index 4e6e64ed..0a12273c 100644 --- a/pkg/protocol/server.go +++ b/pkg/protocol/server.go @@ -1,10 +1,5 @@ package protocol -type WSMessage struct { - Type string `json:"type"` - Payload any `json:"payload,omitempty"` -} - type AgentNodeSystemProfile struct { Hostname string `json:"hostname"` OSName string `json:"os_name"`