refactor(backend): rename OpenFlare directory to lowercase openflare

This commit is contained in:
ryan
2026-08-30 17:43:23 +08:00
parent 06d5fedbfc
commit c93ff6674f
543 changed files with 819 additions and 819 deletions
+26
View File
@@ -0,0 +1,26 @@
# share — 跨插件共享资源层
本目录存放被 **两个及以上插件共同消费**、且无处安放的可复用资源:
- 不能放 `Wavelet/pkg/`:该层是上游通用库,禁止包含业务语义。
- 不能放进某个插件内部:Cordis 规则禁止插件之间互相 import 实现包。
因此 `share/` 是内核(`core/`)之外唯一允许被任意插件直接 import 的共享层。
本目录内的包**禁止** import `Wavelet/openflare/...`(下游业务)与 `Wavelet/plugins/...`(具体插件实现),
只允许 import 标准库、第三方库与 `Wavelet/core`、`Wavelet/pkg`。
## 当前内容
| 包 | 内容 | 消费方 |
| :-- | :-- | :-- |
| `share/protocol` | server 与边缘三进制的控制消息线格式 | `server`、`agent`、`relay`、`flared` |
| `share/geoip` | GeoIP 解析与 IP 工具(含 `iputil` 子包) | `server`、`agent` |
| `share/edge/logging` | 边缘守护进程统一日志初始化 | `agent`、`relay`、`flared` |
## 所有权与上游同步
**本目录由 OpenFlare 拥有,不属于上游 merge 范围**:`git merge wavelet/main` 吸收
`backend/{core,pkg,plugins}` 等上游变更,从不覆盖 `share/` 与 `openflare/`。
计划迁入的候选(随插件化进度):`wsclient`(需先合并 `pkg/util` 的仅存符号)、
`render`(openresty 配置渲染)、`pagesarchive`(Pages 产物解包)。
@@ -0,0 +1,63 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package config provides shared configuration types for edge applications.
package config
import (
"encoding/json"
"fmt"
"strconv"
"strings"
"time"
)
// MillisecondDuration is a time.Duration that marshals to/from JSON as an integer number of milliseconds
// or as a Go duration string (e.g. "1s", "500ms").
type MillisecondDuration time.Duration
// Duration returns the underlying time.Duration value.
func (d MillisecondDuration) Duration() time.Duration {
return time.Duration(d)
}
func (d MillisecondDuration) String() string {
return time.Duration(d).String()
}
// UnmarshalJSON decodes either a numeric millisecond value or a quoted Go duration string.
func (d *MillisecondDuration) UnmarshalJSON(data []byte) error {
raw := strings.TrimSpace(string(data))
if raw == "" || raw == "null" {
*d = 0
return nil
}
if strings.HasPrefix(raw, "\"") {
var text string
if err := json.Unmarshal(data, &text); err != nil {
return err
}
text = strings.TrimSpace(text)
if text == "" {
*d = 0
return nil
}
parsed, err := time.ParseDuration(text)
if err != nil {
return fmt.Errorf("invalid duration string %q: %w", text, err)
}
*d = MillisecondDuration(parsed)
return nil
}
ms, err := strconv.ParseInt(raw, 10, 64)
if err != nil {
return fmt.Errorf("invalid duration milliseconds %q: %w", raw, err)
}
*d = MillisecondDuration(time.Duration(ms) * time.Millisecond)
return nil
}
// MarshalJSON encodes the duration as an integer number of milliseconds.
func (d MillisecondDuration) MarshalJSON() ([]byte, error) {
return json.Marshal(time.Duration(d).Milliseconds())
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config
import (
"encoding/json"
"testing"
"time"
)
func TestMillisecondDurationUnmarshalJSON(t *testing.T) {
tests := []struct {
name string
input string
want time.Duration
wantErr bool
}{
{name: "null", input: "null", want: 0},
{name: "empty", input: `""`, want: 0},
{name: "integer milliseconds", input: "30000", want: 30 * time.Second},
{name: "duration string", input: `"5s"`, want: 5 * time.Second},
{name: "invalid number", input: "not-a-number", wantErr: true},
{name: "invalid duration string", input: `"not-a-duration"`, wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var d MillisecondDuration
err := json.Unmarshal([]byte(tt.input), &d)
if tt.wantErr {
if err == nil {
t.Fatal("expected error")
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if d.Duration() != tt.want {
t.Fatalf("got %s, want %s", d, tt.want)
}
})
}
}
func TestMillisecondDurationMarshalJSON(t *testing.T) {
d := MillisecondDuration(7 * time.Second)
data, err := json.Marshal(d)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(data) != "7000" {
t.Fatalf("unexpected marshaled value: %s", string(data))
}
}
@@ -0,0 +1,47 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package heartbeat handles periodic heartbeat and update checks.
package heartbeat
import (
"context"
"log/slog"
"strings"
edgeupdater "Wavelet/openflare/share/edge/updater"
)
// AutoUpdateSettings defines the settings for automatic edge updates.
type AutoUpdateSettings struct {
AutoUpdate bool
UpdateNow bool
UpdateRepo string
UpdateChannel string
UpdateTag string
}
// TryAutoUpdate attempts to check and apply auto updates for the edge service.
func TryAutoUpdate(ctx context.Context, updater *edgeupdater.Service, settings *AutoUpdateSettings, logLabel string) {
if settings == nil || updater == nil {
return
}
force := settings.UpdateNow
shouldCheck := settings.AutoUpdate || force
if !shouldCheck || strings.TrimSpace(settings.UpdateRepo) == "" {
return
}
channel := "stable"
if force && strings.TrimSpace(settings.UpdateChannel) != "" {
channel = settings.UpdateChannel
}
slog.Info("checking for "+logLabel+" updates", "repo", settings.UpdateRepo, "channel", channel, "force", force)
err := updater.CheckAndUpdate(ctx, settings.UpdateRepo, edgeupdater.UpdateOptions{
Channel: channel,
TagName: settings.UpdateTag,
Force: force,
})
if err != nil {
slog.Error(logLabel+" update check failed", "error", err)
}
}
@@ -0,0 +1,26 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package heartbeat
import (
"context"
"time"
)
// RunLoop invokes fn immediately, then on each interval tick until ctx is cancelled.
func RunLoop(ctx context.Context, interval time.Duration, fn func(context.Context)) {
fn(ctx)
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
fn(ctx)
}
}
}
@@ -0,0 +1,44 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package heartbeat
import (
"context"
"sync/atomic"
"testing"
"time"
)
func TestRunLoopImmediateAndTicker(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var calls atomic.Int32
interval := 20 * time.Millisecond
done := make(chan struct{})
go func() {
RunLoop(ctx, interval, func(context.Context) {
calls.Add(1)
})
close(done)
}()
time.Sleep(5 * time.Millisecond)
if got := calls.Load(); got != 1 {
t.Fatalf("expected 1 immediate call, got %d", got)
}
time.Sleep(35 * time.Millisecond)
if got := calls.Load(); got < 2 {
t.Fatalf("expected at least 2 calls after ticker, got %d", got)
}
cancel()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("RunLoop did not exit after context cancellation")
}
}
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package httpclient provides an authenticated HTTP client for edge services.
package httpclient
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"strings"
"time"
)
// Client is an HTTP client wrapper for communicating with remote HTTP services.
type Client struct {
baseURL string
token string
authHeader string
httpClient *http.Client
}
// New creates a new Client instance.
func New(baseURL, token string, timeout time.Duration, authHeader string) *Client {
return &Client{
baseURL: strings.TrimRight(baseURL, "/"),
token: token,
authHeader: authHeader,
httpClient: &http.Client{Timeout: timeout},
}
}
// SetToken updates the client auth token.
func (c *Client) SetToken(token string) {
c.token = strings.TrimSpace(token)
slog.Debug("http client token updated")
}
// GetJSON sends a GET request and decodes the response body into target.
func (c *Client) GetJSON(ctx context.Context, path string, target any) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil)
if err != nil {
return err
}
c.setAuthHeader(req)
return c.do(req, target)
}
// PostJSON sends a POST request with JSON body and decodes the response body into target.
func (c *Client) PostJSON(ctx context.Context, path string, body any, target any) error {
data, err := json.Marshal(body)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+path, bytes.NewReader(data))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
c.setAuthHeader(req)
return c.do(req, target)
}
// DoRaw performs an HTTP request with custom headers and returns the raw response.
func (c *Client) DoRaw(ctx context.Context, method, path string, headers map[string]string) (*http.Response, error) {
req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, nil)
if err != nil {
return nil, err
}
c.setAuthHeader(req)
for key, value := range headers {
req.Header.Set(key, value)
}
return c.httpClient.Do(req)
}
func (c *Client) setAuthHeader(req *http.Request) {
if c.authHeader != "" {
req.Header.Set(c.authHeader, c.token)
}
}
func (c *Client) do(req *http.Request, target any) error {
res, err := c.httpClient.Do(req)
if err != nil {
slog.Error("http request failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
defer func() {
if err := res.Body.Close(); err != nil {
slog.Error("failed to close response body", "error", err)
}
}()
body, err := io.ReadAll(res.Body)
if err != nil {
slog.Error("http response read failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
if res.StatusCode != http.StatusOK {
slog.Warn("http request returned non-200", "method", req.Method, "path", req.URL.Path, "status", res.Status)
return ReadBodyError(body, res.Status)
}
if target == nil {
return nil
}
if err = json.Unmarshal(body, target); err != nil {
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
return nil
}
// APIError creates a new API error with the given message if it is not empty.
func APIError(msg string) error {
if strings.TrimSpace(msg) == "" {
return nil
}
return errors.New(msg)
}
// ReadBodyError parses the error message from the response body, or returns the fallback message.
func ReadBodyError(body []byte, fallback string) error {
var errBody struct {
ErrorMsg string `json:"error_msg"`
}
if err := json.Unmarshal(body, &errBody); err == nil && strings.TrimSpace(errBody.ErrorMsg) != "" {
return errors.New(errBody.ErrorMsg)
}
return errors.New(fallback)
}
// ReadHTTPError reads the error message from the HTTP response.
func ReadHTTPError(res *http.Response) error {
body, err := io.ReadAll(res.Body)
if err != nil {
return errors.New(res.Status)
}
return ReadBodyError(body, res.Status)
}
@@ -0,0 +1,40 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package logging configures structured logging for edge applications.
package logging
import (
"log/slog"
"os"
"strings"
)
// Options holds configuration options for the structured logger.
type Options struct {
AddSource bool
}
// Setup initialises the default slog handler using the given options and the LOG_LEVEL environment variable.
func Setup(opts Options) {
handlerOpts := &slog.HandlerOptions{
AddSource: opts.AddSource,
Level: ParseLevel(os.Getenv("LOG_LEVEL")),
}
handler := slog.NewTextHandler(os.Stdout, handlerOpts)
slog.SetDefault(slog.New(handler))
}
// ParseLevel converts a log-level string (e.g. "debug", "warn") to the corresponding slog.Level.
func ParseLevel(value string) slog.Level {
switch strings.ToLower(strings.TrimSpace(value)) {
case "debug":
return slog.LevelDebug
case "warn", "warning":
return slog.LevelWarn
case "error":
return slog.LevelError
default:
return slog.LevelInfo
}
}
@@ -0,0 +1,117 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package nodeip detects the preferred public IP address for edge nodes.
package nodeip
import (
"context"
"net"
"sync"
"time"
"Wavelet/openflare/share/geoip"
"Wavelet/openflare/share/geoip/iputil"
)
const (
outboundIPLookupTimeout = 5 * time.Second
publicIPPriorityScore = 2 // matches iputil.Score for public IPv4 addresses
ipCacheTTL = 10 * time.Minute
)
// LookupOutboundIP and LookupLocalIP are the provider functions used to detect the node's outbound/local IP.
// They are package-level variables so they can be overridden in tests.
var (
LookupOutboundIP = geoip.GetOutboundIP
LookupLocalIP = DetectLocal
cacheMu sync.RWMutex
cachedIP string
lastDetected time.Time
)
// Detect returns the best available outbound or local IPv4 address for this node.
func Detect() string {
return DetectWithContext(context.Background())
}
// DetectWithContext returns the best available outbound or local IPv4 address, respecting ctx for cancellation.
func DetectWithContext(ctx context.Context) string {
cacheMu.RLock()
if cachedIP != "" && time.Since(lastDetected) < ipCacheTTL {
ip := cachedIP
cacheMu.RUnlock()
return ip
}
cacheMu.RUnlock()
var ip string
if ip = detectOutbound(ctx); ip == "" {
ip = LookupLocalIP()
}
if ip != "" {
cacheMu.Lock()
cachedIP = ip
lastDetected = time.Now()
cacheMu.Unlock()
}
return ip
}
func detectOutbound(ctx context.Context) string {
ctx, cancel := context.WithTimeout(ctx, outboundIPLookupTimeout)
defer cancel()
ip, err := LookupOutboundIP(ctx)
if err != nil || ip == nil {
return ""
}
return ip.String()
}
// DetectLocal returns the highest-priority non-loopback local IPv4 address found on system interfaces.
func DetectLocal() string {
interfaces, err := net.Interfaces()
if err != nil {
return ""
}
bestIP := ""
bestPriority := -1
for _, iface := range interfaces {
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
continue
}
addrs, err := iface.Addrs()
if err != nil {
continue
}
for _, addr := range addrs {
ipNet, ok := addr.(*net.IPNet)
if !ok || ipNet.IP == nil || ipNet.IP.IsLoopback() {
continue
}
ipv4 := ipNet.IP.To4()
if ipv4 == nil {
continue
}
priority := iputil.Score(ipv4)
if priority > bestPriority {
bestIP = ipv4.String()
bestPriority = priority
}
if bestPriority == publicIPPriorityScore {
return bestIP
}
}
}
return bestIP
}
// ResetCacheForTest clears the cached IP.
func ResetCacheForTest() {
cacheMu.Lock()
cachedIP = ""
lastDetected = time.Time{}
cacheMu.Unlock()
}
@@ -0,0 +1,288 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package observability provides helpers that read Linux /proc and /sys metrics for system monitoring.
package observability
import (
"bufio"
"math"
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"syscall"
)
const (
memInfoMinFieldCount = 2
cpuStatMinFieldCount = 5
netDevMinFieldCount = 16
diskStatsMinFieldCount = 14
)
// ReadLinuxOSRelease returns the OS name and version from /etc/os-release.
func ReadLinuxOSRelease() (string, string) {
file, err := os.Open("/etc/os-release")
if err != nil {
return runtime.GOOS, ""
}
defer func() { _ = file.Close() }()
values := make(map[string]string)
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
key, value, ok := strings.Cut(line, "=")
if !ok {
continue
}
values[key] = strings.Trim(value, `"`)
}
if pretty := strings.TrimSpace(values["PRETTY_NAME"]); pretty != "" {
return pretty, strings.TrimSpace(values["VERSION_ID"])
}
name := strings.TrimSpace(values["NAME"])
if name == "" {
name = runtime.GOOS
}
return name, strings.TrimSpace(values["VERSION_ID"])
}
// ReadLinuxCPUModel returns the CPU model name from /proc/cpuinfo.
func ReadLinuxCPUModel() string {
file, err := os.Open("/proc/cpuinfo")
if err != nil {
return ""
}
defer func() { _ = file.Close() }()
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(strings.ToLower(line), "model name") {
_, value, ok := strings.Cut(line, ":")
if ok {
return strings.TrimSpace(value)
}
}
}
return ""
}
// ReadMemInfo returns total and used memory bytes from /proc/meminfo.
func ReadMemInfo() (int64, int64) {
file, err := os.Open("/proc/meminfo")
if err != nil {
return 0, 0
}
defer func() { _ = file.Close() }()
var memTotalKB int64
var memAvailableKB int64
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
switch {
case strings.HasPrefix(line, "MemTotal:"):
memTotalKB = parseMemInfoValue(line)
case strings.HasPrefix(line, "MemAvailable:"):
memAvailableKB = parseMemInfoValue(line)
}
}
total := memTotalKB * 1024
if total == 0 {
return 0, 0
}
used := max(total-(memAvailableKB*1024), 0)
return total, used
}
func parseMemInfoValue(line string) int64 {
fields := strings.Fields(line)
if len(fields) < memInfoMinFieldCount {
return 0
}
value, err := strconv.ParseInt(fields[1], 10, 64)
if err != nil {
return 0
}
return value
}
// ReadLinuxUptimeSeconds returns system uptime in seconds from /proc/uptime.
func ReadLinuxUptimeSeconds() int64 {
content, err := os.ReadFile("/proc/uptime")
if err != nil {
return 0
}
fields := strings.Fields(string(content))
if len(fields) == 0 {
return 0
}
value, err := strconv.ParseFloat(fields[0], 64)
if err != nil {
return 0
}
return int64(value)
}
// ReadLinuxCPUStat returns aggregate CPU jiffies and idle jiffies from /proc/stat.
func ReadLinuxCPUStat() (uint64, uint64) {
content, err := os.ReadFile("/proc/stat")
if err != nil {
return 0, 0
}
lines := strings.SplitSeq(string(content), "\n")
for line := range lines {
if !strings.HasPrefix(line, "cpu ") {
continue
}
fields := strings.Fields(line)
if len(fields) < cpuStatMinFieldCount {
return 0, 0
}
var total uint64
for i := 1; i < len(fields); i++ {
value, err := strconv.ParseUint(fields[i], 10, 64)
if err != nil {
return 0, 0
}
total += value
}
idle, err := strconv.ParseUint(fields[4], 10, 64)
if err != nil {
return 0, 0
}
return total, idle
}
return 0, 0
}
// ReadLinuxNetworkTotals returns aggregate RX and TX byte totals from /proc/net/dev.
func ReadLinuxNetworkTotals() (int64, int64) {
file, err := os.Open("/proc/net/dev")
if err != nil {
return 0, 0
}
defer func() { _ = file.Close() }()
var rx int64
var tx int64
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if !strings.Contains(line, ":") {
continue
}
name, data, ok := strings.Cut(line, ":")
if !ok {
continue
}
if strings.TrimSpace(name) == "lo" {
continue
}
fields := strings.Fields(data)
if len(fields) < netDevMinFieldCount {
continue
}
rxValue, err := strconv.ParseInt(fields[0], 10, 64)
if err == nil {
rx += rxValue
}
txValue, err := strconv.ParseInt(fields[8], 10, 64)
if err == nil {
tx += txValue
}
}
return rx, tx
}
// ReadLinuxDiskTotals returns aggregate disk read and write byte totals from /proc/diskstats.
func ReadLinuxDiskTotals() (int64, int64) {
file, err := os.Open("/proc/diskstats")
if err != nil {
return 0, 0
}
defer func() { _ = file.Close() }()
var readBytes int64
var writeBytes int64
scanner := bufio.NewScanner(file)
for scanner.Scan() {
fields := strings.Fields(scanner.Text())
if len(fields) < diskStatsMinFieldCount {
continue
}
device := fields[2]
if shouldSkipDiskDevice(device) {
continue
}
readSectors, err := strconv.ParseInt(fields[5], 10, 64)
if err == nil {
readBytes += readSectors * 512
}
writeSectors, err := strconv.ParseInt(fields[9], 10, 64)
if err == nil {
writeBytes += writeSectors * 512
}
}
return readBytes, writeBytes
}
func shouldSkipDiskDevice(device string) bool {
switch {
case device == "":
return true
case strings.HasPrefix(device, "loop"),
strings.HasPrefix(device, "ram"),
strings.HasPrefix(device, "dm-"):
return true
default:
return false
}
}
// StatFilesystem returns total and used bytes for the filesystem containing path.
func StatFilesystem(path string) (int64, int64) {
if strings.TrimSpace(path) == "" {
path = string(os.PathSeparator)
}
absPath := filepath.Clean(path)
var stat syscall.Statfs_t
if err := syscall.Statfs(absPath, &stat); err != nil {
return 0, 0
}
bsize := int64(stat.Bsize) //nolint:unconvert // Statfs_t.Bsize is int64 on Linux and uint32 on Darwin
total := multiplyUint64Int64(stat.Blocks, bsize)
free := multiplyUint64Int64(stat.Bavail, bsize)
used := max(total-free, 0)
return total, used
}
// multiplyUint64Int64 multiplies a uint64 by a positive int64, saturating at
// math.MaxInt64 to avoid int64 overflow.
func multiplyUint64Int64(a uint64, b int64) int64 {
if a == 0 || b <= 0 {
return 0
}
v := a * uint64(b)
if v > math.MaxInt64 {
return math.MaxInt64
}
return int64(v)
}
// ReadFirstLine reads and returns the trimmed first line of a file.
func ReadFirstLine(path string) string {
content, err := os.ReadFile(path) //nolint:gosec // path is a fixed /proc or /sys path from internal callers, not user input
if err != nil {
return ""
}
return strings.TrimSpace(string(content))
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package observability
import (
"os"
"path/filepath"
"testing"
)
func TestParseMemInfoValue(t *testing.T) {
t.Parallel()
tests := []struct {
line string
want int64
}{
{line: "MemTotal: 16384000 kB", want: 16384000},
{line: "MemAvailable: 8192000 kB", want: 8192000},
{line: "invalid", want: 0},
{line: "MemTotal: not-a-number kB", want: 0},
}
for _, tt := range tests {
if got := parseMemInfoValue(tt.line); got != tt.want {
t.Fatalf("parseMemInfoValue(%q) = %d, want %d", tt.line, got, tt.want)
}
}
}
func TestShouldSkipDiskDevice(t *testing.T) {
t.Parallel()
tests := []struct {
device string
want bool
}{
{device: "", want: true},
{device: "loop0", want: true},
{device: "ram0", want: true},
{device: "dm-0", want: true},
{device: "sda", want: false},
{device: "nvme0n1", want: false},
}
for _, tt := range tests {
if got := shouldSkipDiskDevice(tt.device); got != tt.want {
t.Fatalf("shouldSkipDiskDevice(%q) = %v, want %v", tt.device, got, tt.want)
}
}
}
func TestReadFirstLine(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "sample.txt")
if err := os.WriteFile(path, []byte(" first line\nsecond line\n"), 0o644); err != nil {
t.Fatalf("write file: %v", err)
}
if got := ReadFirstLine(path); got != "first line\nsecond line" {
t.Fatalf("ReadFirstLine() = %q, want trimmed first line content", got)
}
if got := ReadFirstLine(filepath.Join(dir, "missing.txt")); got != "" {
t.Fatalf("ReadFirstLine(missing) = %q, want empty string", got)
}
}
func TestStatFilesystem(t *testing.T) {
t.Parallel()
total, used := StatFilesystem(t.TempDir())
if total <= 0 {
t.Fatalf("StatFilesystem() total = %d, want > 0", total)
}
if used < 0 || used > total {
t.Fatalf("StatFilesystem() used = %d, total = %d", used, total)
}
}
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package runner provides shared WebSocket reconnect helpers for edge daemons.
package runner
import (
"context"
"log/slog"
"time"
)
// WSConnection defines the minimum interface required for a WebSocket connection that can be closed.
type WSConnection interface {
Close() error
}
// WSReconnectConfig specifies configuration parameters for the WebSocket reconnect loop.
type WSReconnectConfig struct {
ComponentName string
ConnectBackoff time.Duration
ReconnectDelay time.Duration
OnShutdown func()
}
// RunWSReconnectLoop runs a loop that attempts to keep a WebSocket connection active, automatically reconnecting when closed or failed.
func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
connect func(context.Context) (WSConnection, error),
handle func(context.Context, WSConnection),
) error {
if cfg.ConnectBackoff <= 0 {
cfg.ConnectBackoff = 5 * time.Second
}
if cfg.ReconnectDelay <= 0 {
cfg.ReconnectDelay = 2 * time.Second
}
label := cfg.ComponentName
if label == "" {
label = "edge"
}
for {
select {
case <-ctx.Done():
if cfg.OnShutdown != nil {
cfg.OnShutdown()
}
return ctx.Err()
default:
// Continue reconnect loop
}
conn, err := connect(ctx)
if err != nil {
slog.Error(label+" ws connect failed, will retry", "error", err)
SleepContext(ctx, cfg.ConnectBackoff)
continue
}
handle(ctx, conn)
_ = conn.Close()
slog.Info(label + " ws connection closed, reconnecting...")
SleepContext(ctx, cfg.ReconnectDelay)
}
}
// SleepContext pauses execution for the given duration or until the context is canceled.
func SleepContext(ctx context.Context, d time.Duration) {
t := time.NewTimer(d)
defer t.Stop()
select {
case <-ctx.Done():
case <-t.C:
}
}
@@ -0,0 +1,56 @@
//go:build !windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package updater provides capabilities to check for, download, and apply updates.
package updater
import (
"errors"
"fmt"
"log/slog"
"os"
"syscall"
)
func replaceAndRestart(execPath string, tmpPath string) error {
backupPath := execPath + ".bak"
if err := removeBackupBinary(backupPath); err != nil {
return err
}
if err := os.Rename(execPath, backupPath); err != nil {
renameErr := err
if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) {
slog.Error("remove tmp binary failed", "path", tmpPath, "error", err)
return fmt.Errorf("backup current binary: %w; remove tmp binary: %w", renameErr, err)
}
return fmt.Errorf("backup current binary: %w", renameErr)
}
if err := os.Rename(tmpPath, execPath); err != nil {
replaceErr := err
if err := os.Rename(backupPath, execPath); err != nil {
slog.Error("restore backup binary failed", "path", backupPath, "error", err)
return fmt.Errorf("replace binary: %w; restore backup binary: %w", replaceErr, err)
}
return fmt.Errorf("replace binary: %w", replaceErr)
}
if err := removeBackupBinary(backupPath); err != nil {
return err
}
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { //nolint:gosec // execPath is the validated edge updater binary path
return fmt.Errorf("exec restart: %w", err)
}
return errors.New("unreachable after exec")
}
func removeBackupBinary(path string) error {
if err := os.Remove(path); err != nil {
if os.IsNotExist(err) {
return nil
}
slog.Error("remove backup binary failed", "path", path, "error", err)
return err
}
return nil
}
@@ -0,0 +1,18 @@
//go:build !windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"path/filepath"
"testing"
)
func TestRemoveBackupBinaryIgnoresMissingFile(t *testing.T) {
backupPath := filepath.Join(t.TempDir(), "openflare-agent.bak")
if err := removeBackupBinary(backupPath); err != nil {
t.Fatalf("expected missing backup cleanup to be ignored: %v", err)
}
}
@@ -0,0 +1,56 @@
//go:build windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"fmt"
"os"
"os/exec"
"strings"
)
func replaceAndRestart(execPath string, tmpPath string) error {
backupPath := execPath + ".bak"
scriptPath := execPath + ".update.cmd"
script := fmt.Sprintf(`@echo off
setlocal
:waitloop
move /Y "%s" "%s" >nul 2>nul
if errorlevel 1 (
ping 127.0.0.1 -n 2 >nul
goto waitloop
)
move /Y "%s" "%s" >nul 2>nul
if errorlevel 1 exit /b 1
start "" %s
del /Q "%s" >nul 2>nul
del /Q "%%~f0" >nul 2>nul
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
os.Remove(tmpPath)
return fmt.Errorf("write restart script: %w", err)
}
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
if err := cmd.Start(); err != nil {
os.Remove(scriptPath)
os.Remove(tmpPath)
return fmt.Errorf("schedule restart: %w", err)
}
os.Exit(0)
return nil
}
func buildWindowsCommandLine(execPath string, args []string) string {
parts := []string{quoteWindowsArg(execPath)}
for _, arg := range args {
parts = append(parts, quoteWindowsArg(arg))
}
return strings.Join(parts, " ")
}
func quoteWindowsArg(value string) string {
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
}
@@ -0,0 +1,426 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package updater provides capabilities to check for, download, and apply updates.
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"os"
"runtime"
"strings"
"time"
"Wavelet/openflare/share/ofutil"
)
const (
maxChecksumAssetSize = 64 * 1024
goosWindows = "windows"
updateTmpFilePerm = 0o600
updateBinaryFilePerm = 0o755
)
var replaceAndRestartFunc = replaceAndRestart
// Config defines the configuration for the update service.
type Config struct {
LocalVersion string
AssetPrefix string
LogLabel string
}
// Service handles checking and applying application binary updates.
type Service struct {
httpClient *http.Client
lastCheckKey string
localVersion string
assetPrefix string
logLabel string
}
// New creates a new updater Service with the provided configuration.
func New(cfg Config) *Service {
return &Service{
httpClient: &http.Client{Timeout: 30 * time.Second},
localVersion: cfg.LocalVersion,
assetPrefix: cfg.AssetPrefix,
logLabel: cfg.LogLabel,
}
}
// UpdateOptions specifies parameters for checking and applying updates.
type UpdateOptions struct {
Channel string
TagName string
Force bool
}
type githubRelease struct {
TagName string `json:"tag_name"`
Prerelease bool `json:"prerelease"`
Draft bool `json:"draft"`
Assets []githubAsset `json:"assets"`
}
type githubAsset struct {
Name string `json:"name"`
BrowserDownloadURL string `json:"browser_download_url"`
Digest string `json:"digest"`
}
// CheckAndUpdate checks for a newer release on GitHub and performs an update if available.
func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error {
release, err := s.getRelease(ctx, repo, options)
if err != nil {
return fmt.Errorf("check latest release: %w", err)
}
if release == nil || release.TagName == "" {
return nil
}
remoteVersion := normalizeVersion(release.TagName)
localVersion := normalizeVersion(s.localVersion)
checkKey := buildReleaseCheckKey(options, remoteVersion)
if remoteVersion == localVersion {
return nil
}
if !options.Force && checkKey != "" && checkKey == s.lastCheckKey {
return nil
}
if !isNewer(localVersion, remoteVersion) {
s.lastCheckKey = checkKey
return nil
}
slog.Info(s.logLabel+" update available", "from", localVersion, "to", remoteVersion)
assetName := s.assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
downloadURL, expectedChecksum, err := s.resolveReleaseAsset(ctx, release, assetName)
if err != nil {
if downloadURL == "" {
s.lastCheckKey = checkKey
}
return err
}
execPath, err := os.Executable()
if err != nil {
return fmt.Errorf("get executable path: %w", err)
}
if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil {
return fmt.Errorf("download and restart: %w", err)
}
s.lastCheckKey = checkKey
return nil
}
func (s *Service) getRelease(ctx context.Context, repo string, options UpdateOptions) (*githubRelease, error) {
tagName := strings.TrimSpace(options.TagName)
if tagName != "" {
return s.getReleaseByTag(ctx, repo, tagName)
}
if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") {
return s.getLatestPreviewRelease(ctx, repo)
}
return s.getLatestStableRelease(ctx, repo)
}
func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo)
return s.fetchReleaseFromURL(ctx, url)
}
func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
var releases []githubRelease
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
return nil, err
}
for _, release := range releases {
if release.Draft || !release.Prerelease {
continue
}
releaseCopy := release
return &releaseCopy, nil
}
return nil, nil
}
func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag))
return s.fetchReleaseFromURL(ctx, url)
}
func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubRelease, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, err
}
defer func() {
if err := resp.Body.Close(); err != nil {
slog.Error("failed to close response body", "error", err)
}
}()
if resp.StatusCode == http.StatusNotFound {
return nil, nil
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
return decodeRelease(resp.Body)
}
func decodeRelease(reader io.Reader) (*githubRelease, error) {
var release githubRelease
if err := json.NewDecoder(reader).Decode(&release); err != nil {
return nil, err
}
return &release, nil
}
func (s *Service) resolveReleaseAsset(ctx context.Context, release *githubRelease, assetName string) (downloadURL string, expectedChecksum string, err error) {
checksumAssetName := assetName + ".sha256"
var checksumURL string
for _, asset := range release.Assets {
switch asset.Name {
case assetName:
downloadURL = asset.BrowserDownloadURL
expectedChecksum = normalizeGitHubDigest(asset.Digest)
case checksumAssetName:
checksumURL = asset.BrowserDownloadURL
}
}
if downloadURL == "" {
return "", "", fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
}
if expectedChecksum != "" {
return downloadURL, expectedChecksum, nil
}
if checksumURL == "" {
return downloadURL, "", fmt.Errorf("no sha256 digest or checksum asset %q in release %s", checksumAssetName, release.TagName)
}
expectedChecksum, err = s.downloadChecksum(ctx, checksumURL, assetName)
if err != nil {
return downloadURL, "", fmt.Errorf("download checksum: %w", err)
}
return downloadURL, expectedChecksum, nil
}
func normalizeGitHubDigest(digest string) string {
digest = strings.TrimSpace(digest)
if digest == "" {
return ""
}
const prefix = "sha256:"
if strings.HasPrefix(strings.ToLower(digest), prefix) {
digest = digest[len(prefix):]
}
digest = strings.ToLower(digest)
if isSHA256Hex(digest) {
return digest
}
return ""
}
func (s *Service) downloadChecksum(ctx context.Context, url string, assetName string) (string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return "", err
}
resp, err := s.httpClient.Do(req)
if err != nil {
return "", err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("checksum download returned %s", resp.Status)
}
content, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumAssetSize+1))
if err != nil {
return "", err
}
if len(content) > maxChecksumAssetSize {
return "", fmt.Errorf("checksum asset exceeds %d bytes", maxChecksumAssetSize)
}
checksum, err := parseSHA256Checksum(string(content), assetName)
if err != nil {
return "", err
}
return checksum, nil
}
func parseSHA256Checksum(content string, assetName string) (string, error) {
assetName = strings.TrimSpace(assetName)
for line := range strings.SplitSeq(content, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if checksum, ok := parseSHA256Line(line, assetName); ok {
return checksum, nil
}
}
if assetName == "" {
return "", errors.New("checksum asset does not contain a valid sha256 digest")
}
return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName)
}
func parseSHA256Line(line string, assetName string) (string, bool) {
fields := strings.Fields(line)
if len(fields) == 1 && isSHA256Hex(fields[0]) {
return strings.ToLower(fields[0]), true
}
if len(fields) >= 2 && isSHA256Hex(fields[0]) {
fileName := strings.TrimPrefix(strings.TrimSpace(fields[1]), "*")
if assetName == "" || fileName == assetName {
return strings.ToLower(fields[0]), true
}
}
prefix := "SHA256("
if strings.HasPrefix(line, prefix) {
closing := strings.Index(line, ")")
if closing > len(prefix) && closing+1 < len(line) {
fileName := strings.TrimSpace(line[len(prefix):closing])
rest := strings.TrimSpace(line[closing+1:])
rest = strings.TrimPrefix(rest, "=")
rest = strings.TrimSpace(rest)
if isSHA256Hex(rest) && (assetName == "" || fileName == assetName) {
return strings.ToLower(rest), true
}
}
}
return "", false
}
func isSHA256Hex(value string) bool {
value = strings.TrimSpace(value)
if len(value) != sha256.Size*2 {
return false
}
_, err := hex.DecodeString(value)
return err == nil
}
func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error {
expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum))
if !isSHA256Hex(expectedChecksum) {
return errors.New("invalid expected sha256 checksum")
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
resp, err := s.httpClient.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("download returned %s", resp.Status)
}
tmpPath := targetPath + ".update"
if runtime.GOOS == goosWindows && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") {
tmpPath += ".exe"
}
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, updateTmpFilePerm) //nolint:gosec // tmpPath is derived from the configured updater binary location
if err != nil {
return err
}
hasher := sha256.New()
if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil {
_ = tmpFile.Close()
_ = os.Remove(tmpPath)
return err
}
if err = tmpFile.Close(); err != nil {
_ = os.Remove(tmpPath)
return err
}
actualChecksum := hex.EncodeToString(hasher.Sum(nil))
if actualChecksum != expectedChecksum {
_ = os.Remove(tmpPath)
return fmt.Errorf("sha256 checksum mismatch: expected %s, got %s", expectedChecksum, actualChecksum)
}
if err = os.Chmod(tmpPath, updateBinaryFilePerm); err != nil && runtime.GOOS != goosWindows { //nolint:gosec // downloaded edge binary must remain executable
_ = os.Remove(tmpPath)
return fmt.Errorf("set executable permission: %w", err)
}
slog.Info(s.logLabel + " binary updated, restarting")
return replaceAndRestartFunc(targetPath, tmpPath)
}
func (s *Service) assetNameForGOOSGOARCH(goos string, goarch string) string {
name := fmt.Sprintf("%s-%s-%s", s.assetPrefix, goos, goarch)
if goos == goosWindows {
return name + ".exe"
}
return name
}
func normalizeVersion(v string) string {
v = strings.TrimSpace(v)
v = strings.TrimPrefix(v, "v")
return v
}
func isNewer(local, remote string) bool {
return compareVersions(local, remote) < 0
}
func buildReleaseCheckKey(options UpdateOptions, remoteVersion string) string {
channel := strings.TrimSpace(options.Channel)
if channel == "" {
channel = "stable"
}
if tagName := strings.TrimSpace(options.TagName); tagName != "" {
return channel + ":" + tagName
}
return channel + ":" + remoteVersion
}
func compareVersions(local string, remote string) int {
return ofutil.CompareVersions(local, remote)
}
@@ -0,0 +1,327 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
)
type roundTripFunc func(req *http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func testService(httpClient *http.Client) *Service {
return &Service{
httpClient: httpClient,
localVersion: "v1.0.0",
assetPrefix: "openflare-agent",
logLabel: "agent",
}
}
func TestGetLatestPreviewRelease(t *testing.T) {
service := testService(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases?per_page=20" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`[
{"tag_name":"v1.0.0","prerelease":false},
{"tag_name":"v1.1.0-rc.1","prerelease":true}
]`)),
}, nil
}),
})
release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{Channel: "preview"})
if err != nil {
t.Fatalf("expected preview release query to succeed: %v", err)
}
if release == nil || release.TagName != "v1.1.0-rc.1" {
t.Fatalf("unexpected preview release: %#v", release)
}
}
func TestGetReleaseByTag(t *testing.T) {
service := testService(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v1.1.0-rc.1" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)),
}, nil
}),
})
release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{Channel: "preview", TagName: "v1.1.0-rc.1", Force: true})
if err != nil {
t.Fatalf("expected tag release query to succeed: %v", err)
}
if release == nil || release.TagName != "v1.1.0-rc.1" {
t.Fatalf("unexpected tag release: %#v", release)
}
}
func TestCheckAndUpdateRequiresChecksumSource(t *testing.T) {
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
service := testService(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{
"tag_name":"v1.0.1",
"assets":[
{"name":"` + assetName + `","browser_download_url":"https://example.test/agent"}
]
}`)),
}, nil
}),
})
err := service.CheckAndUpdate(context.Background(), "Rain-kl/OpenFlare", UpdateOptions{})
if err == nil || !strings.Contains(err.Error(), "no sha256 digest or checksum asset") {
t.Fatalf("expected missing checksum source error, got %v", err)
}
}
func TestNormalizeGitHubDigest(t *testing.T) {
checksum := strings.Repeat("a", sha256.Size*2)
testCases := []struct {
name string
input string
want string
}{
{name: "prefixed digest", input: "sha256:" + checksum, want: checksum},
{name: "bare hex", input: checksum, want: checksum},
{name: "empty", input: "", want: ""},
{name: "invalid", input: "sha256:not-a-digest", want: ""},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
if got := normalizeGitHubDigest(testCase.input); got != testCase.want {
t.Fatalf("unexpected digest: got %q want %q", got, testCase.want)
}
})
}
}
func TestResolveReleaseAssetPrefersDigest(t *testing.T) {
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
checksum := strings.Repeat("b", sha256.Size*2)
service := testService(nil)
downloadURL, expectedChecksum, err := service.resolveReleaseAsset(context.Background(), &githubRelease{
TagName: "v1.0.1",
Assets: []githubAsset{
{
Name: assetName,
BrowserDownloadURL: "https://example.test/agent",
Digest: "sha256:" + checksum,
},
{
Name: assetName + ".sha256",
BrowserDownloadURL: "https://example.test/agent.sha256",
},
},
}, assetName)
if err != nil {
t.Fatalf("expected digest resolution to succeed: %v", err)
}
if downloadURL != "https://example.test/agent" {
t.Fatalf("unexpected download url: %s", downloadURL)
}
if expectedChecksum != checksum {
t.Fatalf("unexpected checksum: got %s want %s", expectedChecksum, checksum)
}
}
func TestResolveReleaseAssetFallsBackToChecksumAsset(t *testing.T) {
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
checksum := strings.Repeat("c", sha256.Size*2)
service := testService(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://example.test/agent.sha256" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(checksum + "\n")),
}, nil
}),
})
downloadURL, expectedChecksum, err := service.resolveReleaseAsset(context.Background(), &githubRelease{
TagName: "v1.0.1",
Assets: []githubAsset{
{
Name: assetName,
BrowserDownloadURL: "https://example.test/agent",
},
{
Name: assetName + ".sha256",
BrowserDownloadURL: "https://example.test/agent.sha256",
},
},
}, assetName)
if err != nil {
t.Fatalf("expected checksum fallback to succeed: %v", err)
}
if downloadURL != "https://example.test/agent" {
t.Fatalf("unexpected download url: %s", downloadURL)
}
if expectedChecksum != checksum {
t.Fatalf("unexpected checksum: got %s want %s", expectedChecksum, checksum)
}
}
func TestParseSHA256Checksum(t *testing.T) {
checksum := strings.Repeat("a", sha256.Size*2)
testCases := []struct {
name string
content string
asset string
want string
}{
{name: "single digest", content: checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
{name: "sha256sum format", content: checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
{name: "bsd format", content: "SHA256(openflare-agent-linux-amd64)= " + checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
{name: "selects matching file", content: strings.Repeat("b", sha256.Size*2) + " other\n" + checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
got, err := parseSHA256Checksum(testCase.content, testCase.asset)
if err != nil {
t.Fatalf("expected checksum parse to succeed: %v", err)
}
if got != testCase.want {
t.Fatalf("unexpected checksum: got %s want %s", got, testCase.want)
}
})
}
}
func TestDownloadAndRestartVerifiesChecksum(t *testing.T) {
payload := []byte("new-agent-binary")
sum := sha256.Sum256(payload)
expectedChecksum := hex.EncodeToString(sum[:])
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
t.Fatalf("write target: %v", err)
}
var replacedTarget string
var replacedTemp string
originalReplace := replaceAndRestartFunc
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
replacedTarget = execPath
replacedTemp = tmpPath
return nil
}
t.Cleanup(func() {
replaceAndRestartFunc = originalReplace
})
service := testService(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(string(payload))),
}, nil
}),
})
if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil {
t.Fatalf("expected verified download to succeed: %v", err)
}
if replacedTarget != targetPath {
t.Fatalf("unexpected replace target: %s", replacedTarget)
}
if replacedTemp == "" {
t.Fatal("expected replacement temp path to be recorded")
}
if _, err := os.Stat(replacedTemp); err != nil {
t.Fatalf("expected verified temp binary to remain for replacement: %v", err)
}
}
func TestDownloadAndRestartRejectsChecksumMismatch(t *testing.T) {
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
t.Fatalf("write target: %v", err)
}
originalReplace := replaceAndRestartFunc
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
t.Fatal("replace should not run on checksum mismatch")
return nil
}
t.Cleanup(func() {
replaceAndRestartFunc = originalReplace
})
service := testService(&http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("tampered")),
}, nil
}),
})
err := service.downloadAndRestart(context.Background(), "https://example.test/agent", strings.Repeat("0", sha256.Size*2), targetPath)
if err == nil || !strings.Contains(err.Error(), "sha256 checksum mismatch") {
t.Fatalf("expected checksum mismatch error, got %v", err)
}
if _, err = os.Stat(targetPath + ".update"); !os.IsNotExist(err) {
t.Fatalf("expected temp update file to be removed, stat err=%v", err)
}
}
func TestIsNewerSupportsPrerelease(t *testing.T) {
testCases := []struct {
name string
local string
remote string
expected bool
}{
{name: "stable newer than prerelease", local: "1.2.3-rc.1", remote: "1.2.3", expected: true},
{name: "same stable not newer", local: "1.2.3", remote: "1.2.3-rc.1", expected: false},
{name: "higher prerelease sequence", local: "1.2.3-rc.1", remote: "1.2.3-rc.2", expected: true},
{name: "higher minor", local: "1.2.3", remote: "1.3.0-rc.1", expected: true},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
if actual := isNewer(testCase.local, testCase.remote); actual != testCase.expected {
t.Fatalf("unexpected compare result: local=%s remote=%s actual=%v expected=%v", testCase.local, testCase.remote, actual, testCase.expected)
}
})
}
}
@@ -0,0 +1,159 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package wsclient provides WebSocket client abstractions for edge node communication.
package wsclient
import (
"context"
"time"
pkgprotocol "Wavelet/openflare/share/protocol"
shared "Wavelet/openflare/share/wsclient"
)
// WSMessage is an alias for the shared WebSocket message type.
type WSMessage = shared.WSMessage
// MessageHandler is an alias for the shared WebSocket message handler type.
type MessageHandler = shared.MessageHandler
// Preset represents a predefined connection configuration for a specific edge role.
type Preset int
const (
// PresetAgent is the configuration preset for agent connections.
PresetAgent Preset = iota
// PresetRelay is the configuration preset for relay connections.
PresetRelay
// PresetFlared is the configuration preset for flared (tunnel) connections.
PresetFlared
)
type presetConfig struct {
HeaderKey string
WSPath string
}
var presets = map[Preset]presetConfig{
PresetAgent: {HeaderKey: "X-Agent-Token", WSPath: "/api/v1/agent/ws"},
PresetRelay: {HeaderKey: "X-Agent-Token", WSPath: "/api/v1/relay/ws"},
PresetFlared: {HeaderKey: "X-Tunnel-Token", WSPath: "/api/v1/tunnel/ws"},
}
// PresetHeaderKey returns the HTTP header key used for authentication with the given preset.
func PresetHeaderKey(preset Preset) string {
return presets[preset].HeaderKey
}
// PresetWSPath returns the WebSocket path used for the given preset.
func PresetWSPath(preset Preset) string {
return presets[preset].WSPath
}
// Client is a WebSocket client configured for a specific edge preset.
type Client struct {
sharedClient *shared.Client
}
// New creates a new Client for the given preset, base URL, token, and timeout.
func New(preset Preset, baseURL, token string, timeout time.Duration) *Client {
cfg := presets[preset]
return &Client{
sharedClient: shared.New(shared.Config{
BaseURL: baseURL,
Token: token,
Timeout: timeout,
HeaderKey: cfg.HeaderKey,
WSPath: cfg.WSPath,
}),
}
}
// SetToken updates the authentication token used by the client.
func (c *Client) SetToken(token string) {
c.sharedClient.SetToken(token)
}
// URL returns the fully resolved WebSocket URL for this client.
func (c *Client) URL() string {
return c.sharedClient.URL()
}
// Connection represents an established WebSocket connection to an edge node.
type Connection struct {
sharedConn *shared.Connection
}
// Connect establishes a WebSocket connection using the client configuration.
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
conn, err := c.sharedClient.Connect(ctx)
if err != nil {
return nil, err
}
return &Connection{sharedConn: conn}, nil
}
// AgentConnection is a Connection specialized for agent node communication.
type AgentConnection struct {
Connection
}
// ConnectAgent establishes a WebSocket connection and returns it as an AgentConnection.
func (c *Client) ConnectAgent(ctx context.Context) (*AgentConnection, error) {
conn, err := c.Connect(ctx)
if err != nil {
return nil, err
}
return &AgentConnection{Connection: *conn}, nil
}
// URL returns the resolved WebSocket URL of this connection.
func (conn *Connection) URL() string {
if conn == nil || conn.sharedConn == nil {
return ""
}
return conn.sharedConn.URL
}
// SendPing sends a ping message over the connection.
func (conn *Connection) SendPing() error {
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypePing, nil)
}
// SendPong sends a pong message over the connection.
func (conn *Connection) SendPong() error {
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypePong, nil)
}
// SendMessage sends a typed message with an optional payload over the connection.
func (conn *Connection) SendMessage(msgType string, payload any) error {
return conn.sharedConn.SendMessage(msgType, payload)
}
// Receive reads the next message from the connection.
func (conn *Connection) Receive() (pkgprotocol.WSMessage, error) {
var message pkgprotocol.WSMessage
if err := conn.sharedConn.Receive(&message); err != nil {
return message, err
}
return message, nil
}
// RunReceiveLoop blocks and dispatches incoming messages to the handler until the context is canceled.
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
return conn.sharedConn.RunReceiveLoop(ctx, handler)
}
// Close gracefully closes the WebSocket connection.
func (conn *Connection) Close() error {
if conn == nil || conn.sharedConn == nil {
return nil
}
return conn.sharedConn.Close()
}
// SendStatus sends a node status payload over the agent connection.
func (conn *AgentConnection) SendStatus(payload pkgprotocol.NodePayload) error {
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypeStatus, payload)
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package wsclient
import (
"testing"
"time"
)
func TestPresetConfig(t *testing.T) {
tests := []struct {
preset Preset
headerKey string
wsPath string
}{
{preset: PresetAgent, headerKey: "X-Agent-Token", wsPath: "/api/v1/agent/ws"},
{preset: PresetRelay, headerKey: "X-Agent-Token", wsPath: "/api/v1/relay/ws"},
{preset: PresetFlared, headerKey: "X-Tunnel-Token", wsPath: "/api/v1/tunnel/ws"},
}
for _, tt := range tests {
t.Run(tt.wsPath, func(t *testing.T) {
if got := PresetHeaderKey(tt.preset); got != tt.headerKey {
t.Fatalf("PresetHeaderKey() = %q, want %q", got, tt.headerKey)
}
if got := PresetWSPath(tt.preset); got != tt.wsPath {
t.Fatalf("PresetWSPath() = %q, want %q", got, tt.wsPath)
}
})
}
}
func TestClientURL(t *testing.T) {
tests := []struct {
name string
preset Preset
baseURL string
want string
}{
{
name: "agent https",
preset: PresetAgent,
baseURL: "https://example.com",
want: "wss://example.com/api/v1/agent/ws",
},
{
name: "relay http with path prefix",
preset: PresetRelay,
baseURL: "http://example.com/api",
want: "ws://example.com/api/api/v1/relay/ws",
},
{
name: "flared wss",
preset: PresetFlared,
baseURL: "wss://edge.example.com",
want: "wss://edge.example.com/api/v1/tunnel/ws",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
client := New(tt.preset, tt.baseURL, "token", time.Second)
if got := client.URL(); got != tt.want {
t.Fatalf("URL() = %q, want %q", got, tt.want)
}
})
}
}
@@ -0,0 +1,254 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import "strings"
type countryCentroid struct {
lat float64
lon float64
}
var countryCentroidsByName = map[string]countryCentroid{
"Afghanistan": {lat: 33.833494, lon: 65.992544},
"Albania": {lat: 41.179283, lon: 20.035804},
"Algeria": {lat: 28.158203, lon: 2.617517},
"Angola": {lat: -12.334558, lon: 17.563198},
"Argentina": {lat: -35.178772, lon: -65.156706},
"Armenia": {lat: 40.298846, lon: 44.929380},
"Australia": {lat: -25.574956, lon: 134.361181},
"Austria": {lat: 47.570428, lon: 14.101658},
"Azerbaijan": {lat: 40.331151, lon: 47.636778},
"Bangladesh": {lat: 23.887936, lon: 90.210098},
"Belarus": {lat: 53.534329, lon: 28.048953},
"Belgium": {lat: 50.646495, lon: 4.628130},
"Belize": {lat: 17.170731, lon: -88.720961},
"Benin": {lat: 9.630483, lon: 2.323508},
"Bhutan": {lat: 27.401336, lon: 90.405624},
"Bolivia": {lat: -16.701415, lon: -64.691050},
"Bosnia and Herzegovina": {lat: 44.163989, lon: 17.762462},
"Botswana": {lat: -22.185334, lon: 23.791789},
"Brazil": {lat: -10.827996, lon: -53.113858},
"Brunei": {lat: 4.451218, lon: 114.542293},
"Bulgaria": {lat: 42.777857, lon: 25.225364},
"Burkina Faso": {lat: 12.260370, lon: -1.775743},
"Burundi": {lat: -3.380283, lon: 29.873066},
"Cambodia": {lat: 12.717933, lon: 104.907023},
"Cameroon": {lat: 5.692692, lon: 12.739868},
"Canada": {lat: 57.944575, lon: -102.317776},
"Central African Republic": {lat: 6.563625, lon: 20.490782},
"Chad": {lat: 15.327299, lon: 18.645816},
"Chile": {lat: -35.980286, lon: -71.348917},
"China": {lat: 36.610942, lon: 103.798043},
"Colombia": {lat: 3.923562, lon: -73.077941},
"Costa Rica": {lat: 9.986463, lon: -84.217425},
"Croatia": {lat: 45.185678, lon: 16.458685},
"Cuba": {lat: 21.621368, lon: -78.895403},
"Cyprus": {lat: 34.913044, lon: 33.024092},
"Czech Republic": {lat: 49.701608, lon: 15.330767},
"Democratic Republic of the Congo": {lat: -2.877553, lon: 23.648117},
"Denmark": {lat: 74.725352, lon: -41.290889},
"Djibouti": {lat: 11.745989, lon: 42.567886},
"Dominican Republic": {lat: 18.896879, lon: -70.485959},
"East Timor": {lat: -8.858374, lon: 125.782226},
"Ecuador": {lat: -1.445457, lon: -78.383848},
"Egypt": {lat: 26.492865, lon: 29.867399},
"El Salvador": {lat: 13.734573, lon: -88.865936},
"Equatorial Guinea": {lat: 1.566080, lon: 10.481671},
"Eritrea": {lat: 15.359810, lon: 38.835037},
"Estonia": {lat: 58.693030, lon: 25.811350},
"Ethiopia": {lat: 8.621596, lon: 39.603726},
"Fiji": {lat: -17.827004, lon: 177.984152},
"Finland": {lat: 64.522797, lon: 26.289584},
"France": {lat: -21.285376, lon: 165.438171},
"Gabon": {lat: -0.585980, lon: 11.774049},
"Gambia": {lat: 13.370627, lon: -16.213050},
"Georgia": {lat: 42.175028, lon: 43.486503},
"Germany": {lat: 51.083853, lon: 10.380497},
"Ghana": {lat: 7.940789, lon: -1.215016},
"Greece": {lat: 39.517735, lon: 22.534352},
"Guatemala": {lat: 15.679044, lon: -90.353440},
"Guinea": {lat: 10.440780, lon: -10.943568},
"Guinea Bissau": {lat: 12.050889, lon: -14.929646},
"Guyana": {lat: 4.800330, lon: -58.978019},
"Haiti": {lat: 18.916196, lon: -72.678666},
"Honduras": {lat: 14.821030, lon: -86.656518},
"Hong Kong": {lat: 22.278300, lon: 114.174700},
"Hungary": {lat: 47.168543, lon: 19.410867},
"Iceland": {lat: 65.030350, lon: -18.509671},
"India": {lat: 22.901587, lon: 79.586369},
"Indonesia": {lat: -0.446813, lon: 101.522374},
"Iran": {lat: 32.576164, lon: 54.266247},
"Iraq": {lat: 33.034146, lon: 43.740065},
"Ireland": {lat: 53.234132, lon: -8.117695},
"Israel": {lat: 31.947585, lon: 35.241574},
"Italy": {lat: 43.512811, lon: 12.162338},
"Ivory Coast": {lat: 7.626173, lon: -5.569772},
"Jamaica": {lat: 18.172958, lon: -77.328505},
"Japan": {lat: 36.589925, lon: 137.973525},
"Jordan": {lat: 31.237145, lon: 36.761315},
"Kashmir": {lat: 35.415739, lon: 77.087337},
"Kazakhstan": {lat: 48.155989, lon: 67.278154},
"Kenya": {lat: 0.596719, lon: 37.795079},
"Kosovo": {lat: 42.517832, lon: 20.851519},
"Kuwait": {lat: 29.316145, lon: 47.566177},
"Kyrgyzstan": {lat: 41.483635, lon: 74.582247},
"Laos": {lat: 18.494355, lon: 103.752638},
"Latvia": {lat: 56.853144, lon: 24.894962},
"Lebanon": {lat: 33.915290, lon: 35.885023},
"Lesotho": {lat: -29.573512, lon: 28.243038},
"Liberia": {lat: 6.453996, lon: -9.325868},
"Libya": {lat: 27.031457, lon: 18.008785},
"Lithuania": {lat: 55.314246, lon: 23.872794},
"Luxembourg": {lat: 49.767461, lon: 6.070340},
"Macedonia": {lat: 41.603736, lon: 21.704642},
"Madagascar": {lat: -19.374777, lon: 46.701725},
"Malawi": {lat: -13.200363, lon: 34.288911},
"Malaysia": {lat: 3.577242, lon: 114.692015},
"Mali": {lat: 17.347810, lon: -3.532183},
"Mauritania": {lat: 20.259794, lon: -10.343079},
"Mexico": {lat: 23.943713, lon: -102.518498},
"Moldova": {lat: 47.201792, lon: 28.446473},
"Mongolia": {lat: 46.825609, lon: 103.064235},
"Montenegro": {lat: 42.783955, lon: 19.233365},
"Morocco": {lat: 29.839324, lon: -8.459193},
"Mozambique": {lat: -17.272975, lon: 35.531737},
"Myanmar": {lat: 21.222015, lon: 96.495888},
"Namibia": {lat: -22.132303, lon: 17.208945},
"Nepal": {lat: 28.243336, lon: 83.923410},
"Netherlands": {lat: 52.277341, lon: 5.646194},
"New Zealand": {lat: -43.955382, lon: 170.540471},
"Nicaragua": {lat: 12.835549, lon: -85.032299},
"Niger": {lat: 17.417582, lon: 9.384787},
"Nigeria": {lat: 9.589590, lon: 8.074235},
"North Korea": {lat: 40.144857, lon: 127.215518},
"Northern Cyprus": {lat: 35.225753, lon: 33.390792},
"Norway": {lat: 64.322498, lon: 13.994248},
"Oman": {lat: 20.573738, lon: 56.082719},
"Pakistan": {lat: 29.951586, lon: 69.329881},
"Panama": {lat: 8.521692, lon: -80.059706},
"Papua New Guinea": {lat: -6.607045, lon: 144.228968},
"Paraguay": {lat: -23.223896, lon: -58.399006},
"Peru": {lat: -9.156316, lon: -74.388502},
"Philippines": {lat: 15.948824, lon: 121.457548},
"Poland": {lat: 52.119631, lon: 19.379892},
"Portugal": {lat: 39.685490, lon: -7.977839},
"Qatar": {lat: 25.328766, lon: 51.170816},
"Republic of Serbia": {lat: 44.199984, lon: 20.766250},
"Republic of the Congo": {lat: -0.856400, lon: 15.212074},
"Romania": {lat: 45.845608, lon: 24.971632},
"Russia": {lat: 61.677044, lon: 99.052282},
"Rwanda": {lat: -1.995363, lon: 29.917651},
"Saudi Arabia": {lat: 24.122854, lon: 44.534287},
"Senegal": {lat: 14.338614, lon: -14.472411},
"Sierra Leone": {lat: 8.575093, lon: -11.820132},
"Singapore": {lat: 1.352100, lon: 103.819800},
"Slovakia": {lat: 48.715407, lon: 19.472965},
"Slovenia": {lat: 46.108070, lon: 14.771193},
"Solomon Islands": {lat: -9.628590, lon: 160.156397},
"Somalia": {lat: 4.747961, lon: 45.706540},
"Somaliland": {lat: 9.729995, lon: 46.255133},
"South Africa": {lat: -29.008774, lon: 25.160630},
"South Korea": {lat: 36.475158, lon: 127.872804},
"South Sudan": {lat: 7.307277, lon: 30.253828},
"Spain": {lat: 40.391565, lon: -3.570837},
"Sri Lanka": {lat: 7.649883, lon: 80.700551},
"Sudan": {lat: 15.992059, lon: 29.946244},
"Suriname": {lat: 4.126766, lon: -55.910298},
"Swaziland": {lat: -26.563546, lon: 31.476645},
"Sweden": {lat: 62.835488, lon: 16.752551},
"Switzerland": {lat: 46.805544, lon: 8.205813},
"Syria": {lat: 35.017694, lon: 38.504490},
"Taiwan": {lat: 23.697800, lon: 120.960500},
"Tajikistan": {lat: 38.552044, lon: 70.987127},
"Thailand": {lat: 15.132319, lon: 101.010603},
"The Bahamas": {lat: 24.719920, lon: -78.027339},
"Togo": {lat: 8.576561, lon: 0.949514},
"Trinidad and Tobago": {lat: 10.339930, lon: -61.270459},
"Tunisia": {lat: 34.126731, lon: 9.546833},
"Turkey": {lat: 38.998407, lon: 35.424986},
"Turkmenistan": {lat: 39.109817, lon: 59.392923},
"Uganda": {lat: 1.275770, lon: 32.365755},
"Ukraine": {lat: 49.013453, lon: 31.381193},
"United Arab Emirates": {lat: 23.905698, lon: 54.303458},
"United Kingdom": {lat: -54.468348, lon: -36.367083},
"United Republic of Tanzania": {lat: -6.276146, lon: 34.794901},
"United States of America": {lat: 39.526894, lon: -99.146697},
"Uruguay": {lat: -32.806517, lon: -56.016656},
"Uzbekistan": {lat: 41.750262, lon: 63.148827},
"Vanuatu": {lat: -15.266627, lon: 166.842287},
"Venezuela": {lat: 7.118300, lon: -66.188217},
"Vietnam": {lat: 16.630882, lon: 106.299391},
"Western Sahara": {lat: 24.221136, lon: -12.218085},
"Yemen": {lat: 15.935826, lon: 47.546352},
"Zambia": {lat: -13.460954, lon: 27.775322},
"Zimbabwe": {lat: -19.006775, lon: 29.850564},
}
// CountryCentroidByName returns latitude and longitude for a country name, if found.
// Accepts composite labels such as "Hong Kong, Hong Kong, HK" from GeoIP providers.
func CountryCentroidByName(name string) (lat float64, lon float64, ok bool) {
name = strings.TrimSpace(name)
if name == "" {
return 0, 0, false
}
if v, found := countryCentroidsByName[name]; found {
return v.lat, v.lon, true
}
// Try comma-separated parts (city / region / country / ISO).
for part := range strings.SplitSeq(name, ",") {
part = strings.TrimSpace(part)
if part == "" {
continue
}
if v, found := countryCentroidsByName[part]; found {
return v.lat, v.lon, true
}
if lat, lon, ok = CountryCentroidByISO(part); ok {
return lat, lon, true
}
}
return 0, 0, false
}
// CountryCentroidByISO returns latitude and longitude for an ISO country code, if found.
func CountryCentroidByISO(iso string) (lat float64, lon float64, ok bool) {
name, ok := isoToCountryName[strings.ToUpper(strings.TrimSpace(iso))]
if !ok {
return 0, 0, false
}
return CountryCentroidByName(name)
}
var isoToCountryName = map[string]string{
"AT": "Austria",
"AU": "Australia",
"BR": "Brazil",
"CA": "Canada",
"CH": "Switzerland",
"CN": "China",
"DE": "Germany",
"ES": "Spain",
"FI": "Finland",
"FR": "France",
"GB": "United Kingdom",
"HK": "Hong Kong",
"ID": "Indonesia",
"IE": "Ireland",
"IN": "India",
"IT": "Italy",
"JP": "Japan",
"KR": "South Korea",
"MY": "Malaysia",
"NL": "Netherlands",
"NO": "Norway",
"PL": "Poland",
"RU": "Russia",
"SE": "Sweden",
"SG": "Singapore",
"TH": "Thailand",
"TW": "Taiwan",
"US": "United States of America",
"VN": "Vietnam",
}
@@ -0,0 +1,47 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import "testing"
func TestCountryCentroidGermany(t *testing.T) {
lat, lon, ok := CountryCentroidByISO("DE")
if !ok {
t.Fatal("expected DE centroid")
}
if lat < 47 || lat > 55 || lon < 5 || lon > 15 {
t.Fatalf("unexpected Germany centroid: %f,%f", lat, lon)
}
}
func TestCountryCentroidHongKongSingaporeTaiwan(t *testing.T) {
cases := []struct {
iso string
name string
minLat float64
maxLat float64
minLon float64
maxLon float64
}{
{iso: "HK", name: "Hong Kong", minLat: 22, maxLat: 23, minLon: 113, maxLon: 115},
{iso: "SG", name: "Singapore", minLat: 1, maxLat: 2, minLon: 103, maxLon: 104},
{iso: "TW", name: "Taiwan", minLat: 22, maxLat: 26, minLon: 119, maxLon: 122},
}
for _, tc := range cases {
lat, lon, ok := CountryCentroidByISO(tc.iso)
if !ok {
t.Fatalf("expected %s centroid by ISO", tc.iso)
}
if lat < tc.minLat || lat > tc.maxLat || lon < tc.minLon || lon > tc.maxLon {
t.Fatalf("%s ISO centroid out of range: %f,%f", tc.iso, lat, lon)
}
lat, lon, ok = CountryCentroidByName(tc.name)
if !ok {
t.Fatalf("expected %s centroid by name", tc.name)
}
if lat < tc.minLat || lat > tc.maxLat || lon < tc.minLon || lon > tc.maxLon {
t.Fatalf("%s name centroid out of range: %f,%f", tc.name, lat, lon)
}
}
}
@@ -0,0 +1,37 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"errors"
"net"
)
// EmptyProvider is a no-op GeoIP backend used before a real provider is configured.
type EmptyProvider struct{}
// Name returns the provider identifier.
func (e *EmptyProvider) Name() string {
return "EmptyProvider"
}
// Initialize prepares the empty provider for use.
func (e *EmptyProvider) Initialize() error {
return nil
}
// GetGeoInfo reports that no GeoIP provider has been configured.
func (e *EmptyProvider) GetGeoInfo(_ net.IP) (*GeoInfo, error) {
return nil, errors.New("you are using an empty GeoIP provider, please set a valid provider")
}
// UpdateDatabase reports that no GeoIP provider has been configured.
func (e *EmptyProvider) UpdateDatabase() error {
return errors.New("you are using an empty GeoIP provider, please set a valid provider")
}
// Close releases resources held by the empty provider.
func (e *EmptyProvider) Close() error {
return nil
}
+263
View File
@@ -0,0 +1,263 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package geoip resolves geographic information for IP addresses.
package geoip
import (
"errors"
"fmt"
"log/slog"
"net"
"strings"
"sync"
"time"
"unicode"
ristretto "github.com/dgraph-io/ristretto/v2"
)
// CurrentProvider is the active GeoIP backend used by package-level helpers.
var CurrentProvider Service
var geoCache *providerCache
var providerMutex sync.RWMutex
var providerFactory = newProvider
// Supported GeoIP provider identifiers.
const (
// ProviderDisabled disables GeoIP lookups.
ProviderDisabled = "disabled"
ProviderMaxMind = "mmdb"
ProviderIPAPI = "ip-api"
ProviderGeoJS = "geojs"
ProviderIPInfo = "ipinfo"
defaultGeoCacheDuration = 48 * time.Hour
isoCountryCodeLength = 2
geoipDataDirPerm = 0o750
)
// GeoInfo contains normalized geographic metadata for an IP address.
type GeoInfo struct {
ISOCode string
Name string
Latitude *float64
Longitude *float64
}
func init() {
CurrentProvider = &EmptyProvider{}
geoCache = newProviderCache(defaultGeoCacheDuration)
}
// Service defines the core GeoIP lookup operations implemented by providers.
type Service interface {
Name() string
GetGeoInfo(ip net.IP) (*GeoInfo, error)
UpdateDatabase() error
Close() error
}
type cachedGeoInfo struct {
info *GeoInfo
expiresAt time.Time
}
type providerCache struct {
items *ristretto.Cache[string, cachedGeoInfo]
duration time.Duration
}
func newProviderCache(duration time.Duration) *providerCache {
items, err := ristretto.NewCache(&ristretto.Config[string, cachedGeoInfo]{
NumCounters: 1e5,
MaxCost: 2e4,
BufferItems: 64,
})
if err != nil {
panic(err)
}
return &providerCache{
items: items,
duration: duration,
}
}
func (c *providerCache) Get(key string) (*GeoInfo, bool) {
entry, ok := c.items.Get(key)
if !ok {
return nil, false
}
if time.Now().After(entry.expiresAt) {
c.items.Del(key)
return nil, false
}
return entry.info, true
}
func (c *providerCache) Set(key string, info *GeoInfo) {
c.items.Set(key, cachedGeoInfo{
info: info,
expiresAt: time.Now().Add(c.duration),
}, 1)
c.items.Wait()
}
func (c *providerCache) Flush() {
c.items.Clear()
}
// GetRegionUnicodeEmoji returns the regional indicator emoji for a two-letter ISO code.
func GetRegionUnicodeEmoji(isoCode string) string {
if len(isoCode) != isoCountryCodeLength {
return ""
}
isoCode = strings.ToUpper(isoCode)
if !unicode.IsLetter(rune(isoCode[0])) || !unicode.IsLetter(rune(isoCode[1])) {
return ""
}
rune1 := 0x1F1E6 + (rune(isoCode[0]) - 'A')
rune2 := 0x1F1E6 + (rune(isoCode[1]) - 'A')
return string(rune1) + string(rune2)
}
// InitGeoIP configures the active GeoIP provider.
func InitGeoIP(provider string) {
providerName := normalizeProvider(provider)
nextProvider, err := providerFactory(providerName)
if err != nil {
slog.Error("initialize GeoIP provider failed", "provider", providerName, "error", err)
nextProvider = &EmptyProvider{}
}
setProvider(nextProvider)
if providerName == ProviderDisabled {
slog.Info("GeoIP provider disabled")
return
}
slog.Info("GeoIP provider configured", "provider", CurrentProvider.Name())
}
// GetGeoInfo looks up geographic information for ip using the active provider.
func GetGeoInfo(ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, errors.New("IP address cannot be nil")
}
provider := getProvider()
cacheKey := provider.Name() + ":" + ip.String()
if cachedInfo, found := geoCache.Get(cacheKey); found {
return cachedInfo, nil
}
info, err := provider.GetGeoInfo(ip)
if err == nil && info != nil {
geoCache.Set(cacheKey, info)
}
return info, err
}
// LookupGeoInfoWithProvider looks up geographic information using a temporary provider.
func LookupGeoInfoWithProvider(providerName string, ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, errors.New("IP address cannot be nil")
}
provider, err := providerFactory(normalizeProvider(providerName))
if err != nil {
return nil, err
}
defer func() {
if closeErr := provider.Close(); closeErr != nil {
slog.Warn("close temporary GeoIP provider failed", "provider", provider.Name(), "error", closeErr)
}
}()
return provider.GetGeoInfo(ip)
}
// UpdateDatabase refreshes the active provider database and clears cached lookups.
func UpdateDatabase() error {
err := getProvider().UpdateDatabase()
if err == nil {
geoCache.Flush()
slog.Info("GeoIP cache cleared due to database update.")
}
return err
}
// IsValidProvider reports whether provider names a supported GeoIP backend.
func IsValidProvider(provider string) bool {
switch normalizeProvider(provider) {
case ProviderDisabled, ProviderMaxMind, ProviderIPAPI, ProviderGeoJS, ProviderIPInfo:
return true
default:
return false
}
}
func normalizeProvider(provider string) string {
normalized := strings.TrimSpace(strings.ToLower(provider))
if normalized == "" {
return ProviderDisabled
}
return normalized
}
func newProvider(provider string) (Service, error) {
switch provider {
case ProviderDisabled:
return &EmptyProvider{}, nil
case ProviderMaxMind:
return NewMaxMindGeoIPService()
case ProviderIPAPI:
return NewIPAPIService()
case ProviderGeoJS:
return NewGeoJSService()
case ProviderIPInfo:
return NewIPInfoService()
default:
return nil, fmt.Errorf("unsupported GeoIP provider %q", provider)
}
}
func setProvider(provider Service) {
providerMutex.Lock()
previous := CurrentProvider
CurrentProvider = provider
providerMutex.Unlock()
geoCache.Flush()
if previous != nil && previous != provider {
if err := previous.Close(); err != nil {
slog.Warn("close previous GeoIP provider failed", "error", err)
}
}
}
func getProvider() Service {
providerMutex.RLock()
defer providerMutex.RUnlock()
if CurrentProvider == nil {
return &EmptyProvider{}
}
return CurrentProvider
}
func float64Pointer(value float64) *float64 {
return &value
}
// ProviderFactoryForTest returns the provider factory used by package helpers.
func ProviderFactoryForTest() func(string) (Service, error) {
return providerFactory
}
// SetProviderFactoryForTest replaces the provider factory for tests.
func SetProviderFactoryForTest(factory func(string) (Service, error)) {
if factory == nil {
providerFactory = newProvider
return
}
providerFactory = factory
}
+102
View File
@@ -0,0 +1,102 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"net"
"testing"
)
type fakeProvider struct {
calls int
}
func (f *fakeProvider) Name() string {
return "fake"
}
func (f *fakeProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
f.calls++
return &GeoInfo{
ISOCode: "CN",
Name: "China",
}, nil
}
func (f *fakeProvider) UpdateDatabase() error {
return nil
}
func (f *fakeProvider) Close() error {
return nil
}
func TestGetGeoInfoCachesByProviderAndIP(t *testing.T) {
originalProvider := CurrentProvider
geoCache.Flush()
fake := &fakeProvider{}
CurrentProvider = fake
defer func() {
CurrentProvider = originalProvider
}()
ip := net.ParseIP("8.8.8.8")
record, err := GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if record == nil || record.ISOCode != "CN" {
t.Fatalf("expected cached record, got %#v", record)
}
_, err = GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error on second call, got %v", err)
}
if fake.calls != 1 {
t.Fatalf("expected provider to be called once, got %d", fake.calls)
}
}
func TestUnicodeEmoji(t *testing.T) {
emoji := GetRegionUnicodeEmoji("CN")
if emoji != "🇨🇳" {
t.Errorf("expected emoji for CN, got %s", emoji)
}
}
func TestIsValidProvider(t *testing.T) {
cases := map[string]bool{
"disabled": true,
"mmdb": true,
"ip-api": true,
"geojs": true,
"ipinfo": true,
"unknown": false,
}
for provider, want := range cases {
if got := IsValidProvider(provider); got != want {
t.Fatalf("provider %s validity mismatch: want %v, got %v", provider, want, got)
}
}
}
func TestLookupGeoInfoWithProviderUsesTemporaryProvider(t *testing.T) {
previousFactory := providerFactory
providerFactory = func(provider string) (Service, error) {
return &fakeProvider{}, nil
}
defer func() {
providerFactory = previousFactory
}()
info, err := LookupGeoInfoWithProvider("ipinfo", net.ParseIP("8.8.8.8"))
if err != nil {
t.Fatalf("expected lookup to succeed, got %v", err)
}
if info == nil || info.ISOCode != "CN" || info.Name != "China" {
t.Fatalf("unexpected geo info: %#v", info)
}
}
+92
View File
@@ -0,0 +1,92 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// GeoJSService resolves geographic information using the geojs.io service.
type GeoJSService struct {
Client *http.Client
}
// geoJSResponse 定义了 geojs.io 服务返回的 JSON 响应的结构。
// 我们只定义我们需要的字段。
type geoJSResponse struct {
Country string `json:"country"`
CountryCode string `json:"country_code"`
Latitude float64 `json:"latitude,string"`
Longitude float64 `json:"longitude,string"`
// 可以根据需要添加其他字段,例如:
// City string `json:"city"`
// Region string `json:"region"`
}
// NewGeoJSService 创建并返回一个 GeoJSService 的新实例。
func NewGeoJSService() (*GeoJSService, error) {
return &GeoJSService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置一个合理的超时时间
},
}, nil
}
// Name 返回服务的名称。
func (s *GeoJSService) Name() string {
return "geojs.io"
}
// GetGeoInfo 使用 geojs.io 服务检索给定 IP 地址的地理位置信息。
func (s *GeoJSService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// GeoJS 的 API 端点
apiURL := fmt.Sprintf("https://get.geojs.io/v1/ip/geo/%s.json", ip.String())
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, apiURL, nil)
if err != nil {
return nil, fmt.Errorf("failed to create request for geojs.io: %w", err)
}
resp, err := s.Client.Do(req)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from geojs.io: %w", err)
}
defer func() { _ = resp.Body.Close() }()
// 检查响应状态码
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("geojs.io returned non-200 status code: %d", resp.StatusCode)
}
var apiResp geoJSResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode geojs.io response: %w", err)
}
// 检查国家代码是否为空,因为 geojs 对无效/私有IP可能返回200 OK但内容为空
if apiResp.CountryCode == "" {
return nil, fmt.Errorf("geojs.io returned empty geo info for ip: %s", ip.String())
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Latitude),
Longitude: float64Pointer(apiResp.Longitude),
}, nil
}
// UpdateDatabase 对于 geojs.io 是一个空操作,因为它是一个 Web 服务。
func (s *GeoJSService) UpdateDatabase() error {
return nil
}
// Close 对于 geojs.io 是一个空操作。
func (s *GeoJSService) Close() error {
return nil
}
+95
View File
@@ -0,0 +1,95 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// IPAPIService resolves geographic information using the ip-api.com service.
type IPAPIService struct {
Client *http.Client
}
// ipAPIResponse 定义了 ip-api.com 服务返回的 JSON 响应的结构。
type ipAPIResponse struct {
Status string `json:"status"`
Message string `json:"message"` // 当 status 为 fail 时出现
Country string `json:"country"`
CountryCode string `json:"countryCode"`
Region string `json:"region"`
RegionName string `json:"regionName"`
City string `json:"city"`
Zip string `json:"zip"`
Lat float64 `json:"lat"`
Lon float64 `json:"lon"`
Timezone string `json:"timezone"`
ISP string `json:"isp"`
Org string `json:"org"`
As string `json:"as"`
Query string `json:"query"`
}
// Name returns the provider identifier for the ip-api.com service.
func (s *IPAPIService) Name() string {
return "ip-api.com"
}
// NewIPAPIService 创建并返回一个 IPAPIService 的新实例。
func NewIPAPIService() (*IPAPIService, error) {
return &IPAPIService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置请求超时
},
}, nil
}
// GetGeoInfo 使用 ip-api.com 服务检索给定 IP 地址的地理位置信息。
func (s *IPAPIService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// API URL, 使用 fields 参数来仅请求需要的字段
apiURL := fmt.Sprintf("http://ip-api.com/json/%s?fields=status,message,country,countryCode", ip.String())
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, apiURL, nil)
if err != nil {
return nil, fmt.Errorf("failed to create request for ip-api.com: %w", err)
}
resp, err := s.Client.Do(req)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ip-api.com: %w", err)
}
defer func() { _ = resp.Body.Close() }()
var apiResp ipAPIResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ip-api.com response: %w", err)
}
if apiResp.Status != "success" {
return nil, fmt.Errorf("ip-api.com returned an error: %s", apiResp.Message)
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Lat),
Longitude: float64Pointer(apiResp.Lon),
}, nil
}
// UpdateDatabase 对于 ip-api.com 是一个空操作,因为它是一个 Web 服务。
func (s *IPAPIService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ip-api.com 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPAPIService) Close() error {
// 无需执行任何操作
return nil
}
+138
View File
@@ -0,0 +1,138 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"strconv"
"strings"
"time"
)
// IPInfoService resolves geographic information using the ipinfo.io service.
type IPInfoService struct {
Client *http.Client
// 每天 1000 次请求,限制由 IP 地址的所有人共享。
// APIToken string
}
// ipInfoResponse 定义了 ipinfo.io 服务返回的 JSON 响应的结构,只包含免费额度可用的字段。
type ipInfoResponse struct {
IP string `json:"ip"`
Hostname string `json:"hostname"`
City string `json:"city"`
Region string `json:"region"`
Country string `json:"country"`
CountryCode string `json:"countryCode"` // ipinfo.io 返回 "country" 的 ISO 代码,这里为了与 GeoInfo 保持一致,额外添加一个 CountryCode
Loc string `json:"loc"` // Latitude,Longitude
Org string `json:"org"`
Postal string `json:"postal"`
Timezone string `json:"timezone"`
}
// NewIPInfoService 创建并返回一个 IPInfoService 的新实例。
func NewIPInfoService() (*IPInfoService, error) {
return &IPInfoService{
Client: &http.Client{
Timeout: 5 * time.Second,
},
}, nil
}
// Name 返回服务的名称。
func (s *IPInfoService) Name() string {
return "ipinfo.io"
}
// GetGeoInfo 使用 ipinfo.io 服务检索给定 IP 地址的地理位置信息。
// 免费额度主要提供国家信息。
func (s *IPInfoService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// IPinfo 免费额度不需要 API token 就可以查询基本的 IP 信息。
// API URL: https://ipinfo.io/json (查询自身IP) 或 https://ipinfo.io/YOUR_IP/json
apiURL := fmt.Sprintf("https://ipinfo.io/%s/json", ip.String())
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, apiURL, nil)
if err != nil {
return nil, fmt.Errorf("failed to create request for ipinfo.io: %w", err)
}
resp, err := s.Client.Do(req)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ipinfo.io: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("ipinfo.io returned non-200 status: %d %s", resp.StatusCode, resp.Status)
}
var apiResp ipInfoResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ipinfo.io response: %w", err)
}
latitude, longitude := parseIPInfoCoordinates(apiResp.Loc)
// IPinfo 的 "country" 字段直接返回 ISO 2-letter code,例如 "US", "CN"
// 我们需要将 "country" 字段作为 ISOCode,并尝试获取其对应的国家名称。
// IPinfo 响应中通常不直接提供完整的国家名称,但我们可以通过 CountryCode 映射。
// 为了简化并符合 GeoInfo 结构,我们直接使用 Country 作为 ISOCode,并尝试从 CountryCode 获取名称。
// 实际上,IPinfo 的 'country' 字段就是 ISO 2-letter code。
// 如果需要完整的国家名称,可能需要一个本地的 ISO 代码到名称的映射。
// 为了与 GetRegionUnicodeEmoji 函数兼容,我们直接使用 country 作为 ISOCode。
name := formatIPInfoLocation(apiResp)
return &GeoInfo{
ISOCode: apiResp.Country,
Name: name,
Latitude: latitude,
Longitude: longitude,
}, nil
}
// UpdateDatabase 对于 ipinfo.io 是一个空操作,因为它是一个 Web 服务。
func (s *IPInfoService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ipinfo.io 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPInfoService) Close() error {
// 无需执行任何操作
return nil
}
func parseIPInfoCoordinates(value string) (*float64, *float64) {
parts := strings.Split(strings.TrimSpace(value), ",")
if len(parts) != isoCountryCodeLength {
return nil, nil
}
latitudeValue, latErr := strconv.ParseFloat(strings.TrimSpace(parts[0]), 64)
longitudeValue, lonErr := strconv.ParseFloat(strings.TrimSpace(parts[1]), 64)
if latErr != nil || lonErr != nil {
return nil, nil
}
return float64Pointer(latitudeValue), float64Pointer(longitudeValue)
}
func formatIPInfoLocation(resp ipInfoResponse) string {
parts := make([]string, 0, 3) //nolint:mnd // 3 parts: city, region, country
if city := strings.TrimSpace(resp.City); city != "" {
parts = append(parts, city)
}
if region := strings.TrimSpace(resp.Region); region != "" {
parts = append(parts, region)
}
if country := strings.TrimSpace(resp.Country); country != "" {
parts = append(parts, country)
}
if len(parts) == 0 {
return ""
}
return strings.Join(parts, ", ")
}
@@ -0,0 +1,80 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package iputil provides helpers for parsing, normalizing, and scoring IP addresses.
package iputil
import (
"net"
"strings"
)
const (
scorePublic = 2
scorePrivate = 1
)
// NormalizeIP parses and normalizes a raw IP address string, preferring the IPv4 form for mapped addresses.
func NormalizeIP(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
ip := net.ParseIP(trimmed)
if ip == nil {
return ""
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4.String()
}
return ip.String()
}
// NormalizeRemoteAddr extracts and normalizes the IP address from a host:port remote address string.
func NormalizeRemoteAddr(remoteAddr string) string {
trimmed := strings.TrimSpace(remoteAddr)
if trimmed == "" {
return ""
}
if host, _, err := net.SplitHostPort(trimmed); err == nil {
return NormalizeIP(host)
}
return NormalizeIP(trimmed)
}
// IsPublic reports whether the given IP address is a publicly routable unicast address.
func IsPublic(ip net.IP) bool {
if ip == nil {
return false
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsMulticast() || ip.IsUnspecified() {
return false
}
return true
}
// IsPublicString parses raw and reports whether it represents a publicly routable IP address.
func IsPublicString(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
return IsPublic(ip)
}
// Score returns a preference score for the IP address: 2 for public, 1 for private, -1 for invalid or non-unicast.
func Score(ip net.IP) int {
if ip == nil {
return -1
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsMulticast() || ip.IsUnspecified() {
return -1
}
if IsPublic(ip) {
return scorePublic
}
return scorePrivate
}
@@ -0,0 +1,48 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package iputil
import (
"net"
"testing"
)
func TestNormalizeIP(t *testing.T) {
if got := NormalizeIP(" 8.8.8.8 "); got != "8.8.8.8" {
t.Fatalf("unexpected normalized ipv4: %q", got)
}
if got := NormalizeIP("[::1]"); got != "" {
t.Fatalf("expected invalid bracketed host to be rejected, got %q", got)
}
}
func TestNormalizeRemoteAddr(t *testing.T) {
if got := NormalizeRemoteAddr("203.0.113.10:8443"); got != "203.0.113.10" {
t.Fatalf("unexpected remote addr normalization: %q", got)
}
}
func TestIsPublic(t *testing.T) {
if !IsPublic(net.ParseIP("8.8.8.8")) {
t.Fatal("expected public ip to be detected")
}
if IsPublic(net.ParseIP("10.0.0.8")) {
t.Fatal("expected private ip to be rejected")
}
if IsPublic(net.ParseIP("127.0.0.1")) {
t.Fatal("expected loopback ip to be rejected")
}
}
func TestScore(t *testing.T) {
if got := Score(net.ParseIP("8.8.8.8")); got != 2 {
t.Fatalf("unexpected score for public ip: %d", got)
}
if got := Score(net.ParseIP("10.0.0.8")); got != 1 {
t.Fatalf("unexpected score for private ip: %d", got)
}
if got := Score(net.ParseIP("127.0.0.1")); got != -1 {
t.Fatalf("unexpected score for loopback ip: %d", got)
}
}
+193
View File
@@ -0,0 +1,193 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"context"
"errors"
"fmt"
"io"
"net"
"net/http"
"os"
"path/filepath"
"sync"
"github.com/oschwald/maxminddb-golang"
)
// GeoIPURL is the default download URL for the MaxMind GeoLite2 country database.
var GeoIPURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
// GeoIPFilePath is the default local path for the MaxMind country database file.
var GeoIPFilePath = "./data/GeoLite2-Country.mmdb"
// Record is the MaxMind database record structure for country lookups.
type Record struct {
Country struct {
ISOCode string `maxminddb:"iso_code"`
Names map[string]string `maxminddb:"names"`
} `maxminddb:"country"`
}
// MaxMindGeoIPService resolves geographic information using a local MaxMind MMDB database.
type MaxMindGeoIPService struct {
maxMindDBReader *maxminddb.Reader
dbFilePath string
mu sync.RWMutex
}
// Name returns the provider identifier for the MaxMind database service.
func (s *MaxMindGeoIPService) Name() string {
return "MaxMind"
}
// NewMaxMindGeoIPService creates a MaxMind service using the default database path and URL.
func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
return NewMaxMindGeoIPServiceWithContext(context.Background(), GeoIPFilePath, GeoIPURL)
}
// NewMaxMindGeoIPServiceWithContext creates a MaxMind service with a caller-supplied
// context so the initial database download honors request cancellation.
func NewMaxMindGeoIPServiceWithContext(ctx context.Context, dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) {
if dbFilePath == "" {
dbFilePath = GeoIPFilePath
}
if downloadURL == "" {
downloadURL = GeoIPURL
}
service := &MaxMindGeoIPService{
dbFilePath: dbFilePath,
}
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), geoipDataDirPerm); err != nil {
return nil, fmt.Errorf("failed to create data directory for MaxMind database: %w", err)
}
if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) {
if err := DownloadMaxMindDatabase(ctx, service.dbFilePath, downloadURL); err != nil {
return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err)
}
}
if err := service.initialize(); err != nil {
return nil, fmt.Errorf("failed to initialize MaxMind database: %w", err)
}
return service, nil
}
func (s *MaxMindGeoIPService) initialize() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
_ = s.maxMindDBReader.Close()
s.maxMindDBReader = nil
}
reader, err := maxminddb.Open(s.dbFilePath)
if err != nil {
return fmt.Errorf("error opening MaxMind database at %s: %w", s.dbFilePath, err)
}
s.maxMindDBReader = reader
return nil
}
// GetGeoInfo looks up geographic information for ip in the MaxMind database.
func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
s.mu.RLock()
defer s.mu.RUnlock()
if s.maxMindDBReader == nil {
return nil, errors.New("MaxMind database is not initialized or failed to open")
}
if ip == nil {
return nil, errors.New("IP address cannot be nil")
}
var record Record
if err := s.maxMindDBReader.Lookup(ip, &record); err != nil {
return nil, fmt.Errorf("error looking up IP %s in MaxMind database: %w", ip.String(), err)
}
geoInfo := &GeoInfo{
ISOCode: record.Country.ISOCode,
Name: record.Country.Names["en"],
}
if geoInfo.Name == "" && geoInfo.ISOCode != "" {
geoInfo.Name = geoInfo.ISOCode
}
return geoInfo, nil
}
// UpdateDatabase downloads the latest MaxMind database and reloads the reader.
func (s *MaxMindGeoIPService) UpdateDatabase() error {
if err := DownloadMaxMindDatabase(context.Background(), s.dbFilePath, GeoIPURL); err != nil {
return err
}
return s.initialize()
}
// DownloadMaxMindDatabase downloads the MaxMind database from downloadURL to dbFilePath.
func DownloadMaxMindDatabase(ctx context.Context, dbFilePath string, downloadURL string) error {
if dbFilePath == "" {
dbFilePath = GeoIPFilePath
}
if downloadURL == "" {
downloadURL = GeoIPURL
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil) //nolint:gosec // URL from trusted GeoIP provider config
if err != nil {
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status)
}
if err := os.MkdirAll(filepath.Dir(dbFilePath), geoipDataDirPerm); err != nil {
return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err)
}
tempPath := dbFilePath + ".download"
out, err := os.Create(tempPath) //nolint:gosec // tempPath is derived from configured dbFilePath
if err != nil {
return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err)
}
defer func() {
_ = out.Close()
}()
if _, err = io.Copy(out, resp.Body); err != nil {
return fmt.Errorf("failed to write MaxMind database file: %w", err)
}
if err = out.Close(); err != nil {
return fmt.Errorf("failed to close MaxMind database file: %w", err)
}
if err = os.Rename(tempPath, dbFilePath); err != nil {
return fmt.Errorf("failed to move MaxMind database file into place: %w", err)
}
return nil
}
// Close closes the MaxMind database reader.
func (s *MaxMindGeoIPService) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
err := s.maxMindDBReader.Close()
s.maxMindDBReader = nil
if err != nil {
return fmt.Errorf("error closing MaxMind database: %w", err)
}
}
return nil
}
+258
View File
@@ -0,0 +1,258 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
"Wavelet/openflare/share/geoip/iputil"
)
const defaultOutboundIPLookupTimeout = 5 * time.Second
// OutboundIPStrategy defines a lookup strategy for the current public egress IP.
type OutboundIPStrategy interface {
Name() string
GetOutboundIP(ctx context.Context) (net.IP, error)
}
// OutboundIPAPIAdapter adapts a third-party HTTP API response into an IP value.
type OutboundIPAPIAdapter interface {
Name() string
Endpoint() string
DecodeIP(io.Reader) (net.IP, error)
}
// HTTPOutboundIPStrategy resolves the public egress IP via an HTTP API adapter.
type HTTPOutboundIPStrategy struct {
Client *http.Client
Adapter OutboundIPAPIAdapter
}
// NewHTTPOutboundIPStrategy creates a strategy that queries adapter over HTTP.
func NewHTTPOutboundIPStrategy(adapter OutboundIPAPIAdapter, client *http.Client) *HTTPOutboundIPStrategy {
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
return &HTTPOutboundIPStrategy{
Client: client,
Adapter: adapter,
}
}
// Name returns the strategy or adapter identifier.
func (s *HTTPOutboundIPStrategy) Name() string {
if s == nil || s.Adapter == nil {
return "http-outbound-ip"
}
return s.Adapter.Name()
}
// GetOutboundIP queries the configured HTTP endpoint for the current public IP.
func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
if s == nil || s.Adapter == nil {
return nil, errors.New("outbound IP adapter is nil")
}
if ctx == nil {
return nil, errors.New("context is required")
}
if s.Client != nil {
return s.query(ctx, s.Client)
}
dialer := &net.Dialer{
Timeout: defaultOutboundIPLookupTimeout,
KeepAlive: 30 * time.Second,
}
// Try IPv4 first
ipv4Client := &http.Client{
Timeout: defaultOutboundIPLookupTimeout,
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: func(ctx context.Context, _, addr string) (net.Conn, error) {
return dialer.DialContext(ctx, "tcp4", addr)
},
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
},
}
ip, err := s.query(ctx, ipv4Client)
if err == nil && ip != nil {
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
}
// Fallback to standard client (dual-stack: tcp)
fallbackClient := &http.Client{
Timeout: defaultOutboundIPLookupTimeout,
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: func(ctx context.Context, _, addr string) (net.Conn, error) {
return dialer.DialContext(ctx, "tcp", addr)
},
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
},
}
ip, err = s.query(ctx, fallbackClient)
if err != nil {
return nil, err
}
// Always prioritize IPv4 for relay compatibility
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
// Return IPv6 only if no IPv4 is available
return ip, nil
}
func (s *HTTPOutboundIPStrategy) query(ctx context.Context, client *http.Client) (net.IP, error) {
request, err := http.NewRequestWithContext(ctx, http.MethodGet, s.Adapter.Endpoint(), nil)
if err != nil {
return nil, fmt.Errorf("%s create request failed: %w", s.Name(), err)
}
response, err := client.Do(request)
if err != nil {
return nil, fmt.Errorf("%s request failed: %w", s.Name(), err)
}
defer func() { _ = response.Body.Close() }()
if response.StatusCode != http.StatusOK {
return nil, fmt.Errorf("%s returned non-200 status: %d %s", s.Name(), response.StatusCode, response.Status)
}
ip, err := s.Adapter.DecodeIP(response.Body)
if err != nil {
return nil, fmt.Errorf("%s decode response failed: %w", s.Name(), err)
}
if !iputil.IsPublic(ip) {
return nil, fmt.Errorf("%s returned non-public IP: %s", s.Name(), ip.String())
}
return ip, nil
}
// RealIPCCAdapter decodes public IP responses from realip.cc.
type RealIPCCAdapter struct {
URL string
}
type realIPCCResponse struct {
IP string `json:"ip"`
}
// NewRealIPCCOutboundIPStrategy creates the default realip.cc lookup strategy.
func NewRealIPCCOutboundIPStrategy() *HTTPOutboundIPStrategy {
return NewHTTPOutboundIPStrategy(RealIPCCAdapter{}, nil)
}
// Name returns the realip.cc adapter identifier.
func (a RealIPCCAdapter) Name() string {
return "realip.cc"
}
// Endpoint returns the realip.cc API URL.
func (a RealIPCCAdapter) Endpoint() string {
if strings.TrimSpace(a.URL) != "" {
return strings.TrimSpace(a.URL)
}
return "https://realip.cc"
}
// DecodeIP parses a realip.cc JSON response into a public IP address.
func (a RealIPCCAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
var payload realIPCCResponse
if err := json.NewDecoder(reader).Decode(&payload); err != nil {
return nil, err
}
ip := net.ParseIP(strings.TrimSpace(payload.IP))
if ip == nil {
return nil, fmt.Errorf("invalid IP %q", payload.IP)
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
return ip, nil
}
// PlainTextIPAdapter decodes public IP responses from raw plain text endpoints.
type PlainTextIPAdapter struct {
ProviderName string
URL string
}
// Name returns the provider name.
func (a PlainTextIPAdapter) Name() string {
return a.ProviderName
}
// Endpoint returns the plain text API URL.
func (a PlainTextIPAdapter) Endpoint() string {
return a.URL
}
// DecodeIP parses a plain text response into a public IP address.
func (a PlainTextIPAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
body, err := io.ReadAll(reader)
if err != nil {
return nil, err
}
ipStr := strings.TrimSpace(string(body))
ip := net.ParseIP(ipStr)
if ip == nil {
return nil, fmt.Errorf("invalid plain text IP %q", ipStr)
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
return ip, nil
}
// DefaultOutboundIPStrategies returns the built-in public egress IP lookup strategies.
func DefaultOutboundIPStrategies() []OutboundIPStrategy {
return []OutboundIPStrategy{
NewRealIPCCOutboundIPStrategy(),
NewHTTPOutboundIPStrategy(PlainTextIPAdapter{ProviderName: "ifconfig.me", URL: "https://ifconfig.me"}, nil),
NewHTTPOutboundIPStrategy(PlainTextIPAdapter{ProviderName: "ip.sb", URL: "https://api.ip.sb/ip"}, nil),
NewHTTPOutboundIPStrategy(PlainTextIPAdapter{ProviderName: "icanhazip.com", URL: "https://icanhazip.com"}, nil),
}
}
// GetOutboundIP tries each strategy until one returns a public egress IP.
func GetOutboundIP(ctx context.Context, strategies ...OutboundIPStrategy) (net.IP, error) {
if len(strategies) == 0 {
strategies = DefaultOutboundIPStrategies()
}
var errs []error
for _, strategy := range strategies {
if strategy == nil {
continue
}
ip, err := strategy.GetOutboundIP(ctx)
if err == nil && ip != nil {
return ip, nil
}
if err != nil {
errs = append(errs, fmt.Errorf("%s: %w", strategy.Name(), err))
}
}
if len(errs) == 0 {
return nil, errors.New("no outbound IP lookup strategy configured")
}
return nil, errors.Join(errs...)
}
@@ -0,0 +1,127 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package geoip
import (
"context"
"errors"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
type fakeOutboundIPStrategy struct {
name string
ip net.IP
err error
}
func (f fakeOutboundIPStrategy) Name() string {
return f.name
}
func (f fakeOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
return f.ip, f.err
}
func TestRealIPCCAdapterDecodeIP(t *testing.T) {
ip, err := RealIPCCAdapter{}.DecodeIP(strings.NewReader(`{"ip":"8.8.8.8","country":"United States"}`))
if err != nil {
t.Fatalf("DecodeIP failed: %v", err)
}
if ip.String() != "8.8.8.8" {
t.Fatalf("unexpected IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyUsesAdapter(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
t.Fatalf("unexpected method: %s", r.Method)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"8.8.4.4"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
ip, err := strategy.GetOutboundIP(context.Background())
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "8.8.4.4" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestGetOutboundIPFallsBackToNextStrategy(t *testing.T) {
ip, err := GetOutboundIP(
context.Background(),
fakeOutboundIPStrategy{name: "first", err: errors.New("temporary failure")},
fakeOutboundIPStrategy{name: "second", ip: net.ParseIP("1.1.1.1")},
)
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "1.1.1.1" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyRejectsPrivateIP(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"172.17.0.2"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
if _, err := strategy.GetOutboundIP(context.Background()); err == nil {
t.Fatal("expected private IP to be rejected")
}
}
func TestHTTPOutboundIPStrategyPrioritizesIPv4(t *testing.T) {
tests := []struct {
name string
response string
want string
}{
{
name: "IPv4 address",
response: `{"ip":"8.8.8.8"}`,
want: "8.8.8.8",
},
{
name: "IPv6 address normalized to IPv4",
response: `{"ip":"::ffff:8.8.8.8"}`,
want: "8.8.8.8",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(tt.response))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
ip, err := strategy.GetOutboundIP(context.Background())
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != tt.want {
t.Errorf("GetOutboundIP() = %v, want %v", ip.String(), tt.want)
}
// Ensure it's an IPv4 address (4 bytes)
if ipv4 := ip.To4(); ipv4 == nil {
t.Errorf("expected IPv4 address, got %v", ip)
}
})
}
}
+31
View File
@@ -0,0 +1,31 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ofutil
import "strings"
// UniqueAndCleanStringSlice trims spaces, drops empties, and de-duplicates
// while preserving order. An empty result is nil.
func UniqueAndCleanStringSlice(slice []string) []string {
if slice == nil {
return nil
}
seen := make(map[string]struct{})
result := make([]string, 0)
for _, item := range slice {
trimmed := strings.TrimSpace(item)
if trimmed == "" {
continue
}
if _, ok := seen[trimmed]; ok {
continue
}
seen[trimmed] = struct{}{}
result = append(result, trimmed)
}
if len(result) == 0 {
return nil
}
return result
}
+227
View File
@@ -0,0 +1,227 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ofutil
import (
"strconv"
"strings"
)
const gitDescribeMinIdentifiers = 2
type versionInfo struct {
valid bool
isDev bool
numbers []int
prerelease []string
gitDescribeDistance int
gitDescribeTail []string
}
func parseVersionInfo(version string) versionInfo {
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
if normalized == "" || normalized == "dev" {
return versionInfo{isDev: strings.EqualFold(normalized, "dev")}
}
base := normalized
prerelease := ""
if separator := strings.IndexRune(normalized, '-'); separator >= 0 {
base = normalized[:separator]
prerelease = normalized[separator+1:]
}
segments := strings.Split(base, ".")
parts := make([]int, 0, len(segments))
for _, segment := range segments {
segment = strings.TrimSpace(segment)
if segment == "" {
parts = append(parts, 0)
continue
}
numeric := strings.Builder{}
for _, r := range segment {
if r < '0' || r > '9' {
break
}
numeric.WriteRune(r)
}
if numeric.Len() == 0 {
parts = append(parts, 0)
continue
}
value, err := strconv.Atoi(numeric.String())
if err != nil {
return versionInfo{}
}
parts = append(parts, value)
}
info := versionInfo{valid: len(parts) > 0, numbers: parts}
if prerelease != "" {
identifiers := splitPrereleaseIdentifiers(prerelease)
if distance, tail, ok := parseGitDescribeIdentifiers(identifiers); ok {
info.gitDescribeDistance = distance
info.gitDescribeTail = tail
} else {
info.prerelease = identifiers
}
}
return info
}
func parseGitDescribeIdentifiers(identifiers []string) (int, []string, bool) {
if len(identifiers) < gitDescribeMinIdentifiers {
return 0, nil, false
}
distance, err := strconv.Atoi(strings.TrimSpace(identifiers[0]))
if err != nil || distance <= 0 {
return 0, nil, false
}
commitToken := strings.TrimSpace(identifiers[1])
if commitToken == "" || !strings.HasPrefix(strings.ToLower(commitToken), "g") {
return 0, nil, false
}
return distance, identifiers[1:], true
}
func splitPrereleaseIdentifiers(value string) []string {
parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool {
return r == '.' || r == '-'
})
filtered := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
filtered = append(filtered, part)
}
}
return filtered
}
func CompareVersions(local, remote string) int {
left := parseVersionInfo(local)
right := parseVersionInfo(remote)
if left.isDev {
if right.valid {
return -1
}
return 0
}
if !left.valid || !right.valid {
return 0
}
if result := compareVersionNumbers(left, right); result != 0 {
return result
}
if result := compareGitDescribeDistance(left, right); result != 0 {
return result
}
if left.gitDescribeDistance > 0 || right.gitDescribeDistance > 0 {
return compareGitDescribeTails(left, right)
}
return comparePrereleaseIdentifiers(left, right)
}
func compareVersionNumbers(left, right versionInfo) int {
maxLen := max(len(right.numbers), len(left.numbers))
for index := range maxLen {
leftValue := 0
rightValue := 0
if index < len(left.numbers) {
leftValue = left.numbers[index]
}
if index < len(right.numbers) {
rightValue = right.numbers[index]
}
if leftValue < rightValue {
return -1
}
if leftValue > rightValue {
return 1
}
}
return 0
}
func compareGitDescribeDistance(left, right versionInfo) int {
if left.gitDescribeDistance == right.gitDescribeDistance {
return 0
}
if left.gitDescribeDistance < right.gitDescribeDistance {
return -1
}
return 1
}
func compareGitDescribeTails(left, right versionInfo) int {
maxLen := max(len(right.gitDescribeTail), len(left.gitDescribeTail))
for index := range maxLen {
if index >= len(left.gitDescribeTail) {
return -1
}
if index >= len(right.gitDescribeTail) {
return 1
}
if left.gitDescribeTail[index] < right.gitDescribeTail[index] {
return -1
}
if left.gitDescribeTail[index] > right.gitDescribeTail[index] {
return 1
}
}
return 0
}
func comparePrereleaseIdentifiers(left, right versionInfo) int {
if len(left.prerelease) == 0 && len(right.prerelease) == 0 {
return 0
}
if len(left.prerelease) == 0 {
return 1
}
if len(right.prerelease) == 0 {
return -1
}
maxLen := max(len(right.prerelease), len(left.prerelease))
for index := range maxLen {
if index >= len(left.prerelease) {
return -1
}
if index >= len(right.prerelease) {
return 1
}
if result := comparePrereleasePart(left.prerelease[index], right.prerelease[index]); result != 0 {
return result
}
}
return 0
}
func comparePrereleasePart(leftPart, rightPart string) int {
leftNumber, leftErr := strconv.Atoi(leftPart)
rightNumber, rightErr := strconv.Atoi(rightPart)
switch {
case leftErr == nil && rightErr == nil:
if leftNumber < rightNumber {
return -1
}
if leftNumber > rightNumber {
return 1
}
case leftErr == nil:
return -1
case rightErr == nil:
return 1
default:
if leftPart < rightPart {
return -1
}
if leftPart > rightPart {
return 1
}
}
return 0
}
@@ -0,0 +1,110 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"errors"
"fmt"
"io"
"math"
"os"
"path/filepath"
)
// Limits bounds archive inspection / extraction work.
type Limits struct {
// MaxFiles is the maximum number of regular files allowed.
MaxFiles int
// MaxFileBytes is the maximum size of a single extracted file.
MaxFileBytes int64
// MaxTotalBytes is the maximum sum of all extracted file sizes.
MaxTotalBytes int64
}
// FileEntry is a regular file discovered inside a deployment package.
type FileEntry struct {
Path string
Size int64
// Checksum is retained for API/schema compatibility and is left empty.
// Integrity is enforced via the whole-package SHA-256 on the deployment record.
Checksum string
}
// Manifest is the inspected content of a Pages deployment package.
type Manifest struct {
Files []FileEntry
FileCount int
TotalSize int64
}
// Entry describes one archive member for extraction.
type Entry struct {
// Name is the original path inside the archive.
Name string
// IsDir marks directory entries.
IsDir bool
// IsSymlink marks symbolic links (unsupported for Pages).
IsSymlink bool
// IsHardlink marks hard links (unsupported for Pages).
IsHardlink bool
// IsSpecial marks device, FIFO, socket, and other non-regular entries.
IsSpecial bool
// Size is the archive-declared uncompressed size; 0 means an empty member.
Size uint64
// Open returns a reader for the entry body. Caller must Close it.
Open func() (io.ReadCloser, error)
}
// copyLimited copies actual bytes from src. maxBytes < 0 disables the byte cap;
// maxBytes == 0 permits only an empty stream.
func copyLimited(dst io.Writer, src io.Reader, maxBytes int64) (int64, error) {
if maxBytes < 0 {
return io.Copy(dst, src)
}
readLimit := maxBytes
if maxBytes < math.MaxInt64 {
readLimit++
}
written, err := io.Copy(dst, io.LimitReader(src, readLimit))
if err != nil {
return written, err
}
if written > maxBytes {
return written, errors.New("pages file size out of bounds")
}
return written, nil
}
func copyAndVerifySize(dst io.Writer, src io.Reader, declaredSize uint64, maxBytes int64) (int64, error) {
if declaredSize > uint64(math.MaxInt64) {
return 0, errors.New("pages file size out of bounds")
}
written, err := copyLimited(dst, src, maxBytes)
if err != nil {
return written, err
}
//nolint:gosec // declaredSize is bounded to MaxInt64 above
if written != int64(declaredSize) {
return written, fmt.Errorf("pages declared size %d does not match actual %d", declaredSize, written)
}
return written, nil
}
func writeEntryFile(targetPath string, src io.Reader, declaredSize uint64, maxBytes int64, perm os.FileMode) (int64, error) {
if err := os.MkdirAll(filepath.Dir(targetPath), dirPerm); err != nil {
return 0, err
}
target, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, perm) //nolint:gosec // caller validates path under release dir
if err != nil {
return 0, err
}
written, copyErr := copyAndVerifySize(target, src, declaredSize, maxBytes)
closeErr := target.Close()
if err := errors.Join(copyErr, closeErr); err != nil {
_ = os.Remove(targetPath)
return written, err
}
return written, nil
}
@@ -0,0 +1,255 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"archive/tar"
"bytes"
"errors"
"fmt"
"io"
"os"
"path/filepath"
)
// formatDetectHeadBytes is the sniff window used for archive format detection.
const formatDetectHeadBytes = 512
// ExtractOptions controls package extraction.
type ExtractOptions struct {
// Limits bounds actual files and sizes during extraction when EnforceLimits is true.
Limits Limits
// StripCommonRoot strips a single shared top-level directory when present.
StripCommonRoot bool
// EnforceLimits enables MaxFiles / MaxFileBytes / MaxTotalBytes checks.
// Path, member type, and declared/actual-size validation always remain enabled.
EnforceLimits bool
}
// ExtractBytes extracts a deployment package into destDir.
func ExtractBytes(data []byte, format Format, destDir string, opts ExtractOptions) error {
if format == "" {
var err error
format, err = DetectFormat("", data)
if err != nil {
return err
}
}
return extractFromReaderAt(bytes.NewReader(data), int64(len(data)), format, destDir, opts)
}
// ExtractFile opens path and extracts it into destDir without buffering the
// whole archive or tar member bodies in memory.
func ExtractFile(filePath string, format Format, destDir string, opts ExtractOptions) error {
file, err := os.Open(filePath) //nolint:gosec // controlled path
if err != nil {
return err
}
defer func() { _ = file.Close() }()
info, err := file.Stat()
if err != nil {
return err
}
if format == "" {
head := make([]byte, formatDetectHeadBytes)
n, readErr := file.ReadAt(head, 0)
if readErr != nil && readErr != io.EOF {
return readErr
}
format, err = DetectFormat(filePath, head[:n])
if err != nil {
return err
}
}
return extractFromReaderAt(file, info.Size(), format, destDir, opts)
}
func extractFromReaderAt(ra io.ReaderAt, size int64, format Format, destDir string, opts ExtractOptions) error {
if isTarFamily(format) {
return extractTarFamilyAt(ra, size, format, destDir, opts)
}
entries, err := listRandomAccessEntriesAt(ra, size, format)
if err != nil {
return err
}
return extractEntries(entries, destDir, opts)
}
func extractEntries(entries []Entry, destDir string, opts ExtractOptions) error {
limits := Limits{}
if opts.EnforceLimits {
limits = normalizeLimits(opts.Limits)
}
commonPrefix, err := commonRootForEntries(entries, opts.StripCommonRoot)
if err != nil {
return err
}
measured := &measuredArchive{files: make([]measuredFile, 0, len(entries))}
for _, entry := range entries {
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return err
}
if skip {
continue
}
normalizedPath = StripPrefix(normalizedPath, commonPrefix)
if normalizedPath == "" {
continue
}
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, opts.EnforceLimits); err != nil {
return err
}
if entry.Open == nil {
return fmt.Errorf("%s: pages archive member cannot be opened", normalizedPath)
}
src, err := entry.Open()
if err != nil {
return fmt.Errorf("%s: %w", normalizedPath, err)
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, opts.EnforceLimits)
targetPath, err := safeExtractionTarget(destDir, normalizedPath)
if err != nil {
_ = src.Close()
return err
}
actual, writeErr := writeEntryFile(targetPath, src, entry.Size, maxBytes, filePerm)
closeErr := src.Close()
if err := errors.Join(writeErr, closeErr); err != nil {
return fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
if measured.fileCount == 0 {
return errors.New("pages package is empty")
}
return nil
}
func extractTarFamilyAt(ra io.ReaderAt, size int64, format Format, destDir string, opts ExtractOptions) error {
limits := Limits{}
if opts.EnforceLimits {
limits = normalizeLimits(opts.Limits)
}
firstPass, err := scanTarFamilyAt(ra, size, format, limits, opts.EnforceLimits)
if err != nil {
return err
}
if firstPass.fileCount == 0 {
return errors.New("pages package is empty")
}
commonPrefix := ""
if opts.StripCommonRoot {
paths := make([]string, 0, len(firstPass.files))
for _, file := range firstPass.files {
paths = append(paths, file.path)
}
commonPrefix = FindCommonRootPrefix(paths)
}
tarReader, closeReader, err := openTarFamilyReader(io.NewSectionReader(ra, 0, size), format)
if err != nil {
return err
}
secondPass, extractErr := extractTarReader(tarReader, destDir, commonPrefix, limits, opts.EnforceLimits)
if closeErr := closeReader(); closeErr != nil {
extractErr = errors.Join(extractErr, closeErr)
}
if extractErr != nil {
return extractErr
}
if secondPass.fileCount != firstPass.fileCount || secondPass.totalSize != firstPass.totalSize {
return errors.New("pages tar package changed between validation and extraction")
}
return nil
}
func extractTarReader(
tarReader *tar.Reader,
destDir string,
commonPrefix string,
limits Limits,
enforceLimits bool,
) (*measuredArchive, error) {
measured := &measuredArchive{files: make([]measuredFile, 0)}
for {
header, err := tarReader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return nil, fmt.Errorf("read tar pages package: %w", err)
}
entry := entryFromTarHeader(header)
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return nil, err
}
if skip {
continue
}
normalizedPath = StripPrefix(normalizedPath, commonPrefix)
if normalizedPath == "" {
continue
}
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, enforceLimits); err != nil {
return nil, err
}
targetPath, err := safeExtractionTarget(destDir, normalizedPath)
if err != nil {
return nil, err
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, enforceLimits)
actual, err := writeEntryFile(targetPath, tarReader, entry.Size, maxBytes, filePerm)
if err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
return measured, nil
}
func commonRootForEntries(entries []Entry, strip bool) (string, error) {
paths := make([]string, 0, len(entries))
for _, entry := range entries {
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return "", err
}
if !skip {
paths = append(paths, normalizedPath)
}
}
if !strip {
return "", nil
}
return FindCommonRootPrefix(paths), nil
}
func safeExtractionTarget(destDir, relativePath string) (string, error) {
targetPath := filepath.Join(destDir, filepath.FromSlash(relativePath))
if !isWithinDir(destDir, targetPath) {
return "", fmt.Errorf("pages package path escapes directory: %s", relativePath)
}
return targetPath, nil
}
func isWithinDir(baseDir, targetPath string) bool {
cleanBase := filepath.Clean(baseDir)
cleanTarget := filepath.Clean(targetPath)
rel, err := filepath.Rel(cleanBase, cleanTarget)
if err != nil {
return false
}
return rel != ".." && !hasParentRel(rel)
}
func hasParentRel(rel string) bool {
if rel == ".." {
return true
}
return len(rel) >= 3 && (rel[:3] == "../" || rel[:3] == "..\\")
}
@@ -0,0 +1,191 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package pagesarchive provides multi-format archive detection, inspection and
// extraction helpers for OpenFlare Pages deployment packages.
package pagesarchive
import (
"bytes"
"errors"
"path/filepath"
"strings"
)
// Format identifies a supported Pages deployment package archive format.
type Format string
const (
// FormatZip is a ZIP archive.
FormatZip Format = "zip"
// FormatTar is an uncompressed tar archive.
FormatTar Format = "tar"
// FormatTarGz is a gzip-compressed tar archive.
FormatTarGz Format = "tar.gz"
// FormatTarXz is an xz-compressed tar archive.
FormatTarXz Format = "tar.xz"
// FormatTarBz2 is a bzip2-compressed tar archive.
FormatTarBz2 Format = "tar.bz2"
// FormatSevenZip is a 7z archive.
FormatSevenZip Format = "7z"
ustarMagicOffset = 257
ustarMagicMinLen = 262
)
// DetectFormatFromName returns the archive format inferred from a file name.
// Returns an empty Format and false when the extension is unsupported.
func DetectFormatFromName(fileName string) (Format, bool) {
name := strings.ToLower(strings.TrimSpace(fileName))
switch {
case strings.HasSuffix(name, ".tar.gz"), strings.HasSuffix(name, ".tgz"):
return FormatTarGz, true
case strings.HasSuffix(name, ".tar.xz"), strings.HasSuffix(name, ".txz"):
return FormatTarXz, true
case strings.HasSuffix(name, ".tar.bz2"), strings.HasSuffix(name, ".tbz2"), strings.HasSuffix(name, ".tbz"):
return FormatTarBz2, true
case strings.HasSuffix(name, ".tar"):
return FormatTar, true
case strings.HasSuffix(name, ".7z"):
return FormatSevenZip, true
case strings.HasSuffix(name, ".zip"):
return FormatZip, true
default:
return "", false
}
}
// DetectFormatFromBytes returns the archive format inferred from magic bytes.
// Prefer DetectFormatFromName when a reliable file name is available.
func DetectFormatFromBytes(data []byte) (Format, bool) {
if isZipMagic(data) {
return FormatZip, true
}
if isSevenZipMagic(data) {
return FormatSevenZip, true
}
if isXZMagic(data) {
return FormatTarXz, true
}
if isGzipMagic(data) {
return FormatTarGz, true
}
if isBzip2Magic(data) {
return FormatTarBz2, true
}
if looksLikeTar(data) {
return FormatTar, true
}
return "", false
}
// DetectFormat prefers the file name when present, otherwise magic bytes.
func DetectFormat(fileName string, data []byte) (Format, error) {
if format, ok := DetectFormatFromName(fileName); ok {
return format, nil
}
if format, ok := DetectFormatFromBytes(data); ok {
return format, nil
}
return "", errors.New("unsupported pages package format")
}
// Extension returns the canonical file extension for a format (without leading dot).
func Extension(format Format) string {
switch format {
case FormatZip:
return "zip"
case FormatTar:
return "tar"
case FormatTarGz:
return "tar.gz"
case FormatTarXz:
return "tar.xz"
case FormatTarBz2:
return "tar.bz2"
case FormatSevenZip:
return "7z"
default:
return "bin"
}
}
// MIMEType returns a reasonable content type for the archive format.
func MIMEType(format Format) string {
switch format {
case FormatZip:
return "application/zip"
case FormatTar:
return "application/x-tar"
case FormatTarGz:
return "application/gzip"
case FormatTarXz:
return "application/x-xz"
case FormatTarBz2:
return "application/x-bzip2"
case FormatSevenZip:
return "application/x-7z-compressed"
default:
return "application/octet-stream"
}
}
// SupportedExtensions lists human-readable extensions for UI copy and accept attributes.
func SupportedExtensions() []string {
return []string{".zip", ".tar.gz", ".tgz", ".tar.xz", ".txz", ".tar.bz2", ".tbz2", ".tar", ".7z"}
}
// AcceptAttribute returns a comma-separated accept list for file inputs.
func AcceptAttribute() string {
return strings.Join(SupportedExtensions(), ",")
}
// NormalizeNameExtension returns a storage-safe extension for the given format/name.
func NormalizeNameExtension(fileName string, format Format) string {
if format != "" {
return Extension(format)
}
if formatFromName, ok := DetectFormatFromName(fileName); ok {
return Extension(formatFromName)
}
ext := strings.TrimPrefix(filepath.Ext(fileName), ".")
if ext == "" {
return "bin"
}
return strings.ToLower(ext)
}
func isZipMagic(data []byte) bool {
return len(data) >= 4 &&
data[0] == 0x50 && data[1] == 0x4b &&
(data[2] == 0x03 || data[2] == 0x05 || data[2] == 0x07) &&
(data[3] == 0x04 || data[3] == 0x06 || data[3] == 0x08)
}
func isSevenZipMagic(data []byte) bool {
return len(data) >= 6 &&
data[0] == 0x37 && data[1] == 0x7a && data[2] == 0xbc &&
data[3] == 0xaf && data[4] == 0x27 && data[5] == 0x1c
}
func isXZMagic(data []byte) bool {
return len(data) >= 6 &&
data[0] == 0xfd && data[1] == 0x37 && data[2] == 0x7a &&
data[3] == 0x58 && data[4] == 0x5a && data[5] == 0x00
}
func isGzipMagic(data []byte) bool {
return len(data) >= 2 && data[0] == 0x1f && data[1] == 0x8b
}
func isBzip2Magic(data []byte) bool {
return len(data) >= 3 && data[0] == 0x42 && data[1] == 0x5a && data[2] == 0x68
}
func looksLikeTar(data []byte) bool {
// POSIX ustar magic at offset 257 ("ustar\0" or "ustar ").
if len(data) < ustarMagicMinLen {
return false
}
return bytes.Equal(data[ustarMagicOffset:ustarMagicOffset+5], []byte("ustar"))
}
@@ -0,0 +1,303 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"archive/tar"
"bytes"
"errors"
"fmt"
"io"
"math"
"os"
"path"
)
// InspectOptions controls package inspection.
type InspectOptions struct {
// RootDir is an optional project root subdirectory that must contain EntryFile.
RootDir string
// EntryFile is the required entry file name (e.g. index.html).
EntryFile string
// Limits bounds files and actual extracted sizes.
Limits Limits
// VerifySizes is retained for source compatibility. Inspection now always
// streams regular members and verifies actual bytes against declared sizes.
VerifySizes bool
}
type measuredFile struct {
path string
size int64
}
type measuredArchive struct {
files []measuredFile
fileCount int
totalSize int64
}
// InspectFile opens path and inspects it as a Pages deployment package without
// loading the whole archive or any tar member body into memory.
func InspectFile(filePath string, format Format, opts InspectOptions) (*Manifest, error) {
file, err := os.Open(filePath) //nolint:gosec // filePath is a controlled temp upload path
if err != nil {
return nil, err
}
defer func() { _ = file.Close() }()
info, err := file.Stat()
if err != nil {
return nil, err
}
if format == "" {
head := make([]byte, formatDetectHeadBytes)
n, readErr := file.ReadAt(head, 0)
if readErr != nil && readErr != io.EOF {
return nil, readErr
}
format, err = DetectFormat(filePath, head[:n])
if err != nil {
return nil, err
}
}
return inspectFromReaderAt(file, info.Size(), format, opts)
}
// InspectBytes inspects an in-memory deployment package.
func InspectBytes(data []byte, format Format, opts InspectOptions) (*Manifest, error) {
if format == "" {
var err error
format, err = DetectFormat("", data)
if err != nil {
return nil, err
}
}
return inspectFromReaderAt(bytes.NewReader(data), int64(len(data)), format, opts)
}
func inspectFromReaderAt(ra io.ReaderAt, size int64, format Format, opts InspectOptions) (*Manifest, error) {
limits := normalizeLimits(opts.Limits)
var (
measured *measuredArchive
err error
)
if isTarFamily(format) {
measured, err = scanTarFamilyAt(ra, size, format, limits, true)
} else {
var entries []Entry
entries, err = listRandomAccessEntriesAt(ra, size, format)
if err == nil {
measured, err = inspectRandomAccessEntries(entries, limits)
}
}
if err != nil {
return nil, err
}
return buildMeasuredManifest(measured, opts)
}
func inspectRandomAccessEntries(entries []Entry, limits Limits) (*measuredArchive, error) {
measured := &measuredArchive{files: make([]measuredFile, 0, len(entries))}
for _, entry := range entries {
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return nil, err
}
if skip {
continue
}
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, true); err != nil {
return nil, err
}
if entry.Open == nil {
return nil, fmt.Errorf("%s: pages archive member cannot be opened", normalizedPath)
}
src, err := entry.Open()
if err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, true)
actual, copyErr := copyAndVerifySize(io.Discard, src, entry.Size, maxBytes)
closeErr := src.Close()
if err := errors.Join(copyErr, closeErr); err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
return measured, nil
}
func scanTarFamilyAt(
ra io.ReaderAt,
size int64,
format Format,
limits Limits,
enforceLimits bool,
) (*measuredArchive, error) {
if size < 0 {
return nil, errors.New("invalid pages package size")
}
tarReader, closeReader, err := openTarFamilyReader(io.NewSectionReader(ra, 0, size), format)
if err != nil {
return nil, err
}
measured, scanErr := scanTarReader(tarReader, limits, enforceLimits)
if closeErr := closeReader(); closeErr != nil {
scanErr = errors.Join(scanErr, closeErr)
}
return measured, scanErr
}
func scanTarReader(tarReader *tar.Reader, limits Limits, enforceLimits bool) (*measuredArchive, error) {
measured := &measuredArchive{files: make([]measuredFile, 0)}
for {
header, err := tarReader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return nil, fmt.Errorf("read tar pages package: %w", err)
}
entry := entryFromTarHeader(header)
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return nil, err
}
if skip {
continue
}
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, enforceLimits); err != nil {
return nil, err
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, enforceLimits)
actual, err := copyAndVerifySize(io.Discard, tarReader, entry.Size, maxBytes)
if err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
return measured, nil
}
func buildMeasuredManifest(measured *measuredArchive, opts InspectOptions) (*Manifest, error) {
if measured == nil || measured.fileCount == 0 {
return nil, errors.New("pages package is empty")
}
targetEntryPath, err := resolveTargetEntryPath(opts.RootDir, opts.EntryFile)
if err != nil {
return nil, err
}
paths := make([]string, 0, len(measured.files))
for _, file := range measured.files {
paths = append(paths, file.path)
}
commonPrefix := FindCommonRootPrefix(paths)
manifest := &Manifest{
Files: make([]FileEntry, 0, measured.fileCount),
FileCount: measured.fileCount,
TotalSize: measured.totalSize,
}
entrySeen := false
for _, file := range measured.files {
normalizedPath := StripPrefix(file.path, commonPrefix)
if normalizedPath == targetEntryPath {
entrySeen = true
}
manifest.Files = append(manifest.Files, FileEntry{
Path: normalizedPath,
Size: file.size,
})
}
if !entrySeen {
return nil, fmt.Errorf("pages package is missing entry file %s", targetEntryPath)
}
return manifest, nil
}
func prepareMeasuredFile(measured *measuredArchive, normalizedPath string, declaredSize uint64, limits Limits, enforceLimits bool) error {
if declaredSize > uint64(math.MaxInt64) {
return fmt.Errorf("%s: pages file size out of bounds", normalizedPath)
}
if !enforceLimits {
return nil
}
if measured.fileCount >= limits.MaxFiles {
return fmt.Errorf("pages deployment file count exceeds %d", limits.MaxFiles)
}
if exceedsFileByteLimit(declaredSize, limits.MaxFileBytes) {
return fmt.Errorf("pages file too large: %s", normalizedPath)
}
remaining := limits.MaxTotalBytes - measured.totalSize
if remaining < 0 || declaredSize > uint64(remaining) { //nolint:gosec // remaining is checked non-negative
return errors.New("pages extracted size exceeds limit")
}
return nil
}
func appendMeasuredFile(measured *measuredArchive, normalizedPath string, actual int64) {
measured.files = append(measured.files, measuredFile{path: normalizedPath, size: actual})
measured.fileCount++
measured.totalSize += actual
}
func effectiveFileLimit(limits Limits, totalSize int64, enforceLimits bool) int64 {
if !enforceLimits {
return -1
}
remaining := limits.MaxTotalBytes - totalSize
if remaining < limits.MaxFileBytes {
return remaining
}
return limits.MaxFileBytes
}
func resolveTargetEntryPath(rootDir, entryFile string) (string, error) {
normalizedRoot, err := NormalizeLogicalPath(rootDir, true)
if err != nil {
return "", fmt.Errorf("invalid pages root directory: %w", err)
}
if entryFile == "" {
entryFile = "index.html"
}
normalizedEntry, err := NormalizeLogicalPath(entryFile, false)
if err != nil {
return "", fmt.Errorf("invalid pages entry file: %w", err)
}
if normalizedRoot == "" {
return normalizedEntry, nil
}
return path.Join(normalizedRoot, normalizedEntry), nil
}
func validateArchiveEntry(entry Entry) (string, bool, error) {
normalizedPath, skip, err := NormalizeEntryPath(entry.Name)
if err != nil {
return "", false, err
}
if entry.IsSymlink {
return "", false, fmt.Errorf("pages package contains unsupported symlink: %s", normalizedPath)
}
if entry.IsHardlink {
return "", false, fmt.Errorf("pages package contains unsupported hardlink: %s", normalizedPath)
}
if entry.IsSpecial {
return "", false, fmt.Errorf("pages package contains unsupported special entry: %s", normalizedPath)
}
if skip || entry.IsDir {
return normalizedPath, true, nil
}
return normalizedPath, false, nil
}
func isTarFamily(format Format) bool {
switch format {
case FormatTar, FormatTarGz, FormatTarXz, FormatTarBz2:
return true
case FormatZip, FormatSevenZip:
return false
default:
return false
}
}
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
const (
defaultMaxFiles = 1000
defaultMaxFileBytes = 100 * 1024 * 1024
defaultMaxTotalBytes = 100 * 1024 * 1024
dirPerm = 0o750
filePerm = 0o644
)
// normalizeLimits applies defaults for control-plane inspection.
// Callers that already validated the package should use EnforceLimits=false instead.
func normalizeLimits(limits Limits) Limits {
if limits.MaxFiles <= 0 {
limits.MaxFiles = defaultMaxFiles
}
if limits.MaxFileBytes <= 0 {
limits.MaxFileBytes = defaultMaxFileBytes
}
if limits.MaxTotalBytes <= 0 {
limits.MaxTotalBytes = defaultMaxTotalBytes
}
return limits
}
func exceedsFileByteLimit(size uint64, maxBytes int64) bool {
// maxBytes <= 0 means unlimited (trusted extract path).
if maxBytes <= 0 {
return false
}
if size == 0 {
return false
}
return size > uint64(maxBytes) //nolint:gosec // maxBytes is positive
}
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"archive/tar"
"archive/zip"
"compress/bzip2"
"compress/gzip"
"errors"
"fmt"
"io"
"os"
"github.com/bodgit/sevenzip"
"github.com/ulikunitz/xz"
)
type archiveFile interface {
Name() string
Mode() os.FileMode
IsDir() bool
Size() uint64
Open() (io.ReadCloser, error)
}
type zipArchiveFile struct {
file *zip.File
}
func (z zipArchiveFile) Name() string { return z.file.Name }
func (z zipArchiveFile) Mode() os.FileMode { return z.file.Mode() }
func (z zipArchiveFile) IsDir() bool { return z.file.FileInfo().IsDir() }
func (z zipArchiveFile) Size() uint64 { return z.file.UncompressedSize64 }
func (z zipArchiveFile) Open() (io.ReadCloser, error) {
return z.file.Open()
}
type sevenZipArchiveFile struct {
file *sevenzip.File
}
func (z sevenZipArchiveFile) Name() string { return z.file.Name }
func (z sevenZipArchiveFile) Mode() os.FileMode { return z.file.Mode() }
func (z sevenZipArchiveFile) IsDir() bool { return z.file.FileInfo().IsDir() }
func (z sevenZipArchiveFile) Size() uint64 { return z.file.UncompressedSize }
func (z sevenZipArchiveFile) Open() (io.ReadCloser, error) {
return z.file.Open()
}
// listRandomAccessEntriesAt lists zip/7z members without reading their bodies.
// Tar-family archives use the sequential streaming paths in inspect.go/extract.go.
func listRandomAccessEntriesAt(ra io.ReaderAt, size int64, format Format) ([]Entry, error) {
if size < 0 {
return nil, errors.New("invalid pages package size")
}
switch format {
case FormatZip:
return listZipEntriesAt(ra, size)
case FormatSevenZip:
return listSevenZipEntriesAt(ra, size)
case FormatTar, FormatTarGz, FormatTarXz, FormatTarBz2:
return nil, fmt.Errorf("unsupported random-access pages package format: %s", format)
default:
return nil, fmt.Errorf("unsupported random-access pages package format: %s", format)
}
}
func entriesFromArchiveFiles(files []archiveFile) []Entry {
entries := make([]Entry, 0, len(files))
for _, item := range files {
file := item
mode := file.Mode()
isDir := file.IsDir()
isSymlink := mode&os.ModeSymlink != 0
isSpecial := !isDir && !isSymlink && !mode.IsRegular()
entries = append(entries, Entry{
Name: file.Name(),
IsDir: isDir,
IsSymlink: isSymlink,
IsSpecial: isSpecial,
Size: file.Size(),
Open: file.Open,
})
}
return entries
}
func listZipEntriesAt(ra io.ReaderAt, size int64) ([]Entry, error) {
reader, err := zip.NewReader(ra, size)
if err != nil {
return nil, fmt.Errorf("open zip pages package: %w", err)
}
files := make([]archiveFile, 0, len(reader.File))
for _, item := range reader.File {
files = append(files, zipArchiveFile{file: item})
}
return entriesFromArchiveFiles(files), nil
}
func listSevenZipEntriesAt(ra io.ReaderAt, size int64) ([]Entry, error) {
reader, err := sevenzip.NewReader(ra, size)
if err != nil {
return nil, fmt.Errorf("open 7z pages package: %w", err)
}
files := make([]archiveFile, 0, len(reader.File))
for _, item := range reader.File {
files = append(files, sevenZipArchiveFile{file: item})
}
return entriesFromArchiveFiles(files), nil
}
func openTarFamilyReader(r io.Reader, format Format) (*tar.Reader, func() error, error) {
switch format {
case FormatTar:
return tar.NewReader(r), func() error { return nil }, nil
case FormatTarGz:
gzReader, err := gzip.NewReader(r)
if err != nil {
return nil, nil, fmt.Errorf("open gzip pages package: %w", err)
}
return tar.NewReader(gzReader), gzReader.Close, nil
case FormatTarXz:
xzReader, err := xz.NewReader(r)
if err != nil {
return nil, nil, fmt.Errorf("open xz pages package: %w", err)
}
return tar.NewReader(xzReader), func() error { return nil }, nil
case FormatTarBz2:
return tar.NewReader(bzip2.NewReader(r)), func() error { return nil }, nil
case FormatZip, FormatSevenZip:
return nil, nil, fmt.Errorf("unsupported tar family format: %s", format)
default:
return nil, nil, fmt.Errorf("unsupported tar family format: %s", format)
}
}
func entryFromTarHeader(header *tar.Header) Entry {
entry := Entry{Name: header.Name}
if header.Size > 0 {
entry.Size = uint64(header.Size) //nolint:gosec // archive/tar rejects negative sizes
}
switch header.Typeflag {
case tar.TypeReg, tar.TypeRegA: //nolint:staticcheck // TypeRegA appears in older archives
// Regular file.
case tar.TypeDir:
entry.IsDir = true
case tar.TypeSymlink:
entry.IsSymlink = true
case tar.TypeLink:
entry.IsHardlink = true
default:
entry.IsSpecial = true
}
return entry
}
@@ -0,0 +1,233 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"archive/tar"
"archive/zip"
"bytes"
"compress/gzip"
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/ulikunitz/xz"
)
func TestDetectFormatFromName(t *testing.T) {
cases := map[string]Format{
"site.zip": FormatZip,
"site.TAR.GZ": FormatTarGz,
"site.tgz": FormatTarGz,
"site.tar.xz": FormatTarXz,
"site.txz": FormatTarXz,
"site.tar.bz2": FormatTarBz2,
"site.tar": FormatTar,
"site.7z": FormatSevenZip,
}
for name, want := range cases {
got, ok := DetectFormatFromName(name)
assert.True(t, ok, name)
assert.Equal(t, want, got, name)
}
_, ok := DetectFormatFromName("site.rar")
assert.False(t, ok)
}
func TestInspectAndExtractZip(t *testing.T) {
data := testZip(t, map[string]string{
"dist/index.html": "<html>ok</html>",
"dist/app.js": "console.log(1)",
})
manifest, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{MaxFiles: 100, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
})
require.NoError(t, err)
assert.Equal(t, 2, manifest.FileCount)
paths := make(map[string]struct{}, len(manifest.Files))
for _, file := range manifest.Files {
paths[file.Path] = struct{}{}
assert.Empty(t, file.Checksum, "per-file checksum should not be computed")
assert.Positive(t, file.Size)
}
assert.Contains(t, paths, "index.html")
assert.Contains(t, paths, "app.js")
assert.Equal(t, int64(len("<html>ok</html>")+len("console.log(1)")), manifest.TotalSize)
dest := t.TempDir()
require.NoError(t, ExtractBytes(data, FormatZip, dest, ExtractOptions{
StripCommonRoot: true,
EnforceLimits: true,
Limits: Limits{MaxFiles: 100, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
}))
body, err := os.ReadFile(filepath.Join(dest, "index.html")) //nolint:gosec
require.NoError(t, err)
assert.Equal(t, "<html>ok</html>", string(body))
}
func TestInspectFileUsesDeclaredSizesWithoutHash(t *testing.T) {
data := testZip(t, map[string]string{
"index.html": "<html>disk</html>",
"asset.css": "body{}",
})
path := filepath.Join(t.TempDir(), "site.zip")
require.NoError(t, os.WriteFile(path, data, 0o600))
manifest, err := InspectFile(path, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{MaxFiles: 100, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
})
require.NoError(t, err)
require.Equal(t, 2, manifest.FileCount)
for _, file := range manifest.Files {
assert.Empty(t, file.Checksum)
}
dest := t.TempDir()
require.NoError(t, ExtractFile(path, FormatZip, dest, ExtractOptions{
EnforceLimits: true,
Limits: Limits{MaxFiles: 100, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
}))
body, err := os.ReadFile(filepath.Join(dest, "index.html")) //nolint:gosec
require.NoError(t, err)
assert.Equal(t, "<html>disk</html>", string(body))
}
func TestInspectVerifySizesOptional(t *testing.T) {
data := testZip(t, map[string]string{
"index.html": "verify-me",
})
manifest, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
VerifySizes: true,
Limits: Limits{MaxFiles: 10, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
})
require.NoError(t, err)
require.Len(t, manifest.Files, 1)
assert.Equal(t, int64(len("verify-me")), manifest.Files[0].Size)
assert.Empty(t, manifest.Files[0].Checksum)
}
func TestExtractTrustedSkipsSizeLimits(t *testing.T) {
// Content larger than a tiny limit would fail if limits were enforced.
large := strings.Repeat("x", 64)
data := testZip(t, map[string]string{
"index.html": large,
})
dest := t.TempDir()
require.NoError(t, ExtractBytes(data, FormatZip, dest, ExtractOptions{
// Agent trusts control-plane validation: no size/count re-check.
EnforceLimits: false,
}))
body, err := os.ReadFile(filepath.Join(dest, "index.html")) //nolint:gosec
require.NoError(t, err)
assert.Equal(t, large, string(body))
}
func TestInspectAndExtractTarGz(t *testing.T) {
data := testTarGz(t, map[string]string{
"index.html": "<html>tar</html>",
"style.css": "body{}",
})
format, err := DetectFormat("site.tar.gz", data)
require.NoError(t, err)
assert.Equal(t, FormatTarGz, format)
manifest, err := InspectBytes(data, format, InspectOptions{
EntryFile: "index.html",
Limits: Limits{MaxFiles: 100, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
})
require.NoError(t, err)
assert.Equal(t, 2, manifest.FileCount)
dest := t.TempDir()
require.NoError(t, ExtractBytes(data, format, dest, ExtractOptions{
EnforceLimits: true,
Limits: Limits{MaxFiles: 100, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
}))
body, err := os.ReadFile(filepath.Join(dest, "index.html")) //nolint:gosec
require.NoError(t, err)
assert.Equal(t, "<html>tar</html>", string(body))
}
func TestInspectTarXz(t *testing.T) {
data := testTarXz(t, map[string]string{
"index.html": "<html>xz</html>",
})
manifest, err := InspectBytes(data, FormatTarXz, InspectOptions{
EntryFile: "index.html",
Limits: Limits{MaxFiles: 10, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
})
require.NoError(t, err)
assert.Equal(t, 1, manifest.FileCount)
}
func TestRejectZipSlip(t *testing.T) {
data := testZip(t, map[string]string{
"../evil.txt": "x",
"index.html": "ok",
})
_, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{MaxFiles: 10, MaxFileBytes: 1 << 20, MaxTotalBytes: 1 << 20},
})
require.Error(t, err)
}
func testZip(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
writer := zip.NewWriter(&buffer)
for name, content := range files {
file, err := writer.Create(name)
require.NoError(t, err)
_, err = file.Write([]byte(content))
require.NoError(t, err)
}
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func testTarGz(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
gzWriter := gzip.NewWriter(&buffer)
tarWriter := tar.NewWriter(gzWriter)
for name, content := range files {
require.NoError(t, tarWriter.WriteHeader(&tar.Header{
Name: name,
Mode: 0o644,
Size: int64(len(content)),
}))
_, err := tarWriter.Write([]byte(content))
require.NoError(t, err)
}
require.NoError(t, tarWriter.Close())
require.NoError(t, gzWriter.Close())
return buffer.Bytes()
}
func testTarXz(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
xzWriter, err := xz.NewWriter(&buffer)
require.NoError(t, err)
tarWriter := tar.NewWriter(xzWriter)
for name, content := range files {
require.NoError(t, tarWriter.WriteHeader(&tar.Header{
Name: name,
Mode: 0o644,
Size: int64(len(content)),
}))
_, err := tarWriter.Write([]byte(content))
require.NoError(t, err)
}
require.NoError(t, tarWriter.Close())
require.NoError(t, xzWriter.Close())
return buffer.Bytes()
}
@@ -0,0 +1,142 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"errors"
"fmt"
"path"
"strings"
"unicode"
"unicode/utf8"
)
// NormalizeLogicalPath validates and normalizes a relative POSIX path.
// Empty input is returned unchanged only when allowEmpty is true.
func NormalizeLogicalPath(raw string, allowEmpty bool) (string, error) {
if raw == "" {
if allowEmpty {
return "", nil
}
return "", errors.New("pages path is required")
}
if err := validateLogicalPathText(raw); err != nil {
return "", err
}
cleaned := path.Clean(raw)
if cleaned == "." || cleaned == "" {
if allowEmpty {
return "", nil
}
return "", errors.New("pages path is required")
}
if strings.HasPrefix(cleaned, "/") || cleaned == ".." || strings.HasPrefix(cleaned, "../") {
return "", fmt.Errorf("pages path escapes directory: %s", raw)
}
return cleaned, nil
}
func validateLogicalPathText(raw string) error {
if !utf8.ValidString(raw) {
return errors.New("pages path is not valid UTF-8")
}
if strings.Contains(raw, "\\") {
return fmt.Errorf("pages path must use POSIX separators: %s", raw)
}
if strings.HasPrefix(raw, "/") || path.IsAbs(raw) {
return fmt.Errorf("pages path must be relative: %s", raw)
}
if err := validateLogicalPathRunes(raw); err != nil {
return err
}
return validateLogicalPathSegments(raw)
}
func validateLogicalPathRunes(raw string) error {
for _, r := range raw {
if r == 0 || unicode.IsControl(r) {
return errors.New("pages path contains a control character")
}
if r == '\'' || r == '"' || r == ';' {
return fmt.Errorf("pages path contains an unsupported character: %s", raw)
}
}
return nil
}
func validateLogicalPathSegments(raw string) error {
for segment := range strings.SplitSeq(raw, "/") {
if len(segment) >= 2 && segment[1] == ':' {
return fmt.Errorf("pages path contains a Windows drive: %s", raw)
}
if segment == "." || segment == ".." {
return fmt.Errorf("pages path escapes directory or contains a dot segment: %s", raw)
}
}
return nil
}
// NormalizeEntryPath cleans an archive entry path and rejects zip-slip / absolute paths.
// skip=true means the entry should be ignored (empty path or directory marker).
func NormalizeEntryPath(raw string) (cleaned string, skip bool, err error) {
if raw == "" {
return "", true, nil
}
cleanedPath, normalizeErr := NormalizeLogicalPath(raw, false)
if normalizeErr != nil {
return "", false, fmt.Errorf("invalid pages package path %q: %w", raw, normalizeErr)
}
if strings.HasSuffix(raw, "/") {
return cleanedPath, true, nil
}
return cleanedPath, false, nil
}
// FindCommonRootPrefix returns a trailing-slash directory prefix shared by all file paths.
// When files do not share a single root folder the result is empty.
func FindCommonRootPrefix(paths []string) string {
var firstFilePath string
hasMultipleFiles := false
for _, item := range paths {
normalizedPath, skip, err := NormalizeEntryPath(item)
if err != nil || skip {
continue
}
if firstFilePath == "" {
firstFilePath = normalizedPath
} else {
hasMultipleFiles = true
}
}
if firstFilePath == "" {
return ""
}
parts := strings.Split(firstFilePath, "/")
if len(parts) <= 1 {
return ""
}
commonPrefix := parts[0] + "/"
if !hasMultipleFiles {
return commonPrefix
}
for _, item := range paths {
normalizedPath, skip, err := NormalizeEntryPath(item)
if err != nil || skip {
continue
}
if !strings.HasPrefix(normalizedPath, commonPrefix) {
return ""
}
}
return commonPrefix
}
// StripPrefix removes a directory prefix from a cleaned path when present.
func StripPrefix(normalizedPath, prefix string) string {
if prefix == "" {
return normalizedPath
}
return strings.TrimPrefix(normalizedPath, prefix)
}
@@ -0,0 +1,499 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"archive/tar"
"archive/zip"
"bytes"
"encoding/base64"
"io"
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
var testLimits = Limits{
MaxFiles: 100,
MaxFileBytes: 1 << 20,
MaxTotalBytes: 1 << 20,
}
func TestNormalizeLogicalPathStrict(t *testing.T) {
t.Parallel()
valid := []struct {
name string
raw string
allowEmpty bool
want string
}{
{name: "empty root", allowEmpty: true},
{name: "single file", raw: "index.html", want: "index.html"},
{name: "nested posix", raw: "public/assets/app.js", want: "public/assets/app.js"},
{name: "unicode", raw: "静态/首页.html", want: "静态/首页.html"},
{name: "repeated separator is normalized", raw: "public//app.js", want: "public/app.js"},
}
for _, tt := range valid {
t.Run(tt.name, func(t *testing.T) {
got, err := NormalizeLogicalPath(tt.raw, tt.allowEmpty)
require.NoError(t, err)
assert.Equal(t, tt.want, got)
})
}
invalidUTF8 := string([]byte{'a', '/', 0xff})
invalid := []struct {
name string
raw string
}{
{name: "empty entry"},
{name: "absolute", raw: "/etc/passwd"},
{name: "unc", raw: "//server/share"},
{name: "windows drive", raw: "C:/site/index.html"},
{name: "nested windows drive", raw: "site/C:/index.html"},
{name: "windows separator", raw: `site\index.html`},
{name: "windows unc", raw: `\\server\share`},
{name: "parent segment", raw: "../index.html"},
{name: "nested parent segment", raw: "site/../index.html"},
{name: "current segment", raw: "site/./index.html"},
{name: "nul", raw: "site/\x00index.html"},
{name: "newline", raw: "site/\nindex.html"},
{name: "delete control", raw: "site/\x7findex.html"},
{name: "single quote", raw: "site/'index.html"},
{name: "double quote", raw: `site/"index.html`},
{name: "semicolon", raw: "site/;index.html"},
{name: "invalid utf8", raw: invalidUTF8},
}
for _, tt := range invalid {
t.Run(tt.name, func(t *testing.T) {
_, err := NormalizeLogicalPath(tt.raw, false)
require.Error(t, err)
})
}
cleaned, skip, err := NormalizeEntryPath("")
require.NoError(t, err)
assert.Empty(t, cleaned)
assert.True(t, skip)
cleaned, skip, err = NormalizeEntryPath("assets/")
require.NoError(t, err)
assert.Equal(t, "assets", cleaned)
assert.True(t, skip)
}
func TestSupportedFormatsInspectAndExtract(t *testing.T) {
t.Parallel()
sevenZipData := decodeFixture(t, "N3q8ryccAASgR6WICAAAAAAAAABmAAAAAAAAAN2R8/FiYXIKZm9vCgEEBgACCQQEAAcLAgABAQABAQAMBAQACAoB6bOiBKhlMn4AAAUCGQUAAAAAABERAGIAYQByAAAAZgBvAG8AAAAZAgAAFBIBAACFM3PyY9YBAFgCcvJj1gEVCgEAIICkgSCApIEAAA==")
bzipTarData := decodeFixture(t, "QlpoOTFBWSZTWYp5f6EAAHV//P64A8RQAf/iOm/9cO/v/9AAAgBADlAABAADAAgwAU1RIZJpNCaammmnqbSGTI9Q0BoBpp6mjIaGmmRoaHGRpkxNBkyYTTIGQ0BoDTJoYATQGG1KCntExT01MhoAABoAAHqAAPU9QacVDtN45fA6MmuGVQlWowrijpZgwASITYSPUcpJpoQGMkKq69jMkUR6L86R5j0IySUaZEjazEqhQ9E8vuuxsmWZQLCA84jNsobYNEzuEB1eCPhw8nc2AOz+xrCY5hVxQW1IIokpfSRKi+McvXU+QoYuEg6BD4w8x3K0imi+bULpkLCylCZ4lzoGlTQgibvG67sQcrTCRBTbBCVL7zC0q0qULmK/WOneu94s9cs4s4K98SjY2YvpdZvl42kwtxvvPMheorYQ2pcxyF4sNQYvd4+bgqm5gKXElqnGF3jhxGTeXp9eCUxWVlbi9ikxAik4xxATl7cJrISVWnHwUFiLdhEnKWw0Lhm3ZyKlX7P5Wj7b9TLAmWBaAwH/F3JFOFCQinl/oQ==")
cases := []struct {
name string
format Format
data []byte
entryFile string
wantPath string
}{
{name: "zip", format: FormatZip, data: testZip(t, map[string]string{"bundle/index.html": "zip"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar", format: FormatTar, data: testTar(t, map[string]string{"bundle/index.html": "tar"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar gzip", format: FormatTarGz, data: testTarGz(t, map[string]string{"bundle/index.html": "gzip"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar xz", format: FormatTarXz, data: testTarXz(t, map[string]string{"bundle/index.html": "xz"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar bzip2", format: FormatTarBz2, data: bzipTarData, entryFile: "index.html", wantPath: "index.html"},
{name: "7z", format: FormatSevenZip, data: sevenZipData, entryFile: "foo", wantPath: "foo"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
manifest, err := InspectBytes(tt.data, tt.format, InspectOptions{
EntryFile: tt.entryFile,
Limits: testLimits,
})
require.NoError(t, err)
assert.Positive(t, manifest.FileCount)
assertManifestContains(t, manifest, tt.wantPath)
destDir := t.TempDir()
require.NoError(t, ExtractBytes(tt.data, tt.format, destDir, ExtractOptions{
Limits: testLimits,
StripCommonRoot: true,
EnforceLimits: true,
}))
_, err = os.Stat(filepath.Join(destDir, filepath.FromSlash(tt.wantPath)))
require.NoError(t, err)
})
}
}
func TestExtractFilePreservesCommonRootAndEnforcesLimits(t *testing.T) {
t.Parallel()
data := testTarGz(t, map[string]string{
"repository/dist/index.html": "pages",
"repository/dist/app.js": "javascript",
})
archivePath := filepath.Join(t.TempDir(), "site.tar.gz")
require.NoError(t, os.WriteFile(archivePath, data, 0o600))
manifest, err := InspectFile(archivePath, FormatTarGz, InspectOptions{
RootDir: "dist",
EntryFile: "index.html",
Limits: testLimits,
})
require.NoError(t, err)
assertManifestContains(t, manifest, "dist/index.html")
destDir := t.TempDir()
require.NoError(t, ExtractFile(archivePath, FormatTarGz, destDir, ExtractOptions{
Limits: testLimits,
StripCommonRoot: true,
EnforceLimits: true,
}))
body, err := os.ReadFile(filepath.Join(destDir, "dist", "index.html")) //nolint:gosec
require.NoError(t, err)
assert.Equal(t, "pages", string(body))
err = ExtractFile(archivePath, FormatTarGz, t.TempDir(), ExtractOptions{
Limits: Limits{
MaxFiles: 10,
MaxFileBytes: 5,
MaxTotalBytes: 1 << 20,
},
StripCommonRoot: true,
EnforceLimits: true,
})
require.ErrorContains(t, err, "file too large")
err = ExtractFile(archivePath, FormatTarGz, t.TempDir(), ExtractOptions{
Limits: Limits{
MaxFiles: 10,
MaxFileBytes: 1 << 20,
MaxTotalBytes: int64(len("pages") + len("javascript") - 1),
},
StripCommonRoot: true,
EnforceLimits: true,
})
require.ErrorContains(t, err, "extracted size exceeds limit")
}
func TestRandomAccessMembersVerifyActualSize(t *testing.T) {
t.Parallel()
entry := Entry{
Name: "index.html",
Size: 1,
Open: func() (io.ReadCloser, error) {
return io.NopCloser(strings.NewReader("actual-body")), nil
},
}
_, err := inspectRandomAccessEntries([]Entry{entry}, testLimits)
require.ErrorContains(t, err, "declared size 1 does not match actual 11")
destDir := t.TempDir()
err = extractEntries([]Entry{entry}, destDir, ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
})
require.ErrorContains(t, err, "declared size 1 does not match actual 11")
_, statErr := os.Stat(filepath.Join(destDir, "index.html"))
assert.ErrorIs(t, statErr, os.ErrNotExist, "failed extraction must remove the partial file")
}
func TestActualByteLimitAbortsReaderEarly(t *testing.T) {
t.Parallel()
reader := &countingFillReader{remaining: 1 << 30}
entry := Entry{
Name: "index.html",
Size: 0,
Open: func() (io.ReadCloser, error) {
return io.NopCloser(reader), nil
},
}
_, err := inspectRandomAccessEntries([]Entry{entry}, Limits{
MaxFiles: 1,
MaxFileBytes: 32,
MaxTotalBytes: 32,
})
require.ErrorContains(t, err, "size out of bounds")
assert.LessOrEqual(t, reader.read, int64(33), "inspection must stop after limit+1 actual bytes")
tarData := tarWithDeclaredBodyOnly(t, "index.html", 1<<30)
_, err = InspectBytes(tarData, FormatTar, InspectOptions{
EntryFile: "index.html",
Limits: Limits{
MaxFiles: 1,
MaxFileBytes: 32,
MaxTotalBytes: 32,
},
})
require.ErrorContains(t, err, "file too large")
}
func TestArchiveLimitsUseFilesAndActualTotals(t *testing.T) {
t.Parallel()
data := testZip(t, map[string]string{
"index.html": "1234",
"app.js": "5678",
})
_, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{
MaxFiles: 1,
MaxFileBytes: 8,
MaxTotalBytes: 16,
},
})
require.ErrorContains(t, err, "file count exceeds")
_, err = InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{
MaxFiles: 2,
MaxFileBytes: 8,
MaxTotalBytes: 7,
},
})
require.ErrorContains(t, err, "extracted size exceeds limit")
sevenZipData := decodeFixture(t, "N3q8ryccAASgR6WICAAAAAAAAABmAAAAAAAAAN2R8/FiYXIKZm9vCgEEBgACCQQEAAcLAgABAQABAQAMBAQACAoB6bOiBKhlMn4AAAUCGQUAAAAAABERAGIAYQByAAAAZgBvAG8AAAAZAgAAFBIBAACFM3PyY9YBAFgCcvJj1gEVCgEAIICkgSCApIEAAA==")
_, err = InspectBytes(sevenZipData, FormatSevenZip, InspectOptions{
EntryFile: "foo",
Limits: Limits{
MaxFiles: 10,
MaxFileBytes: 3,
MaxTotalBytes: 32,
},
})
require.ErrorContains(t, err, "file too large")
}
func TestRejectUnsupportedTarMemberTypes(t *testing.T) {
t.Parallel()
cases := []struct {
name string
header tar.Header
wantErr string
}{
{name: "symlink", header: tar.Header{Name: "link", Typeflag: tar.TypeSymlink, Linkname: "index.html"}, wantErr: "unsupported symlink"},
{name: "hardlink", header: tar.Header{Name: "hard", Typeflag: tar.TypeLink, Linkname: "index.html"}, wantErr: "unsupported hardlink"},
{name: "fifo", header: tar.Header{Name: "pipe", Typeflag: tar.TypeFifo, Mode: 0o600}, wantErr: "unsupported special entry"},
{name: "character device", header: tar.Header{Name: "tty", Typeflag: tar.TypeChar, Mode: 0o600, Devmajor: 1, Devminor: 3}, wantErr: "unsupported special entry"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
data := tarWithSpecialEntry(t, &tt.header)
_, err := InspectBytes(data, FormatTar, InspectOptions{
EntryFile: "index.html",
Limits: testLimits,
})
require.ErrorContains(t, err, tt.wantErr)
destDir := t.TempDir()
err = ExtractBytes(data, FormatTar, destDir, ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
})
require.ErrorContains(t, err, tt.wantErr)
_, statErr := os.Stat(filepath.Join(destDir, "index.html"))
assert.ErrorIs(t, statErr, os.ErrNotExist, "tar validation pass must reject before writing files")
})
}
}
func TestRejectUnsupportedZipMemberTypes(t *testing.T) {
t.Parallel()
cases := []struct {
name string
mode os.FileMode
wantErr string
}{
{name: "symlink", mode: os.ModeSymlink | 0o777, wantErr: "unsupported symlink"},
{name: "named pipe", mode: os.ModeNamedPipe | 0o600, wantErr: "unsupported special entry"},
{name: "device", mode: os.ModeDevice | 0o600, wantErr: "unsupported special entry"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
data := zipWithSpecialEntry(t, tt.mode)
_, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: testLimits,
})
require.ErrorContains(t, err, tt.wantErr)
err = ExtractBytes(data, FormatZip, t.TempDir(), ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
})
require.ErrorContains(t, err, tt.wantErr)
})
}
}
func TestTarMetadataHeadersRemainTransparent(t *testing.T) {
t.Parallel()
for _, format := range []tar.Format{tar.FormatPAX, tar.FormatGNU} {
format := format
t.Run(format.String(), func(t *testing.T) {
data := tarWithLongMetadata(t, format)
manifest, err := InspectBytes(data, FormatTar, InspectOptions{
EntryFile: "index.html",
Limits: testLimits,
})
require.NoError(t, err)
assert.Equal(t, 2, manifest.FileCount)
assertManifestContains(t, manifest, "index.html")
destDir := t.TempDir()
require.NoError(t, ExtractBytes(data, FormatTar, destDir, ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
}))
_, err = os.Stat(filepath.Join(destDir, "index.html"))
require.NoError(t, err)
})
}
}
type countingFillReader struct {
remaining int64
read int64
}
func (r *countingFillReader) Read(p []byte) (int, error) {
if r.remaining == 0 {
return 0, io.EOF
}
if int64(len(p)) > r.remaining {
p = p[:r.remaining]
}
for i := range p {
p[i] = 'x'
}
r.remaining -= int64(len(p))
r.read += int64(len(p))
return len(p), nil
}
func decodeFixture(t *testing.T, encoded string) []byte {
t.Helper()
data, err := base64.StdEncoding.DecodeString(encoded)
require.NoError(t, err)
return data
}
func assertManifestContains(t *testing.T, manifest *Manifest, path string) {
t.Helper()
for _, file := range manifest.Files {
if file.Path == path {
return
}
}
require.Failf(t, "manifest path missing", "path %q not found in %#v", path, manifest.Files)
}
func testTar(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
for name, content := range files {
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: name,
Mode: 0o644,
Size: int64(len(content)),
}))
_, err := writer.Write([]byte(content))
require.NoError(t, err)
}
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func tarWithDeclaredBodyOnly(t *testing.T, name string, size int64) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: name,
Mode: 0o644,
Size: size,
}))
// Deliberately omit the body and trailer. The limit must reject from the
// header before archive/tar attempts to stream the declared body.
return buffer.Bytes()
}
func tarWithSpecialEntry(t *testing.T, special *tar.Header) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: "index.html",
Mode: 0o644,
Size: 2,
}))
_, err := writer.Write([]byte("ok"))
require.NoError(t, err)
require.NoError(t, writer.WriteHeader(special))
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func zipWithSpecialEntry(t *testing.T, mode os.FileMode) []byte {
t.Helper()
var buffer bytes.Buffer
writer := zip.NewWriter(&buffer)
index, err := writer.Create("index.html")
require.NoError(t, err)
_, err = index.Write([]byte("ok"))
require.NoError(t, err)
header := &zip.FileHeader{Name: "special"}
header.SetMode(mode)
special, err := writer.CreateHeader(header)
require.NoError(t, err)
if mode&os.ModeSymlink != 0 {
_, err = special.Write([]byte("index.html"))
require.NoError(t, err)
}
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func tarWithLongMetadata(t *testing.T, format tar.Format) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
longName := strings.Repeat("long-segment-", 12) + "asset.js"
header := &tar.Header{
Name: longName,
Mode: 0o644,
Size: 1,
Format: format,
}
if format == tar.FormatPAX {
header.PAXRecords = map[string]string{"comment": "metadata is not a deployable member"}
}
require.NoError(t, writer.WriteHeader(header))
_, err := writer.Write([]byte("x"))
require.NoError(t, err)
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: "index.html",
Mode: 0o644,
Size: 2,
Format: format,
}))
_, err = writer.Write([]byte("ok"))
require.NoError(t, err)
require.NoError(t, writer.Close())
return buffer.Bytes()
}
+227
View File
@@ -0,0 +1,227 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package protocol defines the communication protocol between OpenFlare server, agent, and relay components.
package protocol
import "encoding/json"
// APIResponse is a generic API response wrapper.
type APIResponse[T any] struct {
ErrorMsg string `json:"error_msg"`
Data T `json:"data"`
}
// HeartbeatData is the heartbeat request payload from agent.
type HeartbeatData struct {
AgentSettings *AgentSettings `json:"agent_settings"`
ActiveConfig *ActiveConfigMeta `json:"active_config"`
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
}
// HeartbeatResult is the heartbeat response payload.
type HeartbeatResult struct {
AgentSettings *AgentSettings
ActiveConfig *ActiveConfigMeta
WAFIPGroups []WAFIPGroup
}
// AgentSettings holds agent configuration settings.
type AgentSettings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
RestartOpenrestyNow bool `json:"restart_openresty_now"`
}
// WSMessageType constants define WebSocket message types.
const (
WSMessageTypeStatus = "status"
WSMessageTypeSettings = "settings"
WSMessageTypeActiveConfig = "active_config"
WSMessageTypeForceSyncConfig = "force_sync_config"
WSMessageTypeWAFIPGroups = "waf_ip_groups"
WSMessageTypePing = "ping"
WSMessageTypePong = "pong"
)
// WSMessage represents a WebSocket message.
type WSMessage struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
// WSOutboundMessage represents an outbound WebSocket message.
type WSOutboundMessage struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}
// WebSocketConnection defines the WebSocket connection interface.
type WebSocketConnection interface {
URL() string
SendStatus(payload NodePayload) error
SendPong() error
Receive() (WSMessage, error)
Close() error
}
// OpenrestyStatus constants define OpenResty health status values.
const (
OpenrestyStatusHealthy = "healthy"
OpenrestyStatusUnhealthy = "unhealthy"
OpenrestyStatusUnknown = "unknown"
)
// NodePayload is the agent node registration / heartbeat payload.
// schema_version 2: host_metrics + edge_health + access_logs facts only (no business pre-aggregation).
// Agents are destroy/rebuild upgraded; no wire-level compatibility aliases.
type NodePayload struct {
SchemaVersion int `json:"schema_version,omitempty"`
NodeID string `json:"node_id"`
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
ExtVersion string `json:"ext_version"`
CurrentVersion string `json:"current_version"`
LastError string `json:"last_error"`
OpenrestyStatus string `json:"openresty_status"`
OpenrestyMessage string `json:"openresty_message"`
Profile *NodeSystemProfile `json:"profile,omitempty"`
HostMetrics *NodeMetricSnapshot `json:"host_metrics,omitempty"`
EdgeHealth *NodeEdgeHealth `json:"edge_health,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
Buffered []BufferedObservabilityRecord `json:"buffered,omitempty"`
HealthEvents []NodeHealthEvent `json:"health_events"`
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
}
// NodeEdgeHealth is an instantaneous OpenResty health snapshot (L2).
type NodeEdgeHealth struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
Status string `json:"status"`
Message string `json:"message"`
Connections int64 `json:"connections"`
}
// NodeSystemProfile describes the system profile of a node.
type NodeSystemProfile struct {
Hostname string `json:"hostname"`
OSName string `json:"os_name"`
OSVersion string `json:"os_version"`
KernelVersion string `json:"kernel_version"`
Architecture string `json:"architecture"`
CPUModel string `json:"cpu_model"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAtUnix int64 `json:"reported_at_unix"`
}
// NodeMetricSnapshot is a metric snapshot of a node.
type NodeMetricSnapshot struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
}
// NodeAccessLog is an access log entry from agent (L1 business fact).
type NodeAccessLog struct {
LoggedAtUnix int64 `json:"logged_at_unix"`
RemoteAddr string `json:"remote_addr"`
Host string `json:"host"`
Path string `json:"path"`
UserAgent string `json:"user_agent,omitempty"`
CacheStatus string `json:"cache_status,omitempty"` // $upstream_cache_status
StatusCode int `json:"status_code"`
BytesSent int64 `json:"bytes_sent"` // body bytes = 已提供数据
RequestLength int64 `json:"request_length"` // 接收数据
RequestTimeMs int64 `json:"request_time_ms"` // optional
}
// BufferedObservabilityRecord is a buffered observability record (facts only).
type BufferedObservabilityRecord struct {
CapturedAtUnix int64 `json:"captured_at_unix,omitempty"`
HostMetrics *NodeMetricSnapshot `json:"host_metrics,omitempty"`
EdgeHealth *NodeEdgeHealth `json:"edge_health,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
}
// NodeHealthEvent represents a node health event.
type NodeHealthEvent struct {
EventType string `json:"event_type"`
Severity string `json:"severity"`
Message string `json:"message"`
TriggeredAtUnix int64 `json:"triggered_at_unix"`
Metadata map[string]string `json:"metadata,omitempty"`
}
// RegisterNodeResponse is the node registration response.
type RegisterNodeResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"agent_token"`
Name string `json:"name"`
}
// ActiveConfigResponse is the active configuration response.
type ActiveConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
SourceConfigJSON string `json:"source_config_json"`
SupportFiles []SupportFile `json:"support_files"`
CreatedAt string `json:"created_at"`
}
// WAFIPGroup defines a WAF IP group.
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
Checksum string `json:"checksum"`
}
// WAFIPGroupSyncRequest is a WAF IP group sync request.
type WAFIPGroupSyncRequest struct {
IDs []uint `json:"ids,omitempty"`
Checksums map[string]string `json:"checksums,omitempty"`
}
// WAFIPGroupSyncResponse is a WAF IP group sync response.
type WAFIPGroupSyncResponse struct {
Groups []WAFIPGroup `json:"groups"`
}
// SupportFile represents a support file for relay.
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
// PagesDeploymentHashResponse is the upload SHA-256 hash for a Pages deployment package.
type PagesDeploymentHashResponse struct {
DeploymentID uint `json:"deployment_id"`
Hash string `json:"hash"`
}
// PagesProjectLatestHashResponse is the hash of a project's currently active Pages deployment.
// Agents poll this like a "latest" pointer without caring about historical deployment IDs.
type PagesProjectLatestHashResponse struct {
ProjectID uint `json:"project_id"`
DeploymentID uint `json:"deployment_id"`
Hash string `json:"hash"`
PackageSize int64 `json:"package_size"`
FileCount int `json:"file_count"`
TotalSize int64 `json:"total_size"`
}
@@ -0,0 +1,110 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
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",
},
},
{
name: "PagesProjectLatestHashResponse",
value: PagesProjectLatestHashResponse{},
expected: map[string]string{
"ProjectID": "project_id",
"DeploymentID": "deployment_id",
"Hash": "hash",
"PackageSize": "package_size",
"FileCount": "file_count",
"TotalSize": "total_size",
},
},
}
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)
}
}
+134
View File
@@ -0,0 +1,134 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package protocol
// AgentNodeSystemProfile is an alias for NodeSystemProfile used by server.
type AgentNodeSystemProfile = NodeSystemProfile
// AgentNodeMetricSnapshot is an alias for NodeMetricSnapshot used by server.
type AgentNodeMetricSnapshot = NodeMetricSnapshot
// AgentNodeHealthEvent is an alias for NodeHealthEvent used by server.
type AgentNodeHealthEvent = NodeHealthEvent
// RelayProxyStat holds relay proxy statistics.
type RelayProxyStat struct {
Name string `json:"name"`
Type string `json:"type"`
Status string `json:"status"`
ClientVersion string `json:"client_version"`
LastStartTime string `json:"last_start_time"`
LastCloseTime string `json:"last_close_time"`
ClientAddr string `json:"client_addr"`
}
// RelayHeartbeatPayload is the relay heartbeat payload.
type RelayHeartbeatPayload struct {
Version string `json:"version"`
ExtVersion string `json:"frp_version"`
RelayStatus string `json:"relay_status"`
FrpsConnCount int `json:"frps_connections"`
FrpsProxyCount int `json:"frps_proxy_count"`
FrpsClientCount int `json:"frps_client_count"`
FrpsProxies []RelayProxyStat `json:"frps_proxies,omitempty"`
Name string `json:"name"`
IP string `json:"ip"`
Profile *AgentNodeSystemProfile `json:"profile,omitempty"`
Snapshot *AgentNodeMetricSnapshot `json:"snapshot,omitempty"`
HealthEvents []AgentNodeHealthEvent `json:"health_events,omitempty"`
}
// RelayConfig holds relay configuration.
type RelayConfig struct {
BindPort int `json:"bind_port"`
VhostHTTPPort int `json:"vhost_http_port"`
AuthToken string `json:"auth_token"`
LogLevel string `json:"log_level"`
WebServerEnabled bool `json:"web_server_enabled"`
WebServerPort int `json:"web_server_port"`
}
// RelaySettings holds relay runtime settings.
type RelaySettings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
}
// RelayHeartbeatResponse is the relay heartbeat response.
type RelayHeartbeatResponse struct {
RelayConfig *RelayConfig `json:"relay_config"`
RelaySettings *RelaySettings `json:"relay_settings"`
}
// ActiveConfigMeta holds active configuration metadata.
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
// FlaredConnectedRelay describes a connected relay info for flared.
type FlaredConnectedRelay struct {
RelayNodeID string `json:"relay_node_id"`
Status string `json:"status"`
ProxyCount int `json:"proxy_count"`
}
// FlaredHeartbeatPayload is the flared heartbeat payload.
type FlaredHeartbeatPayload struct {
ClientVersion string `json:"client_version"`
FrpVersion string `json:"frp_version"`
IP string `json:"ip"`
TunnelStatus string `json:"tunnel_status"`
ConnectedRelays []FlaredConnectedRelay `json:"connected_relays"`
CurrentVersion string `json:"current_version"`
CurrentChecksum string `json:"current_checksum"`
}
// FlaredHeartbeatResponse is the flared heartbeat response.
type FlaredHeartbeatResponse struct {
ActiveConfig *ActiveConfigMeta `json:"active_config"`
TunnelSettings *RelaySettings `json:"tunnel_settings"`
}
// FlaredTunnelConfigResponse is the flared tunnel configuration response.
type FlaredTunnelConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
Relays []FlaredRelayInfo `json:"relays"`
Proxies []FlaredProxyEntry `json:"proxies"`
}
// FlaredRelayInfo holds flared relay information.
type FlaredRelayInfo struct {
RelayNodeID string `json:"relay_node_id"`
Address string `json:"address"`
AuthToken string `json:"auth_token"`
ProxyURL string `json:"proxy_url"`
}
// FlaredProxyEntry represents a flared proxy entry.
type FlaredProxyEntry struct {
Name string `json:"name"`
Type string `json:"type"`
LocalAddr string `json:"local_addr"`
LocalPort int `json:"local_port"`
CustomDomains []string `json:"custom_domains"`
}
// ApplyLogPayload is the apply log payload for flared.
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"`
}
+33
View File
@@ -0,0 +1,33 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package protocol
import "strings"
// TOMLQuote renders s as a quoted TOML basic string, escaping characters that
// would otherwise break the document or allow key injection (quotes,
// backslashes, control/newline characters). Use it for every interpolated
// value written into frps/frpc TOML configs.
func TOMLQuote(s string) string {
var b strings.Builder
b.WriteByte('"')
for _, r := range s {
switch r {
case '\\':
b.WriteString(`\\`)
case '"':
b.WriteString(`\"`)
case '\n':
b.WriteString(`\n`)
case '\r':
b.WriteString(`\r`)
case '\t':
b.WriteString(`\t`)
default:
b.WriteRune(r)
}
}
b.WriteByte('"')
return b.String()
}
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package protocol
import "testing"
func TestTOMLQuote(t *testing.T) {
cases := map[string]string{
`plain`: `"plain"`,
`a"b`: `"a\"b"`,
`a\b`: `"a\\b"`,
"injection\"\n": `"injection\"\n"`,
"": `""`,
"a\tb": `"a\tb"`,
}
for in, want := range cases {
if got := TOMLQuote(in); got != want {
t.Errorf("TOMLQuote(%q) = %q, want %q", in, got, want)
}
}
}
@@ -0,0 +1,39 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package protocol
import (
"encoding/json"
"fmt"
)
// MaxWAFIPGroupSnapshotBytes is the maximum serialized size accepted for the
// complete Agent/OpenResty WAF IP group runtime document.
const MaxWAFIPGroupSnapshotBytes = 20 << 20
type wafIPGroupSnapshot struct {
Groups map[string]WAFIPGroup `json:"groups"`
}
// MarshalWAFIPGroupSnapshot serializes the exact document written by the
// Agent to waf_ip_groups.json.
func MarshalWAFIPGroupSnapshot(groups map[string]WAFIPGroup) ([]byte, error) {
if groups == nil {
groups = map[string]WAFIPGroup{}
}
return json.Marshal(wafIPGroupSnapshot{Groups: groups})
}
// ValidateWAFIPGroupSnapshotSize rejects a complete runtime document that
// cannot be published safely to the OpenResty shared-memory snapshot.
func ValidateWAFIPGroupSnapshotSize(groups map[string]WAFIPGroup) error {
data, err := MarshalWAFIPGroupSnapshot(groups)
if err != nil {
return err
}
if len(data) > MaxWAFIPGroupSnapshotBytes {
return fmt.Errorf("WAF IP 组快照大小 %d 字节超过上限 %d 字节", len(data), MaxWAFIPGroupSnapshotBytes)
}
return nil
}
@@ -0,0 +1,57 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package protocol
import (
"strings"
"testing"
)
func TestMarshalWAFIPGroupSnapshotMatchesAgentRuntimeDocument(t *testing.T) {
data, err := MarshalWAFIPGroupSnapshot(map[string]WAFIPGroup{
"7": {ID: 7, Name: "deny", Type: "manual", Enabled: true, IPList: []string{"192.0.2.7"}, Checksum: "sum"},
})
if err != nil {
t.Fatalf("MarshalWAFIPGroupSnapshot failed: %v", err)
}
want := `{"groups":{"7":{"id":7,"name":"deny","type":"manual","enabled":true,"ip_list":["192.0.2.7"],"checksum":"sum"}}}`
if string(data) != want {
t.Fatalf("snapshot = %s, want %s", data, want)
}
}
func TestValidateWAFIPGroupSnapshotSizeBoundary(t *testing.T) {
groups := map[string]WAFIPGroup{
"1": {ID: 1, Type: "manual", Enabled: true, IPList: []string{"192.0.2.1"}, Checksum: strings.Repeat("a", 64)},
}
base, err := MarshalWAFIPGroupSnapshot(groups)
if err != nil {
t.Fatalf("marshal base snapshot: %v", err)
}
groups["1"] = WAFIPGroup{
ID: 1,
Name: strings.Repeat("x", MaxWAFIPGroupSnapshotBytes-len(base)),
Type: "manual",
Enabled: true,
IPList: []string{"192.0.2.1"},
Checksum: strings.Repeat("a", 64),
}
atLimit, err := MarshalWAFIPGroupSnapshot(groups)
if err != nil {
t.Fatalf("marshal boundary snapshot: %v", err)
}
if len(atLimit) != MaxWAFIPGroupSnapshotBytes {
t.Fatalf("boundary snapshot size = %d, want %d", len(atLimit), MaxWAFIPGroupSnapshotBytes)
}
if err := ValidateWAFIPGroupSnapshotSize(groups); err != nil {
t.Fatalf("boundary snapshot rejected: %v", err)
}
group := groups["1"]
group.Name += "x"
groups["1"] = group
if err := ValidateWAFIPGroupSnapshotSize(groups); err == nil {
t.Fatal("oversized snapshot was accepted")
}
}
@@ -0,0 +1,290 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"fmt"
"strconv"
"strings"
)
const (
// OriginErrorPageSupportPath is the SupportFile path for the origin error HTML template.
OriginErrorPageSupportPath = "error_pages/origin_error.html.tmpl"
// OriginErrorPageInternalLocation is the named nginx location that serves the error body.
// OriginErrorPageInternalLocation is the named nginx location that serves the error body
// for the all-methods mode (get_only disabled). Must be a NAMED location (@...), not a
// URI internal redirect: error_page URI redirects rewrite the request method to GET, so a
// method check inside the location could never distinguish POST/PUT. Named locations keep
// the original method and (without the `=` form) the original error status.
//
// When get_only is enabled this location is NOT emitted: GET-only mode replaces the body
// via Lua header/body filters inside the proxy location, so non-GET responses pass through
// with their original status and body.
OriginErrorPageInternalLocation = "@__openflare_origin_error"
defaultOriginErrorPageStatusTag = "500-599"
)
// DefaultOriginErrorPageHTML is the built-in default (aligned with frontend minimalist).
// Placeholders {{status}} and {{host}} are substituted at request time by Lua.
const DefaultOriginErrorPageHTML = `<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>{{status}} | OpenFlare</title>
<style>
* { box-sizing: border-box; margin: 0; padding: 0; }
body {
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
background-color: #ffffff;
color: #333333;
height: 100vh;
display: flex;
flex-direction: column;
justify-content: center;
align-items: center;
text-align: center;
padding: 48px 24px;
-webkit-font-smoothing: antialiased;
}
.container {
max-width: 600px;
width: 100%;
display: flex;
flex-direction: column;
align-items: center;
gap: 24px;
}
.error-code {
font-size: 48px;
font-weight: 700;
color: #333333;
line-height: 1.2;
letter-spacing: -0.02em;
}
.error-description {
font-size: 20px;
line-height: 1.6;
color: #666666;
max-width: 480px;
}
.host {
font-size: 14px;
color: #999999;
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, monospace;
word-break: break-all;
}
.footer {
margin-top: 48px;
display: flex;
align-items: center;
justify-content: center;
gap: 8px;
color: #999999;
font-size: 14px;
font-weight: 500;
}
.brand-icon { width: 24px; height: 24px; fill: currentColor; display: block; }
@media (max-width: 480px) {
.error-code { font-size: 36px; }
.error-description { font-size: 18px; }
}
</style>
</head>
<body>
<div class="container">
<h1 class="error-code" aria-label="HTTP status">{{status}}</h1>
<p class="error-description">
The upstream server is unreachable. Please try again later or contact the site administrator if the problem persists.
</p>
<p class="host">{{host}}</p>
<div class="footer">
<svg class="brand-icon" viewBox="0 0 24 24" xmlns="http://www.w3.org/2000/svg" aria-hidden="true">
<path d="M13 2L3 14H12L11 22L21 10H12L13 2Z" />
</svg>
<span>OpenFlare</span>
</div>
</div>
</body>
</html>
`
// EffectiveOriginErrorPageHTML returns custom HTML when set, otherwise the built-in default.
func EffectiveOriginErrorPageHTML(cfg ConfigSnapshot) string {
if strings.TrimSpace(cfg.OriginErrorPageHTML) == "" {
return DefaultOriginErrorPageHTML
}
return cfg.OriginErrorPageHTML
}
func effectiveOriginErrorPageStatusTags(cfg ConfigSnapshot) []string {
if len(cfg.OriginErrorPageStatusCodes) == 0 {
return []string{defaultOriginErrorPageStatusTag}
}
return cfg.OriginErrorPageStatusCodes
}
func originErrorPageSupportFile(cfg ConfigSnapshot) SupportFile {
return SupportFile{
Path: OriginErrorPageSupportPath,
Content: EffectiveOriginErrorPageHTML(cfg),
}
}
func renderOriginErrorPageIntercept(cfg ConfigSnapshot) string {
if !cfg.OriginErrorPageEnabled {
return ""
}
codes, err := ExpandStatusCodeTags(effectiveOriginErrorPageStatusTags(cfg))
if err != nil || len(codes) == 0 {
return ""
}
if cfg.OriginErrorPageGetOnly {
// GET-only mode must NOT use proxy_intercept_errors: interception discards
// the upstream error body, so non-GET requests could never receive the
// original response (nginx would serve its own default error page instead).
// The body is replaced by Lua header/body filters that only fire for GET;
// non-GET responses pass through with status, headers and body untouched.
return renderOriginErrorPageLuaFilterBlock(codes)
}
// Intercept at the proxy level for all methods. nginx does not allow
// proxy_intercept_errors inside limit_except (only allow/deny are valid
// there), so the custom HTML is served by the named error location.
return " proxy_intercept_errors on;\n"
}
// renderOriginErrorPageLuaFilterBlock emits the GET-only body replacement inside the
// proxy location. header_filter decides whether the response should be replaced and
// reads the template once into ngx.ctx; body_filter swaps the upstream body for the
// custom HTML and forces end-of-body so remaining upstream chunks are discarded.
// Non-GET requests (or statuses outside the configured set) are never touched.
func renderOriginErrorPageLuaFilterBlock(codes []int) string {
codeList := make([]string, len(codes))
for i, code := range codes {
codeList[i] = strconv.Itoa(code)
}
return fmt.Sprintf(` header_filter_by_lua_block {
local codes = {%s}
local function match(code)
for _, c in ipairs(codes) do
if c == code then
return true
end
end
return false
end
local status = ngx.status
if match(status) and ngx.req.get_method() == "GET" then
ngx.header.content_length = nil
ngx.header["Content-Type"] = "text/html; charset=utf-8"
local f = io.open("%s", "r")
local body = f and f:read("*a")
if f then
f:close()
end
if not body then
body = "<!DOCTYPE html><html><head><meta charset=\"utf-8\"><title>" .. tostring(status) .. "</title></head><body><h1>" .. tostring(status) .. "</h1></body></html>"
end
body = body:gsub("{{status}}", function() return tostring(status) end)
body = body:gsub("{{host}}", function() return ngx.var.host or "" end)
ngx.ctx.openflare_error_html = body
end
}
body_filter_by_lua_block {
local html = ngx.ctx.openflare_error_html
if html then
ngx.arg[1] = html
ngx.arg[2] = true
ngx.ctx.openflare_error_html = nil
end
}
`, strings.Join(codeList, ", "), ErrorPageTmplPlaceholder)
}
// renderOriginErrorPageServerBits emits server-level error_page + named error location
// for the all-methods mode. Returns empty string when disabled, expand fails, no codes
// remain, or get_only is enabled (GET-only mode replaces the body via Lua filters inside
// the proxy location, see renderOriginErrorPageIntercept).
//
// IMPORTANT: do NOT use `error_page CODE = @name` (equals without response code).
// That form adopts the status returned by the error URI; content_by_lua defaults
// to 200 and ngx.status is often 0, so clients saw 200 with body "{{status}}"→"0".
// Without `=`, nginx keeps the original error status for the redirect.
func renderOriginErrorPageServerBits(cfg ConfigSnapshot) string {
if !cfg.OriginErrorPageEnabled || cfg.OriginErrorPageGetOnly {
return ""
}
codes, err := ExpandStatusCodeTags(effectiveOriginErrorPageStatusTags(cfg))
if err != nil || len(codes) == 0 {
return ""
}
parts := make([]string, len(codes))
for i, code := range codes {
parts[i] = strconv.Itoa(code)
}
var builder strings.Builder
// No `=` — preserve original error status (502 stays 502).
fmt.Fprintf(&builder, " error_page %s %s;\n", strings.Join(parts, " "), OriginErrorPageInternalLocation)
builder.WriteString(renderOriginErrorPageInternalLocation())
return builder.String()
}
func renderOriginErrorPageInternalLocation() string {
// Resolve status from $status (set by error_page redirect), then
// upstream_status, then ngx.status. Force ngx.status so the client receives
// the real error code. Use function replacers so host/status with `%` are safe.
//
// Note: fmt.Sprintf is used only for the path placeholders; Lua `%` must be
// written as `%%` so Sprintf does not treat them as format verbs.
//
// The location is NAMED (@...), not a URI internal redirect: URI redirects
// (location = /uri) rewrite the request method to GET. Named locations keep
// the original method and (without `=`) the original error status.
return fmt.Sprintf(` location %s {
default_type text/html;
charset utf-8;
content_by_lua_block {
local function resolve_error_status()
local code = tonumber(ngx.var.status)
if code and code >= 400 then
return code
end
local upstream = ngx.var.upstream_status or ""
-- multi-upstream: "502, 502" or failed connect "0"
local first = upstream:match("(%%d+)")
code = tonumber(first)
if code and code >= 400 then
return code
end
code = tonumber(ngx.status)
if code and code >= 400 then
return code
end
return 502
end
local code = resolve_error_status()
ngx.status = code
local f = io.open("%s", "r")
if not f then
ngx.header["Content-Type"] = "text/html; charset=utf-8"
ngx.say("Error ", tostring(code))
return
end
local body = f:read("*a")
f:close()
local status = tostring(code)
local host = ngx.var.host or ""
-- function replacer: plain insert, no percent pattern side effects
body = body:gsub("{{status}}", function() return status end)
body = body:gsub("{{host}}", function() return host end)
ngx.header["Content-Type"] = "text/html; charset=utf-8"
ngx.say(body)
}
}
`, OriginErrorPageInternalLocation, ErrorPageTmplPlaceholder)
}
@@ -0,0 +1,264 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"strings"
"testing"
)
func TestRenderOriginErrorPageEnabled(t *testing.T) {
t.Parallel()
doc := Document{
Routes: []Route{{
ID: 1, SiteName: "ex", Domains: []string{"ex.test"},
OriginURL: "http://127.0.0.1:9", Enabled: true,
}},
OpenRestyConfig: ConfigSnapshot{
OriginErrorPageEnabled: true,
OriginErrorPageStatusCodes: []string{"500-599"},
},
}
out, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(out, "proxy_intercept_errors on") {
t.Fatal("missing intercept")
}
if !strings.Contains(out, "error_page") || !strings.Contains(out, "@__openflare_origin_error") {
t.Fatal("missing error_page")
}
if !strings.Contains(out, "error_page 500") {
t.Fatalf("expected expanded status codes in error_page, got:\n%s", out)
}
// Must NOT use `error_page … = @name` (adopts error-URI status → often 200).
for _, line := range strings.Split(out, "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "error_page ") && strings.Contains(trimmed, " = ") {
t.Fatalf("error_page must not use '=' form, got: %s", trimmed)
}
}
if !strings.Contains(out, "error_page ") || !strings.Contains(out, " @__openflare_origin_error;") {
t.Fatal("error_page must redirect to the named error location without '='")
}
if !strings.Contains(out, "location @__openflare_origin_error {") {
t.Fatal("error location must be a named location (@...) that preserves the request method")
}
if strings.Contains(out, "location = /__openflare_origin_error") {
t.Fatal("error location must NOT be a URI internal redirect (error_page URI redirects rewrite the method to GET, breaking the get_only gate)")
}
if !strings.Contains(out, "resolve_error_status") || !strings.Contains(out, "ngx.status = code") {
t.Fatal("internal location must resolve and set ngx.status to the original error code")
}
if !strings.Contains(out, ErrorPageTmplPlaceholder) {
t.Fatal("missing error page template placeholder")
}
res, err := Render(doc, nil)
if err != nil {
t.Fatal(err)
}
found := false
for _, f := range res.SupportFiles {
if f.Path == OriginErrorPageSupportPath {
found = true
if !strings.Contains(f.Content, "{{status}}") {
t.Fatal("template missing placeholder")
}
if !strings.Contains(f.Content, "{{host}}") {
t.Fatal("template missing host placeholder")
}
}
}
if !found {
t.Fatal("missing support file")
}
}
func TestRenderOriginErrorPageGetOnly(t *testing.T) {
t.Parallel()
doc := Document{
Routes: []Route{{
ID: 1, SiteName: "ex", Domains: []string{"ex.test"},
OriginURL: "http://127.0.0.1:9", Enabled: true,
}},
OpenRestyConfig: ConfigSnapshot{
OriginErrorPageEnabled: true,
OriginErrorPageStatusCodes: []string{"500-599"},
OriginErrorPageGetOnly: true,
},
}
out, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatal(err)
}
// Regression: GET-only must NOT intercept at the proxy level.
// proxy_intercept_errors discards the upstream error body, so non-GET requests
// would receive nginx's own default error page instead of the original
// response (this was the reported bug: POST 503 returned OpenResty's page).
if strings.Contains(out, "proxy_intercept_errors") {
t.Fatal("get_only must not emit proxy_intercept_errors (it discards the upstream body for non-GET)")
}
// The body replacement must happen in Lua filters that only fire for GET.
if !strings.Contains(out, "header_filter_by_lua_block") {
t.Fatal("get_only must emit header_filter_by_lua_block inside the proxy location")
}
if !strings.Contains(out, "body_filter_by_lua_block") {
t.Fatal("get_only must emit body_filter_by_lua_block inside the proxy location")
}
if !strings.Contains(out, `ngx.req.get_method() == "GET"`) {
t.Fatal("Lua filter must replace the body only for GET requests")
}
if !strings.Contains(out, `ngx.ctx.openflare_error_html`) {
t.Fatal("Lua filter must stash the error HTML in ngx.ctx for the body filter")
}
if !strings.Contains(out, `local codes = {500`) {
t.Fatal("Lua filter must carry the expanded status codes")
}
if !strings.Contains(out, ErrorPageTmplPlaceholder) {
t.Fatal("missing error page template placeholder")
}
// No error_page / named location machinery in GET-only mode.
if strings.Contains(out, "error_page") {
t.Fatal("get_only must not emit error_page (named-location path can only serve HTML or an empty status, never the original body)")
}
if strings.Contains(out, "@__openflare_origin_error") {
t.Fatal("get_only must not emit the named error location")
}
// nginx rejects proxy_intercept_errors inside limit_except (only allow/deny
// are valid there); GET-only must rely on Lua filters instead.
if strings.Contains(out, "limit_except") {
t.Fatal("get_only must not emit limit_except (proxy_intercept_errors is not allowed there)")
}
if strings.Contains(out, "location = /__openflare_origin_error") {
t.Fatal("must not use URI internal redirect (rewrites method to GET, breaking the GET gate)")
}
}
func TestRenderOriginErrorPageDisabled(t *testing.T) {
t.Parallel()
doc := Document{
Routes: []Route{{
ID: 1, SiteName: "ex", Domains: []string{"ex.test"},
OriginURL: "http://127.0.0.1:9", Enabled: true,
}},
OpenRestyConfig: ConfigSnapshot{OriginErrorPageEnabled: false},
}
out, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatal(err)
}
if strings.Contains(out, "proxy_intercept_errors") {
t.Fatal("should not intercept when disabled")
}
if strings.Contains(out, "@__openflare_origin_error") {
t.Fatal("should not emit error location when disabled")
}
res, err := Render(doc, nil)
if err != nil {
t.Fatal(err)
}
for _, f := range res.SupportFiles {
if f.Path == OriginErrorPageSupportPath {
t.Fatal("should not emit support file when disabled")
}
}
}
func TestRenderOriginErrorPageDefaultsEmptyHTMLAndStatusCodes(t *testing.T) {
t.Parallel()
doc := Document{
Routes: []Route{{
ID: 1, SiteName: "ex", Domains: []string{"ex.test"},
OriginURL: "http://127.0.0.1:9", Enabled: true,
}},
OpenRestyConfig: ConfigSnapshot{
OriginErrorPageEnabled: true,
},
}
out, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(out, "error_page 500") {
t.Fatalf("empty status codes should default to 500-599, got:\n%s", out)
}
html := EffectiveOriginErrorPageHTML(doc.OpenRestyConfig)
if html != DefaultOriginErrorPageHTML {
t.Fatal("empty HTML should use default template")
}
if !strings.Contains(html, "{{status}}") || !strings.Contains(html, "{{host}}") {
t.Fatal("default HTML must include placeholders")
}
if !strings.Contains(html, "OpenFlare") || !strings.Contains(html, "upstream server is unreachable") {
t.Fatal("default HTML missing minimalist copy")
}
}
func TestRenderOriginErrorPageCustomHTMLInSupportFile(t *testing.T) {
t.Parallel()
custom := "<html><body>custom {{status}} @ {{host}}</body></html>"
doc := Document{
Routes: []Route{{
ID: 1, SiteName: "ex", Domains: []string{"ex.test"},
OriginURL: "http://127.0.0.1:9", Enabled: true,
}},
OpenRestyConfig: ConfigSnapshot{
OriginErrorPageEnabled: true,
OriginErrorPageStatusCodes: []string{"502"},
OriginErrorPageHTML: custom,
},
}
res, err := Render(doc, nil)
if err != nil {
t.Fatal(err)
}
found := false
for _, f := range res.SupportFiles {
if f.Path == OriginErrorPageSupportPath {
found = true
if f.Content != custom {
t.Fatalf("support file content = %q, want custom HTML", f.Content)
}
}
}
if !found {
t.Fatal("missing support file")
}
out, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(out, "error_page 502 @__openflare_origin_error;") {
t.Fatalf("expected single 502 error_page without '=', got:\n%s", out)
}
}
func TestRenderOriginErrorPageSkipsPagesRoutes(t *testing.T) {
t.Parallel()
doc := Document{
Routes: []Route{{
ID: 1, SiteName: "pages", Domains: []string{"pages.test"},
UpstreamType: "pages", Enabled: true,
PagesDeployment: &PagesDeployment{
ProjectID: 1, LocalRoot: PagesDirPlaceholder + "/projects/1/current",
EntryFile: "index.html",
},
}},
OpenRestyConfig: ConfigSnapshot{
OriginErrorPageEnabled: true,
OriginErrorPageStatusCodes: []string{"500-599"},
},
}
out, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatal(err)
}
if strings.Contains(out, "proxy_intercept_errors") {
t.Fatal("pages routes must not get proxy_intercept_errors")
}
if strings.Contains(out, "@__openflare_origin_error") {
t.Fatal("pages routes must not get origin error location")
}
}
@@ -0,0 +1,995 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package openresty renders OpenResty configuration from proxy route definitions.
package openresty
import (
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/hex"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"net/url"
"path"
"regexp"
"sort"
"strconv"
"strings"
)
const (
routeUpstreamTypePages = "pages"
indexHTML = "/index.html"
)
// RenderJSON parses the given JSON string as a Document and renders the full
// OpenResty configuration bundle, injecting the provided certificate support files.
func RenderJSON(sourceJSON string, certificateFiles []SupportFile) (*Result, error) {
var doc Document
if err := json.Unmarshal([]byte(strings.TrimSpace(sourceJSON)), &doc); err != nil {
return nil, fmt.Errorf("openresty source config json is invalid: %w", err)
}
return Render(doc, certificateFiles)
}
// Render produces a complete OpenResty configuration Result from a Document and
// a set of certificate support files.
func Render(doc Document, certificateFiles []SupportFile) (*Result, error) {
mainConfig := RenderMainConfig(doc)
routeConfig, err := RenderRouteConfig(doc, certificateFiles)
if err != nil {
return nil, err
}
wafConfig, err := RenderWAFConfig(doc.WAF)
if err != nil {
return nil, err
}
files := append([]SupportFile(nil), certificateFiles...)
files = append(files, SupportFile{Path: "waf_config.json", Content: wafConfig})
if doc.OpenRestyConfig.OriginErrorPageEnabled {
files = append(files, originErrorPageSupportFile(doc.OpenRestyConfig))
}
if doc.OpenRestyConfig.SWOfflineEnabled && len(doc.OpenRestyConfig.SWOfflineDomains) > 0 {
files = append(files, ServiceWorkerSupportFiles(doc.OpenRestyConfig)...)
}
files = DedupeSupportFiles(files)
return &Result{
MainConfig: mainConfig,
RouteConfig: routeConfig,
SupportFiles: files,
Checksum: ChecksumBundle(mainConfig, routeConfig, files),
}, nil
}
// RenderMainConfig renders the nginx main configuration string from the given
// Document, falling back to the built-in default template when none is set.
// Limit-req zones are derived from each route's effective rate after merge.
func RenderMainConfig(doc Document) string {
cfg := doc.OpenRestyConfig
templateText := cfg.MainConfigTemplate
if strings.TrimSpace(templateText) == "" {
templateText = defaultMainConfigTemplate
}
return renderMainConfigTemplate(templateText, cfg, collectEffectiveLimitReqRates(doc.Routes, cfg))
}
// ValidateMainConfigTemplate checks that the provided template text is non-empty
// and contains all required OpenResty placeholder tokens.
func ValidateMainConfigTemplate(templateText string) error {
trimmed := strings.TrimSpace(templateText)
if trimmed == "" {
return errors.New("OpenRestyMainConfigTemplate 不能为空")
}
for _, placeholder := range requiredMainConfigTemplatePlaceholders {
if !strings.Contains(trimmed, placeholder) {
return fmt.Errorf("OpenRestyMainConfigTemplate 必须保留占位符 %s", placeholder)
}
}
return nil
}
// RenderRouteConfig generates the nginx server-block configuration for all
// routes in the Document, resolving certificate files as needed.
func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, error) {
var builder strings.Builder
builder.WriteString("# This file is generated by OpenFlare. Do not edit manually.\n")
certificates := certificatesByID(certificateFiles)
for _, route := range doc.Routes {
domains := normalizedRouteDomains(route)
if len(domains) == 0 {
return "", fmt.Errorf("route %s domains are invalid", route.SiteName)
}
serverNames := renderServerNames(domains)
displayName := resolveRouteSiteName(route)
cacheConfig := routeCacheConfig{Enabled: route.CacheEnabled, Policy: route.CachePolicy, Rules: route.CacheRules}
limitConfig := mergeRouteLimitConfig(route, doc.OpenRestyConfig)
powEnabled := getPoWConfigForRoute(route.ID, doc.WAF)
if normalizeRouteUpstreamType(route.UpstreamType) == routeUpstreamTypePages {
if err := renderPagesRoute(&builder, route, displayName, serverNames, certificates, limitConfig, powEnabled, doc.OpenRestyConfig); err != nil {
return "", err
}
continue
}
if err := renderProxyRoute(&builder, route, displayName, serverNames, certificates, cacheConfig, limitConfig, powEnabled, doc.OpenRestyConfig); err != nil {
return "", err
}
}
return builder.String(), nil
}
// RenderWAFConfig serialises the WAF runtime configuration (rule groups and
// per-site bindings) as a JSON string consumed by the OpenResty Lua runtime.
func RenderWAFConfig(snapshot WAFDocument) (string, error) {
data, err := json.Marshal(snapshot)
return string(data), err
}
// ChecksumBundle returns a stable SHA-256 hex digest over the combined content
// of the main config, route config, and deduplicated support files, excluding
// the source config JSON file itself.
func ChecksumBundle(mainConfig string, routeConfig string, supportFiles []SupportFile) string {
var builder strings.Builder
builder.WriteString(mainConfig)
builder.WriteString("\n--route-config--\n")
builder.WriteString(routeConfig)
builder.WriteString("\n--support-files--\n")
files := DedupeSupportFiles(supportFiles)
sort.Slice(files, func(i int, j int) bool { return files[i].Path < files[j].Path })
for _, file := range files {
if file.Path == SourceConfigFileName {
continue
}
builder.WriteString(file.Path)
builder.WriteString("\n")
builder.WriteString(file.Content)
builder.WriteString("\n")
}
sum := sha256.Sum256([]byte(builder.String()))
return hex.EncodeToString(sum[:])
}
// DedupeSupportFiles returns a new slice with duplicate paths removed, keeping
// the last occurrence of each path.
func DedupeSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
unique := make(map[string]SupportFile, len(files))
for _, file := range files {
unique[file.Path] = file
}
result := make([]SupportFile, 0, len(unique))
for _, file := range unique {
result = append(result, file)
}
return result
}
func renderMainConfigTemplate(templateText string, cfg ConfigSnapshot, limitReqRates []string) string {
replacer := strings.NewReplacer(
"{{OpenRestyWorkerProcesses}}", cfg.WorkerProcesses,
"{{OpenRestyWorkerConnections}}", strconv.Itoa(cfg.WorkerConnections),
"{{OpenRestyWorkerRlimitNofile}}", strconv.Itoa(cfg.WorkerRlimitNofile),
"{{OpenRestyConnectionUpgradeMap}}", renderConnectionUpgradeMap(),
"{{OpenRestyDefaultServerBlock}}", renderDefaultServerBlock(cfg.DefaultServerReturnStatus, cfg.HTTP3Enabled),
"{{OpenRestyAccessLogPath}}", AccessLogPlaceholder,
"{{OpenRestyErrorLogPath}}", ErrorLogPlaceholder,
"{{OpenRestyEventsUseDirective}}", renderTemplateDirective(cfg.EventsUse != "", fmt.Sprintf("use %s;", cfg.EventsUse)),
"{{OpenRestyEventsMultiAcceptDirective}}", renderTemplateDirective(cfg.EventsMultiAcceptEnabled, "multi_accept on;"),
"{{OpenRestyKeepaliveTimeout}}", strconv.Itoa(cfg.KeepaliveTimeout),
"{{OpenRestyKeepaliveRequests}}", strconv.Itoa(cfg.KeepaliveRequests),
"{{OpenRestyClientHeaderTimeout}}", strconv.Itoa(cfg.ClientHeaderTimeout),
"{{OpenRestyClientBodyTimeout}}", strconv.Itoa(cfg.ClientBodyTimeout),
"{{OpenRestyClientMaxBodySize}}", cfg.ClientMaxBodySize,
"{{OpenRestyLargeClientHeaderBuffers}}", cfg.LargeClientHeaderBuffers,
"{{OpenRestySendTimeout}}", strconv.Itoa(cfg.SendTimeout),
"{{OpenRestyProxyConnectTimeout}}", strconv.Itoa(cfg.ProxyConnectTimeout),
"{{OpenRestyProxySendTimeout}}", strconv.Itoa(cfg.ProxySendTimeout),
"{{OpenRestyProxyReadTimeout}}", strconv.Itoa(cfg.ProxyReadTimeout),
"{{OpenRestyProxyRequestBuffering}}", onOff(cfg.ProxyRequestBuffering),
"{{OpenRestyProxyBuffering}}", onOff(cfg.ProxyBufferingEnabled),
"{{OpenRestyProxyBuffers}}", cfg.ProxyBuffers,
"{{OpenRestyProxyBufferSize}}", cfg.ProxyBufferSize,
"{{OpenRestyProxyBusyBuffersSize}}", cfg.ProxyBusyBuffersSize,
"{{OpenRestyGzip}}", onOff(cfg.GzipEnabled),
"{{OpenRestyGzipMinLength}}", strconv.Itoa(cfg.GzipMinLength),
"{{OpenRestyGzipCompLevel}}", strconv.Itoa(cfg.GzipCompLevel),
"{{OpenRestyResolverDirective}}", renderTemplateDirective(cfg.Resolvers != "", fmt.Sprintf("resolver %s;", cfg.Resolvers)),
"{{OpenRestyCacheBlock}}", renderOpenRestyCacheTemplateBlock(cfg, limitReqRates),
"{{OpenRestyRouteConfigInclude}}", RouteConfigPlaceholder,
)
return replacer.Replace(templateText)
}
func renderTemplateDirective(enabled bool, statement string) string {
if !enabled {
return ""
}
return fmt.Sprintf(" %s\n", statement)
}
func renderOpenRestyCacheTemplateBlock(cfg ConfigSnapshot, limitReqRates []string) string {
lines := []string{renderOpenRestyLimitZoneBlock(limitReqRates)}
if !cfg.CacheEnabled {
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
return strings.Join(lines, "")
}
cachePath := strings.TrimSpace(cfg.CachePath)
if cachePath == "" || strings.HasPrefix(cachePath, "/var/") {
cachePath = ProxyCachePathPlaceholder
}
lines = append(lines, strings.Join([]string{
fmt.Sprintf(" proxy_cache_path %s levels=%s keys_zone=openflare_cache:10m inactive=%s max_size=%s;", cachePath, cfg.CacheLevels, cfg.CacheInactive, cfg.CacheMaxSize),
fmt.Sprintf(" proxy_cache_key \"%s\";", cfg.CacheKeyTemplate),
fmt.Sprintf(" proxy_cache_lock %s;", onOff(cfg.CacheLockEnabled)),
fmt.Sprintf(" proxy_cache_lock_timeout %s;", cfg.CacheLockTimeout),
fmt.Sprintf(" proxy_cache_use_stale %s;", cfg.CacheUseStale),
"",
}, "\n"))
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
return strings.Join(lines, "")
}
func renderOpenRestyLimitZoneBlock(limitReqRates []string) string {
var builder strings.Builder
builder.WriteString(" limit_conn_zone $server_name zone=openflare_conn_per_server:10m;\n")
builder.WriteString(" limit_conn_zone $binary_remote_addr zone=openflare_conn_per_ip:10m;\n")
for _, rate := range limitReqRates {
fmt.Fprintf(
&builder,
" limit_req_zone $openflare_waf_site$binary_remote_addr zone=%s:10m rate=%s;\n",
limitReqZoneName(rate),
rate,
)
}
return builder.String()
}
func collectEffectiveLimitReqRates(routes []Route, cfg ConfigSnapshot) []string {
seen := make(map[string]struct{}, len(routes))
for _, route := range routes {
rate := strings.TrimSpace(mergeRouteLimitConfig(route, cfg).LimitReqPerIP)
if rate == "" {
continue
}
seen[rate] = struct{}{}
}
if len(seen) == 0 {
return nil
}
rates := make([]string, 0, len(seen))
for rate := range seen {
rates = append(rates, rate)
}
sort.Strings(rates)
return rates
}
func limitReqZoneName(rate string) string {
normalized := strings.ToLower(strings.TrimSpace(rate))
normalized = strings.ReplaceAll(normalized, "/", "")
return "openflare_req_" + normalized
}
func renderOpenRestyObservabilityTemplateBlock() string {
return fmt.Sprintf(" lua_shared_dict openflare_observability 10m;\n lua_shared_dict openflare_pow_challenges 10m;\n lua_shared_dict openflare_pow_sessions 10m;\n lua_shared_dict openflare_pow_config 1m;\n lua_shared_dict openflare_waf_config 1m;\n lua_shared_dict openflare_waf_ip_groups 64m;\n init_worker_by_lua_file %s/observability/init.lua;\n log_by_lua_file %s/observability/log.lua;\n\n server {\n listen %s;\n server_name openflare-observability;\n access_log off;\n\n location = /openflare/stub_status {\n stub_status;\n }\n\n location = /openflare/observability {\n default_type application/json;\n content_by_lua_file %s/observability/read.lua;\n }\n }\n\n", LuaDirPlaceholder, LuaDirPlaceholder, ObservabilityListenPlaceholder, LuaDirPlaceholder)
}
func renderHTTPProxyServer(serverNames string, siteName string, originURL string, originHost string, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, _ bool, cfg ConfigSnapshot) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s location / {\n%s%s%s%s%s%s }\n%s%s}\n\n", serverNames, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderOriginErrorPageIntercept(cfg), renderProxyPassBlock(originURL, upstreamConfig), renderOriginErrorPageServerBits(cfg), renderPowStaticLocationBlock(powEnabled))
}
func renderPagesAPIProxyLocationBlock(deployment *PagesDeployment) string {
if deployment == nil || !deployment.APIProxyEnabled {
return ""
}
path := strings.TrimSpace(deployment.APIProxyPath)
pass := strings.TrimSpace(deployment.APIProxyPass)
rewrite := strings.TrimSpace(deployment.APIProxyRewrite)
if path == "" || pass == "" {
return ""
}
if !strings.HasPrefix(path, "/") {
path = "/" + path
}
cleanPath := strings.TrimSuffix(path, "/")
var builder strings.Builder
// 使用 fmt.Fprintf 替代 WriteString(fmt.Sprintf(...))(QF1012)
fmt.Fprintf(&builder, "\n location %s {\n", cleanPath)
if rewrite != "" {
if !strings.HasPrefix(rewrite, "/") {
rewrite = "/" + rewrite
}
cleanRewrite := strings.TrimSuffix(rewrite, "/")
if cleanRewrite == "" {
fmt.Fprintf(&builder, " rewrite ^%s/(.*)$ /$1 break;\n", regexp.QuoteMeta(cleanPath))
fmt.Fprintf(&builder, " rewrite ^%s$ / break;\n", regexp.QuoteMeta(cleanPath))
} else {
fmt.Fprintf(&builder, " rewrite ^%s/(.*)$ %s/$1 break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite)
fmt.Fprintf(&builder, " rewrite ^%s$ %s break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite)
}
}
fmt.Fprintf(&builder, " proxy_pass %s;\n", pass)
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Host $http_host;\n")
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
builder.WriteString(" proxy_set_header Connection $connection_upgrade;\n")
builder.WriteString(" }\n")
return builder.String()
}
func renderHTTPPagesServer(serverNames string, siteName string, deployment *PagesDeployment, limitConfig routeLimitConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, _ bool, _ ConfigSnapshot) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s root %s;\n index %s;%s%s\n\n location / {\n%s%s }\n%s}\n\n", serverNames, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderPagesRootLocationBlock(deployment, limitConfig, basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderHTTPRedirectServer(serverNames string) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", serverNames)
}
func renderHTTPSServer(serverNames string, siteName string, originURL string, originHost string, certificateID uint, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, swEnabled bool, cfg ConfigSnapshot) string {
certPath := fmt.Sprintf("%s/%d.crt", CertDirPlaceholder, certificateID)
keyPath := fmt.Sprintf("%s/%d.key", CertDirPlaceholder, certificateID)
var h3Listen string
var h3Header string
if cfg.HTTP3Enabled {
h3Listen = " listen 443 quic;\n"
h3Header = " add_header Alt-Svc 'h3=\":443\"; ma=86400';\n"
}
if swEnabled {
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s location / {\n%s%s%s%s%s%s }\n%s%s%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlockWithSW(siteName, powEnabled, cfg), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderOriginErrorPageIntercept(cfg), renderProxyPassBlock(originURL, upstreamConfig), renderOriginErrorPageServerBits(cfg), renderPowStaticLocationBlock(powEnabled), renderServiceWorkerChallenger(cfg))
}
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s location / {\n%s%s%s%s%s%s }\n%s%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderOriginErrorPageIntercept(cfg), renderProxyPassBlock(originURL, upstreamConfig), renderOriginErrorPageServerBits(cfg), renderPowStaticLocationBlock(powEnabled))
}
func renderHTTPSPagesServer(serverNames string, siteName string, certificateID uint, deployment *PagesDeployment, limitConfig routeLimitConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, swEnabled bool, cfg ConfigSnapshot) string {
certPath := fmt.Sprintf("%s/%d.crt", CertDirPlaceholder, certificateID)
keyPath := fmt.Sprintf("%s/%d.key", CertDirPlaceholder, certificateID)
var h3Listen string
var h3Header string
if cfg.HTTP3Enabled {
h3Listen = " listen 443 quic;\n"
h3Header = " add_header Alt-Svc 'h3=\":443\"; ma=86400';\n"
}
if swEnabled {
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s root %s;\n index %s;%s%s\n\n location / {\n%s%s }\n%s%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlockWithSW(siteName, powEnabled, cfg), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderPagesRootLocationBlock(deployment, limitConfig, basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled), renderServiceWorkerChallenger(cfg))
}
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s root %s;\n index %s;%s%s\n\n location / {\n%s%s }\n%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderPagesRootLocationBlock(deployment, limitConfig, basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderPagesRootLocationBlock(deployment *PagesDeployment, limitConfig routeLimitConfig, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string) string {
tryFile := pagesRootTryFile(deployment)
var builder strings.Builder
builder.WriteString("\n location = / {\n")
builder.WriteString(renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword))
builder.WriteString(renderRouteLimitBlock(limitConfig))
fmt.Fprintf(&builder, " try_files %s =404;\n", tryFile)
builder.WriteString(" }\n")
return builder.String()
}
func pagesRootTryFile(deployment *PagesDeployment) string {
if deployment != nil && deployment.SPAFallbackEnabled {
return pagesFallbackPath(deployment)
}
return "/" + pagesEntryFile(deployment)
}
func renderPagesLocationBlock(deployment *PagesDeployment, limitConfig routeLimitConfig) string {
var builder strings.Builder
builder.WriteString(renderRouteLimitBlock(limitConfig))
if deployment != nil && deployment.SPAFallbackEnabled {
fmt.Fprintf(&builder, " try_files $uri $uri/ %s;\n", pagesFallbackPath(deployment))
} else {
builder.WriteString(" try_files $uri $uri/ =404;\n")
}
return builder.String()
}
func pagesDeploymentRoot(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.LocalRoot) == "" {
return PagesDirPlaceholder
}
return filepathToNginxPath(deployment.LocalRoot)
}
func pagesEntryFile(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.EntryFile) == "" {
return "index.html"
}
return strings.TrimPrefix(filepathToNginxPath(deployment.EntryFile), "/")
}
func pagesFallbackPath(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.SPAFallbackPath) == "" {
return indexHTML
}
value := filepathToNginxPath(strings.TrimSpace(deployment.SPAFallbackPath))
if !strings.HasPrefix(value, "/") {
value = "/" + value
}
if value == "/" || strings.HasSuffix(value, "/") || strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") || strings.ContainsAny(value, " \t\r\n") {
return indexHTML
}
for segment := range strings.SplitSeq(value, "/") {
if segment == "." || segment == ".." {
return indexHTML
}
}
cleaned := path.Clean(value)
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
return "/index.html"
}
return cleaned
}
func renderProxyHeaderBlock(originURL string, originHost string, customHeaders []CustomHeader, upstreamConfig routeUpstreamConfig, cfg ConfigSnapshot) string {
var builder strings.Builder
if strings.TrimSpace(originHost) != "" {
fmt.Fprintf(&builder, " proxy_set_header Host %s;\n", quoteNginxStringLiteral(originHost))
} else {
builder.WriteString(" proxy_set_header Host $host;\n")
}
if upstreamServerName := resolveUpstreamServerName(originURL, originHost); upstreamServerName != "" {
builder.WriteString(" proxy_ssl_server_name on;\n")
fmt.Fprintf(&builder, " proxy_ssl_name %s;\n", quoteNginxStringLiteral(upstreamServerName))
}
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
if cfg.WebsocketEnabled {
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Connection $connection_upgrade;\n")
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
} else if upstreamConfig.UsesNamedUpstream {
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Connection \"\";\n")
}
for _, header := range customHeaders {
fmt.Fprintf(&builder, " proxy_set_header %s %s;\n", header.Key, quoteNginxStringLiteral(header.Value))
}
return builder.String()
}
func renderAccessBlock(siteName string, powEnabled bool) string {
escapedSiteName := escapeNginxString(siteName)
if !powEnabled {
return fmt.Sprintf(" set $openflare_waf_site \"%s\";\n access_by_lua_file %s/waf/check.lua;\n", escapedSiteName, LuaDirPlaceholder)
}
return fmt.Sprintf(` set $openflare_waf_site "%s";
access_by_lua_block {
if not string.find(package.path, "%s/?.lua", 1, true) then
package.path = "%s/?.lua;%s/?/init.lua;" .. package.path
end
require("waf.runtime").check()
if ngx.ctx.openflare_waf_blocked then
return
end
require("pow.runtime").check()
}
`, escapedSiteName, LuaDirPlaceholder, LuaDirPlaceholder, LuaDirPlaceholder)
}
// renderAccessBlockWithSW emits the access phase directives for a server block,
// merging the Service Worker runtime check into the single access directive.
// nginx runs only the last access_by_lua* directive in a scope, so the SW check
// must never be emitted as a second directive; otherwise it would silently
// override (or be overridden by) the WAF/PoW check.
func renderAccessBlockWithSW(siteName string, powEnabled bool, _ ConfigSnapshot) string {
escapedSiteName := escapeNginxString(siteName)
if !powEnabled {
return fmt.Sprintf(` set $openflare_waf_site "%s";
access_by_lua_block {
if not string.find(package.path, "%s/?.lua", 1, true) then
package.path = "%s/?.lua;%s/?/init.lua;" .. package.path
end
require("waf.runtime").check()
if ngx.ctx.openflare_waf_blocked then
return
end
require("sw.runtime").check()
}
`, escapedSiteName, LuaDirPlaceholder, LuaDirPlaceholder, LuaDirPlaceholder)
}
return fmt.Sprintf(` set $openflare_waf_site "%s";
access_by_lua_block {
if not string.find(package.path, "%s/?.lua", 1, true) then
package.path = "%s/?.lua;%s/?/init.lua;" .. package.path
end
require("waf.runtime").check()
if ngx.ctx.openflare_waf_blocked then
return
end
require("pow.runtime").check()
require("sw.runtime").check()
}
`, escapedSiteName, LuaDirPlaceholder, LuaDirPlaceholder, LuaDirPlaceholder)
}
func renderBasicAuthBlock(enabled bool, username, password string) string {
if !enabled || username == "" || password == "" {
return ""
}
encoded := base64.StdEncoding.EncodeToString([]byte(username + ":" + password))
return fmt.Sprintf(" rewrite_by_lua_block {\n local auth = ngx.var.http_authorization\n if auth ~= \"Basic %s\" then\n ngx.header[\"WWW-Authenticate\"] = 'Basic realm=\"Restricted\"'\n return ngx.exit(401)\n end\n }\n", encoded)
}
func renderPowLocationBlocks(powEnabled bool) string {
if !powEnabled {
return ""
}
return fmt.Sprintf("\n location = %spass-challenge {\n content_by_lua_file %s/pow/verify.lua;\n }\n\n location = %smake-challenge {\n content_by_lua_file %s/pow/challenge.lua;\n }\n\n", anubisAPIPrefix, LuaDirPlaceholder, anubisAPIPrefix, LuaDirPlaceholder)
}
func renderPowStaticLocationBlock(powEnabled bool) string {
if !powEnabled {
return ""
}
return fmt.Sprintf(" location %s {\n alias %s/;\n types {\n text/css css;\n application/javascript js mjs;\n application/json json;\n image/webp webp;\n font/woff2 woff2;\n }\n }\n\n", anubisStaticPrefix, PowStaticDirPlaceholder)
}
func renderRouteCacheBlock(cacheConfig routeCacheConfig, cfg ConfigSnapshot) string {
if !cfg.CacheEnabled || !cacheConfig.Enabled {
return ""
}
var builder strings.Builder
builder.WriteString(" set $openflare_skip_cache 0;\n")
builder.WriteString(" if ($request_method != GET) {\n set $openflare_skip_cache 1;\n }\n")
if condition := renderRouteCachePolicyCondition(cacheConfig); condition != "" {
builder.WriteString(condition)
}
builder.WriteString(" proxy_cache openflare_cache;\n")
builder.WriteString(" proxy_cache_methods GET;\n")
builder.WriteString(" proxy_cache_bypass $openflare_skip_cache;\n")
builder.WriteString(" proxy_no_cache $openflare_skip_cache $upstream_http_set_cookie;\n")
builder.WriteString(" proxy_cache_valid 200 206 301 120m;\n")
builder.WriteString(" proxy_cache_valid 302 303 20m;\n")
builder.WriteString(" proxy_cache_valid 404 410 3m;\n")
return builder.String()
}
func renderRouteLimitBlock(limitConfig routeLimitConfig) string {
var builder strings.Builder
if limitConfig.LimitConnPerServer > 0 {
fmt.Fprintf(&builder, " limit_conn openflare_conn_per_server %d;\n", limitConfig.LimitConnPerServer)
}
if limitConfig.LimitConnPerIP > 0 {
fmt.Fprintf(&builder, " limit_conn openflare_conn_per_ip %d;\n", limitConfig.LimitConnPerIP)
}
if strings.TrimSpace(limitConfig.LimitRate) != "" {
fmt.Fprintf(&builder, " limit_rate %s;\n", limitConfig.LimitRate)
}
if strings.TrimSpace(limitConfig.LimitReqPerIP) != "" {
rate := strings.TrimSpace(limitConfig.LimitReqPerIP)
burst := calculateBurst(rate)
fmt.Fprintf(&builder, " limit_req zone=%s burst=%d nodelay;\n", limitReqZoneName(rate), burst)
fmt.Fprintf(&builder, " limit_req_status 429;\n")
}
return builder.String()
}
func mergeRouteLimitConfig(route Route, cfg ConfigSnapshot) routeLimitConfig {
return routeLimitConfig{
LimitConnPerServer: mergeLimitConn(route.LimitConnPerServer, cfg.DefaultLimitConnPerServer),
LimitConnPerIP: mergeLimitConn(route.LimitConnPerIP, cfg.DefaultLimitConnPerIP),
LimitRate: mergeLimitRate(route.LimitRate, cfg.DefaultLimitRate),
LimitReqPerIP: mergeLimitRate(route.LimitReqPerIP, cfg.DefaultLimitReqPerIP),
}
}
func mergeLimitConn(route, def int) int {
if route == -1 {
return 0
}
if route > 0 {
return route
}
if def > 0 {
return def
}
return 0
}
func mergeLimitRate(route, def string) string {
r := strings.ToLower(strings.TrimSpace(route))
if r == "-1" {
return ""
}
if r != "" && r != "0" {
return r
}
d := strings.ToLower(strings.TrimSpace(def))
if d != "" && d != "0" {
return d
}
return ""
}
func renderRouteCachePolicyCondition(cacheConfig routeCacheConfig) string {
policy := normalizeRenderCachePolicy(cacheConfig.Policy)
switch policy {
case cachePolicyStatic:
return fmt.Sprintf(" if ($uri !~* %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildSuffixMatchPattern(DefaultStaticCacheExtensions)))
case cachePolicySuffix:
return fmt.Sprintf(" if ($uri !~* %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildSuffixMatchPattern(cacheConfig.Rules)))
case cachePolicyPathPrefix:
return fmt.Sprintf(" if ($uri !~ %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildPathPrefixMatchPattern(cacheConfig.Rules)))
case cachePolicyPathExact:
return fmt.Sprintf(" if ($uri !~ %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildPathExactMatchPattern(cacheConfig.Rules)))
case cachePolicyAll, cachePolicyURL:
return ""
default:
// Unknown policy: treat as static for safety (do not cache everything).
return fmt.Sprintf(" if ($uri !~* %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildSuffixMatchPattern(DefaultStaticCacheExtensions)))
}
}
// normalizeRenderCachePolicy maps stored policy for OpenResty generation.
// Legacy empty and "url" mean "all GETs after security bypass" (pre-static default).
// Explicit "static" uses the built-in extension allowlist. Unknown policies fall back to static.
func normalizeRenderCachePolicy(raw string) string {
policy := strings.TrimSpace(strings.ToLower(raw))
switch policy {
case "", cachePolicyURL, cachePolicyAll:
return cachePolicyAll
case cachePolicyStatic:
return cachePolicyStatic
case cachePolicySuffix, cachePolicyPathPrefix, cachePolicyPathExact:
return policy
default:
return policy
}
}
func renderProxyPassBlock(originURL string, upstreamConfig routeUpstreamConfig) string {
parsed, err := url.Parse(originURL)
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
return fmt.Sprintf(" proxy_pass %s;\n", originURL)
}
if upstreamConfig.UsesNamedUpstream {
return fmt.Sprintf(" proxy_pass %s://%s%s;\n", upstreamConfig.Scheme, upstreamConfig.Name, upstreamConfig.ProxyPassURI)
}
return fmt.Sprintf(" proxy_pass %s;\n", originURL)
}
func buildRouteUpstreamConfig(route Route, upstreams []string) routeUpstreamConfig {
if len(upstreams) == 0 {
return routeUpstreamConfig{}
}
if len(upstreams) == 1 {
parsed, err := url.Parse(strings.TrimSpace(upstreams[0]))
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
return routeUpstreamConfig{}
}
return routeUpstreamConfig{Name: buildRouteUpstreamName(route), Scheme: parsed.Scheme, ProxyPassURI: buildUpstreamProxyPassURI(parsed), Servers: []string{parsed.Host}, UsesNamedUpstream: true}
}
servers := make([]string, 0, len(upstreams))
var scheme string
for _, upstream := range upstreams {
parsed, err := url.Parse(strings.TrimSpace(upstream))
if err != nil || parsed.Host == "" || parsed.Scheme == "" || (strings.TrimSpace(parsed.EscapedPath()) != "" && strings.TrimSpace(parsed.EscapedPath()) != "/") || parsed.RawQuery != "" {
return routeUpstreamConfig{}
}
if scheme == "" {
scheme = parsed.Scheme
} else if scheme != parsed.Scheme {
return routeUpstreamConfig{}
}
servers = append(servers, parsed.Host)
}
return routeUpstreamConfig{Name: buildRouteUpstreamName(route), Scheme: scheme, Servers: servers, UsesNamedUpstream: true}
}
func normalizeRouteUpstreamType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case routeUpstreamTypePages:
return routeUpstreamTypePages
default:
return "direct"
}
}
func renderNamedUpstreamBlock(upstreamConfig routeUpstreamConfig) string {
var builder strings.Builder
fmt.Fprintf(&builder, "upstream %s {\n", upstreamConfig.Name)
for _, server := range upstreamConfig.Servers {
fmt.Fprintf(&builder, " server %s max_fails=3 fail_timeout=10s;\n", server)
}
builder.WriteString(" keepalive 128;\n}\n\n")
return builder.String()
}
func resolveRouteSiteName(route Route) string {
if name := strings.TrimSpace(route.SiteName); name != "" {
return name
}
if domains := normalizedRouteDomains(route); len(domains) > 0 {
return domains[0]
}
return ""
}
func buildRouteUpstreamName(route Route) string {
identity := resolveRouteSiteName(route)
sanitized := strings.Map(func(r rune) rune {
switch {
case r >= 'a' && r <= 'z':
return r
case r >= 'A' && r <= 'Z':
return r + ('a' - 'A')
case r >= '0' && r <= '9':
return r
default:
return '_'
}
}, identity)
sanitized = strings.Trim(sanitized, "_")
if sanitized == "" {
sanitized = "backend"
}
return fmt.Sprintf("backend_%s_%d", sanitized, route.ID)
}
func buildUpstreamProxyPassURI(parsed *url.URL) string {
path := parsed.EscapedPath()
if path == "/" {
path = ""
}
if parsed.RawQuery == "" {
return path
}
return fmt.Sprintf("%s?%s", path, parsed.RawQuery)
}
func renderConnectionUpgradeMap() string {
return " map $http_upgrade $connection_upgrade {\n default upgrade;\n '' \"\";\n }\n\n"
}
func renderDefaultServerBlock(statusCode int, http3Enabled bool) string {
if statusCode <= 0 {
statusCode = 421
}
var h3Default string
if http3Enabled {
h3Default = "\n listen 443 quic reuseport default_server;"
}
return strings.Join([]string{
" server {",
" listen 80 default_server;",
" server_name _;",
"",
fmt.Sprintf(" return %d;", statusCode),
" }",
"",
" server {",
" listen 443 ssl default_server;" + h3Default,
" server_name _;",
"",
" ssl_reject_handshake on;",
" }",
"",
}, "\n")
}
func normalizedRouteDomains(route Route) []string {
return route.Domains
}
func certificateIDsFromDomainCertIDs(domainCertIDs []uint) []uint {
seen := make(map[uint]struct{}, len(domainCertIDs))
normalized := make([]uint, 0, len(domainCertIDs))
for _, id := range domainCertIDs {
if id == 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
return normalized
}
func certificatesByID(files []SupportFile) map[uint]string {
result := make(map[uint]string)
for _, file := range files {
if !strings.HasSuffix(file.Path, ".crt") {
continue
}
idText := strings.TrimSuffix(file.Path, ".crt")
var id uint
if _, err := fmt.Sscanf(idText, "%d", &id); err == nil && id != 0 {
result[id] = file.Content
}
}
return result
}
func validateCertificateCoverage(certPEM string, domains []string) error {
block, _ := pem.Decode([]byte(certPEM))
if block == nil {
return errors.New("certificate PEM is invalid")
}
leaf, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return err
}
for _, domain := range domains {
if err := leaf.VerifyHostname(domain); err != nil {
return fmt.Errorf("certificate does not cover domain %s", domain)
}
}
return nil
}
func getPoWConfigForRoute(routeID uint, snapshot WAFDocument) bool {
enabledGroups := make(map[uint]WAFRuleGroup, len(snapshot.RuleGroups))
globalGroupIDs := make([]uint, 0)
for _, group := range snapshot.RuleGroups {
if !group.Enabled {
continue
}
enabledGroups[group.ID] = group
if group.IsGlobal {
globalGroupIDs = append(globalGroupIDs, group.ID)
}
}
var boundGroupIDs []uint
for _, binding := range snapshot.Bindings {
if binding.RouteID != routeID {
continue
}
for _, groupID := range binding.RuleGroupIDs {
if _, ok := enabledGroups[groupID]; ok {
boundGroupIDs = append(boundGroupIDs, groupID)
}
}
break
}
activeGroupIDs := uniqueUintIDs(append(append([]uint{}, globalGroupIDs...), boundGroupIDs...))
for _, groupID := range activeGroupIDs {
group := enabledGroups[groupID]
if graphContainsNodeType(group.Graph, "pow") {
return true
}
}
return false
}
func graphContainsNodeType(graph WAFRuleGraph, nodeType string) bool {
for _, node := range graph.Nodes {
if node.Type == nodeType {
return true
}
}
return false
}
func uniqueUintIDs(values []uint) []uint {
seen := make(map[uint]struct{}, len(values))
result := make([]uint, 0, len(values))
for _, value := range values {
if value == 0 {
continue
}
if _, ok := seen[value]; ok {
continue
}
seen[value] = struct{}{}
result = append(result, value)
}
return result
}
func resolveUpstreamServerName(originURL string, originHost string) string {
parsed, err := url.Parse(originURL)
if err != nil || !strings.EqualFold(parsed.Scheme, "https") {
return ""
}
if strings.TrimSpace(originHost) != "" {
parsedHost, err := url.Parse("//" + originHost)
if err == nil && parsedHost.Hostname() != "" {
return parsedHost.Hostname()
}
return originHost
}
return parsed.Hostname()
}
func renderServerNames(domains []string) string { return strings.Join(domains, " ") }
func onOff(value bool) string {
if value {
return "on"
}
return "off"
}
func quoteNginxStringLiteral(value string) string {
escaped := strings.ReplaceAll(value, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
return fmt.Sprintf(`"%s"`, escaped)
}
func filepathToNginxPath(value string) string {
return strings.ReplaceAll(strings.TrimSpace(value), `\`, `/`)
}
func escapeNginxString(value string) string {
escaped := strings.ReplaceAll(value, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
return escaped
}
func buildSuffixMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
parts = append(parts, regexp.QuoteMeta(rule))
}
return fmt.Sprintf("\\.(?:%s)$", strings.Join(parts, "|"))
}
func buildPathPrefixMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
trimmed := strings.TrimRight(rule, "/")
if trimmed == "" {
trimmed = "/"
}
if trimmed == "/" {
parts = append(parts, "/")
continue
}
parts = append(parts, regexp.QuoteMeta(trimmed)+"(?:/|$)")
}
return fmt.Sprintf("^(?:%s)", strings.Join(parts, "|"))
}
func buildPathExactMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
parts = append(parts, regexp.QuoteMeta(rule))
}
return fmt.Sprintf("^(?:%s)$", strings.Join(parts, "|"))
}
const (
limitReqDefaultBurst = 5
limitReqPerSecondBurstMul = 2
limitReqPerMinuteBurstDiv = 5
)
func calculateBurst(rateStr string) int {
rateStr = strings.ToLower(strings.TrimSpace(rateStr))
if rateStr == "" {
return 0
}
var val int
var unit string
_, err := fmt.Sscanf(rateStr, "%dr/%s", &val, &unit)
if err != nil || val <= 0 {
return limitReqDefaultBurst
}
switch unit {
case "s":
return val * limitReqPerSecondBurstMul
case "m":
b := val / limitReqPerMinuteBurstDiv
if b < limitReqDefaultBurst {
return limitReqDefaultBurst
}
return b
default:
return limitReqDefaultBurst
}
}
@@ -0,0 +1,154 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"fmt"
"strings"
)
type routeCertPartition struct {
httpOnlyDomains []string
domainsByCertID map[uint][]string
}
func partitionRouteDomainsByCert(domains []string, certIDs, domainCertIDs []uint) routeCertPartition {
httpOnlyDomains := make([]string, 0, len(domains))
domainsByCertID := make(map[uint][]string, len(certIDs))
for index, domain := range domains {
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
httpOnlyDomains = append(httpOnlyDomains, domain)
continue
}
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
}
return routeCertPartition{
httpOnlyDomains: httpOnlyDomains,
domainsByCertID: domainsByCertID,
}
}
func validateRouteCertificates(route Route, displayName string, certIDs []uint, partition routeCertPartition, certificates map[uint]string) error {
for _, certID := range certIDs {
assignedDomains := partition.domainsByCertID[certID]
if len(assignedDomains) == 0 {
continue
}
certPEM, ok := certificates[certID]
if !ok {
return fmt.Errorf("route %s certificate %d does not exist", route.SiteName, certID)
}
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
return fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
}
}
return nil
}
func renderPagesRouteHTTPS(
builder *strings.Builder,
serverNames, displayName string,
route Route,
partition routeCertPartition,
certIDs []uint,
limitConfig routeLimitConfig,
powEnabled bool,
cfg ConfigSnapshot,
) {
if route.RedirectHTTP {
if len(partition.httpOnlyDomains) > 0 {
builder.WriteString(renderHTTPPagesServer(renderServerNames(partition.httpOnlyDomains), displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, false, cfg))
}
for _, certID := range certIDs {
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
}
}
} else {
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, false, cfg))
}
for _, certID := range certIDs {
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPSPagesServer(renderServerNames(assignedDomains), displayName, certID, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, routeSWEnabled(assignedDomains, cfg), cfg))
}
}
}
func renderProxyRouteHTTPS(
builder *strings.Builder,
serverNames, displayName string,
route Route,
partition routeCertPartition,
certIDs []uint,
cacheConfig routeCacheConfig,
limitConfig routeLimitConfig,
upstreamConfig routeUpstreamConfig,
powEnabled bool,
cfg ConfigSnapshot,
) {
if route.RedirectHTTP {
if len(partition.httpOnlyDomains) > 0 {
builder.WriteString(renderHTTPProxyServer(renderServerNames(partition.httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, false, cfg))
}
for _, certID := range certIDs {
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
}
}
} else {
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, false, cfg))
}
for _, certID := range certIDs {
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, routeSWEnabled(assignedDomains, cfg), cfg))
}
}
}
func renderPagesRoute(builder *strings.Builder, route Route, displayName, serverNames string, certificates map[uint]string, limitConfig routeLimitConfig, powEnabled bool, cfg ConfigSnapshot) error {
if route.PagesDeployment == nil {
return fmt.Errorf("route %s pages deployment is missing", route.SiteName)
}
if !route.EnableHTTPS {
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, false, cfg))
return nil
}
certIDs := certificateIDsFromDomainCertIDs(route.DomainCertIDs)
domainCertIDs := route.DomainCertIDs
if len(certIDs) == 0 {
return fmt.Errorf("路由 %s 未配置证书", route.SiteName)
}
partition := partitionRouteDomainsByCert(normalizedRouteDomains(route), certIDs, domainCertIDs)
if err := validateRouteCertificates(route, displayName, certIDs, partition, certificates); err != nil {
return err
}
renderPagesRouteHTTPS(builder, serverNames, displayName, route, partition, certIDs, limitConfig, powEnabled, cfg)
return nil
}
func renderProxyRoute(builder *strings.Builder, route Route, displayName, serverNames string, certificates map[uint]string, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, powEnabled bool, cfg ConfigSnapshot) error {
upstreams := route.Upstreams
if len(upstreams) == 0 && strings.TrimSpace(route.OriginURL) != "" {
upstreams = []string{route.OriginURL}
}
upstreamConfig := buildRouteUpstreamConfig(route, upstreams)
if upstreamConfig.UsesNamedUpstream {
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
}
if !route.EnableHTTPS {
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, false, cfg))
return nil
}
certIDs := certificateIDsFromDomainCertIDs(route.DomainCertIDs)
domainCertIDs := route.DomainCertIDs
if len(certIDs) == 0 {
return fmt.Errorf("路由 %s 未配置证书", route.SiteName)
}
partition := partitionRouteDomainsByCert(normalizedRouteDomains(route), certIDs, domainCertIDs)
if err := validateRouteCertificates(route, displayName, certIDs, partition, certificates); err != nil {
return err
}
renderProxyRouteHTTPS(builder, serverNames, displayName, route, partition, certIDs, cacheConfig, limitConfig, upstreamConfig, powEnabled, cfg)
return nil
}
@@ -0,0 +1,667 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"encoding/json"
"strings"
"testing"
)
func TestRenderOpenRestyUsesDedicatedWAFIPGroupSharedDict(t *testing.T) {
block := renderOpenRestyObservabilityTemplateBlock()
if !strings.Contains(block, "lua_shared_dict openflare_waf_config 1m;") {
t.Fatal("expected general WAF coordination dictionary to remain available")
}
if !strings.Contains(block, "lua_shared_dict openflare_waf_ip_groups 64m;") {
t.Fatalf("expected dedicated 64m WAF IP group dictionary, got:\n%s", block)
}
}
func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
doc := Document{
Routes: []Route{
{ID: 1, SiteName: "example.com", Domains: []string{"example.com", "www.example.com"}},
{ID: 2, SiteName: "named-site", Domains: []string{"other.example.com"}},
},
WAF: WAFDocument{
RuleGroups: []WAFRuleGroup{
{
ID: 1, Name: "pow-group", Enabled: true,
Graph: WAFRuleGraph{Entry: "pow", Nodes: map[string]WAFRuleNode{"pow": {Type: "pow"}}},
},
},
Bindings: []WAFBinding{
{RouteID: 1, SiteName: "example.com", RuleGroupIDs: []uint{1}},
{RouteID: 2, SiteName: "named-site", RuleGroupIDs: []uint{1}},
},
},
}
wafConfig, err := RenderWAFConfig(doc.WAF)
if err != nil {
t.Fatalf("RenderWAFConfig() error = %v", err)
}
var decoded WAFDocument
if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if len(decoded.Bindings) != 2 || decoded.Bindings[0].SiteName != "example.com" || decoded.Bindings[1].SiteName != "named-site" {
t.Fatalf("bindings did not preserve route site names: %#v", decoded.Bindings)
}
routeConfig, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
if !strings.Contains(routeConfig, `set $openflare_waf_site "example.com"`) {
t.Fatalf("expected route config to use normalized site name example.com, got:\n%s", routeConfig)
}
if !strings.Contains(routeConfig, `require("pow.runtime").check()`) {
t.Fatalf("expected route config to enable pow runtime, got:\n%s", routeConfig)
}
}
func TestRenderWAFConfigDoesNotSynthesizeLegacyPoWConfig(t *testing.T) {
doc := WAFDocument{
RuleGroups: []WAFRuleGroup{
{
ID: 1,
Name: "global",
Enabled: true,
IsGlobal: true,
PoWEnabled: true,
},
},
Bindings: []WAFBinding{
{RouteID: 1, SiteName: "example.com", RuleGroupIDs: []uint{}},
},
}
wafConfig, err := RenderWAFConfig(doc)
if err != nil {
t.Fatalf("RenderWAFConfig() error = %v", err)
}
var decoded WAFDocument
if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if len(decoded.RuleGroups) != 1 {
t.Fatalf("expected 1 rule group, got %d", len(decoded.RuleGroups))
}
if decoded.RuleGroups[0].PoWConfig != nil {
t.Fatalf("expected renderer not to synthesize legacy PoW config, got %#v", decoded.RuleGroups[0].PoWConfig)
}
}
func TestGetPoWConfigForRouteUsesGlobalGroupWithoutExplicitBinding(t *testing.T) {
snapshot := WAFDocument{
RuleGroups: []WAFRuleGroup{
{
ID: 1, Name: "global", Enabled: true, IsGlobal: true,
Graph: WAFRuleGraph{Entry: "pow", Nodes: map[string]WAFRuleNode{"pow": {Type: "pow"}}},
},
},
Bindings: []WAFBinding{
{RouteID: 42, SiteName: "example.com", RuleGroupIDs: []uint{}},
},
}
enabled := getPoWConfigForRoute(42, snapshot)
if !enabled {
t.Fatal("expected pow to be enabled via global rule group")
}
}
func TestRenderRouteConfigEnablesPoWLocationsFromRuntimeGraph(t *testing.T) {
doc := Document{
Routes: []Route{{ID: 1, SiteName: "pow.example.com", Domains: []string{"pow.example.com"}, OriginURL: "http://127.0.0.1:8080", Enabled: true}},
WAF: WAFDocument{
RuleGroups: []WAFRuleGroup{{
ID: 1, Name: "graph-pow", Enabled: true, IsGlobal: true,
Graph: WAFRuleGraph{Entry: "start", Nodes: map[string]WAFRuleNode{
"start": {Type: "start", Next: map[string]string{"next": "pow"}},
"pow": {Type: "pow", Config: json.RawMessage(`{"algorithm":"fast","difficulty":4,"session_ttl":600,"challenge_ttl":300}`), Next: map[string]string{"next": "allow"}},
"allow": {Type: "allow"},
}},
}},
Bindings: []WAFBinding{{RouteID: 1, SiteName: "pow.example.com", RuleGroupIDs: []uint{}}},
},
}
rendered, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
for _, expected := range []string{
`location = /.within.website/x/cmd/anubis/api/make-challenge`,
`location = /.within.website/x/cmd/anubis/api/pass-challenge`,
`location /.within.website/x/cmd/anubis/static/`,
} {
if !strings.Contains(rendered, expected) {
t.Fatalf("expected graph PoW route to contain %q, got:\n%s", expected, rendered)
}
}
}
func TestRenderWAFConfigPreservesRuntimeGraphAndBindingOrder(t *testing.T) {
doc := WAFDocument{
RuleGroups: []WAFRuleGroup{{
ID: 9, Name: "graph", Enabled: true,
Graph: WAFRuleGraph{Entry: "start", Nodes: map[string]WAFRuleNode{
"start": {Type: "start", Next: map[string]string{"next": "allow"}},
"allow": {Type: "allow"},
}},
}},
Bindings: []WAFBinding{{RouteID: 3, SiteName: "ordered.example.com", RuleGroupIDs: []uint{9, 4, 7}}},
}
raw, err := RenderWAFConfig(doc)
if err != nil {
t.Fatalf("RenderWAFConfig() error = %v", err)
}
var decoded WAFDocument
if err := json.Unmarshal([]byte(raw), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if decoded.RuleGroups[0].Graph.Entry != "start" {
t.Fatalf("runtime graph was not preserved: %#v", decoded.RuleGroups[0].Graph)
}
if got := decoded.Bindings[0].RuleGroupIDs; len(got) != 3 || got[0] != 9 || got[1] != 4 || got[2] != 7 {
t.Fatalf("binding order changed: %#v", got)
}
}
func TestRenderPagesAPIProxyLocationBlock(t *testing.T) {
tests := []struct {
name string
deployment *PagesDeployment
expected []string
unexpected []string
}{
{
name: "nil deployment",
deployment: nil,
expected: []string{""},
},
{
name: "disabled proxy",
deployment: &PagesDeployment{
APIProxyEnabled: false,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
},
expected: []string{""},
},
{
name: "enabled proxy without rewrite",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "",
},
expected: []string{
"location /api {",
"proxy_pass http://127.0.0.1:8080;",
"proxy_http_version 1.1;",
"proxy_set_header Host $http_host;",
},
unexpected: []string{
"rewrite",
},
},
{
name: "enabled proxy with rewrite to root",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "/",
},
expected: []string{
"location /api {",
"rewrite ^/api/(.*)$ /$1 break;",
"rewrite ^/api$ / break;",
"proxy_pass http://127.0.0.1:8080;",
},
},
{
name: "enabled proxy with rewrite to subpath",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "/v2",
},
expected: []string{
"location /api {",
"rewrite ^/api/(.*)$ /v2/$1 break;",
"rewrite ^/api$ /v2 break;",
"proxy_pass http://127.0.0.1:8080;",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := renderPagesAPIProxyLocationBlock(tt.deployment)
if len(tt.expected) == 1 && tt.expected[0] == "" {
if got != "" {
t.Fatalf("expected empty output, got: %q", got)
}
return
}
for _, exp := range tt.expected {
if !strings.Contains(got, exp) {
t.Errorf("expected output to contain %q, but got:\n%s", exp, got)
}
}
for _, unexp := range tt.unexpected {
if strings.Contains(got, unexp) {
t.Errorf("expected output NOT to contain %q, but got:\n%s", unexp, got)
}
}
})
}
}
func TestRenderPagesRootLocationBlock(t *testing.T) {
tests := []struct {
name string
deployment *PagesDeployment
expected []string
unexpected []string
}{
{
name: "spa fallback disabled serves entry file at root",
deployment: &PagesDeployment{
SPAFallbackEnabled: false,
EntryFile: "index.html",
},
expected: []string{
"location = / {",
"try_files /index.html =404;",
},
},
{
name: "spa fallback disabled with custom entry file",
deployment: &PagesDeployment{
SPAFallbackEnabled: false,
EntryFile: "app.html",
},
expected: []string{
"location = / {",
"try_files /app.html =404;",
},
},
{
name: "spa fallback enabled serves fallback file at root",
deployment: &PagesDeployment{
SPAFallbackEnabled: true,
SPAFallbackPath: "/index.html",
},
expected: []string{
"location = / {",
"try_files /index.html =404;",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := renderPagesRootLocationBlock(tt.deployment, routeLimitConfig{}, false, "", "")
if len(tt.expected) == 1 && tt.expected[0] == "" {
if got != "" {
t.Fatalf("expected empty output, got: %q", got)
}
return
}
for _, exp := range tt.expected {
if !strings.Contains(got, exp) {
t.Errorf("expected output to contain %q, but got:\n%s", exp, got)
}
}
for _, unexp := range tt.unexpected {
if strings.Contains(got, unexp) {
t.Errorf("expected output NOT to contain %q, but got:\n%s", unexp, got)
}
}
})
}
}
func TestRenderRouteConfigPagesWithoutSPAFallbackServesRoot(t *testing.T) {
doc := Document{
Routes: []Route{
{
ID: 1,
SiteName: "speedtest.example.com",
Domains: []string{"speedtest.example.com"},
UpstreamType: "pages",
EnableHTTPS: false,
PagesDeployment: &PagesDeployment{
LocalRoot: "/data/var/lib/openflare/pages/projects/1/current",
EntryFile: "index.html",
SPAFallbackEnabled: false,
},
},
},
}
routeConfig, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
if !strings.Contains(routeConfig, "location = / {") {
t.Fatalf("expected root location block, got:\n%s", routeConfig)
}
if !strings.Contains(routeConfig, "try_files /index.html =404;") {
t.Fatalf("expected root try_files for entry file, got:\n%s", routeConfig)
}
if !strings.Contains(routeConfig, "try_files $uri $uri/ =404;") {
t.Fatalf("expected static file try_files in location /, got:\n%s", routeConfig)
}
}
func TestRenderRouteConfigPagesWithSPAFallbackServesRoot(t *testing.T) {
doc := Document{
Routes: []Route{
{
ID: 1,
SiteName: "speedtest.example.com",
Domains: []string{"speedtest.example.com"},
UpstreamType: "pages",
EnableHTTPS: false,
PagesDeployment: &PagesDeployment{
LocalRoot: "/data/var/lib/openflare/pages/projects/1/current",
EntryFile: "index.html",
SPAFallbackEnabled: true,
SPAFallbackPath: "/index.html",
},
},
},
}
routeConfig, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
if !strings.Contains(routeConfig, "location = / {") {
t.Fatalf("expected root location block for spa fallback, got:\n%s", routeConfig)
}
if !strings.Contains(routeConfig, "try_files $uri $uri/ /index.html;") {
t.Fatalf("expected spa fallback try_files in location /, got:\n%s", routeConfig)
}
}
func TestRenderRouteCachePolicyConditionStaticDefault(t *testing.T) {
staticBlock := renderRouteCachePolicyCondition(routeCacheConfig{Enabled: true, Policy: "static"})
if staticBlock == "" {
t.Fatal("static policy should emit a path condition")
}
if !strings.Contains(staticBlock, "css") || !strings.Contains(staticBlock, "woff2") {
t.Fatalf("static policy should include default extensions, got:\n%s", staticBlock)
}
if !strings.Contains(staticBlock, "map") || !strings.Contains(staticBlock, "mjs") {
t.Fatalf("static policy should include map and mjs, got:\n%s", staticBlock)
}
if strings.Contains(staticBlock, "html") {
t.Fatalf("static policy must not include html, got:\n%s", staticBlock)
}
// Pattern is \.(?:css|js|...)$ — reject bare "json" as an alternation token.
if strings.Contains(staticBlock, "|json|") || strings.Contains(staticBlock, "|json)") || strings.Contains(staticBlock, "(?:json|") {
t.Fatalf("static policy must not include json (CF default), got:\n%s", staticBlock)
}
// Legacy empty/url = all (wide cache after method bypass).
emptyPolicy := renderRouteCachePolicyCondition(routeCacheConfig{Enabled: true, Policy: ""})
if emptyPolicy != "" {
t.Fatalf("empty policy should map to all (no path filter), got %q", emptyPolicy)
}
allBlock := renderRouteCachePolicyCondition(routeCacheConfig{Enabled: true, Policy: "all"})
if allBlock != "" {
t.Fatalf("all policy should not add path condition, got %q", allBlock)
}
urlBlock := renderRouteCachePolicyCondition(routeCacheConfig{Enabled: true, Policy: "url"})
if urlBlock != "" {
t.Fatalf("legacy url policy should map to all, got %q", urlBlock)
}
}
func TestRenderRouteCacheBlockAlignsCloudflareDefaults(t *testing.T) {
block := renderRouteCacheBlock(
routeCacheConfig{Enabled: true, Policy: "static"},
ConfigSnapshot{CacheEnabled: true},
)
if !strings.Contains(block, "proxy_cache openflare_cache") {
t.Fatalf("expected proxy_cache, got:\n%s", block)
}
if !strings.Contains(block, "\\.(?:") {
t.Fatalf("expected static suffix pattern, got:\n%s", block)
}
if !strings.Contains(block, "request_method != GET") {
t.Fatalf("expected method bypass for non-GET, got:\n%s", block)
}
if strings.Contains(block, "$http_authorization") {
t.Fatalf("must not bypass on Authorization (CF-aligned), got:\n%s", block)
}
if strings.Contains(block, "$http_cookie") {
t.Fatalf("must not bypass on Cookie (CF-aligned), got:\n%s", block)
}
if strings.Contains(block, "$http_cache_control") {
t.Fatalf("must not bypass on request Cache-Control (CF-aligned), got:\n%s", block)
}
if !strings.Contains(block, "proxy_no_cache $openflare_skip_cache $upstream_http_set_cookie") {
t.Fatalf("expected Set-Cookie no-cache gate, got:\n%s", block)
}
if !strings.Contains(block, "proxy_cache_valid 200 206 301 120m") {
t.Fatalf("expected default Edge TTL for 200/206/301, got:\n%s", block)
}
if !strings.Contains(block, "proxy_cache_valid 302 303 20m") {
t.Fatalf("expected default Edge TTL for 302/303, got:\n%s", block)
}
if !strings.Contains(block, "proxy_cache_valid 404 410 3m") {
t.Fatalf("expected default Edge TTL for 404/410, got:\n%s", block)
}
if !strings.Contains(block, "proxy_cache_bypass $openflare_skip_cache") {
t.Fatalf("expected proxy_cache_bypass on skip flag only, got:\n%s", block)
}
}
func TestMergeRouteLimitConfig(t *testing.T) {
t.Parallel()
cases := []struct {
name string
route Route
cfg ConfigSnapshot
want routeLimitConfig
}{
{
name: "both zero off",
route: Route{},
cfg: ConfigSnapshot{},
want: routeLimitConfig{},
},
{
name: "inherit all defaults",
route: Route{},
cfg: ConfigSnapshot{
DefaultLimitConnPerServer: 100,
DefaultLimitConnPerIP: 10,
DefaultLimitRate: "512k",
DefaultLimitReqPerIP: "10r/s",
},
want: routeLimitConfig{LimitConnPerServer: 100, LimitConnPerIP: 10, LimitRate: "512k", LimitReqPerIP: "10r/s"},
},
{
name: "explicit off ignores default",
route: Route{LimitConnPerServer: -1, LimitConnPerIP: -1, LimitRate: "-1", LimitReqPerIP: "-1"},
cfg: ConfigSnapshot{
DefaultLimitConnPerServer: 100,
DefaultLimitConnPerIP: 10,
DefaultLimitRate: "512k",
DefaultLimitReqPerIP: "10r/s",
},
want: routeLimitConfig{},
},
{
name: "route overrides default",
route: Route{LimitConnPerServer: 50, LimitConnPerIP: 5, LimitRate: "1m"},
cfg: ConfigSnapshot{
DefaultLimitConnPerServer: 100,
DefaultLimitConnPerIP: 10,
DefaultLimitRate: "512k",
DefaultLimitReqPerIP: "10r/s",
},
want: routeLimitConfig{LimitConnPerServer: 50, LimitConnPerIP: 5, LimitRate: "1m", LimitReqPerIP: "10r/s"},
},
{
name: "partial inherit",
route: Route{LimitConnPerServer: 0, LimitConnPerIP: -1, LimitRate: ""},
cfg: ConfigSnapshot{
DefaultLimitConnPerServer: 100,
DefaultLimitConnPerIP: 10,
DefaultLimitRate: "256k",
DefaultLimitReqPerIP: "10r/s",
},
want: routeLimitConfig{LimitConnPerServer: 100, LimitConnPerIP: 0, LimitRate: "256k", LimitReqPerIP: "10r/s"},
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
got := mergeRouteLimitConfig(tc.route, tc.cfg)
if got != tc.want {
t.Fatalf("mergeRouteLimitConfig() = %#v, want %#v", got, tc.want)
}
})
}
}
func TestRenderRouteConfigAppliesDefaultLimits(t *testing.T) {
doc := Document{
Routes: []Route{{
SiteName: "example.com",
Domains: []string{"example.com"},
Enabled: true,
OriginURL: "http://127.0.0.1:8080",
Upstreams: []string{"http://127.0.0.1:8080"},
}},
OpenRestyConfig: ConfigSnapshot{
DefaultLimitConnPerServer: 120,
DefaultLimitConnPerIP: 12,
DefaultLimitRate: "512k",
DefaultLimitReqPerIP: "10r/s",
},
}
rendered, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
for _, want := range []string{
"limit_conn openflare_conn_per_server 120;",
"limit_conn openflare_conn_per_ip 12;",
"limit_rate 512k;",
"limit_req zone=openflare_req_10rs burst=20 nodelay;",
"limit_req_status 429;",
} {
if !strings.Contains(rendered, want) {
t.Fatalf("expected %q in route config, got:\n%s", want, rendered)
}
}
}
func TestRenderMainConfigEmitsLimitReqZonesByEffectiveRate(t *testing.T) {
doc := Document{
Routes: []Route{
{
SiteName: "a.example.com",
Domains: []string{"a.example.com"},
Enabled: true,
OriginURL: "http://127.0.0.1:8080",
Upstreams: []string{"http://127.0.0.1:8080"},
},
{
SiteName: "b.example.com",
Domains: []string{"b.example.com"},
Enabled: true,
OriginURL: "http://127.0.0.1:8081",
Upstreams: []string{"http://127.0.0.1:8081"},
LimitReqPerIP: "5r/s",
},
{
SiteName: "c.example.com",
Domains: []string{"c.example.com"},
Enabled: true,
OriginURL: "http://127.0.0.1:8082",
Upstreams: []string{"http://127.0.0.1:8082"},
LimitReqPerIP: "-1",
},
},
OpenRestyConfig: ConfigSnapshot{
DefaultLimitReqPerIP: "10r/s",
},
}
mainConfig := RenderMainConfig(doc)
for _, want := range []string{
"limit_req_zone $openflare_waf_site$binary_remote_addr zone=openflare_req_10rs:10m rate=10r/s;",
"limit_req_zone $openflare_waf_site$binary_remote_addr zone=openflare_req_5rs:10m rate=5r/s;",
} {
if !strings.Contains(mainConfig, want) {
t.Fatalf("expected %q in main config, got:\n%s", want, mainConfig)
}
}
if strings.Contains(mainConfig, "openflare_req_per_ip") {
t.Fatalf("unexpected legacy zone name in main config:\n%s", mainConfig)
}
routeConfig, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
if !strings.Contains(routeConfig, "limit_req zone=openflare_req_10rs burst=20 nodelay;") {
t.Fatalf("expected inherited zone on route a, got:\n%s", routeConfig)
}
if !strings.Contains(routeConfig, "limit_req zone=openflare_req_5rs burst=10 nodelay;") {
t.Fatalf("expected custom zone on route b, got:\n%s", routeConfig)
}
// route c is off: count limit_req lines should equal 2 routes * (http+https? depends) — assert c server has no limit_req by site name block is hard; ensure -1 route does not force extra zones
if strings.Count(mainConfig, "limit_req_zone") != 2 {
t.Fatalf("expected exactly 2 limit_req_zone lines, got main:\n%s", mainConfig)
}
}
func TestRenderRouteConfigExplicitOffSkipsDefaultLimits(t *testing.T) {
doc := Document{
Routes: []Route{{
SiteName: "example.com",
Domains: []string{"example.com"},
Enabled: true,
OriginURL: "http://127.0.0.1:8080",
Upstreams: []string{"http://127.0.0.1:8080"},
LimitConnPerServer: -1,
LimitConnPerIP: -1,
LimitRate: "-1",
LimitReqPerIP: "-1",
}},
OpenRestyConfig: ConfigSnapshot{
DefaultLimitConnPerServer: 120,
DefaultLimitConnPerIP: 12,
DefaultLimitRate: "512k",
DefaultLimitReqPerIP: "10r/s",
},
}
rendered, err := RenderRouteConfig(doc, nil)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
if strings.Contains(rendered, "limit_conn") || strings.Contains(rendered, "limit_rate") || strings.Contains(rendered, "limit_req") {
t.Fatalf("expected no limit directives, got:\n%s", rendered)
}
}
@@ -0,0 +1,136 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"crypto/sha256"
"encoding/hex"
"strings"
)
// SW location strings and Lua module paths used by the Service Worker offline fallback.
const (
SWJSLocation = "location = /sw.js"
SWOfflineLocation = "location = /offline.html"
SWChallengeLua = "sw/challenge.lua"
SWRuntimeLua = "sw/runtime.lua"
swDirPrefix = "sw/"
)
// DefaultSWOfflineHTML is the built-in contact page shown when the domain is blocked.
const DefaultSWOfflineHTML = `<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>网站暂时无法访问 | 联系站长</title>
<style>
* { box-sizing: border-box; margin: 0; padding: 0; }
body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif; background: #ffffff; color: #333333; height: 100vh; display: flex; flex-direction: column; justify-content: center; align-items: center; text-align: center; padding: 48px 24px; }
h1 { font-size: 28px; font-weight: 700; margin-bottom: 16px; }
p { font-size: 16px; line-height: 1.7; color: #666666; max-width: 520px; }
</style>
</head>
<body>
<h1>网站暂时无法访问</h1>
<p>当前域名暂时无法从网络访问。请通过其他方式联系网站管理员获取最新访问入口。</p>
</body>
</html>
`
// EffectiveSWOfflineHTML returns custom HTML when set, otherwise the built-in default.
func EffectiveSWOfflineHTML(cfg ConfigSnapshot) string {
if strings.TrimSpace(cfg.SWOfflineHTML) == "" {
return DefaultSWOfflineHTML
}
return cfg.SWOfflineHTML
}
// ServiceWorkerSupportFiles returns the sw.js script and offline contact page.
// The sw.js content is derived from the offline HTML (see defaultSWJS) so that
// HTML-only edits change the script, forcing browsers to re-install the worker
// and re-cache the updated page.
func ServiceWorkerSupportFiles(cfg ConfigSnapshot) []SupportFile {
if !cfg.SWOfflineEnabled {
return nil
}
html := EffectiveSWOfflineHTML(cfg)
return []SupportFile{
{Path: swDirPrefix + "sw.js", Content: defaultSWJS(html)},
{Path: swDirPrefix + "offline.html", Content: html},
}
}
// swJSTemplate is the service worker body. The cache name is replaced with a
// version derived from the offline HTML: editing the HTML changes the cache
// name, which changes the sw.js bytes, which makes the browser re-install the
// worker (sw.js is served with Cache-Control: no-cache) and fetch the new
// /offline.html into the fresh cache during install.
const swJSTemplate = `var CACHE = "__CACHE_NAME__";
var OFFLINE = "/offline.html";
self.addEventListener("install", function (e) {
e.waitUntil(caches.open(CACHE).then(function (c) { return c.addAll([OFFLINE]); }));
self.skipWaiting();
});
self.addEventListener("activate", function (e) {
e.waitUntil(caches.keys().then(function (keys) {
return Promise.all(keys.filter(function (k) { return k.indexOf("openflare-offline-") === 0 && k !== CACHE; }).map(function (k) { return caches.delete(k); }));
}));
self.clients.claim();
});
self.addEventListener("fetch", function (e) {
if (e.request.method !== "GET" || e.request.mode !== "navigate") { return; }
e.respondWith(
fetch(e.request).catch(function () {
return caches.match(e.request).then(function (r) { return r || caches.match(OFFLINE); });
})
);
});
`
func defaultSWJS(offlineHTML string) string {
sum := sha256.Sum256([]byte(offlineHTML))
version := hex.EncodeToString(sum[:])[:12]
return strings.ReplaceAll(swJSTemplate, "__CACHE_NAME__", "openflare-offline-"+version)
}
// routeSWEnabled returns true when SW offline fallback applies to this route.
func routeSWEnabled(routeDomains []string, cfg ConfigSnapshot) bool {
if !cfg.SWOfflineEnabled || len(cfg.SWOfflineDomains) == 0 {
return false
}
scope := make(map[string]struct{}, len(cfg.SWOfflineDomains))
for _, d := range cfg.SWOfflineDomains {
scope[d] = struct{}{}
}
for _, d := range routeDomains {
if _, ok := scope[d]; ok {
return true
}
}
return false
}
// renderServiceWorkerChallenger emits SW static locations and the homepage
// challenge intercept for HTTPS server blocks.
func renderServiceWorkerChallenger(_ ConfigSnapshot) string {
var builder strings.Builder
builder.WriteString("\n location = /sw.js {\n")
builder.WriteString(" alias " + SWDirPlaceholder + "/sw.js;\n")
builder.WriteString(" default_type application/javascript;\n")
builder.WriteString(" add_header Service-Worker-Allowed /;\n")
builder.WriteString(" add_header Cache-Control \"no-cache\";\n")
builder.WriteString(" }\n\n")
builder.WriteString(" location = /offline.html {\n")
builder.WriteString(" alias " + SWDirPlaceholder + "/offline.html;\n")
builder.WriteString(" default_type text/html;\n")
builder.WriteString(" add_header Cache-Control \"no-cache\";\n")
builder.WriteString(" }\n\n")
builder.WriteString(" location = /__openflare_sw_challenge {\n")
builder.WriteString(" internal;\n")
builder.WriteString(" # hit when sw.runtime.check() intercepts the homepage in the access phase\n")
builder.WriteString(" content_by_lua_file " + SWDirPlaceholder + "/challenge.lua;\n")
builder.WriteString(" }\n")
return builder.String()
}
@@ -0,0 +1,285 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"fmt"
"math/big"
"strings"
"testing"
"time"
)
func TestEffectiveSWOfflineHTML(t *testing.T) {
if got := EffectiveSWOfflineHTML(ConfigSnapshot{}); got != DefaultSWOfflineHTML {
t.Fatalf("default mismatch")
}
custom := "<html>custom</html>"
if got := EffectiveSWOfflineHTML(ConfigSnapshot{SWOfflineHTML: custom}); got != custom {
t.Fatalf("custom mismatch")
}
}
func TestServiceWorkerSupportFiles(t *testing.T) {
disabled := ServiceWorkerSupportFiles(ConfigSnapshot{})
if disabled != nil {
t.Fatalf("expected nil when disabled, got %v", disabled)
}
enabled := ServiceWorkerSupportFiles(ConfigSnapshot{SWOfflineEnabled: true})
if len(enabled) != 2 {
t.Fatalf("expected 2 support files, got %d", len(enabled))
}
paths := map[string]string{}
for _, f := range enabled {
paths[f.Path] = f.Content
}
if _, ok := paths["sw/sw.js"]; !ok {
t.Fatalf("missing sw/sw.js")
}
if _, ok := paths["sw/offline.html"]; !ok {
t.Fatalf("missing sw/offline.html")
}
if paths["sw/offline.html"] != DefaultSWOfflineHTML {
t.Fatalf("expected built-in offline html, got %q", paths["sw/offline.html"])
}
if !strings.Contains(paths["sw/sw.js"], `var OFFLINE = "/offline.html";`) {
t.Fatalf("offline path must stay stable (exact location match), got:\n%s", paths["sw/sw.js"])
}
}
func TestDefaultSWJSCacheNameTracksOfflineHTML(t *testing.T) {
htmlA := "<html>page-a</html>"
htmlB := "<html>page-b</html>"
jsA := defaultSWJS(htmlA)
jsB := defaultSWJS(htmlB)
if jsA == jsB {
t.Fatal("sw.js content must change when the offline HTML changes")
}
extractCache := func(js string) string {
const prefix = `var CACHE = "`
start := strings.Index(js, prefix)
if start < 0 {
t.Fatalf("missing cache name in:\n%s", js)
}
rest := js[start+len(prefix):]
end := strings.Index(rest, `"`)
if end < 0 {
t.Fatalf("unterminated cache name in:\n%s", js)
}
return rest[:end]
}
cacheA := extractCache(jsA)
cacheB := extractCache(jsB)
if cacheA == cacheB {
t.Fatalf("cache names must differ per HTML, got %q", cacheA)
}
if !strings.HasPrefix(cacheA, "openflare-offline-") {
t.Fatalf("unexpected cache name %q", cacheA)
}
for _, js := range []string{jsA, jsB} {
if strings.Contains(js, "openflare-offline-v1") {
t.Fatalf("static cache name must not remain, got:\n%s", js)
}
if strings.Contains(js, "__CACHE_NAME__") {
t.Fatalf("template placeholder leaked into sw.js:\n%s", js)
}
}
// Same HTML must produce identical sw.js (deterministic checksum).
if defaultSWJS(htmlA) != jsA {
t.Fatal("sw.js must be deterministic for identical HTML")
}
}
func TestRenderAccessBlockWithSWMergesSingleBlock(t *testing.T) {
for _, powEnabled := range []bool{false, true} {
name := "pow-disabled"
if powEnabled {
name = "pow-enabled"
}
t.Run(name, func(t *testing.T) {
got := renderAccessBlockWithSW("example.com", powEnabled, ConfigSnapshot{})
if n := strings.Count(got, "access_by_lua_block"); n != 1 {
t.Fatalf("expected exactly 1 access_by_lua_block, got %d:\n%s", n, got)
}
if !strings.Contains(got, `require("sw.runtime").check()`) {
t.Fatalf("expected sw.runtime check, got:\n%s", got)
}
wafIdx := strings.Index(got, `require("waf.runtime").check()`)
swIdx := strings.Index(got, `require("sw.runtime").check()`)
if wafIdx < 0 || swIdx < 0 || wafIdx > swIdx {
t.Fatalf("expected waf.runtime before sw.runtime, got:\n%s", got)
}
if powEnabled {
powIdx := strings.Index(got, `require("pow.runtime").check()`)
if powIdx < 0 || wafIdx > powIdx || powIdx > swIdx {
t.Fatalf("expected waf.runtime before pow.runtime before sw.runtime, got:\n%s", got)
}
}
})
}
}
func TestRenderServiceWorkerChallengerHTTPExclusion(t *testing.T) {
cfg := ConfigSnapshot{SWOfflineEnabled: true}
for name, rendered := range map[string]string{
"proxy": renderHTTPProxyServer("example.com", "example.com", "http://127.0.0.1:8080", "", nil, routeCacheConfig{}, routeLimitConfig{}, routeUpstreamConfig{}, false, false, "", "", false, cfg),
"pages": renderHTTPPagesServer("example.com", "example.com", nil, routeLimitConfig{}, false, false, "", "", false, cfg),
"https": renderHTTPSServer("example.com", "example.com", "http://127.0.0.1:8080", "", 1, nil, routeCacheConfig{}, routeLimitConfig{}, routeUpstreamConfig{}, false, false, "", "", true, cfg),
"hpages": renderHTTPSPagesServer("example.com", "example.com", 1, nil, routeLimitConfig{}, false, false, "", "", true, cfg),
} {
if strings.Contains(rendered, "access_by_lua_block") && strings.Count(rendered, "access_by_lua_block") != 1 {
t.Fatalf("%s: expected at most one access block, got:\n%s", name, rendered)
}
}
httpProxy := renderHTTPProxyServer("example.com", "example.com", "http://127.0.0.1:8080", "", nil, routeCacheConfig{}, routeLimitConfig{}, routeUpstreamConfig{}, false, false, "", "", false, cfg)
if strings.Contains(httpProxy, "sw.runtime") || strings.Contains(httpProxy, "openflare_sw_challenge") || strings.Contains(httpProxy, "location = /sw.js") {
t.Fatalf("HTTP proxy server must not carry SW intercept, got:\n%s", httpProxy)
}
httpPages := renderHTTPPagesServer("example.com", "example.com", nil, routeLimitConfig{}, false, false, "", "", false, cfg)
if strings.Contains(httpPages, "sw.runtime") || strings.Contains(httpPages, "openflare_sw_challenge") || strings.Contains(httpPages, "location = /sw.js") {
t.Fatalf("HTTP pages server must not carry SW intercept, got:\n%s", httpPages)
}
httpsProxy := renderHTTPSServer("example.com", "example.com", "http://127.0.0.1:8080", "", 1, nil, routeCacheConfig{}, routeLimitConfig{}, routeUpstreamConfig{}, false, false, "", "", true, cfg)
for _, want := range []string{"sw.runtime", "location = /sw.js", "location = /offline.html", "__openflare_sw_challenge"} {
if !strings.Contains(httpsProxy, want) {
t.Fatalf("HTTPS proxy server missing %q, got:\n%s", want, httpsProxy)
}
}
httpsPages := renderHTTPSPagesServer("example.com", "example.com", 1, nil, routeLimitConfig{}, false, false, "", "", true, cfg)
for _, want := range []string{"sw.runtime", "location = /sw.js", "location = /offline.html", "__openflare_sw_challenge"} {
if !strings.Contains(httpsPages, want) {
t.Fatalf("HTTPS pages server missing %q, got:\n%s", want, httpsPages)
}
}
}
func TestRouteSWEnabled(t *testing.T) {
cfgOff := ConfigSnapshot{SWOfflineEnabled: false, SWOfflineDomains: []string{"example.com"}}
if routeSWEnabled([]string{"example.com"}, cfgOff) {
t.Fatal("expected false when master switch off")
}
cfgEmpty := ConfigSnapshot{SWOfflineEnabled: true, SWOfflineDomains: nil}
if routeSWEnabled([]string{"example.com"}, cfgEmpty) {
t.Fatal("expected false when scope empty")
}
cfgHit := ConfigSnapshot{SWOfflineEnabled: true, SWOfflineDomains: []string{"example.com", "other.com"}}
if !routeSWEnabled([]string{"api.example.com", "example.com"}, cfgHit) {
t.Fatal("expected true on single domain intersection")
}
if routeSWEnabled([]string{"api.example.com", "third.com"}, cfgHit) {
t.Fatal("expected false on no intersection")
}
}
func TestRenderHTTPSServerSWScope(t *testing.T) {
render := func(swEnabled bool) string {
return renderHTTPSServer("example.com", "example.com", "http://127.0.0.1:8080", "", 1, nil, routeCacheConfig{}, routeLimitConfig{}, routeUpstreamConfig{}, false, false, "", "", swEnabled, ConfigSnapshot{SWOfflineEnabled: true})
}
hit := render(routeSWEnabled([]string{"example.com"}, ConfigSnapshot{SWOfflineEnabled: true, SWOfflineDomains: []string{"example.com"}}))
for _, want := range []string{`require("sw.runtime").check()`, "location = /sw.js", "location = /offline.html", "__openflare_sw_challenge"} {
if !strings.Contains(hit, want) {
t.Fatalf("scoped HTTPS server missing %q, got:\n%s", want, hit)
}
}
miss := render(routeSWEnabled([]string{"example.com"}, ConfigSnapshot{SWOfflineEnabled: true, SWOfflineDomains: []string{"other.com"}}))
for _, notWant := range []string{`require("sw.runtime").check()`, "location = /sw.js", "location = /offline.html", "__openflare_sw_challenge"} {
if strings.Contains(miss, notWant) {
t.Fatalf("out-of-scope HTTPS server must not carry %q, got:\n%s", notWant, miss)
}
}
if miss != renderHTTPSServer("example.com", "example.com", "http://127.0.0.1:8080", "", 1, nil, routeCacheConfig{}, routeLimitConfig{}, routeUpstreamConfig{}, false, false, "", "", false, ConfigSnapshot{}) {
t.Fatalf("out-of-scope HTTPS server must match pre-feature bytes, got:\n%s", miss)
}
}
func TestRenderRouteConfigSWSCOPEPerCertPartition(t *testing.T) {
doc := Document{
OpenRestyConfig: ConfigSnapshot{
SWOfflineEnabled: true,
SWOfflineDomains: []string{"a.com"},
},
Routes: []Route{{
ID: 1,
SiteName: "multi.example.com",
Domains: []string{"a.com", "b.com"},
OriginURL: "http://127.0.0.1:8080",
EnableHTTPS: true,
DomainCertIDs: []uint{11, 22},
}},
}
certFiles := []SupportFile{
{Path: "11.crt", Content: testCertificatePEMForDomain(t, "a.com")},
{Path: "22.crt", Content: testCertificatePEMForDomain(t, "b.com")},
}
rendered, err := RenderRouteConfig(doc, certFiles)
if err != nil {
t.Fatalf("RenderRouteConfig() error = %v", err)
}
inScope := httpsServerBlockForCert(t, rendered, 11)
if !strings.Contains(inScope, "server_name a.com;") {
t.Fatalf("cert 11 block must serve a.com, got:\n%s", inScope)
}
for _, want := range []string{`require("sw.runtime").check()`, "location = /sw.js", "location = /offline.html", "__openflare_sw_challenge"} {
if !strings.Contains(inScope, want) {
t.Fatalf("in-scope cert partition (a.com) missing %q, got:\n%s", want, inScope)
}
}
outOfScope := httpsServerBlockForCert(t, rendered, 22)
if !strings.Contains(outOfScope, "server_name b.com;") {
t.Fatalf("cert 22 block must serve b.com, got:\n%s", outOfScope)
}
for _, notWant := range []string{`require("sw.runtime").check()`, "location = /sw.js", "location = /offline.html", "__openflare_sw_challenge"} {
if strings.Contains(outOfScope, notWant) {
t.Fatalf("out-of-scope cert partition (b.com) must not carry %q, got:\n%s", notWant, outOfScope)
}
}
}
func httpsServerBlockForCert(t *testing.T, rendered string, certID uint) string {
t.Helper()
marker := fmt.Sprintf("ssl_certificate %s/%d.crt;", CertDirPlaceholder, certID)
for _, block := range strings.Split(rendered, "server {") {
if strings.Contains(block, marker) {
return "server {" + block
}
}
t.Fatalf("no server block found for cert %d in:\n%s", certID, rendered)
return ""
}
func testCertificatePEMForDomain(t *testing.T, domain string) string {
t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatalf("rsa.GenerateKey() error = %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(time.Now().UnixNano()),
Subject: pkix.Name{CommonName: domain},
DNSNames: []string{domain},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(24 * time.Hour),
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
if err != nil {
t.Fatalf("x509.CreateCertificate() error = %v", err)
}
return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}))
}
func TestRenderServiceWorkerChallenger(t *testing.T) {
got := renderServiceWorkerChallenger(ConfigSnapshot{SWOfflineEnabled: true})
for _, want := range []string{"location = /sw.js", "location = /offline.html", "challenge.lua", "content_by_lua"} {
if !strings.Contains(got, want) {
t.Fatalf("challenger missing %q", want)
}
}
}
@@ -0,0 +1,71 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"errors"
"fmt"
"sort"
"strconv"
"strings"
)
const (
// StatusCodeMin is the lowest HTTP status code accepted for origin error pages.
StatusCodeMin = 400
// StatusCodeMax is the highest HTTP status code accepted for origin error pages.
StatusCodeMax = 599
)
// ParseStatusCodeTag parses a single tag such as "502" or "500-599".
// Bounds must fall within StatusCodeMin–StatusCodeMax inclusive.
func ParseStatusCodeTag(tag string) (lo, hi int, err error) {
tag = strings.TrimSpace(tag)
if tag == "" {
return 0, 0, errors.New("状态码标签不能为空")
}
if before, after, ok := strings.Cut(tag, "-"); ok {
lo, err = strconv.Atoi(before)
if err != nil {
return 0, 0, fmt.Errorf("无效状态码区间: %s", tag)
}
hi, err = strconv.Atoi(after)
if err != nil {
return 0, 0, fmt.Errorf("无效状态码区间: %s", tag)
}
} else {
lo, err = strconv.Atoi(tag)
if err != nil {
return 0, 0, fmt.Errorf("无效状态码: %s", tag)
}
hi = lo
}
if lo > hi {
return 0, 0, fmt.Errorf("状态码区间左右端点反序: %s", tag)
}
if lo < StatusCodeMin || hi > StatusCodeMax {
return 0, 0, fmt.Errorf("状态码须在 %d–%d: %s", StatusCodeMin, StatusCodeMax, tag)
}
return lo, hi, nil
}
// ExpandStatusCodeTags expands status code tags into a sorted unique list of integers.
func ExpandStatusCodeTags(tags []string) ([]int, error) {
set := map[int]struct{}{}
for _, tag := range tags {
lo, hi, err := ParseStatusCodeTag(tag)
if err != nil {
return nil, err
}
for c := lo; c <= hi; c++ {
set[c] = struct{}{}
}
}
out := make([]int, 0, len(set))
for c := range set {
out = append(out, c)
}
sort.Ints(out)
return out, nil
}
@@ -0,0 +1,30 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import "testing"
func TestExpandStatusCodeTags(t *testing.T) {
t.Parallel()
codes, err := ExpandStatusCodeTags([]string{"500-502", "522", "501"})
if err != nil {
t.Fatal(err)
}
// want sorted unique: 500,501,502,522
if len(codes) != 4 || codes[0] != 500 || codes[3] != 522 {
t.Fatalf("got %v", codes)
}
_, err = ExpandStatusCodeTags([]string{"399"})
if err == nil {
t.Fatal("expected error")
}
_, err = ExpandStatusCodeTags([]string{"503-500"})
if err == nil {
t.Fatal("expected reverse range error")
}
_, err = ExpandStatusCodeTags([]string{"5xx"})
if err == nil {
t.Fatal("expected syntax error")
}
}
@@ -0,0 +1,403 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openresty
import (
"encoding/json"
"fmt"
)
// Placeholder constants used as sentinel values in rendered OpenResty config
// files; the deploy process replaces them with real paths before reload.
const (
CertDirPlaceholder = "__OPENFLARE_CERT_DIR__"
RouteConfigPlaceholder = "__OPENFLARE_ROUTE_CONFIG__"
AccessLogPlaceholder = "__OPENFLARE_ACCESS_LOG__"
ErrorLogPlaceholder = "__OPENFLARE_ERROR_LOG__"
PIDPathPlaceholder = "__OPENFLARE_PID_PATH__"
NginxCacheDirPlaceholder = "__OPENFLARE_NGINX_CACHE_DIR__"
ProxyCachePathPlaceholder = "__OPENFLARE_PROXY_CACHE_PATH__"
LuaDirPlaceholder = "__OPENFLARE_LUA_DIR__"
ObservabilityListenPlaceholder = "__OPENFLARE_OBSERVABILITY_LISTEN__"
ObservabilityPortPlaceholder = "__OPENFLARE_OBSERVABILITY_PORT__"
PowStaticDirPlaceholder = "__OPENFLARE_POW_STATIC_DIR__"
PagesDirPlaceholder = "__OPENFLARE_PAGES_DIR__"
ErrorPageTmplPlaceholder = "__OPENFLARE_ERROR_PAGE_TMPL__"
SWDirPlaceholder = "__OPENFLARE_SW_DIR__"
SourceConfigFileName = "openresty_config.json"
)
const (
cachePolicyStatic = "static"
cachePolicyAll = "all"
cachePolicyURL = "url" // legacy alias of all
cachePolicySuffix = "suffix"
cachePolicyPathPrefix = "path_prefix"
cachePolicyPathExact = "path_exact"
defaultWAFBlockStatus = 418
anubisStaticPrefix = "/.within.website/x/cmd/anubis/static/"
anubisAPIPrefix = "/.within.website/x/cmd/anubis/api/"
)
// DefaultStaticCacheExtensions is the built-in suffix allowlist for cache_policy=static.
// HTML and JSON are excluded (Cloudflare default). map/mjs/wasm are intentional extras.
var DefaultStaticCacheExtensions = []string{
"css", "js", "mjs", "map",
"ico", "cur", "gif", "jpg", "jpeg", "png", "webp", "avif", "svg", "svgz",
"ttf", "otf", "woff", "woff2", "eot",
"mp3", "mp4", "webm", "ogg", "flac",
"wasm", "pdf",
"zip", "7z", "gz", "tar",
}
// OpenFlareRuntimeUser is the dedicated service account shared by the agent
// process and OpenResty worker processes.
const OpenFlareRuntimeUser = "openflare"
// OpenRestyWorkerUser is kept as an alias for existing call sites.
const OpenRestyWorkerUser = OpenFlareRuntimeUser
const defaultMainConfigTemplate = `# This file is generated by OpenFlare. Do not edit manually.
user ` + OpenFlareRuntimeUser + `;
worker_processes {{OpenRestyWorkerProcesses}};
worker_rlimit_nofile {{OpenRestyWorkerRlimitNofile}};
pid __OPENFLARE_PID_PATH__;
error_log {{OpenRestyErrorLogPath}} warn;
events {
worker_connections {{OpenRestyWorkerConnections}};
{{OpenRestyEventsUseDirective}}{{OpenRestyEventsMultiAcceptDirective}}}
http {
include mime.types;
default_type application/octet-stream;
server_tokens off;
client_body_temp_path __OPENFLARE_NGINX_CACHE_DIR__/client_temp;
proxy_temp_path __OPENFLARE_NGINX_CACHE_DIR__/proxy_temp;
fastcgi_temp_path __OPENFLARE_NGINX_CACHE_DIR__/fastcgi_temp;
uwsgi_temp_path __OPENFLARE_NGINX_CACHE_DIR__/uwsgi_temp;
scgi_temp_path __OPENFLARE_NGINX_CACHE_DIR__/scgi_temp;
{{OpenRestyConnectionUpgradeMap}}{{OpenRestyDefaultServerBlock}} log_format openflare_json escape=json '{"ts":"$time_iso8601","host":"$host","path":"$request_uri","remote_addr":"$remote_addr","status":$status,"request_time":$request_time,"bytes_sent":$body_bytes_sent,"request_length":$request_length,"user_agent":"$http_user_agent","cache_status":"$upstream_cache_status"}';
access_log {{OpenRestyAccessLogPath}} openflare_json;
sendfile on;
tcp_nopush on;
tcp_nodelay on;
keepalive_timeout {{OpenRestyKeepaliveTimeout}};
keepalive_requests {{OpenRestyKeepaliveRequests}};
client_header_timeout {{OpenRestyClientHeaderTimeout}};
client_body_timeout {{OpenRestyClientBodyTimeout}};
client_max_body_size {{OpenRestyClientMaxBodySize}};
large_client_header_buffers {{OpenRestyLargeClientHeaderBuffers}};
send_timeout {{OpenRestySendTimeout}};
proxy_connect_timeout {{OpenRestyProxyConnectTimeout}};
proxy_send_timeout {{OpenRestyProxySendTimeout}};
proxy_read_timeout {{OpenRestyProxyReadTimeout}};
proxy_request_buffering {{OpenRestyProxyRequestBuffering}};
proxy_buffering {{OpenRestyProxyBuffering}};
proxy_buffers {{OpenRestyProxyBuffers}};
proxy_buffer_size {{OpenRestyProxyBufferSize}};
proxy_busy_buffers_size {{OpenRestyProxyBusyBuffersSize}};
gzip {{OpenRestyGzip}};
gzip_min_length {{OpenRestyGzipMinLength}};
gzip_comp_level {{OpenRestyGzipCompLevel}};
{{OpenRestyResolverDirective}}{{OpenRestyCacheBlock}} include {{OpenRestyRouteConfigInclude}};
}
`
// SupportFile represents an auxiliary file (certificate, WAF config, etc.)
// that is written alongside the main OpenResty configuration.
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
// CustomHeader is a key/value pair injected as an additional proxy_set_header
// directive for a specific route.
type CustomHeader struct {
Key string `json:"key"`
Value string `json:"value"`
}
// PoWListConfig holds the IP, CIDR, path, and user-agent lists used by the
// Proof-of-Work whitelist or blacklist filter.
type PoWListConfig struct {
IPs []string `json:"ips"`
IPCidrs []string `json:"ip_cidrs"`
Paths []string `json:"paths"`
PathRegexes []string `json:"path_regexes"`
UserAgents []string `json:"user_agents"`
}
// PoWConfig holds the full Proof-of-Work challenge parameters for a route,
// including difficulty, algorithm, TTLs, and allow/block lists.
type PoWConfig struct {
Difficulty int `json:"difficulty"`
Algorithm string `json:"algorithm"`
SessionTTL int `json:"session_ttl"`
ChallengeTTL int `json:"challenge_ttl"`
Whitelist PoWListConfig `json:"whitelist"`
Blacklist PoWListConfig `json:"blacklist"`
}
// DefaultPoWConfig returns the canonical PoW defaults used when pow_enabled is
// true but no explicit pow_config payload is available.
func DefaultPoWConfig() PoWConfig {
return PoWConfig{
Difficulty: 4,
Algorithm: "fast",
SessionTTL: 600,
ChallengeTTL: 300,
Whitelist: PoWListConfig{IPs: []string{}, IPCidrs: []string{}, Paths: []string{}, PathRegexes: []string{}, UserAgents: []string{}},
Blacklist: PoWListConfig{IPs: []string{}, IPCidrs: []string{}, Paths: []string{}, PathRegexes: []string{}, UserAgents: []string{}},
}
}
// Route describes a single proxy or pages site entry in the OpenFlare config
// document, including upstream, TLS, caching, rate-limiting and WAF settings.
type Route struct {
ID uint `json:"id,omitempty"`
SiteName string `json:"site_name,omitempty"`
Domains []string `json:"domains,omitempty"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host,omitempty"`
Upstreams []string `json:"upstreams,omitempty"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
DomainCertIDs []uint `json:"domain_cert_ids,omitempty"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server,omitempty"`
LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"`
LimitRate string `json:"limit_rate,omitempty"`
LimitReqPerIP string `json:"limit_req_per_ip,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy,omitempty"`
CacheRules []string `json:"cache_rules,omitempty"`
CustomHeaders []CustomHeader `json:"custom_headers,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
BasicAuthEnabled bool `json:"basic_auth_enabled,omitempty"`
BasicAuthUsername string `json:"basic_auth_username,omitempty"`
BasicAuthPassword string `json:"basic_auth_password,omitempty"`
UpstreamType string `json:"upstream_type,omitempty"`
PagesDeployment *PagesDeployment `json:"pages_deployment,omitempty"`
}
// PagesDeployment holds the static-site deployment parameters for a Pages-type
// route, including local root, entry file, SPA fallback, and API proxy options.
//
// LocalRoot is anchored on ProjectID (projects/{id}/current), not a specific
// deployment ID, so Agents can switch active packages without re-publishing
// main config / reloading OpenResty root paths.
type PagesDeployment struct {
ProjectID uint `json:"project_id"`
ProjectSlug string `json:"project_slug"`
DeploymentID uint `json:"deployment_id"`
DeploymentNumber int `json:"deployment_number"`
Checksum string `json:"checksum"`
EntryFile string `json:"entry_file"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
SPAFallbackPath string `json:"spa_fallback_path"`
APIProxyEnabled bool `json:"api_proxy_enabled"`
APIProxyPath string `json:"api_proxy_path"`
APIProxyPass string `json:"api_proxy_pass"`
APIProxyRewrite string `json:"api_proxy_rewrite"`
LocalRoot string `json:"local_root"`
}
// PagesProjectLocalRoot returns the Agent-local root for a Pages project.
func PagesProjectLocalRoot(projectID uint) string {
if projectID == 0 {
return PagesDirPlaceholder
}
return fmt.Sprintf("%s/projects/%d/current", PagesDirPlaceholder, projectID)
}
// WAFRuleGraph is the compact graph executed by the OpenResty WAF runtime.
type WAFRuleGraph struct {
Entry string `json:"entry"`
Nodes map[string]WAFRuleNode `json:"nodes"`
}
// WAFRuleNode contains one compiled node and its handle-to-target edges.
type WAFRuleNode struct {
Type string `json:"type"`
Config json.RawMessage `json:"config,omitempty"`
Next map[string]string `json:"next,omitempty"`
}
// WAFRuleGroup defines one enabled runtime graph. Legacy flattened fields are
// retained only for decoding older stored snapshots during rolling upgrades.
type WAFRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body,omitempty"`
IPWhitelist []string `json:"ip_whitelist,omitempty"`
IPBlacklist []string `json:"ip_blacklist,omitempty"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist,omitempty"`
CountryBlacklist []string `json:"country_blacklist,omitempty"`
RegionWhitelist []string `json:"region_whitelist,omitempty"`
RegionBlacklist []string `json:"region_blacklist,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
Graph WAFRuleGraph `json:"graph"`
}
// WAFIPGroup is a named, reusable list of IP addresses or CIDRs that can be
// referenced by multiple WAF rule groups as a whitelist or blacklist.
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list,omitempty"`
}
// WAFBinding associates a route (by site name) with the WAF rule groups that
// should be enforced for that site.
type WAFBinding struct {
RouteID uint `json:"route_id"`
SiteName string `json:"site_name"`
RuleGroupIDs []uint `json:"rule_group_ids"`
}
// WAFDocument is the top-level WAF configuration snapshot containing rule
// groups, IP groups, and per-site bindings.
type WAFDocument struct {
RuleGroups []WAFRuleGroup `json:"rule_groups"`
IPGroups []WAFIPGroup `json:"ip_groups,omitempty"`
Bindings []WAFBinding `json:"bindings"`
}
// ConfigSnapshot holds the full set of OpenResty tuning parameters that are
// rendered into the nginx main configuration template.
type ConfigSnapshot struct {
DefaultServerReturnStatus int `json:"default_server_return_status"`
WorkerProcesses string `json:"worker_processes"`
WorkerConnections int `json:"worker_connections"`
WorkerRlimitNofile int `json:"worker_rlimit_nofile"`
EventsUse string `json:"events_use,omitempty"`
EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"`
KeepaliveTimeout int `json:"keepalive_timeout"`
KeepaliveRequests int `json:"keepalive_requests"`
ClientHeaderTimeout int `json:"client_header_timeout"`
ClientBodyTimeout int `json:"client_body_timeout"`
ClientMaxBodySize string `json:"client_max_body_size"`
LargeClientHeaderBuffers string `json:"large_client_header_buffers"`
SendTimeout int `json:"send_timeout"`
ProxyConnectTimeout int `json:"proxy_connect_timeout"`
ProxySendTimeout int `json:"proxy_send_timeout"`
ProxyReadTimeout int `json:"proxy_read_timeout"`
WebsocketEnabled bool `json:"websocket_enabled"`
HTTP3Enabled bool `json:"http3_enabled"`
ProxyRequestBuffering bool `json:"proxy_request_buffering"`
ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"`
ProxyBuffers string `json:"proxy_buffers"`
ProxyBufferSize string `json:"proxy_buffer_size"`
ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"`
GzipEnabled bool `json:"gzip_enabled"`
GzipMinLength int `json:"gzip_min_length"`
GzipCompLevel int `json:"gzip_comp_level"`
Resolvers string `json:"resolvers,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePath string `json:"cache_path,omitempty"`
CacheLevels string `json:"cache_levels"`
CacheInactive string `json:"cache_inactive"`
CacheMaxSize string `json:"cache_max_size"`
CacheKeyTemplate string `json:"cache_key_template"`
CacheLockEnabled bool `json:"cache_lock_enabled"`
CacheLockTimeout string `json:"cache_lock_timeout"`
CacheUseStale string `json:"cache_use_stale"`
MainConfigTemplate string `json:"main_config_template,omitempty"`
DefaultLimitConnPerServer int `json:"default_limit_conn_per_server,omitempty"`
DefaultLimitConnPerIP int `json:"default_limit_conn_per_ip,omitempty"`
DefaultLimitRate string `json:"default_limit_rate,omitempty"`
DefaultLimitReqPerIP string `json:"default_limit_req_per_ip,omitempty"`
OriginErrorPageEnabled bool `json:"origin_error_page_enabled"`
OriginErrorPageStatusCodes []string `json:"origin_error_page_status_codes,omitempty"`
OriginErrorPageHTML string `json:"origin_error_page_html,omitempty"`
// OriginErrorPageGetOnly limits custom error HTML to GET requests; other methods pass through.
OriginErrorPageGetOnly bool `json:"origin_error_page_get_only,omitempty"`
// SWOfflineEnabled enables the Service Worker offline fallback for HTTPS routes.
SWOfflineEnabled bool `json:"sw_offline_enabled,omitempty"`
// SWOfflineHTML is the contact-page HTML served offline; empty uses the built-in default.
SWOfflineHTML string `json:"sw_offline_html,omitempty"`
// SWOfflineDomains restricts the offline fallback to matching HTTPS routes.
SWOfflineDomains []string `json:"sw_offline_domains,omitempty"`
}
// Document is the top-level input structure for the OpenResty renderer,
// combining routes, OpenResty tuning, and WAF configuration.
type Document struct {
Routes []Route `json:"routes"`
OpenRestyConfig ConfigSnapshot `json:"openresty_config"`
WAF WAFDocument `json:"waf"`
}
// Result is the output produced by Render, containing the rendered main
// config, route config, support files, and a content checksum.
type Result struct {
MainConfig string
RouteConfig string
SupportFiles []SupportFile
Checksum string
}
type routeCacheConfig struct {
Enabled bool
Policy string
Rules []string
}
type routeLimitConfig struct {
LimitConnPerServer int
LimitConnPerIP int
LimitRate string
LimitReqPerIP string
}
type routeUpstreamConfig struct {
Name string
Scheme string
ProxyPassURI string
Servers []string
UsesNamedUpstream bool
}
var requiredMainConfigTemplatePlaceholders = []string{
"{{OpenRestyWorkerProcesses}}",
"{{OpenRestyWorkerConnections}}",
"{{OpenRestyWorkerRlimitNofile}}",
"{{OpenRestyConnectionUpgradeMap}}",
"{{OpenRestyDefaultServerBlock}}",
"{{OpenRestyAccessLogPath}}",
"{{OpenRestyErrorLogPath}}",
"{{OpenRestyEventsUseDirective}}",
"{{OpenRestyEventsMultiAcceptDirective}}",
"{{OpenRestyKeepaliveTimeout}}",
"{{OpenRestyKeepaliveRequests}}",
"{{OpenRestyClientHeaderTimeout}}",
"{{OpenRestyClientBodyTimeout}}",
"{{OpenRestyClientMaxBodySize}}",
"{{OpenRestyLargeClientHeaderBuffers}}",
"{{OpenRestySendTimeout}}",
"{{OpenRestyProxyConnectTimeout}}",
"{{OpenRestyProxySendTimeout}}",
"{{OpenRestyProxyReadTimeout}}",
"{{OpenRestyProxyRequestBuffering}}",
"{{OpenRestyProxyBuffering}}",
"{{OpenRestyProxyBuffers}}",
"{{OpenRestyProxyBufferSize}}",
"{{OpenRestyProxyBusyBuffersSize}}",
"{{OpenRestyGzip}}",
"{{OpenRestyGzipMinLength}}",
"{{OpenRestyGzipCompLevel}}",
"{{OpenRestyCacheBlock}}",
"{{OpenRestyRouteConfigInclude}}",
}
+252
View File
@@ -0,0 +1,252 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package wsclient provides a WebSocket client for agent/server communication.
package wsclient
import (
"context"
"encoding/json"
"errors"
"log/slog"
"net"
"net/http"
"net/url"
"strings"
"sync"
"time"
"Wavelet/pkg/util"
"golang.org/x/net/websocket"
)
const (
writeDeadlineSecs = 5
defaultReadDeadlineSecs = 75
)
// Config holds the configuration for a WebSocket client connection.
type Config struct {
BaseURL string
Token string
Timeout time.Duration
HeaderKey string // e.g. "X-Agent-Token", "X-Tunnel-Token"
WSPath string // e.g. "/api/relay/ws", "/api/agent/ws", "/api/flared/ws"
}
// Client provides methods to connect and communicate over WebSocket.
type Client struct {
cfg Config
}
// WSMessage represents a typed WebSocket message with an optional JSON payload.
type WSMessage struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
// MessageHandler handles WebSocket connection lifecycle and incoming messages.
type MessageHandler interface {
OnConnect(ctx context.Context) error
HandleMessage(ctx context.Context, msg WSMessage) error
OnClose(err error)
}
// Connection represents an active WebSocket connection.
type Connection struct {
Conn *websocket.Conn
URL string
ReadTimeout time.Duration
writeMu sync.Mutex
}
// New creates a new WebSocket client with the given configuration.
func New(cfg Config) *Client {
cfg.BaseURL = strings.TrimRight(cfg.BaseURL, "/")
cfg.Token = strings.TrimSpace(cfg.Token)
cfg.HeaderKey = strings.TrimSpace(cfg.HeaderKey)
cfg.WSPath = strings.TrimSpace(cfg.WSPath)
return &Client{
cfg: cfg,
}
}
// SetToken updates the authentication token used for the WebSocket connection.
func (c *Client) SetToken(token string) {
c.cfg.Token = strings.TrimSpace(token)
}
// URL returns the WebSocket URL for the configured endpoint, or empty string on error.
func (c *Client) URL() string {
wsURL, err := c.BuildWebsocketURL()
if err != nil {
return ""
}
return wsURL
}
// BuildWebsocketURL constructs the WebSocket URL by converting the base URL scheme and appending the WS path.
func (c *Client) BuildWebsocketURL() (string, error) {
parsed, err := url.Parse(c.cfg.BaseURL)
if err != nil {
return "", err
}
switch parsed.Scheme {
case "http":
parsed.Scheme = "ws"
case "https":
parsed.Scheme = "wss"
case "ws", "wss":
default:
return "", errors.New("server_url scheme must be http, https, ws, or wss")
}
wsPath := c.cfg.WSPath
if !strings.HasPrefix(wsPath, "/") {
wsPath = "/" + wsPath
}
parsed.Path = strings.TrimRight(parsed.Path, "/") + wsPath
parsed.RawQuery = ""
parsed.Fragment = ""
return parsed.String(), nil
}
// Connect establishes a new WebSocket connection to the configured server.
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
wsURL, err := c.BuildWebsocketURL()
if err != nil {
return nil, err
}
if c.cfg.Token == "" {
return nil, errors.New("ws token is empty")
}
origin := c.cfg.BaseURL
if origin == "" {
origin = "http://localhost"
}
config, err := websocket.NewConfig(wsURL, origin)
if err != nil {
return nil, err
}
config.Header = http.Header{}
if c.cfg.HeaderKey != "" {
config.Header.Set(c.cfg.HeaderKey, c.cfg.Token)
}
if c.cfg.Timeout > 0 {
config.Dialer = &net.Dialer{Timeout: c.cfg.Timeout}
}
slog.Debug("ws dialing server", "url", wsURL)
conn, err := config.DialContext(ctx)
if err != nil {
return nil, err
}
slog.Debug("ws dial succeeded", "url", wsURL)
return &Connection{Conn: conn, URL: wsURL, ReadTimeout: websocketReadTimeout(c.cfg.Timeout)}, nil
}
// SendMessage sends a typed message with an optional payload over the WebSocket connection.
func (conn *Connection) SendMessage(msgType string, payload any) error {
if conn == nil || conn.Conn == nil {
return errors.New("ws connection is nil")
}
slog.Debug("ws sending message", "type", msgType)
// Create the outbound message wrapper
message := struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}{
Type: msgType,
Payload: payload,
}
conn.writeMu.Lock()
defer conn.writeMu.Unlock()
_ = conn.Conn.SetWriteDeadline(time.Now().Add(writeDeadlineSecs * time.Second))
return websocket.JSON.Send(conn.Conn, message)
}
// Receive reads a single message from the WebSocket connection into target.
func (conn *Connection) Receive(target any) error {
if conn == nil || conn.Conn == nil {
return errors.New("ws connection is nil")
}
if conn.ReadTimeout > 0 {
_ = conn.Conn.SetReadDeadline(time.Now().Add(conn.ReadTimeout))
}
err := websocket.JSON.Receive(conn.Conn, target)
if err != nil {
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
slog.Debug("ws receive timeout waiting for server message", "timeout", conn.ReadTimeout)
}
return err
}
return nil
}
func websocketReadTimeout(requestTimeout time.Duration) time.Duration {
timeout := requestTimeout * 6
if timeout < defaultReadDeadlineSecs*time.Second {
return defaultReadDeadlineSecs * time.Second
}
return timeout
}
// RunReceiveLoop continuously receives messages and dispatches them to the handler until the context is cancelled.
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
doneChan := make(chan struct{})
defer close(doneChan)
util.Go(func() {
select {
case <-ctx.Done():
_ = conn.Close()
case <-doneChan:
}
})
if err := handler.OnConnect(ctx); err != nil {
handler.OnClose(err)
return err
}
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
}
var raw WSMessage
if err := conn.Receive(&raw); err != nil {
handler.OnClose(err)
return err
}
switch raw.Type {
case "ping":
slog.Debug("ws received ping from server, replying with pong")
if err := conn.SendMessage("pong", nil); err != nil {
slog.Error("ws send pong response failed", "error", err)
}
case "pong":
slog.Debug("ws received pong response from server")
default:
if err := handler.HandleMessage(ctx, raw); err != nil {
slog.Error("ws handler failed to process message", "type", raw.Type, "error", err)
return err
}
}
}
}
// Close gracefully closes the WebSocket connection.
func (conn *Connection) Close() error {
if conn == nil || conn.Conn == nil {
return nil
}
return conn.Conn.Close()
}