mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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}}",
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user