refactor(edge): 抽取边缘运行时共享包并完成 Phase 3 重构

- 新增 internal/apps/edge/,三组件改为薄包装,删除 3000+ 行重复代码
- Agent 心跳周期下沉至 heartbeat/cycle.go
- 协议类型迁入 pkg/protocol/agent.go
- 补充设计文档与 changelog
This commit is contained in:
ryan
2026-06-19 14:56:00 +08:00
parent cc5e53c51e
commit db9a9f98fd
53 changed files with 2091 additions and 2947 deletions
+16 -8
View File
@@ -92,15 +92,23 @@ func main() {
} }
syncService := syncservice.New(client, runtimeManager, stateStore) syncService := syncservice.New(client, runtimeManager, stateStore)
syncService.SetPagesDir(cfg.PagesDir) syncService.SetPagesDir(cfg.PagesDir)
heartbeatService := heartbeat.New(client)
updateService := updater.New()
runner := &agent.Runner{ runner := &agent.Runner{
Config: cfg, Config: cfg,
StateStore: stateStore, StateStore: stateStore,
ObservabilityBuffer: observabilityBuffer, HeartbeatCycle: &heartbeat.Cycle{
HeartbeatService: heartbeat.New(client), Config: cfg,
SyncService: syncService, StateStore: stateStore,
Updater: updater.New(), ObservabilityBuffer: observabilityBuffer,
RuntimeManager: runtimeManager, Heartbeat: heartbeatService,
WebSocketService: wsClient, Sync: syncService,
Updater: updateService,
},
HeartbeatService: heartbeatService,
SyncService: syncService,
RuntimeManager: runtimeManager,
WebSocketService: wsClient,
} }
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
+2 -18
View File
@@ -6,9 +6,9 @@ import (
"log/slog" "log/slog"
"os" "os"
"os/signal" "os/signal"
"strings"
"syscall" "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/config"
"github.com/Rain-kl/Wavelet/internal/apps/flared/flared" "github.com/Rain-kl/Wavelet/internal/apps/flared/flared"
"github.com/Rain-kl/Wavelet/internal/apps/flared/frpc" "github.com/Rain-kl/Wavelet/internal/apps/flared/frpc"
@@ -19,10 +19,7 @@ import (
) )
func main() { func main() {
// Setup simple structured logging edgelogging.Setup(edgelogging.Options{})
slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{
Level: parseLevel(os.Getenv("LOG_LEVEL")),
})))
configPath := flag.String("config", "./flared.json", "flared config path") configPath := flag.String("config", "./flared.json", "flared config path")
flag.Parse() flag.Parse()
@@ -72,16 +69,3 @@ func main() {
} }
slog.Info("flared process stopped") 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
}
}
+2 -18
View File
@@ -6,9 +6,9 @@ import (
"log/slog" "log/slog"
"os" "os"
"os/signal" "os/signal"
"strings"
"syscall" "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/config"
"github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps"
"github.com/Rain-kl/Wavelet/internal/apps/relay/heartbeat" "github.com/Rain-kl/Wavelet/internal/apps/relay/heartbeat"
@@ -19,10 +19,7 @@ import (
) )
func main() { func main() {
// Setup simple structured logging edgelogging.Setup(edgelogging.Options{})
slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{
Level: parseLevel(os.Getenv("LOG_LEVEL")),
})))
configPath := flag.String("config", "./relay.json", "relay config path") configPath := flag.String("config", "./relay.json", "relay config path")
flag.Parse() flag.Parse()
@@ -72,16 +69,3 @@ func main() {
} }
slog.Info("relay process stopped") 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
}
}
+3
View File
@@ -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 的构建上下文为根目录。 - 合并并简化仓库结构:将 `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` 等)。 - 调整子项目结构与包路径:将 `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: - 调整编译产物输出名称与 Makefile:
+121
View File
@@ -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/`。
+1
View File
@@ -189,3 +189,4 @@ OpenFlare 已收敛为**单 monorepo**(Go 模块 `github.com/Rain-kl/Wavelet`
* 发布、同步、回滚与 Agent 模型变化:更新 [Agent 与发布模型](./agent-design.md)。 * 发布、同步、回滚与 Agent 模型变化:更新 [Agent 与发布模型](./agent-design.md)。
* 部署方式变化:更新 [部署说明](../deployment/deployment.md) 与 README。 * 部署方式变化:更新 [部署说明](../deployment/deployment.md) 与 README。
* 配置项变化:更新 [配置项参考](../reference/configuration.md)。 * 配置项变化:更新 [配置项参考](../reference/configuration.md)。
* 边缘组件共享运行时重构:更新 [边缘运行时重构](./edge-runtime-refactor.md)。
+53
View File
@@ -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
```
+1
View File
@@ -19,6 +19,7 @@
| [OpenFlare 前端迁移 — AI 委派](./handover-openflare-frontend-migration.md) | 前端迁移任务队列与验收状态 | | [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) | 多角度迁移验收报告 | | [前端路由验证](./verify-frontend-routes.md) · [Service 验证](./verify-frontend-services.md) · [UI 验证](./verify-frontend-ui.md) · [构建验证](./verify-frontend-build.md) | 多角度迁移验收报告 |
| [文档结构更新 — AI 接手](./handover-docs-restructure-update.md) | 重构后多智能体分析结论与文档批量更新记录 | | [文档结构更新 — AI 接手](./handover-docs-restructure-update.md) | 重构后多智能体分析结论与文档批量更新记录 |
| [边缘运行时 Phase 3 任务拆解](./20260619-edge-phase3-tasks.md) | Batch 1 并行任务(duration / observability / heartbeat loop) |
## 使用建议 ## 使用建议
+22 -217
View File
@@ -9,10 +9,11 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config" "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/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state" "github.com/Rain-kl/Wavelet/internal/apps/agent/state"
"github.com/Rain-kl/Wavelet/internal/apps/agent/wsclient" "github.com/Rain-kl/Wavelet/internal/apps/agent/wsclient"
edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat"
) )
type HeartbeatService interface { type HeartbeatService interface {
@@ -29,10 +30,6 @@ type SyncService interface {
ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error
} }
type Updater interface {
CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error
}
type RuntimeManager interface { type RuntimeManager interface {
CheckHealth(ctx context.Context) error CheckHealth(ctx context.Context) error
Restart(ctx context.Context) error Restart(ctx context.Context) error
@@ -44,32 +41,23 @@ type WebSocketService interface {
URL() string URL() string
} }
type UpdateOptions struct {
Channel string
TagName string
Force bool
}
type Runner struct { type Runner struct {
Config *config.Config Config *config.Config
StateStore *state.Store StateStore *state.Store
ObservabilityBuffer *state.ObservabilityBufferStore HeartbeatCycle *agentheartbeat.Cycle
HeartbeatService HeartbeatService HeartbeatService HeartbeatService
SyncService SyncService SyncService SyncService
Updater Updater
RuntimeManager RuntimeManager RuntimeManager RuntimeManager
WebSocketService WebSocketService WebSocketService WebSocketService
autoUpdate bool
updateNow bool
updateRepo string
updateChan string
updateTag string
restartOpenrestyNow bool restartOpenrestyNow bool
websocketUpgradeEnabled bool websocketUpgradeEnabled bool
} }
func (r *Runner) Run(ctx context.Context) error { func (r *Runner) Run(ctx context.Context) error {
if r.HeartbeatCycle != nil {
r.HeartbeatCycle.RecordSyncError = r.recordSyncError
}
nodeID, err := r.StateStore.EnsureNodeID() nodeID, err := r.StateStore.EnsureNodeID()
if err != nil { if err != nil {
return err 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) { func (r *Runner) performHeartbeatCycle(ctx context.Context, nodeID string, startup bool) (bool, error) {
r.refreshOpenrestyHealth(ctx) r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(nodeID) return r.HeartbeatCycle.Perform(ctx, nodeID, startup, r)
heartbeatResult, err := r.HeartbeatService.Heartbeat(ctx, payload) }
if err != nil {
return false, err func (r *Runner) Apply(settings *protocol.AgentSettings) bool {
} return r.applySettings(settings)
r.ackObservabilityWindows(ackWindows) }
if heartbeatResult == nil {
heartbeatResult = &protocol.HeartbeatResult{} func (r *Runner) RestartOpenrestyIfNeeded(ctx context.Context) {
}
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)
}
r.tryRestartOpenresty(ctx) r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
return changed, nil
} }
func (r *Runner) shouldUseWebSocket() bool { 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) childCtx, cancel := context.WithCancel(ctx)
defer cancel() defer cancel()
// Start status ticker sender in background
go func() { go func() {
for { for {
select { 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 { func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
r.refreshOpenrestyHealth(ctx) r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(nodeID) payload, ackWindows := r.HeartbeatCycle.PrepareHeartbeatPayload(nodeID)
if err := conn.SendStatus(payload); err != nil { if err := conn.SendStatus(payload); err != nil {
return err return err
} }
r.ackObservabilityWindows(ackWindows) r.HeartbeatCycle.AckObservabilityWindows(ackWindows)
return nil return nil
} }
@@ -303,7 +269,7 @@ func (r *Runner) handleWebSocketMessage(ctx context.Context, message protocol.WS
} }
changed := r.applySettings(&settings) changed := r.applySettings(&settings)
r.tryRestartOpenresty(ctx) r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx) edgeheartbeat.TryAutoUpdate(ctx, r.HeartbeatCycle.Updater, agentheartbeat.AgentSettingsToAutoUpdate(&settings), "agent")
if !r.websocketUpgradeEnabled { if !r.websocketUpgradeEnabled {
slog.Debug("agent ws disabled by server settings; falling back to http heartbeat") slog.Debug("agent ws disabled by server settings; falling back to http heartbeat")
return changed, errors.New("websocket upgrade disabled by server") 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) slog.Debug("agent ws waf ip groups decode failed", "error", err)
return false, nil return false, nil
} }
r.applyWAFIPGroups(ctx, groups) r.HeartbeatCycle.ApplyWAFIPGroups(ctx, groups)
return false, nil return false, nil
case protocol.WSMessageTypePing: case protocol.WSMessageTypePing:
slog.Debug("agent ws ping received") 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) slog.Debug("agent websocket upgrade setting updated", "from", r.websocketUpgradeEnabled, "to", settings.WebsocketUpgradeEnabled)
} }
r.websocketUpgradeEnabled = 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 r.restartOpenrestyNow = settings.RestartOpenrestyNow
return changed return changed
} }
@@ -436,37 +397,12 @@ func (r *Runner) tryRestartOpenresty(ctx context.Context) {
r.recordOpenrestyHealthy() 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 { func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
if strings.TrimSpace(r.Config.DiscoveryToken) == "" { if strings.TrimSpace(r.Config.DiscoveryToken) == "" {
return errors.New("agent_token 为空且未配置 discovery_token") return errors.New("agent_token 为空且未配置 discovery_token")
} }
slog.Info("agent discovery registration started") 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 { if err != nil {
return err return err
} }
@@ -493,26 +429,9 @@ func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
*nodeID = response.NodeID *nodeID = response.NodeID
slog.Info("agent discovery registration succeeded", "node_id", response.NodeID) slog.Info("agent discovery registration succeeded", "node_id", response.NodeID)
r.refreshOpenrestyHealth(ctx) r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(*nodeID) if _, err = r.HeartbeatCycle.Perform(ctx, *nodeID, true, r); err != nil {
heartbeatResult, heartbeatErr := r.HeartbeatService.Heartbeat(ctx, payload) slog.Error("agent post-register heartbeat failed", "error", err)
if heartbeatErr != nil {
slog.Error("agent post-register heartbeat failed", "error", heartbeatErr)
return nil
} }
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 return nil
} }
@@ -582,118 +501,4 @@ func (r *Runner) recordOpenrestyUnhealthy(err error, fallbackOnly bool) {
if saveErr := r.StateStore.Save(snapshot); saveErr != nil { if saveErr := r.StateStore.Save(snapshot); saveErr != nil {
slog.Error("save state after recording openresty error failed", "error", saveErr) 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)
}
}
+34 -21
View File
@@ -11,10 +11,24 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config" "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/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state" "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 { type fakeHeartbeatService struct {
mu sync.Mutex mu sync.Mutex
registerCalls int registerCalls int
@@ -192,7 +206,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
syncService := &fakeSyncService{ syncService := &fakeSyncService{
startupErr: errors.New("当前没有激活版本,保持当前 OpenResty 配置"), startupErr: errors.New("当前没有激活版本,保持当前 OpenResty 配置"),
} }
runner := &Runner{ runner := withHeartbeatCycle(&Runner{
Config: &config.Config{ Config: &config.Config{
AccessToken: "agent-token", AccessToken: "agent-token",
NodeName: "edge-01", NodeName: "edge-01",
@@ -204,7 +218,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
StateStore: stateStore, StateStore: stateStore,
HeartbeatService: heartbeatService, HeartbeatService: heartbeatService,
SyncService: syncService, SyncService: syncService,
} }, nil)
err := runner.Run(ctx) err := runner.Run(ctx)
if !errors.Is(err, context.Canceled) { if !errors.Is(err, context.Canceled) {
@@ -245,7 +259,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
} }
}, },
} }
runner := &Runner{ runner := withHeartbeatCycle(&Runner{
Config: &config.Config{ Config: &config.Config{
AccessToken: "agent-token", AccessToken: "agent-token",
NodeName: "edge-01", NodeName: "edge-01",
@@ -257,7 +271,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
StateStore: stateStore, StateStore: stateStore,
HeartbeatService: heartbeatService, HeartbeatService: heartbeatService,
SyncService: syncService, SyncService: syncService,
} }, nil)
err := runner.Run(ctx) err := runner.Run(ctx)
if !errors.Is(err, context.Canceled) { if !errors.Is(err, context.Canceled) {
@@ -303,7 +317,7 @@ func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) {
healthErr: errors.New("docker openresty container is not running"), healthErr: errors.New("docker openresty container is not running"),
clearHealthOnRestart: true, clearHealthOnRestart: true,
} }
runner := &Runner{ runner := withHeartbeatCycle(&Runner{
Config: &config.Config{ Config: &config.Config{
AccessToken: "agent-token", AccessToken: "agent-token",
NodeName: "edge-01", NodeName: "edge-01",
@@ -316,7 +330,7 @@ func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) {
HeartbeatService: heartbeatService, HeartbeatService: heartbeatService,
SyncService: &fakeSyncService{}, SyncService: &fakeSyncService{},
RuntimeManager: runtimeManager, RuntimeManager: runtimeManager,
} }, nil)
err := runner.Run(ctx) err := runner.Run(ctx)
if !errors.Is(err, context.Canceled) { if !errors.Is(err, context.Canceled) {
@@ -357,7 +371,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
t.Fatalf("failed to seed state: %v", err) t.Fatalf("failed to seed state: %v", err)
} }
runner := &Runner{ runner := withHeartbeatCycle(&Runner{
Config: &config.Config{ Config: &config.Config{
NodeName: "edge-observe-1", NodeName: "edge-observe-1",
NodeIP: "10.0.0.51", NodeIP: "10.0.0.51",
@@ -369,7 +383,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond), HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
}, },
StateStore: stateStore, StateStore: stateStore,
} }, nil)
if err := os.MkdirAll(filepath.Dir(runner.Config.AccessLogPath), 0o755); err != nil { if err := os.MkdirAll(filepath.Dir(runner.Config.AccessLogPath), 0o755); err != nil {
t.Fatalf("failed to prepare access log dir: %v", err) 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) t.Fatalf("failed to prepare access log: %v", err)
} }
firstPayload := runner.nodePayload("node-observe") firstPayload := runner.HeartbeatCycle.NodePayload("node-observe")
if firstPayload.Profile == nil { if firstPayload.Profile == nil {
t.Fatal("expected first heartbeat payload to include system profile") 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) 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 { if secondPayload.Profile != nil {
t.Fatal("expected unchanged profile to be omitted on subsequent heartbeat") 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{ Config: &config.Config{
AccessToken: "agent-token", AccessToken: "agent-token",
NodeName: "edge-buffer-01", NodeName: "edge-buffer-01",
@@ -451,11 +465,10 @@ func TestRunnerReplaysBufferedObservabilityAfterHeartbeatRecovery(t *testing.T)
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond), HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
ObservabilityReplayMinutes: 15, ObservabilityReplayMinutes: 15,
}, },
StateStore: stateStore, StateStore: stateStore,
ObservabilityBuffer: bufferStore, HeartbeatService: heartbeatService,
HeartbeatService: heartbeatService, SyncService: &fakeSyncService{},
SyncService: &fakeSyncService{}, }, bufferStore)
}
if err := os.MkdirAll(filepath.Dir(runner.Config.RouteConfigPath), 0o755); err != nil { if err := os.MkdirAll(filepath.Dir(runner.Config.RouteConfigPath), 0o755); err != nil {
t.Fatalf("failed to prepare route config dir: %v", err) t.Fatalf("failed to prepare route config dir: %v", err)
} }
@@ -518,7 +531,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("failed to load config: %v", err) t.Fatalf("failed to load config: %v", err)
} }
runner := &Runner{ runner := withHeartbeatCycle(&Runner{
Config: &config.Config{ Config: &config.Config{
ServerURL: cfg.ServerURL, ServerURL: cfg.ServerURL,
DiscoveryToken: cfg.DiscoveryToken, DiscoveryToken: cfg.DiscoveryToken,
@@ -531,7 +544,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
StateStore: stateStore, StateStore: stateStore,
HeartbeatService: heartbeatService, HeartbeatService: heartbeatService,
SyncService: syncService, SyncService: syncService,
} }, nil)
runner.Config = cfg runner.Config = cfg
runner.Config.Version = config.Version runner.Config.Version = config.Version
runner.Config.ExtVersion = "1.27.1.2" runner.Config.ExtVersion = "1.27.1.2"
@@ -561,7 +574,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) { func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) {
syncService := &fakeSyncService{} syncService := &fakeSyncService{}
runner := &Runner{SyncService: syncService} runner := withHeartbeatCycle(&Runner{SyncService: syncService}, nil)
payload, err := json.Marshal(protocol.ActiveConfigMeta{ payload, err := json.Marshal(protocol.ActiveConfigMeta{
Version: "20260529-001", Version: "20260529-001",
Checksum: "checksum-ws", Checksum: "checksum-ws",
@@ -589,12 +602,12 @@ func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) {
} }
func TestRunnerHandlesWebSocketSettingsDisabled(t *testing.T) { func TestRunnerHandlesWebSocketSettingsDisabled(t *testing.T) {
runner := &Runner{ runner := withHeartbeatCycle(&Runner{
Config: &config.Config{ Config: &config.Config{
HeartbeatInterval: config.MillisecondDuration(10 * time.Second), HeartbeatInterval: config.MillisecondDuration(10 * time.Second),
}, },
websocketUpgradeEnabled: true, websocketUpgradeEnabled: true,
} }, nil)
payload, err := json.Marshal(protocol.AgentSettings{ payload, err := json.Marshal(protocol.AgentSettings{
HeartbeatInterval: 15000, HeartbeatInterval: 15000,
WebsocketUpgradeEnabled: false, WebsocketUpgradeEnabled: false,
+2 -70
View File
@@ -1,11 +1,9 @@
package config package config
import ( import (
"context"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"net"
"os" "os"
pathpkg "path" pathpkg "path"
"path/filepath" "path/filepath"
@@ -13,8 +11,7 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/pkg/geoip" "github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
"github.com/Rain-kl/Wavelet/pkg/utils" "github.com/Rain-kl/Wavelet/pkg/utils"
) )
@@ -35,11 +32,6 @@ const (
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb" defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
) )
var (
lookupOutboundIP = geoip.GetOutboundIP
lookupLocalIP = detectLocalNodeIP
)
type Config struct { type Config struct {
ServerURL string `json:"server_url"` ServerURL string `json:"server_url"`
AccessToken string `json:"agent_token"` AccessToken string `json:"agent_token"`
@@ -166,7 +158,7 @@ func applyDefaults(cfg *Config, baseDir string) {
cfg.NodeName = detectHostname() cfg.NodeName = detectHostname()
} }
if cfg.NodeIP == "" { if cfg.NodeIP == "" {
cfg.NodeIP = detectNodeIP() cfg.NodeIP = nodeip.Detect()
} }
if cfg.MainConfigPath == "" { if cfg.MainConfigPath == "" {
cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultMainConfigRelativePath) cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultMainConfigRelativePath)
@@ -400,64 +392,4 @@ func detectHostname() string {
return strings.TrimSpace(host) 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)
}
+12 -10
View File
@@ -10,7 +10,9 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip"
"github.com/Rain-kl/Wavelet/pkg/geoip" "github.com/Rain-kl/Wavelet/pkg/geoip"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
) )
func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) { func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) {
@@ -248,12 +250,12 @@ func TestLoadUsesEnvConfigWhenFileIsMissing(t *testing.T) {
} }
func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) { func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) {
previousLookup := lookupOutboundIP previousLookup := nodeip.LookupOutboundIP
lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) { nodeip.LookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
return net.ParseIP("8.8.8.8"), nil return net.ParseIP("8.8.8.8"), nil
} }
defer func() { defer func() {
lookupOutboundIP = previousLookup nodeip.LookupOutboundIP = previousLookup
}() }()
dir := t.TempDir() dir := t.TempDir()
@@ -281,17 +283,17 @@ func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) {
} }
func TestLoadFallsBackToLocalIPWhenOutboundLookupFails(t *testing.T) { func TestLoadFallsBackToLocalIPWhenOutboundLookupFails(t *testing.T) {
previousOutboundLookup := lookupOutboundIP previousOutboundLookup := nodeip.LookupOutboundIP
previousLocalLookup := lookupLocalIP previousLocalLookup := nodeip.LookupLocalIP
lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) { nodeip.LookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
return nil, errors.New("realip.cc unavailable") return nil, errors.New("realip.cc unavailable")
} }
lookupLocalIP = func() string { nodeip.LookupLocalIP = func() string {
return "9.9.9.9" return "9.9.9.9"
} }
defer func() { defer func() {
lookupOutboundIP = previousOutboundLookup nodeip.LookupOutboundIP = previousOutboundLookup
lookupLocalIP = previousLocalLookup nodeip.LookupLocalIP = previousLocalLookup
}() }()
dir := t.TempDir() dir := t.TempDir()
@@ -512,7 +514,7 @@ func TestNodeIPPriority(t *testing.T) {
if tt.ip != "" { if tt.ip != "" {
parsed = net.ParseIP(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) t.Fatalf("unexpected priority for %q: got %d want %d", tt.ip, got, tt.expected)
} }
}) })
+2 -51
View File
@@ -1,54 +1,5 @@
package config package config
import ( import edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
"encoding/json"
"fmt"
"strconv"
"strings"
"time"
)
type MillisecondDuration time.Duration type MillisecondDuration = edgeconfig.MillisecondDuration
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())
}
+217
View File
@@ -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)
}
+17 -116
View File
@@ -1,55 +1,44 @@
package httpclient package httpclient
import ( import (
"bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"log/slog"
"net/http" "net/http"
"strings"
"time" "time"
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
) )
type Client struct { type Client struct {
baseURL string base *edgehttp.Client
token string
httpClient *http.Client
} }
func New(baseURL string, token string, timeout time.Duration) *Client { func New(baseURL string, token string, timeout time.Duration) *Client {
return &Client{ return &Client{
baseURL: strings.TrimRight(baseURL, "/"), base: edgehttp.New(baseURL, token, timeout, "X-Agent-Token"),
token: token,
httpClient: &http.Client{
Timeout: timeout,
},
} }
} }
func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) { 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]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err return nil, err
} }
slog.Debug("http register node response", "node_id", resp.Data.NodeID)
return &resp.Data, nil return &resp.Data, nil
} }
func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) { func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) {
resp := protocol.APIResponse[protocol.HeartbeatData]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err return nil, err
} }
return &protocol.HeartbeatResult{ 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) { func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
resp := protocol.APIResponse[protocol.ActiveConfigResponse]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err 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 return &resp.Data, nil
} }
func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error { 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]{} 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 err
} }
return apiError(resp.ErrorMsg) return edgehttp.APIError(resp.ErrorMsg)
} }
func (c *Client) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) { func (c *Client) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
resp := protocol.APIResponse[protocol.WAFIPGroupSyncResponse]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err return nil, err
} }
return &resp.Data, nil return &resp.Data, nil
} }
func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) { 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) 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
}
req.Header.Set("X-Agent-Token", c.token)
res, err := c.httpClient.Do(req)
if err != nil { if err != nil {
return nil, err return nil, err
} }
defer res.Body.Close() defer res.Body.Close()
if res.StatusCode != http.StatusOK { if res.StatusCode != http.StatusOK {
return nil, readHTTPError(res) return nil, edgehttp.ReadHTTPError(res)
} }
return io.ReadAll(res.Body) return io.ReadAll(res.Body)
} }
func (c *Client) SetToken(token string) { func (c *Client) SetToken(token string) {
c.token = strings.TrimSpace(token) c.base.SetToken(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)
}
+3 -25
View File
@@ -1,29 +1,7 @@
package logging package logging
import ( import edgelogging "github.com/Rain-kl/Wavelet/internal/apps/edge/logging"
"log/slog"
"os"
"strings"
)
func Setup() { func Setup() {
opts := &slog.HandlerOptions{ edgelogging.Setup(edgelogging.Options{AddSource: true})
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
}
}
+13 -267
View File
@@ -1,21 +1,18 @@
package observability package observability
import ( import (
"bufio"
"crypto/sha256" "crypto/sha256"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"os" "os"
"path/filepath"
"runtime" "runtime"
"strconv"
"strings" "strings"
"syscall"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config" "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/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state" "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 { 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(), CapturedAtUnix: now.Unix(),
} }
memTotal, memUsed := readMemInfo() memTotal, memUsed := edgeobs.ReadMemInfo()
metric.MemoryTotalBytes = memTotal metric.MemoryTotalBytes = memTotal
metric.MemoryUsedBytes = memUsed metric.MemoryUsedBytes = memUsed
storageTotal, storageUsed := statFilesystem(cfg.DataDir) storageTotal, storageUsed := edgeobs.StatFilesystem(cfg.DataDir)
metric.StorageTotalBytes = storageTotal metric.StorageTotalBytes = storageTotal
metric.StorageUsedBytes = storageUsed metric.StorageUsedBytes = storageUsed
metric.NetworkRxBytes, metric.NetworkTxBytes = readLinuxNetworkTotals() metric.NetworkRxBytes, metric.NetworkTxBytes = edgeobs.ReadLinuxNetworkTotals()
metric.DiskReadBytes, metric.DiskWriteBytes = readLinuxDiskTotals() metric.DiskReadBytes, metric.DiskWriteBytes = edgeobs.ReadLinuxDiskTotals()
if stateStore == nil { if stateStore == nil {
return metric return metric
} }
totalCPU, idleCPU := readLinuxCPUStat() totalCPU, idleCPU := edgeobs.ReadLinuxCPUStat()
snapshot, err := stateStore.Load() snapshot, err := stateStore.Load()
if err != nil { if err != nil {
return metric return metric
@@ -121,12 +118,12 @@ func BuildHealthEvents(snapshot *state.Snapshot) []protocol.NodeHealthEvent {
func collectProfile(cfg *config.Config) *protocol.NodeSystemProfile { func collectProfile(cfg *config.Config) *protocol.NodeSystemProfile {
hostname, _ := os.Hostname() hostname, _ := os.Hostname()
osName, osVersion := readLinuxOSRelease() osName, osVersion := edgeobs.ReadLinuxOSRelease()
kernelVersion := readFirstLine("/proc/sys/kernel/osrelease") kernelVersion := edgeobs.ReadFirstLine("/proc/sys/kernel/osrelease")
cpuModel := readLinuxCPUModel() cpuModel := edgeobs.ReadLinuxCPUModel()
totalMemory, _ := readMemInfo() totalMemory, _ := edgeobs.ReadMemInfo()
totalDisk, _ := statFilesystem(cfg.DataDir) totalDisk, _ := edgeobs.StatFilesystem(cfg.DataDir)
uptimeSeconds := readLinuxUptimeSeconds() uptimeSeconds := edgeobs.ReadLinuxUptimeSeconds()
return &protocol.NodeSystemProfile{ return &protocol.NodeSystemProfile{
Hostname: strings.TrimSpace(hostname), Hostname: strings.TrimSpace(hostname),
@@ -150,255 +147,4 @@ func fingerprintProfile(profile *protocol.NodeSystemProfile) string {
} }
sum := sha256.Sum256(raw) sum := sha256.Sum256(raw)
return hex.EncodeToString(sum[:]) 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))
}
+43
View File
@@ -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
)
@@ -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
}
@@ -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, `"`, `""`) + `"`
}
+9 -358
View File
@@ -1,366 +1,17 @@
package updater package updater
import ( import (
"context" edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
"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"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config" "github.com/Rain-kl/Wavelet/internal/apps/agent/config"
) )
const maxChecksumAssetSize = 64 * 1024 type Service = edgeupdater.Service
type UpdateOptions = edgeupdater.UpdateOptions
var replaceAndRestartFunc = replaceAndRestart
type Service struct {
httpClient *http.Client
lastCheckKey string
}
func New() *Service { func New() *Service {
return &Service{ return edgeupdater.New(edgeupdater.Config{
httpClient: &http.Client{Timeout: 30 * time.Second}, LocalVersion: config.Version,
} AssetPrefix: "openflare-agent",
} LogLabel: "agent",
})
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)
}
+54
View File
@@ -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())
}
@@ -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))
}
}
@@ -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)
}
}
+23
View File
@@ -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)
}
}
}
+41
View File
@@ -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")
}
}
+132
View File
@@ -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)
}
+33
View File
@@ -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
}
}
+69
View File
@@ -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
}
+268
View File
@@ -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))
}
@@ -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)
}
}
+64
View File
@@ -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):
}
}
@@ -48,4 +48,4 @@ func removeBackupBinary(path string) error {
return err return err
} }
return nil return nil
} }
@@ -12,4 +12,4 @@ func TestRemoveBackupBinaryIgnoresMissingFile(t *testing.T) {
if err := removeBackupBinary(backupPath); err != nil { if err := removeBackupBinary(backupPath); err != nil {
t.Fatalf("expected missing backup cleanup to be ignored: %v", err) t.Fatalf("expected missing backup cleanup to be ignored: %v", err)
} }
} }
@@ -50,4 +50,4 @@ func buildWindowsCommandLine(execPath string, args []string) string {
func quoteWindowsArg(value string) string { func quoteWindowsArg(value string) string {
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
} }
+381
View File
@@ -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)
}
@@ -11,9 +11,6 @@ import (
"runtime" "runtime"
"strings" "strings"
"testing" "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) 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) 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) { func TestGetLatestPreviewRelease(t *testing.T) {
service := &Service{ service := testService(&http.Client{
httpClient: &http.Client{ Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
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" {
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())
t.Fatalf("unexpected request url: %s", req.URL.String()) }
} return &http.Response{
return &http.Response{ StatusCode: http.StatusOK,
StatusCode: http.StatusOK, Header: make(http.Header),
Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`[
Body: io.NopCloser(strings.NewReader(`[
{"tag_name":"v1.0.0","prerelease":false}, {"tag_name":"v1.0.0","prerelease":false},
{"tag_name":"v1.1.0-rc.1","prerelease":true} {"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 { if err != nil {
t.Fatalf("expected preview release query to succeed: %v", err) t.Fatalf("expected preview release query to succeed: %v", err)
} }
@@ -51,22 +55,20 @@ func TestGetLatestPreviewRelease(t *testing.T) {
} }
func TestGetReleaseByTag(t *testing.T) { func TestGetReleaseByTag(t *testing.T) {
service := &Service{ service := testService(&http.Client{
httpClient: &http.Client{ Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
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" {
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())
t.Fatalf("unexpected request url: %s", req.URL.String()) }
} return &http.Response{
return &http.Response{ StatusCode: http.StatusOK,
StatusCode: http.StatusOK, Header: make(http.Header),
Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)),
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)), }, nil
}, 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 { if err != nil {
t.Fatalf("expected tag release query to succeed: %v", err) t.Fatalf("expected tag release query to succeed: %v", err)
} }
@@ -76,34 +78,26 @@ func TestGetReleaseByTag(t *testing.T) {
} }
func TestCheckAndUpdateRequiresChecksumAsset(t *testing.T) { func TestCheckAndUpdateRequiresChecksumAsset(t *testing.T) {
originalVersion := config.Version assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
config.Version = "v1.0.0" service := testService(&http.Client{
t.Cleanup(func() { Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
config.Version = originalVersion if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" {
}) t.Fatalf("unexpected request url: %s", req.URL.String())
}
assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH) return &http.Response{
service := &Service{ StatusCode: http.StatusOK,
httpClient: &http.Client{ Header: make(http.Header),
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { Body: io.NopCloser(strings.NewReader(`{
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", "tag_name":"v1.0.1",
"assets":[ "assets":[
{"name":"` + assetName + `","browser_download_url":"https://example.test/agent"} {"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") { if err == nil || !strings.Contains(err.Error(), "no matching checksum asset") {
t.Fatalf("expected missing checksum asset error, got %v", err) t.Fatalf("expected missing checksum asset error, got %v", err)
} }
@@ -157,17 +151,15 @@ func TestDownloadAndRestartVerifiesChecksum(t *testing.T) {
replaceAndRestartFunc = originalReplace replaceAndRestartFunc = originalReplace
}) })
service := &Service{ service := testService(&http.Client{
httpClient: &http.Client{ Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { return &http.Response{
return &http.Response{ StatusCode: http.StatusOK,
StatusCode: http.StatusOK, Header: make(http.Header),
Header: make(http.Header), Body: io.NopCloser(strings.NewReader(string(payload))),
Body: io.NopCloser(strings.NewReader(string(payload))), }, nil
}, nil }),
}), })
},
}
if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil { if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil {
t.Fatalf("expected verified download to succeed: %v", err) t.Fatalf("expected verified download to succeed: %v", err)
@@ -198,17 +190,15 @@ func TestDownloadAndRestartRejectsChecksumMismatch(t *testing.T) {
replaceAndRestartFunc = originalReplace replaceAndRestartFunc = originalReplace
}) })
service := &Service{ service := testService(&http.Client{
httpClient: &http.Client{ Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { return &http.Response{
return &http.Response{ StatusCode: http.StatusOK,
StatusCode: http.StatusOK, Header: make(http.Header),
Header: make(http.Header), Body: io.NopCloser(strings.NewReader("tampered")),
Body: io.NopCloser(strings.NewReader("tampered")), }, nil
}, nil }),
}), })
},
}
err := service.downloadAndRestart(context.Background(), "https://example.test/agent", strings.Repeat("0", sha256.Size*2), targetPath) 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") { if err == nil || !strings.Contains(err.Error(), "sha256 checksum mismatch") {
@@ -239,4 +229,4 @@ func TestIsNewerSupportsPrerelease(t *testing.T) {
} }
}) })
} }
} }
+3 -30
View File
@@ -7,38 +7,11 @@ import (
"path/filepath" "path/filepath"
"strings" "strings"
"time" "time"
edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
) )
type MillisecondDuration time.Duration type MillisecondDuration = edgeconfig.MillisecondDuration
func (d *MillisecondDuration) UnmarshalJSON(b []byte) error {
var v interface{}
if err := json.Unmarshal(b, &v); err != nil {
return err
}
switch value := v.(type) {
case float64:
*d = MillisecondDuration(time.Duration(value) * time.Millisecond)
return nil
case string:
duration, err := time.ParseDuration(value)
if err != nil {
return err
}
*d = MillisecondDuration(duration)
return nil
default:
return errors.New("invalid duration format")
}
}
func (d MillisecondDuration) Duration() time.Duration {
return time.Duration(d)
}
func (d MillisecondDuration) String() string {
return time.Duration(d).String()
}
type Config struct { type Config struct {
ServerURL string `json:"server_url"` ServerURL string `json:"server_url"`
+15 -31
View File
@@ -3,8 +3,8 @@ package flared
import ( import (
"context" "context"
"log/slog" "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/config"
"github.com/Rain-kl/Wavelet/internal/apps/flared/frpc" "github.com/Rain-kl/Wavelet/internal/apps/flared/frpc"
"github.com/Rain-kl/Wavelet/internal/apps/flared/heartbeat" "github.com/Rain-kl/Wavelet/internal/apps/flared/heartbeat"
@@ -23,31 +23,17 @@ type Runner struct {
} }
func (r *Runner) Run(ctx context.Context) error { func (r *Runner) Run(ctx context.Context) error {
// Start background services
go r.HeartbeatService.Run(ctx) go r.HeartbeatService.Run(ctx)
go r.SyncService.Run(ctx) go r.SyncService.Run(ctx)
// WebSocket reconnection loop return edgerunner.RunWSReconnectLoop(ctx, edgerunner.WSReconnectConfig{
for { ComponentName: "flared",
select { OnShutdown: r.FrpcManager.Stop,
case <-ctx.Done(): }, func(ctx context.Context) (edgerunner.WSConnection, error) {
r.FrpcManager.Stop() return r.WebSocketService.Connect(ctx)
return ctx.Err() }, func(ctx context.Context, conn edgerunner.WSConnection) {
default:
}
conn, err := r.WebSocketService.Connect(ctx)
if err != nil {
slog.Error("flared ws connect failed, will retry", "error", err)
r.sleepContext(ctx, 5*time.Second)
continue
}
r.handleConnection(ctx, conn) r.handleConnection(ctx, conn)
_ = conn.Close() })
slog.Info("flared ws connection closed, reconnecting...")
r.sleepContext(ctx, 2*time.Second)
}
} }
type flaredWSHandler struct { type flaredWSHandler struct {
@@ -73,13 +59,11 @@ func (h *flaredWSHandler) OnClose(err error) {
slog.Error("flared ws receive failed", "error", err) slog.Error("flared ws receive failed", "error", err)
} }
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) { func (r *Runner) handleConnection(ctx context.Context, conn edgerunner.WSConnection) {
_ = conn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r}) wsConn, ok := conn.(*wsclient.Connection)
} if !ok {
slog.Error("flared ws connection has unexpected type")
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) { return
select {
case <-ctx.Done():
case <-time.After(d):
} }
} _ = wsConn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r})
}
+16 -103
View File
@@ -3,23 +3,16 @@ package heartbeat
import ( import (
"context" "context"
"log/slog" "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/config"
"github.com/Rain-kl/Wavelet/internal/apps/flared/frpc" "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/httpclient"
"github.com/Rain-kl/Wavelet/internal/apps/flared/updater" "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" service "github.com/Rain-kl/Wavelet/pkg/protocol"
) )
var (
lookupOutboundIP = geoip.GetOutboundIP
lookupLocalIP = detectLocalNodeIP
)
type Service struct { type Service struct {
client *httpclient.Client client *httpclient.Client
frpcManager *frpc.Manager 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) { func (s *Service) Run(ctx context.Context) {
ticker := time.NewTicker(s.config.HeartbeatInterval.Duration()) edgeheartbeat.RunLoop(ctx, s.config.HeartbeatInterval.Duration(), s.doHeartbeat)
defer ticker.Stop()
// initial heartbeat
s.doHeartbeat(ctx)
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
s.doHeartbeat(ctx)
}
}
} }
func (s *Service) doHeartbeat(ctx context.Context) { func (s *Service) doHeartbeat(ctx context.Context) {
slog.Debug("sending flared heartbeat") slog.Debug("sending flared heartbeat")
ip := detectNodeIP()
payload := service.FlaredHeartbeatPayload{ payload := service.FlaredHeartbeatPayload{
ClientVersion: config.Version, ClientVersion: config.Version,
FrpVersion: s.frpcManager.GetVersion(), FrpVersion: s.frpcManager.GetVersion(),
IP: ip, IP: nodeip.Detect(),
TunnelStatus: "running", // TODO implement proper status tracking TunnelStatus: "running",
ConnectedRelays: s.frpcManager.GetConnectedRelays(), ConnectedRelays: s.frpcManager.GetConnectedRelays(),
CurrentVersion: s.frpcManager.GetCurrentConfigVersion(), CurrentVersion: s.frpcManager.GetCurrentConfigVersion(),
CurrentChecksum: s.frpcManager.GetCurrentConfigChecksum(), CurrentChecksum: s.frpcManager.GetCurrentConfigChecksum(),
@@ -76,84 +54,19 @@ func (s *Service) doHeartbeat(ctx context.Context) {
slog.Debug("flared heartbeat succeeded") slog.Debug("flared heartbeat succeeded")
if resp != nil && resp.TunnelSettings != nil { 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) { func tunnelSettingsToAutoUpdate(settings *service.RelaySettings) *edgeheartbeat.AutoUpdateSettings {
if settings == nil || s.updater == nil { if settings == nil {
return return nil
} }
force := settings.UpdateNow return &edgeheartbeat.AutoUpdateSettings{
shouldCheck := settings.AutoUpdate || force AutoUpdate: settings.AutoUpdate,
if !shouldCheck || settings.UpdateRepo == "" { UpdateNow: settings.UpdateNow,
return 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
}
+11 -96
View File
@@ -1,16 +1,10 @@
package httpclient package httpclient
import ( import (
"bytes"
"context" "context"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"strings"
"time" "time"
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
service "github.com/Rain-kl/Wavelet/pkg/protocol" service "github.com/Rain-kl/Wavelet/pkg/protocol"
) )
@@ -20,27 +14,21 @@ type APIResponse[T any] struct {
} }
type Client struct { type Client struct {
baseURL string base *edgehttp.Client
token string
httpClient *http.Client
} }
func New(baseURL string, token string, timeout time.Duration) *Client { func New(baseURL string, token string, timeout time.Duration) *Client {
return &Client{ return &Client{
baseURL: strings.TrimRight(baseURL, "/"), base: edgehttp.New(baseURL, token, timeout, "X-Tunnel-Token"),
token: token,
httpClient: &http.Client{
Timeout: timeout,
},
} }
} }
func (c *Client) Heartbeat(ctx context.Context, payload service.FlaredHeartbeatPayload) (*service.FlaredHeartbeatResponse, error) { func (c *Client) Heartbeat(ctx context.Context, payload service.FlaredHeartbeatPayload) (*service.FlaredHeartbeatResponse, error) {
resp := APIResponse[service.FlaredHeartbeatResponse]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err return nil, err
} }
return &resp.Data, nil 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) { func (c *Client) GetActiveConfig(ctx context.Context) (*service.FlaredTunnelConfigResponse, error) {
resp := APIResponse[service.FlaredTunnelConfigResponse]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err return nil, err
} }
return &resp.Data, nil 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 { func (c *Client) ReportApplyLog(ctx context.Context, payload service.ApplyLogPayload) error {
resp := APIResponse[any]{} 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 err
} }
return apiError(resp.ErrorMsg) return edgehttp.APIError(resp.ErrorMsg)
} }
func (c *Client) SetToken(token string) { func (c *Client) SetToken(token string) {
c.token = strings.TrimSpace(token) c.base.SetToken(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)
}
@@ -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, `"`, `""`) + `"`
}
+9 -363
View File
@@ -1,371 +1,17 @@
package updater package updater
import ( import (
"context" edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"os"
"runtime"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/utils"
"github.com/Rain-kl/Wavelet/internal/apps/flared/config" "github.com/Rain-kl/Wavelet/internal/apps/flared/config"
) )
const maxChecksumAssetSize = 64 * 1024 type Service = edgeupdater.Service
type UpdateOptions = edgeupdater.UpdateOptions
var replaceAndRestartFunc = replaceAndRestart
type Service struct {
httpClient *http.Client
lastCheckKey string
}
func New() *Service { func New() *Service {
return &Service{ return edgeupdater.New(edgeupdater.Config{
httpClient: &http.Client{Timeout: 30 * time.Second}, LocalVersion: config.Version,
} AssetPrefix: "openflared",
} LogLabel: "flared",
})
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)
}
+4 -87
View File
@@ -1,49 +1,18 @@
package config package config
import ( import (
"context"
"encoding/json" "encoding/json"
"errors" "errors"
"net"
"os" "os"
"path/filepath" "path/filepath"
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/pkg/geoip" edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil" "github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip"
) )
type MillisecondDuration time.Duration type MillisecondDuration = edgeconfig.MillisecondDuration
func (d *MillisecondDuration) UnmarshalJSON(b []byte) error {
var v interface{}
if err := json.Unmarshal(b, &v); err != nil {
return err
}
switch value := v.(type) {
case float64:
*d = MillisecondDuration(time.Duration(value) * time.Millisecond)
return nil
case string:
duration, err := time.ParseDuration(value)
if err != nil {
return err
}
*d = MillisecondDuration(duration)
return nil
default:
return errors.New("invalid duration format")
}
}
func (d MillisecondDuration) Duration() time.Duration {
return time.Duration(d)
}
func (d MillisecondDuration) String() string {
return time.Duration(d).String()
}
type Config struct { type Config struct {
ServerURL string `json:"server_url"` ServerURL string `json:"server_url"`
@@ -130,7 +99,7 @@ func applyDefaults(cfg *Config, baseDir string) {
cfg.NodeName = strings.TrimSpace(host) cfg.NodeName = strings.TrimSpace(host)
} }
if cfg.NodeIP == "" { if cfg.NodeIP == "" {
cfg.NodeIP = detectNodeIP() cfg.NodeIP = nodeip.Detect()
} }
if cfg.StatePath == "" { if cfg.StatePath == "" {
cfg.StatePath = filepath.Join(cfg.DataDir, "relay-state.json") 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) 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
}
+12 -36
View File
@@ -3,13 +3,13 @@ package heartbeat
import ( import (
"context" "context"
"log/slog" "log/slog"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/relay/config" "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/frps"
"github.com/Rain-kl/Wavelet/internal/apps/relay/httpclient" "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/observability"
"github.com/Rain-kl/Wavelet/internal/apps/relay/state" "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" "github.com/Rain-kl/Wavelet/internal/apps/relay/updater"
service "github.com/Rain-kl/Wavelet/pkg/protocol" 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) { func (s *Service) Run(ctx context.Context) {
ticker := time.NewTicker(s.config.HeartbeatInterval.Duration()) edgeheartbeat.RunLoop(ctx, s.config.HeartbeatInterval.Duration(), s.doHeartbeat)
defer ticker.Stop()
// initial heartbeat
s.doHeartbeat(ctx)
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
s.doHeartbeat(ctx)
}
}
} }
func (s *Service) doHeartbeat(ctx context.Context) { func (s *Service) doHeartbeat(ctx context.Context) {
@@ -79,30 +66,19 @@ func (s *Service) doHeartbeat(ctx context.Context) {
s.frpsManager.UpdateConfig(resp.RelayConfig) s.frpsManager.UpdateConfig(resp.RelayConfig)
if resp != nil && resp.RelaySettings != nil { 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) { func relaySettingsToAutoUpdate(settings *service.RelaySettings) *edgeheartbeat.AutoUpdateSettings {
if settings == nil || s.updater == nil { if settings == nil {
return return nil
} }
force := settings.UpdateNow return &edgeheartbeat.AutoUpdateSettings{
shouldCheck := settings.AutoUpdate || force AutoUpdate: settings.AutoUpdate,
if !shouldCheck || settings.UpdateRepo == "" { UpdateNow: settings.UpdateNow,
return UpdateRepo: settings.UpdateRepo,
} UpdateChannel: settings.UpdateChannel,
channel := "stable" UpdateTag: settings.UpdateTag,
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)
} }
} }
+7 -92
View File
@@ -1,16 +1,10 @@
package httpclient package httpclient
import ( import (
"bytes"
"context" "context"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"strings"
"time" "time"
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
service "github.com/Rain-kl/Wavelet/pkg/protocol" service "github.com/Rain-kl/Wavelet/pkg/protocol"
) )
@@ -20,105 +14,26 @@ type APIResponse[T any] struct {
} }
type Client struct { type Client struct {
baseURL string base *edgehttp.Client
token string
httpClient *http.Client
} }
func New(baseURL string, token string, timeout time.Duration) *Client { func New(baseURL string, token string, timeout time.Duration) *Client {
return &Client{ return &Client{
baseURL: strings.TrimRight(baseURL, "/"), base: edgehttp.New(baseURL, token, timeout, "X-Agent-Token"),
token: token,
httpClient: &http.Client{
Timeout: timeout,
},
} }
} }
func (c *Client) Heartbeat(ctx context.Context, payload service.RelayHeartbeatPayload) (*service.RelayHeartbeatResponse, error) { func (c *Client) Heartbeat(ctx context.Context, payload service.RelayHeartbeatPayload) (*service.RelayHeartbeatResponse, error) {
resp := APIResponse[service.RelayHeartbeatResponse]{} 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 return nil, err
} }
if err := apiError(resp.ErrorMsg); err != nil { if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err return nil, err
} }
return &resp.Data, nil return &resp.Data, nil
} }
func (c *Client) SetToken(token string) { func (c *Client) SetToken(token string) {
c.token = strings.TrimSpace(token) c.base.SetToken(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)
}
+13 -225
View File
@@ -1,18 +1,15 @@
package observability package observability
import ( import (
"bufio"
"crypto/sha256" "crypto/sha256"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"os" "os"
"path/filepath"
"runtime" "runtime"
"strconv"
"strings" "strings"
"syscall"
"time" "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/config"
"github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps"
"github.com/Rain-kl/Wavelet/internal/apps/relay/state" "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() now := time.Now().UTC()
metric := &service.AgentNodeMetricSnapshot{CapturedAtUnix: now.Unix()} metric := &service.AgentNodeMetricSnapshot{CapturedAtUnix: now.Unix()}
metric.MemoryTotalBytes, metric.MemoryUsedBytes = readMemInfo() metric.MemoryTotalBytes, metric.MemoryUsedBytes = edgeobs.ReadMemInfo()
metric.StorageTotalBytes, metric.StorageUsedBytes = statFilesystem(cfg.DataDir) metric.StorageTotalBytes, metric.StorageUsedBytes = edgeobs.StatFilesystem(cfg.DataDir)
metric.NetworkRxBytes, metric.NetworkTxBytes = readLinuxNetworkTotals() metric.NetworkRxBytes, metric.NetworkTxBytes = edgeobs.ReadLinuxNetworkTotals()
metric.DiskReadBytes, metric.DiskWriteBytes = readLinuxDiskTotals() metric.DiskReadBytes, metric.DiskWriteBytes = edgeobs.ReadLinuxDiskTotals()
if stateStore == nil { if stateStore == nil {
return metric return metric
} }
totalCPU, idleCPU := readLinuxCPUStat() totalCPU, idleCPU := edgeobs.ReadLinuxCPUStat()
snapshot, err := stateStore.Load() snapshot, err := stateStore.Load()
if err != nil { if err != nil {
return metric return metric
@@ -88,20 +85,20 @@ func BuildHealthEvents(status frps.RuntimeStatus) []service.AgentNodeHealthEvent
func collectProfile(cfg *config.Config) *service.AgentNodeSystemProfile { func collectProfile(cfg *config.Config) *service.AgentNodeSystemProfile {
hostname, _ := os.Hostname() hostname, _ := os.Hostname()
osName, osVersion := readLinuxOSRelease() osName, osVersion := edgeobs.ReadLinuxOSRelease()
totalMemory, _ := readMemInfo() totalMemory, _ := edgeobs.ReadMemInfo()
totalDisk, _ := statFilesystem(cfg.DataDir) totalDisk, _ := edgeobs.StatFilesystem(cfg.DataDir)
return &service.AgentNodeSystemProfile{ return &service.AgentNodeSystemProfile{
Hostname: strings.TrimSpace(hostname), Hostname: strings.TrimSpace(hostname),
OSName: osName, OSName: osName,
OSVersion: osVersion, OSVersion: osVersion,
KernelVersion: readFirstLine("/proc/sys/kernel/osrelease"), KernelVersion: edgeobs.ReadFirstLine("/proc/sys/kernel/osrelease"),
Architecture: runtime.GOARCH, Architecture: runtime.GOARCH,
CPUModel: readLinuxCPUModel(), CPUModel: edgeobs.ReadLinuxCPUModel(),
CPUCores: runtime.NumCPU(), CPUCores: runtime.NumCPU(),
TotalMemoryBytes: totalMemory, TotalMemoryBytes: totalMemory,
TotalDiskBytes: totalDisk, TotalDiskBytes: totalDisk,
UptimeSeconds: readLinuxUptimeSeconds(), UptimeSeconds: edgeobs.ReadLinuxUptimeSeconds(),
ReportedAtUnix: time.Now().UTC().Unix(), ReportedAtUnix: time.Now().UTC().Unix(),
} }
} }
@@ -113,213 +110,4 @@ func fingerprintProfile(profile *service.AgentNodeSystemProfile) string {
} }
sum := sha256.Sum256(raw) sum := sha256.Sum256(raw)
return hex.EncodeToString(sum[:]) 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))
}
+15 -31
View File
@@ -4,8 +4,8 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"log/slog" "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/config"
"github.com/Rain-kl/Wavelet/internal/apps/relay/frps" "github.com/Rain-kl/Wavelet/internal/apps/relay/frps"
"github.com/Rain-kl/Wavelet/internal/apps/relay/heartbeat" "github.com/Rain-kl/Wavelet/internal/apps/relay/heartbeat"
@@ -25,30 +25,16 @@ type Runner struct {
} }
func (r *Runner) Run(ctx context.Context) error { func (r *Runner) Run(ctx context.Context) error {
// Start heartbeat loop in background
go r.HeartbeatService.Run(ctx) go r.HeartbeatService.Run(ctx)
// WebSocket reconnection loop return edgerunner.RunWSReconnectLoop(ctx, edgerunner.WSReconnectConfig{
for { ComponentName: "relay",
select { OnShutdown: r.FrpsManager.Stop,
case <-ctx.Done(): }, func(ctx context.Context) (edgerunner.WSConnection, error) {
r.FrpsManager.Stop() return r.WebSocketService.Connect(ctx)
return ctx.Err() }, func(ctx context.Context, conn edgerunner.WSConnection) {
default:
}
conn, err := r.WebSocketService.Connect(ctx)
if err != nil {
slog.Error("relay ws connect failed, will retry", "error", err)
r.sleepContext(ctx, 5*time.Second)
continue
}
r.handleConnection(ctx, conn) r.handleConnection(ctx, conn)
_ = conn.Close() })
slog.Info("relay ws connection closed, reconnecting...")
r.sleepContext(ctx, 2*time.Second)
}
} }
type relayWSHandler struct { type relayWSHandler struct {
@@ -78,13 +64,11 @@ func (h *relayWSHandler) OnClose(err error) {
slog.Error("relay ws receive failed", "error", err) slog.Error("relay ws receive failed", "error", err)
} }
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) { func (r *Runner) handleConnection(ctx context.Context, conn edgerunner.WSConnection) {
_ = conn.RunReceiveLoop(ctx, &relayWSHandler{runner: r}) wsConn, ok := conn.(*wsclient.Connection)
} if !ok {
slog.Error("relay ws connection has unexpected type")
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) { return
select {
case <-ctx.Done():
case <-time.After(d):
} }
} _ = wsConn.RunReceiveLoop(ctx, &relayWSHandler{runner: r})
}
@@ -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
}
+9 -363
View File
@@ -1,371 +1,17 @@
package updater package updater
import ( import (
"context" edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
"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/relay/config" "github.com/Rain-kl/Wavelet/internal/apps/relay/config"
) )
const maxChecksumAssetSize = 64 * 1024 type Service = edgeupdater.Service
type UpdateOptions = edgeupdater.UpdateOptions
var replaceAndRestartFunc = replaceAndRestart
type Service struct {
httpClient *http.Client
lastCheckKey string
}
func New() *Service { func New() *Service {
return &Service{ return edgeupdater.New(edgeupdater.Config{
httpClient: &http.Client{Timeout: 30 * time.Second}, LocalVersion: config.Version,
} AssetPrefix: "openflare-relay",
} LogLabel: "relay",
})
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)
}
@@ -159,17 +159,6 @@ type RegisterNodeResponse struct {
Name string `json:"name"` 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 { type ActiveConfigResponse struct {
Version string `json:"version"` Version string `json:"version"`
Checksum string `json:"checksum"` Checksum string `json:"checksum"`
@@ -178,11 +167,6 @@ type ActiveConfigResponse struct {
CreatedAt string `json:"created_at"` CreatedAt string `json:"created_at"`
} }
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
type WAFIPGroup struct { type WAFIPGroup struct {
ID uint `json:"id"` ID uint `json:"id"`
Name string `json:"name"` Name string `json:"name"`
@@ -204,4 +188,4 @@ type WAFIPGroupSyncResponse struct {
type SupportFile struct { type SupportFile struct {
Path string `json:"path"` Path string `json:"path"`
Content string `json:"content"` Content string `json:"content"`
} }
+95
View File
@@ -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)
}
}
-5
View File
@@ -1,10 +1,5 @@
package protocol package protocol
type WSMessage struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}
type AgentNodeSystemProfile struct { type AgentNodeSystemProfile struct {
Hostname string `json:"hostname"` Hostname string `json:"hostname"`
OSName string `json:"os_name"` OSName string `json:"os_name"`